Files
Mutual_Fund/nl2sql/sql_agent.py
T

90 lines
3.1 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.
"""NL2SQL SQL 生成器:只负责调用 Chat 模型并清理模型输出。"""
from __future__ import annotations
import json
import re
from dataclasses import dataclass
from typing import Any
from sqlglot import exp, parse
from sqlglot.errors import ParseError
from tool.llm import llm
from nl2sql.semantics import build_semantic_context
class SqlGenerationError(ValueError):
"""模型未返回可用的只读 SQL。"""
@dataclass(frozen=True)
class GeneratedSql:
sql: str
def _clean_model_output(output: str) -> str:
value = (output or "").strip()
value = re.sub(r"^```(?:sql)?\s*", "", value, flags=re.IGNORECASE)
value = re.sub(r"\s*```$", "", value).strip()
return value.removesuffix(";").strip()
def _build_prompt(
question: str,
schema: dict[str, Any],
few_shot: list[dict[str, Any]],
semantic_context: dict[str, Any] | None = None,
conversation_context: str = "",
) -> str:
return (
"你是基金业务数据库 SQL 生成器。\n"
"只根据提供的 Schema 生成一条 MySQL SELECT,禁止写操作、跨库访问和未提供的表字段。\n"
"只输出 SQL,不要输出解释、Markdown 或代码围栏。\n"
f"用户问题:{question}\n"
f"Schema:{json.dumps(schema, ensure_ascii=False, sort_keys=True)}\n"
f"业务语义:{json.dumps(semantic_context or {}, ensure_ascii=False, sort_keys=True)}\n"
f"会话上下文:{conversation_context}\n"
f"Few-shot:{json.dumps(few_shot or [], ensure_ascii=False, sort_keys=True)}"
)
async def generate_sql(
question: str,
schema: dict[str, Any],
*,
llm_client=llm,
few_shot: list[dict[str, Any]] | None = None,
semantic_context: dict[str, Any] | None = None,
conversation_context: str = "",
) -> GeneratedSql:
"""调用 Chat 模型生成 SQL,并在返回前确认其为单条 SELECT。"""
if not question or not question.strip():
raise SqlGenerationError("用户问题不能为空")
messages = [
{"role": "system", "content": "你必须严格遵守只输出单条 SELECT SQL。"},
{
"role": "user",
"content": _build_prompt(
question,
schema,
few_shot or [],
semantic_context or build_semantic_context(question, schema),
conversation_context,
),
},
]
try:
output = await llm_client.chat(messages, temperature=0)
except Exception as exc: # noqa: BLE001 统一收敛模型异常
raise SqlGenerationError("SQL 模型调用失败") from exc
sql = _clean_model_output(output)
if not sql:
raise SqlGenerationError("模型未返回 SQL")
try:
statements = parse(sql, read="mysql")
except ParseError as exc:
raise SqlGenerationError("模型返回的 SQL 无法解析") from exc
if len(statements) != 1 or not isinstance(statements[0], exp.Select):
raise SqlGenerationError("模型只允许返回单条 SELECT")
return GeneratedSql(sql=statements[0].sql(dialect="mysql"))