first commit

This commit is contained in:
陈辅元
2026-04-10 16:52:07 +08:00
commit 84fe545640
87 changed files with 20842 additions and 0 deletions
+299
View File
@@ -0,0 +1,299 @@
"""
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条