300 lines
8.9 KiB
Python
300 lines
8.9 KiB
Python
"""
|
||
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条
|