-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathharness.py
More file actions
188 lines (146 loc) · 5.36 KB
/
Copy pathharness.py
File metadata and controls
188 lines (146 loc) · 5.36 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
"""
Harness 层: 约束 + 可观测性 + 审计日志
任何 AI 服务(不只是 Agent)都需要 Harness:
- 约束: 输入校验、频率限制, 防止滥用和异常输入
- 观测: 每次调用的耗时、步骤分解, 方便排查性能问题
- 审计: 完整调用记录写入日志文件, 出问题能回溯
实现方式: 装饰器模式, 不侵入已有代码, 在外面包一层
"""
import json
import os
import time
from datetime import datetime, timezone
# ============ 配置 ============
# 审计日志文件路径
AUDIT_LOG_FILE = "./audit_log.jsonl"
# 输入约束
MAX_QUERY_LENGTH = 1000 # query 最长字符数
MIN_QUERY_LENGTH = 1 # query 最短字符数
# ============ 计时器 ============
class Timer:
"""
简单的计时器, 记录每个步骤的耗时.
用法:
timer = Timer()
timer.start("bm25_search")
... 执行搜索 ...
timer.stop("bm25_search")
print(timer.summary()) # {"bm25_search": 0.12, "total": 0.12}
相当于 C++:
auto start = chrono::high_resolution_clock::now();
... 执行 ...
auto end = chrono::high_resolution_clock::now();
"""
def __init__(self):
self._starts = {} # 步骤名 -> 开始时间
self._durations = {} # 步骤名 -> 耗时(秒)
self._total_start = time.time()
def start(self, step_name: str) -> None:
"""记录某个步骤的开始时间"""
self._starts[step_name] = time.time()
def stop(self, step_name: str) -> float:
"""记录某个步骤的结束时间, 返回耗时秒数"""
if step_name not in self._starts:
return 0.0
duration = time.time() - self._starts[step_name]
self._durations[step_name] = round(duration, 4)
return duration
def summary(self) -> dict:
"""返回所有步骤的耗时汇总"""
result = dict(self._durations)
result["total"] = round(time.time() - self._total_start, 4)
return result
# ============ 审计日志 ============
def write_audit_log(entry: dict) -> None:
"""
写入一条审计日志(JSON Lines 格式, 每行一条 JSON).
JSONL 格式的好处:
- 每行独立, 追加写入不用读整个文件
- 方便用 grep/jq 等工具查询
- 不怕写到一半崩溃导致整个文件损坏(最多丢最后一行)
"""
entry["timestamp"] = datetime.now(timezone.utc).isoformat()
try:
with open(AUDIT_LOG_FILE, "a", encoding="utf-8") as f:
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception as e:
# 审计日志写入失败不应该影响正常功能
print(f"审计日志写入失败: {e}")
# ============ 输入校验 ============
def validate_query(query: str) -> tuple[bool, str]:
"""
校验搜索 query 是否合法.
返回:
(是否合法, 错误信息)
合法时: (True, "")
不合法: (False, "错误原因")
"""
if not query or not query.strip():
return False, "查询内容不能为空"
query = query.strip()
if len(query) < MIN_QUERY_LENGTH:
return False, f"查询内容太短(最少 {MIN_QUERY_LENGTH} 个字符)"
if len(query) > MAX_QUERY_LENGTH:
return False, f"查询内容太长(最多 {MAX_QUERY_LENGTH} 个字符, 当前 {len(query)})"
return True, ""
# ============ 统计汇总 ============
class Stats:
"""
累计统计信息, 运行期间持续更新.
记录总调用次数、总耗时、平均耗时等, 通过 index_status 暴露给用户.
"""
def __init__(self):
self.total_calls = 0
self.total_time = 0.0
self.failed_calls = 0
def record(self, duration: float, success: bool = True) -> None:
self.total_calls += 1
self.total_time += duration
if not success:
self.failed_calls += 1
def summary(self) -> dict:
avg = self.total_time / self.total_calls if self.total_calls > 0 else 0
return {
"total_calls": self.total_calls,
"failed_calls": self.failed_calls,
"total_time": round(self.total_time, 2),
"avg_time": round(avg, 2),
}
# 全局统计实例
search_stats = Stats()
# ============ 测试代码 ============
if __name__ == "__main__":
print("=== 测试 Harness 层 ===\n")
# 测试计时器
print("--- 计时器测试 ---")
timer = Timer()
timer.start("step_a")
time.sleep(0.1)
timer.stop("step_a")
timer.start("step_b")
time.sleep(0.2)
timer.stop("step_b")
print(f"耗时: {timer.summary()}")
# 测试输入校验
print("\n--- 输入校验测试 ---")
test_cases = ["", " ", "a", "正常的搜索query", "x" * 1001]
for q in test_cases:
valid, msg = validate_query(q)
status = "通过" if valid else f"拒绝({msg})"
print(f" '{q[:20]}...' -> {status}" if len(q) > 20 else f" '{q}' -> {status}")
# 测试审计日志
print("\n--- 审计日志测试 ---")
write_audit_log({
"action": "search",
"query": "MCP协议是什么",
"results_count": 3,
"duration": 1.23,
})
print(f"日志已写入: {AUDIT_LOG_FILE}")
# 测试统计
print("\n--- 统计测试 ---")
stats = Stats()
stats.record(1.2)
stats.record(0.8)
stats.record(2.0, success=False)
print(f"统计: {stats.summary()}")