feat:客服agent接入nl2sql
This commit is contained in:
@@ -16,6 +16,10 @@ class QueryTooLongError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
class DataQueryRejected(ValueError):
|
||||
"""NL2SQL 数据查询被拒绝(未开放、无权限或配额不足),message 可直接回复用户。"""
|
||||
|
||||
|
||||
async def _config(config_getter, key: str, default: str):
|
||||
value = config_getter(key, default)
|
||||
if isawaitable(value):
|
||||
@@ -37,6 +41,7 @@ class AnonymousCustomerAgent:
|
||||
generate_answer,
|
||||
audit_writer,
|
||||
config_getter,
|
||||
data_query=None,
|
||||
):
|
||||
self.context = context
|
||||
self.rag_retrieve = rag_retrieve
|
||||
@@ -44,8 +49,18 @@ class AnonymousCustomerAgent:
|
||||
self.generate_answer = generate_answer
|
||||
self.audit_writer = audit_writer
|
||||
self.config_getter = config_getter
|
||||
# 可选 NL2SQL 数据查询依赖:签名 data_query(*, question, customer_id,
|
||||
# session_id, trace_id) -> dict;匿名 runtime 不装配(None),行为不变。
|
||||
self.data_query = data_query
|
||||
|
||||
async def handle(self, session_id: str, query: str, *, trace_id: str) -> dict:
|
||||
async def handle(
|
||||
self,
|
||||
session_id: str,
|
||||
query: str,
|
||||
*,
|
||||
trace_id: str,
|
||||
customer_id: int | None = None,
|
||||
) -> dict:
|
||||
if len(query) > 2000:
|
||||
raise QueryTooLongError("query长度不能超过2000字符")
|
||||
# 先取历史再写入当前问题,保证意图识别拿到的历史不含本轮输入;取不到历史不阻断请求
|
||||
@@ -68,6 +83,7 @@ class AnonymousCustomerAgent:
|
||||
else:
|
||||
intent, search_query = recognized, query
|
||||
sources = []
|
||||
data_query_meta = None
|
||||
if intent == Intent.GUIDE_PURCHASE:
|
||||
answer = await _config(
|
||||
self.config_getter,
|
||||
@@ -89,6 +105,21 @@ class AnonymousCustomerAgent:
|
||||
)
|
||||
elif intent == Intent.CHITCHAT:
|
||||
answer = await self._chitchat(session_id)
|
||||
elif intent == Intent.NL2SQL_REQUEST:
|
||||
if self.data_query is None or customer_id is None:
|
||||
# 匿名会话或未装配数据查询能力:引导登录,不触发任何数据库查询
|
||||
answer = await _config(
|
||||
self.config_getter,
|
||||
"agent.customer.template.nl2sql_unavailable",
|
||||
"数据查询功能需要登录后使用,请先登录再来问我您的持仓和交易信息~",
|
||||
)
|
||||
else:
|
||||
answer, sources, data_query_meta = await self._run_data_query(
|
||||
question=search_query,
|
||||
customer_id=customer_id,
|
||||
session_id=session_id,
|
||||
trace_id=trace_id,
|
||||
)
|
||||
elif intent in (Intent.KNOWLEDGE_QA, Intent.COMPANY_INFO):
|
||||
try:
|
||||
# 用补全指代后的问题检索,省略主语的追问才能命中
|
||||
@@ -129,13 +160,63 @@ class AnonymousCustomerAgent:
|
||||
)
|
||||
|
||||
await self.context.append(session_id, "assistant", answer)
|
||||
return {
|
||||
result = {
|
||||
"answer": answer,
|
||||
"sources": sources,
|
||||
"intent": intent.value,
|
||||
"rewritten_query": search_query,
|
||||
"trace_id": trace_id,
|
||||
}
|
||||
if data_query_meta is not None:
|
||||
result["data_query"] = data_query_meta
|
||||
return result
|
||||
|
||||
async def _run_data_query(
|
||||
self,
|
||||
*,
|
||||
question: str,
|
||||
customer_id: int,
|
||||
session_id: str,
|
||||
trace_id: str,
|
||||
) -> tuple[str, list, dict]:
|
||||
"""调用注入的 NL2SQL 数据查询能力,失败时统一降级为客服话术。"""
|
||||
try:
|
||||
payload = await _maybe_await(
|
||||
self.data_query(
|
||||
question=question,
|
||||
customer_id=customer_id,
|
||||
session_id=session_id,
|
||||
trace_id=trace_id,
|
||||
)
|
||||
)
|
||||
except DataQueryRejected as exc:
|
||||
return str(exc), [], None
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"client data query failed: trace_id=%s session_id=%s customer_id=%s",
|
||||
trace_id,
|
||||
session_id,
|
||||
customer_id,
|
||||
)
|
||||
answer = await _config(
|
||||
self.config_getter,
|
||||
"agent.customer.template.nl2sql_fallback",
|
||||
"暂时无法完成数据查询,请稍后再试或联系人工客服。",
|
||||
)
|
||||
return answer, [], None
|
||||
|
||||
if not isinstance(payload, dict) or not str(payload.get("answer") or "").strip():
|
||||
return await _config(
|
||||
self.config_getter,
|
||||
"agent.customer.template.nl2sql_fallback",
|
||||
"暂时无法完成数据查询,请稍后再试或联系人工客服。",
|
||||
), [], None
|
||||
meta = {
|
||||
key: payload[key]
|
||||
for key in ("query_id", "row_count", "truncated", "chart")
|
||||
if payload.get(key) is not None
|
||||
}
|
||||
return str(payload["answer"]), list(payload.get("sources") or []), meta or None
|
||||
|
||||
async def _chitchat(self, session_id: str) -> str:
|
||||
"""带对话历史调用 LLM 做受限闲聊,失败时退回固定话术。"""
|
||||
|
||||
Reference in New Issue
Block a user