Coverage for agentos/tools/validation.py: 33%
165 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 01:44 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 01:44 +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[field] = {"type": kwargs.get("expected_type")}
99 elif rule == ValidationRule.RANGE_CHECK:
100 self._field_schemas[field] = {"min": kwargs.get("min"), "max": kwargs.get("max")}
101 elif rule == ValidationRule.PATTERN_MATCH:
102 self._field_schemas[field] = {"pattern": kwargs.get("pattern")}
103 elif rule == ValidationRule.ENUM_CHECK:
104 self._field_schemas[field] = {"allowed_values": kwargs.get("allowed_values")}
106 return self
108 def validate(self, tool_result: ToolResult) -> ValidationResult:
109 """验证工具结果。"""
110 result = ValidationResult(is_valid=True)
112 # 检查工具执行是否成功
113 if tool_result.error:
114 result.add_issue(
115 ValidationIssue(
116 rule=ValidationRule.REQUIRED_FIELD,
117 severity=ValidationSeverity.ERROR,
118 message=f"工具执行失败: {tool_result.error}",
119 suggestion="请检查工具参数和依赖环境",
120 )
121 )
122 return result
124 if not tool_result.output:
125 result.add_issue(
126 ValidationIssue(
127 rule=ValidationRule.REQUIRED_FIELD,
128 severity=ValidationSeverity.WARNING,
129 message="工具返回空输出",
130 suggestion="检查工具是否按预期工作",
131 )
132 )
133 return result
135 # 尝试解析输出
136 parsed_output = self._parse_output(tool_result.output)
137 if isinstance(parsed_output, ValidationIssue):
138 result.add_issue(parsed_output)
139 return result
141 result.normalized_output = parsed_output
143 # 应用验证规则
144 self._apply_rules(result, parsed_output)
146 return result
148 def _parse_output(self, output: str) -> Any | ValidationIssue:
149 """解析工具输出。"""
150 # 尝试解析为 JSON
151 try:
152 return json.loads(output)
153 except json.JSONDecodeError:
154 pass
156 # 尝试解析为 Python 字典格式(如 "{'key': 'value'}")
157 try:
158 # 安全地使用 eval 但限制为字面量
159 import ast
161 return ast.literal_eval(output)
162 except (SyntaxError, ValueError):
163 pass
165 # 检查是否为纯文本
166 if output.strip():
167 return {"text": output.strip()}
169 return ValidationIssue(
170 rule=ValidationRule.JSON_FORMAT,
171 severity=ValidationSeverity.ERROR,
172 message="无法解析工具输出",
173 actual=output[:100] if output else "空字符串",
174 suggestion="工具应返回 JSON 或结构化文本",
175 )
177 def _apply_rules(self, result: ValidationResult, data: Any) -> None:
178 """应用验证规则到数据。"""
179 if not isinstance(data, dict):
180 return
182 for f, rules in self._rules.items():
183 if f not in data:
184 if ValidationRule.REQUIRED_FIELD in rules:
185 result.add_issue(
186 ValidationIssue(
187 rule=ValidationRule.REQUIRED_FIELD,
188 severity=ValidationSeverity.ERROR,
189 message=f"缺少必需字段: {field}",
190 field=field,
191 suggestion=f"工具应返回字段 '{field}'",
192 )
193 )
194 continue
196 value = data[field]
198 for rule in rules:
199 if rule == ValidationRule.TYPE_CHECK:
200 expected_type = self._field_schemas[field]["type"]
201 if not isinstance(value, expected_type):
202 result.add_issue(
203 ValidationIssue(
204 rule=ValidationRule.TYPE_CHECK,
205 severity=ValidationSeverity.ERROR,
206 message=f"字段类型错误: {field}",
207 field=field,
208 expected=expected_type.__name__,
209 actual=type(value).__name__,
210 suggestion=f"字段 '{field}' 应为 {expected_type.__name__} 类型",
211 )
212 )
214 elif rule == ValidationRule.RANGE_CHECK:
215 schema = self._field_schemas[field]
216 if "min" in schema and value < schema["min"]:
217 result.add_issue(
218 ValidationIssue(
219 rule=ValidationRule.RANGE_CHECK,
220 severity=ValidationSeverity.WARNING,
221 message=f"字段值过小: {field}",
222 field=field,
223 expected=f">= {schema['min']}",
224 actual=value,
225 suggestion=f"字段 '{field}' 应大于等于 {schema['min']}",
226 )
227 )
228 if "max" in schema and value > schema["max"]:
229 result.add_issue(
230 ValidationIssue(
231 rule=ValidationRule.RANGE_CHECK,
232 severity=ValidationSeverity.WARNING,
233 message=f"字段值过大: {field}",
234 field=field,
235 expected=f"<= {schema['max']}",
236 actual=value,
237 suggestion=f"字段 '{field}' 应小于等于 {schema['max']}",
238 )
239 )
241 elif rule == ValidationRule.PATTERN_MATCH:
242 pattern = self._field_schemas[field]["pattern"]
243 if not re.match(pattern, str(value)):
244 result.add_issue(
245 ValidationIssue(
246 rule=ValidationRule.PATTERN_MATCH,
247 severity=ValidationSeverity.ERROR,
248 message=f"字段格式错误: {field}",
249 field=field,
250 expected=f"匹配模式: {pattern}",
251 actual=value,
252 suggestion=f"字段 '{field}' 应符合正则表达式: {pattern}",
253 )
254 )
256 elif rule == ValidationRule.ENUM_CHECK:
257 allowed = self._field_schemas[field]["allowed_values"]
258 if value not in allowed:
259 result.add_issue(
260 ValidationIssue(
261 rule=ValidationRule.ENUM_CHECK,
262 severity=ValidationSeverity.ERROR,
263 message=f"字段值不在允许范围内: {field}",
264 field=field,
265 expected=allowed,
266 actual=value,
267 suggestion=f"字段 '{field}' 应为以下值之一: {allowed}",
268 )
269 )
272class ToolErrorClassifier:
273 """工具错误分类器。"""
275 @staticmethod
276 def classify(tool_result: ToolResult) -> ErrorCategory:
277 """根据工具结果分类错误。"""
278 if tool_result.error:
279 error_msg = tool_result.error.lower()
281 if any(kw in error_msg for kw in ["permission", "access denied", "forbidden"]):
282 return ErrorCategory.AUTH
283 elif any(kw in error_msg for kw in ["timeout", "timed out"]):
284 return ErrorCategory.TIMEOUT
285 elif any(kw in error_msg for kw in ["network", "connection", "dns"]):
286 return ErrorCategory.NETWORK
287 elif any(kw in error_msg for kw in ["not found", "file not found", "no such file"]):
288 return ErrorCategory.VALIDATION
289 elif any(
290 kw in error_msg
291 for kw in ["authentication", "auth", "api key", "invalid key", "unauthorized"]
292 ):
293 return ErrorCategory.AUTH
294 elif any(kw in error_msg for kw in ["memory", "disk", "resource"]):
295 return ErrorCategory.RESOURCE
296 elif any(kw in error_msg for kw in ["syntax", "invalid", "malformed"]):
297 return ErrorCategory.VALIDATION
298 elif any(kw in error_msg for kw in ["rate limit", "too many", "quota"]):
299 return ErrorCategory.RATE_LIMIT
301 return ErrorCategory.UNKNOWN
303 @staticmethod
304 def get_recovery_suggestions(category: ErrorCategory, tool_name: str) -> list[str]:
305 """获取针对特定工具的错误恢复建议。"""
306 base_suggestions = {
307 ErrorCategory.AUTH: [
308 f"检查 {tool_name} 工具所需的权限",
309 "确认当前用户有足够的访问权限",
310 "检查 API Key 或认证令牌是否有效",
311 ],
312 ErrorCategory.TIMEOUT: [
313 f"增加 {tool_name} 工具的超时时间",
314 "检查目标服务是否正常运行",
315 "考虑使用更轻量的查询参数",
316 ],
317 ErrorCategory.NETWORK: [
318 "检查网络连接",
319 "确认目标服务地址是否正确",
320 "尝试使用代理或 VPN",
321 ],
322 ErrorCategory.VALIDATION: [
323 f"检查 {tool_name} 工具的输入参数",
324 "确认文件路径或资源是否存在",
325 "验证输入数据的格式和类型",
326 ],
327 ErrorCategory.RESOURCE: ["清理磁盘空间", "增加系统内存", "减少并发请求数量"],
328 ErrorCategory.RATE_LIMIT: ["降低请求频率", "使用指数退避重试", "检查 API 配额限制"],
329 ErrorCategory.UNKNOWN: [
330 f"查看 {tool_name} 工具的详细日志",
331 "检查工具依赖是否完整",
332 "尝试重启相关服务",
333 ],
334 }
336 return base_suggestions.get(category, ["请查看详细错误信息"])
339def validate_tool_output(
340 tool_name: str, tool_result: ToolResult, expected_schema: dict | None = None
341) -> ValidationResult:
342 """
343 验证工具输出的便捷函数。
345 Args:
346 tool_name: 工具名称
347 tool_result: 工具执行结果
348 expected_schema: 期望的输出模式(可选)
350 Returns:
351 ValidationResult: 验证结果
352 """
353 validator = ToolOutputValidator(tool_name)
355 if expected_schema:
356 for field, schema in expected_schema.items():
357 if "type" in schema:
358 validator.add_rule(field, ValidationRule.TYPE_CHECK, expected_type=schema["type"])
359 if "required" in schema and schema["required"]:
360 validator.add_rule(field, ValidationRule.REQUIRED_FIELD)
361 if "pattern" in schema:
362 validator.add_rule(field, ValidationRule.PATTERN_MATCH, pattern=schema["pattern"])
363 if "enum" in schema:
364 validator.add_rule(field, ValidationRule.ENUM_CHECK, allowed_values=schema["enum"])
365 if "min" in schema or "max" in schema:
366 validator.add_rule(
367 field, ValidationRule.RANGE_CHECK, min=schema.get("min"), max=schema.get("max")
368 )
370 return validator.validate(tool_result)
373def classify_tool_error(tool_result: ToolResult) -> dict[str, Any]:
374 """
375 分类工具错误并返回结构化信息。
377 Args:
378 tool_result: 工具执行结果
380 Returns:
381 Dict: 包含错误分类和恢复建议的字典
382 """
383 category = ToolErrorClassifier.classify(tool_result)
384 suggestions = ToolErrorClassifier.get_recovery_suggestions(category, "unknown")
386 return {
387 "category": category.name,
388 "error_message": tool_result.error,
389 "suggestions": suggestions,
390 "severity": "error" if category != ErrorCategory.UNKNOWN else "warning",
391 }