207 lines
5.5 KiB
Python
207 lines
5.5 KiB
Python
"""
|
||||
|
|
SQL Generator Agent - SQL生成专家
|
|||
|
|
根据Schema和问题生成高质量SQL
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
import logging
|
|||
|
|
import json
|
|||
|
|
import re
|
|||
|
|
from typing import Dict, Optional
|
|||
|
|
|
|||
|
|
from camel.agents import ChatAgent
|
|||
|
|
from camel.models import ChatModel
|
|||
|
|
from config.prompts import SQL_GENERATOR_SYSTEM
|
|||
|
|
|
|||
|
|
logger = logging.getLogger(__name__)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class SQLGeneratorAgent:
|
|||
|
|
"""
|
|||
|
|
SQL Generator Agent
|
|||
|
|
|
|||
|
|
职责:
|
|||
|
|
- 理解用户问题和Schema结构
|
|||
|
|
- 生成准确的SQL语句
|
|||
|
|
- 处理复杂的JOIN、聚合、子查询
|
|||
|
|
- 遵循金融/证券业务特殊规则
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
def __init__(
|
|||
|
|
self,
|
|||
|
|
model: ChatModel,
|
|||
|
|
system_message: Optional[str] = None,
|
|||
|
|
dialect: str = "tsql"
|
|||
|
|
):
|
|||
|
|
"""
|
|||
|
|
初始化Agent
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
model: CAMEL AI模型实例
|
|||
|
|
system_message: 系统提示词
|
|||
|
|
dialect: SQL方言
|
|||
|
|
"""
|
|||
|
|
self.dialect = dialect
|
|||
|
|
self.system_message = system_message or SQL_GENERATOR_SYSTEM
|
|||
|
|
self.agent = ChatAgent(
|
|||
|
|
system_message=self.system_message,
|
|||
|
|
model=model,
|
|||
|
|
)
|
|||
|
|
logger.info(f"[OK] SQLGeneratorAgent初始化完成 (dialect={dialect})")
|
|||
|
|
|
|||
|
|
def generate(
|
|||
|
|
self,
|
|||
|
|
question: str,
|
|||
|
|
schema_str: str,
|
|||
|
|
examples: Optional[List[Dict]] = None,
|
|||
|
|
dialect: Optional[str] = None
|
|||
|
|
) -> str:
|
|||
|
|
"""
|
|||
|
|
生成SQL
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
question: 用户问题
|
|||
|
|
schema_str: Schema描述字符串
|
|||
|
|
examples: Few-shot示例列表
|
|||
|
|
dialect: 覆盖默认dialect
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
SQL语句字符串
|
|||
|
|
"""
|
|||
|
|
from config.prompts import SQL_GENERATOR_USER
|
|||
|
|
|
|||
|
|
dialect = dialect or self.dialect
|
|||
|
|
prompt = SQL_GENERATOR_USER.format(
|
|||
|
|
schema=schema_str,
|
|||
|
|
question=question,
|
|||
|
|
dialect=dialect
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# 添加few-shot示例(如果有)
|
|||
|
|
if examples:
|
|||
|
|
prompt = self._inject_examples(prompt, examples)
|
|||
|
|
|
|||
|
|
try:
|
|||
|
|
response = self.agent.step(prompt)
|
|||
|
|
sql = self._extract_sql(response.msg.content)
|
|||
|
|
|
|||
|
|
logger.debug(f"生成的SQL: {sql[:200]}...")
|
|||
|
|
return sql
|
|||
|
|
|
|||
|
|
except Exception as e:
|
|||
|
|
logger.error(f"SQL生成失败: {e}")
|
|||
|
|
raise
|
|||
|
|
|
|||
|
|
def _extract_sql(self, content: str) -> str:
|
|||
|
|
"""
|
|||
|
|
从Agent响应中提取SQL语句
|
|||
|
|
|
|||
|
|
处理:
|
|||
|
|
- Markdown代码块 (```sql ... ```)
|
|||
|
|
- 纯SQL文本
|
|||
|
|
- JSON格式 {"sql": "..."}
|
|||
|
|
"""
|
|||
|
|
content = content.strip()
|
|||
|
|
|
|||
|
|
# 尝试提取```sql```块
|
|||
|
|
if "```sql" in content:
|
|||
|
|
start = content.find("```sql") + 6
|
|||
|
|
end = content.find("```", start)
|
|||
|
|
if end != -1:
|
|||
|
|
return content[start:end].strip()
|
|||
|
|
|
|||
|
|
# 尝试提取通用代码块```
|
|||
|
|
if "```" in content:
|
|||
|
|
start = content.find("```") + 3
|
|||
|
|
end = content.find("```", start)
|
|||
|
|
if end != -1:
|
|||
|
|
return content[start:end].strip()
|
|||
|
|
|
|||
|
|
# 尝试解析JSON
|
|||
|
|
try:
|
|||
|
|
data = json.loads(content)
|
|||
|
|
if "sql" in data:
|
|||
|
|
return data["sql"]
|
|||
|
|
except json.JSONDecodeError:
|
|||
|
|
pass
|
|||
|
|
|
|||
|
|
# 返回原始内容(假设是纯SQL)
|
|||
|
|
return content
|
|||
|
|
|
|||
|
|
def _inject_examples(
|
|||
|
|
self,
|
|||
|
|
prompt: str,
|
|||
|
|
examples: List[Dict]
|
|||
|
|
) -> str:
|
|||
|
|
"""
|
|||
|
|
注入few-shot示例到提示词
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
prompt: 原始提示词
|
|||
|
|
examples: 示例列表,每项为 {"question": "...", "schema": "...", "sql": "..."}
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
增强后的提示词
|
|||
|
|
"""
|
|||
|
|
examples_text = []
|
|||
|
|
for ex in examples[:3]: # 最多3个示例
|
|||
|
|
examples_text.append(
|
|||
|
|
f"示例:\n问题:{ex['question']}\n"
|
|||
|
|
f"Schema: {ex['schema'][:200]}...\n"
|
|||
|
|
f"SQL: {ex['sql']}"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
examples_block = "\n\n".join(examples_text)
|
|||
|
|
|
|||
|
|
# 插入到提示词末尾(要求之前)
|
|||
|
|
return f"{prompt}\n\n参考示例:\n{examples_block}\n\n请生成SQL:"
|
|||
|
|
|
|||
|
|
def generate_with_reasoning(
|
|||
|
|
self,
|
|||
|
|
question: str,
|
|||
|
|
schema_str: str,
|
|||
|
|
dialect: Optional[str] = None
|
|||
|
|
) -> Tuple[str, str]:
|
|||
|
|
"""
|
|||
|
|
生成SQL并返回解释
|
|||
|
|
|
|||
|
|
Args:
|
|||
|
|
question: 用户问题
|
|||
|
|
schema_str: Schema描述
|
|||
|
|
dialect: SQL方言
|
|||
|
|
|
|||
|
|
Returns:
|
|||
|
|
(SQL语句, 解释)
|
|||
|
|
"""
|
|||
|
|
from config.prompts import SQL_GENERATOR_USER
|
|||
|
|
|
|||
|
|
dialect = dialect or self.dialect
|
|||
|
|
prompt = f"""
|
|||
|
|
{sql_generator_user.format(schema=schema_str, question=question, dialect=dialect)}
|
|||
|
|
|
|||
|
|
请同时输出SQL和简要解释(JSON格式):
|
|||
|
|
{{
|
|||
|
|
"sql": "SELECT ...",
|
|||
|
|
"explanation": "SQL逻辑说明"
|
|||
|
|
}}
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
try:
|
|||
|
|
response = self.agent.step(prompt)
|
|||
|
|
content = response.msg.content.strip()
|
|||
|
|
|
|||
|
|
# 尝试解析JSON
|
|||
|
|
try:
|
|||
|
|
if "```json" in content:
|
|||
|
|
start = content.find("```json") + 7
|
|||
|
|
end = content.find("```", start)
|
|||
|
|
content = content[start:end].strip()
|
|||
|
|
data = json.loads(content)
|
|||
|
|
return data.get("sql", ""), data.get("explanation", "")
|
|||
|
|
except (json.JSONDecodeError, ValueError):
|
|||
|
|
# 降级:提取SQL,解释为空
|
|||
|
|
return self._extract_sql(content), ""
|
|||
|
|
|
|||
|
|
except Exception as e:
|
|||
|
|
logger.error(f"生成失败: {e}")
|
|||
|
|
raise
|