""" Validator Agent - SQL审核员 验证SQL的正确性和安全性 """ import logging import json from typing import Dict, List, Tuple, Optional from camel.agents import ChatAgent from camel.models import ChatModel from config.prompts import VALIDATOR_SYSTEM from utils.validators import full_validation_pipeline from utils.sql_parser import validate_sql_syntax, validate_schema_consistency logger = logging.getLogger(__name__) class ValidatorAgent: """ Validator Agent 职责: - 语法正确性验证 - Schema一致性检查 - 安全性检查(禁止DML/DDL) - 性能问题识别 - 提供修正建议 """ def __init__( self, model: ChatModel, system_message: Optional[str] = None, schema_manager = None ): """ 初始化Agent Args: model: CAMEL AI模型实例 system_message: 系统提示词 schema_manager: Schema管理器(程序验证用) """ self.schema_manager = schema_manager self.system_message = system_message or VALIDATOR_SYSTEM self.agent = ChatAgent( system_message=self.system_message, model=model, ) logger.info("[OK] ValidatorAgent初始化完成") def validate( self, sql: str, schema_str: str, dialect: str = "tsql", check_dangerous: bool = True ) -> Dict: """ 完整验证流程(程序 + LLM双重验证) Args: sql: SQL语句 schema_str: Schema描述 dialect: SQL方言 check_dangerous: 是否检查危险操作 Returns: 验证结果字典 { "valid": bool, "errors": [...], "warnings": [...], "suggestions": [...] } """ result = { "valid": True, "errors": [], "warnings": [], "suggestions": [] } # === 阶段1:程序验证(快速、确定性) === program_result = self._program_validation( sql, dialect, check_dangerous ) result["errors"].extend(program_result.get("errors", [])) result["warnings"].extend(program_result.get("warnings", [])) result["suggestions"].extend(program_result.get("suggestions", [])) # 如果程序验证已发现致命错误,跳过LLM验证 if program_result.get("fatal", False): result["valid"] = False return result # === 阶段2:LLM语义验证 === try: llm_result = self._llm_validation(sql, schema_str) result["errors"].extend(llm_result.get("errors", [])) result["warnings"].extend(llm_result.get("warnings", [])) result["suggestions"].extend(llm_result.get("suggestions", [])) except Exception as e: logger.warning(f"LLM验证失败,使用程序验证结果: {e}") result["warnings"].append(f"LLM验证异常: {str(e)}") result["valid"] = len(result["errors"]) == 0 return result def _program_validation( self, sql: str, dialect: str, check_dangerous: bool ) -> Dict: """ 程序验证(规则引擎) Returns: {"errors": [], "warnings": [], "suggestions": [], "fatal": bool} """ errors = [] warnings = [] suggestions = [] fatal = False # 1. 语法检查 syntax_ok, syntax_errors = validate_sql_syntax(sql, dialect) if not syntax_ok: errors.extend(syntax_errors) fatal = True # 语法错误无法修复,直接失败 return { "errors": errors, "warnings": warnings, "suggestions": suggestions, "fatal": fatal } # 2. Schema一致性(如果有schema_manager) if self.schema_manager: schema_ok, schema_errors = validate_schema_consistency( sql, self.schema_manager, dialect ) if not schema_ok: errors.extend(schema_errors) # Schema错误通常也是fatal的 fatal = True # 3. 危险操作检查 if check_dangerous: from utils.validators import check_dangerous_operations safe, dangers = check_dangerous_operations(sql) if not safe: errors.append(f"包含危险操作: {', '.join(dangers)}") fatal = True return { "errors": errors, "warnings": warnings, "suggestions": suggestions, "fatal": fatal } def _llm_validation(self, sql: str, schema_str: str) -> Dict: """ LLM语义验证 Args: sql: SQL语句 schema_str: Schema描述 Returns: 验证结果 """ from config.prompts import VALIDATOR_USER prompt = VALIDATOR_USER.format(sql=sql, schema=schema_str) try: response = self.agent.step(prompt) content = response.msg.content.strip() # 解析JSON响应 result = self._parse_json_response(content) # 标准化字段 return { "errors": result.get("errors", []), "warnings": result.get("warnings", []), "suggestions": result.get("suggestions", []), } except Exception as e: logger.error(f"LLM验证异常: {e}") return { "errors": [f"LLM验证失败: {str(e)}"], "warnings": [], "suggestions": [] } def _parse_json_response(self, content: str) -> Dict: """解析JSON响应""" try: # 提取```json```块 if "```json" in content: start = content.find("```json") + 7 end = content.find("```", start) if end != -1: content = content[start:end].strip() return json.loads(content) except json.JSONDecodeError as e: logger.warning(f"JSON解析失败: {e}, content={content[:200]}") return {"errors": ["验证结果解析失败"], "warnings": [], "suggestions": []} def quick_check(self, sql: str) -> Tuple[bool, List[str]]: """ 快速检查(仅语法和危险操作) Returns: (是否通过, 错误列表) """ errors = [] # 语法 syntax_ok, syntax_errors = validate_sql_syntax(sql) if not syntax_ok: errors.extend(syntax_errors) return False, errors # 危险操作 from utils.validators import check_dangerous_operations safe, dangers = check_dangerous_operations(sql) if not safe: errors.append(f"危险操作: {', '.join(dangers)}") return False, errors return True, [] def suggest_fixes( self, sql: str, errors: List[str], schema_str: str ) -> List[str]: """ 根据错误建议修复方案 Args: sql: 原始SQL errors: 错误列表 schema_str: Schema描述 Returns: 修复建议列表 """ suggestions = [] # 常见错误模式匹配 for error in errors: error_lower = error.lower() if "field not exist" in error_lower or "字段不存在" in error_lower: suggestions.append( "建议:检查字段名拼写,或使用schema_manager.get_table(table).column_names查看可用字段" ) if "table not exist" in error_lower or "表不存在" in error_lower: suggestions.append( "建议:检查表名拼写,或使用schema_manager.list_tables()查看所有表" ) if "missing join condition" in error_lower or "缺少on条件" in error_lower: suggestions.append( "建议:为每个JOIN添加明确的ON条件,基于外键关系" ) if "group by" in error_lower: suggestions.append( "建议:SELECT中的所有非聚合字段都必须出现在GROUP BY子句中" ) # LLM补充建议 if len(suggestions) < len(errors): try: prompt = f""" SQL: {sql} 错误: {errors} Schema: {schema_str[:1000]} 请给出2-3条具体的修复建议(简洁明了): """ response = self.agent.step(prompt) llm_suggestions = response.msg.content.strip().split('\n') suggestions.extend([s for s in llm_suggestions if s.strip()]) except Exception as e: logger.warning(f"获取LLM建议失败: {e}") return suggestions[:5] # 最多5条