feat:修复投顾agent功能

This commit is contained in:
2026-09-14 21:41:40 +08:00
parent fc1d74570f
commit 058f45115d
14 changed files with 142 additions and 33 deletions
+18
View File
@@ -13,6 +13,7 @@ from config.settings import settings
from nl2sql.contracts import DataQueryRequest, DataQueryResult from nl2sql.contracts import DataQueryRequest, DataQueryResult
from nl2sql.embedding import EmbeddingError from nl2sql.embedding import EmbeddingError
from nl2sql.retrieval import retrieve_metadata from nl2sql.retrieval import retrieve_metadata
from nl2sql.cache import cache_get, cache_set, build_question_cache_key
from nl2sql.runtime_config import runtime_config from nl2sql.runtime_config import runtime_config
from nl2sql.schema import load_authoritative_schema from nl2sql.schema import load_authoritative_schema
from repositories.customer_relation import CustomerRelationRepo from repositories.customer_relation import CustomerRelationRepo
@@ -154,6 +155,7 @@ async def execute_advisor_data_query(
pass pass
if scope == "customer" and _is_current_holdings_query(question): if scope == "customer" and _is_current_holdings_query(question):
redis = redis or database.redis.client()
return await _query_current_holdings( return await _query_current_holdings(
db, db,
customer_id=customer_id, customer_id=customer_id,
@@ -166,6 +168,21 @@ async def execute_advisor_data_query(
milvus = milvus or database.milvus.client() milvus = milvus or database.milvus.client()
redis = redis or database.redis.client() redis = redis or database.redis.client()
cache_key = build_question_cache_key(
question,
permission=permission,
data_scope={"customer_ids": customer_ids},
page=page,
page_size=page_size,
sort_by=sort_by,
sort_order=sort_order,
)
cached = await cache_get(redis, cache_key)
if cached is not None:
cached["trace_id"] = trace_id
cached["customer_id"] = customer_id
cached["warnings"] = [*cached.get("warnings", []), "cache_hit"]
return cached
llm_client = llm_client or default_llm llm_client = llm_client or default_llm
request = DataQueryRequest( request = DataQueryRequest(
question=question, question=question,
@@ -222,4 +239,5 @@ async def execute_advisor_data_query(
payload["answer"] = f"{existing_answer.rstrip('。')}。客户姓名:{enriched['name_summary']}。" payload["answer"] = f"{existing_answer.rstrip('。')}。客户姓名:{enriched['name_summary']}。"
payload["sql"] = None payload["sql"] = None
payload["customer_id"] = customer_id payload["customer_id"] = customer_id
await cache_set(redis, cache_key, payload, ttl=runtime_config.cache_ttl)
return payload return payload
+11
View File
@@ -108,6 +108,17 @@ async def classify_advisor_intent(
"rule", "rule",
"识别为系统提示文本而非业务查询", "识别为系统提示文本而非业务查询",
) )
fast_intent = recognize_advisor_intent(query)
if (
fast_intent == AGENT_INTENT_DATA_QUERY
and not any(term in query for term in _RECOMMENDATION_TERMS)
):
return IntentClassification(
fast_intent,
0.95,
"rule_fast",
"明显查询类问题",
)
if llm_client is not None: if llm_client is not None:
try: try:
raw = await asyncio.wait_for( raw = await asyncio.wait_for(
+2
View File
@@ -331,6 +331,7 @@ async def chat_stream(
session_id=chat_request.session_id, session_id=chat_request.session_id,
conversation_context=conversation_context, conversation_context=conversation_context,
llm_client=getattr(_advisor_runtime(request), "llm_client", None), llm_client=getattr(_advisor_runtime(request), "llm_client", None),
redis=redis_db.client(),
) )
except QueryServiceError as exc: except QueryServiceError as exc:
payload = agent_failure(ERR_CODE_LLM_ERROR, _data_query_error_message(exc), trace_id=trace_id) payload = agent_failure(ERR_CODE_LLM_ERROR, _data_query_error_message(exc), trace_id=trace_id)
@@ -565,6 +566,7 @@ async def advisor_data_query(
page_size=body.page_size, page_size=body.page_size,
sort_by=body.sort_by, sort_by=body.sort_by,
sort_order=body.sort_order, sort_order=body.sort_order,
redis=redis_db.client(),
) )
except QueryServiceError as exc: except QueryServiceError as exc:
return agent_failure( return agent_failure(
+1
View File
@@ -248,6 +248,7 @@ async def query_data(
db, db,
database=settings.mysql.database, database=settings.mysql.database,
candidate_tables=table_names, candidate_tables=table_names,
redis=redis,
) )
request_contract = DataQueryRequest( request_contract = DataQueryRequest(
+1 -4
View File
@@ -72,10 +72,7 @@ export const assistantApi = {
sessionHistory: (sessionId: string) => sessionHistory: (sessionId: string) =>
apiFetch<AgentMessage[]>(`/advisor-agent/session/${encodeURIComponent(sessionId)}/history`), apiFetch<AgentMessage[]>(`/advisor-agent/session/${encodeURIComponent(sessionId)}/history`),
/** /** POST /chat/stream —— 读取 SSE,并支持调用方实时消费事件。 */
* POST /chat/stream —— 响应是 text/event-stream,但后端整包输出,
* 因此读完整条响应后一次性返回聚合结果(非流式渲染)。
*/
chatOnce( chatOnce(
body: { body: {
query: string; query: string;
+1
View File
@@ -90,6 +90,7 @@ async def lifespan(app: FastAPI):
await app.state.advisor_event_consumer.stop() await app.state.advisor_event_consumer.stop()
if app.state.advisor_scheduler is not None: if app.state.advisor_scheduler is not None:
app.state.advisor_scheduler.shutdown() app.state.advisor_scheduler.shutdown()
await llm_client.aclose()
await database.dispose() await database.dispose()
+26
View File
@@ -49,6 +49,32 @@ def build_cache_key(
return f"nl2sql:{digest}" return f"nl2sql:{digest}"
def build_question_cache_key(
question: str,
*,
permission: dict[str, Any],
data_scope: dict[str, Any] | None = None,
page: int = 1,
page_size: int = 100,
sort_by: str | None = None,
sort_order: str = "asc",
) -> str:
"""生成自然语言查询结果缓存键,覆盖问题、权限和数据范围。"""
payload = {
"question": " ".join((question or "").split()),
"permission": _permission_payload(permission),
"data_scope": data_scope or {},
"page": page,
"page_size": page_size,
"sort_by": sort_by or "",
"sort_order": sort_order,
}
digest = hashlib.sha256(
json.dumps(payload, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode()
).hexdigest()
return f"nl2sql:question:{digest}"
async def cache_get(redis, key: str) -> dict[str, Any] | None: async def cache_get(redis, key: str) -> dict[str, Any] | None:
"""读取缓存,Redis 异常或内容损坏时返回空结果。""" """读取缓存,Redis 异常或内容损坏时返回空结果。"""
try: try:
+6 -1
View File
@@ -16,6 +16,11 @@ _SYSTEM_PROMPT = """你是基金平台 NL2SQL 的问题改写器。
不能扩大权限范围,不能臆造查询结果;上下文不足时保留原问题。 不能扩大权限范围,不能臆造查询结果;上下文不足时保留原问题。
""" """
_REFERENCE_TERMS = (
"他", "她", "它", "他们", "她们", "它们", "这个客户", "该客户", "那个客户",
"这个基金", "该基金", "那只基金", "这只基金", "这个结果", "上一轮", "刚才", "前面",
)
def _rewrite_explicit_customer_identity(question: str) -> str | None: def _rewrite_explicit_customer_identity(question: str) -> str | None:
normalized = re.sub(r"\s+", "", question) normalized = re.sub(r"\s+", "", question)
@@ -68,7 +73,7 @@ async def rewrite_query(
contextual_identity = _rewrite_customer_name_follow_up(original, context) contextual_identity = _rewrite_customer_name_follow_up(original, context)
if contextual_identity: if contextual_identity:
return contextual_identity return contextual_identity
if not context or llm_client is None: if not context or llm_client is None or not any(term in original for term in _REFERENCE_TERMS):
return original return original
prompt = ( prompt = (
+1
View File
@@ -67,6 +67,7 @@ async def render_answer(
{"role": "user", "content": prompt}, {"role": "user", "content": prompt},
], ],
temperature=0, temperature=0,
max_tokens=256,
) )
except Exception: # noqa: BLE001 摘要失败回退固定文本 except Exception: # noqa: BLE001 摘要失败回退固定文本
return fallback return fallback
+10 -3
View File
@@ -1,6 +1,7 @@
"""NL2SQL 元数据召回与授权过滤。""" """NL2SQL 元数据召回与授权过滤。"""
from __future__ import annotations from __future__ import annotations
import asyncio
from collections.abc import Awaitable, Callable, Iterable from collections.abc import Awaitable, Callable, Iterable
from typing import Any from typing import Any
@@ -70,15 +71,21 @@ async def retrieve_metadata(
"is_valid", "is_valid",
"is_deprecated", "is_deprecated",
] ]
results: list[dict[str, Any]] = [] async def search_chunk(chunk_type: str):
for chunk_type in ("table_meta", "field_meta"): return await milvus_client.search(
hits = await milvus_client.search(
collection_name=NL2SQL_COLLECTION, collection_name=NL2SQL_COLLECTION,
data=[vector], data=[vector],
limit=top_k, limit=top_k,
filter=f'chunk_type == "{chunk_type}" and is_valid == true', filter=f'chunk_type == "{chunk_type}" and is_valid == true',
output_fields=output_fields, output_fields=output_fields,
) )
table_hits, field_hits = await asyncio.gather(
search_chunk("table_meta"),
search_chunk("field_meta"),
)
results: list[dict[str, Any]] = []
for hits in (table_hits, field_hits):
for hit in _flatten_hits(hits): for hit in _flatten_hits(hits):
entity = hit.get("entity") or hit entity = hit.get("entity") or hit
if not entity.get("is_valid", True) or entity.get("is_deprecated", False): if not entity.get("is_valid", True) or entity.get("is_deprecated", False):
+31 -2
View File
@@ -1,12 +1,23 @@
"""从 MySQL information_schema 加载并校验 NL2SQL 权威 Schema。""" """从 MySQL information_schema 加载并校验 NL2SQL 权威 Schema。"""
from __future__ import annotations from __future__ import annotations
import hashlib
import json
import logging
from typing import Any from typing import Any
from sqlalchemy import bindparam, text from sqlalchemy import bindparam, text
from nl2sql.metadata import normalize_column_row, normalize_table_row from nl2sql.metadata import normalize_column_row, normalize_table_row
logger = logging.getLogger("nl2sql.schema")
SCHEMA_CACHE_TTL = 60
def _schema_cache_key(database: str, candidates: list[str]) -> str:
payload = json.dumps([database, candidates], ensure_ascii=False, separators=(",", ":"))
return "nl2sql:schema:" + hashlib.sha256(payload.encode()).hexdigest()
TABLES_SQL = text( TABLES_SQL = text(
""" """
@@ -34,13 +45,25 @@ async def load_authoritative_schema(
*, *,
database: str, database: str,
candidate_tables: set[str] | list[str], candidate_tables: set[str] | list[str],
redis=None,
cache_ttl: int = SCHEMA_CACHE_TTL,
) -> dict[str, list[dict[str, Any]]]: ) -> dict[str, list[dict[str, Any]]]:
"""从权威元数据源加载候选表,并剔除不存在的表和孤立字段。""" """从权威元数据源加载候选表,并剔除不存在的表和孤立字段。"""
candidates = {str(name).strip() for name in candidate_tables if str(name).strip()} candidates = {str(name).strip() for name in candidate_tables if str(name).strip()}
if not candidates: if not candidates:
return {"tables": [], "columns": []} return {"tables": [], "columns": []}
params = {"database": database, "candidate_tables": sorted(candidates)} candidate_list = sorted(candidates)
cache_key = _schema_cache_key(database, candidate_list)
if redis is not None:
try:
cached = await redis.get(cache_key)
if cached:
return json.loads(cached)
except Exception: # noqa: BLE001 缓存异常回退权威查询
logger.warning("Schema 缓存读取失败", exc_info=True)
params = {"database": database, "candidate_tables": candidate_list}
table_result = await session.execute(TABLES_SQL, params) table_result = await session.execute(TABLES_SQL, params)
column_result = await session.execute(COLUMNS_SQL, params) column_result = await session.execute(COLUMNS_SQL, params)
@@ -56,4 +79,10 @@ async def load_authoritative_schema(
normalized = normalize_column_row(dict(row)) normalized = normalize_column_row(dict(row))
if normalized["table_name"] in valid_table_names and normalized["field_name"]: if normalized["table_name"] in valid_table_names and normalized["field_name"]:
columns.append(normalized) columns.append(normalized)
return {"tables": tables, "columns": columns} schema = {"tables": tables, "columns": columns}
if redis is not None:
try:
await redis.set(cache_key, json.dumps(schema, ensure_ascii=False), ex=cache_ttl)
except Exception: # noqa: BLE001 缓存异常不阻断查询
logger.warning("Schema 缓存写入失败", exc_info=True)
return schema
+1 -1
View File
@@ -74,7 +74,7 @@ async def generate_sql(
}, },
] ]
try: try:
output = await llm_client.chat(messages, temperature=0) output = await llm_client.chat(messages, temperature=0, max_tokens=256)
except Exception as exc: # noqa: BLE001 统一收敛模型异常 except Exception as exc: # noqa: BLE001 统一收敛模型异常
raise SqlGenerationError("SQL 模型调用失败") from exc raise SqlGenerationError("SQL 模型调用失败") from exc
sql = _clean_model_output(output) sql = _clean_model_output(output)
+1 -1
View File
@@ -304,7 +304,7 @@ def _build_data_query(*, db_session_factory, milvus_client, llm_client, config_g
async def schema_loader(table_names: set[str], _permission: dict): async def schema_loader(table_names: set[str], _permission: dict):
return await load_authoritative_schema( return await load_authoritative_schema(
db, database=schema_database, candidate_tables=table_names db, database=schema_database, candidate_tables=table_names, redis=redis
) )
request = DataQueryRequest( request = DataQueryRequest(
+32 -21
View File
@@ -103,6 +103,21 @@ class LLMClient:
def __init__(self, cfg: LLMCfg | None = None): def __init__(self, cfg: LLMCfg | None = None):
self.cfg = cfg or settings.llm self.cfg = cfg or settings.llm
self.backends = resolve_backends(self.cfg) self.backends = resolve_backends(self.cfg)
self._clients: dict[str, httpx.AsyncClient] = {}
def _client_for(self, backend: Backend) -> httpx.AsyncClient:
clients = getattr(self, "_clients", None)
if clients is None:
clients = {}
self._clients = clients
if backend.name not in clients:
clients[backend.name] = backend.client(self.cfg.timeout)
return clients[backend.name]
async def aclose(self) -> None:
for client in getattr(self, "_clients", {}).values():
await client.aclose()
getattr(self, "_clients", {}).clear()
# ---- 首选后端(健康检查/日志/embed 使用) ----------------------------- # ---- 首选后端(健康检查/日志/embed 使用) -----------------------------
@property @property
@@ -186,9 +201,9 @@ class LLMClient:
"stream": False, "stream": False,
} }
try: try:
async with backend.client(self.cfg.timeout) as client: client = self._client_for(backend)
r = await client.post(url, headers=backend.headers, json=payload) r = await client.post(url, headers=backend.headers, json=payload)
r.raise_for_status() r.raise_for_status()
choice = r.json()["choices"][0] choice = r.json()["choices"][0]
message = choice["message"] message = choice["message"]
content = message.get("content") content = message.get("content")
@@ -250,20 +265,16 @@ class LLMClient:
response = None response = None
max_retries = max(1, int(getattr(self.cfg, "max_retries", 1))) max_retries = max(1, int(getattr(self.cfg, "max_retries", 1)))
retry_backoff_sec = float(getattr(self.cfg, "retry_backoff_sec", 0)) retry_backoff_sec = float(getattr(self.cfg, "retry_backoff_sec", 0))
async with backend.client(self.cfg.timeout) as client: client = self._client_for(backend)
for attempt in range(max_retries): for attempt in range(max_retries):
try: try:
response = await client.post( response = await client.post(url, headers=backend.headers, json=payload)
url, response.raise_for_status()
headers=backend.headers, break
json=payload, except Exception:
) if attempt == max_retries - 1:
response.raise_for_status() raise
break await asyncio.sleep(retry_backoff_sec * (2**attempt))
except Exception:
if attempt == max_retries - 1:
raise
await asyncio.sleep(retry_backoff_sec * (2**attempt))
if response is None: # pragma: no cover - defensive guard if response is None: # pragma: no cover - defensive guard
raise LLMFailError("Embedding 请求未返回响应") raise LLMFailError("Embedding 请求未返回响应")
data = response.json() data = response.json()
@@ -286,10 +297,10 @@ class LLMClient:
url = backend.base_url.removesuffix("/v1") + "/api/tags" url = backend.base_url.removesuffix("/v1") + "/api/tags"
else: else:
url = f"{backend.base_url}/models" url = f"{backend.base_url}/models"
async with backend.client(min(self.cfg.timeout, 10)) as client: client = self._client_for(backend)
r = await client.get(url, headers=backend.headers) r = await client.get(url, headers=backend.headers)
if r.status_code >= 500: if r.status_code >= 500:
r.raise_for_status() r.raise_for_status()
# 全局单例:Agent 统一引入 # 全局单例:Agent 统一引入