feat:客服agent接入nl2sql
This commit is contained in:
@@ -72,6 +72,7 @@ data/output/
|
||||
*.xlsx
|
||||
*.parquet
|
||||
docs/
|
||||
.workbuddy/
|
||||
|
||||
# 本地临时需求文档
|
||||
DK2_客服Agent模块完整开发计划(v1.1).md
|
||||
|
||||
+124
-1
@@ -2,18 +2,127 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from typing import Any
|
||||
|
||||
from nl2sql.embedding import embed_texts
|
||||
from nl2sql.metadata import build_metadata_chunks
|
||||
from nl2sql.milvus_collections import NL2SQL_COLLECTION
|
||||
|
||||
# ORM 已定义但数据库尚未建表的表,在表说明后追加此标注,
|
||||
# 让 NL2SQL 生成侧知道该表暂不可查询(权威 Schema 校验也会兜底拒绝)。
|
||||
MISSING_TABLE_MARKER = "(注意:当前数据库中尚未建表)"
|
||||
|
||||
|
||||
def _escape_filter_value(value: str) -> str:
|
||||
return value.replace("\\", "\\\\").replace('"', '\\"')
|
||||
|
||||
|
||||
def _row_table_name(row: dict[str, Any]) -> str:
|
||||
"""兼容 information_schema 大小写键名,取行内表名。"""
|
||||
value = (
|
||||
row.get("table_name")
|
||||
or row.get("TABLE_NAME")
|
||||
or row.get("Table_name")
|
||||
or ""
|
||||
)
|
||||
return str(value).strip()
|
||||
|
||||
|
||||
def _filter_rows_by_tables(
|
||||
rows: list[dict[str, Any]], allowed_tables: set[str]
|
||||
) -> list[dict[str, Any]]:
|
||||
return [row for row in rows if _row_table_name(row) in allowed_tables]
|
||||
|
||||
|
||||
async def _delete_tables_outside_allowlist(
|
||||
milvus_client, allowed_tables: set[str]
|
||||
) -> list[str]:
|
||||
"""删除集合中不在允许名单内的表元数据 chunk,返回被清理的表名。"""
|
||||
existing_rows = await milvus_client.query(
|
||||
collection_name=NL2SQL_COLLECTION,
|
||||
filter="is_valid == true",
|
||||
output_fields=["table_name"],
|
||||
limit=16384,
|
||||
)
|
||||
existing_tables = {
|
||||
str(row.get("table_name") or "").strip()
|
||||
for row in existing_rows or []
|
||||
if str(row.get("table_name") or "").strip()
|
||||
}
|
||||
removed: list[str] = []
|
||||
for table_name in sorted(existing_tables - allowed_tables):
|
||||
await milvus_client.delete(
|
||||
collection_name=NL2SQL_COLLECTION,
|
||||
filter=f'table_name == "{_escape_filter_value(table_name)}"',
|
||||
)
|
||||
removed.append(table_name)
|
||||
return removed
|
||||
|
||||
|
||||
def build_orm_metadata_rows(
|
||||
tables: Iterable[Any],
|
||||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
"""把 SQLAlchemy Table 定义合成 information_schema 风格的元数据行。
|
||||
|
||||
用于 model/ 已定义、但数据库尚未建表的表:
|
||||
- 表注释取 ORM ``__table_args__`` 的 comment(ORM 无则空串);
|
||||
- 字段注释取列 comment(ORM 通常未标注,则为空串);
|
||||
- 字段类型/可空性从 ORM 列定义推导。
|
||||
"""
|
||||
table_rows: list[dict[str, Any]] = []
|
||||
column_rows: list[dict[str, Any]] = []
|
||||
for table in tables:
|
||||
table_rows.append(
|
||||
{
|
||||
"TABLE_NAME": str(table.name).strip(),
|
||||
"TABLE_COMMENT": str(table.comment or "").strip(),
|
||||
"TABLE_TYPE": "BASE TABLE",
|
||||
}
|
||||
)
|
||||
for position, column in enumerate(table.columns, start=1):
|
||||
column_rows.append(
|
||||
{
|
||||
"TABLE_NAME": str(table.name).strip(),
|
||||
"COLUMN_NAME": str(column.name).strip(),
|
||||
"COLUMN_COMMENT": str(column.comment or "").strip(),
|
||||
"DATA_TYPE": str(column.type).strip().lower(),
|
||||
"IS_NULLABLE": "YES" if column.nullable else "NO",
|
||||
"ORDINAL_POSITION": position,
|
||||
}
|
||||
)
|
||||
return table_rows, column_rows
|
||||
|
||||
|
||||
def merge_orm_metadata_rows(
|
||||
table_rows: list[dict[str, Any]],
|
||||
column_rows: list[dict[str, Any]],
|
||||
*,
|
||||
allowed_tables: set[str],
|
||||
missing_table_objects: Iterable[Any] = (),
|
||||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
"""以 model/ ORM 全量为名单主体合并元数据行。
|
||||
|
||||
- 库里真实存在的 ORM 表:沿用 information_schema 行(列注释更全);
|
||||
- 库里没有的 ORM 表(missing_table_objects):用 ORM 定义合成行,
|
||||
并在表说明追加 MISSING_TABLE_MARKER 标注"尚未建表";
|
||||
- 名单之外(DB-only)的行在此丢弃,交由 sync_metadata 从 Milvus 清理。
|
||||
"""
|
||||
allowed = {str(name).strip() for name in allowed_tables if str(name).strip()}
|
||||
merged_tables = [row for row in table_rows if _row_table_name(row) in allowed]
|
||||
merged_columns = [row for row in column_rows if _row_table_name(row) in allowed]
|
||||
|
||||
missing_table_rows, missing_column_rows = build_orm_metadata_rows(
|
||||
missing_table_objects
|
||||
)
|
||||
for row in missing_table_rows:
|
||||
comment = str(row.get("TABLE_COMMENT") or "").strip()
|
||||
row["TABLE_COMMENT"] = f"{comment}{MISSING_TABLE_MARKER}".strip()
|
||||
merged_tables.extend(missing_table_rows)
|
||||
merged_columns.extend(missing_column_rows)
|
||||
return merged_tables, merged_columns
|
||||
|
||||
|
||||
async def prepare_metadata_rows(
|
||||
chunks: list[dict[str, Any]],
|
||||
*,
|
||||
@@ -39,7 +148,18 @@ async def sync_metadata(
|
||||
*,
|
||||
embedder: Callable[[list[str]], Awaitable[list[list[float]]]] = embed_texts,
|
||||
timestamp: int | None = None,
|
||||
allowed_tables: set[str] | None = None,
|
||||
) -> int:
|
||||
"""构造并 upsert 元数据 chunk。
|
||||
|
||||
allowed_tables 提供时:仅同步名单内的表,并删除集合中名单之外的
|
||||
表元数据( Milvus 里只保留主数据源认可的表,如 model/ ORM 覆盖的表)。
|
||||
"""
|
||||
if allowed_tables is not None:
|
||||
allowed = {str(name).strip() for name in allowed_tables if str(name).strip()}
|
||||
table_rows = _filter_rows_by_tables(table_rows, allowed)
|
||||
column_rows = _filter_rows_by_tables(column_rows, allowed)
|
||||
|
||||
chunks = build_metadata_chunks(table_rows, column_rows)
|
||||
chunks_by_table: dict[str, list[dict[str, Any]]] = {}
|
||||
for chunk in chunks:
|
||||
@@ -80,4 +200,7 @@ async def sync_metadata(
|
||||
if rows:
|
||||
await milvus_client.upsert(collection_name=NL2SQL_COLLECTION, data=rows)
|
||||
updated_count += len(rows)
|
||||
|
||||
if allowed_tables is not None:
|
||||
await _delete_tables_outside_allowlist(milvus_client, allowed)
|
||||
return updated_count
|
||||
|
||||
+4
-1
@@ -52,7 +52,10 @@ INTENT_SYSTEM_PROMPT = (
|
||||
"- knowledge_qa: 用户询问基金相关的知识性问题,如净值、费率、风险、申赎规则等\n"
|
||||
"- company_info: 用户询问华夏科技公司本身的信息,如公司全称、成立时间、牌照、总部地址、"
|
||||
"客服电话、服务时间、官网、投诉渠道等\n"
|
||||
"- nl2sql_request: 用户要求查询具体数据或账户信息\n"
|
||||
"- nl2sql_request: 用户要求查询具体数据或账户信息,"
|
||||
"如“我的持仓有哪些”“我买了多少XX基金”“我的交易记录”"
|
||||
"“最近一周XX基金的净值数据”“XX基金最新的规模/费率数据”等"
|
||||
"要求数据本身而非知识解释的问题\n"
|
||||
"- chitchat: 普通寒暄,如问候、致谢、告别、询问你是谁/你能做什么、在吗等一两句话的闲聊\n"
|
||||
"- off_topic: 用户要求你实质性地处理与金融、基金、公司业务无关的事情,"
|
||||
"如写代码、讲笑话、写作文、问天气、聊政治、情感咨询、做数学题等\n"
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
"""按 fin_product 为指定客户生成 fin_holdings 持仓数据(本地开发/演示用)。
|
||||
|
||||
用法:
|
||||
python scripts/seed_holdings_for_customer.py --customer-id 18
|
||||
python scripts/seed_holdings_for_customer.py --customer-id 18 --count 5 --force
|
||||
|
||||
规则:
|
||||
- 只挑选 fin_product 中 status=在售 的产品,按 id 升序取前 count 只;
|
||||
- 每笔持仓的买入净值 = 当前净值 × (1 - 浮动),浮动由固定 seed 生成,保证可复现;
|
||||
- shares = cost_amount / 买入净值(4 位小数),current_value = shares × 当前净值;
|
||||
- profit_loss / profit_ratio 由上述字段推导;
|
||||
- 客户已有持仓时默认拒绝,需 --force 才会追加。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import random
|
||||
import sys
|
||||
from decimal import Decimal, ROUND_HALF_UP
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from config.database.mysql import get_session_factory
|
||||
from model.fin_holdings import FinHoldings
|
||||
from model.fin_product import FinProduct
|
||||
from model.sys_user import SysUser
|
||||
|
||||
_MONEY = Decimal("0.01")
|
||||
_SHARES = Decimal("0.0001")
|
||||
_RATIO = Decimal("0.0001")
|
||||
|
||||
|
||||
def _money(value: Decimal) -> Decimal:
|
||||
return value.quantize(_MONEY, rounding=ROUND_HALF_UP)
|
||||
|
||||
|
||||
async def seed(customer_id: int, *, count: int, force: bool) -> list[dict]:
|
||||
session_factory = get_session_factory()
|
||||
async with session_factory() as db:
|
||||
user = await db.get(SysUser, customer_id)
|
||||
if user is None or user.user_type != "CUSTOMER":
|
||||
raise SystemExit(f"customer_id={customer_id} 不是客户用户或不存在")
|
||||
|
||||
existing = (
|
||||
await db.execute(
|
||||
select(FinHoldings).where(FinHoldings.customer_id == customer_id)
|
||||
)
|
||||
).scalars().all()
|
||||
if existing and not force:
|
||||
raise SystemExit(
|
||||
f"客户 {customer_id} 已有 {len(existing)} 条持仓,如需追加请加 --force"
|
||||
)
|
||||
|
||||
products = (
|
||||
await db.execute(
|
||||
select(FinProduct)
|
||||
.where(FinProduct.status == "在售")
|
||||
.order_by(FinProduct.id)
|
||||
.limit(count)
|
||||
)
|
||||
).scalars().all()
|
||||
if not products:
|
||||
raise SystemExit("fin_product 中没有在售产品")
|
||||
|
||||
rng = random.Random(f"holdings-{customer_id}")
|
||||
created: list[dict] = []
|
||||
for product in products:
|
||||
nav = product.nav or Decimal("1.000000")
|
||||
cost_amount = _money(Decimal(rng.randint(5_000, 50_000)))
|
||||
# 买入净值在当前净值的 88%~97% 之间浮动,形成有涨有跌的持仓
|
||||
buy_nav = _money(nav * Decimal(str(round(rng.uniform(0.88, 0.97), 6))))
|
||||
if buy_nav <= 0:
|
||||
buy_nav = Decimal("1.0000")
|
||||
shares = (cost_amount / buy_nav).quantize(_SHARES, rounding=ROUND_HALF_UP)
|
||||
current_value = _money(shares * nav)
|
||||
profit_loss = _money(current_value - cost_amount)
|
||||
profit_ratio = (profit_loss / cost_amount).quantize(_RATIO, rounding=ROUND_HALF_UP)
|
||||
|
||||
db.add(
|
||||
FinHoldings(
|
||||
customer_id=customer_id,
|
||||
product_id=product.id,
|
||||
shares=shares,
|
||||
cost_amount=cost_amount,
|
||||
current_value=current_value,
|
||||
profit_loss=profit_loss,
|
||||
profit_ratio=profit_ratio,
|
||||
status="持有中",
|
||||
)
|
||||
)
|
||||
created.append(
|
||||
{
|
||||
"product_id": product.id,
|
||||
"product_name": product.product_name,
|
||||
"buy_nav": str(buy_nav),
|
||||
"nav": str(nav),
|
||||
"shares": str(shares),
|
||||
"cost_amount": str(cost_amount),
|
||||
"current_value": str(current_value),
|
||||
"profit_loss": str(profit_loss),
|
||||
"profit_ratio": str(profit_ratio),
|
||||
}
|
||||
)
|
||||
|
||||
await db.commit()
|
||||
return created
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="按 fin_product 生成客户持仓数据")
|
||||
parser.add_argument("--customer-id", type=int, required=True)
|
||||
parser.add_argument("--count", type=int, default=8, help="生成持仓条数(默认 8)")
|
||||
parser.add_argument(
|
||||
"--force", action="store_true", help="客户已有持仓时仍允许追加"
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
created = asyncio.run(seed(args.customer_id, count=args.count, force=args.force))
|
||||
total_cost = sum(Decimal(item["cost_amount"]) for item in created)
|
||||
total_value = sum(Decimal(item["current_value"]) for item in created)
|
||||
print(f"客户 {args.customer_id} 新增持仓 {len(created)} 条:")
|
||||
for item in created:
|
||||
print(
|
||||
f" [{item['product_id']}] {item['product_name'][:24]}… "
|
||||
f"买入净值={item['buy_nav']} 当前净值={item['nav']} "
|
||||
f"份额={item['shares']} 成本={item['cost_amount']} "
|
||||
f"市值={item['current_value']} 盈亏={item['profit_loss']} "
|
||||
f"({item['profit_ratio']})"
|
||||
)
|
||||
print(f"合计:成本 {total_cost} 元,市值 {total_value} 元,"
|
||||
f"盈亏 {total_value - total_cost} 元")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,7 +1,18 @@
|
||||
"""读取当前 MySQL 元数据并同步到 Milvus。"""
|
||||
"""同步 NL2SQL 元数据到 Milvus(表名单以 model/ 目录的 ORM 为主)。
|
||||
|
||||
主数据源策略:
|
||||
- 表名单 = model/ 下 SQLAlchemy ORM 定义的全部表(Base.metadata),
|
||||
不再与数据库求交集——数据库尚未建表的 ORM 表也纳入元数据,
|
||||
其表/字段信息从 ORM 定义合成,并在表说明标注"尚未建表";
|
||||
- 数据库里真实存在的 ORM 表,表/列中文注释仍取自 information_schema
|
||||
(列注释只存在于库中,ORM 代码里没有字段级注释);
|
||||
- 名单之外(库里有、model/ 没定义)的表元数据 chunk 会被从 Milvus 删除。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import importlib
|
||||
import pkgutil
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
@@ -14,7 +25,8 @@ from config.database import mysql
|
||||
from config.database.milvus import client as milvus_client
|
||||
from config.database.mysql import get_session_factory
|
||||
from config.settings import settings
|
||||
from nl2sql.metadata_sync import sync_metadata
|
||||
from model.base import Base
|
||||
from nl2sql.metadata_sync import merge_orm_metadata_rows, sync_metadata
|
||||
|
||||
|
||||
TABLES_SQL = text(
|
||||
@@ -36,6 +48,15 @@ COLUMNS_SQL = text(
|
||||
)
|
||||
|
||||
|
||||
def collect_orm_tables() -> set[str]:
|
||||
"""导入 model/ 全部模块,从 Base.metadata 收集 ORM 表名。"""
|
||||
import model
|
||||
|
||||
for module_info in pkgutil.iter_modules(model.__path__):
|
||||
importlib.import_module(f"model.{module_info.name}")
|
||||
return set(Base.metadata.tables.keys())
|
||||
|
||||
|
||||
async def load_information_schema() -> tuple[list[dict], list[dict]]:
|
||||
async with get_session_factory()() as session:
|
||||
tables = [
|
||||
@@ -49,14 +70,43 @@ async def load_information_schema() -> tuple[list[dict], list[dict]]:
|
||||
return tables, columns
|
||||
|
||||
|
||||
async def synchronize() -> int:
|
||||
def _table_name(row: dict) -> str:
|
||||
return str(row.get("TABLE_NAME") or row.get("table_name") or "").strip()
|
||||
|
||||
|
||||
async def synchronize() -> tuple[int, set[str], set[str], set[str]]:
|
||||
"""返回 (upserted, allowed_tables, 名单外表, 未建表的 ORM 表)。"""
|
||||
orm_tables = collect_orm_tables()
|
||||
try:
|
||||
tables, columns = await load_information_schema()
|
||||
return await sync_metadata(milvus_client(), tables, columns)
|
||||
db_tables = {_table_name(row) for row in tables}
|
||||
missing_tables = sorted(orm_tables - db_tables)
|
||||
table_rows, column_rows = merge_orm_metadata_rows(
|
||||
tables,
|
||||
columns,
|
||||
allowed_tables=orm_tables,
|
||||
missing_table_objects=(
|
||||
Base.metadata.tables[name] for name in missing_tables
|
||||
),
|
||||
)
|
||||
upserted = await sync_metadata(
|
||||
milvus_client(),
|
||||
table_rows,
|
||||
column_rows,
|
||||
allowed_tables=orm_tables,
|
||||
)
|
||||
dropped = db_tables - orm_tables
|
||||
return upserted, orm_tables, dropped, set(missing_tables)
|
||||
finally:
|
||||
await mysql.dispose()
|
||||
await milvus_db.dispose()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print(f"upserted {asyncio.run(synchronize())} NL2SQL metadata chunks")
|
||||
upserted, allowed, dropped, missing = asyncio.run(synchronize())
|
||||
print(f"ORM 表名单(model/ 全量): {len(allowed)} 张")
|
||||
print(f"其中数据库尚未建表(用 ORM 定义合成元数据): {len(missing)} 张")
|
||||
if missing:
|
||||
print(f" {sorted(missing)}")
|
||||
print(f"已排除的非 ORM 表: {sorted(dropped) if dropped else '无'}")
|
||||
print(f"upserted {upserted} NL2SQL metadata chunks")
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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 做受限闲聊,失败时退回固定话术。"""
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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()),
|
||||
|
||||
Reference in New Issue
Block a user