Coverage for agentos/tools/validation.py: 32%
169 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 07:12 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 07:12 +0800
1"""
2v1.15.0 — 工具输出验证层:结构化结果验证 + 错误分类 + 自动修复建议。
4核心功能:
51. 验证工具返回结果是否符合预期格式
62. 自动分类工具执行错误
73. 提供可操作的修复建议
84. 集成到 ToolExecutor 中,提升 Agent 鲁棒性
9"""
11from __future__ import annotations
13import json
14import re
15from dataclasses import dataclass, field
16from enum import StrEnum
17from typing import Any
19from ..errors.handler import ErrorCategory
20from .base import ToolResult
23class ValidationSeverity(StrEnum):
24 """验证结果严重性等级。"""
26 INFO = "info" # 信息性提示
27 WARNING = "warning" # 警告,可能有问题但可继续
28 ERROR = "error" # 错误,需要修复
29 CRITICAL = "critical" # 严重错误,必须修复
32class ValidationRule(StrEnum):
33 """验证规则类型。"""
35 JSON_FORMAT = "json_format" # JSON 格式验证
36 REQUIRED_FIELD = "required_field" # 必需字段检查
37 TYPE_CHECK = "type_check" # 类型检查
38 RANGE_CHECK = "range_check" # 范围检查
39 PATTERN_MATCH = "pattern_match" # 正则匹配
40 LENGTH_CHECK = "length_check" # 长度检查
41 ENUM_CHECK = "enum_check" # 枚举值检查
42 STRUCTURE_CHECK = "structure_check" # 结构检查
45@dataclass
46class ValidationIssue:
47 """验证问题。"""
49 rule: ValidationRule
50 severity: ValidationSeverity
51 message: str
52 field: str | None = None
53 expected: Any | None = None
54 actual: Any | None = None
55 suggestion: str = ""
58@dataclass
59class ValidationResult:
60 """验证结果。"""
62 is_valid: bool
63 issues: list[ValidationIssue] = field(default_factory=list)
64 normalized_output: Any | None = None
66 @property
67 def has_errors(self) -> bool:
68 return any(
69 issue.severity in (ValidationSeverity.ERROR, ValidationSeverity.CRITICAL)
70 for issue in self.issues
71 )
73 @property
74 def has_warnings(self) -> bool:
75 return any(issue.severity == ValidationSeverity.WARNING for issue in self.issues)
77 def add_issue(self, issue: ValidationIssue) -> None:
78 self.issues.append(issue)
79 if issue.severity in (ValidationSeverity.ERROR, ValidationSeverity.CRITICAL):
80 self.is_valid = False
83class ToolOutputValidator:
84 """工具输出验证器。"""
86 def __init__(self, tool_name: str):
87 self.tool_name = tool_name
88 self._rules: dict[str, list[ValidationRule]] = {}
89 self._field_schemas: dict[str, dict] = {}
91 def add_rule(self, field: str, rule: ValidationRule, **kwargs) -> ToolOutputValidator:
92 """为指定字段添加验证规则。"""
93 if field not in self._rules:
94 self._rules[field] = []
95 self._rules[field].append(rule)
97 if rule == ValidationRule.TYPE_CHECK:
98 self._field_schemas.setdefault(field, {})["type"] = kwargs.get("expected_type")
99 elif rule == ValidationRule.RANGE_CHECK:
100 entry = self._field_schemas.setdefault(field, {})
101 if kwargs.get("min") is not None:
102 entry["min"] = kwargs["min"]
103 if kwargs.get("max") is not None:
104 entry["max"] = kwargs["max"]
105 elif rule == ValidationRule.PATTERN_MATCH:
106 self._field_schemas.setdefault(field, {})["pattern"] = kwargs.get("pattern")
107 elif rule == ValidationRule.ENUM_CHECK:
108 self._field_schemas.setdefault(field, {})["allowed_values"] = kwargs.get("allowed_values")
110 return self
112 def validate(self, tool_result: ToolResult) -> ValidationResult:
113 """验证工具结果。"""
114 result = ValidationResult(is_valid=True)
116 # 检查工具执行是否成功
117 if tool_result.error:
118 result.add_issue(
119 ValidationIssue(
120 rule=ValidationRule.REQUIRED_FIELD,
121 severity=ValidationSeverity.ERROR,
122 message=f"工具执行失败: {tool_result.error}",
123 suggestion="请检查工具参数和依赖环境",
124 )
125 )
126 return result
128 if not tool_result.output:
129 result.add_issue(
130 ValidationIssue(
131 rule=ValidationRule.REQUIRED_FIELD,
132 severity=ValidationSeverity.WARNING,
133 message="工具返回空输出",
134 suggestion="检查工具是否按预期工作",
135 )
136 )
137 return result
139 # 尝试解析输出
140 parsed_output = self._parse_output(tool_result.output)
141 if isinstance(parsed_output, ValidationIssue):
142 result.add_issue(parsed_output)
143 return result
145 result.normalized_output = parsed_output
147 # 应用验证规则
148 self._apply_rules(result, parsed_output)
150 return result
152 def _parse_output(self, output: str) -> Any | ValidationIssue:
153 """解析工具输出。"""
154 # 尝试解析为 JSON
155 try:
156 return json.loads(output)
157 except json.JSONDecodeError:
158 pass
160 # 尝试解析为 Python 字典格式(如 "{'key': 'value'}")
161 try:
162 # 安全地使用 eval 但限制为字面量
163 import ast
165 return ast.literal_eval(output)
166 except (SyntaxError, ValueError):
167 pass
169 # 检查是否为纯文本
170 if output.strip():
171 return {"text": output.strip()}
173 return ValidationIssue(
174 rule=ValidationRule.JSON_FORMAT,
175 severity=ValidationSeverity.ERROR,
176 message="无法解析工具输出",
177 actual=output[:100] if output else "空字符串",
178 suggestion="工具应返回 JSON 或结构化文本",
179 )
181 def _apply_rules(self, result: ValidationResult, data: Any) -> None:
182 """应用验证规则到数据。"""
183 if not isinstance(data, dict):
184 return
186 for fname, rules in self._rules.items():
187 if fname not in data:
188 if ValidationRule.REQUIRED_FIELD in rules:
189 result.add_issue(
190 ValidationIssue(
191 rule=ValidationRule.REQUIRED_FIELD,
192 severity=ValidationSeverity.ERROR,
193 message=f"缺少必需字段: {fname}",
194 field=fname,
195 suggestion=f"工具应返回字段 '{fname}'",
196 )
197 )
198 continue
200 value = data[fname]
202 for rule in rules:
203 if rule == ValidationRule.TYPE_CHECK:
204 expected_type = self._field_schemas[fname]["type"]
205 if not isinstance(value, expected_type):
206 result.add_issue(
207 ValidationIssue(
208 rule=ValidationRule.TYPE_CHECK,
209 severity=ValidationSeverity.ERROR,
210 message=f"字段类型错误: {fname}",
211 field=fname,
212 expected=expected_type.__name__,
213 actual=type(value).__name__,
214 suggestion=f"字段 '{fname}' 应为 {expected_type.__name__} 类型",
215 )
216 )
218 elif rule == ValidationRule.RANGE_CHECK:
219 schema = self._field_schemas[fname]
220 if "min" in schema and schema["min"] is not None and value < schema["min"]:
221 result.add_issue(
222 ValidationIssue(
223 rule=ValidationRule.RANGE_CHECK,
224 severity=ValidationSeverity.WARNING,
225 message=f"字段值过小: {fname}",
226 field=fname,
227 expected=f">= {schema['min']}",
228 actual=value,
229 suggestion=f"字段 '{fname}' 应大于等于 {schema['min']}",
230 )
231 )
232 if "max" in schema and schema["max"] is not None and value > schema["max"]:
233 result.add_issue(
234 ValidationIssue(
235 rule=ValidationRule.RANGE_CHECK,
236 severity=ValidationSeverity.WARNING,
237 message=f"字段值过大: {fname}",
238 field=fname,
239 expected=f"<= {schema['max']}",
240 actual=value,
241 suggestion=f"字段 '{fname}' 应小于等于 {schema['max']}",
242 )
243 )
245 elif rule == ValidationRule.PATTERN_MATCH:
246 pattern = self._field_schemas[fname]["pattern"]
247 if not re.match(pattern, str(value)):
248 result.add_issue(
249 ValidationIssue(
250 rule=ValidationRule.PATTERN_MATCH,
251 severity=ValidationSeverity.ERROR,
252 message=f"字段格式错误: {fname}",
253 field=fname,
254 expected=f"匹配模式: {pattern}",
255 actual=value,
256 suggestion=f"字段 '{fname}' 应符合正则表达式: {pattern}",
257 )
258 )
260 elif rule == ValidationRule.ENUM_CHECK:
261 allowed = self._field_schemas[fname]["allowed_values"]
262 if value not in allowed:
263 result.add_issue(
264 ValidationIssue(
265 rule=ValidationRule.ENUM_CHECK,
266 severity=ValidationSeverity.ERROR,
267 message=f"字段值不在允许范围内: {fname}",
268 field=fname,
269 expected=allowed,
270 actual=value,
271 suggestion=f"字段 '{fname}' 应为以下值之一: {allowed}",
272 )
273 )
276class ToolErrorClassifier:
277 """工具错误分类器。"""
279 @staticmethod
280 def classify(tool_result: ToolResult) -> ErrorCategory:
281 """根据工具结果分类错误。"""
282 if tool_result.error:
283 error_msg = tool_result.error.lower()
285 if any(kw in error_msg for kw in ["permission", "access denied", "forbidden"]):
286 return ErrorCategory.AUTH
287 elif any(kw in error_msg for kw in ["timeout", "timed out"]):
288 return ErrorCategory.TIMEOUT
289 elif any(kw in error_msg for kw in ["network", "connection", "dns"]):
290 return ErrorCategory.NETWORK
291 elif any(kw in error_msg for kw in ["not found", "file not found", "no such file"]):
292 return ErrorCategory.VALIDATION
293 elif any(
294 kw in error_msg
295 for kw in ["authentication", "auth", "api key", "invalid key", "unauthorized"]
296 ):
297 return ErrorCategory.AUTH
298 elif any(kw in error_msg for kw in ["memory", "disk", "resource"]):
299 return ErrorCategory.RESOURCE
300 elif any(kw in error_msg for kw in ["syntax", "invalid", "malformed"]):
301 return ErrorCategory.VALIDATION
302 elif any(kw in error_msg for kw in ["rate limit", "too many", "quota"]):
303 return ErrorCategory.RATE_LIMIT
305 return ErrorCategory.UNKNOWN
307 @staticmethod
308 def get_recovery_suggestions(category: ErrorCategory, tool_name: str) -> list[str]:
309 """获取针对特定工具的错误恢复建议。"""
310 base_suggestions = {
311 ErrorCategory.AUTH: [
312 f"检查 {tool_name} 工具所需的权限",
313 "确认当前用户有足够的访问权限",
314 "检查 API Key 或认证令牌是否有效",
315 ],
316 ErrorCategory.TIMEOUT: [
317 f"增加 {tool_name} 工具的超时时间",
318 "检查目标服务是否正常运行",
319 "考虑使用更轻量的查询参数",
320 ],
321 ErrorCategory.NETWORK: [
322 "检查网络连接",
323 "确认目标服务地址是否正确",
324 "尝试使用代理或 VPN",
325 ],
326 ErrorCategory.VALIDATION: [
327 f"检查 {tool_name} 工具的输入参数",
328 "确认文件路径或资源是否存在",
329 "验证输入数据的格式和类型",
330 ],
331 ErrorCategory.RESOURCE: ["清理磁盘空间", "增加系统内存", "减少并发请求数量"],
332 ErrorCategory.RATE_LIMIT: ["降低请求频率", "使用指数退避重试", "检查 API 配额限制"],
333 ErrorCategory.UNKNOWN: [
334 f"查看 {tool_name} 工具的详细日志",
335 "检查工具依赖是否完整",
336 "尝试重启相关服务",
337 ],
338 }
340 return base_suggestions.get(category, ["请查看详细错误信息"])
343def validate_tool_output(
344 tool_name: str, tool_result: ToolResult, expected_schema: dict | None = None
345) -> ValidationResult:
346 """
347 验证工具输出的便捷函数。
349 Args:
350 tool_name: 工具名称
351 tool_result: 工具执行结果
352 expected_schema: 期望的输出模式(可选)
354 Returns:
355 ValidationResult: 验证结果
356 """
357 validator = ToolOutputValidator(tool_name)
359 if expected_schema:
360 for field, schema in expected_schema.items():
361 if "type" in schema:
362 validator.add_rule(field, ValidationRule.TYPE_CHECK, expected_type=schema["type"])
363 if "required" in schema and schema["required"]:
364 validator.add_rule(field, ValidationRule.REQUIRED_FIELD)
365 if "pattern" in schema:
366 validator.add_rule(field, ValidationRule.PATTERN_MATCH, pattern=schema["pattern"])
367 if "enum" in schema:
368 validator.add_rule(field, ValidationRule.ENUM_CHECK, allowed_values=schema["enum"])
369 if "min" in schema or "max" in schema:
370 validator.add_rule(
371 field, ValidationRule.RANGE_CHECK, min=schema.get("min"), max=schema.get("max")
372 )
374 return validator.validate(tool_result)
377def classify_tool_error(tool_result: ToolResult) -> dict[str, Any]:
378 """
379 分类工具错误并返回结构化信息。
381 Args:
382 tool_result: 工具执行结果
384 Returns:
385 Dict: 包含错误分类和恢复建议的字典
386 """
387 category = ToolErrorClassifier.classify(tool_result)
388 suggestions = ToolErrorClassifier.get_recovery_suggestions(category, "unknown")
390 return {
391 "category": category.name,
392 "error_message": tool_result.error,
393 "suggestions": suggestions,
394 "severity": "error" if category != ErrorCategory.UNKNOWN else "warning",
395 }