Files
group_fqcd_jr/nl2sql_yc.py
T
张胜宇 39ca52449a fix(W29-c): nl2sql_yc 的 GROUP BY 拼装缺陷(会签组 6)+ 更正一处误判
## 更正:上一个提交里有一句话是错的

docs/49 组 5 与 D2.1 §1.6 原写「nl2sql_yc.py 同样缺写意图守卫」—— 该表述有误。
立组 6 时按纪律**先取证、再立单**,实测 8 条写意图问句对它 **8/8 已被拒绝**
(它没有宽泛关键词路由:「删除所有客户的持仓记录」不命中它的「当前持仓」
⇒ 计划落 unknown ⇒ 被 _validate_plan 拒绝)⇒ **它不需要写意图守卫**。

三处文档的错误表述已当日更正并**保留更正痕迹**(docs/49 组 5 的更正块、
D2.1 §1.6 的更正块),不静默抹掉。取证日志见报告 §10.1。

## 顺带量出的真缺陷:GROUP BY 拼装

取证同时发现一条**完全正常**的只读问句直接报错:

    查询近30天净值 -> status=error
    (pymysql.err.OperationalError) (1055, "Expression #3 of SELECT list is not in
     GROUP BY clause and contains nonaggregated column 'jr_agent.n.nav'
     ... incompatible with sql_mode=only_full_group_by")

根因(_compile_sql):SELECT 里既有 dimensions(已进 GROUP BY)又有 metrics,
而 metrics 里的**非聚合**列(n.nav / m.close_price / h.market_value …)**没进 GROUP BY**。

修法(仅 +24 / −1 行,守住会签单的「最小化边界」):
① 没有聚合函数就不加 GROUP BY;② 有聚合时把**所有非聚合 select 列**一并纳入,
而不是只放 dimensions(漏掉非聚合 metric 正是 1055 的成因)。

## 影响面

踩:_mock_plan 的 4 条路由(收盘/行情、净值(已实证)、当前持仓、账户余额);
不踩:SUM / COUNT 聚合路由,以及 _offsite_mock_plan 全部 5 条(均无 dimensions)。
=> 场外主链路不受影响;受影响的是通用兜底路由的 4 类问句(演示与离线联调走的正是这条)。

## 验证

- 实测(真连库):净值/行情查询由 error 转 success;资金流水(SUM)GROUP BY 语义不变;error 归零
- 新增测试 5 条:**断言 SQL 结构而非「跑得通」** —— 缺陷只在真实 MySQL 的
  only_full_group_by 下现形,而单测跑在 sqlite(不做该检查);断言"能跑"在修复前后
  都是绿的,等于没测。其中一条是**通用不变量**:有 GROUP BY 时 SELECT 里每个非聚合列
  都必须在组内 —— 将来往 metric_map 加非聚合指标时会**先红**,而不是等真实库冒 1055
- offsite 相关 13 passed;全量 2545 passed / 3 skipped / 0 failed(2540 + 5)
- 金标 55 条与 W27 基线判分逐项零差异
- ruff:新文件 0 告警;nl2sql_yc.py **零新增告警**
  (既存 39 条 E501/F601/UP035/F401/B905 未清理 —— 不在会签单「最小化边界」内)

会签:docs/49(A-10)组 6 · 会签 19(☑ 2026-09-22 受理);白名单已登记 docs/48 类 3 表。
2026-09-22 10:31:24 +08:00

746 lines
36 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 MVP。
查询数据只允许使用参数化 SELECT;审计记录可通过回调写入
conversation_message.tool_calls。该文件可被运营和投顾 Agent 直接导入调用。
"""
from __future__ import annotations
import json
import logging
import os
import re
import uuid
from dataclasses import asdict, dataclass, field
from datetime import datetime, time, timedelta
from pathlib import Path
from typing import Any, Callable
logger = logging.getLogger(__name__)
ALLOWED_TABLES = {
"sys_customer_assignment", "fin_customer_profile", "fin_risk_assessment",
"fin_product", "fin_fee_rule", "fin_market_price", "fin_nav_history",
"fin_holding", "fin_transaction", "fin_sim_order", "fin_sim_account",
"fin_cash_ledger", "client_facing_content",
}
DOMAINS = {
"market_nav": {"fin_market_price", "fin_nav_history", "fin_product"},
"customer_risk": {
"sys_customer_assignment", "fin_customer_profile", "fin_risk_assessment",
},
"product_fee": {"fin_product", "fin_fee_rule"},
"trading_account": {
"fin_sim_order", "fin_transaction", "fin_sim_account",
"fin_cash_ledger", "fin_holding", "fin_product",
},
"client_content": {"client_facing_content"},
}
TABLE_COLUMNS = {
"sys_customer_assignment": {
"id", "customer_id", "employee_id", "employee_role", "assigned_at", "unassigned_at",
},
"fin_customer_profile": {
"customer_id", "trade_account", "real_name", "birth_date", "occupation",
"mobile_masked", "investor_type", "investment_horizon", "preferred_asset_class",
"trading_frequency", "last_active_at", "total_asset", "behavior_score",
"risk_tags", "opened_at", "updated_at",
},
"fin_risk_assessment": {
"id", "customer_id", "questionnaire_version", "answers", "total_score",
"investor_type", "assessed_at", "valid_until", "created_at",
},
"fin_product": {
"id", "product_code", "product_name", "exchange_code", "product_category",
"risk_level", "fund_manager", "currency", "lot_size", "price_tick",
"current_nav", "current_nav_at", "min_amount", "open_start_at", "open_end_at",
"open_period_start", "open_period_end", "transaction_fee_rate",
"single_investor_max_holding_ratio", "management_fee_rate", "custodian_fee_rate",
"risk_disclosure_required", "second_confirmation_required", "recording_required",
"status", "created_at", "updated_at",
},
"fin_fee_rule": {
"id", "rule_code", "product_id", "exchange_code", "order_side", "customer_tier",
"min_trade_amount", "max_trade_amount", "fee_rate", "minimum_fee", "fixed_fee",
"priority", "effective_from", "effective_until", "status", "created_at", "updated_at",
},
"fin_market_price": {
"id", "product_id", "trade_date", "open_price", "high_price", "low_price",
"close_price", "volume", "turnover_amount", "total_fund_shares", "source",
"source_updated_at", "created_at",
},
"fin_nav_history": {"id", "product_id", "nav_date", "nav", "created_at"},
"fin_holding": {
"id", "customer_id", "trade_account", "product_id", "total_quantity",
"shares", "available_quantity", "frozen_quantity", "average_cost", "cost_amount",
"market_value", "current_value", "profit_loss", "profit_loss_ratio", "status",
"first_acquired_at", "version", "updated_at",
},
"fin_transaction": {
"id", "transaction_no", "order_id", "work_order_id", "customer_id", "account_id",
"product_id", "order_side", "transaction_type", "executed_price", "nav", "executed_quantity",
"shares", "gross_amount", "amount", "fee_rule_id", "fee_rate_snapshot", "fee_amount",
"fee", "net_amount", "quote_at", "quote_source", "executed_at", "confirmed_at",
"confirmed_by", "auto_confirmed", "created_at",
},
"fin_sim_order": {
"id", "order_no", "customer_id", "account_id", "product_id", "order_side",
"price_type", "quantity", "limit_price", "quote_price", "quote_at", "quote_source",
"channel", "advisor_id", "filled_quantity", "average_executed_price", "status",
"risk_rule_hits", "risk_disclosure_ack_at", "second_confirmation_at",
"recording_reference", "ops_handler_id", "ops_handled_at", "compliance_handler_id",
"compliance_handled_at", "reject_reason", "submitted_at", "cancelled_at",
"created_at", "updated_at",
},
"fin_sim_account": {
"id", "account_no", "customer_id", "currency", "cash_balance", "available_cash",
"frozen_cash", "initial_balance", "status", "version", "created_at", "updated_at",
},
"fin_cash_ledger": {
"id", "ledger_no", "account_id", "transaction_id", "entry_type", "amount",
"balance_after", "available_cash_after", "frozen_cash_after", "idempotency_key",
"occurred_at", "created_at",
},
"client_facing_content": {
"id", "customer_id", "content_type", "draft_content", "generated_by_portal",
"review_status", "reviewer_user_id", "reviewed_at", "published_at",
"created_at", "updated_at",
},
}
JOIN_SQL = {
("fin_transaction", "fin_product"): "t.product_id = p.id",
("fin_holding", "fin_product"): "h.product_id = p.id",
("fin_market_price", "fin_product"): "m.product_id = p.id",
("fin_nav_history", "fin_product"): "n.product_id = p.id",
("fin_sim_order", "fin_product"): "o.product_id = p.id",
("fin_fee_rule", "fin_product"): "f.product_id = p.id",
("fin_transaction", "fin_customer_profile"): "t.customer_id = cp.customer_id",
("fin_holding", "fin_customer_profile"): "h.customer_id = cp.customer_id",
("fin_sim_account", "fin_customer_profile"): "a.customer_id = cp.customer_id",
("fin_cash_ledger", "fin_sim_account"): "l.account_id = a.id",
("fin_sim_account", "fin_customer_profile"): "a.customer_id = cp.customer_id",
("fin_cash_ledger", "fin_sim_account"): "l.account_id = a.id",
("fin_transaction", "fin_sim_account"): "t.account_id = a.id",
("fin_transaction", "sys_customer_assignment"): "t.customer_id = ca.customer_id",
("fin_holding", "sys_customer_assignment"): "h.customer_id = ca.customer_id",
("fin_customer_profile", "sys_customer_assignment"): "cp.customer_id = ca.customer_id",
}
def _load_local_config() -> dict[str, str]:
"""加载标准环境变量,并兼容 0901/.evn 的中文标签密钥格式。"""
values = dict(os.environ)
path = Path(__file__).with_name(".evn")
if not path.exists():
return values
raw_lines = [line.strip() for line in path.read_text(encoding="utf-8").splitlines() if line.strip()]
for line in raw_lines:
if "=" in line:
key, value = line.split("=", 1)
values.setdefault(key.strip(), value.strip().strip('"').strip("'"))
if raw_lines:
first_value = raw_lines[0].split(":", 1)[-1].strip() if ":" in raw_lines[0] else raw_lines[0]
values.setdefault("DEEPSEEK_API_KEY_NL2SQL", first_value)
if len(raw_lines) > 1:
second_value = raw_lines[1].split(":", 1)[-1].strip() if ":" in raw_lines[1] else raw_lines[1]
values.setdefault("ALIYUN_API_KEY", second_value)
return values
CONFIG = _load_local_config()
DATABASE_URL = CONFIG.get("DATABASE_URL", "")
LLM_BASE_URL = CONFIG.get("DEEPSEEK_BASE_URL") or "https://api.deepseek.com"
LLM_MODEL = CONFIG.get("DEEPSEEK_NL2SQL_MODEL") or CONFIG.get("DEEPSEEK_MODEL") or "deepseek-chat"
LLM_KEY = CONFIG.get("DEEPSEEK_API_KEY_NL2SQL") or CONFIG.get("DEEPSEEK_API_KEY", "")
@dataclass
class AuthContext:
user_id: int | None = None
roles: list[str] = field(default_factory=list)
allowed_domains: list[str] = field(default_factory=lambda: list(DOMAINS))
customer_scope: str = "all"
allowed_fields: list[str] | None = None
masked_fields: list[str] = field(default_factory=list)
max_rows: int = 500
max_query_seconds: int = 10
@dataclass
class QueryRequest:
question: str
auth_context: AuthContext
conversation_id: str | None = None
request_id: str | None = None
timezone: str = "Asia/Shanghai"
confirmation: str | None = None
@dataclass
class QueryPlan:
intent: str
domains: list[str]
tables: list[str]
time_mode: str = "none"
time_column: str | None = None
start: str | None = None
end: str | None = None
metrics: list[str] = field(default_factory=list)
dimensions: list[str] = field(default_factory=list)
filters: list[dict[str, Any]] = field(default_factory=list)
sort: list[dict[str, str]] = field(default_factory=list)
limit: int = 50
confidence: float = 0.0
needs_confirmation: bool = False
confirmation_question: str | None = None
unsupported_reason: str | None = None
def _json_object(text_value: str) -> dict[str, Any]:
cleaned = re.sub(r"```(?:json)?|```", "", text_value or "").strip()
match = re.search(r"\{.*\}", cleaned, flags=re.S)
if not match:
raise ValueError("模型未返回结构化查询计划")
value = json.loads(match.group(0))
if not isinstance(value, dict):
raise ValueError("查询计划必须是 JSON 对象")
return value
def _days_range(days: int) -> tuple[str, str]:
end = datetime.now()
start = end - timedelta(days=days)
return start.strftime("%Y-%m-%d %H:%M:%S"), end.strftime("%Y-%m-%d %H:%M:%S")
def _catalog_text(tables: set[str]) -> str:
lines = ["只允许使用以下表和字段:"]
for table in sorted(tables):
lines.append(f"{table}: {', '.join(sorted(TABLE_COLUMNS[table]))}")
lines.append("指标口径:交易金额=gross_amount;净交易金额=net_amount;收盘价=close_price;基金净值=nav;")
lines.append("成交价格=executed_price;当前持仓市值=market_value;历史资金变化=fin_cash_ledger.amount。")
return "\n".join(lines)
def _route_keywords(question: str) -> list[str]:
routes = []
if any(word in question for word in ("行情", "收盘", "开盘", "最高价", "最低价")):
routes.append("market_nav")
if "净值" in question:
routes.append("market_nav")
if any(word in question for word in ("客户", "风险", "测评", "画像", "投顾", "运营")):
routes.append("customer_risk")
if any(word in question for word in ("产品", "基金", "费率", "适配", "风险等级")):
routes.append("product_fee")
if any(word in question for word in ("委托", "成交", "交易", "持仓", "账户", "资金", "现金", "盈亏")):
routes.append("trading_account")
if any(word in question for word in ("内容", "审核", "发布")):
routes.append("client_content")
return list(dict.fromkeys(routes)) or ["trading_account"]
def _llm_plan(request: QueryRequest, domains: list[str]) -> dict[str, Any]:
try:
from openai import OpenAI
except ImportError as exc:
raise RuntimeError("缺少 openai 依赖,无法调用模型") from exc
tables = set().union(*(DOMAINS[d] for d in domains))
system = f"""你是金融查询计划生成器。只输出 JSON,不输出解释。
用户身份已由后端确认,customer_scope={request.auth_context.customer_scope}。
{_catalog_text(tables)}
返回字段:intent, domains, tables, time_mode, time_column, start, end, metrics,
dimensions, filters, sort, limit, confidence, needs_confirmation, confirmation_question。
只从给定表和字段中选择;无法准确判断时降低 confidence。历史时点查询不能使用当前
fin_holding、fin_sim_account 或当前归属表冒充历史数据。"""
client = OpenAI(api_key=LLM_KEY, base_url=LLM_BASE_URL)
response = client.chat.completions.create(
model=LLM_MODEL,
messages=[{"role": "system", "content": system}, {"role": "user", "content": request.question}],
temperature=0,
response_format={"type": "json_object"},
timeout=request.auth_context.max_query_seconds,
)
return _json_object(response.choices[0].message.content)
def _normalize_plan(raw: dict[str, Any], request: QueryRequest, domains: list[str]) -> QueryPlan:
tables = [name for name in raw.get("tables", []) if name in ALLOWED_TABLES]
if not tables:
tables = sorted(set().union(*(DOMAINS[d] for d in domains)))
confidence = float(raw.get("confidence", 0.0) or 0.0)
confidence = max(0.0, min(confidence, 1.0))
plan = QueryPlan(
intent=str(raw.get("intent") or "unknown"),
domains=[d for d in raw.get("domains", domains) if d in DOMAINS] or domains,
tables=tables,
time_mode=str(raw.get("time_mode") or raw.get("temporal", {}).get("mode") or "none"),
time_column=raw.get("time_column") or raw.get("temporal", {}).get("time_column"),
start=raw.get("start") or raw.get("temporal", {}).get("start"),
end=raw.get("end") or raw.get("temporal", {}).get("end"),
metrics=list(raw.get("metrics") or []),
dimensions=list(raw.get("dimensions") or []),
filters=list(raw.get("filters") or []),
sort=list(raw.get("sort") or []),
limit=min(max(int(raw.get("limit", 50) or 50), 1), request.auth_context.max_rows),
confidence=confidence,
needs_confirmation=bool(raw.get("needs_confirmation", False)),
confirmation_question=raw.get("confirmation_question"),
unsupported_reason=raw.get("unsupported_reason"),
)
if plan.confidence < 0.60:
plan.needs_confirmation = True
plan.confirmation_question = plan.confirmation_question or "请补充客户、产品、时间范围或指标口径。"
elif plan.confidence < 0.85:
plan.needs_confirmation = True
plan.confirmation_question = plan.confirmation_question or "请确认我对查询范围和指标口径的理解是否正确。"
return plan
def _validate_plan(plan: QueryPlan, auth: AuthContext) -> tuple[bool, str]:
allowed_tables = set().union(*(DOMAINS[d] for d in auth.allowed_domains if d in DOMAINS))
if not set(plan.tables).issubset(allowed_tables):
return False, "查询包含当前角色未授权的数据表"
if len(plan.tables) > 8:
return False, "查询涉及的数据表过多"
if plan.time_mode in {"range", "as_of"} and not plan.time_column:
return False, "历史查询缺少时间字段"
temporal_current = {"fin_holding", "fin_sim_account", "sys_customer_assignment"}
if plan.time_mode == "as_of" and temporal_current.intersection(plan.tables):
return False, "当前表缺少历史时点来源,暂不支持该历史时点查询"
if plan.start and plan.end and plan.start > plan.end:
return False, "查询开始时间不能晚于结束时间"
return True, "计划校验通过"
def _field_allowed(table: str, field_name: str, auth: AuthContext) -> bool:
if field_name not in TABLE_COLUMNS.get(table, set()):
return False
if auth.allowed_fields is None:
return True
return f"{table}.{field_name}" in auth.allowed_fields or field_name in auth.allowed_fields
def _safe_sql_check(sql: str, params: dict[str, Any], plan: QueryPlan, auth: AuthContext) -> tuple[bool, str]:
normalized = re.sub(r"\s+", " ", sql.strip())
upper = normalized.upper()
if not upper.startswith("SELECT "):
return False, "仅支持 SELECT 查询"
if ";" in normalized.rstrip(";"):
return False, "禁止执行多语句"
if re.search(r"\b(INSERT|UPDATE|DELETE|DROP|ALTER|TRUNCATE|CREATE|GRANT|REVOKE|CALL|INTO\s+OUTFILE)\b", upper):
return False, "检测到禁止的数据库操作"
if "*" in normalized:
return False, "禁止使用 SELECT *"
table_refs = set(re.findall(r"\b(?:FROM|JOIN)\s+([A-Za-z_][A-Za-z0-9_]*)", normalized, flags=re.I))
if not table_refs.issubset(set(plan.tables)):
return False, "SQL 使用了计划外的数据表"
for table in table_refs:
aliases = re.findall(rf"\b{re.escape(table)}\s+(?:AS\s+)?([A-Za-z_][A-Za-z0-9_]*)", normalized, flags=re.I)
del aliases
if plan.time_mode in {"range", "as_of"} and plan.time_column and plan.time_column not in normalized:
return False, "历史查询缺少计划中的时间条件"
if len(params) > 30:
return False, "查询参数过多"
return True, "SQL 安全校验通过"
#: 聚合函数探测(`W29-c` · 会签项 19)。用途有二:
#: ① 判断 `GROUP BY` 是否**必要** —— 没有聚合就不该分组;
#: ② 判断哪些 select 列**必须**进 `GROUP BY` —— 非聚合的那些。
_AGGREGATE_CALL = re.compile(r"\b(SUM|COUNT|AVG|MIN|MAX)\s*\(", re.IGNORECASE)
def _alias(table: str) -> str:
return {
"fin_transaction": "t", "fin_product": "p", "fin_holding": "h",
"fin_market_price": "m", "fin_nav_history": "n", "fin_sim_order": "o",
"fin_sim_account": "a", "fin_cash_ledger": "l",
"fin_customer_profile": "cp", "sys_customer_assignment": "ca",
"fin_risk_assessment": "ra", "fin_fee_rule": "f",
"client_facing_content": "c",
}.get(table, table[:1])
def _compile_sql(plan: QueryPlan, auth: AuthContext) -> tuple[str, dict[str, Any]]:
if plan.unsupported_reason:
raise ValueError(plan.unsupported_reason)
primary = plan.tables[0]
alias = _alias(primary)
select_parts = []
group_parts = []
for dimension in plan.dimensions:
if "." not in dimension:
dimension = f"{primary}.{dimension}"
table, column = dimension.split(".", 1)
if not _field_allowed(table, column, auth):
raise PermissionError(f"字段未授权:{dimension}")
select_parts.append(f"{_alias(table)}.{column} AS {column}")
group_parts.append(f"{_alias(table)}.{column}")
metric_map = {
"交易金额": "SUM(t.gross_amount) AS gross_amount",
"gross_amount": "SUM(t.gross_amount) AS gross_amount",
"净交易金额": "SUM(t.net_amount) AS net_amount",
"net_amount": "SUM(t.net_amount) AS net_amount",
"成交数量": "SUM(t.executed_quantity) AS executed_quantity",
"持仓数量": "h.total_quantity AS total_quantity",
"持仓市值": "h.market_value AS market_value",
"浮动盈亏": "h.profit_loss AS profit_loss",
"可用份额": "h.available_quantity AS available_quantity",
"可用现金": "a.available_cash AS available_cash",
"现金余额": "a.cash_balance AS cash_balance",
"历史资金变化": "SUM(l.amount) AS cash_change",
"收盘价": "m.close_price AS close_price",
"基金总份额": "m.total_fund_shares AS total_fund_shares",
"基金净值": "n.nav AS nav",
"成交价格": "t.executed_price AS executed_price",
"客户数": "COUNT(DISTINCT cp.customer_id) AS customer_count",
}
for metric in plan.metrics or ["客户数"]:
expression = metric_map.get(str(metric))
if not expression:
raise ValueError(f"暂不支持指标:{metric}")
select_parts.append(expression)
if not select_parts:
raise ValueError("查询计划没有可返回的字段")
params: dict[str, Any] = {}
where = ["1=1"]
if primary == "fin_transaction":
where.append("t.order_side IN ('买入', '卖出')")
if plan.time_column and plan.start:
params["start_time"] = plan.start
where.append(f"{alias}.{plan.time_column} >= :start_time")
if plan.time_column and plan.end:
params["end_time"] = plan.end
where.append(f"{alias}.{plan.time_column} <= :end_time")
for index, item in enumerate(plan.filters):
field_name = item.get("field")
operator = str(item.get("operator", "=")).upper()
if not isinstance(field_name, str) or "." not in field_name or operator not in {"=", "!=", ">", ">=", "<", "<=", "LIKE"}:
raise ValueError("查询筛选条件不合法")
table, column = field_name.split(".", 1)
if not _field_allowed(table, column, auth):
raise PermissionError(f"字段未授权:{field_name}")
key = f"filter_{index}"
params[key] = item.get("value")
where.append(f"{_alias(table)}.{column} {operator} :{key}")
from_sql = f"{primary} {alias}"
pending = set(plan.tables) - {primary}
joined = {primary}
while pending:
progress = False
for table in sorted(pending):
join = None
for existing in joined:
join = JOIN_SQL.get((existing, table)) or JOIN_SQL.get((table, existing))
if join:
break
if not join:
continue
from_sql += f" JOIN {table} {_alias(table)} ON {join}"
joined.add(table)
pending.remove(table)
progress = True
if not progress:
break
if pending:
raise ValueError(f"缺少合法 Join 路径:{', '.join(sorted(pending))}")
sql = f"SELECT {', '.join(select_parts)} FROM {from_sql} WHERE {' AND '.join(where)}"
# `W29-c`(会签项 19):`GROUP BY` 只在**确实需要分组**时才加。
#
# 原实现只要 `dimensions` 非空就加 `GROUP BY`,但 `SELECT` 里还有 `metrics` ——
# 其中**非聚合**的那些(`n.nav` / `h.market_value` / `m.close_price` …)没被放进
# `GROUP BY`,MySQL `only_full_group_by`(本项目默认 `sql_mode`)直接拒绝:
# ERROR 1055 Expression #3 of SELECT list is not in GROUP BY clause ...
# 症状是**一条完全正常的只读问句返回 `status=error`**(实测「查询近30天净值」必现)。
#
# 两条修正:
# ① 没有聚合函数就不需要分组 —— 加了只会引入这类错误;
# ② 有聚合时,把**所有非聚合的 select 列**一并纳入 `GROUP BY`,
# 而不是只放 `dimensions`(漏掉非聚合 metric 正是 1055 的成因)。
if group_parts and any(_AGGREGATE_CALL.search(part) for part in select_parts):
for part in select_parts:
expression = part.split(" AS ")[0]
if _AGGREGATE_CALL.search(part) or expression in group_parts:
continue
group_parts.append(expression)
sql += f" GROUP BY {', '.join(group_parts)}"
if plan.sort:
sort_items = []
for item in plan.sort:
field = str(item.get("field", "")).split(".")[-1]
direction = "DESC" if str(item.get("direction", "desc")).lower() == "desc" else "ASC"
if field not in {part.split(" AS ")[-1] for part in select_parts}:
continue
sort_items.append(f"{field} {direction}")
if sort_items:
sql += " ORDER BY " + ", ".join(sort_items)
sql += f" LIMIT {min(plan.limit, auth.max_rows)}"
return sql, params
def _mock_plan(request: QueryRequest, domains: list[str]) -> dict[str, Any]:
"""无模型或测试时的保守规则,便于单元测试和离线联调。"""
question = request.question
if "近30天" in question or "最近30天" in question:
start, end = _days_range(30)
else:
start = end = None
if "收盘" in question or "行情" in question:
return {"intent": "market_price", "domains": ["market_nav"], "tables": ["fin_market_price", "fin_product"],
"time_mode": "range", "time_column": "trade_date", "start": start, "end": end,
"metrics": ["收盘价"], "dimensions": ["fin_product.product_name", "fin_market_price.trade_date"],
"confidence": 0.88}
if "净值" in question:
return {"intent": "nav_history", "domains": ["market_nav"], "tables": ["fin_nav_history", "fin_product"],
"time_mode": "range", "time_column": "nav_date", "start": start, "end": end,
"metrics": ["基金净值"], "dimensions": ["fin_product.product_name", "fin_nav_history.nav_date"],
"confidence": 0.88}
if "当前持仓" in question or "持仓市值" in question:
return {"intent": "current_holding", "domains": ["customer_risk", "trading_account"],
"tables": ["fin_holding", "fin_product", "fin_customer_profile"],
"metrics": ["持仓市值"], "dimensions": ["fin_customer_profile.real_name", "fin_product.product_name"],
"confidence": 0.87}
if "账户余额" in question or "可用现金" in question:
return {"intent": "current_account", "domains": ["trading_account"],
"tables": ["fin_sim_account", "fin_customer_profile"], "metrics": ["可用现金"],
"dimensions": ["fin_customer_profile.real_name"], "confidence": 0.87}
if "资金变动" in question or "资金流水" in question:
return {"intent": "cash_ledger", "domains": ["trading_account"],
"tables": ["fin_cash_ledger", "fin_sim_account"], "time_mode": "range",
"time_column": "occurred_at", "start": start, "end": end,
"metrics": ["历史资金变化"], "dimensions": ["fin_cash_ledger.occurred_at"], "confidence": 0.87}
return {"intent": "unknown", "domains": domains, "tables": sorted(set().union(*(DOMAINS[d] for d in domains))),
"confidence": 0.45, "needs_confirmation": True,
"confirmation_question": "请明确要查询的业务对象、指标和时间范围。"}
def _offsite_mock_plan(request: QueryRequest, domains: list[str]) -> dict[str, Any] | None:
"""为场外核对补充确定性查询计划,仍复用统一 SQL 安全校验。"""
question = request.question
if "基金代码" not in question:
return None
fund_match = re.search(r"基金代码(?:为|是)?\s*([A-Za-z0-9_-]+)", question)
if fund_match is None:
return None
filters: list[dict[str, Any]] = [{
"field": "fin_product.product_code",
"operator": "=",
"value": fund_match.group(1),
}]
account_match = re.search(r"账户标识(?:为|是)?\s*([^,,。;;\s]+)", question)
account_filter = {
"field": "fin_holding.trade_account",
"operator": "=",
"value": account_match.group(1),
} if account_match is not None else None
common = {
"domains": ["market_nav", "trading_account"],
"filters": filters,
"confidence": 0.99,
"needs_confirmation": False,
"limit": 1,
}
if "最新总份额" in question and "申请前持有份额" in question:
holding_filters = filters + ([account_filter] if account_filter else [])
return {
**common,
"intent": "offsite_subscription_holding_check",
"tables": ["fin_market_price", "fin_nav_history", "fin_holding", "fin_product"],
"metrics": ["基金总份额", "基金净值", "持仓数量"],
"filters": holding_filters,
}
if "最新总份额" in question and "最新净值" in question:
return {
**common,
"intent": "offsite_fund_limit_check",
"tables": ["fin_market_price", "fin_nav_history", "fin_product"],
"metrics": ["基金总份额", "基金净值"],
}
if "当前最新可用份额" in question or "可用份额" in question:
holding_filters = filters + ([account_filter] if account_filter else [])
return {
**common,
"intent": "offsite_redemption_available_check",
"tables": ["fin_holding", "fin_product"],
"metrics": ["可用份额"],
"filters": holding_filters,
}
if "最新净值" in question:
return {
**common,
"intent": "offsite_latest_nav",
"tables": ["fin_nav_history", "fin_product"],
"metrics": ["基金净值"],
}
if "最新总份额" in question:
return {
**common,
"intent": "offsite_latest_total_shares",
"tables": ["fin_market_price", "fin_product"],
"metrics": ["基金总份额"],
}
return None
def _audit_payload(request: QueryRequest, plan: QueryPlan, sql: str | None, params: dict[str, Any],
validation: dict[str, Any], execution: dict[str, Any]) -> dict[str, Any]:
return {
"tool": "nl2sql", "engine_version": "mvp-1", "request_id": request.request_id,
"conversation_id": request.conversation_id, "intent": plan.intent,
"confidence": plan.confidence, "domains": plan.domains,
"query_plan": asdict(plan), "generated_sql": sql,
"parameters": {key: "<redacted>" if "password" in key.lower() else value for key, value in params.items()},
"authorized_context": {"user_id": request.auth_context.user_id, "roles": request.auth_context.roles,
"customer_scope": request.auth_context.customer_scope},
"validation": validation, "execution": execution,
"created_at": datetime.now().isoformat(timespec="seconds"),
}
def write_conversation_tool_calls(connection: Any, request: QueryRequest, audit: dict[str, Any]) -> None:
"""将审计对象写入已有 conversation_message;宿主也可使用 audit_writer 自行落库。"""
from sqlalchemy import text
connection.execute(
text("""INSERT INTO conversation_message
(session_id, customer_id, portal, role, content, tool_calls, intent, confidence, created_at)
VALUES (:session_id, :customer_id, :portal, 'assistant', :content, :tool_calls,
:intent, :confidence, :created_at)"""),
{
"session_id": request.conversation_id or request.request_id,
"customer_id": request.auth_context.user_id,
"portal": "nl2sql",
"content": audit.get("execution", {}).get("status", "nl2sql"),
"tool_calls": json.dumps(audit, ensure_ascii=False, default=str),
"intent": audit.get("intent"),
"confidence": audit.get("confidence"),
"created_at": datetime.now(),
},
)
def query(request: QueryRequest, *, db_engine: Any = None,
audit_writer: Callable[[dict[str, Any]], None] | None = None,
use_llm: bool = True, persist_audit: bool = False) -> dict[str, Any]:
"""统一入口:返回结构化结果,供运营和投顾 Agent 直接调用。"""
if not request.question or not request.question.strip():
return {"status": "rejected", "message": "查询问题不能为空"}
request.request_id = request.request_id or str(uuid.uuid4())
domains = _route_keywords(request.question)
try:
raw = (
_llm_plan(request, domains)
if use_llm and LLM_KEY
else _offsite_mock_plan(request, domains) or _mock_plan(request, domains)
)
plan = _normalize_plan(raw, request, domains)
if request.confirmation and plan.needs_confirmation:
plan.needs_confirmation = False
plan.confidence = max(plan.confidence, 0.85)
valid, message = _validate_plan(plan, request.auth_context)
if not valid:
audit = _audit_payload(request, plan, None, {}, {"valid": False, "message": message}, {"status": "rejected"})
if audit_writer:
audit_writer(audit)
if persist_audit and db_engine is not None:
with db_engine.begin() as connection:
write_conversation_tool_calls(connection, request, audit)
return {"status": "rejected", "message": message, "audit": audit}
if plan.needs_confirmation:
audit = _audit_payload(request, plan, None, {}, {"valid": True, "message": "等待确认"}, {"status": "waiting"})
if audit_writer:
audit_writer(audit)
if persist_audit and db_engine is not None:
with db_engine.begin() as connection:
write_conversation_tool_calls(connection, request, audit)
return {"status": "need_confirmation", "message": plan.confirmation_question, "query_plan": asdict(plan), "audit": audit}
sql, params = _compile_sql(plan, request.auth_context)
safe, safety_message = _safe_sql_check(sql, params, plan, request.auth_context)
if not safe:
raise PermissionError(safety_message)
if db_engine is None:
execution = {"status": "not_executed", "reason": "未提供数据库连接,仅返回已校验 SQL"}
audit = _audit_payload(request, plan, sql, params, {"valid": True, "message": safety_message}, execution)
if audit_writer:
audit_writer(audit)
if persist_audit and db_engine is not None:
with db_engine.begin() as connection:
write_conversation_tool_calls(connection, request, audit)
return {"status": "ready", "sql": sql, "parameters": params, "query_plan": asdict(plan), "audit": audit}
from sqlalchemy import text
with db_engine.connect() as connection:
result = connection.execute(text(sql), params)
columns = list(result.keys())
rows = [dict(zip(columns, row)) for row in result.fetchmany(request.auth_context.max_rows)]
execution = {"status": "success", "row_count": len(rows), "truncated": len(rows) >= request.auth_context.max_rows}
audit = _audit_payload(request, plan, sql, params, {"valid": True, "message": safety_message}, execution)
if audit_writer:
audit_writer(audit)
if persist_audit:
with db_engine.begin() as connection:
write_conversation_tool_calls(connection, request, audit)
return {"status": "success", "data": {"total": len(rows), "rows": rows},
"query_plan": asdict(plan), "sql": sql, "audit": audit}
except (ValueError, PermissionError, RuntimeError) as exc:
logger.warning("NL2SQL业务失败:%s", exc)
return {"status": "error", "message": str(exc), "request_id": request.request_id}
except Exception:
logger.exception("NL2SQL执行失败")
return {"status": "error", "message": "查询执行失败,请稍后重试", "request_id": request.request_id}
def build_request(question: str, auth_context: dict[str, Any], **kwargs: Any) -> QueryRequest:
"""将 Agent 的字典请求转换为统一请求对象。"""
return QueryRequest(question=question, auth_context=AuthContext(**auth_context), **kwargs)
def query_dict(question: str, auth_context: dict[str, Any], *,
db_engine: Any = None,
audit_writer: Callable[[dict[str, Any]], None] | None = None,
use_llm: bool = True, persist_audit: bool = False,
**request_kwargs: Any) -> dict[str, Any]:
"""给 Agent 使用的字典式快捷入口。"""
return query(
build_request(question, auth_context, **request_kwargs),
db_engine=db_engine,
audit_writer=audit_writer,
use_llm=use_llm,
persist_audit=persist_audit,
)
def get_tool_definition() -> dict[str, Any]:
"""返回可注册到运营或投顾 Agent 的统一工具定义。"""
return {
"name": "financial_nl2sql",
"description": "对金融业务数据执行只读自然语言查询;低置信度时先确认。",
"input_schema": {
"type": "object",
"required": ["question", "auth_context"],
"properties": {
"question": {"type": "string"},
"auth_context": {
"type": "object",
"required": ["roles", "customer_scope"],
"properties": {
"user_id": {"type": ["integer", "null"]},
"roles": {"type": "array", "items": {"type": "string"}},
"customer_scope": {"type": "string", "enum": ["self", "own_customers", "all"]},
"allowed_domains": {"type": "array", "items": {"type": "string"}},
},
},
"conversation_id": {"type": ["string", "null"]},
"request_id": {"type": ["string", "null"]},
"timezone": {"type": "string"},
"confirmation": {"type": ["string", "null"]},
},
},
}
if __name__ == "__main__":
logging.basicConfig(level=logging.INFO)
demo = build_request("查询最近30天的行情收盘价", {"roles": ["advisor"], "customer_scope": "all"})
print(json.dumps(query(demo, use_llm=False), ensure_ascii=False, indent=2, default=str))