Files
ai-g3sb-backman2.0/agents/sql_generator.py
T
2026-04-10 16:52:07 +08:00

207 lines
5.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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