feat:客服agent接入nl2sql

This commit is contained in:
2026-09-13 20:48:42 +08:00
parent b133dd58c4
commit dd281b3361
11 changed files with 778 additions and 18 deletions
+3
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
from config.database.milvus import client as milvus_client
from config.database.mysql import get_session_factory
from config.database.redis import client as redis_client
from config.settings import settings as app_settings
from service.client_agent.runtime import build_client_runtime
from service.customer_agent.config import DatabaseConfigProvider
from tool.llm import llm as llm_client
@@ -19,6 +20,8 @@ def build_default_runtime():
llm_client=llm_client,
config_getter=provider.get,
audit_writer=provider.write_audit,
db_session_factory=get_session_factory(),
schema_database=app_settings.mysql.database,
)
+147 -2
View File
@@ -7,15 +7,28 @@ import logging
import uuid
from types import SimpleNamespace
from nl2sql.contracts import DataQueryRequest
from nl2sql.history import archive_query_safely
from nl2sql.limits import QueryLimiter
from nl2sql.retrieval import retrieve_metadata
from nl2sql.runtime_config import runtime_config
from nl2sql.schema import load_authoritative_schema
from agent.client_agent.session import ClientSessionService
from rag.embedding import embed_texts
from rag.generation import generate_answer
from rag.intent import intent_recognize
from rag.retrieve import rag_retrieve
from service.customer_agent.chat import AnonymousCustomerAgent
from service.customer_agent.chat import (
AnonymousCustomerAgent,
DataQueryRejected,
)
from service.client_agent.memory_extractor import DialogueMemoryExtractor
from service.memory.facade import MemoryService
from service.memory.schemas import CustomerMemoryContext, MemoryUnitDTO, ShortTermMessage
from service.nl2sql.answer_render import render_query_answer
from service.nl2sql.customer_permission import load_customer_query_permission
from service.nl2sql.query_service import QueryServiceError, execute_query
_active_customer = contextvars.ContextVar("client_agent_customer", default=None)
@@ -91,7 +104,9 @@ class MemoryAwareClientAgent:
warnings_token = _active_warnings.set(warnings)
context_token = _active_memory_context.set(memory_context)
try:
result = await self.agent.handle(session_id, query, trace_id=trace_id)
result = await self.agent.handle(
session_id, query, trace_id=trace_id, customer_id=customer_id
)
if self.extractor is not None:
await self._save_candidates(
customer_id,
@@ -180,6 +195,122 @@ class MemoryAwareClientAgent:
raise TypeError(f"unsupported memory context item: {type(item).__name__}")
def _build_data_query(*, db_session_factory, milvus_client, llm_client, config_getter, redis, schema_database):
"""构造登录客户的 NL2SQL 数据查询依赖。
身份与行级范围由服务端强制注入:data_scope 只带当前登录客户自己的
customer_id,配合权限快照的 row_scopes 在 SQL 层兜底,杜绝水平越权。
"""
async def data_query(*, question, customer_id, session_id, trace_id):
async with db_session_factory() as db:
permission = await load_customer_query_permission(
db, customer_id, config_getter=config_getter
)
if not permission.get("can_query"):
raise DataQueryRejected("数据查询功能暂未开放,您可以先咨询基金知识或开户流程~")
limiter = QueryLimiter(redis)
if not await limiter.acquire(
customer_id,
daily_quota=permission.get("daily_quota", 0) or 20,
max_concurrent=1,
rate_limit=10,
):
raise DataQueryRejected("您今天的数据查询次数已达上限,请明天再来吧~")
try:
return await _execute_customer_query(
db=db,
permission=permission,
question=question,
customer_id=customer_id,
session_id=session_id,
trace_id=trace_id,
)
finally:
await limiter.release(customer_id)
async def _execute_customer_query(*, db, permission, question, customer_id, session_id, trace_id):
query_id = uuid.uuid4().hex
async def permission_loader(_user_id: int):
return permission
async def metadata_retriever(retrieval_question: str):
return await retrieve_metadata(
retrieval_question, milvus_client, top_k=runtime_config.retrieval_top_k
)
async def schema_loader(table_names: set[str], _permission: dict):
return await load_authoritative_schema(
db, database=schema_database, candidate_tables=table_names
)
request = DataQueryRequest(
question=question,
user_id=customer_id,
trace_id=trace_id,
session_id=session_id,
caller_agent="client_agent",
# 行级范围只允许是登录客户本人,不接受任何外部输入
data_scope={"customer_ids": [customer_id]},
include_sql=False,
max_rows=min(permission.get("max_rows") or 200, runtime_config.max_rows),
)
try:
result = await execute_query(
request,
session=db,
query_id=query_id,
permission_loader=permission_loader,
metadata_retriever=metadata_retriever,
schema_loader=schema_loader,
llm_client=llm_client,
masks=permission.get("masks") or {},
summary_llm=llm_client,
)
except QueryServiceError as exc:
await archive_query_safely(
db,
query_id=query_id,
user_id=customer_id,
question=question,
status="blocked",
error_message=str(exc),
trace_id=trace_id,
session_id=session_id,
caller_agent="client_agent",
)
# 不向用户暴露 SQL 和内部异常细节
raise DataQueryRejected(
"这个问题我暂时查不了,您可以换个问法,或者联系人工客服帮您处理~"
) from exc
await archive_query_safely(
db,
query_id=query_id,
user_id=customer_id,
question=question,
status="success",
row_count=result.row_count,
truncated=result.truncated,
elapsed_ms=result.elapsed_ms,
trace_id=trace_id,
session_id=session_id,
caller_agent="client_agent",
)
return {
"answer": render_query_answer(result),
"sources": [],
"query_id": result.query_id,
"row_count": result.row_count,
"truncated": result.truncated,
"chart": result.chart,
}
return data_query
def build_client_runtime(
*,
redis,
@@ -187,6 +318,8 @@ def build_client_runtime(
llm_client,
config_getter,
audit_writer,
db_session_factory=None,
schema_database=None,
memory_service=None,
memory_extractor=None,
):
@@ -233,6 +366,17 @@ def build_client_runtime(
config_getter=config_getter,
)
data_query = None
if db_session_factory is not None and schema_database:
data_query = _build_data_query(
db_session_factory=db_session_factory,
milvus_client=milvus_client,
llm_client=llm_client,
config_getter=config_getter,
redis=redis,
schema_database=schema_database,
)
agent = AnonymousCustomerAgent(
context=context,
rag_retrieve=retrieve,
@@ -240,6 +384,7 @@ def build_client_runtime(
generate_answer=generate,
audit_writer=audit_writer,
config_getter=config_getter,
data_query=data_query,
)
extractor = memory_extractor or DialogueMemoryExtractor(llm_client)
wrapped_agent = MemoryAwareClientAgent(