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/api/advisor/audit.py b/api/advisor/audit.py index 611378d..cea0bd0 100644 --- a/api/advisor/audit.py +++ b/api/advisor/audit.py @@ -22,7 +22,7 @@ async def ledger( start: datetime | None = Query(None), end: datetime | None = Query(None), page: int = Query(1, ge=1), - page_size: int = Query(20, ge=1, le=100), + page_size: int = Query(10, ge=1, le=100), user: SysUser = Depends(require_advisor), db: AsyncSession = Depends(get_db), ): diff --git a/api/advisor/customers.py b/api/advisor/customers.py index d18e9e2..1fbfcdb 100644 --- a/api/advisor/customers.py +++ b/api/advisor/customers.py @@ -17,7 +17,7 @@ async def list_customers( status: str | None = Query(None, max_length=16), keyword: str | None = Query(None, max_length=64), page: int = Query(1, ge=1), - page_size: int = Query(20, ge=1, le=100), + page_size: int = Query(10, ge=1, le=100), user: SysUser = Depends(require_advisor), db: AsyncSession = Depends(get_db), ): @@ -51,7 +51,7 @@ async def get_holdings( async def get_reports( customer_id: int, page: int = Query(1, ge=1), - page_size: int = Query(20, ge=1, le=100), + page_size: int = Query(10, ge=1, le=100), user: SysUser = Depends(require_advisor), db: AsyncSession = Depends(get_db), ): diff --git a/api/advisor/drafts.py b/api/advisor/drafts.py index ff9791a..bb97af8 100644 --- a/api/advisor/drafts.py +++ b/api/advisor/drafts.py @@ -19,7 +19,7 @@ async def list_drafts( customer_id: int | None = Query(None), status: str | None = Query(None, max_length=16), page: int = Query(1, ge=1), - page_size: int = Query(20, ge=1, le=100), + page_size: int = Query(10, ge=1, le=100), user: SysUser = Depends(require_advisor), db: AsyncSession = Depends(get_db), ): diff --git a/api/advisor/todos.py b/api/advisor/todos.py index 1fba54d..23db5d8 100644 --- a/api/advisor/todos.py +++ b/api/advisor/todos.py @@ -17,7 +17,7 @@ async def list_todos( status: str | None = Query(None, max_length=16), todo_type: str | None = Query(None, max_length=32), page: int = Query(1, ge=1), - page_size: int = Query(20, ge=1, le=100), + page_size: int = Query(10, ge=1, le=100), user: SysUser = Depends(require_advisor), db: AsyncSession = Depends(get_db), ): diff --git a/api/advisor/visits.py b/api/advisor/visits.py index a92e2e4..e1a69c3 100644 --- a/api/advisor/visits.py +++ b/api/advisor/visits.py @@ -16,7 +16,7 @@ router = APIRouter() async def list_visits( customer_id: int | None = Query(None), page: int = Query(1, ge=1), - page_size: int = Query(20, ge=1, le=100), + page_size: int = Query(10, ge=1, le=100), user: SysUser = Depends(require_advisor), db: AsyncSession = Depends(get_db), ): @@ -64,7 +64,7 @@ async def talk_templates(user: SysUser = Depends(require_advisor)): async def touch_logs( customer_id: int = Query(...), page: int = Query(1, ge=1), - page_size: int = Query(20, ge=1, le=100), + page_size: int = Query(10, ge=1, le=100), user: SysUser = Depends(require_advisor), db: AsyncSession = Depends(get_db), ): diff --git a/api/routers/advisor_agent.py b/api/routers/advisor_agent.py index 0542a78..cbb1a12 100644 --- a/api/routers/advisor_agent.py +++ b/api/routers/advisor_agent.py @@ -2,11 +2,11 @@ from __future__ import annotations import json +import re from typing import Literal from fastapi import APIRouter, BackgroundTasks, Depends, Query, Request from fastapi.responses import StreamingResponse -from pydantic import ValidationError from sqlalchemy.ext.asyncio import AsyncSession from agent.advisor_agent.auth import ensure_customer_access @@ -40,6 +40,7 @@ from config.deps import get_db from config.database import mysql, redis as redis_db from model.sys_user import SysUser from repositories.advisor_draft import AdvisorDraftRepo +from repositories.customer_relation import CustomerRelationRepo from service.advisor_agent.context import ( load_fund_analysis_context, load_customer_risk, @@ -92,6 +93,48 @@ def _advisor_runtime(request: Request): return getattr(getattr(app, "state", None), "advisor_agent_runtime", None) +def _infer_chat_intent(query: str) -> str | None: + """从自然语言问题推断投顾意图;无法确定时保留通用问答。""" + if any(word in query for word in ("调仓", "再平衡", "组合偏离")): + return "rebalance" + if any(word in query for word in ("沟通话术", "怎么和客户说", "解释给客户")): + return "dialogue-script" + if any(word in query for word in ("基金分析", "分析这只基金", "分析产品")): + return "fund_analysis" + if any(word in query for word in ("推荐", "产品建议", "买什么基金", "适合的基金")): + return AGENT_INTENT_RECOMMEND + return None + + +async def _resolve_customer_from_query(db, *, advisor_id: int, query: str) -> tuple[int | None, str | None]: + """解析问题中的客户编号或姓名,并限制在当前投顾客户范围内。""" + number_match = re.search(r"(?:客户|用户)\s*[#编号号:]?\s*(\d+)", query) + relation_repo = CustomerRelationRepo(db) + if number_match: + customer_id = int(number_match.group(1)) + relation = await relation_repo.get_active_relation( + customer_id=customer_id, + advisor_id=advisor_id, + ) + if relation is None: + return None, "问题中的客户不在当前投顾的授权范围内" + return customer_id, None + + rows = await relation_repo.list_customer_rows(advisor_id=advisor_id, limit=100) + matched = { + int(account.id) + for _relation, account, _profile in rows + if account.real_name and account.real_name in query + } + if len(matched) == 1: + return next(iter(matched)), None + if len(matched) > 1: + return None, "问题中的客户姓名无法唯一确定,请补充客户编号" + if "客户" in query or "用户" in query: + return None, "请在问题中补充客户编号或客户姓名" + return None, None + + async def _recall_advisor_memories( request: Request, *, customer_id: int, query: str ) -> list[dict]: @@ -148,19 +191,80 @@ async def _run_rebalance_background( @router.post("/chat/stream") async def chat_stream( request: Request, - body: dict, + body: AdvisorChatReq, user: SysUser = Depends(audited_advisor), db: AsyncSession = Depends(get_db), ): trace_id = _trace_id(request) - try: - chat_request = AdvisorChatReq.model_validate(body) - except ValidationError: - payload = agent_failure( - ERR_CODE_FORBIDDEN_CUSTOMER, - "对话请求缺少有效客户范围或参数", - trace_id=trace_id, + chat_request = body + customer_id = chat_request.customer_id + inferred_intent = chat_request.intent or _infer_chat_intent(chat_request.query) + + # 请求体只传问题时,从问题中解析客户;解析结果仍必须经过投顾关系授权校验。 + if customer_id is None: + customer_id, resolve_error = await _resolve_customer_from_query( + db, + advisor_id=user.id, + query=chat_request.query, ) + if resolve_error: + payload = agent_failure( + ERR_CODE_FORBIDDEN_CUSTOMER, + resolve_error, + trace_id=trace_id, + ) + async def resolve_error_events(): + yield f"data: {json.dumps({'type': SSE_EVENT_TYPE_ERROR, **payload}, ensure_ascii=False)}\n\n" + + return StreamingResponse( + resolve_error_events(), + media_type="text/event-stream", + headers={"X-Trace-Id": trace_id}, + ) + else: + payload = None + + # 不带客户编号时只提供通用基金问答,不读取客户画像,也不生成个性化草稿。 + if customer_id is None: + if inferred_intent in { + AGENT_INTENT_RECOMMEND, + "rebalance", + "fund_analysis", + "dialogue-script", + }: + if payload is None: + payload = agent_failure( + ERR_CODE_FORBIDDEN_CUSTOMER, + "个性化投顾分析需要在问题中明确客户编号或姓名", + trace_id=trace_id, + ) + else: + runtime = _advisor_runtime(request) + llm_client = getattr(runtime, "llm_client", None) + if llm_client is None: + answer = "已收到问题。当前未配置通用投顾模型,请选择客户后使用个性化分析,或联系管理员配置 Agent 服务。" + else: + answer = await generate_text( + llm_client, + system_prompt="你是基金投顾助手,只回答通用基金知识和产品分析问题,不读取或推断任何客户信息,不承诺收益,不代客交易。", + user_prompt=chat_request.query, + fallback=lambda: "当前模型暂时不可用,请稍后重试。", + timeout=5.0, + ) + + async def events(): + for event in ( + {"type": SSE_EVENT_TYPE_META, "intent": "general_question"}, + {"type": SSE_EVENT_TYPE_TEXT, "content": answer}, + {"type": SSE_EVENT_TYPE_DONE}, + ): + yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n" + + return StreamingResponse( + events(), + media_type="text/event-stream", + headers={"X-Trace-Id": trace_id}, + ) else: customer_id = chat_request.customer_id resolved_intent = recognize_advisor_intent( @@ -170,11 +274,12 @@ async def chat_stream( relation = await ensure_customer_access( db, advisor_id=user.id, customer_id=int(customer_id) ) + if inferred_intent == AGENT_INTENT_RECOMMEND: if resolved_intent == AGENT_INTENT_RECOMMEND: memories = await _recall_advisor_memories( request, customer_id=int(customer_id), - query=chat_request.query or "", + query=chat_request.query, ) runtime = _advisor_runtime(request) context = await load_recommendation_context( @@ -296,7 +401,7 @@ async def list_drafts( customer_id: int | None = Query(default=None), status: Literal[DRAFT_STATUS_DRAFT, DRAFT_STATUS_DISCARDED] | None = Query(default=None), page: int = Query(default=1, ge=1), - page_size: int = Query(default=20, ge=1, le=100), + page_size: int = Query(default=10, ge=1, le=100), user: SysUser = Depends(audited_advisor), db: AsyncSession = Depends(get_db), ): diff --git a/api/routers/nl2sql.py b/api/routers/nl2sql.py index 889f780..000ad61 100644 --- a/api/routers/nl2sql.py +++ b/api/routers/nl2sql.py @@ -57,6 +57,7 @@ from tool.llm import llm from utils.exceptions import ForbiddenError, NotFoundError, ParamError from utils.request_id import get_request_id, new_request_id from utils.response import success +from utils.pagination import normalize_pagination router = APIRouter() @@ -474,10 +475,11 @@ async def list_query_history( ): """分页读取当前员工自己的查询历史。""" ensure_query_employee(user) + page, page_size, offset = normalize_pagination(page, page_size) rows = await Nl2SqlPermissionRepo(db).list_query_history( user.id, limit=page_size, - offset=(page - 1) * page_size, + offset=offset, status=status, start_time=start_time, end_time=end_time, diff --git a/api/routers/nl2sql_admin.py b/api/routers/nl2sql_admin.py index 3c62dea..963ea75 100644 --- a/api/routers/nl2sql_admin.py +++ b/api/routers/nl2sql_admin.py @@ -5,7 +5,7 @@ import json import time from datetime import datetime, timedelta, timezone -from fastapi import APIRouter, Depends +from fastapi import APIRouter, Depends, Query from fastapi.responses import PlainTextResponse from sqlalchemy.ext.asyncio import AsyncSession @@ -275,7 +275,7 @@ async def run_admin_job( @router.get("/nl2sql/admin/jobs/history") async def admin_job_history( page: int = 1, - page_size: int = 20, + page_size: int = Query(10, ge=1, le=10), status: str | None = None, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db), diff --git a/api/routers/risk.py b/api/routers/risk.py index 20ac821..0d5748c 100644 --- a/api/routers/risk.py +++ b/api/routers/risk.py @@ -1,5 +1,5 @@ """风控处置路由:预警列表 + 放行/拦截/冻结(仅风控专员)。""" -from fastapi import APIRouter, Depends +from fastapi import APIRouter, Depends, Query from sqlalchemy.ext.asyncio import AsyncSession from api.deps import require_risk_officer @@ -14,11 +14,12 @@ router = APIRouter(prefix="/risk", tags=["风控"]) @router.get("/alert/list", summary="预警列表") async def list_alerts( status: str | None = None, + page: int = Query(1, ge=1), + page_size: int = Query(10, ge=1, le=10), user: SysUser = Depends(require_risk_officer), db: AsyncSession = Depends(get_db), ): - alerts = await risk_handle.list_alerts(db, status) - return success([a.model_dump(mode="json") for a in alerts]) + return success(await risk_handle.list_alerts(db, status, page=page, page_size=page_size)) @router.post("/alert/{alert_id}/release", summary="放行") diff --git a/api/routers/work_order.py b/api/routers/work_order.py index 4ad730d..cec183f 100644 --- a/api/routers/work_order.py +++ b/api/routers/work_order.py @@ -1,5 +1,5 @@ """业务工单路由:列表 / 详情 / 认领 / 提交审核 / 复核(仅风控专员)。""" -from fastapi import APIRouter, Depends +from fastapi import APIRouter, Depends, Query from sqlalchemy.ext.asyncio import AsyncSession from api.deps import require_risk_officer @@ -15,11 +15,16 @@ router = APIRouter(prefix="/work-order", tags=["工单"]) @router.get("/list", summary="工单列表") async def list_work_orders( status: str | None = None, + page: int = Query(1, ge=1), + page_size: int = Query(10, ge=1, le=10), user: SysUser = Depends(require_risk_officer), db: AsyncSession = Depends(get_db), ): - orders = await work_order_service.list_work_orders(db, status) - return success([o.model_dump(mode="json") for o in orders]) + return success( + await work_order_service.list_work_orders( + db, status, page=page, page_size=page_size + ) + ) @router.get("/{work_order_id}", summary="工单详情") diff --git a/nl2sql/job_history.py b/nl2sql/job_history.py index 750d696..7763c4b 100644 --- a/nl2sql/job_history.py +++ b/nl2sql/job_history.py @@ -7,6 +7,7 @@ from datetime import datetime from sqlalchemy import text from nl2sql.audit import _sanitize +from utils.pagination import normalize_pagination, pagination_result async def record_job_history(db, result, *, elapsed_ms: float, parameter_summary: dict | None = None) -> None: @@ -49,9 +50,15 @@ async def record_job_history_safely(db, result, *, elapsed_ms: float, parameter_ return True -async def list_job_history(db, *, page: int = 1, page_size: int = 20, status: str | None = None) -> list[dict]: +async def list_job_history(db, *, page: int = 1, page_size: int = 10, status: str | None = None) -> dict: """分页查询任务历史,只返回执行摘要,不返回 detail 明细。""" + page, page_size, offset = normalize_pagination(page, page_size) conditions = "WHERE (:status IS NULL OR status = :status)" + count_result = await db.execute( + text(f"SELECT COUNT(*) AS total FROM nl2sql_job_history {conditions}"), + {"status": status}, + ) + total = int(count_result.scalar() or 0) statement = text( f""" SELECT id, job_name, status, attempts, error_type, elapsed_ms, create_time @@ -65,8 +72,13 @@ async def list_job_history(db, *, page: int = 1, page_size: int = 20, status: st statement, { "status": status, - "limit": max(1, min(page_size, 100)), - "offset": max(0, (page - 1) * page_size), + "limit": page_size, + "offset": offset, }, ) - return [dict(row) for row in result.mappings().all()] + return pagination_result( + [dict(row) for row in result.mappings().all()], + total, + page=page, + page_size=page_size, + ) 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 2637a49..788077b 100644 --- a/rag/intent.py +++ b/rag/intent.py @@ -1,9 +1,12 @@ """Customer-service intent recognition contract.""" from __future__ import annotations +import json import logging import re +from dataclasses import dataclass from enum import StrEnum +from inspect import isawaitable logger = logging.getLogger("rag.intent") @@ -26,25 +29,58 @@ INTENT_VALUES = frozenset(item.value for item in Intent) _INTENT_PATTERN = re.compile( "|".join(re.escape(value) for value in sorted(INTENT_VALUES, key=len, reverse=True)) ) +_JSON_PATTERN = re.compile(r"\{.*\}", re.S) + +# 改写句允许的最大长度:相对原句的倍数与绝对下限取大者,防止模型把历史大段塞进改写句 +_REWRITE_MAX_RATIO = 4 +_REWRITE_MIN_LIMIT = 200 INTENT_SYSTEM_PROMPT = ( "你是华夏科技(一家基金代销金融机构)智能客服的意图分类器。" - "请根据用户输入,从以下选项中选择最匹配的意图,并仅输出对应的英文标签(不输出任何其他内容):\n" + "你会看到最近几轮对话(可能为空)和用户的【当前输入】。\n" + "任务:\n" + "1. 只对【当前输入】判定意图;【最近对话】仅用于理解当前输入中省略的主语或指代" + "(如“它”“那个”“呢”“还有吗”)。\n" + "2. 若当前输入依赖上文才能理解,请把它改写为一句不依赖上文的完整问题," + "只能补全上文已明确出现的对象,不得添加新信息、不得改变原意;" + "若当前输入本身已完整,query 原样返回。\n" + "3. 当前输入若明显切换到新话题,以当前输入为准,不要延续上一轮的意图。\n" + "4. 致谢、告别、寒暄即使出现在知识问答之后也归为 chitchat。\n" + "意图选项:\n" "- guide_purchase: 用户询问如何购买基金、开户、注册等引导类问题\n" "- want_advisor: 用户希望获得个性化基金推荐或投资顾问服务\n" "- knowledge_qa: 用户询问基金相关的知识性问题,如净值、费率、风险、申赎规则等\n" "- company_info: 用户询问华夏科技公司本身的信息,如公司全称、成立时间、牌照、总部地址、" "客服电话、服务时间、官网、投诉渠道等\n" - "- nl2sql_request: 用户要求查询具体数据或账户信息\n" + "- nl2sql_request: 用户要求查询具体数据或账户信息," + "如“我的持仓有哪些”“我买了多少XX基金”“我的交易记录”" + "“最近一周XX基金的净值数据”“XX基金最新的规模/费率数据”等" + "要求数据本身而非知识解释的问题\n" "- chitchat: 普通寒暄,如问候、致谢、告别、询问你是谁/你能做什么、在吗等一两句话的闲聊\n" "- off_topic: 用户要求你实质性地处理与金融、基金、公司业务无关的事情," "如写代码、讲笑话、写作文、问天气、聊政治、情感咨询、做数学题等\n" "- no_match: 无法归入以上任何一类\n" "注意:寒暄性质的一两句话归为 chitchat;一旦用户提出金融之外的实质性请求,归为 off_topic。\n" - "仅输出一个小写英文标签,不要输出解释、标点或换行。" + "只输出一行 JSON,格式:{\"intent\": \"<小写英文标签>\", \"query\": \"<完整问题>\"}," + "不要输出解释、Markdown 或其他内容。" ) +@dataclass(frozen=True) +class IntentResult: + intent: Intent + # 结合历史补全后的独立问题;无需补全或解析失败时等于原句 + query: str + used_history: bool = False + + +async def _config_value(config_getter, key: str, default): + value = config_getter(key, str(default)) + if isawaitable(value): + value = await value + return type(default)(value) + + def parse_intent(raw: str | None) -> Intent: """从模型原始输出中提取第一个合法标签,提取不到则返回 NO_MATCH。""" if not raw: @@ -53,16 +89,77 @@ def parse_intent(raw: str | None) -> Intent: return Intent(match.group(0)) if match else Intent.NO_MATCH -async def intent_recognize(query: str, *, llm_client) -> Intent: - if not query or not query.strip(): - return Intent.NO_MATCH - messages = [ +def parse_intent_result(raw: str | None, *, fallback_query: str, used_history: bool = False) -> IntentResult: + """优先按 JSON 解析意图与改写句;JSON 不可用时退回正则抽标签、原句作为 query。""" + match = _JSON_PATTERN.search(raw) if raw else None + data = None + if match: + try: + data = json.loads(match.group(0)) + except ValueError: + data = None + if not isinstance(data, dict): + return IntentResult(parse_intent(raw), fallback_query, used_history) + intent = parse_intent(str(data.get("intent") or "")) + rewritten = str(data.get("query") or "").strip() + limit = max(len(fallback_query) * _REWRITE_MAX_RATIO, _REWRITE_MIN_LIMIT) + query = rewritten if rewritten and len(rewritten) <= limit else fallback_query + return IntentResult(intent, query, used_history) + + +def render_history(history, *, max_turns: int, max_chars: int) -> str: + """把最近 max_turns 轮 user/assistant 消息折叠成一段文本,每条截断到 max_chars。""" + if max_turns <= 0 or not history: + return "" + recent = [ + message for message in history + if isinstance(message, dict) + and message.get("role") in ("user", "assistant") + and message.get("content") + ][-max_turns * 2:] + lines = [] + for message in recent: + role = "用户" if message["role"] == "user" else "客服" + content = str(message["content"]).replace("\n", " ").strip() + if len(content) > max_chars: + content = content[:max_chars] + "…" + lines.append(f"{role}:{content}") + return "\n".join(lines) + + +def build_intent_messages(query: str, rendered_history: str) -> list[dict]: + if rendered_history: + user_content = f"【最近对话】\n{rendered_history}\n\n【当前输入】\n{query}" + else: + user_content = f"【当前输入】\n{query}" + return [ {"role": "system", "content": INTENT_SYSTEM_PROMPT}, - {"role": "user", "content": query}, + {"role": "user", "content": user_content}, ] + + +async def intent_recognize( + query: str, + *, + llm_client, + history: list[dict] | None = None, + config_getter=None, +) -> IntentResult: + """结合最近几轮对话识别当前输入的意图,并给出补全指代后的独立问题。 + + history 为空或 agent.customer.intent.history_turns 配置为 0 时退化为仅看当前句。 + """ + if not query or not query.strip(): + return IntentResult(Intent.NO_MATCH, query or "") + max_turns, max_chars = 3, 200 + if config_getter is not None: + max_turns = await _config_value(config_getter, "agent.customer.intent.history_turns", max_turns) + max_chars = await _config_value(config_getter, "agent.customer.intent.history_max_chars", max_chars) + rendered = render_history(history, max_turns=max_turns, max_chars=max_chars) + messages = build_intent_messages(query, rendered) try: raw = await llm_client.chat(messages) except Exception: logger.exception("intent recognition failed") - return Intent.NO_MATCH - return parse_intent(raw) + return IntentResult(Intent.NO_MATCH, query, bool(rendered)) + return parse_intent_result(raw, fallback_query=query, used_history=bool(rendered)) diff --git a/repositories/biz_work_order.py b/repositories/biz_work_order.py index aa645e6..be0acad 100644 --- a/repositories/biz_work_order.py +++ b/repositories/biz_work_order.py @@ -1,7 +1,7 @@ """biz_work_order 工单仓储:查工单 + 流转条件更新(防并发,不 commit)。""" from __future__ import annotations -from sqlalchemy import select, update +from sqlalchemy import func, select, update from model.biz_work_order import BizWorkOrder from repositories.base import BaseRepository @@ -21,6 +21,8 @@ class BizWorkOrderRepo(BaseRepository): handler_id: int | None = None, status: str | None = None, customer_id: int | None = None, + limit: int = 10, + offset: int = 0, ) -> list[BizWorkOrder]: stmt = select(BizWorkOrder) if handler_id is not None: @@ -29,9 +31,26 @@ class BizWorkOrderRepo(BaseRepository): stmt = stmt.where(BizWorkOrder.status == status) if customer_id is not None: stmt = stmt.where(BizWorkOrder.customer_id == customer_id) - stmt = stmt.order_by(BizWorkOrder.id.desc()) + stmt = stmt.order_by(BizWorkOrder.id.desc()).limit(limit).offset(offset) return list((await self.db.scalars(stmt)).all()) + async def count_with_filter( + self, + *, + handler_id: int | None = None, + status: str | None = None, + customer_id: int | None = None, + ) -> int: + """统计工单筛选结果总数,供分页响应使用。""" + stmt = select(func.count()).select_from(BizWorkOrder) + if handler_id is not None: + stmt = stmt.where(BizWorkOrder.handler_id == handler_id) + if status is not None: + stmt = stmt.where(BizWorkOrder.status == status) + if customer_id is not None: + stmt = stmt.where(BizWorkOrder.customer_id == customer_id) + return int((await self.db.scalar(stmt)) or 0) + async def conditional_transition( self, work_order_id: int, *, from_status: str, to_status: str, **fields ) -> bool: diff --git a/repositories/fin_risk_alert.py b/repositories/fin_risk_alert.py index 78233c5..dd880c8 100644 --- a/repositories/fin_risk_alert.py +++ b/repositories/fin_risk_alert.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import datetime -from sqlalchemy import select, update +from sqlalchemy import func, select, update from model.fin_risk_alert import FinRiskAlert from repositories.base import BaseRepository @@ -13,16 +13,32 @@ class FinRiskAlertRepo(BaseRepository): model = FinRiskAlert async def list_by_status( - self, status: str | None = None, customer_id: int | None = None + self, + status: str | None = None, + customer_id: int | None = None, + *, + limit: int = 10, + offset: int = 0, ) -> list[FinRiskAlert]: stmt = select(FinRiskAlert) if status is not None: stmt = stmt.where(FinRiskAlert.status == status) if customer_id is not None: stmt = stmt.where(FinRiskAlert.customer_id == customer_id) - stmt = stmt.order_by(FinRiskAlert.id.desc()) + stmt = stmt.order_by(FinRiskAlert.id.desc()).limit(limit).offset(offset) return list((await self.db.scalars(stmt)).all()) + async def count_by_status( + self, status: str | None = None, customer_id: int | None = None + ) -> int: + """统计筛选条件下的预警总数,供后端分页响应使用。""" + stmt = select(func.count()).select_from(FinRiskAlert) + if status is not None: + stmt = stmt.where(FinRiskAlert.status == status) + if customer_id is not None: + stmt = stmt.where(FinRiskAlert.customer_id == customer_id) + return int((await self.db.scalar(stmt)) or 0) + async def conditional_handle( self, alert_id: int, diff --git a/repositories/memory_unit.py b/repositories/memory_unit.py index 61dee88..6f6ebb8 100644 --- a/repositories/memory_unit.py +++ b/repositories/memory_unit.py @@ -68,6 +68,36 @@ class MemoryUnitRepo(BaseRepository): ) return list((await self.db.scalars(statement)).all()) + async def list_for_customer_by_ids( + self, + customer_id: int, + ids: list[int], + *, + memory_type: str | None = None, + tag: str | None = None, + now: datetime | None = None, + ) -> list[MemoryUnit]: + """按 ID 集合取回本人有效记忆,用于向量召回命中后的主体回表。""" + if not ids: + return [] + now = now or datetime.now() + conditions = [ + MemoryUnit.customer_id == customer_id, + MemoryUnit.id.in_(ids), + MemoryUnit.status.in_(ACTIVE_STATUSES), + (MemoryUnit.valid_until.is_(None) | (MemoryUnit.valid_until > now)), + ] + if memory_type: + conditions.append(MemoryUnit.memory_type == memory_type) + if tag: + conditions.append(MemoryUnit.tag == tag) + statement = ( + select(MemoryUnit) + .where(*conditions) + .order_by(MemoryUnit.update_time.desc(), MemoryUnit.id.desc()) + ) + return list((await self.db.scalars(statement)).all()) + async def merge_evidence( self, memory: MemoryUnit, diff --git a/requirements.txt b/requirements.txt index a62e987..991220f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -12,6 +12,7 @@ pymilvus neo4j python-dotenv pypdf +sqlglot fastapi~=0.141.1 sqlalchemy~=2.0.52 @@ -34,4 +35,4 @@ neo4j~=6.3.0 redis~=8.1.0 pymilvus~=3.0.1 pydantic-settings~=2.15.0 -pyjwt~=2.13.0 \ No newline at end of file +pyjwt~=2.13.0 diff --git a/schemas/advisor_agent.py b/schemas/advisor_agent.py index 9d106e3..427189b 100644 --- a/schemas/advisor_agent.py +++ b/schemas/advisor_agent.py @@ -33,7 +33,8 @@ class AdvisorRebalanceRunReq(BaseModel): class AdvisorChatReq(BaseModel): - customer_id: int = Field(gt=0) + # 通用投顾问答不需要客户上下文;个性化推荐时再传入客户编号。 + customer_id: int | None = Field(default=None, gt=0) intent: Literal[ AGENT_INTENT_RECOMMEND, AGENT_INTENT_REBALANCE, @@ -41,7 +42,7 @@ class AdvisorChatReq(BaseModel): AGENT_INTENT_DIALOGUE_SCRIPT, AGENT_INTENT_DATA_QUERY, ] | None = None - query: str | None = Field(default=None, max_length=4000) + query: str = Field(min_length=1, max_length=4000) class AdvisorDataQueryReq(BaseModel): diff --git a/scripts/apply_memory_unit_upgrade.py b/scripts/apply_memory_unit_upgrade.py new file mode 100644 index 0000000..10f9b4a --- /dev/null +++ b/scripts/apply_memory_unit_upgrade.py @@ -0,0 +1,138 @@ +"""幂等执行 memory_unit 表结构升级(对齐 model/memory_unit.py)。 + +用法: + python scripts/apply_memory_unit_upgrade.py + +对应 SQL 版本见 sql/memory_unit_upgrade_20260913.sql。 +重复执行安全:ADD COLUMN 前检查 information_schema,MODIFY/UPDATE 本身幂等。 +""" +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from sqlalchemy import text + +from config.database import mysql as mysql_db +from config.database.mysql import get_session_factory +from config.settings import settings + +TABLE = "memory_unit" + +ADD_COLUMNS = { + "session_id": "session_id VARCHAR(64) NULL COMMENT '产生记忆的会话ID'", + "agent_run_id": "agent_run_id VARCHAR(64) NULL COMMENT '产生记忆的Agent运行ID'", + "evidence_ref": "evidence_ref JSON NULL COMMENT '证据引用列表(会话/消息溯源)'", + "historical_accuracy": "historical_accuracy DECIMAL(5,2) NOT NULL DEFAULT 0.50 COMMENT '历史准确率'", + "confidence_version": "confidence_version VARCHAR(32) NULL COMMENT '置信度算法版本'", + "confidence_reason": "confidence_reason VARCHAR(255) NULL COMMENT '置信度评分原因'", + "confidence_update_time": "confidence_update_time DATETIME NULL COMMENT '置信度更新时间'", + "memory_version": "memory_version INT NOT NULL DEFAULT 1 COMMENT '记忆版本号'", + "last_verified_at": "last_verified_at DATETIME NULL COMMENT '最近验证时间'", + "milvus_id": "milvus_id VARCHAR(128) NULL COMMENT 'Milvus向量主键'", + "graph_node_id": "graph_node_id VARCHAR(128) NULL COMMENT 'Neo4j图谱节点ID'", + "milvus_sync_status": "milvus_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '向量同步状态'", + "neo4j_sync_status": "neo4j_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '图谱同步状态'", + "sync_retry_count": "sync_retry_count INT NOT NULL DEFAULT 0 COMMENT '同步重试次数'", + "last_sync_error": "last_sync_error VARCHAR(500) NULL COMMENT '最近同步错误'", + "next_retry_at": "next_retry_at DATETIME NULL COMMENT '下次重试时间'", + "last_synced_at": "last_synced_at DATETIME NULL COMMENT '最近成功同步时间'", +} + + +async def columns_of(session, database: str) -> dict[str, str]: + rows = ( + await session.execute( + text( + "SELECT COLUMN_NAME, DATA_TYPE FROM information_schema.columns " + "WHERE TABLE_SCHEMA = :d AND TABLE_NAME = :t" + ), + {"d": database, "t": TABLE}, + ) + ).mappings().all() + return {str(r["COLUMN_NAME"]): str(r["DATA_TYPE"]).lower() for r in rows} + + +async def main() -> None: + async with get_session_factory()() as session: + db = settings.mysql.database + cols = await columns_of(session, db) + + added = [] + for name, ddl in ADD_COLUMNS.items(): + if name in cols: + print(f"[skip] 列已存在: {name}") + continue + await session.execute(text(f"ALTER TABLE {TABLE} ADD COLUMN {ddl}")) + added.append(name) + print(f"[ok] 新增列: {name}") + await session.commit() + + cols = await columns_of(session, db) + + if cols.get("valid_from") == "date" or cols.get("valid_until") == "date": + await session.execute( + text( + f"ALTER TABLE {TABLE} " + "MODIFY COLUMN valid_from DATETIME NULL COMMENT '生效起始时间', " + "MODIFY COLUMN valid_until DATETIME NULL COMMENT '失效时间'" + ) + ) + await session.commit() + print("[ok] valid_from/valid_until: DATE -> DATETIME") + else: + print("[skip] valid_from/valid_until 已是 DATETIME") + + result = await session.execute( + text(f"UPDATE {TABLE} SET status = 'candidate' WHERE status = 'active'") + ) + await session.commit() + print(f"[ok] status active->candidate: {result.rowcount} 行") + + await session.execute( + text( + f"ALTER TABLE {TABLE} MODIFY COLUMN status VARCHAR(16) NOT NULL " + "DEFAULT 'candidate' COMMENT '记忆状态(candidate/confirmed/expired/rejected/archived)'" + ) + ) + await session.commit() + print("[ok] status 默认值改为 candidate") + + result = await session.execute( + text( + f"UPDATE {TABLE} SET memory_type = 'SERVICE_FACT' " + "WHERE memory_type IS NULL OR memory_type = 'FACT'" + ) + ) + await session.commit() + print(f"[ok] memory_type NULL/FACT -> SERVICE_FACT: {result.rowcount} 行") + + await session.execute( + text( + f"ALTER TABLE {TABLE} MODIFY COLUMN memory_type VARCHAR(32) NOT NULL " + "COMMENT '记忆业务类型'" + ) + ) + await session.commit() + print("[ok] memory_type 收紧为 NOT NULL") + + final = await columns_of(session, db) + missing = [name for name in ADD_COLUMNS if name not in final] + if missing: + print(f"[fail] 仍有缺失列: {missing}") + raise SystemExit(1) + null_types = ( + await session.execute( + text(f"SELECT COUNT(*) AS n FROM {TABLE} WHERE memory_type IS NULL") + ) + ).scalar() + print(f"[done] 迁移完成,共新增 {len(added)} 列;memory_type NULL 残留: {null_types}") + + await mysql_db.dispose() + + +if __name__ == "__main__": + asyncio.run(main()) 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/advisor/agent_client.py b/service/advisor/agent_client.py index 416f272..791546b 100644 --- a/service/advisor/agent_client.py +++ b/service/advisor/agent_client.py @@ -138,7 +138,7 @@ class AdvisorAgentClient: customer_id: int | None = None, status: str | None = None, page: int = 1, - page_size: int = 20, + page_size: int = 10, ) -> dict: params = {"advisor_id": advisor_id, "page": page, "page_size": page_size} if customer_id is not None: diff --git a/service/advisor/audit.py b/service/advisor/audit.py index be67d6b..f9519a6 100644 --- a/service/advisor/audit.py +++ b/service/advisor/audit.py @@ -9,6 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from model.sys_user import SysUser from repositories.audit_log import AuditLogRepo +from utils.pagination import normalize_pagination, pagination_result def _audit_item(a) -> dict: @@ -34,17 +35,20 @@ async def list_ledger( start: datetime | None = None, end: datetime | None = None, page: int = 1, - page_size: int = 20, + page_size: int = 10, ) -> dict: + page, page_size, offset = normalize_pagination(page, page_size) repo = AuditLogRepo(db) items = await repo.list_by_advisor( user_id=user.id, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end, - limit=page_size, offset=(page - 1) * page_size, + limit=page_size, offset=offset, ) total = await repo.count_by_advisor( user_id=user.id, action=action, keyword=keyword, start=start, end=end ) - return {"total": total, "page": page, "page_size": page_size, "items": [_audit_item(a) for a in items]} + return pagination_result( + [_audit_item(a) for a in items], total, page=page, page_size=page_size + ) def _to_csv(rows: list[dict]) -> str: diff --git a/service/advisor/customers.py b/service/advisor/customers.py index 18d8a05..cbda198 100644 --- a/service/advisor/customers.py +++ b/service/advisor/customers.py @@ -28,6 +28,7 @@ from service.advisor.audit_writer import write_audit from service.advisor.masking import mask_name, mask_phone from service.advisor.permissions import ensure_customer_owned, require_owned_relation from utils.exceptions import ForbiddenError, NotFoundError +from utils.pagination import normalize_pagination, pagination_result # 持仓中状态(与 service/holdings.py 口径一致) _HOLDING_STATUS = "持有中" @@ -57,12 +58,13 @@ async def list_customers( status: str | None = None, keyword: str | None = None, page: int = 1, - page_size: int = 20, + page_size: int = 10, ) -> dict: + page, page_size, offset = normalize_pagination(page, page_size) repo = CustomerRelationRepo(db) rows = await repo.list_customer_rows( advisor_id=user.id, status=status, keyword=keyword, - limit=page_size, offset=(page - 1) * page_size, + limit=page_size, offset=offset, ) total = await repo.count_customer_rows(advisor_id=user.id, status=status, keyword=keyword) items = [] @@ -80,7 +82,7 @@ async def list_customers( else None, } ) - return {"total": total, "page": page, "page_size": page_size, "items": items} + return pagination_result(items, total, page=page, page_size=page_size) async def get_customer( @@ -133,19 +135,19 @@ async def get_customer_holdings( async def get_customer_reports( - db: AsyncSession, user: SysUser, customer_id: int, *, page: int = 1, page_size: int = 20 + db: AsyncSession, user: SysUser, customer_id: int, *, page: int = 1, page_size: int = 10 ) -> dict: + page, page_size, offset = normalize_pagination(page, page_size) await ensure_customer_owned(db, user.id, customer_id) repo = AdvisorReportRepo(db) items = await repo.list_by_customer( customer_id, advisor_id=user.id, limit=page_size, - offset=(page - 1) * page_size, + offset=offset, ) - return { - "total": await repo.count_by_advisor(advisor_id=user.id, customer_id=customer_id), - "items": [ + return pagination_result( + [ { "report_id": r.report_id, "intent": r.intent, @@ -155,7 +157,10 @@ async def get_customer_reports( } for r in items ], - } + await repo.count_by_advisor(advisor_id=user.id, customer_id=customer_id), + page=page, + page_size=page_size, + ) async def update_relation( diff --git a/service/advisor/drafts.py b/service/advisor/drafts.py index c57e5a7..d12b259 100644 --- a/service/advisor/drafts.py +++ b/service/advisor/drafts.py @@ -144,7 +144,7 @@ async def list_drafts( customer_id: int | None = None, status: str | None = None, page: int = 1, - page_size: int = 20, + page_size: int = 10, ) -> dict: result = await get_agent_client().draft_list( auth_header=auth_header, diff --git a/service/advisor/todos.py b/service/advisor/todos.py index 9753b80..3399f8e 100644 --- a/service/advisor/todos.py +++ b/service/advisor/todos.py @@ -11,6 +11,7 @@ from model.sys_user import SysUser from repositories.advisor_todo import AdvisorTodoRepo from schemas.advisor import TodoHandleReq from utils.exceptions import ForbiddenError, NotFoundError, ParamError +from utils.pagination import normalize_pagination, pagination_result def _todo_item(t: AdvisorTodo) -> dict: @@ -36,15 +37,18 @@ async def list_todos( status: str | None = None, todo_type: str | None = None, page: int = 1, - page_size: int = 20, + page_size: int = 10, ) -> dict: + page, page_size, offset = normalize_pagination(page, page_size) repo = AdvisorTodoRepo(db) items = await repo.list_by_advisor( advisor_id=user.id, status=status, todo_type=todo_type, - limit=page_size, offset=(page - 1) * page_size, + limit=page_size, offset=offset, ) total = await repo.count_by_advisor(advisor_id=user.id, status=status, todo_type=todo_type) - return {"total": total, "page": page, "page_size": page_size, "items": [_todo_item(t) for t in items]} + return pagination_result( + [_todo_item(t) for t in items], total, page=page, page_size=page_size + ) async def handle_todo(db: AsyncSession, user: SysUser, todo_id: int, req: TodoHandleReq) -> dict: diff --git a/service/advisor/visits.py b/service/advisor/visits.py index f1a0429..8437691 100644 --- a/service/advisor/visits.py +++ b/service/advisor/visits.py @@ -15,6 +15,7 @@ from schemas.advisor import VisitCreateReq, VisitUpdateReq from service.advisor.permissions import ensure_customer_owned from service.advisor.audit_writer import write_audit from utils.exceptions import NotFoundError, ParamError +from utils.pagination import normalize_pagination, pagination_result # 合规话术库(内置,标准化投教/市场解读/调仓沟通)。 # 假设:V1.0 内置常量,后续可迁移 sys_config 运营化;内容不含承诺收益等敏感词。 @@ -40,17 +41,20 @@ def _visit_item(v: AdvisorVisitRecord) -> dict: async def list_visits( db: AsyncSession, user: SysUser, *, customer_id: int | None = None, - page: int = 1, page_size: int = 20, + page: int = 1, page_size: int = 10, ) -> dict: + page, page_size, offset = normalize_pagination(page, page_size) if customer_id is not None: await ensure_customer_owned(db, user.id, customer_id) repo = AdvisorVisitRecordRepo(db) items = await repo.list_by_advisor( advisor_id=user.id, customer_id=customer_id, - limit=page_size, offset=(page - 1) * page_size, + limit=page_size, offset=offset, ) total = await repo.count_by_advisor(advisor_id=user.id, customer_id=customer_id) - return {"total": total, "page": page, "page_size": page_size, "items": [_visit_item(v) for v in items]} + return pagination_result( + [_visit_item(v) for v in items], total, page=page, page_size=page_size + ) async def create_visit(db: AsyncSession, user: SysUser, req: VisitCreateReq) -> dict: @@ -114,19 +118,20 @@ def list_talk_templates() -> list[dict]: async def list_touch_logs( - db: AsyncSession, user: SysUser, customer_id: int, *, page: int = 1, page_size: int = 20 + db: AsyncSession, user: SysUser, customer_id: int, *, page: int = 1, page_size: int = 10 ) -> dict: """触达留痕:站内信(发送触达)+ 回访记录(人工沟通)聚合。 假设:V1.0 触达通道仅站内信与回访;短信/企微/电话统一落在回访记录。 """ + page, page_size, offset = normalize_pagination(page, page_size) await ensure_customer_owned(db, user.id, customer_id) messages = await SysMessageRepo(db).list_by_user( - customer_id, limit=page_size, offset=(page - 1) * page_size + customer_id, limit=page_size, offset=offset ) visits = await AdvisorVisitRecordRepo(db).list_by_advisor( advisor_id=user.id, customer_id=customer_id, - limit=page_size, offset=(page - 1) * page_size, + limit=page_size, offset=offset, ) return { "messages": [ diff --git a/service/advisor_agent/draft.py b/service/advisor_agent/draft.py index 0866e0a..79d50a4 100644 --- a/service/advisor_agent/draft.py +++ b/service/advisor_agent/draft.py @@ -18,6 +18,7 @@ from service.advisor_agent.compliance import ensure_safe_content from repositories.sensitive_word import SensitiveWordRepo from model.advisor_draft import AdvisorDraft from utils.exceptions import ApiError +from utils.pagination import normalize_pagination, pagination_result def build_generated_content(content: str) -> str: @@ -133,18 +134,19 @@ async def list_drafts( customer_id: int | None = None, status: str | None = None, page: int = 1, - page_size: int = 20, + page_size: int = 10, ) -> dict: - page = max(1, page) - page_size = min(100, max(1, page_size)) + page, page_size, offset = normalize_pagination(page, page_size) total, items = await repo.list_drafts( advisor_id=advisor_id, customer_id=customer_id, status=status, limit=page_size, - offset=(page - 1) * page_size, + offset=offset, + ) + return pagination_result( + [summarize_draft(item) for item in items], total, page=page, page_size=page_size ) - return {"total": total, "items": [summarize_draft(item) for item in items]} async def save_draft( 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 2cc2804..0d14f54 100644 --- a/service/client_agent/runtime.py +++ b/service/client_agent/runtime.py @@ -2,20 +2,34 @@ from __future__ import annotations +import asyncio import contextvars 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) @@ -59,13 +73,20 @@ class MemoryConversationContext: class MemoryAwareClientAgent: - """在现有客服 Agent 外包裹记忆召回、候选保存和降级处理。""" + """在现有客服 Agent 外包裹记忆召回、候选保存和降级处理。 - def __init__(self, *, agent, memory_service, context, extractor=None): + 候选记忆保存默认放入后台任务执行(background_saves=True),不阻塞 + 客服响应;测试或需要确定性顺序的场景可设为 False 改回同步执行。 + """ + + def __init__(self, *, agent, memory_service, context, extractor=None, + background_saves: bool = True): self.agent = agent self.memory_service = memory_service self.context = context self.extractor = extractor + self.background_saves = background_saves + self._pending_saves: set[asyncio.Task] = set() async def handle(self, session_id: str, query: str, *, trace_id: str, customer_id: int) -> dict: """执行记忆召回、客服回答、消息写入和候选记忆保存。""" @@ -91,8 +112,10 @@ 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) - if self.extractor is not None: + result = await self.agent.handle( + session_id, query, trace_id=trace_id, customer_id=customer_id + ) + if self.extractor is not None and not self.background_saves: await self._save_candidates( customer_id, session_id, @@ -102,6 +125,15 @@ class MemoryAwareClientAgent: trace_id=trace_id, ) result["memory_warnings"] = list(warnings) + if self.extractor is not None and self.background_saves: + self._spawn_save( + customer_id, + session_id, + query, + memory_context, + warnings, + trace_id=trace_id, + ) return result finally: _active_customer.reset(message_token) @@ -109,6 +141,35 @@ class MemoryAwareClientAgent: _active_warnings.reset(warnings_token) _active_memory_context.reset(context_token) + def _spawn_save( + self, + customer_id, + session_id, + query, + context, + warnings, + *, + trace_id: str | None = None, + ) -> None: + """把候选记忆保存放入后台任务;任务异常已自捕获,不击穿响应。""" + task = asyncio.create_task( + self._save_candidates( + customer_id, + session_id, + query, + context, + warnings, + trace_id=trace_id, + ) + ) + self._pending_saves.add(task) + task.add_done_callback(self._pending_saves.discard) + + async def wait_for_pending_saves(self) -> None: + """等待全部后台保存完成,供测试与优雅退出使用。""" + if self._pending_saves: + await asyncio.gather(*list(self._pending_saves), return_exceptions=True) + async def _save_candidates( self, customer_id, @@ -142,9 +203,24 @@ class MemoryAwareClientAgent: try: evidence_count = 1 if candidate.get("signal_type") == "interest_query": - # 兴趣主题首次出现即保存为候选,后续由长期记忆按精确内容合并证据。 + # 兴趣主题先计数:达到阈值才固化为长期记忆,避免单次关注污染画像。 candidate["memory_type"] = "CUSTOMER_PREFERENCE" candidate["source"] = "dialogue_inferred" + try: + _, reached = await self.memory_service.record_interest_signal( + customer_id=customer_id, tag=candidate["tag"] + ) + except Exception as exc: + logger.exception( + "client interest signal failed: trace_id=%s customer_id=%s tag=%s", + trace_id, + customer_id, + candidate.get("tag"), + ) + warnings.append(f"interest_signal_failed:{type(exc).__name__}") + continue + if not reached: + continue memory = MemoryUnitDTO( customer_id=customer_id, session_id=session_id, @@ -180,6 +256,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 +379,8 @@ def build_client_runtime( llm_client, config_getter, audit_writer, + db_session_factory=None, + schema_database=None, memory_service=None, memory_extractor=None, ): @@ -207,8 +401,13 @@ def build_client_runtime( config_getter=config_getter, ) - async def recognize(query): - return await intent_recognize(query, llm_client=llm_client) + async def recognize(query, history=None): + return await intent_recognize( + query, + llm_client=llm_client, + history=history, + config_getter=config_getter, + ) async def generate(messages): memory_context = _active_memory_context.get() @@ -228,6 +427,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, @@ -235,6 +445,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 0138830..d99d7b6 100644 --- a/service/customer_agent/chat.py +++ b/service/customer_agent/chat.py @@ -2,16 +2,24 @@ from __future__ import annotations import json +import logging import re from inspect import isawaitable -from rag.intent import Intent +from rag.intent import Intent, IntentResult + + +logger = logging.getLogger(__name__) 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): @@ -33,6 +41,7 @@ class AnonymousCustomerAgent: generate_answer, audit_writer, config_getter, + data_query=None, ): self.context = context self.rag_retrieve = rag_retrieve @@ -40,10 +49,26 @@ 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字符") + # 先取历史再写入当前问题,保证意图识别拿到的历史不含本轮输入;取不到历史不阻断请求 + try: + history = await self.context.get(session_id) + except Exception: + logger.exception("load conversation history failed: session_id=%s", session_id) + history = [] await self.context.append(session_id, "user", query) if self._contains_sensitive_input(query): await _maybe_await(self.audit_writer( @@ -52,8 +77,13 @@ class AnonymousCustomerAgent: session_id=session_id, )) - intent = await _maybe_await(self.intent_recognize(query)) + recognized = await _maybe_await(self.intent_recognize(query, history)) + if isinstance(recognized, IntentResult): + intent, search_query = recognized.intent, recognized.query + else: + intent, search_query = recognized, query sources = [] + data_query_meta = None if intent == Intent.GUIDE_PURCHASE: answer = await _config( self.config_getter, @@ -75,9 +105,25 @@ 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: - sources = await _maybe_await(self.rag_retrieve(query, None)) + # 用补全指代后的问题检索,省略主语的追问才能命中 + sources = await _maybe_await(self.rag_retrieve(search_query, None)) except Exception: sources = [] if not sources: @@ -114,12 +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/customer_agent/runtime.py b/service/customer_agent/runtime.py index 45b5b0a..038c2fe 100644 --- a/service/customer_agent/runtime.py +++ b/service/customer_agent/runtime.py @@ -36,8 +36,13 @@ def build_anonymous_runtime( config_getter=config_getter, ) - async def recognize(query): - return await intent_recognize(query, llm_client=llm_client) + async def recognize(query, history=None): + return await intent_recognize( + query, + llm_client=llm_client, + history=history, + config_getter=config_getter, + ) async def generate(messages): return await generate_answer( diff --git a/service/memory/facade.py b/service/memory/facade.py index 868558d..98f2112 100644 --- a/service/memory/facade.py +++ b/service/memory/facade.py @@ -97,7 +97,7 @@ class MemoryService: warnings.append(f"customer_product_recall_failed:{type(exc).__name__}") try: memories, memory_warnings = await self.long_term.recall( - db, customer_id, limit=max(limit, 10) + db, customer_id, limit=max(limit, 10), query=query ) warnings.extend(memory_warnings) except Exception as exc: diff --git a/service/memory/long_term.py b/service/memory/long_term.py index 5106a3c..6b1f90e 100644 --- a/service/memory/long_term.py +++ b/service/memory/long_term.py @@ -51,10 +51,10 @@ class LongTermMemoryService: ) entity = existing else: - # final_score 是召回重排阶段的临时分数,不属于 MySQL 主体字段。 + # final_score/semantic_similarity 是召回重排阶段的临时分数,不属于 MySQL 主体字段。 values = memory.model_dump( mode="json", - exclude={"id", "milvus_id", "graph_node_id", "final_score"}, + exclude={"id", "milvus_id", "graph_node_id", "final_score", "semantic_similarity"}, ) now = datetime.now() values["last_verified_at"] = values.get("last_verified_at") or now @@ -147,12 +147,50 @@ class LongTermMemoryService: memory_type: str | None = None, tag: str | None = None, limit: int = 100, + query: str | None = None, ) -> tuple[list[MemoryUnitDTO], list[str]]: - """按客户、类型、标签和有效期召回主体记忆。""" + """召回主体记忆;带 query 时叠加 Milvus 语义召回并标注相似度。""" entities = await self.repository_factory(db).list_for_customer( customer_id, memory_type=memory_type, tag=tag, limit=limit ) - return [self._to_dto(entity) for entity in entities], [] + dtos = [self._to_dto(entity) for entity in entities] + if query is None or not str(query).strip(): + return dtos, [] + + warnings: list[str] = [] + try: + vector = self.embedder(str(query)) + if isawaitable(vector): + vector = await vector + hits = await self.milvus_store.search(vector, customer_id, limit=max(limit, 1)) + except Exception as exc: + return dtos, [f"semantic_recall_failed:{type(exc).__name__}"] + + similarities: dict[int, float] = {} + for hit in hits: + raw_id = str(hit.get("memory_id", "")) + if raw_id.lstrip("-").isdigit(): + similarities[int(raw_id)] = float(hit.get("distance") or 0.0) + + matched: set[int] = set() + for dto in dtos: + if dto.id in similarities: + dto.semantic_similarity = similarities[dto.id] + matched.add(dto.id) + + missing_ids = [mid for mid in similarities if mid not in matched] + if missing_ids: + try: + extra_entities = await self.repository_factory(db).list_for_customer_by_ids( + customer_id, missing_ids, memory_type=memory_type, tag=tag + ) + for entity in extra_entities: + dto = self._to_dto(entity) + dto.semantic_similarity = similarities.get(dto.id) + dtos.append(dto) + except Exception as exc: + warnings.append(f"semantic_recall_fetch_failed:{type(exc).__name__}") + return dtos, warnings async def refresh_confidence( self, db, customer_id: int, *, limit: int = 500 diff --git a/service/memory/milvus_memory.py b/service/memory/milvus_memory.py index 40a041a..3016f19 100644 --- a/service/memory/milvus_memory.py +++ b/service/memory/milvus_memory.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import Any + from pymilvus import AsyncMilvusClient, DataType from config.database.milvus import client as configured_client @@ -87,15 +89,43 @@ class MilvusMemoryStore: return memory_id async def search(self, vector: list[float], customer_id: int, *, limit: int = 10) -> list[dict]: - """按客户 ID 过滤向量查询结果。""" + """按客户 ID 过滤向量查询,返回归一化的命中列表。""" await self.ensure_collection() - return await self.client.search( + raw = await self.client.search( collection_name=self.collection_name, data=[vector], limit=limit, filter=f"customer_id == {int(customer_id)}", output_fields=["memory_id", "customer_id", "memory_type", "tag", "content", "status"], ) + return self._normalize_hits(raw) + + @staticmethod + def _normalize_hits(raw: Any) -> list[dict]: + """将 pymilvus 返回结构收敛为 memory_id + distance 的扁平列表。""" + hits: list[dict] = [] + for batch in raw or []: + for hit in batch or []: + if not isinstance(hit, dict): + continue + entity = hit.get("entity") or {} + memory_id = entity.get("memory_id") or hit.get("id") + if memory_id is None: + continue + try: + distance = float(hit.get("distance", 0.0)) + except (TypeError, ValueError): + distance = 0.0 + hits.append( + { + "memory_id": str(memory_id), + "distance": max(0.0, min(1.0, distance)), + "tag": entity.get("tag"), + "content": entity.get("content"), + "status": entity.get("status"), + } + ) + return hits async def delete(self, memory_id: int | str) -> None: """删除一条客户记忆向量。""" diff --git a/service/memory/schemas.py b/service/memory/schemas.py index 47a909a..d3f3be2 100644 --- a/service/memory/schemas.py +++ b/service/memory/schemas.py @@ -79,6 +79,7 @@ class MemoryUnitDTO(BaseModel): confidence_reason: str | None = Field(default=None, max_length=255) confidence_update_time: datetime | None = None final_score: float | None = Field(default=None, ge=0.0, le=1.0) + semantic_similarity: float | None = Field(default=None, ge=0.0, le=1.0) evidence_count: int = Field(default=0, ge=0) recall_count: int = Field(default=0, ge=0) status: MemoryStatus = MemoryStatus.CANDIDATE 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()), diff --git a/service/product.py b/service/product.py index 0bb05b4..24c1993 100644 --- a/service/product.py +++ b/service/product.py @@ -9,6 +9,7 @@ from model.fin_product import FinProduct from repositories.fund_nav import FundNavRepo from repositories.product import ProductRepo from utils.exceptions import NotFoundError, ParamError +from utils.pagination import normalize_pagination, pagination_result def _to_item(p: FinProduct) -> dict: @@ -44,10 +45,7 @@ async def list_products( sort_by: str = "create_time", sort_order: str = "desc", ) -> dict: - if page < 1: - raise ParamError("page 必须 >= 1") - if page_size < 1 or page_size > 100: - raise ParamError("page_size 必须在 1-100 之间") + page, page_size, offset = normalize_pagination(page, page_size, max_page_size=100) if sort_order not in ("asc", "desc"): raise ParamError("sort_order 仅支持 asc/desc") @@ -59,14 +57,11 @@ async def list_products( sort_by=sort_by, sort_order=sort_order, limit=page_size, - offset=(page - 1) * page_size, + offset=offset, + ) + return pagination_result( + [_to_item(p) for p in items], total, page=page, page_size=page_size ) - return { - "total": total, - "page": page, - "page_size": page_size, - "items": [_to_item(p) for p in items], - } def _calc_metrics(rows: list) -> dict: diff --git a/service/risk/handle.py b/service/risk/handle.py index 340948b..7bd4ae0 100644 --- a/service/risk/handle.py +++ b/service/risk/handle.py @@ -24,6 +24,7 @@ from repositories.trade_order import TradeOrderRepo from schemas.risk import RiskAlertResp from service.risk.settle import settle from utils.exceptions import NotFoundError, ParamError +from utils.pagination import normalize_pagination, pagination_result from utils.order_no import gen_order_no # 预警级别 → 工单优先级 @@ -253,6 +254,24 @@ async def freeze(db: AsyncSession, handler: SysUser, alert_id: int) -> dict: raise -async def list_alerts(db: AsyncSession, status: str | None = None) -> list[RiskAlertResp]: - alerts = await FinRiskAlertRepo(db).list_by_status(status) - return [_alert_resp(a) for a in alerts] +async def list_alerts( + db: AsyncSession, + status: str | None = None, + *, + page: int = 1, + page_size: int = 10, +) -> dict: + """分页查询风控预警,默认每页 10 条。""" + page, page_size, offset = normalize_pagination(page, page_size) + repo = FinRiskAlertRepo(db) + alerts = await repo.list_by_status( + status, + limit=page_size, + offset=offset, + ) + return pagination_result( + [_alert_resp(a).model_dump(mode="json") for a in alerts], + await repo.count_by_status(status), + page=page, + page_size=page_size, + ) diff --git a/service/work_order.py b/service/work_order.py index 6bf18a8..84b9a50 100644 --- a/service/work_order.py +++ b/service/work_order.py @@ -15,6 +15,7 @@ from model.sys_user import SysUser from repositories.biz_work_order import BizWorkOrderRepo from schemas.work_order import WorkOrderResp from utils.exceptions import NotFoundError, ParamError +from utils.pagination import normalize_pagination, pagination_result def _audit(db: AsyncSession, user: SysUser, module: str, action: str, target, detail: str) -> None: @@ -49,9 +50,27 @@ def _resp(wo: BizWorkOrder) -> WorkOrderResp: ) -async def list_work_orders(db: AsyncSession, status: str | None = None) -> list[WorkOrderResp]: - orders = await BizWorkOrderRepo(db).list_with_filter(status=status) - return [_resp(o) for o in orders] +async def list_work_orders( + db: AsyncSession, + status: str | None = None, + *, + page: int = 1, + page_size: int = 10, +) -> dict: + """分页查询公共工单,默认每页 10 条。""" + page, page_size, offset = normalize_pagination(page, page_size) + repo = BizWorkOrderRepo(db) + orders = await repo.list_with_filter( + status=status, + limit=page_size, + offset=offset, + ) + return pagination_result( + [_resp(o).model_dump(mode="json") for o in orders], + await repo.count_with_filter(status=status), + page=page, + page_size=page_size, + ) async def get_work_order(db: AsyncSession, work_order_id: int) -> WorkOrderResp: diff --git a/sql/memory_unit_upgrade_20260913.sql b/sql/memory_unit_upgrade_20260913.sql new file mode 100644 index 0000000..022a7b9 --- /dev/null +++ b/sql/memory_unit_upgrade_20260913.sql @@ -0,0 +1,45 @@ +-- memory_unit 表结构升级:对齐 model/memory_unit.py(34 列) +-- 日期:2026-09-13 执行方式:scripts/apply_memory_unit_upgrade.py(幂等)或本文件手工执行 +-- 背景:DB 为旧版 20 列结构,缺 session_id/evidence_ref/置信度/同步状态等 17 列, +-- 导致客服 Agent 长期记忆保存与召回抛 OperationalError(memory_warnings 来源)。 + +-- 1) 新增 17 列 +ALTER TABLE memory_unit + ADD COLUMN session_id VARCHAR(64) NULL COMMENT '产生记忆的会话ID', + ADD COLUMN agent_run_id VARCHAR(64) NULL COMMENT '产生记忆的Agent运行ID', + ADD COLUMN evidence_ref JSON NULL COMMENT '证据引用列表(会话/消息溯源)', + ADD COLUMN historical_accuracy DECIMAL(5,2) NOT NULL DEFAULT 0.50 COMMENT '历史准确率', + ADD COLUMN confidence_version VARCHAR(32) NULL COMMENT '置信度算法版本', + ADD COLUMN confidence_reason VARCHAR(255) NULL COMMENT '置信度评分原因', + ADD COLUMN confidence_update_time DATETIME NULL COMMENT '置信度更新时间', + ADD COLUMN memory_version INT NOT NULL DEFAULT 1 COMMENT '记忆版本号', + ADD COLUMN last_verified_at DATETIME NULL COMMENT '最近验证时间', + ADD COLUMN milvus_id VARCHAR(128) NULL COMMENT 'Milvus向量主键', + ADD COLUMN graph_node_id VARCHAR(128) NULL COMMENT 'Neo4j图谱节点ID', + ADD COLUMN milvus_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '向量同步状态', + ADD COLUMN neo4j_sync_status VARCHAR(16) NOT NULL DEFAULT 'pending' COMMENT '图谱同步状态', + ADD COLUMN sync_retry_count INT NOT NULL DEFAULT 0 COMMENT '同步重试次数', + ADD COLUMN last_sync_error VARCHAR(500) NULL COMMENT '最近同步错误', + ADD COLUMN next_retry_at DATETIME NULL COMMENT '下次重试时间', + ADD COLUMN last_synced_at DATETIME NULL COMMENT '最近成功同步时间'; + +-- 2) 类型对齐:DATE 无法承载时间语义,扩为 DATETIME +ALTER TABLE memory_unit + MODIFY COLUMN valid_from DATETIME NULL COMMENT '生效起始时间', + MODIFY COLUMN valid_until DATETIME NULL COMMENT '失效时间'; + +-- 3) 状态语义迁移:旧默认 active → 新枚举 candidate,并更新表默认值 +UPDATE memory_unit SET status = 'candidate' WHERE status = 'active'; +ALTER TABLE memory_unit + MODIFY COLUMN status VARCHAR(16) NOT NULL DEFAULT 'candidate' COMMENT '记忆状态(candidate/confirmed/expired/rejected/archived)'; + +-- 4) memory_type 旧枚举归并:NULL 与 'FACT' → 'SERVICE_FACT',再收紧 NOT NULL +UPDATE memory_unit SET memory_type = 'SERVICE_FACT' + WHERE memory_type IS NULL OR memory_type = 'FACT'; +ALTER TABLE memory_unit + MODIFY COLUMN memory_type VARCHAR(32) NOT NULL COMMENT '记忆业务类型'; + +-- 说明: +-- - id/customer_id 为 BIGINT UNSIGNED,与 ORM BigInteger 兼容,不做变更; +-- - 遗留列 conflict_count/dimension/polarity 无 ORM 映射,保留不动; +-- - NUMERIC(5,2) 在 MySQL 中即 DECIMAL(5,2),探针报告的差异为别名,非真实差异。 diff --git a/utils/pagination.py b/utils/pagination.py new file mode 100644 index 0000000..d9fba89 --- /dev/null +++ b/utils/pagination.py @@ -0,0 +1,39 @@ +"""分页通用工具:统一页码校验、偏移量和列表响应结构。""" +from __future__ import annotations + +from collections.abc import Sequence +from math import ceil +from typing import Any + + +def normalize_pagination( + page: int = 1, + page_size: int = 10, + *, + max_page_size: int = 10, +) -> tuple[int, int, int]: + """规范化分页参数,返回 ``(page, page_size, offset)``。""" + page = max(1, int(page)) + page_size = min(max_page_size, max(1, int(page_size))) + return page, page_size, (page - 1) * page_size + + +def pagination_result( + items: Sequence[Any], + total: int, + *, + page: int, + page_size: int, +) -> dict[str, Any]: + """构造统一分页响应,包含总页数供前端直接计算翻页状态。""" + total = max(0, int(total)) + return { + "items": list(items), + "total": total, + "page": page, + "page_size": page_size, + "total_pages": ceil(total / page_size) if total else 0, + } + + +__all__ = ["normalize_pagination", "pagination_result"]