feat:客服agent接入nl2sql
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user