Files
ai-g3sb-backman2.0/agents/validator.py
T
2026-04-10 16:52:07 +08:00

300 lines
8.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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条