diff --git a/.gitignore b/.gitignore index 52d0c43..20a29a3 100644 --- a/.gitignore +++ b/.gitignore @@ -72,6 +72,7 @@ data/output/ *.xlsx *.parquet docs/ +.workbuddy/ # 本地临时需求文档 DK2_客服Agent模块完整开发计划(v1.1).md diff --git a/nl2sql/metadata_sync.py b/nl2sql/metadata_sync.py index 0c3ed9c..34fb20a 100644 --- a/nl2sql/metadata_sync.py +++ b/nl2sql/metadata_sync.py @@ -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 diff --git a/rag/intent.py b/rag/intent.py index acd8786..788077b 100644 --- a/rag/intent.py +++ b/rag/intent.py @@ -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" diff --git a/scripts/seed_holdings_for_customer.py b/scripts/seed_holdings_for_customer.py new file mode 100644 index 0000000..42c36bd --- /dev/null +++ b/scripts/seed_holdings_for_customer.py @@ -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() diff --git a/scripts/sync_nl2sql_metadata.py b/scripts/sync_nl2sql_metadata.py index 7c3ba9a..56c3596 100644 --- a/scripts/sync_nl2sql_metadata.py +++ b/scripts/sync_nl2sql_metadata.py @@ -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") diff --git a/service/client_agent/bootstrap.py b/service/client_agent/bootstrap.py index eec3b6a..5d0a2ac 100644 --- a/service/client_agent/bootstrap.py +++ b/service/client_agent/bootstrap.py @@ -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, ) diff --git a/service/client_agent/runtime.py b/service/client_agent/runtime.py index 04b2cc0..3449bf3 100644 --- a/service/client_agent/runtime.py +++ b/service/client_agent/runtime.py @@ -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( diff --git a/service/customer_agent/chat.py b/service/customer_agent/chat.py index 393776d..d99d7b6 100644 --- a/service/customer_agent/chat.py +++ b/service/customer_agent/chat.py @@ -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 做受限闲聊,失败时退回固定话术。""" diff --git a/service/nl2sql/answer_render.py b/service/nl2sql/answer_render.py new file mode 100644 index 0000000..ed982de --- /dev/null +++ b/service/nl2sql/answer_render.py @@ -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 diff --git a/service/nl2sql/customer_permission.py b/service/nl2sql/customer_permission.py new file mode 100644 index 0000000..165b43e --- /dev/null +++ b/service/nl2sql/customer_permission.py @@ -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, + } diff --git a/service/nl2sql/query_service.py b/service/nl2sql/query_service.py index 73897ef..6ee2d8a 100644 --- a/service/nl2sql/query_service.py +++ b/service/nl2sql/query_service.py @@ -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()),