first commit
This commit is contained in:
@@ -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条
|
||||
Reference in New Issue
Block a user