diff --git a/agent/advisor_agent/__init__.py b/agent/advisor_agent/__init__.py new file mode 100644 index 0000000..e50a803 --- /dev/null +++ b/agent/advisor_agent/__init__.py @@ -0,0 +1,8 @@ +"""投顾 Agent 领域层。 + +该模块只负责生成投顾内部草稿,不负责客户触达、交易下单或交易状态回流。 +""" + +package_name = "advisor_agent" + +__all__ = ["package_name"] diff --git a/agent/advisor_agent/auth.py b/agent/advisor_agent/auth.py new file mode 100644 index 0000000..4987af9 --- /dev/null +++ b/agent/advisor_agent/auth.py @@ -0,0 +1,34 @@ +"""投顾 Agent 的角色与客户关系授权规则。""" +from __future__ import annotations + +from common.common_const import ( + CUSTOMER_REL_STATUS_SIGNED, + CUSTOMER_REL_STATUS_UNSIGNED, + EMPLOYEE_ROLE_ADVISOR, + ERR_CODE_FORBIDDEN_CUSTOMER, +) +from utils.exceptions import ApiError +from repositories.customer_relation import CustomerRelationRepo + + +def ensure_advisor_role(user) -> bool: + if ( + getattr(user, "user_type", None) != "EMPLOYEE" + or getattr(user, "employee_role", None) != EMPLOYEE_ROLE_ADVISOR + ): + raise ApiError(ERR_CODE_FORBIDDEN_CUSTOMER, "无权操作该客户数据") + return True + + +def relation_allows_access(status: str) -> bool: + return status in {CUSTOMER_REL_STATUS_UNSIGNED, CUSTOMER_REL_STATUS_SIGNED} + + +async def ensure_customer_access(db, *, advisor_id: int, customer_id: int): + relation = await CustomerRelationRepo(db).get_active_relation( + customer_id=customer_id, + advisor_id=advisor_id, + ) + if relation is None: + raise ApiError(ERR_CODE_FORBIDDEN_CUSTOMER, "无权操作该客户数据") + return relation diff --git a/agent/advisor_agent/drafts.py b/agent/advisor_agent/drafts.py new file mode 100644 index 0000000..aa68b49 --- /dev/null +++ b/agent/advisor_agent/drafts.py @@ -0,0 +1,27 @@ +"""投顾 Agent 草稿状态机基础规则。""" +from __future__ import annotations + +from common.common_const import ( + DRAFT_STATUS_DISCARDED, + DRAFT_STATUS_DRAFT, + ERR_CODE_DRAFT_NOT_FOUND, +) + + +class DraftStateError(Exception): + def __init__(self, code: int, message: str): + self.code = code + self.message = message + super().__init__(message) + + +def transition_draft_status(status: str, operation: str) -> str: + if status == DRAFT_STATUS_DRAFT and operation == "discard": + return DRAFT_STATUS_DISCARDED + if status == DRAFT_STATUS_DISCARDED: + raise DraftStateError(ERR_CODE_DRAFT_NOT_FOUND, "草稿不存在或者已废弃") + if status != DRAFT_STATUS_DRAFT: + raise DraftStateError(ERR_CODE_DRAFT_NOT_FOUND, "草稿状态无效") + if operation == "save": + return DRAFT_STATUS_DRAFT + raise DraftStateError(ERR_CODE_DRAFT_NOT_FOUND, "不支持的草稿操作") diff --git a/agent/advisor_agent/fallback.py b/agent/advisor_agent/fallback.py new file mode 100644 index 0000000..bff93d9 --- /dev/null +++ b/agent/advisor_agent/fallback.py @@ -0,0 +1,52 @@ +"""投顾 Agent 的统一超时与降级执行器。""" +from __future__ import annotations + +import asyncio +import inspect +import logging +from dataclasses import dataclass +from typing import Awaitable, Callable, TypeVar + +from utils.exceptions import ApiError + + +T = TypeVar("T") +logger = logging.getLogger("advisor_agent.fallback") + + +@dataclass(frozen=True) +class FallbackResult: + value: object + degraded: bool + code: int | None = None + + +async def call_with_fallback( + primary: Callable[[], Awaitable[T]], + secondary: Callable[[], Awaitable[T]] | None, + *, + timeout: float, + degraded_code: int, + on_degraded: Callable[[Exception], None] | None = None, +) -> FallbackResult: + try: + return FallbackResult( + value=await asyncio.wait_for(primary(), timeout=timeout), + degraded=False, + ) + except Exception as primary_error: + logger.warning("advisor dependency degraded; using fallback", exc_info=primary_error) + if on_degraded is not None: + try: + callback_result = on_degraded(primary_error) + if inspect.isawaitable(callback_result): + await callback_result + except Exception: + logger.warning("advisor degradation audit callback failed", exc_info=True) + if secondary is None: + raise ApiError(degraded_code, "Agent核心服务调用失败") from primary_error + try: + value = await asyncio.wait_for(secondary(), timeout=timeout) + except Exception as secondary_error: + raise ApiError(degraded_code, "Agent降级服务调用失败") from secondary_error + return FallbackResult(value=value, degraded=True, code=degraded_code) diff --git a/agent/advisor_agent/intent/__init__.py b/agent/advisor_agent/intent/__init__.py new file mode 100644 index 0000000..4b1c513 --- /dev/null +++ b/agent/advisor_agent/intent/__init__.py @@ -0,0 +1 @@ +"""投顾 Agent 的业务意图实现。""" diff --git a/agent/advisor_agent/intent/draft_generation.py b/agent/advisor_agent/intent/draft_generation.py new file mode 100644 index 0000000..2e28b2c --- /dev/null +++ b/agent/advisor_agent/intent/draft_generation.py @@ -0,0 +1,97 @@ +"""将意图计算结果组装为可持久化的 Agent 草稿。""" +from __future__ import annotations + +from decimal import Decimal + +from agent.advisor_agent.intent.rebalance import build_rebalance_plan +from agent.advisor_agent.intent.recommend import recommend_candidates +from common.common_const import ( + AGENT_INTENT_REBALANCE, + AGENT_INTENT_RECOMMEND, + DRAFT_STATUS_DRAFT, +) +from service.advisor_agent.draft import build_generated_content + + +def _recommend_markdown(items: list[dict]) -> str: + lines = ["# 基金推荐草稿", ""] + for item in items: + lines.append( + f"- {item.get('product_name', item.get('product_code'))}" + f"({item.get('product_code')},风险等级 {item.get('risk_level')})" + ) + return "\n".join(lines) + + +def build_recommendation_draft( + *, + customer_id: int, + advisor_id: int, + customer_risk: str, + candidates: list[dict], + relation_status: str, + memories: list[dict] | None = None, +) -> dict: + items = recommend_candidates( + customer_risk, + candidates, + memories=memories or [], + ) + return { + "customer_id": customer_id, + "advisor_id": advisor_id, + "intent": AGENT_INTENT_RECOMMEND, + "title": "基金推荐草稿", + "status": DRAFT_STATUS_DRAFT, + "structured_data": { + "customer_risk": customer_risk, + "items": items, + "relation_status": relation_status, + }, + "content": build_generated_content(_recommend_markdown(items)), + "disclaimer_ok": True, + } + + +def build_rebalance_draft( + *, + customer_id: int, + advisor_id: int, + customer_risk: str, + relation_status: str, + holdings: list[dict], + target_allocation: dict[str, int | float | Decimal], + threshold: Decimal, + candidates: list[dict], +) -> dict | None: + plan = build_rebalance_plan( + relation_status=relation_status, + customer_risk=customer_risk, + holdings=holdings, + target_allocation=target_allocation, + threshold=threshold, + candidates=candidates, + ) + if plan is None: + return None + structured_data = dict(plan) + structured_data["customer_risk"] = customer_risk + content = ["# 组合调仓建议草稿", "", "## 赎回清单"] + content.extend( + f"- {item['product_code']}:{item['amount']} 元" for item in plan["sell"] + ) + content.append("\n## 申购清单") + content.extend( + f"- {item['product_code']}:{item['amount']} 元" for item in plan["buy"] + ) + return { + "customer_id": customer_id, + "advisor_id": advisor_id, + "intent": AGENT_INTENT_REBALANCE, + "title": "组合调仓建议草稿", + "status": DRAFT_STATUS_DRAFT, + "structured_data": structured_data, + "deviation": max((abs(value) for value in plan["deviation"].values()), default=Decimal("0")), + "content": build_generated_content("\n".join(content)), + "disclaimer_ok": True, + } diff --git a/agent/advisor_agent/intent/fund_analysis.py b/agent/advisor_agent/intent/fund_analysis.py new file mode 100644 index 0000000..65ecc49 --- /dev/null +++ b/agent/advisor_agent/intent/fund_analysis.py @@ -0,0 +1,61 @@ +"""基金深度分析的数据整理层。""" +from __future__ import annotations + +from decimal import Decimal +from typing import Iterable + + +def _number(value): + if isinstance(value, Decimal): + return float(value) + return value + + +def _display_value(value): + if value is None: + return "暂无" + if isinstance(value, float): + return f"{value:g}" + return str(value) + + +def _build_analysis_text(fund: dict, metrics: list[dict]) -> str: + fund_name = fund.get("fund_name") or fund.get("fund_code") or "该基金" + if not metrics: + return f"{fund_name}暂无足够业绩数据,暂无法形成完整解读。" + + latest = metrics[-1] + period = _display_value(latest.get("period")) + return_rate = _display_value(latest.get("return_rate")) + max_drawdown = _display_value(latest.get("max_drawdown")) + sharpe = _display_value(latest.get("sharpe")) + return ( + f"{fund_name}在{period}的历史收益率为{return_rate}%," + f"最大回撤为{max_drawdown}%,夏普比率为{sharpe}。" + "以上仅基于历史业绩数据,不代表未来收益。" + ) + + +def build_fund_analysis(fund: dict, performance_rows: Iterable[dict]) -> dict: + metrics = [] + for row in performance_rows: + metrics.append( + { + key: _number(value) + for key, value in row.items() + } + ) + return { + "fund_code": fund.get("fund_code"), + "fund_name": fund.get("fund_name"), + "risk_level": fund.get("risk_level"), + "metrics": metrics, + "analysis_text": _build_analysis_text(fund, metrics), + "chart_data": { + "periods": [row.get("period") for row in metrics], + "return_rate": [row.get("return_rate") for row in metrics], + "annual_volatility": [row.get("annual_volatility") for row in metrics], + "max_drawdown": [row.get("max_drawdown") for row in metrics], + "sharpe": [row.get("sharpe") for row in metrics], + }, + } diff --git a/agent/advisor_agent/intent/generation_flow.py b/agent/advisor_agent/intent/generation_flow.py new file mode 100644 index 0000000..d8e82d6 --- /dev/null +++ b/agent/advisor_agent/intent/generation_flow.py @@ -0,0 +1,97 @@ +"""投顾意图结果到草稿与事件的编排。""" +from __future__ import annotations + +from decimal import Decimal +import json + +from agent.advisor_agent.intent.draft_generation import ( + build_rebalance_draft, + build_recommendation_draft, +) +from common.common_const import EVENT_ADVISOR_REBALANCE_DRAFT_CREATED +from service.advisor_agent.draft import create_draft +from agent.advisor_agent.llm import generate_text + + +async def generate_rebalance_draft( + *, + draft_repo, + publish, + customer_id: int, + advisor_id: int, + customer_risk: str, + relation_status: str, + holdings: list[dict], + target_allocation: dict[str, int | float | Decimal], + threshold: Decimal, + candidates: list[dict], + trace_id: str, +) -> dict | None: + draft_data = build_rebalance_draft( + customer_id=customer_id, + advisor_id=advisor_id, + customer_risk=customer_risk, + relation_status=relation_status, + holdings=holdings, + target_allocation=target_allocation, + threshold=threshold, + candidates=candidates, + ) + if draft_data is None: + return None + + draft = await create_draft(draft_repo, draft_data) + event_id = await publish( + event_name=EVENT_ADVISOR_REBALANCE_DRAFT_CREATED, + trace_id=trace_id, + trigger_user_id=advisor_id, + customer_id=customer_id, + payload={ + "draft_id": draft.draft_id, + "customer_id": customer_id, + "advisor_id": advisor_id, + "deviation": float(draft.deviation or 0), + "created_at": draft.create_time.isoformat() + if draft.create_time + else None, + }, + ) + return {"draft": draft, "event_id": event_id} + + +async def generate_recommendation_draft( + *, + draft_repo, + customer_id: int, + advisor_id: int, + customer_risk: str, + relation_status: str, + candidates: list[dict], + trace_id: str, + memories: list[dict] | None = None, + llm_client=None, + llm_timeout: float = 5.0, +): + draft_data = build_recommendation_draft( + customer_id=customer_id, + advisor_id=advisor_id, + customer_risk=customer_risk, + candidates=candidates, + relation_status=relation_status, + memories=memories, + ) + if llm_client is not None: + draft_data["content"] = await generate_text( + llm_client, + system_prompt="你是基金投顾助手,只生成内部投顾草稿说明,不下单、不承诺收益。", + user_prompt=( + "请根据以下候选基金和客户记忆生成简洁推荐说明:" + + json.dumps( + {"candidates": candidates, "memories": memories or []}, + ensure_ascii=False, + ) + ), + fallback=lambda: draft_data["content"], + timeout=llm_timeout, + ) + return await create_draft(draft_repo, draft_data) diff --git a/agent/advisor_agent/intent/rebalance.py b/agent/advisor_agent/intent/rebalance.py new file mode 100644 index 0000000..64d70f4 --- /dev/null +++ b/agent/advisor_agent/intent/rebalance.py @@ -0,0 +1,107 @@ +"""组合偏离度与调仓建议计算。""" +from __future__ import annotations + +from collections import defaultdict +from decimal import Decimal, ROUND_HALF_UP +from typing import Iterable + +from common.common_const import ( + CUSTOMER_REL_STATUS_SIGNED, + ERR_CODE_NOT_SIGNED_REBALANCE, +) +from common.suitability import check_suitability +from utils.exceptions import ApiError + + +_MONEY = Decimal("0.01") +_PERCENT = Decimal("100") + + +def _money(value: Decimal) -> Decimal: + return value.quantize(_MONEY, rounding=ROUND_HALF_UP) + + +def build_rebalance_plan( + *, + relation_status: str, + customer_risk: str, + holdings: Iterable[dict], + target_allocation: dict[str, int | float | Decimal], + threshold: Decimal, + candidates: Iterable[dict], +) -> dict | None: + if relation_status != CUSTOMER_REL_STATUS_SIGNED: + raise ApiError(ERR_CODE_NOT_SIGNED_REBALANCE, "客户尚未签约,禁止生成调仓草稿") + + values: dict[str, Decimal] = defaultdict(Decimal) + holdings_by_class: dict[str, list[dict]] = defaultdict(list) + for holding in holdings: + asset_class = str(holding.get("asset_class", "")) + value = Decimal(str(holding.get("market_value", 0) or 0)) + values[asset_class] += value + holdings_by_class[asset_class].append(holding) + + total = sum(values.values(), Decimal("0")) + if total <= 0: + return None + + target = { + asset_class: Decimal(str(weight)) for asset_class, weight in target_allocation.items() + } + deviation: dict[str, Decimal] = {} + for asset_class in target: + actual = values.get(asset_class, Decimal("0")) / total * _PERCENT + deviation[asset_class] = (actual - target[asset_class]).quantize( + Decimal("0.01"), rounding=ROUND_HALF_UP + ) + + if not any(abs(value) > threshold for value in deviation.values()): + return None + + sell: list[dict] = [] + buy: list[dict] = [] + for asset_class, drift in deviation.items(): + if drift > threshold: + target_value = total * target[asset_class] / _PERCENT + excess = _money(values.get(asset_class, Decimal("0")) - target_value) + remaining = excess + for holding in holdings_by_class.get(asset_class, []): + amount = min( + remaining, + _money(Decimal(str(holding.get("market_value", 0) or 0))), + ) + if amount > 0: + sell.append( + { + "product_code": holding.get("product_code"), + "asset_class": asset_class, + "amount": amount, + } + ) + remaining -= amount + if remaining <= 0: + break + elif drift < -threshold: + target_value = total * target[asset_class] / _PERCENT + amount = _money(target_value - values.get(asset_class, Decimal("0"))) + for candidate in candidates: + if candidate.get("asset_class") != asset_class: + continue + if not check_suitability( + customer_risk, candidate.get("risk_level", "") + ).ok: + continue + buy.append( + { + "product_code": candidate.get("product_code"), + "asset_class": asset_class, + "amount": amount, + } + ) + break + + return { + "deviation": deviation, + "sell": sell, + "buy": buy, + } diff --git a/agent/advisor_agent/intent/recommend.py b/agent/advisor_agent/intent/recommend.py new file mode 100644 index 0000000..ce7928c --- /dev/null +++ b/agent/advisor_agent/intent/recommend.py @@ -0,0 +1,45 @@ +"""基金推荐候选过滤与排序。""" +from __future__ import annotations + +from collections import defaultdict +from typing import Iterable + +from common.common_const import MEMORY_INFO_TYPE_OPINION +from common.suitability import check_suitability + + +def _opinion_bonuses(memories: Iterable[dict]) -> dict[str, float]: + bonuses: dict[str, float] = defaultdict(float) + for memory in memories: + if memory.get("info_type") != MEMORY_INFO_TYPE_OPINION: + continue + product_code = memory.get("product_code") + if product_code: + bonuses[str(product_code)] += float(memory.get("score", 0.0) or 0.0) + return bonuses + + +def recommend_candidates( + customer_risk: str, + candidates: Iterable[dict], + *, + memories: Iterable[dict] = (), +) -> list[dict]: + """硬过滤不适当产品,再按业绩和主观观点排序。""" + bonuses = _opinion_bonuses(memories) + result = [] + for candidate in candidates: + suitability = check_suitability(customer_risk, candidate.get("risk_level", "")) + if not suitability.ok: + continue + item = dict(candidate) + base_score = float(item.get("performance_score", 0.0) or 0.0) + item["recommendation_score"] = base_score + bonuses.get( + str(item.get("product_code")), 0.0 + ) + result.append(item) + return sorted( + result, + key=lambda item: item["recommendation_score"], + reverse=True, + ) diff --git a/agent/advisor_agent/intent/talk_script.py b/agent/advisor_agent/intent/talk_script.py new file mode 100644 index 0000000..7e3fbfb --- /dev/null +++ b/agent/advisor_agent/intent/talk_script.py @@ -0,0 +1,24 @@ +"""投顾沟通话术安全兜底模板。""" +from __future__ import annotations + +from common.common_const import ( + TALK_SCENE_CUSTOMER_COMPLAINT, + TALK_SCENE_MARKET_FLUCTUATION, + TALK_SCENE_PORTFOLIO_DIVERGENCE, + TALK_SCENE_RISK_BLOCK_ORDER, +) + + +_TEMPLATES = { + TALK_SCENE_RISK_BLOCK_ORDER: "{name},这笔交易正在进行风险审核。请先查看审核结果,后续是否交易由您结合自身情况自行决定。", + TALK_SCENE_MARKET_FLUCTUATION: "{name},近期市场波动可能放大短期净值变化。建议先关注组合风险和自身资金安排,再审慎决定是否调整。", + TALK_SCENE_PORTFOLIO_DIVERGENCE: "{name},当前组合与既定配置基准存在偏离。我可以向您说明偏离来源和可选调整方向,具体决定请结合您的风险承受能力。", + TALK_SCENE_CUSTOMER_COMPLAINT: "{name},很抱歉给您带来不好的体验。我会记录您的问题并协助核实处理进展,具体结果以核查信息为准。", +} + + +def build_talk_script(scene_type: str, *, customer_name: str = "客户") -> dict: + template = _TEMPLATES.get(scene_type) + if template is None: + raise ValueError("不支持的话术场景") + return {"scene_type": scene_type, "content": template.format(name=customer_name)} diff --git a/agent/advisor_agent/llm.py b/agent/advisor_agent/llm.py new file mode 100644 index 0000000..7e2aaa9 --- /dev/null +++ b/agent/advisor_agent/llm.py @@ -0,0 +1,40 @@ +"""投顾 Agent 对共享 LLM 客户端的安全调用封装。""" +from __future__ import annotations + +import inspect + +from agent.advisor_agent.fallback import call_with_fallback + + +async def generate_text( + llm_client, + *, + system_prompt: str, + user_prompt: str, + fallback, + timeout: float = 5.0, +) -> str: + """调用共享 LLM,超时或异常时返回本地安全兜底文本。""" + + async def primary(): + return await llm_client.chat( + [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": user_prompt}, + ] + ) + + async def secondary(): + value = fallback() + return await value if inspect.isawaitable(value) else value + + result = await call_with_fallback( + primary, + secondary, + timeout=timeout, + degraded_code=50001, + ) + return str(result.value) + + +__all__ = ["generate_text"] diff --git a/agent/advisor_agent/memory.py b/agent/advisor_agent/memory.py new file mode 100644 index 0000000..9740749 --- /dev/null +++ b/agent/advisor_agent/memory.py @@ -0,0 +1,58 @@ +"""投顾 Agent 对现有统一记忆系统的适配边界。""" +from __future__ import annotations + +from typing import Protocol + + +class MemoryProvider(Protocol): + async def recall(self, *, customer_id: int, query: str) -> list[dict]: + """召回客户相关记忆,返回结构化记忆单元。""" + + +class EmptyMemoryProvider: + """记忆系统未接入时的安全默认实现。""" + + async def recall(self, *, customer_id: int, query: str) -> list[dict]: + return [] + + +class MemoryServiceProvider: + """将现有 MemoryService 的统一上下文转换为投顾侧记忆列表。""" + + def __init__(self, memory_service, *, session_prefix: str = "advisor-agent"): + self.memory_service = memory_service + self.session_prefix = session_prefix + self.last_warnings: list[str] = [] + + async def recall(self, *, customer_id: int, query: str) -> list[dict]: + self.last_warnings = [] + try: + context = await self.memory_service.recall( + customer_id=customer_id, + session_id=f"{self.session_prefix}:{customer_id}", + query=query, + ) + except Exception as exc: + self.last_warnings = [f"advisor_memory_recall_failed:{type(exc).__name__}"] + return [] + + self.last_warnings.extend(getattr(context, "warnings", []) or []) + memories = getattr(context, "long_term_memories", context) + return [self._to_dict(memory) for memory in memories] + + @staticmethod + def _to_dict(memory) -> dict: + if hasattr(memory, "model_dump"): + data = memory.model_dump(mode="json") + else: + data = dict(memory) + return { + "customer_id": data.get("customer_id"), + "tag": data.get("tag"), + "content": data.get("content"), + "info_type": data.get("info_type", "FACT"), + "memory_type": data.get("memory_type"), + } + + +__all__ = ["EmptyMemoryProvider", "MemoryProvider", "MemoryServiceProvider"] diff --git a/agent/advisor_agent/protocol.py b/agent/advisor_agent/protocol.py new file mode 100644 index 0000000..70e4710 --- /dev/null +++ b/agent/advisor_agent/protocol.py @@ -0,0 +1,30 @@ +"""投顾 Agent 与工作台之间的响应协议。""" +from __future__ import annotations + +from typing import Any + +from common.common_const import ERR_CODE_OK + + +def agent_success(data: Any = None, *, trace_id: str | None = None) -> dict: + return { + "code": ERR_CODE_OK, + "message": "success", + "data": data, + "trace_id": trace_id, + } + + +def agent_failure( + code: int, + message: str, + data: Any = None, + *, + trace_id: str | None = None, +) -> dict: + return { + "code": code, + "message": message, + "data": data, + "trace_id": trace_id, + } diff --git a/agent/advisor_agent/runtime.py b/agent/advisor_agent/runtime.py new file mode 100644 index 0000000..ee5f1a4 --- /dev/null +++ b/agent/advisor_agent/runtime.py @@ -0,0 +1,36 @@ +"""投顾 Agent 运行时依赖组装。""" +from __future__ import annotations + +from dataclasses import dataclass + +from agent.advisor_agent.memory import ( + EmptyMemoryProvider, + MemoryProvider, + MemoryServiceProvider, +) + + +@dataclass +class AdvisorAgentRuntime: + memory_provider: MemoryProvider + llm_client: object | None = None + graph_tool: object | None = None + redis: object | None = None + + +def build_default_runtime( + *, + memory_provider: MemoryProvider | None = None, + llm_client: object | None = None, + graph_tool: object | None = None, + redis: object | None = None, + memory_service: object | None = None, +) -> AdvisorAgentRuntime: + if memory_provider is None and memory_service is not None: + memory_provider = MemoryServiceProvider(memory_service) + return AdvisorAgentRuntime( + memory_provider=memory_provider or EmptyMemoryProvider(), + llm_client=llm_client, + graph_tool=graph_tool, + redis=redis, + ) diff --git a/agent/data_query/__init__.py b/agent/data_query/__init__.py new file mode 100644 index 0000000..91f087d --- /dev/null +++ b/agent/data_query/__init__.py @@ -0,0 +1,5 @@ +"""数据查询 Agent 适配层。""" + +from .agent import DataQueryAgent + +__all__ = ["DataQueryAgent"] diff --git a/agent/data_query/agent.py b/agent/data_query/agent.py new file mode 100644 index 0000000..fb5a9d7 --- /dev/null +++ b/agent/data_query/agent.py @@ -0,0 +1,27 @@ +"""面向其他 Agent 的数据查询适配器。""" +from __future__ import annotations + +from uuid import uuid4 + +from nl2sql.contracts import DataQueryRequest +from nl2sql.streaming import stream_query_events +from service.nl2sql.query_service import execute_query + + +class DataQueryAgent: + """只依赖公共查询服务,不直接调用其他 Agent 或基础设施。""" + + async def query(self, request: DataQueryRequest, **dependencies): + """转发查询请求,缺省生成内部查询 ID。""" + dependencies.setdefault("query_id", uuid4().hex) + return await execute_query(request, **dependencies) + + async def stream(self, request: DataQueryRequest, **dependencies): + """以结构化事件流转发查询,不暴露 SQL、结果行和异常消息。""" + dependencies.setdefault("query_id", uuid4().hex) + async for event in stream_query_events( + request, + query_runner=execute_query, + **dependencies, + ): + yield event diff --git a/api/advisor/audit.py b/api/advisor/audit.py index 113628f..611378d 100644 --- a/api/advisor/audit.py +++ b/api/advisor/audit.py @@ -17,6 +17,7 @@ router = APIRouter() @router.get("/audit/ledger", summary="本人审计台账筛选") async def ledger( action: str | None = Query(None, max_length=64), + customer_id: int | None = Query(None, gt=0), keyword: str | None = Query(None, max_length=128), start: datetime | None = Query(None), end: datetime | None = Query(None), @@ -27,7 +28,7 @@ async def ledger( ): return success( await audit_service.list_ledger( - db, user, action=action, keyword=keyword, start=start, end=end, + db, user, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end, page=page, page_size=page_size, ) ) @@ -36,6 +37,7 @@ async def ledger( @router.get("/audit/export", summary="导出本人审计台账(CSV)") async def export( action: str | None = Query(None, max_length=64), + customer_id: int | None = Query(None, gt=0), keyword: str | None = Query(None, max_length=128), start: datetime | None = Query(None), end: datetime | None = Query(None), @@ -43,7 +45,7 @@ async def export( db: AsyncSession = Depends(get_db), ): csv_text = await audit_service.export_ledger( - db, user, action=action, keyword=keyword, start=start, end=end + db, user, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end ) return Response( content=csv_text, diff --git a/api/advisor/visits.py b/api/advisor/visits.py index e52b156..a92e2e4 100644 --- a/api/advisor/visits.py +++ b/api/advisor/visits.py @@ -5,7 +5,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from api.deps import require_advisor from config.deps import get_db from model.sys_user import SysUser -from schemas.advisor import VisitCreateReq +from schemas.advisor import VisitCreateReq, VisitUpdateReq from service.advisor import visits as visits_service from utils.response import success @@ -36,6 +36,25 @@ async def create_visit( return success(await visits_service.create_visit(db, user, req)) +@router.get("/visits/{visit_id}", summary="回访记录详情") +async def get_visit( + visit_id: int, + user: SysUser = Depends(require_advisor), + db: AsyncSession = Depends(get_db), +): + return success(await visits_service.get_visit(db, user, visit_id)) + + +@router.put("/visits/{visit_id}", summary="编辑回访留痕") +async def update_visit( + visit_id: int, + req: VisitUpdateReq, + user: SysUser = Depends(require_advisor), + db: AsyncSession = Depends(get_db), +): + return success(await visits_service.update_visit(db, user, visit_id, req)) + + @router.get("/talk-templates", summary="合规话术库(内置,投顾参考)") async def talk_templates(user: SysUser = Depends(require_advisor)): return success(visits_service.list_talk_templates()) diff --git a/api/deps.py b/api/deps.py index c0fe658..dc61133 100644 --- a/api/deps.py +++ b/api/deps.py @@ -7,6 +7,7 @@ from model.sys_user import SysUser from repositories.sys_user import SysUserRepo from service.auth import decode_token from utils.exceptions import AuthError, ForbiddenError +from service.advisor_agent.audit import write_advisor_audit, audit_action_for_path KNOWLEDGE_OPERATOR_ROLES = {"ADMIN", "KNOWLEDGE_ADMIN", "KNOWLEDGE_OPERATOR", "运营"} @@ -53,3 +54,28 @@ async def require_risk_officer(user: SysUser = Depends(get_current_user)) -> Sys if user.user_type != "EMPLOYEE" or user.employee_role != "风控专员": raise ForbiddenError("仅风控专员可执行风控处置") return user + +async def require_advisor(user: SysUser = Depends(get_current_user)) -> SysUser: + """Allow only advisor employees to access the advisor workbench.""" + if user.user_type != "EMPLOYEE" or user.employee_role not in {"投顾", "ADMIN"}: + raise ForbiddenError("仅投顾人员可以访问投顾工作台") + return user + + +async def audited_advisor(request: Request, db: AsyncSession = Depends(get_db)): + """Authenticate an advisor Agent request and leave a failure audit trail.""" + try: + user = await get_current_user(request, db) + if user.user_type != "EMPLOYEE" or user.employee_role not in {"投顾", "ADMIN"}: + raise ForbiddenError("仅投顾人员可以访问 advisor_agent") + yield user + except Exception: + await write_advisor_audit( + db, + user=None, + action=audit_action_for_path(request.url.path), + target=None, + trace_id=request.headers.get("X-Trace-Id", ""), + status="失败", + ) + raise \ No newline at end of file diff --git a/api/router.py b/api/router.py index 90d5018..7d7772e 100644 --- a/api/router.py +++ b/api/router.py @@ -6,6 +6,8 @@ from fastapi import APIRouter from api.chat import client_agent, customer_agent, knowledge from api.routers import product, questionnaire from api.routers import account, auth, holdings, purchase, redeem, risk, work_order +from api.routers import advisor_agent, health, nl2sql, nl2sql_admin +from api.advisor import audit, customers, dashboard, diagnosis, drafts, report, todos, visits api_router = APIRouter() api_router.include_router(auth.router, prefix="/api", tags=["认证"]) @@ -20,3 +22,18 @@ api_router.include_router(client_agent.router, prefix="/api/agent/client", tags= api_router.include_router(knowledge.router, prefix="/api/knowledge", tags=["知识库"]) api_router.include_router(product.router, prefix="/api", tags=["产品"]) api_router.include_router(questionnaire.router, prefix="/api", tags=["问卷"]) +api_router.include_router(advisor_agent.router, prefix="/api", tags=["投顾Agent"]) +for workbench_router in ( + dashboard.router, + customers.router, + diagnosis.router, + drafts.router, + todos.router, + visits.router, + audit.router, + report.router, +): + api_router.include_router(workbench_router, prefix="/api/advisor") +api_router.include_router(health.router, prefix="/api", tags=["健康检查"]) +api_router.include_router(nl2sql.router, prefix="/api", tags=["NL2SQL"]) +api_router.include_router(nl2sql_admin.router, prefix="/api", tags=["NL2SQL管理"]) diff --git a/api/routers/advisor_agent.py b/api/routers/advisor_agent.py new file mode 100644 index 0000000..aae2121 --- /dev/null +++ b/api/routers/advisor_agent.py @@ -0,0 +1,407 @@ +"""投顾 Agent HTTP 契约入口与本地业务编排。""" +from __future__ import annotations + +import json +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 +from agent.advisor_agent.intent.fund_analysis import build_fund_analysis +from agent.advisor_agent.intent.talk_script import build_talk_script +from agent.advisor_agent.llm import generate_text +from agent.advisor_agent.intent.generation_flow import ( + generate_rebalance_draft, + generate_recommendation_draft, +) +from agent.advisor_agent.protocol import agent_failure, agent_success +from common.common_const import ( + AGENT_INTENT_RECOMMEND, + CUSTOMER_REL_STATUS_SIGNED, + ERR_CODE_DRAFT_NOT_FOUND, + ERR_CODE_FORBIDDEN_CUSTOMER, + ERR_CODE_LLM_ERROR, + ERR_CODE_NOT_SIGNED_REBALANCE, + DRAFT_STATUS_DISCARDED, + DRAFT_STATUS_DRAFT, + SSE_EVENT_TYPE_DONE, + SSE_EVENT_TYPE_ERROR, + SSE_EVENT_TYPE_META, + SSE_EVENT_TYPE_TEXT, +) +from api.deps import audited_advisor +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 service.advisor_agent.context import ( + load_fund_analysis_context, + load_customer_risk, + load_rebalance_context, + load_recommendation_context, +) +from service.event_publisher import publish_event +from schemas.advisor_agent import ( + AdvisorDraftOperateReq, + AdvisorDraftSaveReq, + AdvisorChatReq, + AdvisorFundAnalysisReq, + AdvisorRebalanceRunReq, + AdvisorTalkScriptReq, +) +from service.advisor_agent.draft import ( + detail_draft, + discard_draft, + ensure_draft_owner, + get_draft as get_draft_service, + list_drafts as list_drafts_service, + save_draft as save_draft_service, +) +from utils.exceptions import ApiError +from utils.request_id import get_request_id, new_request_id +from utils.logger import get_logger + + +router = APIRouter(prefix="/advisor-agent") +logger = get_logger("advisor_agent.router") + +_NOT_READY_CODE = ERR_CODE_LLM_ERROR +_NOT_READY_MESSAGE = "投顾 Agent 核心能力尚未初始化" +_DRAFT_NOT_FOUND_CODE = ERR_CODE_DRAFT_NOT_FOUND +_DRAFT_NOT_FOUND_MESSAGE = "草稿不存在或者已废弃" + + +def _trace_id(request: Request) -> str: + return request.headers.get("X-Trace-Id") or get_request_id() or new_request_id() + + +def _not_ready(request: Request): + return agent_failure(_NOT_READY_CODE, _NOT_READY_MESSAGE, trace_id=_trace_id(request)) + + +def _advisor_runtime(request: Request): + app = request.scope.get("app") + return getattr(getattr(app, "state", None), "advisor_agent_runtime", None) + + +async def _recall_advisor_memories( + request: Request, *, customer_id: int, query: str +) -> list[dict]: + """从应用共享 runtime 读取记忆;未装配或故障时安全降级为空。""" + app = request.scope.get("app") + state = getattr(app, "state", None) + runtime = getattr(state, "advisor_agent_runtime", None) + provider = getattr(runtime, "memory_provider", None) + if provider is None: + return [] + try: + return await provider.recall(customer_id=customer_id, query=query) + except Exception: + logger.warning( + "advisor memory recall failed: customer_id=%s", customer_id, exc_info=True + ) + return [] + + +async def _ensure_draft_access(db, draft, advisor_id: int) -> None: + ensure_draft_owner(draft, advisor_id=advisor_id) + await ensure_customer_access( + db, + advisor_id=advisor_id, + customer_id=draft.customer_id, + ) + + +async def _run_rebalance_background( + *, advisor_id: int, customer_id: int, trace_id: str +) -> None: + """后台任务使用独立会话,避免请求返回后复用已关闭的请求会话。""" + try: + async with mysql.get_session_factory()() as task_db: + relation = await ensure_customer_access( + task_db, advisor_id=advisor_id, customer_id=customer_id + ) + if relation.status != CUSTOMER_REL_STATUS_SIGNED: + return + context = await load_rebalance_context(task_db, customer_id=customer_id) + if context is None: + return + await generate_rebalance_draft( + draft_repo=AdvisorDraftRepo(task_db), + publish=lambda **kwargs: publish_event(redis_db.client(), task_db, **kwargs), + advisor_id=advisor_id, + trace_id=trace_id, + **context, + ) + except Exception: + logger.warning("advisor rebalance background task failed", exc_info=True) + + +@router.post("/chat/stream") +async def chat_stream( + request: Request, + body: dict, + 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, + ) + else: + customer_id = chat_request.customer_id + relation = await ensure_customer_access( + db, advisor_id=user.id, customer_id=int(customer_id) + ) + if chat_request.intent == AGENT_INTENT_RECOMMEND: + memories = await _recall_advisor_memories( + request, + customer_id=int(customer_id), + query=chat_request.query or "", + ) + runtime = _advisor_runtime(request) + context = await load_recommendation_context( + db, customer_id=int(customer_id) + ) + if context is not None: + draft = await generate_recommendation_draft( + draft_repo=AdvisorDraftRepo(db), + advisor_id=user.id, + relation_status=relation.status, + trace_id=trace_id, + memories=memories, + llm_client=getattr(runtime, "llm_client", None), + **context, + ) + + async def events(): + for event in ( + { + "type": SSE_EVENT_TYPE_META, + "draft_id": draft.draft_id, + "intent": draft.intent, + "status": draft.status, + }, + {"type": SSE_EVENT_TYPE_TEXT, "content": draft.content}, + {"type": SSE_EVENT_TYPE_DONE, "draft_id": draft.draft_id}, + ): + 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}, + ) + payload = agent_failure(_NOT_READY_CODE, _NOT_READY_MESSAGE, trace_id=trace_id) + + async def events(): + yield f"data: {json.dumps({'type': SSE_EVENT_TYPE_ERROR, **payload}, ensure_ascii=False)}\n\n" + + return StreamingResponse( + events(), + media_type="text/event-stream", + headers={"X-Trace-Id": trace_id}, + ) + + +@router.get("/draft/list") +async def list_drafts( + request: Request, + advisor_id: int | None = Query(default=None), + 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), + user: SysUser = Depends(audited_advisor), + db: AsyncSession = Depends(get_db), +): + if advisor_id is not None and advisor_id != user.id: + raise ApiError(ERR_CODE_FORBIDDEN_CUSTOMER, "无权操作该客户数据") + if customer_id is not None: + await ensure_customer_access(db, advisor_id=user.id, customer_id=customer_id) + result = await list_drafts_service( + AdvisorDraftRepo(db), + advisor_id=user.id, + customer_id=customer_id, + status=status, + page=page, + page_size=page_size, + ) + return agent_success(result, trace_id=_trace_id(request)) + + +@router.get("/draft/{draft_id}") +async def get_draft( + draft_id: str, + request: Request, + user: SysUser = Depends(audited_advisor), + db: AsyncSession = Depends(get_db), +): + draft = await get_draft_service(AdvisorDraftRepo(db), draft_id) + await _ensure_draft_access(db, draft, user.id) + return agent_success(detail_draft(draft), trace_id=_trace_id(request)) + + +@router.put("/draft/{draft_id}/save") +async def save_draft( + draft_id: str, + request: Request, + body: AdvisorDraftSaveReq, + user: SysUser = Depends(audited_advisor), + db: AsyncSession = Depends(get_db), +): + repo = AdvisorDraftRepo(db) + draft = await get_draft_service(repo, draft_id) + await _ensure_draft_access(db, draft, user.id) + customer_risk = await load_customer_risk(db, customer_id=draft.customer_id) + result = await save_draft_service( + repo, + draft_id, + title=body.title, + content=body.content, + structured_data=body.structured_data, + customer_risk=customer_risk, + ) + return agent_success(result, trace_id=_trace_id(request)) + + +@router.post("/draft/{draft_id}/operate") +async def operate_draft( + draft_id: str, + request: Request, + body: AdvisorDraftOperateReq, + user: SysUser = Depends(audited_advisor), + db: AsyncSession = Depends(get_db), +): + if body.operation != "discard": + raise ApiError(_DRAFT_NOT_FOUND_CODE, "不支持的草稿操作") + repo = AdvisorDraftRepo(db) + draft = await get_draft_service(repo, draft_id) + await _ensure_draft_access(db, draft, user.id) + result = await discard_draft(repo, draft_id) + return agent_success(detail_draft(result), trace_id=_trace_id(request)) + + +@router.post("/rebalance/run") +async def run_rebalance( + request: Request, + body: AdvisorRebalanceRunReq, + user: SysUser = Depends(audited_advisor), + db: AsyncSession = Depends(get_db), + background_tasks: BackgroundTasks = None, +): + customer_id = body.customer_id + relation = await ensure_customer_access( + db, advisor_id=user.id, customer_id=int(customer_id) + ) + if relation.status != CUSTOMER_REL_STATUS_SIGNED: + return agent_failure( + ERR_CODE_NOT_SIGNED_REBALANCE, + "客户尚未签约,禁止生成调仓草稿", + trace_id=_trace_id(request), + ) + if background_tasks is None: + background_tasks = BackgroundTasks() + background_tasks.add_task( + _run_rebalance_background, + advisor_id=user.id, + customer_id=customer_id, + trace_id=_trace_id(request), + ) + return agent_success( + {"accepted": True, "status": "queued"}, + trace_id=_trace_id(request), + ) + + +@router.post("/fund-analysis") +async def fund_analysis( + request: Request, + body: AdvisorFundAnalysisReq, + user: SysUser = Depends(audited_advisor), + db: AsyncSession = Depends(get_db), +): + memories: list[dict] = [] + if body.customer_id is not None: + await ensure_customer_access( + db, advisor_id=user.id, customer_id=body.customer_id + ) + memories = await _recall_advisor_memories( + request, + customer_id=body.customer_id, + query=f"基金分析 {','.join(body.fund_codes)}", + ) + fund = body.fund + performance = body.performance + if fund is None or performance is None: + contexts = await load_fund_analysis_context( + db, + fund_codes=[str(code) for code in body.fund_codes], + ) + if not contexts: + return _not_ready(request) + if len(contexts) == 1: + fund = contexts[0]["fund"] + performance = contexts[0]["performance"] + else: + return agent_success( + {"items": [build_fund_analysis(item["fund"], item["performance"]) for item in contexts]}, + trace_id=_trace_id(request), + ) + result = build_fund_analysis(fund, performance) + runtime = _advisor_runtime(request) + if getattr(runtime, "llm_client", None) is not None: + result["analysis_text"] = await generate_text( + runtime.llm_client, + system_prompt="你是基金投顾助手,只基于给定历史数据生成谨慎的内部分析,不承诺收益。", + user_prompt=json.dumps( + {"analysis": result, "customer_memories": memories}, + ensure_ascii=False, + ), + fallback=lambda: result["analysis_text"], + timeout=5.0, + ) + return agent_success(result, trace_id=_trace_id(request)) + + +@router.post("/generate-talk-script") +async def generate_talk_script( + request: Request, + body: AdvisorTalkScriptReq, + user: SysUser = Depends(audited_advisor), + db: AsyncSession = Depends(get_db), +): + await ensure_customer_access(db, advisor_id=user.id, customer_id=body.customer_id) + memories = await _recall_advisor_memories( + request, + customer_id=body.customer_id, + query=f"沟通话术 {body.scene_type}", + ) + try: + result = build_talk_script( + body.scene_type, + customer_name=body.customer_name, + ) + except ValueError as exc: + raise ApiError(ERR_CODE_LLM_ERROR, str(exc)) from exc + runtime = _advisor_runtime(request) + if getattr(runtime, "llm_client", None) is not None: + result["content"] = await generate_text( + runtime.llm_client, + system_prompt="你是合规的基金投顾助手,只生成谨慎沟通话术,不承诺收益、不代客交易。", + user_prompt=json.dumps( + {"script": result, "customer_memories": memories}, + ensure_ascii=False, + ), + fallback=lambda: result["content"], + timeout=5.0, + ) + return agent_success(result, trace_id=_trace_id(request)) diff --git a/api/routers/health.py b/api/routers/health.py new file mode 100644 index 0000000..0ec7f6a --- /dev/null +++ b/api/routers/health.py @@ -0,0 +1,25 @@ +"""基础设施健康检查路由。""" +from __future__ import annotations + +from fastapi import APIRouter +from fastapi.responses import JSONResponse + +from config import database +from utils.response import fail, success + + +router = APIRouter() + + +@router.get("/health/ready") +async def ready(): + """按需探测四库,并返回适合负载均衡器使用的 HTTP 就绪状态。""" + databases = await database.check_ready_detail() + is_ready = all(item.get("status") == "ok" for item in databases.values()) + payload = {"ready": is_ready, "databases": databases} + if is_ready: + return JSONResponse(status_code=200, content=success(payload).model_dump()) + return JSONResponse( + status_code=503, + content=fail(503, "基础设施尚未就绪", payload).model_dump(), + ) diff --git a/api/routers/nl2sql.py b/api/routers/nl2sql.py new file mode 100644 index 0000000..889f780 --- /dev/null +++ b/api/routers/nl2sql.py @@ -0,0 +1,614 @@ +"""NL2SQL 查询 HTTP 接口。""" +from __future__ import annotations + +import csv +import json +from dataclasses import asdict, replace +from io import StringIO +from uuid import uuid4 + +from datetime import datetime + +from fastapi import APIRouter, Depends, Query, Request +from fastapi.responses import Response, StreamingResponse +from sqlalchemy import text +from sqlalchemy.ext.asyncio import AsyncSession + +from api.deps import get_current_user +from config import database +from config.deps import get_db, get_milvus, get_redis +from config.settings import settings +from model.sys_user import SysUser +from nl2sql.contracts import DataQueryRequest +from nl2sql.cache import ( + DEFAULT_CACHE_TTL, + build_cache_key, + cache_get, + cache_set, + invalidate_tables, + result_from_cache, + result_to_cache, +) +from nl2sql.audit import write_nl2sql_audit_safely +from nl2sql.diagnostics import diagnostic_registry +from nl2sql.executor import execute_readonly_sql +from nl2sql.explain import build_explain_sql +from nl2sql.history import archive_query_safely +from nl2sql.health import check_nl2sql_health +from nl2sql.query_experience import ( + build_clarification, + render_csv, +) +from nl2sql.retrieval import retrieve_metadata +from nl2sql.result import build_chart_config, summarize_result +from nl2sql.runtime import kill_mysql_query, query_runtime_registry +from nl2sql.runtime_config import runtime_config +from nl2sql.session_context import SessionContextStore, build_conversation_context +from nl2sql.limits import QueryLimiter +from nl2sql.metrics import query_metrics +from nl2sql.schema import load_authoritative_schema +from nl2sql.supervisor import SessionBusyError, SessionLock +from nl2sql.sql_security import SqlSecurityError, validate_select_sql +from repositories.nl2sql_permission import Nl2SqlPermissionRepo +from schemas.nl2sql import DataCacheInvalidateReq, DataExplainReq, DataKillReq, DataQueryReq +from service.nl2sql.query_service import QueryServiceError, query as build_query +from service.nl2sql.permission_service import load_query_permission +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 + + +router = APIRouter() + + +def ensure_query_employee(user: SysUser) -> None: + """NL2SQL 仅允许已登录员工账号使用。""" + if user.user_type != "EMPLOYEE": + raise ForbiddenError("仅登录员工可以使用数据查询") + + +def ensure_query_admin(user: SysUser) -> None: + """只允许系统管理员管理运行中查询。""" + if user.user_type == "ADMIN": + return + if user.user_type == "EMPLOYEE" and user.employee_role in {"ADMIN", "系统管理员"}: + return + raise ForbiddenError("仅管理员可以管理运行中查询") + + +def build_history_payload(*, query_id, request_data, user, trace_id, status, result, query_result=None): + """构造查询历史字段,明确排除查询结果行。""" + return { + "query_id": query_id, + "user_id": user.id, + "session_id": request_data.session_id, + "caller_agent": request_data.caller_agent, + "question": request_data.question, + "generated_sql": getattr(result, "sql", None), + "access_tables": getattr(result, "access_tables", set()), + "status": status, + "row_count": getattr(query_result or result, "row_count", 0), + "truncated": getattr(query_result or result, "truncated", False), + "elapsed_ms": getattr(query_result or result, "elapsed_ms", None), + "trace_id": trace_id, + } + + +def history_payload(history) -> dict: + """将查询历史模型转换为不含结果行和连接信息的响应。""" + return { + "query_id": history.query_id, + "user_id": history.user_id, + "session_id": history.session_id, + "caller_agent": history.caller_agent, + "question": history.question, + "generated_sql": history.generated_sql, + "access_tables": history.access_tables or [], + "status": history.status, + "error_code": history.error_code, + "error_message": history.error_message, + "row_count": history.row_count, + "truncated": history.truncated, + "elapsed_ms": history.elapsed_ms, + "trace_id": history.trace_id, + "create_time": history.create_time, + } + + +async def enrich_query_result(question: str, result, *, llm_client): + """为已脱敏结果补充摘要和安全的基础图表配置。""" + return replace( + result, + summary=await summarize_result( + question, + result.columns, + result.rows, + llm_client=llm_client, + ), + chart=build_chart_config(result.columns, result.rows), + ) + + +async def _load_permission(db: AsyncSession, user: SysUser) -> dict: + return await load_query_permission(db, user.id) + + +def format_sse_event(event: str, payload: dict) -> str: + """将结构化事件编码为标准 SSE 文本。""" + return f"event: {event}\ndata: {json.dumps(payload, ensure_ascii=False)}\n\n" + + +@router.post("/nl2sql/query/stream") +async def stream_query_data( + body: DataQueryReq, + request: Request, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), + milvus=Depends(get_milvus), + redis=Depends(get_redis), +): + """以 SSE 返回查询进度和脱敏统计,不改变原查询链路。""" + ensure_query_employee(user) + if body.output_format != "json": + raise ParamError("流式查询只支持 JSON 输出") + stream_id = uuid4().hex + trace_id = request.headers.get("X-Trace-Id") or get_request_id() or new_request_id() + + async def events(): + yield format_sse_event( + "started", + {"query_id": stream_id, "trace_id": trace_id}, + ) + try: + response = await query_data(body, request, user, db, milvus, redis) + payload = response.model_dump() if hasattr(response, "model_dump") else {} + data = payload.get("data") or {} + yield format_sse_event( + "completed", + { + "query_id": data.get("query_id", stream_id), + "trace_id": data.get("trace_id", trace_id), + "row_count": data.get("row_count", 0), + "truncated": bool(data.get("truncated", False)), + }, + ) + except Exception as exc: # noqa: BLE001 流式错误只返回异常类型 + yield format_sse_event( + "failed", + { + "query_id": stream_id, + "trace_id": trace_id, + "error_type": type(exc).__name__, + }, + ) + + return StreamingResponse( + events(), + media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, + ) + + +@router.post("/nl2sql/query") +async def query_data( + body: DataQueryReq, + request: Request, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), + milvus=Depends(get_milvus), + redis=Depends(get_redis), +): + """执行员工自然语言查询并返回数据结果。""" + ensure_query_employee(user) + trace_id = request.headers.get("X-Trace-Id") or get_request_id() or new_request_id() + query_id = uuid4().hex + clarification = build_clarification(body.question) + if clarification is not None: + return success({"status": "clarification_required", "clarification": clarification}) + permission = await _load_permission(db, user) + limiter = QueryLimiter(redis) + quota_acquired = await limiter.acquire( + user.id, + daily_quota=permission.get("daily_quota", 0), + max_concurrent=1, + rate_limit=30, + ) + if not quota_acquired: + query_metrics.record_rate_limited() + await write_nl2sql_audit_safely( + db, + user_id=user.id, + username=user.username, + action="limit_rejected", + target=query_id, + trace_id=trace_id, + detail={"question_length": len(body.question)}, + status="失败", + ) + raise ParamError("查询配额或并发限制已达上限") + session_lock = SessionLock(redis, body.session_id) if body.session_id else None + if session_lock is not None: + try: + await session_lock.acquire() + except SessionBusyError as exc: + raise ParamError("同一会话已有查询运行") from exc + + async def permission_loader(_user_id: int): + return permission + + async def metadata_retriever(question: str): + return await retrieve_metadata(question, milvus, top_k=runtime_config.retrieval_top_k) + + async def schema_loader(table_names: set[str], _permission: dict): + return await load_authoritative_schema( + db, + database=settings.mysql.database, + candidate_tables=table_names, + ) + + request_contract = DataQueryRequest( + question=body.question, + user_id=user.id, + trace_id=trace_id, + session_id=body.session_id, + caller_agent=body.caller_agent, + data_scope=body.data_scope, + max_rows=min(body.max_rows or runtime_config.max_rows, runtime_config.max_rows), + include_sql=body.include_sql, + page=body.page, + page_size=body.page_size, + sort_by=body.sort_by, + sort_order=body.sort_order, + ) + context_store = SessionContextStore(redis) + conversation_context = build_conversation_context( + await context_store.load(user.id, body.session_id) + ) + validated_sql = None + cache_hit = False + try: + validated_sql = await build_query( + request_contract, + permission_loader=permission_loader, + metadata_retriever=metadata_retriever, + schema_loader=schema_loader, + conversation_context=conversation_context, + ) + cache_key = build_cache_key(validated_sql.sql, permission=permission) + cached_payload = await cache_get(redis, cache_key) + if cached_payload is not None: + try: + cached_result = result_from_cache(cached_payload) + except Exception: # noqa: BLE001 缓存内容损坏时回退数据库 + cached_result = None + if cached_result is not None: + cache_hit = True + await write_nl2sql_audit_safely( + db, + user_id=user.id, + username=user.username, + action="cache_hit", + target=query_id, + trace_id=trace_id, + detail={"tables": sorted(validated_sql.access_tables)}, + ) + result = replace( + cached_result, + query_id=query_id, + trace_id=trace_id, + warnings=[*cached_result.warnings, "cache_hit"], + ) + else: + result = await execute_readonly_sql( + db, + validated_sql, + query_id=query_id, + trace_id=trace_id, + user_id=user.id, + masks=permission.get("masks"), + ) + else: + result = await execute_readonly_sql( + db, + validated_sql, + query_id=query_id, + trace_id=trace_id, + user_id=user.id, + masks=permission.get("masks"), + ) + if result.summary is None: + result = await enrich_query_result(body.question, result, llm_client=llm) + await cache_set( + redis, + cache_key, + result_to_cache(result), + ttl=runtime_config.cache_ttl, + access_tables=validated_sql.access_tables, + ) + await write_nl2sql_audit_safely( + db, + user_id=user.id, + username=user.username, + action="query", + target=query_id, + trace_id=trace_id, + detail={"tables": sorted(validated_sql.access_tables), "row_count": result.row_count}, + ) + await context_store.append(user.id, body.session_id, body.question, "success") + diagnostic_registry.record( + query_id=query_id, + status="success", + access_tables=validated_sql.access_tables, + sql=validated_sql.sql, + model_elapsed_ms=result.elapsed_ms, + ) + query_metrics.record( + status="success", + elapsed_ms=result.elapsed_ms, + cache_hit=cache_hit, + ) + except Exception as exc: # noqa: BLE001 查询失败统一归档后再转业务异常 + await archive_query_safely( + db, + query_id=query_id, + user_id=user.id, + question=body.question, + generated_sql=getattr(validated_sql, "sql", None), + access_tables=getattr(validated_sql, "access_tables", set()), + status="failed", + error_message="查询处理失败", + trace_id=trace_id, + session_id=body.session_id, + caller_agent=body.caller_agent, + ) + await write_nl2sql_audit_safely( + db, + user_id=user.id, + username=user.username, + action="query_failed", + target=query_id, + trace_id=trace_id, + detail={"error_type": type(exc).__name__}, + status="失败", + ) + diagnostic_registry.record( + query_id=query_id, + status="failed", + access_tables=getattr(validated_sql, "access_tables", set()), + sql=getattr(validated_sql, "sql", None), + security_rule=type(exc).__name__, + ) + await limiter.release(user.id) + query_metrics.record( + status="timeout" if type(exc).__name__ == "QueryExecutionError" and "超时" in str(exc) else "failed", + failure_reason=type(exc).__name__, + ) + if session_lock is not None: + await session_lock.release() + if isinstance(exc, QueryServiceError): + raise ParamError(str(exc)) from exc + raise + await archive_query_safely( + db, + **build_history_payload( + query_id=query_id, + request_data=body, + user=user, + trace_id=trace_id, + status="success", + result=validated_sql, + query_result=result, + ), + ) + if not body.include_sql: + result = replace(result, sql=None) + if session_lock is not None: + await session_lock.release() + await limiter.release(user.id) + if body.output_format == "csv": + await write_nl2sql_audit_safely( + db, + user_id=user.id, + username=user.username, + action="result_export", + target=query_id, + trace_id=trace_id, + detail={"format": "csv", "row_count": result.row_count}, + ) + return Response( + content="\ufeff" + render_csv(result.columns, result.rows), + media_type="text/csv; charset=utf-8", + headers={"Content-Disposition": f'attachment; filename="nl2sql-{query_id}.csv"'}, + ) + return success(asdict(result)) + + +@router.get("/nl2sql/health") +async def nl2sql_health(user: SysUser = Depends(get_current_user)): + """返回 NL2SQL 依赖状态,仅对已登录员工开放。""" + ensure_query_employee(user) + result = await check_nl2sql_health( + { + "mysql": database.mysql.check_health, + "redis": database.redis.check_health, + "milvus": database.milvus.check_health, + "llm": llm.check_health, + } + ) + return success(result) + + +@router.post("/nl2sql/query/explain") +async def explain_query( + body: DataExplainReq, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + """对当前员工有权限的 SELECT 返回执行计划。""" + ensure_query_employee(user) + permission = await _load_permission(db, user) + max_rows = permission.get("max_rows") or 1000 + try: + validated = validate_select_sql( + body.sql, + authorized_tables=permission.get("tables", set()), + authorized_columns=permission.get("columns"), + max_rows=max_rows, + ) + result = await db.execute(text(build_explain_sql(validated))) + return success([dict(row) for row in result.mappings().all()]) + except SqlSecurityError as exc: + raise ForbiddenError("SQL 未通过安全校验") from exc + + +@router.get("/nl2sql/query-history") +async def list_query_history( + page: int = Query(1, ge=1), + page_size: int = Query(20, ge=1, le=100), + status: str | None = Query(None, max_length=16), + start_time: datetime | None = Query(None), + end_time: datetime | None = Query(None), + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + """分页读取当前员工自己的查询历史。""" + ensure_query_employee(user) + rows = await Nl2SqlPermissionRepo(db).list_query_history( + user.id, + limit=page_size, + offset=(page - 1) * page_size, + status=status, + start_time=start_time, + end_time=end_time, + ) + return success({"page": page, "page_size": page_size, "items": [history_payload(row) for row in rows]}) + + +@router.get("/nl2sql/query-history/{query_id}") +async def get_query_history( + query_id: str, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + """读取当前员工自己的查询历史详情。""" + ensure_query_employee(user) + history = await Nl2SqlPermissionRepo(db).get_query_history(user.id, query_id) + if history is None: + raise NotFoundError("查询历史不存在") + return success(history_payload(history)) + + +@router.get("/nl2sql/query/running") +async def list_running_queries(user: SysUser = Depends(get_current_user)): + """管理员查看当前进程登记的运行中查询。""" + ensure_query_admin(user) + return success([item.to_dict() for item in query_runtime_registry.list_active()]) + + +@router.post("/nl2sql/query/kill") +async def kill_running_query( + body: DataKillReq, + user: SysUser = Depends(get_current_user), +): + """管理员尝试中止运行中查询,并返回是否已验证。""" + ensure_query_admin(user) + item = query_runtime_registry.get(body.query_id) + if item is None or item.status != "running" or item.connection_id is None: + return success({"query_id": body.query_id, "verified": False}) + try: + verified = await kill_mysql_query(item.connection_id) + except Exception: # noqa: BLE001 中止失败返回未验证,不暴露连接异常 + verified = False + if verified: + await query_runtime_registry.mark_killed(body.query_id) + return success({"query_id": body.query_id, "verified": verified}) + + +@router.get("/nl2sql/query-history/{query_id}/export") +async def export_query_history( + query_id: str, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + """导出查询元数据 CSV,不导出结果行、凭据或连接信息。""" + ensure_query_employee(user) + history = await Nl2SqlPermissionRepo(db).get_query_history(user.id, query_id) + if history is None: + raise NotFoundError("查询历史不存在") + await write_nl2sql_audit_safely( + db, + user_id=user.id, + username=user.username, + action="history_export", + target=query_id, + trace_id=get_request_id() or new_request_id(), + detail={"format": "csv"}, + ) + output = StringIO() + fieldnames = [ + "query_id", "question", "generated_sql", "access_tables", "status", + "row_count", "truncated", "elapsed_ms", "trace_id", "create_time", + ] + writer = csv.DictWriter(output, fieldnames=fieldnames) + writer.writeheader() + payload = history_payload(history) + writer.writerow({ + "query_id": payload["query_id"], + "question": payload["question"], + "generated_sql": payload["generated_sql"] or "", + "access_tables": ",".join(payload["access_tables"]), + "status": payload["status"], + "row_count": payload["row_count"], + "truncated": payload["truncated"], + "elapsed_ms": payload["elapsed_ms"] or "", + "trace_id": payload["trace_id"] or "", + "create_time": payload["create_time"] or "", + }) + return Response( + content="\ufeff" + output.getvalue(), + media_type="text/csv; charset=utf-8", + headers={"Content-Disposition": f'attachment; filename="nl2sql-{query_id}.csv"'}, + ) + + +@router.get("/nl2sql/query/diagnostics") +async def query_diagnostics(user: SysUser = Depends(get_current_user)): + """管理员查看脱敏运行诊断信息。""" + ensure_query_admin(user) + health = await check_nl2sql_health( + { + "mysql": database.mysql.check_health, + "redis": database.redis.check_health, + "milvus": database.milvus.check_health, + "llm": llm.check_health, + } + ) + return success({ + "health": health, + "running": [item.to_dict() for item in query_runtime_registry.list_active()], + "metrics": query_metrics.snapshot(), + "recent": diagnostic_registry.list_recent(), + }) + + +@router.post("/nl2sql/cache/invalidate") +async def invalidate_query_cache( + body: DataCacheInvalidateReq, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), + redis=Depends(get_redis), +): + """管理员按表失效查询结果缓存。""" + ensure_query_admin(user) + deleted = await invalidate_tables(redis, set(body.table_names)) + await write_nl2sql_audit_safely( + db, + user_id=user.id, + username=user.username, + action="cache_invalidate", + target=",".join(body.table_names), + trace_id=get_request_id() or new_request_id(), + detail={"deleted": deleted}, + ) + return success({"table_names": sorted(set(body.table_names)), "deleted": deleted}) diff --git a/api/routers/nl2sql_admin.py b/api/routers/nl2sql_admin.py new file mode 100644 index 0000000..3c62dea --- /dev/null +++ b/api/routers/nl2sql_admin.py @@ -0,0 +1,323 @@ +"""NL2SQL 管理接口。""" +from __future__ import annotations + +import json +import time +from datetime import datetime, timedelta, timezone + +from fastapi import APIRouter, Depends +from fastapi.responses import PlainTextResponse +from sqlalchemy.ext.asyncio import AsyncSession + +from api.deps import get_current_user +from config.deps import get_db, get_redis +from model.sys_user import SysUser +from nl2sql.audit import write_nl2sql_audit_safely +from nl2sql.metrics import load_history_metrics_safely, query_metrics, render_prometheus +from nl2sql.runtime_config import runtime_config +from nl2sql.job_history import list_job_history, record_job_history_safely +from nl2sql.jobs import ( + run_consistency_check, + run_history_cleanup, + run_metadata_sync, + run_vector_cleanup, +) +from nl2sql.semantics import get_semantic_catalog_info, refresh_semantic_catalog +from repositories.nl2sql_permission import Nl2SqlPermissionRepo +from schemas.nl2sql_admin import ( + ColumnPermissionCreateReq, + ColumnPermissionUpdateReq, + RoleCreateReq, + RoleUpdateReq, + SensitiveFieldCreateReq, + SensitiveFieldUpdateReq, + TablePermissionCreateReq, + TablePermissionUpdateReq, + MaintenanceJobReq, + RuntimeConfigUpdateReq, +) +from service.nl2sql.admin_service import ( + column_permission_payload, + create_role, + role_payload, + sensitive_field_payload, + table_permission_payload, + validate_mask_type, + validate_table_permission, +) +from utils.exceptions import NotFoundError, ParamError +from utils.request_id import get_request_id, new_request_id +from utils.response import success + +from api.routers.nl2sql import ensure_query_admin + +router = APIRouter() + + +async def _audit( + db, + user: SysUser, + action: str, + target: str, + detail: dict | None = None, + status: str = "成功", +): + await write_nl2sql_audit_safely( + db, + user_id=user.id, + username=user.username, + action=action, + target=target, + trace_id=get_request_id() or new_request_id(), + detail=detail, + status=status, + ) + + +@router.get("/nl2sql/admin/roles") +async def list_roles(user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + rows = await Nl2SqlPermissionRepo(db).list_roles(include_inactive=True) + return success([role_payload(row) for row in rows]) + + +@router.post("/nl2sql/admin/roles") +async def add_role(body: RoleCreateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + try: + role = await create_role(db, body.model_dump()) + except ValueError as exc: + raise ParamError(str(exc)) from exc + await _audit(db, user, "permission_role_create", str(role.id), {"role_code": role.role_code}) + return success(role_payload(role)) + + +@router.patch("/nl2sql/admin/roles/{role_id}") +async def edit_role(role_id: int, body: RoleUpdateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + role = await Nl2SqlPermissionRepo(db).update_role(role_id, **body.model_dump(exclude_none=True)) + if role is None: + raise NotFoundError("NL2SQL 角色不存在") + await _audit(db, user, "permission_role_update", str(role_id)) + return success(role_payload(role)) + + +@router.get("/nl2sql/admin/roles/{role_id}/tables") +async def list_tables(role_id: int, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + rows = await Nl2SqlPermissionRepo(db).list_role_table_permissions(role_id, include_inactive=True) + return success([table_permission_payload(row) for row in rows]) + + +@router.post("/nl2sql/admin/roles/{role_id}/tables") +async def add_table(role_id: int, body: TablePermissionCreateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + if await Nl2SqlPermissionRepo(db).get_role(role_id) is None: + raise NotFoundError("NL2SQL 角色不存在") + try: + values = validate_table_permission({**body.model_dump(), "role_id": role_id}) + item = await Nl2SqlPermissionRepo(db).add_table_permission(role_id=role_id, **values) + except ValueError as exc: + raise ParamError(str(exc)) from exc + await _audit(db, user, "permission_table_create", str(item.id), {"table_name": item.table_name}) + return success(table_permission_payload(item)) + + +@router.patch("/nl2sql/admin/table-permissions/{permission_id}") +async def edit_table(permission_id: int, body: TablePermissionUpdateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + item = await Nl2SqlPermissionRepo(db).update_table_permission(permission_id, **body.model_dump(exclude_none=True)) + if item is None: + raise NotFoundError("NL2SQL 表权限不存在") + await _audit(db, user, "permission_table_update", str(permission_id)) + return success(table_permission_payload(item)) + + +@router.get("/nl2sql/admin/roles/{role_id}/columns") +async def list_columns(role_id: int, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + rows = await Nl2SqlPermissionRepo(db).list_role_column_permissions(role_id, include_inactive=True) + return success([column_permission_payload(row) for row in rows]) + + +@router.post("/nl2sql/admin/roles/{role_id}/columns") +async def add_column(role_id: int, body: ColumnPermissionCreateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + if await Nl2SqlPermissionRepo(db).get_role(role_id) is None: + raise NotFoundError("NL2SQL 角色不存在") + if body.access_mode == "mask": + try: + validate_mask_type(body.mask_type) + except ValueError as exc: + raise ParamError(str(exc)) from exc + item = await Nl2SqlPermissionRepo(db).add_column_permission(role_id=role_id, **body.model_dump()) + await _audit(db, user, "permission_column_create", str(item.id), {"table_name": item.table_name}) + return success(column_permission_payload(item)) + + +@router.patch("/nl2sql/admin/column-permissions/{permission_id}") +async def edit_column(permission_id: int, body: ColumnPermissionUpdateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + values = body.model_dump(exclude_none=True) + if "mask_type" in values: + try: + validate_mask_type(values["mask_type"]) + except ValueError as exc: + raise ParamError(str(exc)) from exc + item = await Nl2SqlPermissionRepo(db).update_column_permission(permission_id, **values) + if item is None: + raise NotFoundError("NL2SQL 字段权限不存在") + await _audit(db, user, "permission_column_update", str(permission_id)) + return success(column_permission_payload(item)) + + +@router.get("/nl2sql/admin/sensitive-fields") +async def list_sensitive(user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + rows = await Nl2SqlPermissionRepo(db).list_sensitive_fields_admin(include_inactive=True) + return success([sensitive_field_payload(row) for row in rows]) + + +@router.post("/nl2sql/admin/sensitive-fields") +async def add_sensitive(body: SensitiveFieldCreateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + item = await Nl2SqlPermissionRepo(db).add_sensitive_field(**body.model_dump()) + await _audit(db, user, "permission_sensitive_create", str(item.id), {"table_name": item.table_name}) + return success(sensitive_field_payload(item)) + + +@router.patch("/nl2sql/admin/sensitive-fields/{field_id}") +async def edit_sensitive(field_id: int, body: SensitiveFieldUpdateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + item = await Nl2SqlPermissionRepo(db).update_sensitive_field(field_id, **body.model_dump(exclude_none=True)) + if item is None: + raise NotFoundError("NL2SQL 敏感字段不存在") + await _audit(db, user, "permission_sensitive_update", str(field_id)) + return success(sensitive_field_payload(item)) + + +@router.get("/nl2sql/admin/metrics") +async def admin_metrics(user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + ensure_query_admin(user) + return success({ + "runtime": query_metrics.snapshot(), + "history": await load_history_metrics_safely(db), + }) + + +@router.get("/nl2sql/admin/runtime-config") +async def get_runtime_config(user: SysUser = Depends(get_current_user)): + """管理员查看当前进程内 NL2SQL 运行参数。""" + ensure_query_admin(user) + return success(runtime_config.model_dump()) + + +@router.patch("/nl2sql/admin/runtime-config") +async def update_runtime_config( + body: RuntimeConfigUpdateReq, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + """管理员更新当前进程内 NL2SQL 非敏感运行参数。""" + ensure_query_admin(user) + try: + values = runtime_config.update(**body.model_dump(exclude_none=True)) + except ValueError as exc: + raise ParamError("运行参数不合法") from exc + await _audit(db, user, "runtime_config_update", "nl2sql", {"fields": sorted(body.model_dump(exclude_none=True))}) + return success(values) + + +@router.get("/nl2sql/admin/metrics/prometheus", response_class=PlainTextResponse) +async def admin_metrics_prometheus(user: SysUser = Depends(get_current_user)): + """管理员读取聚合 Prometheus 指标,不返回查询明细。""" + ensure_query_admin(user) + return render_prometheus() + + +@router.post("/nl2sql/admin/jobs") +async def run_admin_job( + body: MaintenanceJobReq, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), + redis=Depends(get_redis), +): + """管理员手动执行一个幂等运维任务并保存执行历史。""" + ensure_query_admin(user) + started = time.perf_counter() + if body.task == "metadata_sync": + result = await run_metadata_sync(redis=redis) + elif body.task == "vector_cleanup": + result = await run_vector_cleanup(redis=redis) + elif body.task == "consistency_check": + from scripts.check_nl2sql_consistency import collect_consistency + + result = await run_consistency_check(redis=redis, worker=collect_consistency) + else: + before = datetime.now(timezone.utc).replace(tzinfo=None) - timedelta(days=body.before_days) + result = await run_history_cleanup(db, before, redis=redis) + await record_job_history_safely( + db, + result, + elapsed_ms=(time.perf_counter() - started) * 1000, + parameter_summary={"before_days": body.before_days} if body.task == "history_cleanup" else {}, + ) + await _audit(db, user, "maintenance_job_run", result.name, {"status": result.status}) + return success({ + "name": result.name, + "status": result.status, + "attempts": result.attempts, + "detail": result.detail, + "error_type": result.error_type, + }) + + +@router.get("/nl2sql/admin/jobs/history") +async def admin_job_history( + page: int = 1, + page_size: int = 20, + status: str | None = None, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + """管理员查询运维任务执行历史摘要。""" + ensure_query_admin(user) + return success(await list_job_history(db, page=page, page_size=page_size, status=status)) + + +@router.get("/nl2sql/admin/semantics") +async def admin_semantics(user: SysUser = Depends(get_current_user)): + """管理员查看当前语义目录版本和规模摘要。""" + ensure_query_admin(user) + return success(get_semantic_catalog_info()) + + +@router.post("/nl2sql/admin/semantics/refresh") +async def refresh_semantics(user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)): + """管理员刷新语义目录缓存,目录无效时继续使用旧缓存。""" + ensure_query_admin(user) + try: + result = refresh_semantic_catalog() + except (OSError, ValueError, json.JSONDecodeError) as exc: + await _audit( + db, + user, + "semantic_catalog_refresh", + "default", + {"error_type": type(exc).__name__}, + status="失败", + ) + raise ParamError("语义目录刷新失败") from exc + await _audit( + db, + user, + "semantic_catalog_refresh", + "default", + { + "version": result["version"], + "previous_version": result["previous_version"], + "changed": result["changed"], + "digest": result["digest"], + }, + ) + return success(result) diff --git a/common/__init__.py b/common/__init__.py new file mode 100644 index 0000000..89b34ef --- /dev/null +++ b/common/__init__.py @@ -0,0 +1 @@ +"""跨投顾 Agent 与投顾工作台复用的公共契约和业务工具。""" diff --git a/common/common_const.py b/common/common_const.py new file mode 100644 index 0000000..7313563 --- /dev/null +++ b/common/common_const.py @@ -0,0 +1,82 @@ +"""docs/common_const.md 的可执行 Python 投影。 + +文档仍是公共契约的维护来源;业务代码统一从本模块引用具体常量。 +""" + +CUSTOMER_REL_STATUS_UNSIGNED = "已分配" +CUSTOMER_REL_STATUS_SIGNED = "已签约" +CUSTOMER_REL_STATUS_CLOSED = "已结束" + +DRAFT_STATUS_DRAFT = "draft" +DRAFT_STATUS_DISCARDED = "discarded" + +MEMORY_INFO_TYPE_FACT = "FACT" +MEMORY_INFO_TYPE_OPINION = "OPINION" + +AGENT_INTENT_RECOMMEND = "recommend" +AGENT_INTENT_REBALANCE = "rebalance" +AGENT_INTENT_FUND_ANALYSIS = "fund_analysis" +AGENT_INTENT_DIALOGUE_SCRIPT = "dialogue-script" + +TALK_SCENE_RISK_BLOCK_ORDER = "risk_block_order" +TALK_SCENE_MARKET_FLUCTUATION = "market_fluctuation" +TALK_SCENE_PORTFOLIO_DIVERGENCE = "portfolio_divergence" +TALK_SCENE_CUSTOMER_COMPLAINT = "customer_complaint" + +ERR_CODE_OK = 0 +ERR_CODE_FORBIDDEN_CUSTOMER = 40001 +ERR_CODE_SUITABILITY_INVALID = 40020 +ERR_CODE_NOT_SIGNED_REBALANCE = 40030 +ERR_CODE_DRAFT_NOT_FOUND = 40401 +ERR_CODE_LLM_ERROR = 50001 +ERR_CODE_GRAPH_ERROR = 50002 + +SYS_KEY_REBALANCE_DEVIATION_THRESHOLD = "rebalance.deviation.threshold" +SYS_KEY_HIGH_NET_ASSET_THRESHOLD = "customer.high_net.asset.threshold" +SYS_KEY_RISK_QUESTIONNAIRE_EXPIRE_DAY = "risk.questionnaire.expire.day" +SYS_KEY_LARGE_FLOW_THRESHOLD = "customer.large_flow.threshold" + +C_RISK_C1 = "C1" +C_RISK_C2 = "C2" +C_RISK_C3 = "C3" +C_RISK_C4 = "C4" +C_RISK_C5 = "C5" + +PROD_RISK_R1 = "R1" +PROD_RISK_R2 = "R2" +PROD_RISK_R3 = "R3" +PROD_RISK_R4 = "R4" +PROD_RISK_R5 = "R5" + +EVENT_ADVISOR_REBALANCE_DRAFT_CREATED = "event:rebalance_draft_created" +EVENT_PROFILE_UPDATE = "event:profile_update" +EVENT_WORK_ORDER_CHANGE = "event:work_order_change" +EVENT_RISK_ALERT = "event:risk_alert" + +REPORT_DISCLAIMER = ( + "【免责声明】本报告由AI辅助生成,仅供持牌投顾内部参考,不构成任何投资建议。" + "基金有风险,投资需谨慎。所有投资决策请结合自身风险承受能力审慎判断。" +) + +SSE_EVENT_TYPE_TEXT = "text" +SSE_EVENT_TYPE_META = "meta" +SSE_EVENT_TYPE_DONE = "done" +SSE_EVENT_TYPE_ERROR = "error" + +AUDIT_AGENT_CHAT_CALL = "agent_chat_call" +AUDIT_DRAFT_SAVE = "draft_save" +AUDIT_DRAFT_DISCARD = "draft_discard" +AUDIT_CUSTOMER_SIGN = "customer_sign" +AUDIT_VIEW_SENSITIVE = "view_sensitive" + +EMPLOYEE_ROLE_ADVISOR = "投顾" + +CRON_MEMORY_MAINTENANCE = "0 2 * * 1" +CRON_PORTFOLIO_REBALANCE = "0 1 * * *" +CRON_FUND_NAV_UPDATE = "30 1 * * *" + +MSG_TYPE_RECOMMEND = "推荐" +MSG_TYPE_REBALANCE = "调仓" +MSG_TYPE_RISK = "风控" +MSG_TYPE_SIGN = "签约" +MSG_TYPE_SYSTEM = "系统" diff --git a/common/suitability.py b/common/suitability.py new file mode 100644 index 0000000..5416658 --- /dev/null +++ b/common/suitability.py @@ -0,0 +1,29 @@ +"""客户 C 级与产品 R 级的公共适当性校验。""" +from __future__ import annotations + +from dataclasses import dataclass +import re + + +@dataclass(frozen=True) +class SuitabilityResult: + ok: bool + reason: str + + +_RISK_PATTERN = re.compile(r"^[CR]([1-5])$") + + +def check_suitability(customer_risk: str, product_risk: str) -> SuitabilityResult: + customer_match = _RISK_PATTERN.fullmatch(customer_risk or "") + product_match = _RISK_PATTERN.fullmatch(product_risk or "") + if not customer_match or not product_match: + return SuitabilityResult(False, "客户或产品风险等级无效") + if customer_risk[0] != "C" or product_risk[0] != "R": + return SuitabilityResult(False, "客户或产品风险等级类型无效") + + customer_level = int(customer_match.group(1)) + product_level = int(product_match.group(1)) + if product_level > customer_level: + return SuitabilityResult(False, "产品风险等级高于客户风险等级") + return SuitabilityResult(True, "适当性校验通过") diff --git a/common_const.py b/common_const.py index 541ba3e..ab7d717 100644 --- a/common_const.py +++ b/common_const.py @@ -11,9 +11,9 @@ from __future__ import annotations # --------------------------------------------------------------------------- # §1 客户-投顾关系状态(customer_relation.status) # --------------------------------------------------------------------------- -CUSTOMER_REL_STATUS_UNSIGNED = "unsigned" # 未签约,仅分配,未开通投顾正式服务 -CUSTOMER_REL_STATUS_SIGNED = "signed" # 已签约,可下发推荐/调仓方案给客户 -CUSTOMER_REL_STATUS_CLOSED = "closed" # 服务已终止 +CUSTOMER_REL_STATUS_UNSIGNED = "已分配" # 未签约,仅分配,未开通投顾正式服务 +CUSTOMER_REL_STATUS_SIGNED = "已签约" # 已签约,可下发推荐/调仓方案给客户 +CUSTOMER_REL_STATUS_CLOSED = "已结束" # 服务已终止 # --------------------------------------------------------------------------- # §2 草稿状态(advisor_draft.status,Agent 侧;不可物理删除,仅状态流转) @@ -92,8 +92,9 @@ EVENT_PROFILE_UPDATE = "event:profile_update" ADVISOR_EVENTS = (EVENT_REBALANCE_DRAFT_CREATED, EVENT_PROFILE_UPDATE) # 事件消费状态(event_log.status) -EVENT_STATUS_PENDING = "待消费" -EVENT_STATUS_CONSUMED = "已消费" +EVENT_STATUS_PENDING = "pending" +EVENT_STATUS_CONSUMED = "consumed" +EVENT_STATUS_FAILED = "failed" # --------------------------------------------------------------------------- # §10 固定文本模板 diff --git a/config/settings.py b/config/settings.py index 3549b8c..adc42b3 100644 --- a/config/settings.py +++ b/config/settings.py @@ -108,6 +108,32 @@ class LLMCfg(BaseSettings): model_config = SettingsConfigDict(env_prefix="LLM_", env_file=_ENV_FILE, extra="ignore") +class AdvisorAgentCfg(BaseSettings): + """投顾工作台调用 Agent 服务的配置。""" + + base_url: str = "" + timeout: float = 1.0 + request_timeout: float = 1.0 + retry: int = 1 + llm_timeout: float = 5.0 + graph_timeout: float = 2.0 + milvus_timeout: float = 2.0 + + model_config = SettingsConfigDict( + env_prefix="ADVISOR_AGENT_", env_file=_ENV_FILE, extra="ignore" + ) + + +class AdvisorCfg(BaseSettings): + """投顾工作台后台任务开关。""" + + event_consumer_enabled: bool = False + event_retry_interval_sec: float = 30.0 + scheduler_enabled: bool = False + + model_config = SettingsConfigDict(env_prefix="ADVISOR_", env_file=_ENV_FILE, extra="ignore") + + class JwtCfg(BaseSettings): """JWT 鉴权配置(service/auth.py)。""" secret: str # 签名密钥,生产必须换强随机值 @@ -136,6 +162,8 @@ class Settings(BaseSettings): neo4j: Neo4jCfg = Neo4jCfg() milvus: MilvusCfg = MilvusCfg() llm: LLMCfg = LLMCfg() + advisor_agent: AdvisorAgentCfg = AdvisorAgentCfg() + advisor: AdvisorCfg = AdvisorCfg() model_config = SettingsConfigDict(env_file=_ENV_FILE, extra="ignore") diff --git a/main.py b/main.py index cac9545..738f56c 100644 --- a/main.py +++ b/main.py @@ -7,42 +7,96 @@ from fastapi import FastAPI from api.router import api_router from config import database -from rag.milvus_collections import ensure_collections +from config.database import redis as redis_db +from config.settings import settings +from agent.advisor_agent.runtime import build_default_runtime as advisor_build_default_runtime from service.customer_agent.bootstrap import ( build_default_knowledge_upload_service, - build_default_runtime, + build_default_runtime as customer_build_default_runtime, ) from service.client_agent.bootstrap import build_default_runtime as build_client_runtime from service.client_agent.idle_archive_worker import IdleArchiveWorker +from service.advisor.event_consumer import EventConsumer +from service.advisor.scheduler import AdvisorScheduler +from service.memory.facade import MemoryService +from tool.llm import llm as llm_client from utils.exceptions import register_exception_handlers from utils.logger import setup_logging from utils.request_id import RequestIdMiddleware +from utils.performance import PerformanceMiddleware + + +class LazyResource: + """Construct an application resource on first use and reuse it thereafter.""" + + def __init__(self, factory): + self._factory = factory + self._value = None + self._initialized = False + + @property + def initialized(self): + return self._initialized + + def _get(self): + if not self._initialized: + self._value = self._factory() + self._initialized = True + return self._value + + def get(self, key, default=None): + return self._get().get(key, default) + + def __getattr__(self, name): + return getattr(self._get(), name) @asynccontextmanager async def lifespan(app: FastAPI): setup_logging() - await ensure_collections() - app.state.customer_agent_runtime = build_default_runtime() - app.state.client_agent_runtime = build_client_runtime() - app.state.client_agent_archive_worker = IdleArchiveWorker( - redis=app.state.client_agent_runtime.redis, - memory_service=app.state.client_agent_runtime.memory_service, + app.state.customer_agent_runtime = LazyResource(customer_build_default_runtime) + app.state.client_agent_runtime = LazyResource(build_client_runtime) + app.state.advisor_agent_runtime = LazyResource( + lambda: advisor_build_default_runtime( + llm_client=llm_client, + memory_service=MemoryService(), + ) ) - await app.state.client_agent_archive_worker.recover_due_index() - app.state.client_agent_archive_task = asyncio.create_task( - app.state.client_agent_archive_worker.run() + app.state.client_agent_archive_worker = LazyResource( + lambda: IdleArchiveWorker( + redis=app.state.client_agent_runtime.redis, + memory_service=app.state.client_agent_runtime.memory_service, + ) + ) + app.state.client_agent_archive_task = None + app.state.advisor_event_consumer = None + if settings.advisor.event_consumer_enabled: + app.state.advisor_event_consumer = EventConsumer(redis_db.client()) + app.state.advisor_event_consumer.start() + app.state.advisor_scheduler = None + if settings.advisor.scheduler_enabled: + app.state.advisor_scheduler = AdvisorScheduler() + app.state.advisor_scheduler.start() + app.state.knowledge_upload_service = LazyResource( + build_default_knowledge_upload_service ) - app.state.knowledge_upload_service = build_default_knowledge_upload_service() yield - await app.state.client_agent_archive_worker.stop() - await app.state.client_agent_archive_task + worker = app.state.client_agent_archive_worker + if worker.initialized: + await worker.stop() + if app.state.client_agent_archive_task is not None: + await app.state.client_agent_archive_task + if app.state.advisor_event_consumer is not None: + await app.state.advisor_event_consumer.stop() + if app.state.advisor_scheduler is not None: + app.state.advisor_scheduler.shutdown() await database.dispose() app = FastAPI(title="智能公募基金系统", version="0.1.0", lifespan=lifespan) app.add_middleware(RequestIdMiddleware) +app.add_middleware(PerformanceMiddleware) register_exception_handlers(app) app.include_router(api_router) diff --git a/model/advisor_draft.py b/model/advisor_draft.py new file mode 100644 index 0000000..49f56d7 --- /dev/null +++ b/model/advisor_draft.py @@ -0,0 +1,41 @@ +"""投顾 Agent 草稿 ORM 模型。""" +from __future__ import annotations + +from datetime import datetime +from decimal import Decimal + +from sqlalchemy import BigInteger, CheckConstraint, DateTime, JSON, Numeric, String, Text, func +from sqlalchemy.orm import Mapped, mapped_column + +from common.common_const import DRAFT_STATUS_DRAFT +from model.base import Base + + +class AdvisorDraft(Base): + __tablename__ = "advisor_draft" + __table_args__ = ( + CheckConstraint( + "status IN ('draft', 'discarded')", + name="ck_advisor_draft_status", + ), + {"comment": "投顾 Agent 草稿(不存 sent,不生成交易指令)"}, + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + draft_id: Mapped[str] = mapped_column(String(64), unique=True, index=True) + customer_id: Mapped[int] = mapped_column(BigInteger, index=True) + advisor_id: Mapped[int] = mapped_column(BigInteger, index=True) + intent: Mapped[str] = mapped_column(String(32), index=True) + title: Mapped[str] = mapped_column(String(128)) + content: Mapped[str] = mapped_column(Text) + structured_data: Mapped[dict | None] = mapped_column(JSON) + status: Mapped[str] = mapped_column(String(16), default=DRAFT_STATUS_DRAFT, index=True) + deviation: Mapped[Decimal | None] = mapped_column(Numeric(10, 4)) + disclaimer_ok: Mapped[bool] = mapped_column(default=False) + warning: Mapped[str | None] = mapped_column(String(512)) + create_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), index=True + ) + update_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), onupdate=func.now() + ) diff --git a/model/customer_relation.py b/model/customer_relation.py index 24ddf2d..dd29d7c 100644 --- a/model/customer_relation.py +++ b/model/customer_relation.py @@ -1,7 +1,7 @@ """customer_relation 客户-投顾关系表 ORM 模型。 同时服务于: -- 投顾工作台(advisor):签约状态 unsigned/signed/closed,工作台为唯一写入方; +- 投顾工作台(advisor):签约状态 已分配/已签约/已结束,工作台为唯一写入方; - 记忆/client_agent 模块:读取客户与投顾的当前或历史关系。 """ @@ -9,15 +9,24 @@ from __future__ import annotations from datetime import datetime -from sqlalchemy import BigInteger, DateTime, String, func +from sqlalchemy import BigInteger, CheckConstraint, DateTime, String, func from sqlalchemy.orm import Mapped, mapped_column +from common.common_const import ( + CUSTOMER_REL_STATUS_UNSIGNED, +) from model.base import Base class CustomerRelation(Base): __tablename__ = "customer_relation" - __table_args__ = {"comment": "客户-投顾关系表(状态驱动:签约后投顾Agent方案正式触达客户)"} + __table_args__ = ( + CheckConstraint( + "status IN ('已分配', '已签约', '已结束')", + name="ck_customer_relation_status", + ), + {"comment": "客户-投顾关系表(状态驱动:签约后投顾Agent方案正式触达客户)"}, + ) id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) customer_id: Mapped[int] = mapped_column(BigInteger, nullable=False) @@ -28,6 +37,6 @@ class CustomerRelation(Base): signed_time: Mapped[datetime | None] = mapped_column(DateTime) end_time: Mapped[datetime | None] = mapped_column(DateTime) status: Mapped[str] = mapped_column( - String(16), nullable=False, server_default="unsigned" + String(16), nullable=False, server_default=CUSTOMER_REL_STATUS_UNSIGNED ) reason: Mapped[str | None] = mapped_column(String(128)) diff --git a/model/event_log.py b/model/event_log.py index 4a6c387..293f4fa 100644 --- a/model/event_log.py +++ b/model/event_log.py @@ -19,6 +19,6 @@ class EventLog(Base): event_name: Mapped[str] = mapped_column(String(64)) payload: Mapped[dict[str, Any] | None] = mapped_column(JSON) trace_id: Mapped[str | None] = mapped_column(String(64)) - status: Mapped[str] = mapped_column(String(16), server_default="待消费") + status: Mapped[str] = mapped_column(String(16), server_default="pending") create_time: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) consume_time: Mapped[datetime | None] = mapped_column(DateTime) diff --git a/model/fund_performance.py b/model/fund_performance.py new file mode 100644 index 0000000..1f053b4 --- /dev/null +++ b/model/fund_performance.py @@ -0,0 +1,24 @@ +"""基金业绩指标缓存 ORM 模型。""" +from __future__ import annotations + +from datetime import date, datetime +from decimal import Decimal + +from sqlalchemy import BigInteger, Date, DateTime, Numeric, String, func +from sqlalchemy.orm import Mapped, mapped_column + +from model.base import Base + + +class FundPerformance(Base): + __tablename__ = "fund_performance" + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + product_id: Mapped[int] = mapped_column(BigInteger, index=True) + period: Mapped[str] = mapped_column(String(16)) + return_rate: Mapped[Decimal | None] = mapped_column(Numeric(10, 4)) + annual_volatility: Mapped[Decimal | None] = mapped_column(Numeric(10, 4)) + max_drawdown: Mapped[Decimal | None] = mapped_column(Numeric(10, 4)) + sharpe: Mapped[Decimal | None] = mapped_column(Numeric(10, 4)) + calc_date: Mapped[date] = mapped_column(Date) + create_time: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) diff --git a/model/nl2sql_permission.py b/model/nl2sql_permission.py new file mode 100644 index 0000000..c474499 --- /dev/null +++ b/model/nl2sql_permission.py @@ -0,0 +1,129 @@ +"""NL2SQL 查询权限和查询历史 ORM 模型。""" +from __future__ import annotations + +from datetime import datetime + +from sqlalchemy import BigInteger, Boolean, CheckConstraint, DateTime, JSON, String, Text, func +from sqlalchemy.orm import Mapped, mapped_column + +from model.base import Base + + +class Nl2SqlQueryRole(Base): + __tablename__ = "nl2sql_query_role" + __table_args__ = ( + CheckConstraint("can_query IN (0, 1)", name="ck_nl2sql_role_can_query"), + CheckConstraint("status IN ('active', 'inactive')", name="ck_nl2sql_role_status"), + {"comment": "NL2SQL 查询角色,按 sys_user.employee_role 映射"}, + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + role_code: Mapped[str] = mapped_column(String(64), unique=True) + role_name: Mapped[str] = mapped_column(String(128)) + employee_role: Mapped[str] = mapped_column(String(32), unique=True) + can_query: Mapped[bool] = mapped_column(Boolean, default=False) + max_rows: Mapped[int] = mapped_column(BigInteger, default=1000) + daily_quota: Mapped[int] = mapped_column(BigInteger, default=0) + status: Mapped[str] = mapped_column(String(16), default="active", index=True) + create_time: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) + update_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), onupdate=func.now() + ) + + +class Nl2SqlRoleTablePermission(Base): + __tablename__ = "nl2sql_role_table_permission" + __table_args__ = ( + CheckConstraint("permission = 'SELECT'", name="ck_nl2sql_table_permission"), + CheckConstraint( + "row_scope_type IN ('none', 'customer_ids', 'product_ids')", + name="ck_nl2sql_row_scope_type", + ), + CheckConstraint("status IN ('active', 'inactive')", name="ck_nl2sql_table_status"), + {"comment": "NL2SQL 查询角色的表和行级权限"}, + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + role_id: Mapped[int] = mapped_column(BigInteger, index=True) + table_name: Mapped[str] = mapped_column(String(128)) + permission: Mapped[str] = mapped_column(String(16), default="SELECT") + row_scope_type: Mapped[str] = mapped_column(String(32), default="none") + row_scope_column: Mapped[str | None] = mapped_column(String(128)) + status: Mapped[str] = mapped_column(String(16), default="active", index=True) + create_time: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) + update_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), onupdate=func.now() + ) + + +class Nl2SqlRoleColumnPermission(Base): + __tablename__ = "nl2sql_role_column_permission" + __table_args__ = ( + CheckConstraint( + "access_mode IN ('allow', 'deny', 'mask')", + name="ck_nl2sql_column_access_mode", + ), + CheckConstraint("status IN ('active', 'inactive')", name="ck_nl2sql_column_status"), + {"comment": "NL2SQL 查询角色的字段权限"}, + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + role_id: Mapped[int] = mapped_column(BigInteger, index=True) + table_name: Mapped[str] = mapped_column(String(128)) + column_name: Mapped[str] = mapped_column(String(128)) + access_mode: Mapped[str] = mapped_column(String(16), default="allow") + mask_type: Mapped[str | None] = mapped_column(String(32)) + status: Mapped[str] = mapped_column(String(16), default="active", index=True) + create_time: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) + update_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), onupdate=func.now() + ) + + +class Nl2SqlSensitiveField(Base): + __tablename__ = "nl2sql_sensitive_field" + __table_args__ = ( + CheckConstraint("status IN ('active', 'inactive')", name="ck_nl2sql_sensitive_status"), + {"comment": "NL2SQL 全局敏感字段规则"}, + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + table_name: Mapped[str] = mapped_column(String(128)) + column_name: Mapped[str] = mapped_column(String(128)) + mask_type: Mapped[str] = mapped_column(String(32), default="partial") + status: Mapped[str] = mapped_column(String(16), default="active", index=True) + description: Mapped[str | None] = mapped_column(String(255)) + create_time: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) + update_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), onupdate=func.now() + ) + + +class Nl2SqlQueryHistory(Base): + __tablename__ = "nl2sql_query_history" + __table_args__ = ( + CheckConstraint( + "status IN ('success', 'failed', 'blocked', 'timeout')", + name="ck_nl2sql_history_status", + ), + {"comment": "NL2SQL 查询归档,不保存完整结果行"}, + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + query_id: Mapped[str] = mapped_column(String(64), unique=True, index=True) + user_id: Mapped[int] = mapped_column(BigInteger, index=True) + session_id: Mapped[str | None] = mapped_column(String(64), index=True) + caller_agent: Mapped[str | None] = mapped_column(String(64), index=True) + question: Mapped[str] = mapped_column(Text) + generated_sql: Mapped[str | None] = mapped_column(Text) + access_tables: Mapped[list | None] = mapped_column(JSON) + status: Mapped[str] = mapped_column(String(16), index=True) + error_code: Mapped[str | None] = mapped_column(String(64)) + error_message: Mapped[str | None] = mapped_column(String(512)) + row_count: Mapped[int] = mapped_column(BigInteger, default=0) + truncated: Mapped[bool] = mapped_column(Boolean, default=False) + elapsed_ms: Mapped[float | None] = mapped_column() + trace_id: Mapped[str | None] = mapped_column(String(64), index=True) + create_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), index=True + ) diff --git a/nl2sql/audit.py b/nl2sql/audit.py new file mode 100644 index 0000000..1f5ecf7 --- /dev/null +++ b/nl2sql/audit.py @@ -0,0 +1,66 @@ +"""NL2SQL 审计日志适配器,复用现有 audit_log 表。""" +from __future__ import annotations + +import json + +from sqlalchemy import text + + +_SENSITIVE_KEYS = {"password", "token", "secret", "api_key", "connection", "credential"} + + +def _sanitize(value): + """递归移除连接凭据等不应进入审计详情的字段。""" + if isinstance(value, dict): + return { + key: "[REDACTED]" if str(key).lower() in _SENSITIVE_KEYS else _sanitize(item) + for key, item in value.items() + } + if isinstance(value, (list, tuple)): + return [_sanitize(item) for item in value] + return value + + +async def write_nl2sql_audit( + db, + *, + user_id: int | None, + username: str | None, + action: str, + target: str | None, + trace_id: str | None, + detail: dict | None = None, + status: str = "成功", +) -> None: + """写入 NL2SQL 审计事件,调用方负责在失败时降级。""" + statement = text( + """ + INSERT INTO audit_log + (user_id, username, module, action, target, detail, trace_id, status) + VALUES + (:user_id, :username, :module, :action, :target, :detail, :trace_id, :status) + """ + ) + await db.execute( + statement, + { + "user_id": user_id, + "username": username, + "module": "nl2sql", + "action": action, + "target": target, + "detail": json.dumps(_sanitize(detail or {}), ensure_ascii=False), + "trace_id": trace_id, + "status": status, + }, + ) + await db.commit() + + +async def write_nl2sql_audit_safely(db, **kwargs) -> bool: + """审计写入失败时记录 False,不阻断查询主流程。""" + try: + await write_nl2sql_audit(db, **kwargs) + except Exception: # noqa: BLE001 审计故障不能影响查询能力 + return False + return True diff --git a/nl2sql/cache.py b/nl2sql/cache.py new file mode 100644 index 0000000..f3176d8 --- /dev/null +++ b/nl2sql/cache.py @@ -0,0 +1,121 @@ +"""NL2SQL 查询缓存适配器,Redis 不可用时自动降级。""" +from __future__ import annotations + +import hashlib +import json +import logging +from typing import Any + +from sqlglot import parse_one + +from nl2sql.contracts import DataQueryResult + + +logger = logging.getLogger("nl2sql.cache") +DEFAULT_CACHE_TTL = 300 + + +def normalize_sql(sql: str) -> str: + """统一 SQL 的关键字、空白和末尾分号。""" + return parse_one(sql, read="mysql").sql(dialect="mysql") + + +def _permission_payload(permission: dict[str, Any]) -> dict[str, Any]: + """将权限集合转换成稳定、可哈希的结构。""" + return { + "tables": sorted(permission.get("tables", set())), + "columns": { + table: sorted(columns) + for table, columns in sorted(permission.get("columns", {}).items()) + }, + } + + +def build_cache_key( + sql: str, + *, + permission: dict[str, Any], + data_version: str = "", +) -> str: + """生成包含标准化 SQL、权限范围和数据版本的缓存键。""" + payload = { + "sql": normalize_sql(sql), + "permission": _permission_payload(permission), + "data_version": data_version, + } + digest = hashlib.sha256( + json.dumps(payload, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode() + ).hexdigest() + return f"nl2sql:{digest}" + + +async def cache_get(redis, key: str) -> dict[str, Any] | None: + """读取缓存,Redis 异常或内容损坏时返回空结果。""" + try: + value = await redis.get(key) + if not value: + return None + return json.loads(value) + except Exception: # noqa: BLE001 缓存故障必须降级直连数据库 + logger.warning("NL2SQL 缓存读取失败", exc_info=True) + return None + + +async def cache_set( + redis, + key: str, + value: dict[str, Any], + *, + ttl: int, + access_tables: set[str] | None = None, +) -> bool: + """写入缓存并建立访问表索引,Redis 异常时返回 False。""" + try: + await redis.set(key, json.dumps(value, ensure_ascii=False), ex=ttl) + for table_name in access_tables or set(): + await redis.sadd(f"nl2sql:table-keys:{table_name}", key) + except Exception: # noqa: BLE001 缓存故障不能阻断直连查询 + logger.warning("NL2SQL 缓存写入失败", exc_info=True) + return False + return True + + +async def invalidate_tables(redis, table_names: set[str]) -> int: + """删除指定表关联的全部缓存 key。""" + deleted = 0 + try: + for table_name in table_names: + index_key = f"nl2sql:table-keys:{table_name}" + keys = await redis.smembers(index_key) + if keys: + await redis.delete(*keys) + deleted += len(keys) + await redis.delete(index_key) + except Exception: # noqa: BLE001 缓存失效失败只记录并返回已处理数量 + logger.warning("NL2SQL 缓存失效失败", exc_info=True) + return deleted + + +def result_to_cache(result: DataQueryResult) -> dict[str, Any]: + """将统一查询结果转换为可安全序列化的缓存载荷。""" + return { + "query_id": result.query_id, + "trace_id": result.trace_id, + "columns": result.columns, + "rows": result.rows, + "row_count": result.row_count, + "truncated": result.truncated, + "summary": result.summary, + "markdown": result.markdown, + "chart": result.chart, + "metric_definitions": result.metric_definitions, + "query_plan": result.query_plan, + "sql": result.sql, + "elapsed_ms": result.elapsed_ms, + "warnings": result.warnings, + } + + +def result_from_cache(payload: dict[str, Any]) -> DataQueryResult: + """将缓存载荷恢复为统一查询结果。""" + return DataQueryResult(**payload) diff --git a/nl2sql/contracts.py b/nl2sql/contracts.py new file mode 100644 index 0000000..3e5c32a --- /dev/null +++ b/nl2sql/contracts.py @@ -0,0 +1,53 @@ +"""NL2SQL 服务与其他 Agent 共享的公共契约。""" +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + + +@dataclass(frozen=True) +class DataQueryRequest: + question: str + user_id: int + trace_id: str + session_id: str | None = None + caller_agent: str | None = None + data_scope: dict[str, Any] | None = None + max_rows: int | None = None + include_sql: bool = False + page: int = 1 + page_size: int = 100 + sort_by: str | None = None + sort_order: str = "asc" + + def __post_init__(self) -> None: + if not isinstance(self.question, str) or not self.question.strip(): + raise ValueError("question must not be empty") + if self.user_id <= 0: + raise ValueError("user_id must be positive") + if not self.trace_id.strip(): + raise ValueError("trace_id must not be empty") + if self.max_rows is not None and self.max_rows <= 0: + raise ValueError("max_rows must be positive") + if self.page < 1 or self.page_size < 1: + raise ValueError("page and page_size must be positive") + if self.sort_order not in {"asc", "desc"}: + raise ValueError("sort_order must be asc or desc") + + +@dataclass(frozen=True) +class DataQueryResult: + query_id: str + trace_id: str = "" + columns: list[str] = field(default_factory=list) + rows: list[dict[str, Any]] = field(default_factory=list) + row_count: int = 0 + truncated: bool = False + summary: str | None = None + markdown: str | None = None + chart: dict[str, Any] | None = None + metric_definitions: list[dict[str, Any]] = field(default_factory=list) + query_plan: dict[str, Any] | None = None + sql: str | None = None + elapsed_ms: float = 0.0 + warnings: list[str] = field(default_factory=list) diff --git a/nl2sql/diagnostics.py b/nl2sql/diagnostics.py new file mode 100644 index 0000000..6c34cbb --- /dev/null +++ b/nl2sql/diagnostics.py @@ -0,0 +1,50 @@ +"""NL2SQL 管理员诊断数据构造。""" +from __future__ import annotations + +from collections import deque +from threading import Lock + +def build_diagnostic_snapshot( + *, + query_id: str, + status: str, + access_tables: set[str] | list[str], + security_rule: str | None = None, + model_elapsed_ms: float | None = None, + sql: str | None = None, +) -> dict: + """构造不含原始 SQL 和敏感结果的诊断快照。""" + return { + "query_id": query_id, + "status": status, + "access_tables": sorted(set(access_tables)), + "security_rule": security_rule, + "model_elapsed_ms": model_elapsed_ms, + "has_sql": bool(sql), + } + + +class DiagnosticRegistry: + """保存最近的脱敏诊断摘要,进程重启后自动清空。""" + + def __init__(self, *, max_items: int = 100): + self._items = deque(maxlen=max_items) + self._lock = Lock() + + def record(self, **kwargs) -> None: + """记录诊断摘要并丢弃未定义的敏感字段。""" + allowed = { + "query_id", "status", "access_tables", "security_rule", + "model_elapsed_ms", "sql", + } + snapshot = build_diagnostic_snapshot(**{key: value for key, value in kwargs.items() if key in allowed}) + with self._lock: + self._items.appendleft(snapshot) + + def list_recent(self) -> list[dict]: + """返回最近诊断摘要的副本。""" + with self._lock: + return [dict(item) for item in self._items] + + +diagnostic_registry = DiagnosticRegistry() diff --git a/nl2sql/embedding.py b/nl2sql/embedding.py new file mode 100644 index 0000000..bc820a0 --- /dev/null +++ b/nl2sql/embedding.py @@ -0,0 +1,35 @@ +"""NL2SQL 向量化适配器,复用公共 LLM 客户端。""" +from __future__ import annotations + +import logging + +from config.settings import settings +from tool.llm import llm + +logger = logging.getLogger("nl2sql.embedding") +EMBEDDING_BATCH_SIZE = 10 + + +class EmbeddingError(RuntimeError): + """向量服务返回无效数据时抛出的异常。""" + + +async def embed_texts(texts: list[str], *, client=None) -> list[list[float]]: + if not texts: + return [] + provider = client or llm + dimension = settings.llm.embed_dimensions + vectors: list[list[float]] = [] + for start in range(0, len(texts), EMBEDDING_BATCH_SIZE): + batch = texts[start : start + EMBEDDING_BATCH_SIZE] + try: + batch_vectors = await provider.embed(batch) + except Exception as exc: # noqa: BLE001 向量服务异常统一转换 + logger.exception("NL2SQL embedding provider failed") + raise EmbeddingError("Embedding service unavailable") from exc + if len(batch_vectors) != len(batch) or any( + len(vector) != dimension for vector in batch_vectors + ): + raise EmbeddingError(f"Embedding dimension must be {dimension}") + vectors.extend(batch_vectors) + return vectors diff --git a/nl2sql/evaluation.py b/nl2sql/evaluation.py new file mode 100644 index 0000000..6ee81a7 --- /dev/null +++ b/nl2sql/evaluation.py @@ -0,0 +1,111 @@ +"""NL2SQL 离线 Golden Case 评测。""" +from __future__ import annotations + +import json +from pathlib import Path + +from sqlglot import parse_one + +from nl2sql.sql_security import SqlSecurityError, validate_select_sql + + +def load_cases(path: str | Path) -> list[dict]: + """读取结构化评测案例,不读取真实查询结果。""" + payload = json.loads(Path(path).read_text(encoding="utf-8")) + if not isinstance(payload, list): + raise ValueError("Golden Case 必须是数组") + return payload + + +def evaluate_case(case: dict) -> dict: + """评估单个案例的表召回和 SQL 安全结果。""" + expected = set(case.get("expected_tables", [])) + retrieved = set(case.get("retrieved_tables", [])) + recall = 1.0 if not expected else len(expected & retrieved) / len(expected) + result = { + "case_id": str(case.get("case_id", "unknown")), + "table_recall": recall, + "security_pass": False, + } + try: + validated = validate_select_sql( + case.get("sql", ""), + authorized_tables=set(case.get("authorized_tables", [])), + authorized_columns=case.get("authorized_columns"), + max_rows=int(case.get("max_rows", 1000)), + ) + result["security_pass"] = True + result["access_tables"] = sorted(validated.access_tables) + except SqlSecurityError as exc: + result["error_type"] = type(exc).__name__ + except (TypeError, ValueError, KeyError) as exc: + result["error_type"] = type(exc).__name__ + if "expected_sql" in case and "sql" in case: + result["sql_match"] = compare_sql(case["sql"], case["expected_sql"]) + if "expected_result" in case and "actual_result" in case: + result["result_match"] = compare_result_set(case["expected_result"], case["actual_result"]) + return result + + +def evaluate_cases(cases: list[dict]) -> list[dict]: + return [evaluate_case(case) for case in cases] + + +def summarize_evaluation(results: list[dict]) -> dict: + total = len(results) + passed = sum(1 for item in results if item.get("security_pass")) + recall = sum(float(item.get("table_recall", 0.0)) for item in results) + summary = { + "total": total, + "security_passed": passed, + "security_pass_rate": passed / total if total else 0.0, + "average_table_recall": recall / total if total else 0.0, + "failed_case_ids": [item["case_id"] for item in results if not item.get("security_pass")], + } + sql_results = [item["sql_match"] for item in results if "sql_match" in item] + result_results = [item["result_match"] for item in results if "result_match" in item] + if sql_results: + summary["sql_match_rate"] = sum(sql_results) / len(sql_results) + if result_results: + summary["result_match_rate"] = sum(result_results) / len(result_results) + return summary + + +def compare_sql(actual_sql: str, expected_sql: str) -> bool: + """使用 AST 规范化后比较 SQL,不连接数据库。""" + try: + return parse_one(actual_sql, read="mysql").sql(dialect="mysql") == parse_one( + expected_sql, read="mysql" + ).sql(dialect="mysql") + except (TypeError, ValueError): + return False + + +def compare_result_set(expected: dict, actual: dict) -> bool: + """比较假结果集的列名和行数据,不执行 SQL。""" + if not isinstance(expected, dict) or not isinstance(actual, dict): + return False + return ( + expected.get("columns", []) == actual.get("columns", []) + and expected.get("rows", []) == actual.get("rows", []) + ) + + +def build_evaluation_report( + cases: list[dict], + *, + prompt_version: str = "unknown", + semantic_version: str = "unknown", + model_version: str = "unknown", +) -> dict: + """构建带版本元数据的离线评测报告。""" + results = evaluate_cases(cases) + return { + "metadata": { + "prompt_version": prompt_version, + "semantic_version": semantic_version, + "model_version": model_version, + }, + "summary": summarize_evaluation(results), + "cases": results, + } diff --git a/nl2sql/executor.py b/nl2sql/executor.py new file mode 100644 index 0000000..e571588 --- /dev/null +++ b/nl2sql/executor.py @@ -0,0 +1,92 @@ +"""执行已通过安全校验的只读 SQL。""" +from __future__ import annotations + +import asyncio +import time + +from sqlalchemy import text + +from nl2sql.contracts import DataQueryResult +from nl2sql.masking import mask_rows +from nl2sql.rendering import render_markdown +from nl2sql.runtime import QueryRuntimeRegistry, query_runtime_registry +from nl2sql.runtime import kill_mysql_query +from nl2sql.sql_security import ValidatedSql + + +class QueryExecutionError(RuntimeError): + """只读 SQL 执行失败。""" + + +async def execute_readonly_sql( + session, + validated_sql: ValidatedSql, + *, + query_id: str, + trace_id: str, + masks: dict[tuple[str, str], str] | None = None, + max_rows: int | None = None, + timeout_seconds: float | None = None, + user_id: int = 0, + connection_id: int | None = None, + kill_query=None, + runtime_registry: QueryRuntimeRegistry | None = query_runtime_registry, +) -> DataQueryResult: + """执行安全 SQL 并转换为统一查询结果,不接收未校验的原始 SQL。""" + started_at = time.perf_counter() + if runtime_registry is not None: + if connection_id is None and hasattr(session, "connection"): + try: + connection = await session.connection() + connection_result = await connection.execute(text("SELECT CONNECTION_ID()")) + connection_id = int(connection_result.scalar_one()) + except Exception: # noqa: BLE001 连接号仅用于运维,不影响正常查询 + connection_id = None + await runtime_registry.register( + query_id=query_id, + user_id=user_id, + sql=validated_sql.sql, + connection_id=connection_id, + ) + try: + execution = session.execute(text(validated_sql.sql)) + result = ( + await asyncio.wait_for(execution, timeout_seconds) + if timeout_seconds is not None + else await execution + ) + columns = [str(key) for key in result.keys()] + rows = [dict(row) for row in result.mappings().all()] + if masks: + rows = mask_rows(rows, masks) + truncated = max_rows is not None and len(rows) > max_rows + if truncated: + rows = rows[:max_rows] + except asyncio.TimeoutError as exc: + killer = kill_query or kill_mysql_query + if connection_id is not None: + try: + await killer(connection_id) + except Exception: # noqa: BLE001 中止失败仍需返回统一超时错误 + pass + if runtime_registry is not None: + await runtime_registry.complete(query_id, status="timeout") + raise QueryExecutionError("查询超时") from exc + except Exception as exc: # noqa: BLE001 统一收敛数据库异常 + if runtime_registry is not None: + await runtime_registry.complete(query_id, status="failed") + raise QueryExecutionError("查询执行失败") from exc + elapsed_ms = (time.perf_counter() - started_at) * 1000 + if runtime_registry is not None: + await runtime_registry.complete(query_id) + return DataQueryResult( + query_id=query_id, + trace_id=trace_id, + columns=columns, + rows=rows, + row_count=len(rows), + truncated=bool(truncated), + markdown=render_markdown(columns, rows), + sql=validated_sql.sql, + elapsed_ms=elapsed_ms, + ) diff --git a/nl2sql/explain.py b/nl2sql/explain.py new file mode 100644 index 0000000..8d28a36 --- /dev/null +++ b/nl2sql/explain.py @@ -0,0 +1,11 @@ +"""NL2SQL 安全 EXPLAIN 辅助函数。""" +from __future__ import annotations + +from nl2sql.sql_security import ValidatedSql + + +def build_explain_sql(validated_sql: ValidatedSql) -> str: + """只为已通过安全校验的 SELECT 构造 EXPLAIN 语句。""" + if not isinstance(validated_sql, ValidatedSql): + raise TypeError("EXPLAIN 只接受已校验 SQL") + return f"EXPLAIN {validated_sql.sql}" diff --git a/nl2sql/few_shot.py b/nl2sql/few_shot.py new file mode 100644 index 0000000..a7ff492 --- /dev/null +++ b/nl2sql/few_shot.py @@ -0,0 +1,66 @@ +"""NL2SQL Few-shot 示例召回。""" +from __future__ import annotations + +import logging +from collections.abc import Awaitable, Callable +from typing import Any + +from nl2sql.embedding import embed_texts +from nl2sql.milvus_collections import NL2SQL_COLLECTION + + +logger = logging.getLogger("nl2sql.few_shot") + + +def _flatten(results: Any): + for batch in results or []: + if isinstance(batch, dict): + yield batch + else: + yield from batch or [] + + +async def retrieve_few_shot( + query: str, + milvus_client, + *, + embedder: Callable[[list[str]], Awaitable[list[list[float]]]] = embed_texts, + top_k: int = 3, + threshold: float = 0.75, +) -> list[dict[str, Any]]: + """召回已标记为有效的 Few-shot 示例,异常时返回空列表。""" + if not query or not query.strip() or top_k <= 0: + return [] + try: + vector = (await embedder([query]))[0] + hits = await milvus_client.search( + collection_name=NL2SQL_COLLECTION, + data=[vector], + limit=top_k, + filter='chunk_type == "few_shot_example" and is_valid == true and is_deprecated == false', + output_fields=["query", "correct_sql", "explanation", "case_id", "is_valid", "is_deprecated"], + ) + except Exception: # noqa: BLE001 Few-shot 失败不阻断主查询 + logger.warning("NL2SQL Few-shot 召回失败", exc_info=True) + return [] + + examples = [] + for hit in _flatten(hits): + entity = hit.get("entity") or hit + if not entity.get("is_valid", True) or entity.get("is_deprecated", False): + continue + distance = hit.get("distance") + score = 1.0 - distance if distance is not None else hit.get("score") + if score is None or score < threshold: + continue + example = { + "query": entity.get("query", ""), + "correct_sql": entity.get("correct_sql", ""), + "score": score, + } + if entity.get("explanation"): + example["explanation"] = entity["explanation"] + if entity.get("case_id"): + example["case_id"] = entity["case_id"] + examples.append(example) + return examples diff --git a/nl2sql/health.py b/nl2sql/health.py new file mode 100644 index 0000000..77216db --- /dev/null +++ b/nl2sql/health.py @@ -0,0 +1,29 @@ +"""NL2SQL 依赖健康检查。""" +from __future__ import annotations + +import asyncio +import time +from collections.abc import Awaitable, Callable + + +async def _check_one(checker: Callable[[], Awaitable[None]]) -> dict: + started = time.perf_counter() + try: + await checker() + except Exception as exc: # noqa: BLE001 对外只返回异常类型 + return { + "status": f"down: {type(exc).__name__}", + "ms": round((time.perf_counter() - started) * 1000, 1), + } + return {"status": "ok", "ms": round((time.perf_counter() - started) * 1000, 1)} + + +async def check_nl2sql_health(checkers: dict[str, Callable[[], Awaitable[None]]]) -> dict: + """并行检查依赖,返回不含异常详情的状态摘要。""" + names = ("mysql", "redis", "milvus", "llm") + values = await asyncio.gather( + *(_check_one(checkers[name]) for name in names if name in checkers) + ) + result = {name: value for name, value in zip((name for name in names if name in checkers), values)} + result["ready"] = bool(result) and all(item["status"] == "ok" for item in result.values()) + return result diff --git a/nl2sql/history.py b/nl2sql/history.py new file mode 100644 index 0000000..5319480 --- /dev/null +++ b/nl2sql/history.py @@ -0,0 +1,62 @@ +"""NL2SQL 查询历史归档适配器。""" +from __future__ import annotations + +from collections.abc import Iterable +import logging + +from model.nl2sql_permission import Nl2SqlQueryHistory + + +_STATUSES = {"success", "failed", "blocked", "timeout"} +logger = logging.getLogger("nl2sql.history") + + +async def archive_query( + db, + *, + query_id: str, + user_id: int, + question: str, + generated_sql: str | None = None, + access_tables: Iterable[str] = (), + status: str, + error_code: str | None = None, + error_message: str | None = None, + row_count: int = 0, + truncated: bool = False, + elapsed_ms: float | None = None, + trace_id: str | None = None, + session_id: str | None = None, + caller_agent: str | None = None, +) -> None: + """归档查询元数据,明确不写入结果行。""" + if status not in _STATUSES: + raise ValueError("无效的查询历史状态") + history = Nl2SqlQueryHistory( + query_id=query_id, + user_id=user_id, + session_id=session_id, + caller_agent=caller_agent, + question=question, + generated_sql=generated_sql, + access_tables=sorted(set(access_tables)), + status=status, + error_code=error_code, + error_message=error_message, + row_count=row_count, + truncated=truncated, + elapsed_ms=elapsed_ms, + trace_id=trace_id, + ) + db.add(history) + await db.commit() + + +async def archive_query_safely(db, **kwargs) -> bool: + """尝试归档查询,归档存储异常时记录日志并返回 False。""" + try: + await archive_query(db, **kwargs) + except Exception: # noqa: BLE001 归档失败不能阻断查询主流程 + logger.warning("NL2SQL 查询历史归档失败", exc_info=True) + return False + return True diff --git a/nl2sql/job_history.py b/nl2sql/job_history.py new file mode 100644 index 0000000..750d696 --- /dev/null +++ b/nl2sql/job_history.py @@ -0,0 +1,72 @@ +"""NL2SQL 运维任务执行历史适配器。""" +from __future__ import annotations + +import json +from datetime import datetime + +from sqlalchemy import text + +from nl2sql.audit import _sanitize + + +async def record_job_history(db, result, *, elapsed_ms: float, parameter_summary: dict | None = None) -> None: + """保存任务结果摘要,不保存原始参数和连接凭据。""" + statement = text( + """ + INSERT INTO nl2sql_job_history + (job_name, status, attempts, detail, error_type, parameter_summary, elapsed_ms, create_time) + VALUES + (:job_name, :status, :attempts, :detail, :error_type, :parameter_summary, :elapsed_ms, :create_time) + """ + ) + await db.execute( + statement, + { + "job_name": result.name, + "status": result.status, + "attempts": result.attempts, + "detail": json.dumps(_sanitize(result.detail or {}), ensure_ascii=False), + "error_type": result.error_type, + "parameter_summary": json.dumps(_sanitize(parameter_summary or {}), ensure_ascii=False), + "elapsed_ms": elapsed_ms, + "create_time": datetime.now(), + }, + ) + await db.commit() + + +async def record_job_history_safely(db, result, *, elapsed_ms: float, parameter_summary: dict | None = None) -> bool: + """任务历史写入失败时降级,不影响任务结果返回。""" + try: + await record_job_history( + db, + result, + elapsed_ms=elapsed_ms, + parameter_summary=parameter_summary, + ) + except Exception: # noqa: BLE001 历史故障不能阻断运维任务 + return False + return True + + +async def list_job_history(db, *, page: int = 1, page_size: int = 20, status: str | None = None) -> list[dict]: + """分页查询任务历史,只返回执行摘要,不返回 detail 明细。""" + conditions = "WHERE (:status IS NULL OR status = :status)" + statement = text( + f""" + SELECT id, job_name, status, attempts, error_type, elapsed_ms, create_time + FROM nl2sql_job_history + {conditions} + ORDER BY id DESC + LIMIT :limit OFFSET :offset + """ + ) + result = await db.execute( + statement, + { + "status": status, + "limit": max(1, min(page_size, 100)), + "offset": max(0, (page - 1) * page_size), + }, + ) + return [dict(row) for row in result.mappings().all()] diff --git a/nl2sql/jobs.py b/nl2sql/jobs.py new file mode 100644 index 0000000..27b83f2 --- /dev/null +++ b/nl2sql/jobs.py @@ -0,0 +1,124 @@ +"""NL2SQL 可由外部调度器调用的运维任务封装。""" +from __future__ import annotations + +import asyncio +import logging +import uuid +from dataclasses import dataclass +from datetime import datetime +from typing import Awaitable, Callable + +logger = logging.getLogger("nl2sql.jobs") + + +@dataclass(frozen=True) +class JobResult: + """运维任务的稳定结果结构。""" + + name: str + status: str + attempts: int + detail: dict | None = None + error_type: str | None = None + + +async def run_metadata_sync(*, redis=None, worker=None, retries: int = 2) -> JobResult: + """执行数据字典增量同步。""" + if worker is None: + from scripts.sync_nl2sql_metadata import synchronize + + worker = synchronize + return await run_job("metadata_sync", worker, redis=redis, retries=retries) + + +async def run_vector_cleanup(*, redis=None, worker=None, retries: int = 2) -> JobResult: + """执行失效向量清理。""" + if worker is None: + from config.database.milvus import client + from nl2sql.operations import cleanup_invalid_vectors + + worker = lambda: _cleanup_vectors(cleanup_invalid_vectors, client) + return await run_job("vector_cleanup", worker, redis=redis, retries=retries) + + +async def run_consistency_check(*, redis=None, worker, retries: int = 1, backoff: float = 0.5) -> JobResult: + """执行元数据与向量一致性巡检。""" + return await run_job( + "consistency_check", worker, redis=redis, retries=retries, backoff=backoff + ) + + +async def run_history_cleanup( + db, + before: datetime, + *, + redis=None, + repo_factory=None, + retries: int = 1, + backoff: float = 0.5, +) -> JobResult: + """删除保留期之前的查询历史。""" + if repo_factory is None: + from repositories.nl2sql_permission import Nl2SqlPermissionRepo + + repo_factory = Nl2SqlPermissionRepo + + async def worker(): + deleted = await repo_factory(db).delete_history_before(before) + return {"deleted": deleted} + + return await run_job( + "history_cleanup", worker, redis=redis, retries=retries, backoff=backoff + ) + + +async def _cleanup_vectors(cleanup, client_factory): + return {"deleted": await cleanup(client_factory())} + + +async def run_job( + name: str, + worker: Callable[[], Awaitable[dict] | dict], + *, + redis=None, + retries: int = 2, + backoff: float = 0.5, + lock_ttl: int = 900, +) -> JobResult: + """以 Redis 锁和有限重试执行一个幂等任务。""" + lock_key = f"nl2sql:job:{name}:lock" + token = uuid.uuid4().hex + locked = True + if redis is not None: + try: + locked = bool(await redis.set(lock_key, token, ex=lock_ttl, nx=True)) + except Exception: # noqa: BLE001 Redis 故障不应伪造锁成功 + logger.warning("NL2SQL 任务锁不可用,继续执行任务:%s", name) + if not locked: + return JobResult(name=name, status="skipped", attempts=0, detail={"reason": "lock_held"}) + + try: + for attempt in range(1, max(0, retries) + 2): + try: + result = worker() + if asyncio.iscoroutine(result): + result = await result + return JobResult(name=name, status="success", attempts=attempt, detail=result or {}) + except Exception as exc: # noqa: BLE001 任务失败返回类型化摘要 + if attempt > retries: + logger.warning("NL2SQL 任务失败:%s", name, exc_info=True) + return JobResult( + name=name, + status="failed", + attempts=attempt, + error_type=type(exc).__name__, + ) + if backoff > 0: + await asyncio.sleep(backoff * (2 ** (attempt - 1))) + finally: + if redis is not None and locked: + try: + await redis.delete(lock_key) + except Exception: # noqa: BLE001 释放锁失败仅记录日志 + logger.warning("NL2SQL 任务锁释放失败:%s", name, exc_info=True) + raise RuntimeError("NL2SQL 任务执行流程异常") diff --git a/nl2sql/limits.py b/nl2sql/limits.py new file mode 100644 index 0000000..16c9cb2 --- /dev/null +++ b/nl2sql/limits.py @@ -0,0 +1,58 @@ +"""NL2SQL 查询配额、并发和频率限制。""" +from __future__ import annotations + +import time + + +class QueryLimiter: + """基于 Redis 计数器的请求限制器,Redis 故障时允许降级放行。""" + + def __init__(self, redis, *, rate_window_seconds: int = 60): + self.redis = redis + self.rate_window_seconds = rate_window_seconds + + async def acquire( + self, + user_id: int, + *, + daily_quota: int = 0, + max_concurrent: int = 1, + rate_limit: int = 0, + ) -> bool: + """检查并占用一次配额、并发和频率额度,0 表示不限制。""" + try: + if daily_quota: + daily_key = f"nl2sql:daily:{user_id}:{time.strftime('%Y%m%d')}" + daily_count = await self.redis.incr(daily_key) + await self.redis.expire(daily_key, 86400) + if daily_count > daily_quota: + await self.redis.decr(daily_key) + return False + concurrent_key = f"nl2sql:concurrent:{user_id}" + concurrent_count = await self.redis.incr(concurrent_key) + await self.redis.expire(concurrent_key, 3600) + if max_concurrent and concurrent_count > max_concurrent: + await self.redis.decr(concurrent_key) + if daily_quota: + await self.redis.decr(daily_key) + return False + if rate_limit: + rate_key = f"nl2sql:rate:{user_id}:{int(time.time()) // self.rate_window_seconds}" + rate_count = await self.redis.incr(rate_key) + await self.redis.expire(rate_key, self.rate_window_seconds + 1) + if rate_count > rate_limit: + await self.redis.decr(rate_key) + await self.redis.decr(concurrent_key) + if daily_quota: + await self.redis.decr(daily_key) + return False + return True + except Exception: # noqa: BLE001 Redis 故障时降级为不限制 + return True + + async def release(self, user_id: int) -> None: + """释放一次并发额度,失败时静默降级。""" + try: + await self.redis.decr(f"nl2sql:concurrent:{user_id}") + except Exception: # noqa: BLE001 Redis 故障不能影响查询结果 + return diff --git a/nl2sql/masking.py b/nl2sql/masking.py new file mode 100644 index 0000000..8f322e3 --- /dev/null +++ b/nl2sql/masking.py @@ -0,0 +1,41 @@ +"""NL2SQL 查询结果脱敏。""" +from __future__ import annotations + +import hashlib +from typing import Any + + +def _mask_partial(value: Any) -> Any: + if value is None: + return None + text = str(value) + if len(text) == 1: + return "*" + if len(text) <= 4: + return f"{text[0]}{'*' * (len(text) - 2)}{text[-1]}" + return f"{text[:3]}{'*' * max(1, len(text) - 5)}{text[-2:]}" + + +def _mask_value(value: Any, mask_type: str) -> Any: + if value is None: + return None + if mask_type == "hash": + return hashlib.sha256(str(value).encode("utf-8")).hexdigest() + return _mask_partial(value) + + +def mask_rows( + rows: list[dict[str, Any]], + masks: dict[tuple[str, str], str], + *, + table_name: str | None = None, +) -> list[dict[str, Any]]: + """复制并脱敏结果行,不修改数据库查询返回的原始对象。""" + output = [dict(row) for row in rows] + for (table, column), mask_type in masks.items(): + if table_name is not None and table != table_name: + continue + for row in output: + if column in row: + row[column] = _mask_value(row[column], mask_type) + return output diff --git a/nl2sql/metadata.py b/nl2sql/metadata.py new file mode 100644 index 0000000..c43c72c --- /dev/null +++ b/nl2sql/metadata.py @@ -0,0 +1,113 @@ +"""将 MySQL information_schema 结果标准化为 NL2SQL 元数据 chunk。""" +from __future__ import annotations + +import hashlib +import json +from typing import Any + + +def _value(row: dict[str, Any], name: str, default: Any = None) -> Any: + if name in row: + return row[name] + upper = name.upper() + lower = name.lower() + return row.get(upper, row.get(lower, default)) + + +def normalize_table_row(row: dict[str, Any]) -> dict[str, Any] | None: + table_name = str(_value(row, "table_name", "") or "").strip() + table_type = str(_value(row, "table_type", "BASE TABLE") or "").upper() + if not table_name or table_type != "BASE TABLE": + return None + return { + "table_name": table_name, + "table_comment": str(_value(row, "table_comment", "") or "").strip(), + "is_valid": True, + } + + +def normalize_column_row(row: dict[str, Any]) -> dict[str, Any]: + nullable = str(_value(row, "is_nullable", "NO") or "NO").upper() == "YES" + return { + "table_name": str(_value(row, "table_name", "") or "").strip(), + "field_name": str( + _value(row, "column_name", _value(row, "field_name", "")) or "" + ).strip(), + "column_comment": str(_value(row, "column_comment", "") or "").strip(), + "data_type": str(_value(row, "data_type", "") or "").strip(), + "is_nullable": nullable, + "ordinal_position": int(_value(row, "ordinal_position", 0) or 0), + } + + +def _hash_text(text: str) -> str: + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + +def _chunk_id(chunk_type: str, table_name: str, field_name: str = "") -> str: + raw = f"{chunk_type}:{table_name}:{field_name}" + return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:64] + + +def build_metadata_chunks( + tables: list[dict[str, Any]], columns: list[dict[str, Any]] +) -> list[dict[str, Any]]: + valid_tables = sorted( + (table for table in (normalize_table_row(row) for row in tables) if table), + key=lambda item: item["table_name"], + ) + valid_table_names = {table["table_name"] for table in valid_tables} + normalized_columns = sorted( + ( + column + for column in (normalize_column_row(row) for row in columns) + if column["table_name"] in valid_table_names and column["field_name"] + ), + key=lambda item: (item["table_name"], item["ordinal_position"], item["field_name"]), + ) + + chunks: list[dict[str, Any]] = [] + for table in valid_tables: + text = f"表名:{table['table_name']}\n表说明:{table['table_comment']}" + chunks.append( + { + "id": _chunk_id("table_meta", table["table_name"]), + "chunk_type": "table_meta", + "table_name": table["table_name"], + "field_name": "", + "text": text, + "content_hash": _hash_text(text), + "is_valid": table["is_valid"], + "is_deprecated": False, + } + ) + + for column in normalized_columns: + text = ( + f"表名:{column['table_name']}\n" + f"字段名:{column['field_name']}\n" + f"字段说明:{column['column_comment']}\n" + f"字段类型:{column['data_type']}\n" + f"允许为空:{'是' if column['is_nullable'] else '否'}" + ) + chunks.append( + { + "id": _chunk_id( + "field_meta", column["table_name"], column["field_name"] + ), + "chunk_type": "field_meta", + "table_name": column["table_name"], + "field_name": column["field_name"], + "text": text, + "content_hash": _hash_text(text), + "is_valid": True, + "is_deprecated": False, + } + ) + return chunks + + +def metadata_signature(chunk: dict[str, Any]) -> str: + """返回用于比较 chunk 内容的稳定签名字符串。""" + payload = {key: chunk[key] for key in ("chunk_type", "table_name", "field_name", "text")} + return json.dumps(payload, ensure_ascii=False, sort_keys=True, separators=(",", ":")) diff --git a/nl2sql/metadata_sync.py b/nl2sql/metadata_sync.py new file mode 100644 index 0000000..0c3ed9c --- /dev/null +++ b/nl2sql/metadata_sync.py @@ -0,0 +1,83 @@ +"""构造并 upsert NL2SQL 元数据行。""" +from __future__ import annotations + +import time +from collections.abc import Awaitable, Callable +from typing import Any + +from nl2sql.embedding import embed_texts +from nl2sql.metadata import build_metadata_chunks +from nl2sql.milvus_collections import NL2SQL_COLLECTION + + +def _escape_filter_value(value: str) -> str: + return value.replace("\\", "\\\\").replace('"', '\\"') + + +async def prepare_metadata_rows( + chunks: list[dict[str, Any]], + *, + embedder: Callable[[list[str]], Awaitable[list[list[float]]]] = embed_texts, + timestamp: int | None = None, +) -> list[dict[str, Any]]: + if not chunks: + return [] + vectors = await embedder([chunk["text"] for chunk in chunks]) + if len(vectors) != len(chunks): + raise ValueError("Embedding result count must match metadata chunk count") + created_at = int(time.time()) if timestamp is None else timestamp + return [ + {**chunk, "vector": vector, "created_at": created_at} + for chunk, vector in zip(chunks, vectors) + ] + + +async def sync_metadata( + milvus_client, + table_rows: list[dict[str, Any]], + column_rows: list[dict[str, Any]], + *, + embedder: Callable[[list[str]], Awaitable[list[list[float]]]] = embed_texts, + timestamp: int | None = None, +) -> int: + chunks = build_metadata_chunks(table_rows, column_rows) + chunks_by_table: dict[str, list[dict[str, Any]]] = {} + for chunk in chunks: + chunks_by_table.setdefault(chunk["table_name"], []).append(chunk) + + updated_count = 0 + for table_name, table_chunks in chunks_by_table.items(): + escaped_name = _escape_filter_value(table_name) + existing_rows = await milvus_client.query( + collection_name=NL2SQL_COLLECTION, + filter=f'table_name == "{escaped_name}"', + output_fields=["id", "content_hash", "is_valid", "is_deprecated"], + ) + desired_signature = { + (chunk["id"], chunk["content_hash"], chunk["is_valid"], chunk["is_deprecated"]) + for chunk in table_chunks + } + existing_signature = { + ( + row.get("id"), + row.get("content_hash"), + row.get("is_valid", True), + row.get("is_deprecated", False), + ) + for row in existing_rows or [] + } + if desired_signature == existing_signature: + continue + + if existing_rows: + await milvus_client.delete( + collection_name=NL2SQL_COLLECTION, + filter=f'table_name == "{escaped_name}"', + ) + rows = await prepare_metadata_rows( + table_chunks, embedder=embedder, timestamp=timestamp + ) + if rows: + await milvus_client.upsert(collection_name=NL2SQL_COLLECTION, data=rows) + updated_count += len(rows) + return updated_count diff --git a/nl2sql/metrics.py b/nl2sql/metrics.py new file mode 100644 index 0000000..15df109 --- /dev/null +++ b/nl2sql/metrics.py @@ -0,0 +1,158 @@ +"""NL2SQL 查询运行指标,供管理员诊断和后续监控采集。""" +from __future__ import annotations + +from collections import Counter +from threading import Lock +import logging +from datetime import datetime + +from sqlalchemy import text + + +logger = logging.getLogger("nl2sql.metrics") + +HISTORY_METRICS_SQL = text( + """ + SELECT + COUNT(*) AS total, + COALESCE(SUM(status = 'success'), 0) AS success, + COALESCE(SUM(status = 'failed'), 0) AS failed, + COALESCE(SUM(status = 'timeout'), 0) AS timeout, + COALESCE(AVG(elapsed_ms), 0) AS average_elapsed_ms + FROM nl2sql_query_history + WHERE (:since IS NULL OR create_time >= :since) + """ +) + + +class QueryMetrics: + """进程内查询计数器,重启后清零。""" + + def __init__(self, *, slow_threshold_ms: float = 500): + self.slow_threshold_ms = slow_threshold_ms + self._lock = Lock() + self._total = 0 + self._success = 0 + self._failed = 0 + self._cache_hits = 0 + self._timeouts = 0 + self._slow = 0 + self._rate_limited = 0 + self._elapsed_total = 0.0 + self._failure_reasons: Counter[str] = Counter() + + def record( + self, + *, + status: str, + elapsed_ms: float | None = None, + cache_hit: bool = False, + failure_reason: str | None = None, + ) -> None: + with self._lock: + self._total += 1 + if status == "success": + self._success += 1 + else: + self._failed += 1 + if failure_reason: + self._failure_reasons[failure_reason] += 1 + if status == "timeout": + self._timeouts += 1 + if cache_hit: + self._cache_hits += 1 + if elapsed_ms is not None and elapsed_ms >= self.slow_threshold_ms: + self._slow += 1 + if elapsed_ms is not None: + self._elapsed_total += elapsed_ms + + def record_rate_limited(self) -> None: + """记录一次因配额、并发或频率限制而拒绝的请求。""" + with self._lock: + self._rate_limited += 1 + + def snapshot(self) -> dict[str, float | int]: + with self._lock: + total = self._total + return { + "total": total, + "success": self._success, + "failed": self._failed, + "success_rate": self._success / total if total else 0.0, + "failure_rate": self._failed / total if total else 0.0, + "cache_hits": self._cache_hits, + "cache_hit_rate": self._cache_hits / total if total else 0.0, + "timeout_count": self._timeouts, + "slow_query_count": self._slow, + "rate_limited_count": self._rate_limited, + "average_elapsed_ms": self._elapsed_total / total if total else 0.0, + "failure_reasons": dict(sorted(self._failure_reasons.items())), + } + + +query_metrics = QueryMetrics() + + +def render_prometheus(metrics: dict | None = None) -> str: + """将聚合指标导出为无外部依赖的 Prometheus 文本格式。""" + snapshot = metrics or query_metrics.snapshot() + lines = [ + "# HELP nl2sql_queries_total NL2SQL 查询总数", + "# TYPE nl2sql_queries_total counter", + f"nl2sql_queries_total {snapshot.get('total', 0)}", + "# HELP nl2sql_queries_success_total NL2SQL 成功查询数", + "# TYPE nl2sql_queries_success_total counter", + f"nl2sql_queries_success_total {snapshot.get('success', 0)}", + "# HELP nl2sql_queries_failed_total NL2SQL 失败查询数", + "# TYPE nl2sql_queries_failed_total counter", + f"nl2sql_queries_failed_total {snapshot.get('failed', 0)}", + "# HELP nl2sql_queries_timeout_total NL2SQL 超时查询数", + "# TYPE nl2sql_queries_timeout_total counter", + f"nl2sql_queries_timeout_total {snapshot.get('timeout_count', 0)}", + "# HELP nl2sql_queries_cache_hits_total NL2SQL 缓存命中数", + "# TYPE nl2sql_queries_cache_hits_total counter", + f"nl2sql_queries_cache_hits_total {snapshot.get('cache_hits', 0)}", + "# HELP nl2sql_queries_rate_limited_total NL2SQL 限流拒绝数", + "# TYPE nl2sql_queries_rate_limited_total counter", + f"nl2sql_queries_rate_limited_total {snapshot.get('rate_limited_count', 0)}", + "# HELP nl2sql_queries_slow_total NL2SQL 慢查询数", + "# TYPE nl2sql_queries_slow_total counter", + f"nl2sql_queries_slow_total {snapshot.get('slow_query_count', 0)}", + "# HELP nl2sql_queries_average_elapsed_ms NL2SQL 平均耗时毫秒", + "# TYPE nl2sql_queries_average_elapsed_ms gauge", + f"nl2sql_queries_average_elapsed_ms {snapshot.get('average_elapsed_ms', 0.0)}", + ] + for reason, count in sorted((snapshot.get("failure_reasons") or {}).items()): + safe_reason = str(reason).replace("\\", "\\\\").replace('"', '\\"').replace("\n", "\\n") + lines.append(f'nl2sql_query_failures_total{{reason="{safe_reason}"}} {count}') + return "\n".join(lines) + "\n" + + +async def load_history_metrics(db, *, since: datetime | None = None) -> dict[str, int | float]: + """从查询归档表聚合指标,不读取问题、SQL 或结果行。""" + result = await db.execute(HISTORY_METRICS_SQL, {"since": since}) + row = result.mappings().one() + return { + "total": int(row.get("total", 0) or 0), + "success": int(row.get("success", 0) or 0), + "failed": int(row.get("failed", 0) or 0), + "timeout": int(row.get("timeout", 0) or 0), + "average_elapsed_ms": float(row.get("average_elapsed_ms", 0) or 0), + } + + +async def load_history_metrics_safely(db, *, since: datetime | None = None) -> dict[str, int | float | bool]: + """历史指标读取失败时返回明确的不可用状态。""" + try: + result = await load_history_metrics(db, since=since) + return {**result, "available": True} + except Exception: # noqa: BLE001 指标故障不能阻断管理接口 + logger.warning("NL2SQL 历史指标读取失败", exc_info=True) + return { + "total": 0, + "success": 0, + "failed": 0, + "timeout": 0, + "average_elapsed_ms": 0.0, + "available": False, + } diff --git a/nl2sql/milvus_collections.py b/nl2sql/milvus_collections.py new file mode 100644 index 0000000..018c296 --- /dev/null +++ b/nl2sql/milvus_collections.py @@ -0,0 +1,86 @@ +"""NL2SQL 元数据向量的 Milvus 数据库和集合初始化。""" +from __future__ import annotations + +from pymilvus import AsyncMilvusClient, DataType + +from config.settings import settings + +NL2SQL_DATABASE_NAME = settings.milvus.db_name or "mutual_fund" +NL2SQL_COLLECTION = "data_agent_chunks" + + +def build_nl2sql_schema(): + schema = AsyncMilvusClient.create_schema(auto_id=False, enable_dynamic_field=False) + schema.add_field("id", DataType.VARCHAR, is_primary=True, max_length=128) + schema.add_field("vector", DataType.FLOAT_VECTOR, dim=settings.llm.embed_dimensions) + schema.add_field("chunk_type", DataType.VARCHAR, max_length=32) + schema.add_field("table_name", DataType.VARCHAR, max_length=128) + schema.add_field("field_name", DataType.VARCHAR, max_length=128) + schema.add_field("text", DataType.VARCHAR, max_length=65535) + schema.add_field("is_valid", DataType.BOOL) + schema.add_field("is_deprecated", DataType.BOOL) + schema.add_field("content_hash", DataType.VARCHAR, max_length=128) + schema.add_field("created_at", DataType.INT64) + # Few-shot 字段允许元数据 chunk 缺省,避免影响现有表/字段向量写入。 + schema.add_field("query", DataType.VARCHAR, max_length=4096, nullable=True) + schema.add_field("correct_sql", DataType.VARCHAR, max_length=16384, nullable=True) + schema.add_field("explanation", DataType.VARCHAR, max_length=65535, nullable=True) + schema.add_field("case_id", DataType.VARCHAR, max_length=128, nullable=True) + return schema + + +def build_nl2sql_index_params(): + params = AsyncMilvusClient.prepare_index_params() + params.add_index( + field_name="vector", + index_type="HNSW", + metric_type="COSINE", + params={"M": 16, "efConstruction": 200}, + ) + return params + + +async def ensure_nl2sql_database(milvus_client: AsyncMilvusClient | None = None) -> None: + if milvus_client is None: + from config.database.milvus import client as configured_milvus_client + + milvus_client = configured_milvus_client() + client = milvus_client + databases = await client.list_databases() + if NL2SQL_DATABASE_NAME not in databases: + await client.create_database(NL2SQL_DATABASE_NAME) + + +async def ensure_nl2sql_collection(milvus_client: AsyncMilvusClient | None = None) -> None: + if milvus_client is None: + from config.database.milvus import client as configured_milvus_client + + milvus_client = configured_milvus_client() + client = milvus_client + await ensure_nl2sql_database(client) + if await client.has_collection(NL2SQL_COLLECTION): + return + await client.create_collection( + collection_name=NL2SQL_COLLECTION, + schema=build_nl2sql_schema(), + index_params=build_nl2sql_index_params(), + ) + + +async def recreate_nl2sql_collection( + milvus_client: AsyncMilvusClient | None = None, +) -> None: + """删除并按当前 Embedding 维度重新创建 NL2SQL 集合。""" + if milvus_client is None: + from config.database.milvus import client as configured_milvus_client + + milvus_client = configured_milvus_client() + client = milvus_client + await ensure_nl2sql_database(client) + if await client.has_collection(NL2SQL_COLLECTION): + await client.drop_collection(collection_name=NL2SQL_COLLECTION) + await client.create_collection( + collection_name=NL2SQL_COLLECTION, + schema=build_nl2sql_schema(), + index_params=build_nl2sql_index_params(), + ) diff --git a/nl2sql/operations.py b/nl2sql/operations.py new file mode 100644 index 0000000..02877e3 --- /dev/null +++ b/nl2sql/operations.py @@ -0,0 +1,23 @@ +"""NL2SQL Milvus 运维操作。""" +from __future__ import annotations + +from nl2sql.milvus_collections import NL2SQL_COLLECTION + + +INVALID_VECTOR_FILTER = "is_valid == false or is_deprecated == true" + + +async def cleanup_invalid_vectors(milvus_client) -> int: + """删除无效或已废弃的元数据向量,并返回删除前的数量。""" + rows = await milvus_client.query( + collection_name=NL2SQL_COLLECTION, + filter=INVALID_VECTOR_FILTER, + output_fields=["id"], + ) + count = len(rows or []) + if count: + await milvus_client.delete( + collection_name=NL2SQL_COLLECTION, + filter=INVALID_VECTOR_FILTER, + ) + return count diff --git a/nl2sql/permission.py b/nl2sql/permission.py new file mode 100644 index 0000000..2f983e3 --- /dev/null +++ b/nl2sql/permission.py @@ -0,0 +1,78 @@ +"""不依赖外部组件的查询角色映射辅助函数。""" +from __future__ import annotations + + +_ROLE_MAP = { + "ADMIN": "admin", + "KNOWLEDGE_ADMIN": "knowledge_admin", + "KNOWLEDGE_OPERATOR": "knowledge_operator", + "投顾": "advisor", + "理财顾问": "advisor", + "风控专员": "risk", + "客户经理": "customer_manager", +} + + +def map_employee_role(employee_role: str | None) -> str | None: + if not employee_role: + return None + return _ROLE_MAP.get(employee_role, employee_role.lower()) + + +def build_query_permission( + user, + role, + table_permissions, + column_permissions, + sensitive_fields, +) -> dict: + """将用户、角色、表权限和字段权限合并成一次请求的权限快照。""" + permission = { + "can_query": False, + "role": map_employee_role(getattr(user, "employee_role", None)), + "tables": set(), + "columns": {}, + "masks": {}, + "row_scopes": {}, + "max_rows": getattr(role, "max_rows", 0) if role else 0, + "daily_quota": getattr(role, "daily_quota", 0) if role else 0, + } + if getattr(user, "user_type", None) != "EMPLOYEE" or not role: + return permission + if not getattr(role, "can_query", False): + return permission + + permission["can_query"] = True + for item in table_permissions or []: + if getattr(item, "permission", "") == "SELECT" and getattr(item, "table_name", None): + permission["tables"].add(item.table_name) + scope_type = getattr(item, "row_scope_type", "none") + scope_column = getattr(item, "row_scope_column", None) + if scope_type != "none" and scope_column: + permission["row_scopes"][item.table_name] = { + "type": scope_type, + "column": scope_column, + } + + for item in column_permissions or []: + table_name = getattr(item, "table_name", None) + column_name = getattr(item, "column_name", None) + if table_name not in permission["tables"] or not column_name: + continue + if getattr(item, "access_mode", "allow") in {"allow", "mask"}: + permission["columns"].setdefault(table_name, set()).add(column_name) + if getattr(item, "access_mode", "allow") == "mask": + permission["masks"][(table_name, column_name)] = getattr( + item, "mask_type", None + ) or "partial" + elif getattr(item, "access_mode", "allow") == "deny": + permission["columns"].setdefault(table_name, set()).discard(column_name) + + for item in sensitive_fields or []: + table_name = getattr(item, "table_name", None) + column_name = getattr(item, "column_name", None) + if table_name in permission["tables"] and column_name in permission["columns"].get(table_name, set()): + permission["masks"][(table_name, column_name)] = getattr( + item, "mask_type", None + ) or "partial" + return permission diff --git a/nl2sql/query_experience.py b/nl2sql/query_experience.py new file mode 100644 index 0000000..1cf0c0f --- /dev/null +++ b/nl2sql/query_experience.py @@ -0,0 +1,77 @@ +"""NL2SQL 查询体验辅助能力。""" +from __future__ import annotations + +import csv +import re +from io import StringIO + +from sqlglot import exp, parse_one + +from nl2sql.semantics import resolve_semantics + +_IDENTIFIER = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") +_TIME_MARKERS = ("近", "本年", "今年", "去年", "上月", "本月", "季度", "截至", "至今") +_METRIC_MARKERS = ("收益率", "最大回撤", "回撤", "收益", "波动率", "夏普") + + +def build_clarification(question: str) -> dict | None: + """识别缺少统计时间范围的指标问题,返回结构化反问。""" + text = (question or "").strip() + if any(marker in text for marker in _METRIC_MARKERS) and not any( + marker in text for marker in _TIME_MARKERS + ): + return { + "required": True, + "missing": ["time_range"], + "message": "请补充收益率的统计时间范围。", + } + return None + + +def apply_query_options( + sql: str, + *, + page: int = 1, + page_size: int = 100, + sort_by: str | None = None, + sort_order: str = "asc", +) -> str: + """通过 SQL AST 添加分页和白名单标识符排序。""" + if page < 1 or page_size < 1: + raise ValueError("分页参数必须为正数") + if sort_order not in {"asc", "desc"}: + raise ValueError("排序方向无效") + if sort_by is not None and not _IDENTIFIER.fullmatch(sort_by): + raise ValueError("排序字段无效") + statement = parse_one(sql, read="mysql") + if sort_by: + statement = statement.order_by( + exp.Ordered(this=exp.column(sort_by), desc=sort_order == "desc") + ) + statement = statement.limit(page_size) + if page > 1: + statement = statement.offset((page - 1) * page_size) + return statement.sql(dialect="mysql") + + +def render_csv(columns: list[str], rows: list[dict]) -> str: + """将已脱敏的查询结果渲染为 CSV。""" + output = StringIO() + writer = csv.DictWriter(output, fieldnames=columns) + writer.writeheader() + writer.writerows({column: row.get(column) for column in columns} for row in rows) + return output.getvalue() + + +def build_query_explanation(question: str, sql: str) -> dict: + """返回指标口径和不含 SQL 文本的查询计划摘要。""" + statement = parse_one(sql, read="mysql") + tables = sorted({table.name for table in statement.find_all(exp.Table)}) + return { + "metrics": resolve_semantics(question).get("metrics", []), + "plan": { + "tables": tables, + "operation": "SELECT", + "join_count": max(0, len(tables) - 1), + }, + } diff --git a/nl2sql/rendering.py b/nl2sql/rendering.py new file mode 100644 index 0000000..b18e074 --- /dev/null +++ b/nl2sql/rendering.py @@ -0,0 +1,24 @@ +"""NL2SQL 查询结果的基础 Markdown 渲染。""" +from __future__ import annotations + +from typing import Any + + +def _cell(value: Any) -> str: + """将单元格转换为不会破坏 Markdown 表格的文本。""" + if value is None: + return "" + return str(value).replace("|", "\\|").replace("\r", " ").replace("\n", " ") + + +def render_markdown(columns: list[str], rows: list[dict[str, Any]]) -> str: + """将列和行渲染成 Markdown 表格,空结果返回固定提示。""" + if not rows: + return "暂无数据。" + header = "| " + " | ".join(_cell(column) for column in columns) + " |" + divider = "| " + " | ".join("---" for _ in columns) + " |" + body = [ + "| " + " | ".join(_cell(row.get(column)) for column in columns) + " |" + for row in rows + ] + return "\n".join([header, divider, *body]) diff --git a/nl2sql/result.py b/nl2sql/result.py new file mode 100644 index 0000000..195fc41 --- /dev/null +++ b/nl2sql/result.py @@ -0,0 +1,52 @@ +"""NL2SQL 结果摘要和基础图表配置。""" +from __future__ import annotations + +import json +from typing import Any + + +def build_chart_config(columns: list[str], rows: list[dict[str, Any]]) -> dict[str, Any] | None: + """识别简单的分类数值结果,返回安全的柱状图配置。""" + if not rows or len(columns) < 2: + return None + numeric_column = next( + ( + column + for column in columns + if all(isinstance(row.get(column), (int, float)) and not isinstance(row.get(column), bool) for row in rows) + ), + None, + ) + category_column = next((column for column in columns if column != numeric_column), None) + if numeric_column is None or category_column is None: + return None + return {"type": "bar", "category": category_column, "value": numeric_column} + + +async def summarize_result( + question: str, + columns: list[str], + rows: list[dict[str, Any]], + *, + llm_client, + max_prompt_rows: int = 20, +) -> str: + """使用脱敏后的受控结果生成摘要,模型失败时返回固定话术。""" + fallback = f"查询完成,共返回 {len(rows)} 条记录。" + prompt = ( + "请用简洁中文总结查询结果,只基于提供的数据,不要猜测。\n" + f"问题:{question}\n" + f"列名:{json.dumps(columns, ensure_ascii=False)}\n" + f"结果:{json.dumps(rows[:max_prompt_rows], ensure_ascii=False, default=str)}" + ) + try: + answer = await llm_client.chat( + [ + {"role": "system", "content": "你是数据查询结果摘要助手。"}, + {"role": "user", "content": prompt}, + ], + temperature=0, + ) + except Exception: # noqa: BLE001 摘要失败回退固定文本 + return fallback + return answer.strip() or fallback diff --git a/nl2sql/retrieval.py b/nl2sql/retrieval.py new file mode 100644 index 0000000..3d81a1b --- /dev/null +++ b/nl2sql/retrieval.py @@ -0,0 +1,92 @@ +"""NL2SQL 元数据召回与授权过滤。""" +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Iterable +from typing import Any + +from nl2sql.embedding import embed_texts +from nl2sql.milvus_collections import NL2SQL_COLLECTION + + +def _flatten_hits(results: Any) -> Iterable[dict[str, Any]]: + """兼容 Milvus 返回的批次列表和单条字典结构。""" + for batch in results or []: + if isinstance(batch, dict): + yield batch + else: + yield from batch or [] + + +def filter_authorized_metadata( + tables: list[dict[str, Any]], + columns: list[dict[str, Any]], + permission: dict[str, Any], +) -> dict[str, list[dict[str, Any]]]: + """过滤失效实体、无表权限实体和无字段权限实体。""" + authorized_tables = set(permission.get("tables", set())) + authorized_columns = permission.get("columns", {}) + valid_table_names = { + table.get("table_name") + for table in tables + if table.get("is_valid", True) and table.get("table_name") in authorized_tables + } + filtered_tables = [ + table + for table in tables + if table.get("is_valid", True) + and table.get("table_name") in valid_table_names + ] + filtered_columns = [ + column + for column in columns + if column.get("is_valid", True) + and column.get("table_name") in valid_table_names + and column.get("field_name") + in set(authorized_columns.get(column.get("table_name"), set())) + ] + return {"tables": filtered_tables, "columns": filtered_columns} + + +async def retrieve_metadata( + query: str, + milvus_client, + *, + embedder: Callable[[list[str]], Awaitable[list[list[float]]]] = embed_texts, + top_k: int = 5, +) -> list[dict[str, Any]]: + """分别召回表级和字段级元数据,并过滤无效向量。""" + if not query or not query.strip(): + return [] + if top_k <= 0: + return [] + + vector = (await embedder([query]))[0] + output_fields = [ + "id", + "chunk_type", + "table_name", + "field_name", + "text", + "is_valid", + "is_deprecated", + ] + results: list[dict[str, Any]] = [] + for chunk_type in ("table_meta", "field_meta"): + hits = await milvus_client.search( + collection_name=NL2SQL_COLLECTION, + data=[vector], + limit=top_k, + filter=f'chunk_type == "{chunk_type}" and is_valid == true', + output_fields=output_fields, + ) + for hit in _flatten_hits(hits): + entity = hit.get("entity") or hit + if not entity.get("is_valid", True) or entity.get("is_deprecated", False): + continue + results.append( + { + **entity, + "distance": hit.get("distance", hit.get("score")), + } + ) + return results diff --git a/nl2sql/row_scope.py b/nl2sql/row_scope.py new file mode 100644 index 0000000..fd8ade3 --- /dev/null +++ b/nl2sql/row_scope.py @@ -0,0 +1,42 @@ +"""NL2SQL 行级权限条件构造与注入。""" +from __future__ import annotations + +from sqlglot import exp, parse_one +from sqlglot.errors import ParseError + + +class RowScopeError(ValueError): + """行级权限范围缺失或格式不合法。""" + + +def _values(data_scope: dict, scope_type: str) -> list[int]: + values = data_scope.get(scope_type) + if not isinstance(values, list) or not values or any( + isinstance(value, bool) or not isinstance(value, int) for value in values + ): + raise RowScopeError(f"缺少有效的 {scope_type} 行范围") + return sorted(set(values)) + + +def apply_row_scope(sql: str, permission: dict, data_scope: dict | None) -> str: + """按照服务端权限快照为每个受限表注入 IN 条件。""" + scopes = permission.get("row_scopes", {}) + if not scopes: + return sql + if not isinstance(data_scope, dict): + raise RowScopeError("缺少行范围参数") + try: + statement = parse_one(sql, read="mysql") + except ParseError as exc: + raise RowScopeError("SQL 解析失败") from exc + for table in statement.find_all(exp.Table): + scope = scopes.get(table.name) + if not scope: + continue + values = _values(data_scope, scope["type"]) + condition = exp.In( + this=exp.column(scope["column"]), + expressions=[exp.Literal.number(value) for value in values], + ) + statement = statement.where(condition) + return statement.sql(dialect="mysql") diff --git a/nl2sql/runtime.py b/nl2sql/runtime.py new file mode 100644 index 0000000..9e44896 --- /dev/null +++ b/nl2sql/runtime.py @@ -0,0 +1,92 @@ +"""NL2SQL 运行中查询登记和管理员中止状态。""" +from __future__ import annotations + +from dataclasses import asdict, dataclass +from datetime import datetime, timezone +from threading import RLock + +from sqlalchemy import text + + +@dataclass +class QueryRuntime: + """单条运行中查询的最小运行信息。""" + + query_id: str + user_id: int + sql: str + connection_id: int | None + started_at: datetime + status: str = "running" + + def to_dict(self) -> dict: + data = asdict(self) + data["started_at"] = self.started_at.isoformat() + return data + + +class QueryRuntimeRegistry: + """进程内运行查询登记表,跨进程场景由网关保证路由到同一实例。""" + + def __init__(self): + self._items: dict[str, QueryRuntime] = {} + self._lock = RLock() + + async def register( + self, + *, + query_id: str, + user_id: int, + sql: str, + connection_id: int | None = None, + ) -> QueryRuntime: + item = QueryRuntime( + query_id=query_id, + user_id=user_id, + sql=sql, + connection_id=connection_id, + started_at=datetime.now(timezone.utc), + ) + with self._lock: + self._items[query_id] = item + return item + + def get(self, query_id: str) -> QueryRuntime | None: + with self._lock: + return self._items.get(query_id) + + def list_active(self) -> list[QueryRuntime]: + with self._lock: + return [item for item in self._items.values() if item.status == "running"] + + async def complete(self, query_id: str, *, status: str = "completed") -> bool: + with self._lock: + item = self._items.get(query_id) + if item is None: + return False + item.status = status + self._items.pop(query_id, None) + return True + + async def mark_killed(self, query_id: str) -> bool: + with self._lock: + item = self._items.get(query_id) + if item is None or item.status != "running": + return False + item.status = "killed" + return True + + +async def kill_mysql_query(connection_id: int) -> bool: + """通过独立连接执行 KILL QUERY,避免占用被中止的连接。""" + from config.database.mysql import get_session_factory + + if int(connection_id) <= 0: + return False + async with get_session_factory()() as session: + await session.execute(text(f"KILL QUERY {int(connection_id)}")) + await session.commit() + return True + + +query_runtime_registry = QueryRuntimeRegistry() diff --git a/nl2sql/runtime_config.py b/nl2sql/runtime_config.py new file mode 100644 index 0000000..622b7a6 --- /dev/null +++ b/nl2sql/runtime_config.py @@ -0,0 +1,26 @@ +"""NL2SQL 进程内运行时配置。""" +from __future__ import annotations + +from pydantic import BaseModel, ConfigDict, Field + + +class Nl2SqlRuntimeConfig(BaseModel): + """可由管理员动态调整的非敏感查询参数。""" + + model_config = ConfigDict(extra="forbid", validate_assignment=True) + + cache_ttl: int = Field(default=300, ge=30, le=86400) + retrieval_top_k: int = Field(default=5, ge=1, le=50) + retrieval_threshold: float = Field(default=0.0, ge=0.0, le=1.0) + max_rows: int = Field(default=1000, ge=1, le=100000) + max_join_depth: int = Field(default=3, ge=0, le=10) + + def update(self, **values: object) -> dict: + """校验并更新配置,返回当前完整配置。""" + updated = self.model_copy(update=values) + for key, value in updated.model_dump().items(): + setattr(self, key, value) + return self.model_dump() + + +runtime_config = Nl2SqlRuntimeConfig() diff --git a/nl2sql/schema.py b/nl2sql/schema.py new file mode 100644 index 0000000..772b43f --- /dev/null +++ b/nl2sql/schema.py @@ -0,0 +1,59 @@ +"""从 MySQL information_schema 加载并校验 NL2SQL 权威 Schema。""" +from __future__ import annotations + +from typing import Any + +from sqlalchemy import bindparam, text + +from nl2sql.metadata import normalize_column_row, normalize_table_row + + +TABLES_SQL = text( + """ + SELECT TABLE_NAME, TABLE_COMMENT, TABLE_TYPE + FROM information_schema.tables + WHERE TABLE_SCHEMA = :database + AND TABLE_NAME IN :candidate_tables + """ +).bindparams(bindparam("candidate_tables", expanding=True)) + +COLUMNS_SQL = text( + """ + SELECT TABLE_NAME, COLUMN_NAME, COLUMN_COMMENT, DATA_TYPE, + IS_NULLABLE, ORDINAL_POSITION + FROM information_schema.columns + WHERE TABLE_SCHEMA = :database + AND TABLE_NAME IN :candidate_tables + ORDER BY TABLE_NAME, ORDINAL_POSITION + """ +).bindparams(bindparam("candidate_tables", expanding=True)) + + +async def load_authoritative_schema( + session, + *, + database: str, + candidate_tables: set[str] | list[str], +) -> dict[str, list[dict[str, Any]]]: + """从权威元数据源加载候选表,并剔除不存在的表和孤立字段。""" + candidates = {str(name).strip() for name in candidate_tables if str(name).strip()} + if not candidates: + return {"tables": [], "columns": []} + + params = {"database": database, "candidate_tables": sorted(candidates)} + table_result = await session.execute(TABLES_SQL, params) + column_result = await session.execute(COLUMNS_SQL, params) + + tables = [ + normalized + for row in table_result.mappings() + if (normalized := normalize_table_row(dict(row))) is not None + and normalized["table_name"] in candidates + ] + valid_table_names = {table["table_name"] for table in tables} + columns = [] + for row in column_result.mappings(): + normalized = normalize_column_row(dict(row)) + if normalized["table_name"] in valid_table_names and normalized["field_name"]: + columns.append(normalized) + return {"tables": tables, "columns": columns} diff --git a/nl2sql/semantic_catalog.json b/nl2sql/semantic_catalog.json new file mode 100644 index 0000000..b5cec37 --- /dev/null +++ b/nl2sql/semantic_catalog.json @@ -0,0 +1,67 @@ +{ + "version": "2026-09-12", + "terms": [ + { + "term": "基金产品", + "enabled": true, + "aliases": ["基金产品", "产品", "基金", "基金名称", "产品名称"], + "tables": ["fin_product"], + "fields": ["product_code", "product_name", "product_type", "risk_level", "status"] + }, + { + "term": "股票型基金", + "enabled": true, + "aliases": ["股票型", "股票基金", "权益类"], + "tables": ["fin_product"], + "fields": ["product_type"], + "value_hint": "product_type 通常使用业务库中的股票型标准值,不能自行创造枚举值" + }, + { + "term": "客户", + "enabled": true, + "aliases": ["客户", "投资人", "持有人"], + "tables": ["fin_customer_profile", "fin_holdings", "fin_risk_assessment"], + "fields": ["customer_id", "risk_level", "customer_level"] + }, + { + "term": "持仓", + "enabled": true, + "aliases": ["持仓", "持有基金", "资产配置"], + "tables": ["fin_holdings"], + "fields": ["customer_id", "product_id", "shares", "current_value", "profit_loss", "profit_ratio", "status"] + }, + { + "term": "收益率", + "enabled": true, + "aliases": ["收益率", "收益", "回报率", "近一年收益"], + "tables": ["fund_performance"], + "fields": ["return_rate", "period", "calc_date"], + "metric_hint": "收益率优先使用 fund_performance.return_rate,并结合 period 或 calc_date 限定统计区间" + }, + { + "term": "最大回撤", + "enabled": true, + "aliases": ["最大回撤", "回撤"], + "tables": ["fund_performance"], + "fields": ["max_drawdown", "period", "calc_date"], + "metric_hint": "最大回撤使用 fund_performance.max_drawdown,不与收益率字段混用" + } + ], + "relationships": [ + { + "left": "fin_holdings.product_id", + "right": "fin_product.id", + "meaning": "持仓通过 product_id 关联基金产品" + }, + { + "left": "fund_performance.product_id", + "right": "fin_product.id", + "meaning": "基金业绩通过 product_id 关联基金产品" + }, + { + "left": "fin_holdings.customer_id", + "right": "fin_customer_profile.customer_id", + "meaning": "持仓通过 customer_id 关联客户画像" + } + ] +} diff --git a/nl2sql/semantics.py b/nl2sql/semantics.py new file mode 100644 index 0000000..f4b0cd6 --- /dev/null +++ b/nl2sql/semantics.py @@ -0,0 +1,162 @@ +"""基金业务语义目录:把自然语言术语映射为受权威 Schema 约束的数据库概念。""" +from __future__ import annotations + +import json +import hashlib +from functools import lru_cache +from pathlib import Path +from typing import Any + +DEFAULT_CATALOG_PATH = Path(__file__).with_name("semantic_catalog.json") + + +@lru_cache(maxsize=8) +def _load_catalog_cached(path: str) -> dict[str, Any]: + payload = json.loads(Path(path).read_text(encoding="utf-8")) + return validate_semantic_catalog(payload) + + +def load_semantic_catalog(path: str | Path | None = None) -> dict[str, Any]: + """加载可替换的基金业务语义目录。""" + return _load_catalog_cached(str(Path(path or DEFAULT_CATALOG_PATH).resolve())) + + +def validate_semantic_catalog(catalog: dict[str, Any]) -> dict[str, Any]: + """校验语义目录结构,保证目录错误在加载阶段暴露。""" + if not isinstance(catalog, dict) or not str(catalog.get("version", "")).strip(): + raise ValueError("语义目录必须包含 version") + terms = catalog.get("terms") + if not isinstance(terms, list): + raise ValueError("语义目录必须包含 terms 数组") + for item in terms: + if not isinstance(item, dict) or not str(item.get("term", "")).strip(): + raise ValueError("语义目录 term 无效") + if "enabled" in item and not isinstance(item["enabled"], bool): + raise ValueError("语义目录 enabled 必须是布尔值") + for key in ("aliases", "tables", "fields"): + if not isinstance(item.get(key), list) or not item[key]: + raise ValueError(f"语义目录 {key} 无效") + relationships = catalog.get("relationships", []) + if not isinstance(relationships, list): + raise ValueError("语义目录 relationships 必须是数组") + for item in relationships: + if not isinstance(item, dict) or not all( + str(item.get(key, "")).strip() for key in ("left", "right", "meaning") + ): + raise ValueError("语义目录关联关系无效") + return catalog + + +def clear_semantic_catalog_cache() -> None: + """清理目录缓存,使下一次请求重新读取文件。""" + _load_catalog_cached.cache_clear() + + +def refresh_semantic_catalog(path: str | Path | None = None) -> dict[str, Any]: + """校验并刷新语义目录;新目录无效时保留当前缓存。""" + catalog_path = Path(path or DEFAULT_CATALOG_PATH).resolve() + payload = json.loads(catalog_path.read_text(encoding="utf-8")) + catalog = validate_semantic_catalog(payload) + previous = get_semantic_catalog_info(catalog_path) + _load_catalog_cached.cache_clear() + _load_catalog_cached(str(catalog_path)) + current = _build_semantic_catalog_info(catalog, catalog_path) + return { + **current, + "previous_version": previous["version"], + "changed": previous["digest"] != current["digest"], + } + + +def _build_semantic_catalog_info(catalog: dict[str, Any], path: Path) -> dict[str, Any]: + """生成不含业务明细的目录摘要。""" + digest = hashlib.sha256( + json.dumps(catalog, ensure_ascii=False, sort_keys=True).encode("utf-8") + ).hexdigest() + enabled_count = sum(item.get("enabled", True) for item in catalog["terms"]) + return { + "version": catalog["version"], + "term_count": len(catalog["terms"]), + "enabled_term_count": enabled_count, + "disabled_term_count": len(catalog["terms"]) - enabled_count, + "relationship_count": len(catalog.get("relationships", [])), + "digest": digest, + "path": str(path), + } + + +def get_semantic_catalog_info(path: str | Path | None = None) -> dict[str, Any]: + """返回不含业务明细的语义目录摘要。""" + catalog = load_semantic_catalog(path) + return _build_semantic_catalog_info(catalog, Path(path or DEFAULT_CATALOG_PATH).resolve()) + + +def resolve_semantics(question: str, *, catalog: dict[str, Any] | None = None) -> dict[str, Any]: + """根据问题匹配业务术语,返回未经过 Schema 过滤的候选语义。""" + text = (question or "").strip() + catalog = catalog or load_semantic_catalog() + matched = [ + item + for item in catalog["terms"] + if item.get("enabled", True) + if any(alias in text for alias in item.get("aliases", [])) + ] + tables = sorted({table for item in matched for table in item["tables"]}) + fields = sorted({field for item in matched for field in item["fields"]}) + metrics = [ + { + "term": item["term"], + "field": item["fields"][0], + "hint": item.get("metric_hint", item.get("value_hint", "")), + } + for item in matched + if "metric_hint" in item + ] + return { + "terms": [item["term"] for item in matched], + "tables": tables, + "fields": fields, + "metrics": metrics, + "relationships": list(catalog.get("relationships", [])), + } + + +def build_semantic_context(question: str, schema: dict[str, Any]) -> dict[str, Any]: + """只保留当前权威 Schema 中存在的表、字段和关联,避免语义目录越权扩张。""" + resolved = resolve_semantics(question) + schema_tables = { + str(item.get("table_name")) + for item in schema.get("tables", []) + if item.get("table_name") + } + schema_fields = { + (str(item.get("table_name")), str(item.get("field_name"))) + for item in schema.get("columns", []) + if item.get("table_name") and item.get("field_name") + } + tables = sorted(set(resolved["tables"]) & schema_tables) + fields = sorted( + { + field + for table, field in schema_fields + if table in tables and field in set(resolved["fields"]) + } + ) + metrics = [ + metric + for metric in resolved["metrics"] + if any(metric["field"] == field for _, field in schema_fields if _ in tables) + ] + relationships = [ + relation + for relation in resolved["relationships"] + if relation["left"].split(".", 1)[0] in tables + and relation["right"].split(".", 1)[0] in tables + ] + return { + "terms": resolved["terms"], + "tables": tables, + "fields": fields, + "metrics": metrics, + "relationships": relationships, + } diff --git a/nl2sql/session_context.py b/nl2sql/session_context.py new file mode 100644 index 0000000..e53ad66 --- /dev/null +++ b/nl2sql/session_context.py @@ -0,0 +1,82 @@ +"""NL2SQL 会话上下文的 Redis 存储与 Prompt 格式化。""" +from __future__ import annotations + +import json +import logging +from typing import Any + + +logger = logging.getLogger("nl2sql.session_context") + + +class SessionContextStore: + """按员工隔离并限制大小的短期查询上下文。""" + + def __init__(self, redis, *, ttl: int = 1800, max_messages: int = 8, max_chars: int = 4000): + self.redis = redis + self.ttl = ttl + self.max_messages = max_messages + self.max_chars = max_chars + + def _key(self, session_id: str) -> str: + return f"nl2sql:session:{session_id}:context" + + def _owner_key(self, session_id: str) -> str: + return f"nl2sql:session:{session_id}:owner" + + async def _belongs_to(self, user_id: int, session_id: str) -> bool: + owner_key = self._owner_key(session_id) + owner = await self.redis.get(owner_key) + if owner is None: + return bool(await self.redis.set(owner_key, str(user_id), nx=True, ex=self.ttl)) + if isinstance(owner, bytes): + owner = owner.decode() + return str(owner) == str(user_id) + + async def load(self, user_id: int, session_id: str | None) -> list[dict[str, str]]: + """读取当前员工的上下文,异常或跨员工访问时返回空列表。""" + if not session_id: + return [] + try: + if not await self._belongs_to(user_id, session_id): + return [] + raw = await self.redis.get(self._key(session_id)) + if not raw: + return [] + if isinstance(raw, bytes): + raw = raw.decode() + value = json.loads(raw) + return value if isinstance(value, list) else [] + except Exception: # noqa: BLE001 上下文故障不阻断查询 + logger.warning("NL2SQL 会话上下文读取失败", exc_info=True) + return [] + + async def append(self, user_id: int, session_id: str | None, question: str, status: str) -> bool: + """追加问题和状态摘要,不保存 SQL、结果行或敏感数据。""" + if not session_id: + return False + try: + if not await self._belongs_to(user_id, session_id): + return False + messages = await self.load(user_id, session_id) + messages.append({"role": "user", "content": question[:1000]}) + messages.append({"role": "assistant", "content": status[:100]}) + messages = messages[-self.max_messages :] + while len(json.dumps(messages, ensure_ascii=False)) > self.max_chars and messages: + messages.pop(0) + await self.redis.set(self._key(session_id), json.dumps(messages, ensure_ascii=False), ex=self.ttl) + return True + except Exception: # noqa: BLE001 上下文故障不阻断查询 + logger.warning("NL2SQL 会话上下文写入失败", exc_info=True) + return False + + +def build_conversation_context(messages: list[dict[str, Any]], *, max_chars: int = 2000) -> str: + """构造有界 Prompt 上下文,只保留角色和文本内容。""" + lines: list[str] = [] + for message in messages: + role = str(message.get("role", ""))[:20] + content = str(message.get("content", ""))[:500] + if role and content: + lines.append(f"{role}: {content}") + return "\n".join(lines)[-max_chars:] diff --git a/nl2sql/sql_agent.py b/nl2sql/sql_agent.py new file mode 100644 index 0000000..e3a544f --- /dev/null +++ b/nl2sql/sql_agent.py @@ -0,0 +1,89 @@ +"""NL2SQL SQL 生成器:只负责调用 Chat 模型并清理模型输出。""" +from __future__ import annotations + +import json +import re +from dataclasses import dataclass +from typing import Any + +from sqlglot import exp, parse +from sqlglot.errors import ParseError + +from tool.llm import llm +from nl2sql.semantics import build_semantic_context + + +class SqlGenerationError(ValueError): + """模型未返回可用的只读 SQL。""" + + +@dataclass(frozen=True) +class GeneratedSql: + sql: str + + +def _clean_model_output(output: str) -> str: + value = (output or "").strip() + value = re.sub(r"^```(?:sql)?\s*", "", value, flags=re.IGNORECASE) + value = re.sub(r"\s*```$", "", value).strip() + return value.removesuffix(";").strip() + + +def _build_prompt( + question: str, + schema: dict[str, Any], + few_shot: list[dict[str, Any]], + semantic_context: dict[str, Any] | None = None, + conversation_context: str = "", +) -> str: + return ( + "你是基金业务数据库 SQL 生成器。\n" + "只根据提供的 Schema 生成一条 MySQL SELECT,禁止写操作、跨库访问和未提供的表字段。\n" + "只输出 SQL,不要输出解释、Markdown 或代码围栏。\n" + f"用户问题:{question}\n" + f"Schema:{json.dumps(schema, ensure_ascii=False, sort_keys=True)}\n" + f"业务语义:{json.dumps(semantic_context or {}, ensure_ascii=False, sort_keys=True)}\n" + f"会话上下文:{conversation_context}\n" + f"Few-shot:{json.dumps(few_shot or [], ensure_ascii=False, sort_keys=True)}" + ) + + +async def generate_sql( + question: str, + schema: dict[str, Any], + *, + llm_client=llm, + few_shot: list[dict[str, Any]] | None = None, + semantic_context: dict[str, Any] | None = None, + conversation_context: str = "", +) -> GeneratedSql: + """调用 Chat 模型生成 SQL,并在返回前确认其为单条 SELECT。""" + if not question or not question.strip(): + raise SqlGenerationError("用户问题不能为空") + messages = [ + {"role": "system", "content": "你必须严格遵守只输出单条 SELECT SQL。"}, + { + "role": "user", + "content": _build_prompt( + question, + schema, + few_shot or [], + semantic_context or build_semantic_context(question, schema), + conversation_context, + ), + }, + ] + try: + output = await llm_client.chat(messages, temperature=0) + except Exception as exc: # noqa: BLE001 统一收敛模型异常 + raise SqlGenerationError("SQL 模型调用失败") from exc + sql = _clean_model_output(output) + if not sql: + raise SqlGenerationError("模型未返回 SQL") + try: + statements = parse(sql, read="mysql") + except ParseError as exc: + raise SqlGenerationError("模型返回的 SQL 无法解析") from exc + if len(statements) != 1 or not isinstance(statements[0], exp.Select): + raise SqlGenerationError("模型只允许返回单条 SELECT") + return GeneratedSql(sql=statements[0].sql(dialect="mysql")) diff --git a/nl2sql/sql_security.py b/nl2sql/sql_security.py new file mode 100644 index 0000000..9c50390 --- /dev/null +++ b/nl2sql/sql_security.py @@ -0,0 +1,97 @@ +"""NL2SQL 只读 SQL 的 AST 安全校验。""" +from __future__ import annotations + +from dataclasses import dataclass + +from sqlglot import exp, parse +from sqlglot.errors import ParseError + + +_DANGEROUS_FUNCTIONS = { + "LOAD_FILE", + "UUID_FILE_NAME", + "BENCHMARK", + "SLEEP", +} + +class SqlSecurityError(ValueError): + """SQL 未通过只读和权限校验。""" + + +@dataclass(frozen=True) +class ValidatedSql: + sql: str + access_tables: set[str] + + +def _read_limit(statement: exp.Expression) -> int | None: + limit = statement.args.get("limit") + if limit is None: + return None + expression = limit.args.get("expression") + if not isinstance(expression, exp.Literal) or not expression.is_number: + raise SqlSecurityError("只允许使用数字 LIMIT") + value = int(expression.this) + if value < 0: + raise SqlSecurityError("LIMIT 不能为负数") + return value + + +def validate_select_sql( + sql: str, + *, + authorized_tables: set[str], + authorized_columns: dict[str, set[str]] | None = None, + max_rows: int, + max_joins: int = 5, + max_columns: int = 100, +) -> ValidatedSql: + """校验单条 SELECT,检查授权表并将 LIMIT 控制在最大行数内。""" + if not isinstance(sql, str) or not sql.strip(): + raise SqlSecurityError("SQL 不能为空") + if max_rows <= 0: + raise SqlSecurityError("最大行数必须为正数") + if max_joins < 0 or max_columns <= 0: + raise SqlSecurityError("SQL 复杂度限制参数无效") + try: + statements = parse(sql, read="mysql") + except ParseError as exc: + raise SqlSecurityError("SQL 解析失败") from exc + if len(statements) != 1 or not isinstance(statements[0], exp.Select): + raise SqlSecurityError("只允许执行单条 SELECT") + + statement = statements[0] + join_count = len(list(statement.find_all(exp.Join))) + if join_count > max_joins: + raise SqlSecurityError("SQL JOIN 深度超过限制") + if len(statement.expressions) > max_columns: + raise SqlSecurityError("SQL 返回列数超过限制") + for function in statement.find_all(exp.Func): + function_name = getattr(function, "name", "") or function.sql_name() + if function_name.upper() in _DANGEROUS_FUNCTIONS: + raise SqlSecurityError("SQL 包含危险函数") + access_tables: set[str] = set() + aliases: dict[str, str] = {} + for table in statement.find_all(exp.Table): + if table.db or table.catalog: + raise SqlSecurityError("禁止跨库访问") + access_tables.add(table.name) + aliases[table.alias_or_name] = table.name + if not access_tables.issubset(authorized_tables): + raise SqlSecurityError("SQL 访问了未授权表") + if authorized_columns is not None: + if statement.find(exp.Star): + raise SqlSecurityError("配置字段权限时禁止 SELECT *") + default_table = next(iter(access_tables), None) if len(access_tables) == 1 else None + for column in statement.find_all(exp.Column): + table_name = aliases.get(column.table, column.table) or default_table + if table_name is None or column.name not in authorized_columns.get(table_name, set()): + raise SqlSecurityError("SQL 访问了未授权字段") + + current_limit = _read_limit(statement) + if current_limit is None or current_limit > max_rows: + statement = statement.limit(max_rows) + return ValidatedSql( + sql=statement.sql(dialect="mysql"), + access_tables=access_tables, + ) diff --git a/nl2sql/streaming.py b/nl2sql/streaming.py new file mode 100644 index 0000000..e82aa0b --- /dev/null +++ b/nl2sql/streaming.py @@ -0,0 +1,41 @@ +"""NL2SQL 查询流式事件适配。""" +from __future__ import annotations + +from collections.abc import AsyncIterator, Awaitable, Callable +from typing import Any + +from nl2sql.contracts import DataQueryRequest, DataQueryResult + + +async def stream_query_events( + request: DataQueryRequest, + *, + query_runner: Callable[..., Awaitable[DataQueryResult]], + **dependencies: Any, +) -> AsyncIterator[dict[str, Any]]: + """复用查询服务并输出不泄露敏感内容的结构化事件。""" + query_id = dependencies.get("query_id") + if not query_id: + raise ValueError("query_id must be provided") + yield { + "event": "started", + "query_id": query_id, + "trace_id": request.trace_id, + } + try: + result = await query_runner(request, **dependencies) + except Exception as exc: # noqa: BLE001 流式失败只输出异常类型 + yield { + "event": "failed", + "query_id": query_id, + "trace_id": request.trace_id, + "error_type": type(exc).__name__, + } + return + yield { + "event": "completed", + "query_id": result.query_id, + "trace_id": result.trace_id or request.trace_id, + "row_count": result.row_count, + "truncated": result.truncated, + } diff --git a/nl2sql/supervisor.py b/nl2sql/supervisor.py new file mode 100644 index 0000000..03154e1 --- /dev/null +++ b/nl2sql/supervisor.py @@ -0,0 +1,64 @@ +"""NL2SQL Supervisor 的意图和会话并发控制。""" +from __future__ import annotations + +from uuid import uuid4 + + +class UnsupportedIntent(ValueError): + """当前版本不支持的查询意图。""" + + +class SessionBusyError(RuntimeError): + """同一会话已有查询运行。""" + + +def detect_intent(question: str) -> str: + """根据有限关键词识别当前版本的任务类型。""" + text = question.lower() + if "etl" in text or "数据任务" in text: + return "etl_build" + if "口径" in text or "caliber" in text: + return "caliber_maintain" + if "血缘" in text or "lineage" in text: + return "lineage_query" + return "query" + + +def ensure_query_intent(question: str) -> None: + """M1 只允许 query 意图进入 SQL 主链路。""" + intent = detect_intent(question) + if intent != "query": + raise UnsupportedIntent(f"当前不支持 {intent} 意图") + + +class SessionLock: + """基于 Redis NX 的同会话互斥锁。""" + + def __init__(self, redis, session_id: str, *, ttl: int = 120): + self.redis = redis + self.key = f"nl2sql:session:{session_id}:lock" + self.token = uuid4().hex + self.ttl = ttl + self.acquired = False + + async def acquire(self) -> None: + if not await self.redis.set(self.key, self.token, nx=True, ex=self.ttl): + raise SessionBusyError("同一会话已有查询运行") + self.acquired = True + + async def release(self) -> None: + if not self.acquired: + return + release_script = """ + if redis.call('get', KEYS[1]) == ARGV[1] then + return redis.call('del', KEYS[1]) + end + return 0 + """ + if hasattr(self.redis, "eval"): + await self.redis.eval(release_script, 1, self.key, self.token) + else: + current = await self.redis.get(self.key) + if current in {self.token, self.token.encode()}: + await self.redis.delete(self.key) + self.acquired = False diff --git a/repositories/advisor_draft.py b/repositories/advisor_draft.py new file mode 100644 index 0000000..4012507 --- /dev/null +++ b/repositories/advisor_draft.py @@ -0,0 +1,62 @@ +"""投顾 Agent 草稿仓储。""" +from __future__ import annotations + +from sqlalchemy import select + +from agent.advisor_agent.drafts import ( + DRAFT_STATUS_DISCARDED, + DRAFT_STATUS_DRAFT, +) +from model.advisor_draft import AdvisorDraft +from repositories.base import BaseRepository + + +class AdvisorDraftRepo(BaseRepository): + model = AdvisorDraft + + async def get_by_draft_id(self, draft_id: str) -> AdvisorDraft | None: + return await self.db.scalar( + select(AdvisorDraft).where(AdvisorDraft.draft_id == draft_id) + ) + + async def save(self, draft: AdvisorDraft) -> AdvisorDraft: + self.db.add(draft) + await self.db.commit() + await self.db.refresh(draft) + return draft + + async def list_drafts( + self, + *, + advisor_id: int | None = None, + customer_id: int | None = None, + status: str | None = None, + limit: int = 20, + offset: int = 0, + ) -> tuple[int, list[AdvisorDraft]]: + filters = [] + if advisor_id is not None: + filters.append(AdvisorDraft.advisor_id == advisor_id) + if customer_id is not None: + filters.append(AdvisorDraft.customer_id == customer_id) + if status is not None: + filters.append(AdvisorDraft.status == status) + + total = await self.count(where=filters) + items = await self.list( + where=filters, + order_by=AdvisorDraft.update_time.desc(), + limit=limit, + offset=offset, + ) + return total, items + + async def discard(self, draft: AdvisorDraft) -> AdvisorDraft: + if draft.status == DRAFT_STATUS_DISCARDED: + return draft + if draft.status != DRAFT_STATUS_DRAFT: + raise ValueError("草稿状态无效") + draft.status = DRAFT_STATUS_DISCARDED + await self.db.commit() + await self.db.refresh(draft) + return draft diff --git a/repositories/advisor_report.py b/repositories/advisor_report.py index e896f7a..6defb37 100644 --- a/repositories/advisor_report.py +++ b/repositories/advisor_report.py @@ -60,12 +60,20 @@ class AdvisorReportRepo(BaseRepository): return (await self.db.scalar(stmt)) or 0 async def list_by_customer( - self, customer_id: int, limit: int = 100, offset: int = 0 + self, + customer_id: int, + *, + advisor_id: int | None = None, + limit: int = 100, + offset: int = 0, ) -> list[AdvisorReport]: """某客户的历史建议报告(360 视图用)。""" + conditions = [AdvisorReport.customer_id == customer_id] + if advisor_id is not None: + conditions.append(AdvisorReport.advisor_id == advisor_id) stmt = ( select(AdvisorReport) - .where(AdvisorReport.customer_id == customer_id) + .where(*conditions) .order_by(AdvisorReport.id.desc()) .limit(limit) .offset(offset) diff --git a/repositories/advisor_visit_record.py b/repositories/advisor_visit_record.py index 8293e61..27cc0a6 100644 --- a/repositories/advisor_visit_record.py +++ b/repositories/advisor_visit_record.py @@ -10,6 +10,15 @@ from repositories.base import BaseRepository class AdvisorVisitRecordRepo(BaseRepository): model = AdvisorVisitRecord + async def get_by_advisor( + self, visit_id: int, advisor_id: int + ) -> AdvisorVisitRecord | None: + stmt = select(AdvisorVisitRecord).where( + AdvisorVisitRecord.id == visit_id, + AdvisorVisitRecord.advisor_id == advisor_id, + ) + return await self.db.scalar(stmt) + async def list_by_advisor( self, *, diff --git a/repositories/audit_log.py b/repositories/audit_log.py index f728274..29b00ae 100644 --- a/repositories/audit_log.py +++ b/repositories/audit_log.py @@ -12,10 +12,19 @@ from repositories.base import BaseRepository class AuditLogRepo(BaseRepository): model = AuditLog - def _conds(self, *, user_id, module, action, keyword, start, end): + def _conds(self, *, user_id, module, action, customer_id, keyword, start, end): conds = [AuditLog.user_id == user_id, AuditLog.module == module] if action: conds.append(AuditLog.action == action) + if customer_id is not None: + customer = str(customer_id) + conds.append( + or_( + AuditLog.target == customer, + AuditLog.detail.like(f'%"customer_id": {customer}%'), + AuditLog.detail.like(f'%"customer_id":{customer}%'), + ) + ) if start: conds.append(AuditLog.create_time >= start) if end: @@ -31,6 +40,7 @@ class AuditLogRepo(BaseRepository): user_id: int, module: str = "advisor", action: str | None = None, + customer_id: int | None = None, keyword: str | None = None, start: datetime | None = None, end: datetime | None = None, @@ -39,7 +49,7 @@ class AuditLogRepo(BaseRepository): ) -> list[AuditLog]: stmt = ( select(AuditLog) - .where(*self._conds(user_id=user_id, module=module, action=action, keyword=keyword, start=start, end=end)) + .where(*self._conds(user_id=user_id, module=module, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end)) .order_by(AuditLog.id.desc()) .limit(limit) .offset(offset) @@ -52,6 +62,7 @@ class AuditLogRepo(BaseRepository): user_id: int, module: str = "advisor", action: str | None = None, + customer_id: int | None = None, keyword: str | None = None, start: datetime | None = None, end: datetime | None = None, @@ -59,6 +70,6 @@ class AuditLogRepo(BaseRepository): stmt = ( select(func.count()) .select_from(AuditLog) - .where(*self._conds(user_id=user_id, module=module, action=action, keyword=keyword, start=start, end=end)) + .where(*self._conds(user_id=user_id, module=module, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end)) ) return (await self.db.scalar(stmt)) or 0 diff --git a/repositories/customer_relation.py b/repositories/customer_relation.py index dad3410..afc4329 100644 --- a/repositories/customer_relation.py +++ b/repositories/customer_relation.py @@ -11,6 +11,7 @@ from decimal import Decimal from sqlalchemy import func, or_, select +from common.common_const import CUSTOMER_REL_STATUS_SIGNED, CUSTOMER_REL_STATUS_UNSIGNED from model.customer_relation import CustomerRelation from model.fin_customer_profile import FinCustomerProfile from model.fin_holdings import FinHoldings @@ -24,6 +25,23 @@ _HOLDING_STATUS = "持有中" class CustomerRelationRepo(BaseRepository): model = CustomerRelation + async def get_active_relation( + self, *, customer_id: int, advisor_id: int + ) -> CustomerRelation | None: + """读取投顾可访问的当前关系,排除已结束关系。""" + return await self.db.scalar( + select(CustomerRelation) + .where( + CustomerRelation.customer_id == customer_id, + CustomerRelation.advisor_id == advisor_id, + CustomerRelation.status.in_( + [CUSTOMER_REL_STATUS_UNSIGNED, CUSTOMER_REL_STATUS_SIGNED] + ), + ) + .order_by(CustomerRelation.id.desc()) + .limit(1) + ) + async def get_by_customer_advisor( self, customer_id: int, advisor_id: int ) -> CustomerRelation | None: diff --git a/repositories/fund_performance.py b/repositories/fund_performance.py new file mode 100644 index 0000000..2ebf4c9 --- /dev/null +++ b/repositories/fund_performance.py @@ -0,0 +1,31 @@ +"""基金业绩指标仓储。""" +from __future__ import annotations + +from sqlalchemy import select + +from model.fund_performance import FundPerformance +from repositories.base import BaseRepository + + +class FundPerformanceRepo(BaseRepository): + model = FundPerformance + + async def get_latest_for_product(self, product_id: int) -> FundPerformance | None: + stmt = ( + select(FundPerformance) + .where(FundPerformance.product_id == product_id) + .order_by( + FundPerformance.calc_date.desc(), + FundPerformance.id.desc(), + ) + .limit(1) + ) + return (await self.db.scalars(stmt)).first() + + async def list_for_product(self, product_id: int) -> list[FundPerformance]: + stmt = ( + select(FundPerformance) + .where(FundPerformance.product_id == product_id) + .order_by(FundPerformance.period.asc()) + ) + return list((await self.db.scalars(stmt)).all()) diff --git a/repositories/nl2sql_permission.py b/repositories/nl2sql_permission.py new file mode 100644 index 0000000..cb904c7 --- /dev/null +++ b/repositories/nl2sql_permission.py @@ -0,0 +1,184 @@ +"""NL2SQL 查询权限仓储。""" +from __future__ import annotations + +from datetime import datetime + +from sqlalchemy import delete, select + +from model.nl2sql_permission import ( + Nl2SqlQueryRole, + Nl2SqlRoleColumnPermission, + Nl2SqlRoleTablePermission, + Nl2SqlSensitiveField, + Nl2SqlQueryHistory, +) +from repositories.base import BaseRepository + + +class Nl2SqlPermissionRepo(BaseRepository): + async def list_roles(self, *, include_inactive: bool = False): + stmt = select(Nl2SqlQueryRole).order_by(Nl2SqlQueryRole.id) + if not include_inactive: + stmt = stmt.where(Nl2SqlQueryRole.status == "active") + return list((await self.db.scalars(stmt)).all()) + + async def get_role(self, role_id: int): + return await self.db.get(Nl2SqlQueryRole, role_id) + + async def add_role(self, **kwargs): + return await self.add(Nl2SqlQueryRole(**kwargs)) + + async def update_role(self, role_id: int, **kwargs): + obj = await self.get_role(role_id) + if obj is None: + return None + for key, value in kwargs.items(): + if value is not None: + setattr(obj, key, value) + await self.db.commit() + await self.db.refresh(obj) + return obj + + async def list_role_table_permissions(self, role_id: int, *, include_inactive: bool = False): + stmt = select(Nl2SqlRoleTablePermission).where( + Nl2SqlRoleTablePermission.role_id == role_id + ) + if not include_inactive: + stmt = stmt.where(Nl2SqlRoleTablePermission.status == "active") + return list((await self.db.scalars(stmt.order_by(Nl2SqlRoleTablePermission.id))).all()) + + async def get_table_permission(self, permission_id: int): + return await self.db.get(Nl2SqlRoleTablePermission, permission_id) + + async def add_table_permission(self, **kwargs): + return await self.add(Nl2SqlRoleTablePermission(**kwargs)) + + async def update_table_permission(self, permission_id: int, **kwargs): + obj = await self.get_table_permission(permission_id) + if obj is None: + return None + for key, value in kwargs.items(): + if value is not None: + setattr(obj, key, value) + await self.db.commit() + await self.db.refresh(obj) + return obj + + async def list_role_column_permissions(self, role_id: int, *, include_inactive: bool = False): + stmt = select(Nl2SqlRoleColumnPermission).where( + Nl2SqlRoleColumnPermission.role_id == role_id + ) + if not include_inactive: + stmt = stmt.where(Nl2SqlRoleColumnPermission.status == "active") + return list((await self.db.scalars(stmt.order_by(Nl2SqlRoleColumnPermission.id))).all()) + + async def get_column_permission(self, permission_id: int): + return await self.db.get(Nl2SqlRoleColumnPermission, permission_id) + + async def add_column_permission(self, **kwargs): + return await self.add(Nl2SqlRoleColumnPermission(**kwargs)) + + async def update_column_permission(self, permission_id: int, **kwargs): + obj = await self.get_column_permission(permission_id) + if obj is None: + return None + for key, value in kwargs.items(): + if value is not None: + setattr(obj, key, value) + await self.db.commit() + await self.db.refresh(obj) + return obj + + async def list_sensitive_fields_admin(self, *, include_inactive: bool = False): + stmt = select(Nl2SqlSensitiveField).order_by(Nl2SqlSensitiveField.id) + if not include_inactive: + stmt = stmt.where(Nl2SqlSensitiveField.status == "active") + return list((await self.db.scalars(stmt)).all()) + + async def get_sensitive_field(self, field_id: int): + return await self.db.get(Nl2SqlSensitiveField, field_id) + + async def add_sensitive_field(self, **kwargs): + return await self.add(Nl2SqlSensitiveField(**kwargs)) + + async def update_sensitive_field(self, field_id: int, **kwargs): + obj = await self.get_sensitive_field(field_id) + if obj is None: + return None + for key, value in kwargs.items(): + if value is not None: + setattr(obj, key, value) + await self.db.commit() + await self.db.refresh(obj) + return obj + + async def delete_history_before(self, before: datetime) -> int: + result = await self.db.execute( + delete(Nl2SqlQueryHistory).where(Nl2SqlQueryHistory.create_time < before) + ) + await self.db.commit() + return int(result.rowcount or 0) + + async def get_role_by_employee_role(self, employee_role: str): + return await self.db.scalar( + select(Nl2SqlQueryRole).where( + Nl2SqlQueryRole.employee_role == employee_role, + Nl2SqlQueryRole.status == "active", + ) + ) + + async def list_table_permissions(self, role_id: int): + result = await self.db.scalars( + select(Nl2SqlRoleTablePermission).where( + Nl2SqlRoleTablePermission.role_id == role_id, + Nl2SqlRoleTablePermission.status == "active", + ) + ) + return list(result.all()) + + async def list_column_permissions(self, role_id: int): + result = await self.db.scalars( + select(Nl2SqlRoleColumnPermission).where( + Nl2SqlRoleColumnPermission.role_id == role_id, + Nl2SqlRoleColumnPermission.status == "active", + ) + ) + return list(result.all()) + + async def list_sensitive_fields(self): + result = await self.db.scalars( + select(Nl2SqlSensitiveField).where(Nl2SqlSensitiveField.status == "active") + ) + return list(result.all()) + + async def list_query_history( + self, + user_id: int, + *, + limit: int = 20, + offset: int = 0, + status: str | None = None, + start_time: datetime | None = None, + end_time: datetime | None = None, + ): + """按员工自身范围分页读取查询历史。""" + stmt = select(Nl2SqlQueryHistory).where(Nl2SqlQueryHistory.user_id == user_id) + if status: + stmt = stmt.where(Nl2SqlQueryHistory.status == status) + if start_time: + stmt = stmt.where(Nl2SqlQueryHistory.create_time >= start_time) + if end_time: + stmt = stmt.where(Nl2SqlQueryHistory.create_time <= end_time) + result = await self.db.scalars( + stmt.order_by(Nl2SqlQueryHistory.create_time.desc()).limit(limit).offset(offset) + ) + return list(result.all()) + + async def get_query_history(self, user_id: int, query_id: str): + """按员工自身范围读取单条查询历史。""" + return await self.db.scalar( + select(Nl2SqlQueryHistory).where( + Nl2SqlQueryHistory.user_id == user_id, + Nl2SqlQueryHistory.query_id == query_id, + ) + ) diff --git a/repositories/portfolio_benchmark.py b/repositories/portfolio_benchmark.py index 2535759..6d66331 100644 --- a/repositories/portfolio_benchmark.py +++ b/repositories/portfolio_benchmark.py @@ -24,3 +24,15 @@ class PortfolioBenchmarkRepo(BaseRepository): ) ).all() ) + + async def get_active_by_risk(self, risk_level: str) -> PortfolioBenchmark | None: + """Return the enabled benchmark used by the Agent rebalance flow.""" + return await self.db.scalar( + select(PortfolioBenchmark) + .where( + PortfolioBenchmark.status == _ACTIVE_STATUS, + PortfolioBenchmark.risk_level == risk_level, + ) + .order_by(PortfolioBenchmark.id.desc()) + .limit(1) + ) diff --git a/repositories/risk_assessment.py b/repositories/risk_assessment.py index 20b0ced..2ebc59f 100644 --- a/repositories/risk_assessment.py +++ b/repositories/risk_assessment.py @@ -1,6 +1,8 @@ """风评域仓储:风评记录 + 客户画像(画像主键为 customer_id)。""" from __future__ import annotations +from datetime import date + from sqlalchemy import select from model.fin_customer_profile import FinCustomerProfile @@ -11,6 +13,21 @@ from repositories.base import BaseRepository class RiskAssessmentRepo(BaseRepository): model = FinRiskAssessment + async def get_current_by_customer(self, customer_id: int) -> FinRiskAssessment | None: + """Return the latest currently valid risk assessment for a customer.""" + return await self.db.scalar( + select(FinRiskAssessment) + .where( + FinRiskAssessment.customer_id == customer_id, + FinRiskAssessment.valid_until >= date.today(), + ) + .order_by( + FinRiskAssessment.assessment_date.desc(), + FinRiskAssessment.id.desc(), + ) + .limit(1) + ) + class CustomerProfileRepo(BaseRepository): """客户画像仓储。注意:主键是 customer_id 而非 id,不适用基类 get(pk)。""" diff --git a/repositories/sensitive_word.py b/repositories/sensitive_word.py index 5b07778..3e03aaa 100644 --- a/repositories/sensitive_word.py +++ b/repositories/sensitive_word.py @@ -22,3 +22,7 @@ class SensitiveWordRepo(BaseRepository): ) ).all() ) + + async def list_active_words(self) -> list[str]: + """Return enabled words as strings for Agent content scanning.""" + return [item if isinstance(item, str) else item.word for item in await self.list_active()] diff --git a/repositories/sys_message.py b/repositories/sys_message.py index 65cd11b..383226a 100644 --- a/repositories/sys_message.py +++ b/repositories/sys_message.py @@ -10,6 +10,15 @@ from repositories.base import BaseRepository class SysMessageRepo(BaseRepository): model = SysMessage + async def get_by_biz_id(self, biz_id: str, *, user_id: int) -> SysMessage | None: + """按收件人和业务号查询已写入的站内信,供发送重试幂等使用。""" + return await self.db.scalar( + select(SysMessage).where( + SysMessage.biz_id == biz_id, + SysMessage.user_id == user_id, + ) + ) + async def list_by_user( self, user_id: int, limit: int = 100, offset: int = 0 ) -> list[SysMessage]: diff --git a/schemas/advisor.py b/schemas/advisor.py index 6d9e810..25c8b07 100644 --- a/schemas/advisor.py +++ b/schemas/advisor.py @@ -2,8 +2,14 @@ from __future__ import annotations from datetime import datetime -from typing import Any +from typing import Any, Literal +from common_const import ( + TALK_SCENE_CUSTOMER_COMPLAINT, + TALK_SCENE_MARKET_FLUCTUATION, + TALK_SCENE_PORTFOLIO_DIVERGENCE, + TALK_SCENE_RISK_BLOCK_ORDER, +) from pydantic import BaseModel, Field @@ -31,7 +37,12 @@ class TalkScriptReq(BaseModel): """生成沟通话术草稿(同步,超时由工作台降级为「稍后重试」)。""" customer_id: int = Field(gt=0) - scene_type: str = Field(min_length=1, max_length=64) + scene_type: Literal[ + TALK_SCENE_RISK_BLOCK_ORDER, + TALK_SCENE_MARKET_FLUCTUATION, + TALK_SCENE_PORTFOLIO_DIVERGENCE, + TALK_SCENE_CUSTOMER_COMPLAINT, + ] class RelationReq(BaseModel): @@ -51,6 +62,15 @@ class VisitCreateReq(BaseModel): audio_url: str | None = Field(default=None, max_length=512) +class VisitUpdateReq(BaseModel): + """回访记录可编辑字段;不提供字段时拒绝空更新。""" + + visit_type: str | None = Field(default=None, min_length=1, max_length=32) + visit_time: datetime | None = None + summary: str | None = None + audio_url: str | None = Field(default=None, max_length=512) + + class TodoHandleReq(BaseModel): """待办处理入参(process=开始处理 / done=完成)。""" diff --git a/schemas/advisor_agent.py b/schemas/advisor_agent.py new file mode 100644 index 0000000..f48a238 --- /dev/null +++ b/schemas/advisor_agent.py @@ -0,0 +1,60 @@ +"""投顾 Agent HTTP 请求 DTO。""" +from __future__ import annotations + +from typing import Any, Literal + +from pydantic import BaseModel, Field + +from common.common_const import ( + AGENT_INTENT_DIALOGUE_SCRIPT, + AGENT_INTENT_FUND_ANALYSIS, + AGENT_INTENT_REBALANCE, + AGENT_INTENT_RECOMMEND, + TALK_SCENE_CUSTOMER_COMPLAINT, + TALK_SCENE_MARKET_FLUCTUATION, + TALK_SCENE_PORTFOLIO_DIVERGENCE, + TALK_SCENE_RISK_BLOCK_ORDER, +) + + +class AdvisorDraftSaveReq(BaseModel): + title: str | None = Field(default=None, max_length=128) + content: str | None = None + structured_data: dict[str, Any] | None = None + + +class AdvisorDraftOperateReq(BaseModel): + operation: Literal["discard"] + + +class AdvisorRebalanceRunReq(BaseModel): + customer_id: int = Field(gt=0) + + +class AdvisorChatReq(BaseModel): + customer_id: int = Field(gt=0) + intent: Literal[ + AGENT_INTENT_RECOMMEND, + AGENT_INTENT_REBALANCE, + AGENT_INTENT_FUND_ANALYSIS, + AGENT_INTENT_DIALOGUE_SCRIPT, + ] | None = None + query: str | None = Field(default=None, max_length=4000) + + +class AdvisorFundAnalysisReq(BaseModel): + customer_id: int | None = Field(default=None, gt=0) + fund_codes: list[str] = Field(min_length=1) + fund: dict[str, Any] | None = None + performance: list[dict[str, Any]] | None = None + + +class AdvisorTalkScriptReq(BaseModel): + customer_id: int = Field(gt=0) + scene_type: Literal[ + TALK_SCENE_RISK_BLOCK_ORDER, + TALK_SCENE_MARKET_FLUCTUATION, + TALK_SCENE_PORTFOLIO_DIVERGENCE, + TALK_SCENE_CUSTOMER_COMPLAINT, + ] + customer_name: str = Field(default="客户", max_length=32) diff --git a/schemas/nl2sql.py b/schemas/nl2sql.py new file mode 100644 index 0000000..7ff639f --- /dev/null +++ b/schemas/nl2sql.py @@ -0,0 +1,40 @@ +"""NL2SQL HTTP 请求模型。""" +from __future__ import annotations + +from typing import Literal + +from pydantic import BaseModel, Field + + +class DataQueryReq(BaseModel): + """自然语言查询请求。""" + + question: str = Field(min_length=1, max_length=2000) + session_id: str | None = Field(default=None, max_length=64) + caller_agent: str | None = Field(default=None, max_length=64) + data_scope: dict | None = None + max_rows: int | None = Field(default=None, gt=0, le=10000) + include_sql: bool = False + page: int = Field(default=1, ge=1, le=100000) + page_size: int = Field(default=100, ge=1, le=10000) + sort_by: str | None = Field(default=None, max_length=128) + sort_order: Literal["asc", "desc"] = "asc" + output_format: Literal["json", "csv"] = "json" + + +class DataExplainReq(BaseModel): + """安全 EXPLAIN 请求。""" + + sql: str = Field(min_length=1, max_length=20000) + + +class DataKillReq(BaseModel): + """管理员中止运行中查询请求。""" + + query_id: str = Field(min_length=1, max_length=64) + + +class DataCacheInvalidateReq(BaseModel): + """管理员按业务表失效查询缓存。""" + + table_names: list[str] = Field(min_length=1, max_length=50) diff --git a/schemas/nl2sql_admin.py b/schemas/nl2sql_admin.py new file mode 100644 index 0000000..59e5523 --- /dev/null +++ b/schemas/nl2sql_admin.py @@ -0,0 +1,76 @@ +"""NL2SQL 管理接口请求模型。""" +from __future__ import annotations + +from typing import Literal + +from pydantic import BaseModel, Field + + +class RoleCreateReq(BaseModel): + role_code: str = Field(min_length=1, max_length=64) + role_name: str = Field(min_length=1, max_length=128) + employee_role: str = Field(min_length=1, max_length=32) + can_query: bool = False + max_rows: int = Field(default=1000, gt=0, le=100000) + daily_quota: int = Field(default=0, ge=0, le=1000000) + + +class RoleUpdateReq(BaseModel): + role_name: str | None = Field(default=None, min_length=1, max_length=128) + can_query: bool | None = None + max_rows: int | None = Field(default=None, gt=0, le=100000) + daily_quota: int | None = Field(default=None, ge=0, le=1000000) + status: Literal["active", "inactive"] | None = None + + +class TablePermissionCreateReq(BaseModel): + table_name: str = Field(min_length=1, max_length=128) + row_scope_type: Literal["none", "customer_ids", "product_ids"] = "none" + row_scope_column: str | None = Field(default=None, max_length=128) + + +class TablePermissionUpdateReq(BaseModel): + row_scope_type: Literal["none", "customer_ids", "product_ids"] | None = None + row_scope_column: str | None = Field(default=None, max_length=128) + status: Literal["active", "inactive"] | None = None + + +class ColumnPermissionCreateReq(BaseModel): + table_name: str = Field(min_length=1, max_length=128) + column_name: str = Field(min_length=1, max_length=128) + access_mode: Literal["allow", "deny", "mask"] = "allow" + mask_type: Literal["partial", "hash"] | None = None + + +class ColumnPermissionUpdateReq(BaseModel): + access_mode: Literal["allow", "deny", "mask"] | None = None + mask_type: Literal["partial", "hash"] | None = None + status: Literal["active", "inactive"] | None = None + + +class SensitiveFieldCreateReq(BaseModel): + table_name: str = Field(min_length=1, max_length=128) + column_name: str = Field(min_length=1, max_length=128) + mask_type: Literal["partial", "hash"] = "partial" + description: str | None = Field(default=None, max_length=255) + + +class SensitiveFieldUpdateReq(BaseModel): + mask_type: Literal["partial", "hash"] | None = None + description: str | None = Field(default=None, max_length=255) + status: Literal["active", "inactive"] | None = None + + +class MaintenanceJobReq(BaseModel): + task: Literal["metadata_sync", "vector_cleanup", "consistency_check", "history_cleanup"] + before_days: int = Field(default=180, ge=1, le=3650) + + +class RuntimeConfigUpdateReq(BaseModel): + """管理员动态调整 NL2SQL 非敏感运行参数。""" + + cache_ttl: int | None = Field(default=None, ge=30, le=86400) + retrieval_top_k: int | None = Field(default=None, ge=1, le=50) + retrieval_threshold: float | None = Field(default=None, ge=0.0, le=1.0) + max_rows: int | None = Field(default=None, ge=1, le=100000) + max_join_depth: int | None = Field(default=None, ge=0, le=10) diff --git a/scripts/check_advisor_llm.py b/scripts/check_advisor_llm.py new file mode 100644 index 0000000..cd33df9 --- /dev/null +++ b/scripts/check_advisor_llm.py @@ -0,0 +1,35 @@ +"""Read-only smoke check for the configured LLM chat and embedding endpoints.""" +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from tool.llm import llm + + +async def check() -> None: + try: + answer = await llm.chat( + [ + {"role": "system", "content": "你是一个合规的基金投顾助手。"}, + {"role": "user", "content": "只回复:连接正常"}, + ], + temperature=0, + max_tokens=16, + ) + print("chat: ok", answer[:32]) + except Exception as exc: + print(f"chat: {type(exc).__name__}: {exc}") + + try: + vectors = await llm.embed(["投顾记忆联调"]) + print("embedding: ok", len(vectors), len(vectors[0]) if vectors else 0) + except Exception as exc: + print(f"embedding: {type(exc).__name__}: {exc}") + + +if __name__ == "__main__": + asyncio.run(check()) diff --git a/scripts/check_advisor_storage.py b/scripts/check_advisor_storage.py new file mode 100644 index 0000000..e15ed6d --- /dev/null +++ b/scripts/check_advisor_storage.py @@ -0,0 +1,39 @@ +"""Read-only check for投顾 Agent/工作台 MySQL tables and Redis connectivity.""" +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path + +from sqlalchemy import text + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from config.database import mysql, redis + + +async def check() -> None: + try: + async with mysql.get_engine().connect() as connection: + for table in ( + "advisor_draft", + "event_log", + "sensitive_word", + "advisor_report", + "advisor_todo", + "advisor_visit_record", + ): + result = await connection.execute( + text("SHOW TABLES LIKE :table"), {"table": table} + ) + print(f"{table}: {'present' if result.first() else 'missing'}") + print("mysql: ok") + await redis.client().ping() + print("redis: ok") + finally: + await mysql.dispose() + await redis.dispose() + + +if __name__ == "__main__": + asyncio.run(check()) diff --git a/scripts/check_nl2sql_consistency.py b/scripts/check_nl2sql_consistency.py new file mode 100644 index 0000000..877613d --- /dev/null +++ b/scripts/check_nl2sql_consistency.py @@ -0,0 +1,46 @@ +"""检查 MySQL 数据字典和 Milvus NL2SQL 向量数量的一致性。""" +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 import database +from config.database.milvus import client +from config.database.mysql import get_session_factory +from config.settings import settings +from nl2sql.metadata import build_metadata_chunks +from nl2sql.milvus_collections import NL2SQL_COLLECTION + + +async def collect_consistency() -> dict: + table_sql = text( + "SELECT TABLE_NAME, TABLE_COMMENT, TABLE_TYPE FROM information_schema.tables WHERE TABLE_SCHEMA=:db" + ) + column_sql = text( + "SELECT TABLE_NAME, COLUMN_NAME, COLUMN_COMMENT, DATA_TYPE, IS_NULLABLE, ORDINAL_POSITION " + "FROM information_schema.columns WHERE TABLE_SCHEMA=:db ORDER BY TABLE_NAME, ORDINAL_POSITION" + ) + try: + async with get_session_factory()() as session: + tables = [dict(row) for row in (await session.execute(table_sql, {"db": settings.mysql.database})).mappings()] + columns = [dict(row) for row in (await session.execute(column_sql, {"db": settings.mysql.database})).mappings()] + expected = len(build_metadata_chunks(tables, columns)) + rows = await client().query( + collection_name=NL2SQL_COLLECTION, + filter="is_valid == true and is_deprecated == false", + output_fields=["id"], + ) + actual = len(rows or []) + return {"expected_chunks": expected, "actual_valid_vectors": actual, "consistent": expected == actual} + finally: + await database.mysql.dispose() + await database.milvus.dispose() + + +if __name__ == "__main__": + print(asyncio.run(collect_consistency())) diff --git a/scripts/check_nl2sql_e2e.py b/scripts/check_nl2sql_e2e.py new file mode 100644 index 0000000..75da4a2 --- /dev/null +++ b/scripts/check_nl2sql_e2e.py @@ -0,0 +1,147 @@ +"""执行 NL2SQL 真实依赖和元数据链路验收。""" +from __future__ import annotations + +import asyncio +import argparse +import sys +from pathlib import Path +from uuid import uuid4 + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from config import database +from config.database.milvus import client as milvus_client +from config.settings import settings +from nl2sql.contracts import DataQueryRequest +from nl2sql.health import check_nl2sql_health +from nl2sql.milvus_collections import NL2SQL_COLLECTION +from nl2sql.retrieval import retrieve_metadata +from nl2sql.schema import load_authoritative_schema +from service.nl2sql.query_service import execute_query +from service.nl2sql.permission_service import load_query_permission +from tool.llm import llm + + +async def run_real_query(user_id: int, question: str, *, query_id: str | None = None) -> dict: + """使用已有员工权限执行一条真实查询,只输出脱敏统计。""" + from config.database.mysql import get_session_factory + + query_id = query_id or uuid4().hex + try: + async with get_session_factory()() as db: + permission = await load_query_permission(db, user_id) + if not permission.get("can_query"): + return {"query_status": "permission_denied", "user_id": user_id} + + async def permission_loader(_user_id: int): + return permission + + async def metadata_retriever(text_value: str): + return await retrieve_metadata(text_value, milvus_client()) + + async def schema_loader(table_names: set[str], _permission: dict): + return await load_authoritative_schema( + db, + database=settings.mysql.database, + candidate_tables=table_names, + ) + + result = await execute_query( + DataQueryRequest( + question=question, + user_id=user_id, + trace_id=f"nl2sql-e2e-{query_id}", + include_sql=False, + ), + session=db, + query_id=query_id, + permission_loader=permission_loader, + metadata_retriever=metadata_retriever, + schema_loader=schema_loader, + llm_client=llm, + summary_llm=llm, + masks=permission.get("masks"), + ) + return { + "query_status": "success", + "row_count": result.row_count, + "truncated": result.truncated, + "columns": result.columns, + "has_summary": bool(result.summary), + } + except Exception as exc: # noqa: BLE001 验收脚本输出结构化失败摘要 + return {"query_status": "failed", "error_type": type(exc).__name__} + + +def summarize_real_queries(results: list[dict]) -> dict: + """汇总真实查询结果,只保留状态和数量统计。""" + success = [item for item in results if item.get("query_status") == "success"] + return { + "total": len(results), + "success": len(success), + "failed": sum(1 for item in results if item.get("query_status") == "failed"), + "non_empty_success": sum(1 for item in success if item.get("row_count", 0) > 0), + "all_success": bool(results) and len(success) == len(results), + } + + +async def main( + user_id: int | None = None, + question: str | None = None, + questions: list[str] | None = None, +) -> None: + health = await check_nl2sql_health( + { + "mysql": database.mysql.check_health, + "redis": database.redis.check_health, + "milvus": database.milvus.check_health, + "llm": llm.check_health, + } + ) + client = milvus_client() + description = await client.describe_collection(collection_name=NL2SQL_COLLECTION) + rows = await client.query( + collection_name=NL2SQL_COLLECTION, + filter="is_valid == true and is_deprecated == false", + output_fields=["id"], + ) + vector_dim = next( + field["params"]["dim"] + for field in description["fields"] + if field["name"] == "vector" + ) + embedding_dim = None + embedding_error = None + try: + embedding_dim = len(await llm.embed_one("NL2SQL 验收测试")) + except Exception as exc: # noqa: BLE001 记录类型,不输出密钥或请求内容 + embedding_error = type(exc).__name__ + output = { + "health": health, + "collection": NL2SQL_COLLECTION, + "configured_dimension": settings.llm.embed_dimensions, + "collection_dimension": vector_dim, + "valid_vector_count": len(rows or []), + "embedding_dimension": embedding_dim, + "embedding_error": embedding_error, + } + query_questions = questions or ([question] if question else []) + if user_id is not None and query_questions: + query_results = [ + await run_real_query(user_id, item, query_id=f"nl2sql-e2e-{index}-{uuid4().hex[:8]}") + for index, item in enumerate(query_questions, start=1) + ] + output["queries"] = query_results + output["query_summary"] = summarize_real_queries(query_results) + if len(query_results) == 1: + output["query"] = query_results[0] + print(output) + await database.dispose() + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="NL2SQL 真实依赖和查询验收") + parser.add_argument("--user-id", type=int, help="已配置 NL2SQL 权限的员工 ID") + parser.add_argument("--question", action="append", help="真实业务查询问题,可重复传入") + args = parser.parse_args() + asyncio.run(main(args.user_id, questions=args.question)) diff --git a/scripts/check_nl2sql_health.py b/scripts/check_nl2sql_health.py new file mode 100644 index 0000000..6927b65 --- /dev/null +++ b/scripts/check_nl2sql_health.py @@ -0,0 +1,31 @@ +"""命令行检查 NL2SQL 依赖健康状态。""" +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from config import database +from nl2sql.health import check_nl2sql_health +from tool.llm import llm + + +async def check() -> dict: + """检查 MySQL、Redis、Milvus 和 LLM 首选端点。""" + try: + return await check_nl2sql_health( + { + "mysql": database.mysql.check_health, + "redis": database.redis.check_health, + "milvus": database.milvus.check_health, + "llm": llm.check_health, + } + ) + finally: + await database.dispose() + + +if __name__ == "__main__": + print(asyncio.run(check())) diff --git a/scripts/check_nl2sql_permissions.py b/scripts/check_nl2sql_permissions.py new file mode 100644 index 0000000..77002f0 --- /dev/null +++ b/scripts/check_nl2sql_permissions.py @@ -0,0 +1,46 @@ +"""只读检查可用于 NL2SQL 真实验收的员工权限配置。""" +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 import database +from config.database.mysql import get_session_factory + + +async def main() -> None: + queries = { + "employees": text( + "SELECT id, employee_role, status FROM sys_user " + "WHERE user_type = 'EMPLOYEE' ORDER BY id LIMIT 50" + ), + "roles": text( + "SELECT id, employee_role, can_query, max_rows, daily_quota, status " + "FROM nl2sql_query_role ORDER BY id" + ), + "table_permissions": text( + "SELECT role_id, table_name, permission, status " + "FROM nl2sql_role_table_permission ORDER BY role_id, table_name LIMIT 200" + ), + "column_permission_counts": text( + "SELECT role_id, table_name, COUNT(*) AS column_count " + "FROM nl2sql_role_column_permission WHERE status = 'active' " + "GROUP BY role_id, table_name ORDER BY role_id, table_name" + ), + } + try: + async with get_session_factory()() as session: + for name, query in queries.items(): + rows = [dict(row) for row in (await session.execute(query)).mappings()] + print(f"{name}: {rows}") + finally: + await database.mysql.dispose() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/scripts/cleanup_nl2sql_vectors.py b/scripts/cleanup_nl2sql_vectors.py new file mode 100644 index 0000000..a34f2de --- /dev/null +++ b/scripts/cleanup_nl2sql_vectors.py @@ -0,0 +1,23 @@ +"""清理 Milvus 中无效或已废弃的 NL2SQL 向量。""" +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from config import database +from config.database.milvus import client +from nl2sql.operations import cleanup_invalid_vectors + + +async def main() -> None: + try: + print(f"deleted {await cleanup_invalid_vectors(client())} invalid NL2SQL vectors") + finally: + await database.milvus.dispose() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/scripts/evaluate_nl2sql.py b/scripts/evaluate_nl2sql.py new file mode 100644 index 0000000..9b7f884 --- /dev/null +++ b/scripts/evaluate_nl2sql.py @@ -0,0 +1,35 @@ +"""执行 NL2SQL 离线 Golden Case 评测。""" +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from nl2sql.evaluation import build_evaluation_report, load_cases + + +def main() -> None: + parser = argparse.ArgumentParser(description="NL2SQL 离线质量评测") + parser.add_argument( + "--cases", + default=str(Path(__file__).resolve().parents[1] / "tests" / "fixtures" / "nl2sql_golden_cases.json"), + help="Golden Case JSON 文件路径", + ) + parser.add_argument("--prompt-version", default="unknown") + parser.add_argument("--semantic-version", default="unknown") + parser.add_argument("--model-version", default="unknown") + args = parser.parse_args() + report = build_evaluation_report( + load_cases(args.cases), + prompt_version=args.prompt_version, + semantic_version=args.semantic_version, + model_version=args.model_version, + ) + print(json.dumps(report, ensure_ascii=False)) + + +if __name__ == "__main__": + main() diff --git a/scripts/init_nl2sql_milvus.py b/scripts/init_nl2sql_milvus.py new file mode 100644 index 0000000..cc7bf95 --- /dev/null +++ b/scripts/init_nl2sql_milvus.py @@ -0,0 +1,19 @@ +"""初始化 NL2SQL 的 Milvus 数据库和集合。""" +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from config.database.milvus import client +from nl2sql.milvus_collections import ensure_nl2sql_collection + + +async def initialize() -> None: + await ensure_nl2sql_collection(client()) + + +if __name__ == "__main__": + asyncio.run(initialize()) diff --git a/scripts/recreate_nl2sql_milvus.py b/scripts/recreate_nl2sql_milvus.py new file mode 100644 index 0000000..6769391 --- /dev/null +++ b/scripts/recreate_nl2sql_milvus.py @@ -0,0 +1,20 @@ +"""按当前 Embedding 配置重建 NL2SQL Milvus 集合。""" +from __future__ import annotations + +import asyncio +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from config.database.milvus import client +from nl2sql.milvus_collections import recreate_nl2sql_collection + + +async def recreate() -> None: + """删除旧集合并创建当前配置对应的新集合。""" + await recreate_nl2sql_collection(client()) + + +if __name__ == "__main__": + asyncio.run(recreate()) diff --git a/scripts/run_nl2sql_maintenance.py b/scripts/run_nl2sql_maintenance.py new file mode 100644 index 0000000..26eb9ff --- /dev/null +++ b/scripts/run_nl2sql_maintenance.py @@ -0,0 +1,58 @@ +"""运行可由 Cron 或任务平台调用的 NL2SQL 运维任务。""" +from __future__ import annotations + +import argparse +import asyncio +import json +import sys +from datetime import datetime, timedelta, timezone +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from config.deps import get_redis +from nl2sql.jobs import run_consistency_check, run_history_cleanup, run_metadata_sync, run_vector_cleanup + + +async def run(task: str) -> dict: + redis = next(get_redis()) + if task == "metadata_sync": + result = await run_metadata_sync(redis=redis) + elif task == "vector_cleanup": + result = await run_vector_cleanup(redis=redis) + elif task == "consistency_check": + from scripts.check_nl2sql_consistency import collect_consistency + + result = await run_consistency_check(redis=redis, worker=collect_consistency) + elif task == "history_cleanup": + from config.database.mysql import get_session_factory + + async with get_session_factory()() as db: + result = await run_history_cleanup( + db, + datetime.now(timezone.utc).replace(tzinfo=None) - timedelta(days=180), + redis=redis, + ) + else: + raise ValueError("不支持的 NL2SQL 运维任务") + return { + "name": result.name, + "status": result.status, + "attempts": result.attempts, + "detail": result.detail, + "error_type": result.error_type, + } + + +def main() -> None: + parser = argparse.ArgumentParser(description="NL2SQL 运维任务") + parser.add_argument( + "task", + choices=["metadata_sync", "vector_cleanup", "consistency_check", "history_cleanup"], + ) + args = parser.parse_args() + print(json.dumps(asyncio.run(run(args.task)), ensure_ascii=False)) + + +if __name__ == "__main__": + main() diff --git a/scripts/seed_nl2sql_acceptance.py b/scripts/seed_nl2sql_acceptance.py new file mode 100644 index 0000000..82ba92b --- /dev/null +++ b/scripts/seed_nl2sql_acceptance.py @@ -0,0 +1,53 @@ +"""执行 NL2SQL 验收账号和最小权限种子 SQL。""" +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.mysql import dispose, get_session_factory + + +SQL_PATH = Path(__file__).resolve().parents[1] / "sql" / "nl2sql_acceptance_seed.sql" + + +async def main() -> None: + """执行幂等种子文件并输出非敏感核验信息。""" + statements = [item.strip() for item in SQL_PATH.read_text(encoding="utf-8").split(";") if item.strip()] + async with get_session_factory()() as db: + for statement in statements: + await db.execute(text(statement)) + await db.commit() + user_result = await db.execute( + text( + "SELECT id, username, user_type, employee_role, status " + "FROM sys_user WHERE username = 'nl2sql_acceptance'" + ) + ) + role_result = await db.execute( + text( + "SELECT id, role_code, can_query, status " + "FROM nl2sql_query_role WHERE role_code = 'nl2sql_acceptance'" + ) + ) + permission_result = await db.execute( + text( + "SELECT table_name, COUNT(*) AS column_count " + "FROM nl2sql_role_column_permission " + "WHERE role_id = (SELECT id FROM nl2sql_query_role " + "WHERE role_code = 'nl2sql_acceptance') AND status = 'active' " + "GROUP BY table_name ORDER BY table_name" + ) + ) + print({"user": [dict(row) for row in user_result.mappings().all()]}) + print({"role": [dict(row) for row in role_result.mappings().all()]}) + print({"permissions": [dict(row) for row in permission_result.mappings().all()]}) + await dispose() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/scripts/sync_nl2sql_metadata.py b/scripts/sync_nl2sql_metadata.py new file mode 100644 index 0000000..7c3ba9a --- /dev/null +++ b/scripts/sync_nl2sql_metadata.py @@ -0,0 +1,62 @@ +"""读取当前 MySQL 元数据并同步到 Milvus。""" +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 milvus as milvus_db +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 + + +TABLES_SQL = text( + """ + SELECT TABLE_NAME, TABLE_COMMENT, TABLE_TYPE + FROM information_schema.tables + WHERE TABLE_SCHEMA = :database + """ +) + +COLUMNS_SQL = text( + """ + SELECT TABLE_NAME, COLUMN_NAME, COLUMN_COMMENT, DATA_TYPE, + IS_NULLABLE, ORDINAL_POSITION + FROM information_schema.columns + WHERE TABLE_SCHEMA = :database + ORDER BY TABLE_NAME, ORDINAL_POSITION + """ +) + + +async def load_information_schema() -> tuple[list[dict], list[dict]]: + async with get_session_factory()() as session: + tables = [ + dict(row) + for row in (await session.execute(TABLES_SQL, {"database": settings.mysql.database})).mappings() + ] + columns = [ + dict(row) + for row in (await session.execute(COLUMNS_SQL, {"database": settings.mysql.database})).mappings() + ] + return tables, columns + + +async def synchronize() -> int: + try: + tables, columns = await load_information_schema() + return await sync_metadata(milvus_client(), tables, columns) + finally: + await mysql.dispose() + await milvus_db.dispose() + + +if __name__ == "__main__": + print(f"upserted {asyncio.run(synchronize())} NL2SQL metadata chunks") diff --git a/service/advisor/agent_client.py b/service/advisor/agent_client.py index 2a91693..d7e20bd 100644 --- a/service/advisor/agent_client.py +++ b/service/advisor/agent_client.py @@ -169,7 +169,7 @@ class AdvisorAgentClient: ) -> dict: return await self._request( "POST", f"/draft/{draft_id}/operate", auth_header=auth_header, - trace_id=trace_id, json={"action": action}, + trace_id=trace_id, json={"operation": action}, ) async def rebalance_run( diff --git a/service/advisor/audit.py b/service/advisor/audit.py index 28c0c61..be67d6b 100644 --- a/service/advisor/audit.py +++ b/service/advisor/audit.py @@ -29,6 +29,7 @@ async def list_ledger( user: SysUser, *, action: str | None = None, + customer_id: int | None = None, keyword: str | None = None, start: datetime | None = None, end: datetime | None = None, @@ -37,7 +38,7 @@ async def list_ledger( ) -> dict: repo = AuditLogRepo(db) items = await repo.list_by_advisor( - user_id=user.id, action=action, keyword=keyword, start=start, end=end, + user_id=user.id, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end, limit=page_size, offset=(page - 1) * page_size, ) total = await repo.count_by_advisor( @@ -62,6 +63,7 @@ async def export_ledger( user: SysUser, *, action: str | None = None, + customer_id: int | None = None, keyword: str | None = None, start: datetime | None = None, end: datetime | None = None, @@ -69,7 +71,7 @@ async def export_ledger( """导出本人审计台账为 CSV 文本(全量,不限分页大小,上限 10000 条)。""" repo = AuditLogRepo(db) items = await repo.list_by_advisor( - user_id=user.id, action=action, keyword=keyword, start=start, end=end, + user_id=user.id, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end, limit=10000, offset=0, ) return _to_csv([_audit_item(a) for a in items]) diff --git a/service/advisor/customers.py b/service/advisor/customers.py index bed7e9f..18d8a05 100644 --- a/service/advisor/customers.py +++ b/service/advisor/customers.py @@ -138,7 +138,10 @@ async def get_customer_reports( await ensure_customer_owned(db, user.id, customer_id) repo = AdvisorReportRepo(db) items = await repo.list_by_customer( - customer_id, limit=page_size, offset=(page - 1) * page_size + customer_id, + advisor_id=user.id, + limit=page_size, + offset=(page - 1) * page_size, ) return { "total": await repo.count_by_advisor(advisor_id=user.id, customer_id=customer_id), diff --git a/service/advisor/dashboard.py b/service/advisor/dashboard.py index ce69dbf..0eb0849 100644 --- a/service/advisor/dashboard.py +++ b/service/advisor/dashboard.py @@ -21,6 +21,17 @@ from repositories.advisor_visit_record import AdvisorVisitRecordRepo from repositories.customer_relation import CustomerRelationRepo +def _todo_summary(todo) -> dict: + return { + "id": todo.id, + "todo_type": todo.todo_type, + "customer_id": todo.customer_id, + "priority": todo.priority, + "status": todo.status, + "due_at": todo.due_at.isoformat() if todo.due_at else None, + } + + async def get_dashboard(db: AsyncSession, user: SysUser) -> dict: relation_repo = CustomerRelationRepo(db) todo_repo = AdvisorTodoRepo(db) @@ -28,6 +39,9 @@ async def get_dashboard(db: AsyncSession, user: SysUser) -> dict: visit_repo = AdvisorVisitRecordRepo(db) pending = await todo_repo.count_by_advisor(advisor_id=user.id, status=TODO_STATUS_PENDING) + pending_items = await todo_repo.list_by_advisor( + advisor_id=user.id, status=TODO_STATUS_PENDING, limit=10, offset=0 + ) total = await relation_repo.count_customer_rows(advisor_id=user.id) signed = await relation_repo.count_customer_rows( advisor_id=user.id, status=CUSTOMER_REL_STATUS_SIGNED @@ -54,7 +68,10 @@ async def get_dashboard(db: AsyncSession, user: SysUser) -> dict: unknown += cnt return { - "todos": {"pending": pending}, + "todos": { + "pending": pending, + "items": [_todo_summary(todo) for todo in pending_items], + }, "overview": { "total_customers": total, "signed_customers": signed, diff --git a/service/advisor/drafts.py b/service/advisor/drafts.py index 53e3fdb..c57e5a7 100644 --- a/service/advisor/drafts.py +++ b/service/advisor/drafts.py @@ -29,6 +29,7 @@ from repositories.advisor_report import AdvisorReportRepo from repositories.product import ProductRepo from repositories.risk_assessment import CustomerProfileRepo from repositories.sensitive_word import SensitiveWordRepo +from repositories.sys_message import SysMessageRepo from schemas.advisor import DraftSaveReq, RebalanceRunReq, TalkScriptReq from service.advisor.agent_client import get_agent_client from service.advisor.permissions import ensure_customer_owned, require_owned_relation @@ -60,20 +61,30 @@ def _build_report_from_detail(detail: dict, advisor_id: int) -> AdvisorReport: ) -def _extract_buy_codes(suggestions: Any) -> list[str]: - """从建议清单提取「低配申购」产品代码(发送终审适当性校验对象)。 +def _extract_buy_codes(detail: Any) -> list[str]: + """从 Agent 草稿详情提取发送终审需要校验的产品代码。 - 假设:结构为 dict,申购侧键见 _BUY_KEYS;项为 {product_code} 或 {code}/{fund_code}。 - 解析失败返回空列表(视为无结构化产品建议,适当性空过,不误拦)。 + 兼容旧的顶层 ``suggestions``,以及 Agent 当前的 ``structured_data``:调仓只取 + ``buy``,推荐取 ``items``。解析失败返回空列表,交由其它终审规则继续处理。 """ - if not isinstance(suggestions, dict): + if not isinstance(detail, dict): return [] + + structured_data = detail.get("structured_data") + data = structured_data if isinstance(structured_data, dict) else detail + suggestions = data.get("suggestions") + if isinstance(suggestions, dict): + data = suggestions + items: list | None = None for key in _BUY_KEYS: - value = suggestions.get(key) + value = data.get(key) if isinstance(value, list): items = value break + if items is None and isinstance(data.get("items"), list): + items = data["items"] + codes: list[str] = [] for it in items or []: if isinstance(it, dict): @@ -104,7 +115,7 @@ async def _resolve_product_risks( result = await get_agent_client().draft_detail( draft_id, auth_header=auth_header, trace_id=trace_id ) - codes = _extract_buy_codes((result["data"] or {}).get("suggestions")) + codes = _extract_buy_codes(result["data"] or {}) product_repo = ProductRepo(db) risks: list[str | None] = [] for code in codes: @@ -214,7 +225,11 @@ async def save_draft( # 1) 先取草稿确认归属(避免对无权限草稿执行写操作),并拿到 intent/customer_id detail = await get_draft(db, user, auth_header=auth_header, trace_id=trace_id, draft_id=draft_id) # 2) 调 Agent 保存(Agent 重新适当性校验,违规 40020;缺免责仅告警不阻断) - payload = {"title": req.title, "content": req.content, "suggestions": req.suggestions} + payload = { + "title": req.title, + "content": req.content, + "structured_data": req.suggestions, + } result = await get_agent_client().draft_save( draft_id, payload, auth_header=auth_header, trace_id=trace_id ) @@ -262,6 +277,9 @@ async def send_draft( elif report.advisor_id != user.id: raise ForbiddenError("无权操作该客户数据") + if report.send_status == REPORT_SEND_STATUS_DISCARDED: + raise ParamError("已废弃的报告不可发送") + # 幂等:已发送直接返回(重复点击不重复发站内信) if report.send_status == REPORT_SEND_STATUS_SENT: return {**_sent_payload(report), "duplicated": True} @@ -283,6 +301,19 @@ async def send_draft( sensitive_words=sensitive_words, ) + # 提交结果不确定后重试时,先复用已经写入的站内信,避免重复触达客户。 + existing_message = await SysMessageRepo(db).get_by_biz_id( + report.report_id, user_id=report.customer_id + ) + if existing_message is not None: + report.send_status = REPORT_SEND_STATUS_SENT + report.send_time = report.send_time or existing_message.create_time + report.send_by = report.send_by or user.id + report.msg_id = existing_message.id + db.add(report) + await db.commit() + return {**_sent_payload(report), "duplicated": True} + # 写站内信 + 报告置 sent,同事务(失败可重试,避免假送达) msg_type = INTENT_TO_MSG_TYPE.get(report.intent, MSG_TYPE_RECOMMEND) message = SysMessage( @@ -297,9 +328,13 @@ async def send_draft( report.send_by = user.id db.add(message) db.add(report) # 已跟踪对象时无副作用,新对象时入 session - await db.flush() # 生成 message.id,供回填 msg_id - report.msg_id = message.id - await db.commit() + try: + await db.flush() # 生成 message.id,供回填 msg_id + report.msg_id = message.id + await db.commit() + except Exception: + await db.rollback() + raise return _sent_payload(report) diff --git a/service/advisor/event_consumer.py b/service/advisor/event_consumer.py index 25518e5..b2ec86f 100644 --- a/service/advisor/event_consumer.py +++ b/service/advisor/event_consumer.py @@ -18,20 +18,31 @@ from common_const import ( TODO_SOURCE_AGENT_EVENT, TODO_TYPE_NEW_REBALANCE_DRAFT, ) +from config.database.redis import client as redis_client +from config.settings import settings from model.advisor_todo import AdvisorTodo from repositories.advisor_todo import AdvisorTodoRepo from repositories.event_log import EventLogRepo +from service.memory.profile import CustomerProfileMemory logger = logging.getLogger("service.advisor.event_consumer") +def _payload_log_summary(payload: dict) -> str: + """日志只记录字段名,避免事件扩展后误打客户敏感值。""" + return "{" + ", ".join(sorted(str(key) for key in payload)) + "}" + + async def handle_rebalance_created(db, payload: dict) -> None: """消费 rebalance_draft_created → 生成待办(uk_todo 三元组去重,幂等)。""" advisor_id = payload.get("advisor_id") customer_id = payload.get("customer_id") draft_id = payload.get("draft_id") - if not advisor_id: - logger.warning("rebalance_draft_created 缺少 advisor_id,忽略: %s", payload) + if not advisor_id or customer_id is None or not draft_id: + logger.warning( + "rebalance_draft_created 缺少 advisor_id/customer_id/draft_id,忽略字段: %s", + _payload_log_summary(payload), + ) return repo = AdvisorTodoRepo(db) existing = await repo.get_by_unique(TODO_TYPE_NEW_REBALANCE_DRAFT, customer_id, draft_id) @@ -47,16 +58,23 @@ async def handle_rebalance_created(db, payload: dict) -> None: await repo.add(todo) # add 内 commit + refresh -async def handle_profile_update(db, payload: dict) -> None: - """消费 profile_update → 失效该客户 Redis 画像缓存(不落库,画像只读)。 +async def handle_profile_update(db, payload: dict, *, redis=None) -> None: + """消费 profile_update,失效该客户 Redis 画像缓存。""" + customer_id = payload.get("customer_id") + if customer_id is None: + logger.warning("profile_update 缺少 customer_id,忽略字段: %s", _payload_log_summary(payload)) + return - 注:V1.0 工作台尚未建设 profile:{customer_id} 画像热缓存读取,此处仅占位;后续 - 360 读画像接入缓存后生效。事件仍需标记已消费。 - """ - return + warnings = await CustomerProfileMemory(redis=redis or redis_client()).invalidate( + int(customer_id) + ) + for warning in warnings: + logger.warning("profile_update cache invalidation warning: %s", warning) -async def process_event(db, *, event_id: str | None, event_name: str | None, payload: dict) -> None: +async def process_event( + db, *, event_id: str | None, event_name: str | None, payload: dict, redis=None +) -> None: """处理单条事件(幂等入口,实时订阅与补拉共用)。""" # 幂等兜底:event_log 已消费则跳过(实时订阅与补拉并发时防重) if event_id: @@ -67,7 +85,7 @@ async def process_event(db, *, event_id: str | None, event_name: str | None, pay if event_name == EVENT_REBALANCE_DRAFT_CREATED: await handle_rebalance_created(db, payload) elif event_name == EVENT_PROFILE_UPDATE: - await handle_profile_update(db, payload) + await handle_profile_update(db, payload, redis=redis) if event_id: await EventLogRepo(db).mark_consumed(event_id) @@ -92,12 +110,16 @@ def parse_message(raw: str) -> tuple[str | None, str | None, dict] | None: return event_id, event_name, payload -async def pull_pending_events(db) -> int: +async def pull_pending_events(db, *, redis=None) -> int: """补拉 event_log 中未消费的投顾域事件(Pub/Sub 丢消息兜底),返回处理条数。""" events = await EventLogRepo(db).list_pending(ADVISOR_EVENTS, limit=100) for ev in events: await process_event( - db, event_id=ev.event_id, event_name=ev.event_name, payload=ev.payload or {} + db, + event_id=ev.event_id, + event_name=ev.event_name, + payload=ev.payload or {}, + redis=redis, ) return len(events) @@ -105,9 +127,34 @@ async def pull_pending_events(db) -> int: class EventConsumer: """Redis 订阅循环(后台任务,dev 默认关闭,避免 --reload 重复订阅)。""" - def __init__(self, redis): + def __init__(self, redis, *, retry_interval_sec: float | None = None): self.redis = redis self._task: asyncio.Task | None = None + self._retry_task: asyncio.Task | None = None + self.retry_interval_sec = ( + settings.advisor.event_retry_interval_sec + if retry_interval_sec is None + else retry_interval_sec + ) + + async def retry_pending_once(self) -> int: + """补拉并处理一批尚未消费事件;异常保留 pending,供下一轮重试。""" + from config.database.mysql import get_session_factory + + async with get_session_factory()() as session: + return await pull_pending_events(session, redis=self.redis) + + async def _retry_loop(self) -> None: + while True: + await asyncio.sleep(self.retry_interval_sec) + try: + count = await self.retry_pending_once() + if count: + logger.info("advisor event pending retry processed: %s", count) + except asyncio.CancelledError: + raise + except Exception: + logger.exception("advisor event pending retry failed") async def _run(self) -> None: from config.database.mysql import get_session_factory @@ -128,7 +175,11 @@ class EventConsumer: async with get_session_factory()() as session: try: await process_event( - session, event_id=event_id, event_name=event_name, payload=payload + session, + event_id=event_id, + event_name=event_name, + payload=payload, + redis=self.redis, ) except Exception: logger.exception("process event failed: %s %s", event_name, event_id) @@ -139,12 +190,15 @@ class EventConsumer: def start(self) -> None: if self._task is None: self._task = asyncio.create_task(self._run()) + self._retry_task = asyncio.create_task(self._retry_loop()) async def stop(self) -> None: - if self._task is not None: - self._task.cancel() - try: - await self._task - except asyncio.CancelledError: - pass - self._task = None + for task in (self._task, self._retry_task): + if task is not None: + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + self._task = None + self._retry_task = None diff --git a/service/advisor/visits.py b/service/advisor/visits.py index 8c5f143..f1a0429 100644 --- a/service/advisor/visits.py +++ b/service/advisor/visits.py @@ -11,9 +11,10 @@ from model.advisor_visit_record import AdvisorVisitRecord from model.sys_user import SysUser from repositories.advisor_visit_record import AdvisorVisitRecordRepo from repositories.sys_message import SysMessageRepo -from schemas.advisor import VisitCreateReq +from schemas.advisor import VisitCreateReq, VisitUpdateReq from service.advisor.permissions import ensure_customer_owned -from utils.exceptions import ParamError +from service.advisor.audit_writer import write_audit +from utils.exceptions import NotFoundError, ParamError # 合规话术库(内置,标准化投教/市场解读/调仓沟通)。 # 假设:V1.0 内置常量,后续可迁移 sys_config 运营化;内容不含承诺收益等敏感词。 @@ -63,9 +64,50 @@ async def create_visit(db: AsyncSession, user: SysUser, req: VisitCreateReq) -> audio_url=req.audio_url, ) record = await AdvisorVisitRecordRepo(db).add(record) + await write_audit( + db, + user_id=user.id, + username=user.username, + module="advisor", + action="visit_create", + target=str(record.id), + detail={"customer_id": req.customer_id}, + ) return {"visit_id": record.id} +async def get_visit(db: AsyncSession, user: SysUser, visit_id: int) -> dict: + record = await AdvisorVisitRecordRepo(db).get_by_advisor(visit_id, user.id) + if record is None: + raise NotFoundError("回访记录不存在") + return _visit_item(record) + + +async def update_visit( + db: AsyncSession, user: SysUser, visit_id: int, req: VisitUpdateReq +) -> dict: + changes = req.model_dump(exclude_unset=True) + if not changes: + raise ParamError("至少提供一项回访记录修改内容") + record = await AdvisorVisitRecordRepo(db).get_by_advisor(visit_id, user.id) + if record is None: + raise NotFoundError("回访记录不存在") + for field, value in changes.items(): + setattr(record, field, value) + await db.commit() + await db.refresh(record) + await write_audit( + db, + user_id=user.id, + username=user.username, + module="advisor", + action="visit_update", + target=str(record.id), + detail={"customer_id": record.customer_id, "fields": list(changes)}, + ) + return _visit_item(record) + + def list_talk_templates() -> list[dict]: """合规话术库(内置,投顾参考;不自动发送)。""" return TALK_TEMPLATES diff --git a/service/advisor_agent/__init__.py b/service/advisor_agent/__init__.py new file mode 100644 index 0000000..1dd717d --- /dev/null +++ b/service/advisor_agent/__init__.py @@ -0,0 +1 @@ +"""投顾 Agent 业务服务。""" diff --git a/service/advisor_agent/audit.py b/service/advisor_agent/audit.py new file mode 100644 index 0000000..4553b06 --- /dev/null +++ b/service/advisor_agent/audit.py @@ -0,0 +1,50 @@ +"""投顾 Agent 审计日志写入。""" +from __future__ import annotations + +import json + +from sqlalchemy import text + +from common.common_const import AUDIT_AGENT_CHAT_CALL, AUDIT_DRAFT_DISCARD, AUDIT_DRAFT_SAVE + + +def audit_action_for_path(path: str) -> str: + if path.endswith("/save"): + return AUDIT_DRAFT_SAVE + if path.endswith("/operate"): + return AUDIT_DRAFT_DISCARD + return AUDIT_AGENT_CHAT_CALL + + +async def write_advisor_audit( + db, + *, + user, + action: str, + target: str | None, + trace_id: str, + detail: dict | None = None, + status: str = "成功", +) -> None: + statement = text( + """ + INSERT INTO audit_log + (user_id, username, module, action, target, detail, trace_id, status) + VALUES + (:user_id, :username, :module, :action, :target, :detail, :trace_id, :status) + """ + ) + await db.execute( + statement, + { + "user_id": getattr(user, "id", None), + "username": getattr(user, "username", None), + "module": "advisor_agent", + "action": action, + "target": target, + "detail": json.dumps(detail or {}, ensure_ascii=False), + "trace_id": trace_id, + "status": status, + }, + ) + await db.commit() diff --git a/service/advisor_agent/compliance.py b/service/advisor_agent/compliance.py new file mode 100644 index 0000000..734391e --- /dev/null +++ b/service/advisor_agent/compliance.py @@ -0,0 +1,16 @@ +"""投顾 Agent 输出合规校验。""" +from __future__ import annotations + +from common.common_const import ERR_CODE_LLM_ERROR +from utils.exceptions import ApiError + + +def find_sensitive_words(content: str, words: list[str]) -> list[str]: + return [word for word in dict.fromkeys(words) if word and word in content] + + +def ensure_safe_content(content: str, words: list[str]) -> bool: + matched = find_sensitive_words(content, words) + if matched: + raise ApiError(ERR_CODE_LLM_ERROR, "AI输出包含敏感或违规表述") + return True diff --git a/service/advisor_agent/context.py b/service/advisor_agent/context.py new file mode 100644 index 0000000..7be9791 --- /dev/null +++ b/service/advisor_agent/context.py @@ -0,0 +1,177 @@ +"""从现有业务表聚合投顾意图所需的本地上下文。""" +from __future__ import annotations + +from decimal import Decimal, InvalidOperation + +from common.common_const import ( + CUSTOMER_REL_STATUS_SIGNED, + SYS_KEY_REBALANCE_DEVIATION_THRESHOLD, +) +from repositories.fin_customer_profile import FinCustomerProfileRepo +from repositories.fin_holdings import FinHoldingsRepo +from repositories.fin_product import FinProductRepo +from repositories.portfolio_benchmark import PortfolioBenchmarkRepo +from repositories.risk_assessment import RiskAssessmentRepo +from repositories.sys_config import SysConfigRepo +from model.fin_product import FinProduct +from service.advisor_agent.data import to_holding_input, to_product_candidate +from service.advisor_agent.data import to_fund_performance_row + + +async def _load_customer_risk(db, customer_id: int, profile_repo_cls, risk_repo_cls): + assessment = await risk_repo_cls(db).get_current_by_customer(customer_id) + if assessment is not None and assessment.risk_level: + return assessment.risk_level + profile = await profile_repo_cls(db).get_by_customer_id(customer_id) + return profile.risk_level if profile is not None else None + + +async def load_customer_risk( + db, + *, + customer_id: int, + profile_repo_cls=FinCustomerProfileRepo, + risk_repo_cls=RiskAssessmentRepo, +) -> str | None: + return await _load_customer_risk(db, customer_id, profile_repo_cls, risk_repo_cls) + + +async def load_rebalance_context( + db, + *, + customer_id: int, + profile_repo_cls=FinCustomerProfileRepo, + risk_repo_cls=RiskAssessmentRepo, + holdings_repo_cls=FinHoldingsRepo, + product_repo_cls=FinProductRepo, + benchmark_repo_cls=PortfolioBenchmarkRepo, + sys_config_repo_cls=SysConfigRepo, +) -> dict | None: + """聚合画像、持仓、在售产品和组合基准,供调仓引擎使用。 + + 关系签约状态由调用方读取并校验,本函数只负责客户投资数据。 + """ + customer_risk = await _load_customer_risk( + db, customer_id, profile_repo_cls, risk_repo_cls + ) + if not customer_risk: + return None + + benchmark = await benchmark_repo_cls(db).get_active_by_risk(customer_risk) + if benchmark is None: + return None + + threshold = benchmark.drift_threshold + if threshold is None: + raw_threshold = await sys_config_repo_cls(db).get_value( + SYS_KEY_REBALANCE_DEVIATION_THRESHOLD + ) + try: + threshold = Decimal(str(raw_threshold)) + except (InvalidOperation, TypeError, ValueError): + return None + if not threshold.is_finite() or threshold < 0: + return None + + product_repo = product_repo_cls(db) + products = await product_repo.list( + where=[FinProduct.status == "在售"], + limit=1000, + ) + products_by_id = {product.id: product for product in products} + + holdings = await holdings_repo_cls(db).list_by_customer(customer_id, status="持有中") + holding_inputs = [] + for holding in holdings: + product = products_by_id.get(holding.product_id) + if product is None: + product = await product_repo.get(holding.product_id) + if product is not None: + holding_inputs.append(to_holding_input(holding, product)) + + return { + "customer_id": customer_id, + "customer_risk": customer_risk, + "relation_status": CUSTOMER_REL_STATUS_SIGNED, + "holdings": holding_inputs, + "target_allocation": benchmark.target_allocation, + "threshold": threshold, + "candidates": [to_product_candidate(product) for product in products], + } + + +async def load_fund_analysis_context( + db, + *, + fund_codes: list[str], + product_repo_cls=FinProductRepo, + performance_repo_cls=None, +) -> list[dict]: + if performance_repo_cls is None: + from repositories.fund_performance import FundPerformanceRepo + + performance_repo_cls = FundPerformanceRepo + + product_repo = product_repo_cls(db) + performance_repo = performance_repo_cls(db) + result = [] + for code in fund_codes: + product = await product_repo.get_by_code(code) + if product is None: + continue + performance = await performance_repo.list_for_product(product.id) + result.append( + { + "fund": { + "fund_code": product.product_code, + "fund_name": product.product_name, + "risk_level": product.risk_level, + }, + "performance": [to_fund_performance_row(row) for row in performance], + } + ) + return result + + +async def load_recommendation_context( + db, + *, + customer_id: int, + profile_repo_cls=FinCustomerProfileRepo, + product_repo_cls=FinProductRepo, + performance_repo_cls=None, + risk_repo_cls=RiskAssessmentRepo, +) -> dict | None: + customer_risk = await _load_customer_risk( + db, customer_id, profile_repo_cls, risk_repo_cls + ) + if not customer_risk: + return None + + products = await product_repo_cls(db).list( + where=[FinProduct.status == "在售"], + limit=1000, + ) + if performance_repo_cls is None: + from repositories.fund_performance import FundPerformanceRepo + + performance_repo_cls = FundPerformanceRepo + performance_repo = performance_repo_cls(db) + candidates = [] + for product in products: + candidate = to_product_candidate(product) + latest = None + if hasattr(performance_repo, "get_latest_for_product"): + latest = await performance_repo.get_latest_for_product(product.id) + else: + rows = await performance_repo.list_for_product(product.id) + latest = rows[-1] if rows else None + return_rate = getattr(latest, "return_rate", None) + if return_rate is not None: + candidate["performance_score"] = float(return_rate) + candidates.append(candidate) + return { + "customer_id": customer_id, + "customer_risk": customer_risk, + "candidates": candidates, + } diff --git a/service/advisor_agent/data.py b/service/advisor_agent/data.py new file mode 100644 index 0000000..2f41dc4 --- /dev/null +++ b/service/advisor_agent/data.py @@ -0,0 +1,39 @@ +"""现有业务 ORM 到投顾意图输入的适配。""" +from __future__ import annotations + +from decimal import Decimal + + +def _number(value): + return float(value) if isinstance(value, Decimal) else value + + +def to_fund_performance_row(row) -> dict: + return { + "period": getattr(row, "period", None), + "return_rate": _number(getattr(row, "return_rate", None)), + "annual_volatility": _number(getattr(row, "annual_volatility", None)), + "max_drawdown": _number(getattr(row, "max_drawdown", None)), + "sharpe": _number(getattr(row, "sharpe", None)), + } + + +def to_holding_input(holding, product) -> dict: + return { + "product_id": holding.product_id, + "product_code": product.product_code, + "asset_class": product.product_type, + "market_value": holding.current_value or Decimal("0"), + } + + +def to_product_candidate(product) -> dict: + expected_return = product.expected_return or Decimal("0") + return { + "product_id": product.id, + "product_code": product.product_code, + "product_name": product.product_name, + "asset_class": product.product_type, + "risk_level": product.risk_level, + "performance_score": float(expected_return), + } diff --git a/service/advisor_agent/draft.py b/service/advisor_agent/draft.py new file mode 100644 index 0000000..0866e0a --- /dev/null +++ b/service/advisor_agent/draft.py @@ -0,0 +1,179 @@ +"""投顾 Agent 草稿服务。""" +from __future__ import annotations + +from decimal import Decimal +from typing import Any +import uuid + +from common.common_const import ( + DRAFT_STATUS_DISCARDED, + DRAFT_STATUS_DRAFT, + ERR_CODE_DRAFT_NOT_FOUND, + ERR_CODE_FORBIDDEN_CUSTOMER, + ERR_CODE_SUITABILITY_INVALID, + REPORT_DISCLAIMER, +) +from common.suitability import check_suitability +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 + + +def build_generated_content(content: str) -> str: + """生成阶段强制补齐免责声明,避免重复拼接。""" + if REPORT_DISCLAIMER in content: + return content + return f"{content.rstrip()}\n\n{REPORT_DISCLAIMER}" + + +def _disclaimer_warning(content: str) -> tuple[bool, str | None]: + if REPORT_DISCLAIMER in content: + return True, None + return False, "草稿缺少完整免责声明,工作台发送前必须补齐并重新校验" + + +def _not_found() -> ApiError: + return ApiError(ERR_CODE_DRAFT_NOT_FOUND, "草稿不存在或者已废弃") + + +async def _resolve_sensitive_words(repo, explicit_words: list[str] | None) -> list[str]: + if explicit_words is not None: + return explicit_words + if hasattr(repo, "db"): + return await SensitiveWordRepo(repo.db).list_active_words() + return [] + + +def _validate_structured_suitability( + structured_data: dict | None, customer_risk: str | None = None +) -> None: + if not structured_data: + return + customer_risk = customer_risk or structured_data.get("customer_risk") + if not customer_risk: + return + for item in structured_data.get("items", []): + result = check_suitability(customer_risk, item.get("risk_level", "")) + if not result.ok: + raise ApiError(ERR_CODE_SUITABILITY_INVALID, result.reason) + + +async def create_draft(repo, data: dict) -> AdvisorDraft: + content = data.get("content", "") + sensitive_words = await _resolve_sensitive_words(repo, data.get("sensitive_words")) + ensure_safe_content(content, sensitive_words) + disclaimer_ok, warning = _disclaimer_warning(content) + draft = AdvisorDraft( + draft_id=uuid.uuid4().hex, + customer_id=data["customer_id"], + advisor_id=data["advisor_id"], + intent=data["intent"], + title=data["title"], + content=content, + structured_data=data.get("structured_data"), + status=DRAFT_STATUS_DRAFT, + deviation=data.get("deviation"), + disclaimer_ok=disclaimer_ok, + warning=warning, + ) + return await repo.add(draft) + + +def ensure_draft_owner(draft: Any, *, advisor_id: int) -> None: + if draft.advisor_id != advisor_id: + raise ApiError(ERR_CODE_FORBIDDEN_CUSTOMER, "无权操作该客户数据") + + +def summarize_draft(draft: Any) -> dict: + deviation = getattr(draft, "deviation", None) + if isinstance(deviation, Decimal): + deviation = float(deviation) + return { + "draft_id": draft.draft_id, + "customer_id": draft.customer_id, + "advisor_id": draft.advisor_id, + "intent": getattr(draft, "intent", None), + "title": getattr(draft, "title", None), + "status": draft.status, + "deviation": deviation, + "disclaimer_ok": bool(getattr(draft, "disclaimer_ok", False)), + "created_at": draft.create_time.isoformat() + if getattr(draft, "create_time", None) + else None, + "update_time": draft.update_time.isoformat() + if getattr(draft, "update_time", None) + else None, + } + + +def detail_draft(draft: Any) -> dict: + result = summarize_draft(draft) + result.update( + { + "content": draft.content, + "structured_data": getattr(draft, "structured_data", None), + "warning": getattr(draft, "warning", None), + } + ) + return result + + +async def get_draft(repo, draft_id: str): + draft = await repo.get_by_draft_id(draft_id) + if draft is None or draft.status == DRAFT_STATUS_DISCARDED: + raise _not_found() + return draft + + +async def list_drafts( + repo, + *, + advisor_id: int | None = None, + customer_id: int | None = None, + status: str | None = None, + page: int = 1, + page_size: int = 20, +) -> dict: + page = max(1, page) + page_size = min(100, max(1, 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, + ) + return {"total": total, "items": [summarize_draft(item) for item in items]} + + +async def save_draft( + repo, + draft_id: str, + *, + title: str | None = None, + content: str | None = None, + structured_data: dict | None = None, + customer_risk: str | None = None, + sensitive_words: list[str] | None = None, +): + draft = await get_draft(repo, draft_id) + if draft.status != DRAFT_STATUS_DRAFT: + raise _not_found() + if title is not None: + draft.title = title + if content is not None: + draft.content = content + if structured_data is not None: + draft.structured_data = structured_data + resolved_sensitive_words = await _resolve_sensitive_words(repo, sensitive_words) + ensure_safe_content(draft.content, resolved_sensitive_words) + _validate_structured_suitability(draft.structured_data, customer_risk) + draft.disclaimer_ok, draft.warning = _disclaimer_warning(draft.content) + saved = await repo.save(draft) + return detail_draft(saved) + + +async def discard_draft(repo, draft_id: str): + draft = await get_draft(repo, draft_id) + return await repo.discard(draft) diff --git a/service/event_publisher.py b/service/event_publisher.py new file mode 100644 index 0000000..abf052e --- /dev/null +++ b/service/event_publisher.py @@ -0,0 +1,63 @@ +"""跨模块事件双写发布器。""" +from __future__ import annotations + +import json +import uuid + +from sqlalchemy import text + +from common_const import EVENT_STATUS_PENDING + + +async def publish_event( + redis, + db, + *, + event_id: str | None = None, + event_name: str, + trace_id: str, + trigger_user_id: int | None, + customer_id: int | None, + payload: dict, +) -> str: + requested_event_id = event_id + event_id = event_id or uuid.uuid4().hex + envelope = { + "event_id": event_id, + "event_name": event_name, + "trace_id": trace_id, + "trigger_user_id": trigger_user_id, + "customer_id": customer_id, + "payload": payload, + } + if requested_event_id: + existing = await db.scalar( + text("SELECT event_id FROM event_log WHERE event_id = :event_id"), + {"event_id": event_id}, + ) + if existing: + await redis.publish(event_name, json.dumps(envelope, ensure_ascii=False)) + return event_id + statement = text( + """ + INSERT INTO event_log + (event_id, event_name, trace_id, trigger_user_id, customer_id, payload, status) + VALUES + (:event_id, :event_name, :trace_id, :trigger_user_id, :customer_id, :payload, :status) + """ + ) + await db.execute( + statement, + { + "event_id": event_id, + "event_name": event_name, + "trace_id": trace_id, + "trigger_user_id": trigger_user_id, + "customer_id": customer_id, + "payload": json.dumps(payload, ensure_ascii=False), + "status": EVENT_STATUS_PENDING, + }, + ) + await db.commit() + await redis.publish(event_name, json.dumps(envelope, ensure_ascii=False)) + return event_id diff --git a/service/nl2sql/__init__.py b/service/nl2sql/__init__.py new file mode 100644 index 0000000..ed5642d --- /dev/null +++ b/service/nl2sql/__init__.py @@ -0,0 +1 @@ +"""NL2SQL 查询服务。""" diff --git a/service/nl2sql/admin_service.py b/service/nl2sql/admin_service.py new file mode 100644 index 0000000..2ec89ab --- /dev/null +++ b/service/nl2sql/admin_service.py @@ -0,0 +1,105 @@ +"""NL2SQL 管理能力服务。""" +from __future__ import annotations + +from model.nl2sql_permission import ( + Nl2SqlQueryRole, + Nl2SqlRoleColumnPermission, + Nl2SqlRoleTablePermission, + Nl2SqlSensitiveField, +) +from repositories.nl2sql_permission import Nl2SqlPermissionRepo + +VALID_MASK_TYPES = {"partial", "hash"} +VALID_ROW_SCOPE_TYPES = {"none", "customer_ids", "product_ids"} + + +def _clean_text(value: str, name: str) -> str: + cleaned = value.strip() + if not cleaned: + raise ValueError(f"{name}不能为空") + return cleaned + + +def validate_mask_type(mask_type: str | None) -> str | None: + if mask_type is not None and mask_type not in VALID_MASK_TYPES: + raise ValueError("脱敏类型必须为 partial 或 hash") + return mask_type + + +def validate_table_permission(payload: dict) -> dict: + permission = payload.get("permission", "SELECT") + if permission != "SELECT": + raise ValueError("表权限只允许 SELECT") + row_scope_type = payload.get("row_scope_type", "none") + if row_scope_type not in VALID_ROW_SCOPE_TYPES: + raise ValueError("行级范围类型无效") + return { + "table_name": _clean_text(payload["table_name"], "表名"), + "permission": "SELECT", + "row_scope_type": row_scope_type, + "row_scope_column": payload.get("row_scope_column"), + "status": payload.get("status", "active"), + } + + +async def create_role(db, payload: dict, *, repo_factory=Nl2SqlPermissionRepo): + values = { + "role_code": _clean_text(payload["role_code"], "角色编码"), + "role_name": _clean_text(payload["role_name"], "角色名称"), + "employee_role": _clean_text(payload["employee_role"], "员工角色"), + "can_query": bool(payload.get("can_query", False)), + "max_rows": int(payload.get("max_rows", 1000)), + "daily_quota": int(payload.get("daily_quota", 0)), + "status": "active", + } + if values["max_rows"] <= 0 or values["daily_quota"] < 0: + raise ValueError("配额参数无效") + return await repo_factory(db).add_role(**values) + + +def role_payload(role: Nl2SqlQueryRole) -> dict: + return { + "id": role.id, + "role_code": role.role_code, + "role_name": role.role_name, + "employee_role": role.employee_role, + "can_query": role.can_query, + "max_rows": role.max_rows, + "daily_quota": role.daily_quota, + "status": role.status, + } + + +def table_permission_payload(item: Nl2SqlRoleTablePermission) -> dict: + return { + "id": item.id, + "role_id": item.role_id, + "table_name": item.table_name, + "permission": item.permission, + "row_scope_type": item.row_scope_type, + "row_scope_column": item.row_scope_column, + "status": item.status, + } + + +def column_permission_payload(item: Nl2SqlRoleColumnPermission) -> dict: + return { + "id": item.id, + "role_id": item.role_id, + "table_name": item.table_name, + "column_name": item.column_name, + "access_mode": item.access_mode, + "mask_type": item.mask_type, + "status": item.status, + } + + +def sensitive_field_payload(item: Nl2SqlSensitiveField) -> dict: + return { + "id": item.id, + "table_name": item.table_name, + "column_name": item.column_name, + "mask_type": item.mask_type, + "description": item.description, + "status": item.status, + } diff --git a/service/nl2sql/permission_service.py b/service/nl2sql/permission_service.py new file mode 100644 index 0000000..0db2d98 --- /dev/null +++ b/service/nl2sql/permission_service.py @@ -0,0 +1,46 @@ +"""NL2SQL 请求级权限快照服务。""" +from __future__ import annotations + +from nl2sql.permission import build_query_permission +from model.sys_user import SysUser +from repositories.nl2sql_permission import Nl2SqlPermissionRepo +from repositories.sys_user import SysUserRepo + + +def _denied_permission(user_id: int) -> dict: + """构造默认拒绝快照,避免把不存在用户的信息暴露给调用方。""" + return build_query_permission( + SysUser(id=user_id, user_type="UNKNOWN", employee_role=None, status="异常"), + None, + [], + [], + [], + ) + + +async def load_query_permission( + db, + user_id: int, + *, + user_repo_factory=SysUserRepo, + permission_repo_factory=Nl2SqlPermissionRepo, +) -> dict: + """每次调用重新加载用户和 NL2SQL 权限,返回请求级权限快照。""" + user = await user_repo_factory(db).get(user_id) + if user is None or getattr(user, "status", None) != "正常": + return _denied_permission(user_id) + if getattr(user, "user_type", None) != "EMPLOYEE": + return _denied_permission(user_id) + + repo = permission_repo_factory(db) + role = await repo.get_role_by_employee_role(user.employee_role) + table_permissions = await repo.list_table_permissions(role.id) if role else [] + column_permissions = await repo.list_column_permissions(role.id) if role else [] + sensitive_fields = await repo.list_sensitive_fields() + return build_query_permission( + user, + role, + table_permissions, + column_permissions, + sensitive_fields, + ) diff --git a/service/nl2sql/query_service.py b/service/nl2sql/query_service.py new file mode 100644 index 0000000..73897ef --- /dev/null +++ b/service/nl2sql/query_service.py @@ -0,0 +1,170 @@ +"""供其他 Agent 复用的 NL2SQL 查询编排入口。""" +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from dataclasses import replace +from typing import Any + +from nl2sql.contracts import DataQueryRequest +from nl2sql.executor import QueryExecutionError, execute_readonly_sql +from nl2sql.row_scope import RowScopeError, apply_row_scope +from nl2sql.result import build_chart_config, summarize_result +from nl2sql.query_experience import apply_query_options, build_query_explanation +from nl2sql.supervisor import UnsupportedIntent, ensure_query_intent +from nl2sql.sql_agent import generate_sql +from nl2sql.sql_security import SqlSecurityError, validate_select_sql + + +class QueryServiceError(RuntimeError): + """查询编排失败或当前用户没有查询权限。""" + + +async def query( + request: DataQueryRequest, + *, + permission_loader: Callable[[int], Awaitable[dict[str, Any]]], + metadata_retriever: Callable[[str], Awaitable[list[dict[str, Any]]]] | None, + schema_loader: Callable[[set[str], dict[str, Any]], Awaitable[dict[str, Any]]] | None, + llm_client=None, + few_shot_retriever: Callable[[str], Awaitable[list[dict[str, Any]]]] | None = None, + conversation_context: str = "", +): + """执行权限、召回、权威 Schema、生成和安全校验,返回可执行 SQL。""" + try: + ensure_query_intent(request.question) + except UnsupportedIntent as exc: + raise QueryServiceError(str(exc)) from exc + permission = await permission_loader(request.user_id) + if not permission.get("can_query", False): + raise QueryServiceError("当前用户没有 NL2SQL 查询权限") + if metadata_retriever is None or schema_loader is None: + raise QueryServiceError("NL2SQL 查询依赖未配置") + + hits = await metadata_retriever(request.question) + candidate_tables = { + hit.get("table_name") + for hit in hits + if hit.get("table_name") in permission.get("tables", set()) + } + if not candidate_tables: + raise QueryServiceError("未找到有权限的业务表") + schema = await schema_loader(candidate_tables, permission) + if not schema.get("tables"): + raise QueryServiceError("候选表未通过权威 Schema 校验") + + few_shot = [] + if few_shot_retriever is not None: + try: + few_shot = await few_shot_retriever(request.question) + except Exception: # noqa: BLE001 Few-shot 故障不阻断主查询 + few_shot = [] + generated = await generate_sql( + request.question, + schema, + llm_client=llm_client, + few_shot=few_shot, + conversation_context=conversation_context, + ) + max_rows = request.max_rows or permission.get("max_rows") or 1000 + try: + validated = validate_select_sql( + generated.sql, + authorized_tables=permission.get("tables", set()), + authorized_columns=permission.get("columns"), + max_rows=max_rows, + ) + try: + scoped_sql = apply_row_scope( + validated.sql, + permission, + request.data_scope, + ) + except RowScopeError as exc: + raise QueryServiceError("行级权限范围无效") from exc + option_sql = apply_query_options( + scoped_sql, + page=request.page, + page_size=request.page_size, + 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"]) + return validate_select_sql( + option_sql, + authorized_tables=permission.get("tables", set()), + authorized_columns=final_columns, + max_rows=max_rows, + ) + except SqlSecurityError as exc: + raise QueryServiceError("生成的 SQL 未通过安全校验") from exc + + +async def execute_query( + request: DataQueryRequest, + *, + session, + query_id: str, + permission_loader: Callable[[int], Awaitable[dict[str, Any]]], + metadata_retriever: Callable[[str], Awaitable[list[dict[str, Any]]]] | None, + schema_loader: Callable[[set[str], dict[str, Any]], Awaitable[dict[str, Any]]] | None, + llm_client=None, + masks: dict[tuple[str, str], str] | None = None, + timeout_seconds: float | None = None, + summary_llm=None, + few_shot_retriever: Callable[[str], Awaitable[list[dict[str, Any]]]] | None = None, + conversation_context: str = "", +): + """完成 SQL 编排、只读执行和统一结果返回。""" + validated = await query( + request, + permission_loader=permission_loader, + metadata_retriever=metadata_retriever, + schema_loader=schema_loader, + llm_client=llm_client, + few_shot_retriever=few_shot_retriever, + conversation_context=conversation_context, + ) + try: + result = await execute_readonly_sql( + session, + validated, + query_id=query_id, + trace_id=request.trace_id, + user_id=request.user_id, + masks=masks, + max_rows=request.max_rows, + timeout_seconds=timeout_seconds, + ) + except QueryExecutionError as exc: + raise QueryServiceError("查询执行失败") from exc + result = replace( + result, + summary=( + await summarize_result( + request.question, + result.columns, + result.rows, + llm_client=summary_llm, + ) + if summary_llm is not None + else None + ), + chart=build_chart_config(result.columns, result.rows), + metric_definitions=build_query_explanation(request.question, validated.sql)["metrics"], + query_plan={ + **build_query_explanation(request.question, validated.sql)["plan"], + "page": request.page, + "page_size": request.page_size, + "sort_by": request.sort_by, + "sort_order": request.sort_order, + }, + ) + if not request.include_sql: + result = replace(result, sql=None) + return result diff --git a/sql/schema.sql b/sql/schema.sql index 82540d4..2b6b9c3 100644 --- a/sql/schema.sql +++ b/sql/schema.sql @@ -341,6 +341,7 @@ CREATE TABLE IF NOT EXISTS customer_relation ( signed_time DATETIME NULL COMMENT '签约时间', end_time DATETIME NULL COMMENT '关系结束时间', status VARCHAR(16) NOT NULL DEFAULT '已分配' COMMENT '已分配/已签约/已结束', + CONSTRAINT ck_customer_relation_status CHECK (status IN ('已分配', '已签约', '已结束')), reason VARCHAR(128) NULL COMMENT '结束/变更原因', KEY idx_advisor (advisor_id, status), KEY idx_customer (customer_id) diff --git a/tool/llm.py b/tool/llm.py index 9ca55fb..029bbaa 100644 --- a/tool/llm.py +++ b/tool/llm.py @@ -198,18 +198,41 @@ class LLMClient: LLM_EMBED_DIMENSIONS),MRL 模型(qwen3-embedding 等)会截断并重新归一化到目标维度。 """ backend = self.primary - if not backend.embed_model: + embed_model = ( + getattr(self.cfg, "api_embed_model", "") + if backend.name == "api" + else backend.embed_model + ) or backend.embed_model + if not embed_model: raise LLMFailError(f"backend={backend.name} 未配置 embed 模型") - url = f"{backend.base_url}/embeddings" - payload = { - "model": backend.embed_model, - "input": texts, - # "dimensions": self.cfg.embed_dimensions, - } + is_dashscope_compatible = "/compatible-mode/" in backend.base_url.lower() + if is_dashscope_compatible: + base_url = backend.base_url.rstrip("/") + marker = "/compatible-mode/v1" + if base_url.lower().endswith(marker): + base_url = base_url[: -len(marker)] + url = f"{base_url}/api/v1/services/embeddings/multimodal-embedding/multimodal-embedding" + payload = { + "model": embed_model, + "input": {"contents": [{"text": value} for value in texts]}, + } + else: + url = f"{backend.base_url}/embeddings" + payload = { + "model": embed_model, + "input": texts, + # "dimensions": self.cfg.embed_dimensions, + } async with backend.client(self.cfg.timeout) as client: r = await client.post(url, headers=backend.headers, json=payload) r.raise_for_status() data = r.json() + if is_dashscope_compatible: + embeddings = data["output"]["embeddings"] + return [ + item["embedding"] + for item in sorted(embeddings, key=lambda item: item["index"]) + ] return [item["embedding"] for item in data["data"]] async def embed_one(self, text: str) -> list[float]: diff --git a/utils/performance.py b/utils/performance.py new file mode 100644 index 0000000..93011cd --- /dev/null +++ b/utils/performance.py @@ -0,0 +1,31 @@ +"""请求耗时观测中间件。""" +from __future__ import annotations + +import logging +import time + +from starlette.middleware.base import BaseHTTPMiddleware + +logger = logging.getLogger("api.performance") + + +class PerformanceMiddleware(BaseHTTPMiddleware): + """记录接口耗时;日志只包含路径、状态码和耗时,不包含请求参数。""" + + def __init__(self, app, *, slow_ms: float = 500.0): + super().__init__(app) + self.slow_ms = slow_ms + + async def dispatch(self, request, call_next): + started = time.perf_counter() + response = await call_next(request) + elapsed_ms = (time.perf_counter() - started) * 1000 + if elapsed_ms >= self.slow_ms: + logger.warning( + "slow request path=%s method=%s status=%s elapsed_ms=%.2f", + request.url.path, + request.method, + response.status_code, + elapsed_ms, + ) + return response diff --git a/投顾工作台开发文档/common_const.md b/投顾工作台开发文档/common_const.md deleted file mode 100644 index bdbd0f3..0000000 --- a/投顾工作台开发文档/common_const.md +++ /dev/null @@ -1,282 +0,0 @@ -# common_const.md - -> 公共常量定义 用途:投顾工作台PRD(文档A)、投顾Agent需求文档(文档B)共同引用;后端、AI组件统一一份,避免枚举/字符串硬编码不一致。 版本:v1.2(架构评审决议修订) 维护方:产品 + 后端 + AI开发 注意:业务代码禁止直接写死字符串字面量,全部引用本文件定义常量。 - ------- - -## 1. customer_relation 客户‑投顾关系状态 - -表:`customer_relation` - -``` -# 客户投顾签约状态 -CUSTOMER_REL_STATUS_UNSIGNED = "unsigned" # 未签约,仅分配,未开通投顾正式服务 -CUSTOMER_REL_STATUS_SIGNED = "signed" # 已签约,可下发推荐/调仓方案给客户 -CUSTOMER_REL_STATUS_CLOSED = "closed" # 服务已终止,投顾不再提供服务 -``` - -说明: - -1. 只有 `signed` 状态,允许工作台将草稿方案下发至客户Web端; -2. `unsigned`:Agent的rebalance接口直接拦截;recommend可以生成草稿供内部预览,但工作台禁止发送; -3. `closed`:禁止调用投顾Agent为该客户生成对外方案; -4. 签约状态**唯一写入方为工作台**(投顾认领/签约/结束服务),发送前工作台必须实时查询 `status` 兜底校验; -5. 存储值统一为英文 `unsigned/signed/closed`,数据库与代码禁止使用中文「已分配/已签约/已结束」。 - ------- - -## 2. Agent草稿 draft 状态枚举 - -> 存储:Agent内部draft草稿模型,持久化MySQL;草稿不能物理删除,仅做状态流转归档 - -``` -DRAFT_STATUS_DRAFT = "draft" # 草稿初始状态,可编辑修改 -DRAFT_STATUS_DISCARDED = "discarded"# 投顾废弃该草稿,归档保留,不可编辑、不可下发 -``` - -> 说明:`sent`为工作台本地状态,Agent不维护。 状态流转: - -- draft → discarded -- discarded **不允许回退到draft**。 - ------- - -## 3. memory_unit 记忆单元类型 - -表:`memory_unit` - -``` -MEMORY_INFO_TYPE_FACT = "FACT" # 客观事实(问卷、交易行为、基础身份信息,作为业务硬规则依据) -MEMORY_INFO_TYPE_OPINION = "OPINION"# 主观观点/情绪(客户自述感受,仅做排序参考,不能绕过适当性校验) -``` - -## 4. Agent意图枚举(intent) - -> 投顾Agent内部意图路由识别结果,SSE‑meta、草稿元数据中使用 - -``` -AGENT_INTENT_RECOMMEND = "recommend" # 基金推荐意图 -AGENT_INTENT_REBALANCE = "rebalance" # 持仓诊断&调仓再平衡意图 -AGENT_INTENT_FUND_ANALYSIS = "fund_analysis"# 基金深度分析意图 -AGENT_INTENT_DIALOGUE_SCRIPT = "dialogue-script" # 生成沟通话术意图 -``` - -## 5. 沟通话术场景 scene_type(generate‑talk‑script接口入参) - -``` -TALK_SCENE_RISK_BLOCK_ORDER = "risk_block_order" # 订单被风控拦截场景 -TALK_SCENE_MARKET_FLUCTUATION = "market_fluctuation"# 市场波动安抚客户 -TALK_SCENE_PORTFOLIO_DIVERGENCE = "portfolio_divergence" # 组合大幅偏离基准 -TALK_SCENE_CUSTOMER_COMPLAINT = "customer_complaint"# 客户投诉场景 -``` - -## 6. Agent业务错误码定义 - -> HTTP返回 `code` 字段,业务层错误码;http状态码统一200,靠code区分业务异常 - -``` -ERR_CODE_OK = 0 # 成功 -ERR_CODE_FORBIDDEN_CUSTOMER = 40001 # customer_id不属于当前投顾,越权访问 -ERR_CODE_SUITABILITY_INVALID = 40020 # 适当性校验不通过:方案包含高于客户风险等级的产品 -ERR_CODE_NOT_SIGNED_REBALANCE = 40030 # 客户未签约,禁止生成rebalance调仓草稿 -ERR_CODE_DRAFT_NOT_FOUND = 40401 # draft_id不存在或已废弃 -ERR_CODE_LLM_ERROR = 50001 # Agent内部LLM调用异常 -ERR_CODE_GRAPH_ERROR = 50002 # GraphRAG查询异常,触发降级 -``` - -| code | message(默认提示文案) | -| ----- | ------------------------------------------------ | -| 0 | success | -| 40001 | 无权操作该客户数据 | -| 40020 | 方案适当性校验不通过,包含超出客户风险等级的产品 | -| 40030 | 客户尚未签约,禁止生成调仓草稿 | -| 40401 | 草稿不存在或者已废弃 | -| 50001 | AI服务调用异常,请稍后重试 | -| 50002 | 图谱查询异常,已降级返回部分结果 | - -> message 为后端默认文案;业务逻辑判断必须使用 `code`,禁止依赖 message 文案。 - -> **错误码域边界(重要)**:投顾Agent 返回的 `code`(0/40001/40020/40030/40401/50001/50002)为 Agent 自有错误码域;投顾工作台后端自有错误码域为 200/400/401/403/404/500/1001-1005(成功码为 200)。工作台调用 Agent 后**不得将 Agent 内部码原样透传给上层调用方**,须在工作台 service 层将 Agent 码映射为工作台业务分支 + 可读提示(例如 40030 → 提示"客户尚未签约,不支持生成调仓建议")。 - ------- - -## 7. 风险等级常量 - -> 风险测评问卷输出,客户风险等级C1‑C5;产品风险等级R1‑R5 - -``` -# 客户风险等级(C‑Customer) -C_RISK_C1 = "C1" -C_RISK_C2 = "C2" -C_RISK_C3 = "C3" -C_RISK_C4 = "C4" -C_RISK_C5 = "C5" - -# 基金产品风险等级(R‑Risk) -PROD_RISK_R1 = "R1" -PROD_RISK_R2 = "R2" -PROD_RISK_R3 = "R3" -PROD_RISK_R4 = "R4" -PROD_RISK_R5 = "R5" -``` - -> 适当性匹配规则:客户Cn,可购买产品R ≤ n;**硬规则在B端业务层、Agent层两处都要校验,双重防护**,且两处必须复用同一公共适当性校验函数、读取同一数据源(客户C级取自 `fin_risk_assessment.risk_level`,产品R级取自 `fin_product.risk_level`)。 - -> **共享包落地(评审决议)**:公共适当性校验函数沉淀为独立共享包 `suitability`(如 `common/suitability`),工作台与 Agent 均以内部依赖引入,**禁止各自复制实现**。函数契约:`check_suitability(customer_risk: str, product_risk: str) -> {ok: bool, reason: str}`;判定规则 `R_n ≤ C_n`。 - -> **存储编码统一**:客户风险等级统一存 `C1-C5`,产品风险等级统一存 `R1-R5`。历史中文等级「保守/稳健/平衡/进取/激进」仅允许存在于展示层,与 C 级一一对应(保守=C1、稳健=C2、平衡=C3、进取=C4、激进=C5)。数据库字段 `fin_customer_profile.risk_level`、`fin_risk_assessment.risk_level`、`portfolio_benchmark.risk_level` 必须存 `C1-C5`,`fin_product.risk_level` 存 `R1-R5`。 - ------- - -## 8. 系统配置 sys_config key常量 - -> sys_config表 KV配置,调仓、Agent相关参数key,代码不要写死字符串key - -``` -SYS_KEY_REBALANCE_DEVIATION_THRESHOLD = "rebalance.deviation.threshold" # 组合再平衡偏离阈值(全局兜底),浮点数,例0.1代表±10% -SYS_KEY_HIGH_NET_ASSET_THRESHOLD = "customer.high_net.asset.threshold" # 高净值客户资产门槛 -SYS_KEY_RISK_QUESTIONNAIRE_EXPIRE_DAY = "risk.questionnaire.expire.day" # 风险测评过期天数 -SYS_KEY_LARGE_FLOW_THRESHOLD = "customer.large_flow.threshold" # 大额申赎阈值(工作台本地定时任务扫描 fin_transaction) -``` - ------- - -## 9. 事件总线 Pub/Sub 事件名称 - -> Redis Pub/Sub事件频道名,业务同时双写落库,保证消息不丢失 - -``` -EVENT_ADVISOR_REBALANCE_DRAFT_CREATED = "event:rebalance_draft_created" # rebalance草稿生成(工作台消费→待办) -EVENT_PROFILE_UPDATE = "event:profile_update" # 客户画像/记忆更新(工作台只读刷新) -EVENT_WORK_ORDER_CHANGE = "event:work_order_change" # 工单状态变更(工单域,非投顾工作台事件契约) -EVENT_RISK_ALERT = "event:risk_alert" # 风控预警事件(风控域,非投顾工作台事件契约) - -> 投顾工作台与投顾Agent 的事件契约仅包含前两者:`event:rebalance_draft_created`、`event:profile_update`;`work_order_change`、`risk_alert` 属其他域事件。 -``` - -### 事件消息公共字段约定(JSON payload) - -> 事件结构 = 外层公共字段 + 内层 payload。以下为外层公共字段: - -``` -{ - "event_name": "", - "trace_id": "", - "trigger_user_id": "", - "customer_id": "", - "payload": {} -} -``` - -> 内层 payload 按事件名定义。`event:rebalance_draft_created` 的内层 payload 结构如下(`advisor_id` 放内层): - -``` -{ - "draft_id": "", - "customer_id": "", - "advisor_id": "", - "deviation": 0.0, - "created_at": "" -} -``` - -> 事件双写 `event_log` 表(event_name + event_id 幂等 + payload + trace_id + 消费状态),Pub/Sub 仅作实时通知;工作台消费以 event_id 幂等去重。 - ------- - -## 10. 固定文本模板 - -### 10.1 Agent输出报告强制免责声明 - -> ⚠️ Agent生成markdown报告必须拼接此文本;草稿保存接口校验,如果缺失该声明返回警告;**真正拦截发送发生在工作台发送接口**。 - -> **完整性判定规则(决议)**:报告正文必须**逐字包含**上述免责声明完整原文(精确匹配);编辑器中免责声明为**只读区**,投顾改动导致不匹配时,工作台发送环节直接拦截。 - -``` -【免责声明】本报告由AI辅助生成,仅供持牌投顾内部参考,不构成任何投资建议。基金有风险,投资需谨慎。所有投资决策请结合自身风险承受能力审慎判断。 -``` - -### 10.2 SSE事件类型常量 - -> Agent流式SSE推送事件type - -``` -SSE_EVENT_TYPE_TEXT = "text" # 增量文本片段 -SSE_EVENT_TYPE_META = "meta" # 结构化元数据 -SSE_EVENT_TYPE_DONE = "done" # 会话正常结束 -SSE_EVENT_TYPE_ERROR = "error" # 会话异常中断 -``` - ------- - -## 11. 定时任务相关常量 - -``` -# 每周记忆维护任务:置信度重算、矛盾检测、记忆遗忘归档 -CRON_MEMORY_MAINTENANCE = "0 2 * * 1" -# 每日客户组合再平衡任务【由工作台后端调度,Agent不调度】 -CRON_PORTFOLIO_REBALANCE = "0 1 * * *" -# 基金净值更新定时任务 -CRON_FUND_NAV_UPDATE = "30 1 * * *" -``` - ------- - -## 12. 数据库审计日志 action 动作常量 audit_log.action - -``` -AUDIT_AGENT_CHAT_CALL = "agent_chat_call" # 调用投顾Agent对话 -AUDIT_DRAFT_SAVE = "draft_save" # 草稿保存 -AUDIT_DRAFT_DISCARD = "draft_discard" # 草稿废弃 -AUDIT_CUSTOMER_SIGN = "customer_sign" # 客户签约 -AUDIT_VIEW_SENSITIVE = "view_sensitive" # 投顾查看脱敏字段完整值(留痕) -``` - ------- - -## 13. 投顾角色常量(sys_user.employee_role) - -> 工作台 RBAC 与投顾Agent 鉴权统一使用本值,「理财顾问」为同岗位历史表述,不再参与鉴权。 - -``` -EMPLOYEE_ROLE_ADVISOR = "投顾" # 投顾(工作台唯一业务角色) -``` - -## 14. 站内信消息类型常量(sys_message.msg_type) - -``` -MSG_TYPE_RECOMMEND = "推荐" # 基金推荐报告 -MSG_TYPE_REBALANCE = "调仓" # 调仓建议报告 -MSG_TYPE_RISK = "风控" # 风控预警通知 -MSG_TYPE_SIGN = "签约" # 签约相关通知 -MSG_TYPE_SYSTEM = "系统" # 系统通知 -``` - -> 草稿 intent → msg_type 映射:`recommend`→`推荐`,`rebalance`→`调仓`。 - -## 15. 新增数据表(本次评审决议) - -| 表名 | 归属 | 用途 | -| ---- | ---- | ---- | -| `advisor_draft` | 投顾Agent | draft 草稿持久化(draft/discarded,不存 sent) | -| `advisor_report` | 投顾工作台 | 建议报告本地镜像:sent 状态 + 编辑/发送留痕 | -| `advisor_todo` | 投顾工作台 | 待办任务(事件/定时/计算三类来源统一承载) | -| `advisor_visit_record` | 投顾工作台 | 人工回访纪要留痕归档 | -| `event_log` | 共享 | 事件双写落库(event_id 幂等,保证不丢) | -| `sensitive_word` | 共享 | 敏感词库(工作台与 Agent 同源,合规拦截唯一来源) | - -> 字段结构以 `sql/schema.sql` 为准;`advisor_draft.status` 仅 `draft/discarded`,`advisor_report.send_status` 为 `draft/sent/discarded`。 - ------- - -## 16. 敏感词库数据源约定 - -> 敏感词唯一数据源为 `sensitive_word` 表;工作台发送终审与 Agent 生成校验均读取该表,禁止各自维护词库。V1.0 由后端灌种子数据 + 管理脚本维护,运营端维护界面随 Phase 5 补。 - ------- - -# 使用规范(重要) - -1. 所有 Python/后端代码,import 本常量文件,禁止硬编码字符串; -2. 所有状态、错误码、事件名称修改,必须同步更新此文档,通知 AI、B 端后端; -3. 业务规则(如适当性匹配逻辑)写在业务代码中,不在此文档写业务逻辑,本文件只存放**常量字符串、枚举、模板文本**。 \ No newline at end of file