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