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

169 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 21:26 +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.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") 

109 

110 return self 

111 

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

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

114 result = ValidationResult(is_valid=True) 

115 

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 

127 

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 

138 

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 

144 

145 result.normalized_output = parsed_output 

146 

147 # 应用验证规则 

148 self._apply_rules(result, parsed_output) 

149 

150 return result 

151 

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 

159 

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

161 try: 

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

163 import ast 

164 

165 return ast.literal_eval(output) 

166 except (SyntaxError, ValueError): 

167 pass 

168 

169 # 检查是否为纯文本 

170 if output.strip(): 

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

172 

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 ) 

180 

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

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

183 if not isinstance(data, dict): 

184 return 

185 

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 

199 

200 value = data[fname] 

201 

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 ) 

217 

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 ) 

244 

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 ) 

259 

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 ) 

274 

275 

276class ToolErrorClassifier: 

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

278 

279 @staticmethod 

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

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

282 if tool_result.error: 

283 error_msg = tool_result.error.lower() 

284 

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 

304 

305 return ErrorCategory.UNKNOWN 

306 

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 } 

339 

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

341 

342 

343def validate_tool_output( 

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

345) -> ValidationResult: 

346 """ 

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

348 

349 Args: 

350 tool_name: 工具名称 

351 tool_result: 工具执行结果 

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

353 

354 Returns: 

355 ValidationResult: 验证结果 

356 """ 

357 validator = ToolOutputValidator(tool_name) 

358 

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 ) 

373 

374 return validator.validate(tool_result) 

375 

376 

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

378 """ 

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

380 

381 Args: 

382 tool_result: 工具执行结果 

383 

384 Returns: 

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

386 """ 

387 category = ToolErrorClassifier.classify(tool_result) 

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

389 

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 }