649 lines
32 KiB
Python
649 lines
32 KiB
Python
"""金融 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 安全校验通过"
|
||
|
||
|
||
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",
|
||
"可用现金": "a.available_cash AS available_cash",
|
||
"现金余额": "a.cash_balance AS cash_balance",
|
||
"历史资金变化": "SUM(l.amount) AS cash_change",
|
||
"收盘价": "m.close_price AS close_price",
|
||
"基金净值": "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)}"
|
||
if group_parts:
|
||
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 _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 _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))
|