Coverage for agentos/tools/validation.py: 33%

165 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 17:01 +0800

1""" 

2v1.15.0 — 工具输出验证层:结构化结果验证 + 错误分类 + 自动修复建议。 

3 

4核心功能: 

51. 验证工具返回结果是否符合预期格式 

62. 自动分类工具执行错误 

73. 提供可操作的修复建议 

84. 集成到 ToolExecutor 中,提升 Agent 鲁棒性 

9""" 

10 

11from __future__ import annotations 

12 

13import json 

14import re 

15from dataclasses import dataclass, field 

16from enum import StrEnum 

17from typing import Any 

18 

19from ..errors.handler import ErrorCategory 

20from .base import ToolResult 

21 

22 

23class ValidationSeverity(StrEnum): 

24 """验证结果严重性等级。""" 

25 

26 INFO = "info" # 信息性提示 

27 WARNING = "warning" # 警告,可能有问题但可继续 

28 ERROR = "error" # 错误,需要修复 

29 CRITICAL = "critical" # 严重错误,必须修复 

30 

31 

32class ValidationRule(StrEnum): 

33 """验证规则类型。""" 

34 

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" # 结构检查 

43 

44 

45@dataclass 

46class ValidationIssue: 

47 """验证问题。""" 

48 

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 = "" 

56 

57 

58@dataclass 

59class ValidationResult: 

60 """验证结果。""" 

61 

62 is_valid: bool 

63 issues: list[ValidationIssue] = field(default_factory=list) 

64 normalized_output: Any | None = None 

65 

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 ) 

72 

73 @property 

74 def has_warnings(self) -> bool: 

75 return any(issue.severity == ValidationSeverity.WARNING for issue in self.issues) 

76 

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 

81 

82 

83class ToolOutputValidator: 

84 """工具输出验证器。""" 

85 

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] = {} 

90 

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) 

96 

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")} 

105 

106 return self 

107 

108 def validate(self, tool_result: ToolResult) -> ValidationResult: 

109 """验证工具结果。""" 

110 result = ValidationResult(is_valid=True) 

111 

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 

123 

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 

134 

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 

140 

141 result.normalized_output = parsed_output 

142 

143 # 应用验证规则 

144 self._apply_rules(result, parsed_output) 

145 

146 return result 

147 

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 

155 

156 # 尝试解析为 Python 字典格式(如 "{'key': 'value'}") 

157 try: 

158 # 安全地使用 eval 但限制为字面量 

159 import ast 

160 

161 return ast.literal_eval(output) 

162 except (SyntaxError, ValueError): 

163 pass 

164 

165 # 检查是否为纯文本 

166 if output.strip(): 

167 return {"text": output.strip()} 

168 

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 ) 

176 

177 def _apply_rules(self, result: ValidationResult, data: Any) -> None: 

178 """应用验证规则到数据。""" 

179 if not isinstance(data, dict): 

180 return 

181 

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 

195 

196 value = data[field] 

197 

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 ) 

213 

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 ) 

240 

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 ) 

255 

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 ) 

270 

271 

272class ToolErrorClassifier: 

273 """工具错误分类器。""" 

274 

275 @staticmethod 

276 def classify(tool_result: ToolResult) -> ErrorCategory: 

277 """根据工具结果分类错误。""" 

278 if tool_result.error: 

279 error_msg = tool_result.error.lower() 

280 

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 

300 

301 return ErrorCategory.UNKNOWN 

302 

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 } 

335 

336 return base_suggestions.get(category, ["请查看详细错误信息"]) 

337 

338 

339def validate_tool_output( 

340 tool_name: str, tool_result: ToolResult, expected_schema: dict | None = None 

341) -> ValidationResult: 

342 """ 

343 验证工具输出的便捷函数。 

344 

345 Args: 

346 tool_name: 工具名称 

347 tool_result: 工具执行结果 

348 expected_schema: 期望的输出模式(可选) 

349 

350 Returns: 

351 ValidationResult: 验证结果 

352 """ 

353 validator = ToolOutputValidator(tool_name) 

354 

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 ) 

369 

370 return validator.validate(tool_result) 

371 

372 

373def classify_tool_error(tool_result: ToolResult) -> dict[str, Any]: 

374 """ 

375 分类工具错误并返回结构化信息。 

376 

377 Args: 

378 tool_result: 工具执行结果 

379 

380 Returns: 

381 Dict: 包含错误分类和恢复建议的字典 

382 """ 

383 category = ToolErrorClassifier.classify(tool_result) 

384 suggestions = ToolErrorClassifier.get_recovery_suggestions(category, "unknown") 

385 

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 }