feat:客服agent接入nl2sql

This commit is contained in:
2026-09-13 20:48:42 +08:00
parent b133dd58c4
commit dd281b3361
11 changed files with 778 additions and 18 deletions
+3
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
from config.database.milvus import client as milvus_client
from config.database.mysql import get_session_factory
from config.database.redis import client as redis_client
from config.settings import settings as app_settings
from service.client_agent.runtime import build_client_runtime
from service.customer_agent.config import DatabaseConfigProvider
from tool.llm import llm as llm_client
@@ -19,6 +20,8 @@ def build_default_runtime():
llm_client=llm_client,
config_getter=provider.get,
audit_writer=provider.write_audit,
db_session_factory=get_session_factory(),
schema_database=app_settings.mysql.database,
)
+147 -2
View File
@@ -7,15 +7,28 @@ import logging
import uuid
from types import SimpleNamespace
from nl2sql.contracts import DataQueryRequest
from nl2sql.history import archive_query_safely
from nl2sql.limits import QueryLimiter
from nl2sql.retrieval import retrieve_metadata
from nl2sql.runtime_config import runtime_config
from nl2sql.schema import load_authoritative_schema
from agent.client_agent.session import ClientSessionService
from rag.embedding import embed_texts
from rag.generation import generate_answer
from rag.intent import intent_recognize
from rag.retrieve import rag_retrieve
from service.customer_agent.chat import AnonymousCustomerAgent
from service.customer_agent.chat import (
AnonymousCustomerAgent,
DataQueryRejected,
)
from service.client_agent.memory_extractor import DialogueMemoryExtractor
from service.memory.facade import MemoryService
from service.memory.schemas import CustomerMemoryContext, MemoryUnitDTO, ShortTermMessage
from service.nl2sql.answer_render import render_query_answer
from service.nl2sql.customer_permission import load_customer_query_permission
from service.nl2sql.query_service import QueryServiceError, execute_query
_active_customer = contextvars.ContextVar("client_agent_customer", default=None)
@@ -91,7 +104,9 @@ class MemoryAwareClientAgent:
warnings_token = _active_warnings.set(warnings)
context_token = _active_memory_context.set(memory_context)
try:
result = await self.agent.handle(session_id, query, trace_id=trace_id)
result = await self.agent.handle(
session_id, query, trace_id=trace_id, customer_id=customer_id
)
if self.extractor is not None:
await self._save_candidates(
customer_id,
@@ -180,6 +195,122 @@ class MemoryAwareClientAgent:
raise TypeError(f"unsupported memory context item: {type(item).__name__}")
def _build_data_query(*, db_session_factory, milvus_client, llm_client, config_getter, redis, schema_database):
"""构造登录客户的 NL2SQL 数据查询依赖。
身份与行级范围由服务端强制注入:data_scope 只带当前登录客户自己的
customer_id,配合权限快照的 row_scopes 在 SQL 层兜底,杜绝水平越权。
"""
async def data_query(*, question, customer_id, session_id, trace_id):
async with db_session_factory() as db:
permission = await load_customer_query_permission(
db, customer_id, config_getter=config_getter
)
if not permission.get("can_query"):
raise DataQueryRejected("数据查询功能暂未开放,您可以先咨询基金知识或开户流程~")
limiter = QueryLimiter(redis)
if not await limiter.acquire(
customer_id,
daily_quota=permission.get("daily_quota", 0) or 20,
max_concurrent=1,
rate_limit=10,
):
raise DataQueryRejected("您今天的数据查询次数已达上限,请明天再来吧~")
try:
return await _execute_customer_query(
db=db,
permission=permission,
question=question,
customer_id=customer_id,
session_id=session_id,
trace_id=trace_id,
)
finally:
await limiter.release(customer_id)
async def _execute_customer_query(*, db, permission, question, customer_id, session_id, trace_id):
query_id = uuid.uuid4().hex
async def permission_loader(_user_id: int):
return permission
async def metadata_retriever(retrieval_question: str):
return await retrieve_metadata(
retrieval_question, milvus_client, top_k=runtime_config.retrieval_top_k
)
async def schema_loader(table_names: set[str], _permission: dict):
return await load_authoritative_schema(
db, database=schema_database, candidate_tables=table_names
)
request = DataQueryRequest(
question=question,
user_id=customer_id,
trace_id=trace_id,
session_id=session_id,
caller_agent="client_agent",
# 行级范围只允许是登录客户本人,不接受任何外部输入
data_scope={"customer_ids": [customer_id]},
include_sql=False,
max_rows=min(permission.get("max_rows") or 200, runtime_config.max_rows),
)
try:
result = await execute_query(
request,
session=db,
query_id=query_id,
permission_loader=permission_loader,
metadata_retriever=metadata_retriever,
schema_loader=schema_loader,
llm_client=llm_client,
masks=permission.get("masks") or {},
summary_llm=llm_client,
)
except QueryServiceError as exc:
await archive_query_safely(
db,
query_id=query_id,
user_id=customer_id,
question=question,
status="blocked",
error_message=str(exc),
trace_id=trace_id,
session_id=session_id,
caller_agent="client_agent",
)
# 不向用户暴露 SQL 和内部异常细节
raise DataQueryRejected(
"这个问题我暂时查不了,您可以换个问法,或者联系人工客服帮您处理~"
) from exc
await archive_query_safely(
db,
query_id=query_id,
user_id=customer_id,
question=question,
status="success",
row_count=result.row_count,
truncated=result.truncated,
elapsed_ms=result.elapsed_ms,
trace_id=trace_id,
session_id=session_id,
caller_agent="client_agent",
)
return {
"answer": render_query_answer(result),
"sources": [],
"query_id": result.query_id,
"row_count": result.row_count,
"truncated": result.truncated,
"chart": result.chart,
}
return data_query
def build_client_runtime(
*,
redis,
@@ -187,6 +318,8 @@ def build_client_runtime(
llm_client,
config_getter,
audit_writer,
db_session_factory=None,
schema_database=None,
memory_service=None,
memory_extractor=None,
):
@@ -233,6 +366,17 @@ def build_client_runtime(
config_getter=config_getter,
)
data_query = None
if db_session_factory is not None and schema_database:
data_query = _build_data_query(
db_session_factory=db_session_factory,
milvus_client=milvus_client,
llm_client=llm_client,
config_getter=config_getter,
redis=redis,
schema_database=schema_database,
)
agent = AnonymousCustomerAgent(
context=context,
rag_retrieve=retrieve,
@@ -240,6 +384,7 @@ def build_client_runtime(
generate_answer=generate,
audit_writer=audit_writer,
config_getter=config_getter,
data_query=data_query,
)
extractor = memory_extractor or DialogueMemoryExtractor(llm_client)
wrapped_agent = MemoryAwareClientAgent(
+83 -2
View File
@@ -16,6 +16,10 @@ class QueryTooLongError(ValueError):
pass
class DataQueryRejected(ValueError):
"""NL2SQL 数据查询被拒绝(未开放、无权限或配额不足),message 可直接回复用户。"""
async def _config(config_getter, key: str, default: str):
value = config_getter(key, default)
if isawaitable(value):
@@ -37,6 +41,7 @@ class AnonymousCustomerAgent:
generate_answer,
audit_writer,
config_getter,
data_query=None,
):
self.context = context
self.rag_retrieve = rag_retrieve
@@ -44,8 +49,18 @@ class AnonymousCustomerAgent:
self.generate_answer = generate_answer
self.audit_writer = audit_writer
self.config_getter = config_getter
# 可选 NL2SQL 数据查询依赖:签名 data_query(*, question, customer_id,
# session_id, trace_id) -> dict;匿名 runtime 不装配(None),行为不变。
self.data_query = data_query
async def handle(self, session_id: str, query: str, *, trace_id: str) -> dict:
async def handle(
self,
session_id: str,
query: str,
*,
trace_id: str,
customer_id: int | None = None,
) -> dict:
if len(query) > 2000:
raise QueryTooLongError("query长度不能超过2000字符")
# 先取历史再写入当前问题,保证意图识别拿到的历史不含本轮输入;取不到历史不阻断请求
@@ -68,6 +83,7 @@ class AnonymousCustomerAgent:
else:
intent, search_query = recognized, query
sources = []
data_query_meta = None
if intent == Intent.GUIDE_PURCHASE:
answer = await _config(
self.config_getter,
@@ -89,6 +105,21 @@ class AnonymousCustomerAgent:
)
elif intent == Intent.CHITCHAT:
answer = await self._chitchat(session_id)
elif intent == Intent.NL2SQL_REQUEST:
if self.data_query is None or customer_id is None:
# 匿名会话或未装配数据查询能力:引导登录,不触发任何数据库查询
answer = await _config(
self.config_getter,
"agent.customer.template.nl2sql_unavailable",
"数据查询功能需要登录后使用,请先登录再来问我您的持仓和交易信息~",
)
else:
answer, sources, data_query_meta = await self._run_data_query(
question=search_query,
customer_id=customer_id,
session_id=session_id,
trace_id=trace_id,
)
elif intent in (Intent.KNOWLEDGE_QA, Intent.COMPANY_INFO):
try:
# 用补全指代后的问题检索,省略主语的追问才能命中
@@ -129,13 +160,63 @@ class AnonymousCustomerAgent:
)
await self.context.append(session_id, "assistant", answer)
return {
result = {
"answer": answer,
"sources": sources,
"intent": intent.value,
"rewritten_query": search_query,
"trace_id": trace_id,
}
if data_query_meta is not None:
result["data_query"] = data_query_meta
return result
async def _run_data_query(
self,
*,
question: str,
customer_id: int,
session_id: str,
trace_id: str,
) -> tuple[str, list, dict]:
"""调用注入的 NL2SQL 数据查询能力,失败时统一降级为客服话术。"""
try:
payload = await _maybe_await(
self.data_query(
question=question,
customer_id=customer_id,
session_id=session_id,
trace_id=trace_id,
)
)
except DataQueryRejected as exc:
return str(exc), [], None
except Exception:
logger.exception(
"client data query failed: trace_id=%s session_id=%s customer_id=%s",
trace_id,
session_id,
customer_id,
)
answer = await _config(
self.config_getter,
"agent.customer.template.nl2sql_fallback",
"暂时无法完成数据查询,请稍后再试或联系人工客服。",
)
return answer, [], None
if not isinstance(payload, dict) or not str(payload.get("answer") or "").strip():
return await _config(
self.config_getter,
"agent.customer.template.nl2sql_fallback",
"暂时无法完成数据查询,请稍后再试或联系人工客服。",
), [], None
meta = {
key: payload[key]
for key in ("query_id", "row_count", "truncated", "chart")
if payload.get(key) is not None
}
return str(payload["answer"]), list(payload.get("sources") or []), meta or None
async def _chitchat(self, session_id: str) -> str:
"""带对话历史调用 LLM 做受限闲聊,失败时退回固定话术。"""
+53
View File
@@ -0,0 +1,53 @@
"""NL2SQL 查询结果 → 客服口吻回复的渲染器。
不调用 LLM:自然语言摘要由 execute_query 的 summary_llm 生成(result.summary),
这里只负责把摘要 + Markdown 表格组装成客服回复,保证确定性降级。
"""
from __future__ import annotations
from typing import Any
from nl2sql.contracts import DataQueryResult
# 表格最多渲染的行数:超出部分提示"仅展示前 N 条",避免回复过长
_MAX_TABLE_ROWS = 20
_EMPTY_ANSWER = "暂时没有查到相关数据,您可以换个问法,或者问我基金知识、开户流程~"
def render_markdown_table(columns: list[str], rows: list[dict[str, Any]]) -> str:
"""把结果行列渲染为 Markdown 表格;无数据返回空串。"""
if not columns or not rows:
return ""
shown = rows[:_MAX_TABLE_ROWS]
header = "| " + " | ".join(str(column) for column in columns) + " |"
separator = "| " + " | ".join("---" for _ in columns) + " |"
lines = [header, separator]
for row in shown:
cells = [str(row.get(column, "")) for column in columns]
lines.append("| " + " | ".join(cells) + " |")
return "\n".join(lines)
def render_query_answer(result: DataQueryResult) -> str:
"""组装最终客服回复:摘要开头 + 数据表格 + 截断/收尾提示。"""
if result.row_count == 0 or not result.rows:
return _EMPTY_ANSWER
parts: list[str] = []
summary = (result.summary or "").strip()
if summary:
parts.append(summary)
table = render_markdown_table(result.columns, result.rows)
if table:
parts.append(table)
if result.truncated or result.row_count > len(result.rows):
shown = min(len(result.rows), _MAX_TABLE_ROWS)
parts.append(f"结果较多,本次为您展示 {shown} 条(共 {result.row_count} 条),您可以缩小查询范围再看。")
if not summary and len(result.rows) <= _MAX_TABLE_ROWS:
parts.append(f"共为您查到 {result.row_count} 条记录。")
return "\n\n".join(part for part in parts if part).strip() or _EMPTY_ANSWER
+157
View File
@@ -0,0 +1,157 @@
"""登录客户(CUSTOMER)的 NL2SQL 权限快照服务。
与员工路径(permission_service.load_query_permission,按 nl2sql_query_role
配置)不同,客户权限不落库、不做管理后台:表白名单通过 sys_config 配置管理,
且只能从内置白名单中做"减法",行级范围由服务端强制注入 customer_ids,
保证客户永远只能查询自己的数据。
列级校验:快照时从 information_schema 加载白名单表的真实列清单写入
columns,validate_select_sql 据此在执行前拦截 LLM 幻觉列(避免把
Unknown column 错误漏到执行期)。
"""
from __future__ import annotations
import logging
from inspect import isawaitable
from sqlalchemy import bindparam, text
logger = logging.getLogger(__name__)
# 客户可查询的内置表白名单(配置只能在其中做减法,不能新增表)
CUSTOMER_DEFAULT_TABLES: tuple[str, ...] = (
"fin_holdings",
"fin_transaction",
"fin_product",
"fund_nav_history",
"fund_performance",
)
# 行级隔离:出现这些表的 SQL 会被强制注入 customer_id IN (<登录用户>) 条件
CUSTOMER_ROW_SCOPES: dict[str, dict[str, str]] = {
"fin_holdings": {"type": "customer_ids", "column": "customer_id"},
"fin_transaction": {"type": "customer_ids", "column": "customer_id"},
}
# 客户路径首期不开放敏感档案表;后续开放时在此配置 (table, column) -> mask_type
CUSTOMER_MASKS: dict[tuple[str, str], str] = {}
_TRUTHY = {"1", "true", "yes", "on"}
_COLUMNS_SQL = text(
"SELECT TABLE_NAME AS table_name, COLUMN_NAME AS column_name "
"FROM information_schema.columns "
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME IN :tables "
"ORDER BY TABLE_NAME, ORDINAL_POSITION"
).bindparams(bindparam("tables", expanding=True))
async def _load_real_columns(db, tables: set[str]) -> dict[str, set[str]] | None:
"""加载白名单表的真实列名;db 为 None 时返回 None(仅测试路径)。"""
if db is None:
return None
result = await db.execute(_COLUMNS_SQL, {"tables": sorted(tables)})
columns: dict[str, set[str]] = {}
for row in result.mappings():
table_name = str(row["table_name"] or "").strip()
column_name = str(row["column_name"] or "").strip()
if table_name and column_name:
columns.setdefault(table_name, set()).add(column_name)
return columns
def _denied_permission() -> dict:
return {
"can_query": False,
"role": "customer_self",
"tables": set(),
"columns": None,
"masks": {},
"row_scopes": {},
"max_rows": 0,
"daily_quota": 0,
}
async def _config(config_getter, key: str, default: str) -> str:
value = config_getter(key, default)
if isawaitable(value):
value = await value
if value is None or str(value).strip() == "":
return default
return str(value)
def _parse_allowed_tables(raw: str) -> set[str]:
"""解析表白名单配置;非法表名直接忽略,只允许内置白名单的子集。"""
known = set(CUSTOMER_DEFAULT_TABLES)
names = {
item.strip().lower()
for item in str(raw).replace(";", ",").replace(";", ",").split(",")
if item.strip()
}
tables = names & known
return tables
async def load_customer_query_permission(
db,
user_id: int,
*,
config_getter,
) -> dict:
"""每次请求重建客户权限快照。
客户身份已由 API 层(require_customer)和会话归属校验保证,
快照不依赖数据库中的角色配置;db 用于加载白名单表的真实列清单
(传入 None 时跳过列清单,columns 保持 None,仅限测试路径)。
"""
del user_id # 权限与具体请求上下文无关,签名对齐 execute_query 的 permission_loader
if config_getter is None:
return _denied_permission()
enabled = (await _config(config_getter, "nl2sql.customer.enabled", "true")).lower()
if enabled not in _TRUTHY:
return _denied_permission()
raw_tables = await _config(
config_getter,
"nl2sql.customer.allowed_tables",
",".join(CUSTOMER_DEFAULT_TABLES),
)
tables = _parse_allowed_tables(raw_tables)
if not tables:
return _denied_permission()
try:
max_rows = int(await _config(config_getter, "nl2sql.customer.max_rows", "200"))
daily_quota = int(
await _config(config_getter, "nl2sql.customer.daily_quota", "20")
)
except ValueError:
max_rows, daily_quota = 200, 20
max_rows = max(1, max_rows)
daily_quota = max(0, daily_quota)
# 列级校验用真实列清单:拦截 LLM 幻觉列,避免执行期 Unknown column。
# 信息读取失败时按"拒绝"处理(fail-closed),不让无列校验的快照放行。
try:
real_columns = await _load_real_columns(db, tables)
except Exception:
logger.exception("load customer nl2sql columns failed")
return _denied_permission()
return {
"can_query": True,
"role": "customer_self",
"tables": tables,
# 真实列清单(db=None 的测试路径保持 None = 不限列)
"columns": real_columns,
"masks": dict(CUSTOMER_MASKS),
"row_scopes": {
table: dict(scope)
for table, scope in CUSTOMER_ROW_SCOPES.items()
if table in tables
},
"max_rows": max_rows,
"daily_quota": daily_quota,
}
+12 -7
View File
@@ -88,13 +88,18 @@ async def query(
sort_by=request.sort_by,
sort_order=request.sort_order,
)
final_columns = {
table: set(columns)
for table, columns in (permission.get("columns") or {}).items()
}
for table, scope in (permission.get("row_scopes") or {}).items():
if scope.get("column"):
final_columns.setdefault(table, set()).add(scope["column"])
# columns 为 None 表示不做列级限制;仅当配置了列权限时才需要
# 保证行级范围列可访问。保持 dict(含空 dict)行为不变。
if permission.get("columns") is None:
final_columns = None
else:
final_columns = {
table: set(columns)
for table, columns in permission["columns"].items()
}
for table, scope in (permission.get("row_scopes") or {}).items():
if scope.get("column"):
final_columns.setdefault(table, set()).add(scope["column"])
return validate_select_sql(
option_sql,
authorized_tables=permission.get("tables", set()),