feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1,89 @@
|
||||
"""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"))
|
||||
Reference in New Issue
Block a user