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条
|