Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9660323ff2 | ||
|
|
b75ab5ecdc | ||
|
|
058f45115d | ||
|
|
fc1d74570f | ||
|
|
0835335978 | ||
|
|
d9c2c5e207 | ||
|
|
3c0ddbdbd9 |
@@ -13,6 +13,7 @@ from config.settings import settings
|
||||
from nl2sql.contracts import DataQueryRequest, DataQueryResult
|
||||
from nl2sql.embedding import EmbeddingError
|
||||
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.schema import load_authoritative_schema
|
||||
from repositories.customer_relation import CustomerRelationRepo
|
||||
@@ -55,7 +56,9 @@ def enrich_customer_rows(
|
||||
|
||||
def _is_current_holdings_query(question: str) -> bool:
|
||||
text = "".join((question or "").split())
|
||||
return any(term in text for term in ("当前持仓", "目前持仓", "现有持仓", "持仓明细", "持仓情况"))
|
||||
if any(term in text for term in ("当前持仓", "目前持仓", "现有持仓", "持仓明细", "持仓情况")):
|
||||
return True
|
||||
return "持仓" in text and not any(term in text for term in ("历史", "曾经", "已卖出", "交易记录"))
|
||||
|
||||
|
||||
async def _query_current_holdings(
|
||||
@@ -154,6 +157,7 @@ async def execute_advisor_data_query(
|
||||
pass
|
||||
|
||||
if scope == "customer" and _is_current_holdings_query(question):
|
||||
redis = redis or database.redis.client()
|
||||
return await _query_current_holdings(
|
||||
db,
|
||||
customer_id=customer_id,
|
||||
@@ -166,6 +170,21 @@ async def execute_advisor_data_query(
|
||||
|
||||
milvus = milvus or database.milvus.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
|
||||
request = DataQueryRequest(
|
||||
question=question,
|
||||
@@ -222,4 +241,5 @@ async def execute_advisor_data_query(
|
||||
payload["answer"] = f"{existing_answer.rstrip('。')}。客户姓名:{enriched['name_summary']}。"
|
||||
payload["sql"] = None
|
||||
payload["customer_id"] = customer_id
|
||||
await cache_set(redis, cache_key, payload, ttl=runtime_config.cache_ttl)
|
||||
return payload
|
||||
|
||||
@@ -35,20 +35,58 @@ class IntentClassification:
|
||||
reason: str = ""
|
||||
|
||||
|
||||
_CLASSIFIER_PROMPT = """你是基金投顾工作台的意图分类器,只负责分类,不回答用户问题。
|
||||
只能从以下分类中选择一个:
|
||||
- recommend:基金推荐、组合配置或投资方案
|
||||
- rebalance:组合偏离、调仓、再平衡、仓位调整
|
||||
- fund_analysis:单只基金分析、净值、收益、回撤、波动率、夏普比率
|
||||
- dialogue-script:给客户准备沟通话术、解释、安抚、投诉或风险提醒
|
||||
- data_query:查询客户持仓、资产、余额、收益、交易、账户明细
|
||||
- casual_chat:问候、闲聊、感谢、身份询问或无法归入业务分类的内容
|
||||
_CLASSIFIER_PROMPT = """# 任务:投顾助手意图识别
|
||||
你是投顾系统意图分类器,对客户输入文本做意图识别,**只能输出一个类别名称,禁止额外解释**。
|
||||
|
||||
## 类别说明
|
||||
1. 基金推荐:客户希望推荐、筛选基金产品
|
||||
2. 调仓建议:客户询问是否买卖、加减仓、更换基金,寻求调仓交易建议
|
||||
3. 客户持仓基金分析:分析客户现有持仓组合、风险、收益情况,不涉及买卖操作建议
|
||||
4. 跟客户的沟通话术:投顾需要生成一段发给客户的话术文案。【注意:客户发起提问,不会是该类别】
|
||||
5. 数据查询:查询基金客观数据,如净值、基金经理、规模、持仓、费率等事实信息,不做推荐、诊断
|
||||
6. 普通聊天:日常问候、单纯情绪吐槽,无明确业务诉求
|
||||
|
||||
## 判定优先级(同时存在多个诉求时,取优先级最高)
|
||||
调仓建议 > 基金推荐 > 客户持仓基金分析 > 数据查询 > 跟客户的沟通话术 > 普通聊天
|
||||
|
||||
## 示例(xx表示我名下的客户名字)
|
||||
输入:帮我名下XX推荐几只适合养老的基金
|
||||
输出:基金推荐
|
||||
|
||||
输入:XX手上的某某基金现在要不要卖出?
|
||||
输出:调仓建议
|
||||
|
||||
输入:帮我看看xx的基金组合风险高不高
|
||||
输出:客户持仓基金分析
|
||||
|
||||
输入:帮我查一下XX基金最新规模
|
||||
输出:数据查询
|
||||
|
||||
输入:帮我写一段话安抚客户,解释近期回撤
|
||||
输出:跟客户的沟通话术
|
||||
|
||||
输入:今天天气不错
|
||||
输出:普通聊天
|
||||
|
||||
输入:xx持有的这几只基金波动很大,要不要减仓?
|
||||
输出:调仓建议
|
||||
|
||||
输入:帮我看下xx持仓的基金,它们的基金经理是谁
|
||||
输出:数据查询
|
||||
|
||||
现在开始分类
|
||||
输入:{{user_query}}
|
||||
输出:
|
||||
|
||||
只输出 JSON,不要 Markdown,不要额外文字:
|
||||
{"intent":"分类值","confidence":0到1之间的数字,"reason":"不超过30字的原因"}
|
||||
"""
|
||||
|
||||
_RECOMMENDATION_TERMS = ("推荐", "组合建议", "配置建议", "买什么", "适合配置", "筛选基金", "投资方案")
|
||||
_REBALANCE_TERMS = ("调仓", "再平衡", "组合调整", "配置偏离", "偏离目标")
|
||||
_DIALOGUE_TERMS = ("话术", "沟通", "怎么跟客户说", "如何向客户解释", "安抚客户", "投诉处理")
|
||||
_CUSTOMER_DATA_TERMS = ("持仓", "资产", "余额", "交易", "账户", "份额", "市值", "盈亏", "资金")
|
||||
_REFERENCE_TERMS = ("他", "她", "它", "这个客户", "该客户", "那个客户", "刚才", "上一轮")
|
||||
_SCOPE_ERROR_MESSAGE = "投顾范围查询仅支持客户数据查询"
|
||||
|
||||
|
||||
@@ -83,11 +121,31 @@ def _parse_model_result(raw: str) -> IntentClassification | None:
|
||||
)
|
||||
|
||||
|
||||
def resolve_advisor_intent(
|
||||
query: str | None,
|
||||
classified_intent: str,
|
||||
*,
|
||||
customer_resolved: bool,
|
||||
) -> str:
|
||||
"""结合实体解析结果做最终业务路由,处理一句话中的明确优先意图。"""
|
||||
text = re.sub(r"\s+", "", query or "")
|
||||
if any(term in text for term in _REBALANCE_TERMS):
|
||||
return AGENT_INTENT_REBALANCE
|
||||
if any(term in text for term in _RECOMMENDATION_TERMS):
|
||||
return AGENT_INTENT_RECOMMEND
|
||||
if any(term in text for term in _DIALOGUE_TERMS):
|
||||
return AGENT_INTENT_DIALOGUE_SCRIPT
|
||||
if customer_resolved and any(term in text for term in _CUSTOMER_DATA_TERMS):
|
||||
return AGENT_INTENT_DATA_QUERY
|
||||
return classified_intent
|
||||
|
||||
|
||||
async def classify_advisor_intent(
|
||||
query: str | None,
|
||||
llm_client=None,
|
||||
*,
|
||||
explicit_intent: str | None = None,
|
||||
conversation_context: str = "",
|
||||
timeout: float = 2.0,
|
||||
) -> IntentClassification:
|
||||
"""分类用户意图;显式意图兼容旧客户端,模型失败时安全回退规则。"""
|
||||
@@ -102,13 +160,39 @@ async def classify_advisor_intent(
|
||||
"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",
|
||||
"明显查询类问题",
|
||||
)
|
||||
normalized_query = re.sub(r"\s+", "", query)
|
||||
if (
|
||||
any(term in normalized_query for term in _CUSTOMER_DATA_TERMS)
|
||||
and conversation_context
|
||||
and any(term in normalized_query for term in _REFERENCE_TERMS)
|
||||
):
|
||||
return IntentClassification(
|
||||
AGENT_INTENT_DATA_QUERY,
|
||||
0.9,
|
||||
"rule_context",
|
||||
"上下文中的客户数据追问",
|
||||
)
|
||||
if llm_client is not None:
|
||||
try:
|
||||
prompt = query.strip()
|
||||
if conversation_context:
|
||||
prompt = f"会话上下文:\n{conversation_context[:2000]}\n\n当前问题:\n{prompt}"
|
||||
raw = await asyncio.wait_for(
|
||||
llm_client.chat(
|
||||
[
|
||||
{"role": "system", "content": _CLASSIFIER_PROMPT},
|
||||
{"role": "user", "content": query.strip()},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
temperature=0,
|
||||
max_tokens=120,
|
||||
@@ -137,4 +221,4 @@ async def classify_advisor_intent(
|
||||
return _rule_fallback(query)
|
||||
|
||||
|
||||
__all__ = ["IntentClassification", "classify_advisor_intent"]
|
||||
__all__ = ["IntentClassification", "classify_advisor_intent", "resolve_advisor_intent"]
|
||||
|
||||
@@ -26,6 +26,7 @@ _QUERY_ACTIONS = (
|
||||
_DATA_TERMS = (
|
||||
"持仓",
|
||||
"资产",
|
||||
"资金",
|
||||
"收益",
|
||||
"交易记录",
|
||||
"申购",
|
||||
@@ -41,6 +42,15 @@ _FUND_ANALYSIS_TERMS = ("基金分析", "分析基金", "基金表现", "净值
|
||||
_DIALOGUE_TERMS = ("话术", "沟通", "怎么跟客户说", "如何向客户解释", "安抚客户", "投诉处理")
|
||||
_RECOMMEND_TERMS = ("推荐", "组合建议", "配置建议", "买什么", "适合配置", "筛选基金", "投资方案")
|
||||
_CUSTOMER_IDENTITY_TERMS = ("是谁", "姓名", "实名", "基本信息", "联系方式", "手机号")
|
||||
_CUSTOMER_ROSTER_TERMS = ("名单", "列表", "几个")
|
||||
_CUSTOMER_RISK_TERMS = (
|
||||
"风险评级", "风险等级", "风险类型", "激进型", "进取型", "平衡型", "稳健型", "保守型",
|
||||
"C1", "C2", "C3", "C4", "C5",
|
||||
)
|
||||
_IMPLICIT_QUERY_TERMS = (
|
||||
"谁", "哪个", "哪些", "多少", "最大", "最多", "最小", "最少", "最高", "最低",
|
||||
"合计", "总额", "总资产", "排名", "前几", "前十", "超过", "低于", "是否", "有没有",
|
||||
)
|
||||
_SCOPE_ERROR_MESSAGE = "投顾范围查询仅支持客户数据查询"
|
||||
|
||||
|
||||
@@ -79,7 +89,23 @@ def recognize_advisor_intent(query: str | None, explicit_intent: str | None = No
|
||||
)
|
||||
if has_customer_identity:
|
||||
return AGENT_INTENT_DATA_QUERY
|
||||
if has_data and (has_action or "客户" in text or "近一年" in text or "本月" in text):
|
||||
# 客户名册类提问(我名下有哪些客户/客户名单)也属于数据查询。
|
||||
has_customer_roster = (
|
||||
"客户" in text
|
||||
and (has_action or _contains_any(text, _CUSTOMER_ROSTER_TERMS))
|
||||
and not _contains_any(text, _RECOMMEND_TERMS)
|
||||
)
|
||||
if has_customer_roster:
|
||||
return AGENT_INTENT_DATA_QUERY
|
||||
if "客户" in text and _contains_any(text, _CUSTOMER_RISK_TERMS) and not _contains_any(text, _RECOMMEND_TERMS):
|
||||
return AGENT_INTENT_DATA_QUERY
|
||||
if has_data and (
|
||||
has_action
|
||||
or "客户" in text
|
||||
or "近一年" in text
|
||||
or "本月" in text
|
||||
or _contains_any(text, _IMPLICIT_QUERY_TERMS)
|
||||
):
|
||||
return AGENT_INTENT_DATA_QUERY
|
||||
if _contains_any(text, _RECOMMEND_TERMS):
|
||||
return AGENT_INTENT_RECOMMEND
|
||||
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from agent.advisor_agent.auth import ensure_customer_access
|
||||
from agent.advisor_agent.data_query import execute_advisor_data_query
|
||||
from agent.advisor_agent.intent.fund_analysis import build_fund_analysis
|
||||
from agent.advisor_agent.intent.classifier import classify_advisor_intent
|
||||
from agent.advisor_agent.intent.classifier import classify_advisor_intent, resolve_advisor_intent
|
||||
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 (
|
||||
@@ -69,6 +69,7 @@ from schemas.advisor_agent import (
|
||||
AdvisorDataQueryReq,
|
||||
)
|
||||
from service.nl2sql.query_service import QueryServiceError
|
||||
from nl2sql.query_rewriter import rewrite_query
|
||||
from service.advisor_agent.draft import (
|
||||
detail_draft,
|
||||
discard_draft,
|
||||
@@ -248,10 +249,20 @@ async def chat_stream(
|
||||
)
|
||||
|
||||
runtime = _advisor_runtime(request)
|
||||
classification = await classify_advisor_intent(
|
||||
context_store = SessionContextStore(redis_db.client())
|
||||
conversation_context = build_conversation_context(
|
||||
await context_store.load(user.id, chat_request.session_id)
|
||||
)
|
||||
effective_question = await rewrite_query(
|
||||
chat_request.query,
|
||||
conversation_context,
|
||||
llm_client=getattr(runtime, "llm_client", None),
|
||||
)
|
||||
classification = await classify_advisor_intent(
|
||||
effective_question,
|
||||
getattr(runtime, "llm_client", None),
|
||||
explicit_intent=None,
|
||||
conversation_context=conversation_context,
|
||||
)
|
||||
inferred_intent = classification.intent
|
||||
scope = chat_request.scope
|
||||
@@ -277,11 +288,16 @@ async def chat_stream(
|
||||
await ensure_customer_access(db, advisor_id=user.id, customer_id=int(customer_id))
|
||||
else:
|
||||
resolved_customer_id, _resolve_error = await _resolve_customer_from_query(
|
||||
db, advisor_id=user.id, query=chat_request.query
|
||||
db, advisor_id=user.id, query=effective_question
|
||||
)
|
||||
if resolved_customer_id is not None:
|
||||
customer_id = resolved_customer_id
|
||||
scope = "customer"
|
||||
inferred_intent = resolve_advisor_intent(
|
||||
effective_question,
|
||||
inferred_intent,
|
||||
customer_resolved=customer_id is not None,
|
||||
)
|
||||
|
||||
if customer_id is None:
|
||||
if inferred_intent == AGENT_INTENT_DATA_QUERY:
|
||||
@@ -317,20 +333,17 @@ async def chat_stream(
|
||||
if inferred_intent == AGENT_INTENT_DATA_QUERY:
|
||||
effective_scope = "customer" if customer_id is not None else "advisor"
|
||||
try:
|
||||
context_store = SessionContextStore(redis_db.client())
|
||||
conversation_context = build_conversation_context(
|
||||
await context_store.load(user.id, chat_request.session_id)
|
||||
)
|
||||
result = await execute_advisor_data_query(
|
||||
db,
|
||||
advisor_id=user.id,
|
||||
customer_id=int(customer_id) if customer_id is not None else None,
|
||||
scope=effective_scope,
|
||||
question=chat_request.query,
|
||||
question=effective_question,
|
||||
trace_id=trace_id,
|
||||
session_id=chat_request.session_id,
|
||||
conversation_context=conversation_context,
|
||||
llm_client=getattr(_advisor_runtime(request), "llm_client", None),
|
||||
redis=redis_db.client(),
|
||||
)
|
||||
except QueryServiceError as exc:
|
||||
payload = agent_failure(ERR_CODE_LLM_ERROR, _data_query_error_message(exc), trace_id=trace_id)
|
||||
@@ -565,6 +578,7 @@ async def advisor_data_query(
|
||||
page_size=body.page_size,
|
||||
sort_by=body.sort_by,
|
||||
sort_order=body.sort_order,
|
||||
redis=redis_db.client(),
|
||||
)
|
||||
except QueryServiceError as exc:
|
||||
return agent_failure(
|
||||
|
||||
@@ -248,6 +248,7 @@ async def query_data(
|
||||
db,
|
||||
database=settings.mysql.database,
|
||||
candidate_tables=table_names,
|
||||
redis=redis,
|
||||
)
|
||||
|
||||
request_contract = DataQueryRequest(
|
||||
|
||||
@@ -354,16 +354,19 @@
|
||||
<div class="ca-head">
|
||||
<div class="row" style="gap:10px">
|
||||
<div class="brand-logo" style="width:34px;height:34px;font-size:15px">🤖</div>
|
||||
<div><div class="t">华夏基金智能助手</div><div class="s">在线为您提供基金信息服务</div></div>
|
||||
<div><div class="t">华夏基金智能助手</div><div class="s" id="ca-state">在线为您提供基金信息服务</div></div>
|
||||
</div>
|
||||
<button class="btn btn-ghost btn-sm" style="color:#bfdbfe" onclick="toggleCA()">✕</button>
|
||||
<button class="btn btn-ghost btn-sm" style="color:#bfdbfe" onclick="closeCA()">✕</button>
|
||||
</div>
|
||||
<div class="ca-body">
|
||||
<div class="m bot"><div class="bl">您好,我是华夏基金智能助手。您可以咨询基金产品、风险等级和净值信息。</div></div>
|
||||
<div class="m usr"><div class="bl">稳健型的我能买股票型基金吗?</div></div>
|
||||
<div class="m bot"><div class="bl">风险等级 C2(保守型)通常不建议购买 R4 股票型基金,
|
||||
两者风险等级不匹配。您可以关注 R2 债券型或 R3 混合型产品。<br><br>
|
||||
<span class="faint xs">以上信息仅供参考,不构成投资建议。</span></div></div>
|
||||
<div id="ca-connecting" class="xs" style="color:#94a3b8;padding:6px 0">正在建立会话...</div>
|
||||
<div id="ca-messages" style="display:none">
|
||||
<div class="m bot"><div class="bl">您好,我是华夏基金智能助手。您可以咨询基金产品、风险等级和净值信息。</div></div>
|
||||
<div class="m usr"><div class="bl">稳健型的我能买股票型基金吗?</div></div>
|
||||
<div class="m bot"><div class="bl">风险等级 C2(保守型)通常不建议购买 R4 股票型基金,
|
||||
两者风险等级不匹配。您可以关注 R2 债券型或 R3 混合型产品。<br><br>
|
||||
<span class="faint xs">以上信息仅供参考,不构成投资建议。</span></div></div>
|
||||
</div>
|
||||
</div>
|
||||
<div class="ca-foot">
|
||||
<input class="input" placeholder="输入您的问题">
|
||||
@@ -1020,11 +1023,35 @@ function showScreen(id){
|
||||
/* 客户/游客助手挂在全局 layout,后台不显示(后台用投顾助手悬浮球) */
|
||||
document.getElementById("client-agent").style.display =
|
||||
NO_CLIENT_AGENT.includes(id) ? "none" : "block";
|
||||
document.getElementById("ca-win").style.display = "none";
|
||||
closeCA();
|
||||
}
|
||||
/* 打开面板 = 先建会话(session/create),拿到会话号后才允许对话 */
|
||||
function openCA(){
|
||||
const w = document.getElementById("ca-win");
|
||||
if (w.style.display !== "none") return;
|
||||
w.style.display = "flex";
|
||||
const state = document.getElementById("ca-state");
|
||||
document.getElementById("ca-connecting").style.display = "block";
|
||||
document.getElementById("ca-messages").style.display = "none";
|
||||
state.textContent = "正在建立会话...";
|
||||
clearTimeout(window.__caTimer);
|
||||
window.__caTimer = setTimeout(()=>{
|
||||
state.textContent = "会话 #" + Math.random().toString(16).slice(2,8) + " · 已接入您的账户数据";
|
||||
document.getElementById("ca-connecting").style.display = "none";
|
||||
document.getElementById("ca-messages").style.display = "block";
|
||||
}, 700);
|
||||
}
|
||||
/* 关闭面板 = 结束会话(session/end),丢弃服务端上下文 */
|
||||
function closeCA(){
|
||||
const w = document.getElementById("ca-win");
|
||||
if (w.style.display === "none") return;
|
||||
w.style.display = "none";
|
||||
clearTimeout(window.__caTimer);
|
||||
document.getElementById("ca-state").textContent = "会话已结束";
|
||||
}
|
||||
function toggleCA(){
|
||||
const w = document.getElementById("ca-win");
|
||||
w.style.display = w.style.display==="none" ? "flex" : "none";
|
||||
w.style.display === "none" ? openCA() : closeCA();
|
||||
}
|
||||
function go(id){
|
||||
const sec = ADMIN_SECTIONS.find(s=>s.id===id);
|
||||
|
||||
@@ -58,7 +58,8 @@ export function AssistantLauncher() {
|
||||
const [open, setOpen] = useState(false);
|
||||
const [tab, setTab] = useState<"chat" | "drafts">("chat");
|
||||
|
||||
const [sessionId, setSessionId] = useState(() => newSessionId());
|
||||
/* 会话在「打开抽屉」时才建立(见 openAssistant),关闭时丢弃,不在挂载时就占用上下文 */
|
||||
const [sessionId, setSessionId] = useState<string | null>(null);
|
||||
const [customers, setCustomers] = useState<Customer[]>([]);
|
||||
const [customerId, setCustomerId] = useState<number | null>(null);
|
||||
const [messages, setMessages] = useState<ChatMessage[]>([]);
|
||||
@@ -87,7 +88,7 @@ export function AssistantLauncher() {
|
||||
return () => clearInterval(timer);
|
||||
}, [streaming]);
|
||||
|
||||
const sessionTag = useMemo(() => sessionId.slice(-6), [sessionId]);
|
||||
const sessionTag = useMemo(() => (sessionId ? sessionId.slice(-6) : "--------"), [sessionId]);
|
||||
|
||||
/* 客户列表只用于会话上下文选择,失败不阻断助手 */
|
||||
useEffect(() => {
|
||||
@@ -101,8 +102,18 @@ export function AssistantLauncher() {
|
||||
})();
|
||||
}, []);
|
||||
|
||||
/* 会话 ID 变化时恢复该会话的历史;新会话为空属正常,不报错。 */
|
||||
/*
|
||||
* 会话握手:打开抽屉(sessionId 就绪)后拉一次该会话的记录。
|
||||
* 投顾 Agent 没有独立的 session/create,后端按 (advisor_id, session_id) 惰性建立上下文,
|
||||
* 所以这个请求就是「打开助手时发出的会话请求」;关闭抽屉会把 sessionId 置空、丢弃上下文。
|
||||
*/
|
||||
useEffect(() => {
|
||||
if (!sessionId) {
|
||||
setMessages([]);
|
||||
setHistoryError(null);
|
||||
setHistoryLoading(false);
|
||||
return;
|
||||
}
|
||||
let alive = true;
|
||||
void (async () => {
|
||||
setHistoryLoading(true);
|
||||
@@ -182,7 +193,7 @@ export function AssistantLauncher() {
|
||||
const result = await assistantApi.chatOnce(
|
||||
{
|
||||
query,
|
||||
session_id: sessionId,
|
||||
session_id: sessionId ?? undefined,
|
||||
scope: customerId ? "customer" : "advisor",
|
||||
customer_id: customerId,
|
||||
},
|
||||
@@ -241,7 +252,7 @@ export function AssistantLauncher() {
|
||||
setStreaming(true);
|
||||
try {
|
||||
if (tool === "data-query") {
|
||||
const result = await assistantApi.dataQuery({ query, session_id: sessionId });
|
||||
const result = await assistantApi.dataQuery({ query, session_id: sessionId ?? undefined });
|
||||
pushPair(query, result.answer ?? result.summary ?? "查询完成,无文本摘要。", "data_query");
|
||||
} else if (tool === "fund-analysis") {
|
||||
const result = await assistantApi.fundAnalysis(query);
|
||||
@@ -272,7 +283,30 @@ export function AssistantLauncher() {
|
||||
[customerId, input, loadDrafts, pushPair, sessionId, toast]
|
||||
);
|
||||
|
||||
/** 结束会话:中止生成、清空上下文并启用新的 session_id。 */
|
||||
/**
|
||||
* 打开抽屉 = 建立会话。
|
||||
* 投顾 Agent 后端只有 chat/stream 与 session/{id}/history,没有独立的 session/create,
|
||||
* 所以这里生成会话键,由上面的 useEffect 调 history 完成握手(后端按需建立上下文)。
|
||||
*/
|
||||
const openAssistant = useCallback(() => {
|
||||
setOpen(true);
|
||||
setSessionId((current) => current ?? newSessionId());
|
||||
}, []);
|
||||
|
||||
/**
|
||||
* 关闭抽屉 = 结束会话:中止生成、清空消息并丢弃上下文。
|
||||
* 投顾侧后端未提供 session/end,因此结束只能在前端完成(下次打开是新的 session_id)。
|
||||
*/
|
||||
const closeAssistant = useCallback(() => {
|
||||
abortRef.current?.abort();
|
||||
setOpen(false);
|
||||
setMessages([]);
|
||||
setInput("");
|
||||
setStreaming(false);
|
||||
setSessionId(null);
|
||||
}, []);
|
||||
|
||||
/** 抽屉内「结束会话」:不等关闭,立即换一个全新会话并重新握手。 */
|
||||
const endSession = useCallback(() => {
|
||||
abortRef.current?.abort();
|
||||
setMessages([]);
|
||||
@@ -328,7 +362,7 @@ export function AssistantLauncher() {
|
||||
<>
|
||||
{!open && (
|
||||
<button
|
||||
onClick={() => setOpen(true)}
|
||||
onClick={openAssistant}
|
||||
aria-label="打开投顾助手"
|
||||
className="fixed bottom-6 right-6 z-40 flex h-14 w-14 items-center justify-center rounded-full bg-indigo-600 text-white shadow-lg shadow-indigo-600/30 transition-transform hover:scale-105 hover:bg-indigo-700"
|
||||
>
|
||||
@@ -345,7 +379,7 @@ export function AssistantLauncher() {
|
||||
<>
|
||||
<div
|
||||
className="fixed inset-0 z-40 bg-slate-900/10"
|
||||
onClick={() => setOpen(false)}
|
||||
onClick={closeAssistant}
|
||||
aria-hidden
|
||||
/>
|
||||
<aside className="fixed inset-y-0 right-0 z-50 flex w-full max-w-[560px] flex-col bg-white shadow-2xl">
|
||||
@@ -370,7 +404,7 @@ export function AssistantLauncher() {
|
||||
结束会话
|
||||
</Button>
|
||||
<button
|
||||
onClick={() => setOpen(false)}
|
||||
onClick={closeAssistant}
|
||||
aria-label="收起助手"
|
||||
className="rounded-lg p-2 text-slate-400 transition-colors hover:bg-slate-100 hover:text-slate-600"
|
||||
>
|
||||
@@ -414,7 +448,7 @@ export function AssistantLauncher() {
|
||||
<>
|
||||
<div ref={scrollRef} className="flex-1 space-y-4 overflow-y-auto px-5 py-4">
|
||||
{historyLoading && messages.length === 0 && (
|
||||
<p className="py-8 text-center text-sm text-slate-400">正在恢复会话...</p>
|
||||
<p className="py-8 text-center text-sm text-slate-400">正在建立会话...</p>
|
||||
)}
|
||||
{historyError && (
|
||||
<p className="rounded-lg bg-red-50 px-3 py-2 text-xs text-red-600">{historyError}</p>
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
"use client";
|
||||
|
||||
import { Bot, MessageCircle, Send, X } from "lucide-react";
|
||||
import { useState } from "react";
|
||||
import { apiFetch, apiStream } from "@/lib/api";
|
||||
|
||||
interface AgentChatProps {
|
||||
employeeMode?: boolean;
|
||||
}
|
||||
import { Bot, MessageCircle, RotateCw, Send, X } from "lucide-react";
|
||||
import { useCallback, useEffect, useRef, useState } from "react";
|
||||
import { apiFetch } from "@/lib/api";
|
||||
import { getToken } from "@/lib/auth";
|
||||
import { chatAgentApi, type ChatAgentMode } from "@/lib/chat-agent-api";
|
||||
import { USER_TYPE } from "@/lib/roles";
|
||||
|
||||
interface ChatMessage {
|
||||
id: number;
|
||||
@@ -14,72 +13,233 @@ interface ChatMessage {
|
||||
content: string;
|
||||
}
|
||||
|
||||
export function AgentChat({ employeeMode = false }: AgentChatProps) {
|
||||
const GREETING = "您好,我是华夏基金智能助手。您可以咨询基金产品、风险等级和净值信息。";
|
||||
|
||||
/**
|
||||
* 选择用哪个 Agent:
|
||||
* - 已登录客户 → /agent/client/*(能读本人持仓与交易)
|
||||
* - 游客 / 员工 → /agent/customer/*(匿名客服,不需要 token)
|
||||
* 探测失败一律退回匿名客服,不阻断用户。
|
||||
*/
|
||||
async function resolveMode(): Promise<ChatAgentMode> {
|
||||
if (!getToken()) return "customer";
|
||||
try {
|
||||
const result = await apiFetch<{ user?: { user_type?: string | null } }>("/auth/me");
|
||||
return result?.user?.user_type === USER_TYPE.CUSTOMER ? "client" : "customer";
|
||||
} catch {
|
||||
return "customer";
|
||||
}
|
||||
}
|
||||
|
||||
export function AgentChat() {
|
||||
const [open, setOpen] = useState(false);
|
||||
const [input, setInput] = useState("");
|
||||
const [loading, setLoading] = useState(false);
|
||||
const [sending, setSending] = useState(false);
|
||||
const [connecting, setConnecting] = useState(false);
|
||||
const [sessionId, setSessionId] = useState<string | null>(null);
|
||||
const [mode, setMode] = useState<ChatAgentMode>("customer");
|
||||
const [error, setError] = useState<string | null>(null);
|
||||
const [messages, setMessages] = useState<ChatMessage[]>([
|
||||
{ id: 1, role: "agent", content: "您好,我是华夏基金智能助手。您可以咨询基金产品、风险等级和净值信息。" },
|
||||
{ id: 1, role: "agent", content: GREETING },
|
||||
]);
|
||||
|
||||
async function ensureSession() {
|
||||
if (employeeMode) return "";
|
||||
if (sessionId) return sessionId;
|
||||
const result = await apiFetch<{ session_id: string }>("/agent/customer/session/create", { method: "POST" });
|
||||
setSessionId(result.session_id);
|
||||
return result.session_id;
|
||||
}
|
||||
const abortRef = useRef<AbortController | null>(null);
|
||||
const listRef = useRef<HTMLDivElement | null>(null);
|
||||
|
||||
async function sendMessage() {
|
||||
const content = input.trim();
|
||||
if (!content || loading) return;
|
||||
setInput("");
|
||||
setMessages((current) => [...current, { id: Date.now(), role: "user", content }]);
|
||||
setLoading(true);
|
||||
useEffect(() => () => abortRef.current?.abort(), []);
|
||||
useEffect(() => {
|
||||
listRef.current?.scrollTo({ top: listRef.current.scrollHeight, behavior: "smooth" });
|
||||
}, [messages, connecting]);
|
||||
|
||||
/** 建立会话:探测 Agent 类型 → session/create 换取后端签发的 session_id。 */
|
||||
const connect = useCallback(async () => {
|
||||
setConnecting(true);
|
||||
setError(null);
|
||||
try {
|
||||
const activeSessionId = await ensureSession();
|
||||
const response = employeeMode
|
||||
? await apiStream("/advisor-agent/chat/stream", { query: content })
|
||||
: await apiStream("/agent/customer/chat", { session_id: activeSessionId, query: content });
|
||||
setMessages((current) => [...current, { id: Date.now() + 1, role: "agent", content: response || "暂时没有找到合适的回答,请稍后再试。" }]);
|
||||
} catch {
|
||||
setMessages((current) => [...current, { id: Date.now() + 1, role: "agent", content: "当前服务暂时不可用,请稍后再试。" }]);
|
||||
const resolved = await resolveMode();
|
||||
const session = await chatAgentApi.createSession(resolved);
|
||||
setMode(resolved);
|
||||
setSessionId(session.session_id);
|
||||
} catch (err) {
|
||||
setError(err instanceof Error ? err.message : "会话建立失败,请重试");
|
||||
} finally {
|
||||
setLoading(false);
|
||||
setConnecting(false);
|
||||
}
|
||||
}
|
||||
}, []);
|
||||
|
||||
async function closeChat() {
|
||||
/** 打开面板即建立会话(不是等用户第一次提问才建)。 */
|
||||
const openChat = useCallback(() => {
|
||||
setOpen(true);
|
||||
if (!sessionId && !connecting) void connect();
|
||||
}, [connect, connecting, sessionId]);
|
||||
|
||||
/** 关闭面板即结束会话:通知后端清理上下文与限流键,本地一并重置。 */
|
||||
const closeChat = useCallback(async () => {
|
||||
abortRef.current?.abort();
|
||||
setOpen(false);
|
||||
if (!sessionId || employeeMode) return;
|
||||
try { await apiFetch("/agent/customer/session/end", { method: "POST", body: JSON.stringify({ session_id: sessionId }) }); } catch { /* session cleanup is best effort */ }
|
||||
setInput("");
|
||||
setError(null);
|
||||
setMessages([{ id: Date.now(), role: "agent", content: GREETING }]);
|
||||
const current = sessionId;
|
||||
setSessionId(null);
|
||||
}
|
||||
if (!current) return;
|
||||
try {
|
||||
await chatAgentApi.endSession(mode, current);
|
||||
} catch {
|
||||
/* 会话清理失败不阻断用户 */
|
||||
}
|
||||
}, [mode, sessionId]);
|
||||
|
||||
const sendMessage = useCallback(async () => {
|
||||
const content = input.trim();
|
||||
if (!content || sending || connecting) return;
|
||||
if (!sessionId) {
|
||||
setError("会话尚未建立,请点击重试。");
|
||||
return;
|
||||
}
|
||||
setInput("");
|
||||
setError(null);
|
||||
setMessages((prev) => [...prev, { id: Date.now(), role: "user", content }]);
|
||||
setSending(true);
|
||||
|
||||
const controller = new AbortController();
|
||||
abortRef.current = controller;
|
||||
try {
|
||||
const result = await chatAgentApi.chat(
|
||||
mode,
|
||||
{ session_id: sessionId, query: content },
|
||||
controller.signal
|
||||
);
|
||||
setMessages((prev) => [
|
||||
...prev,
|
||||
{
|
||||
id: Date.now() + 1,
|
||||
role: "agent",
|
||||
content: result.error || result.text || "暂时没有找到合适的回答,请稍后再试。",
|
||||
},
|
||||
]);
|
||||
} catch (err) {
|
||||
if ((err as Error)?.name === "AbortError") return;
|
||||
setMessages((prev) => [
|
||||
...prev,
|
||||
{
|
||||
id: Date.now() + 1,
|
||||
role: "agent",
|
||||
content: err instanceof Error ? err.message : "当前服务暂时不可用,请稍后再试。",
|
||||
},
|
||||
]);
|
||||
} finally {
|
||||
setSending(false);
|
||||
abortRef.current = null;
|
||||
}
|
||||
}, [connecting, input, mode, sending, sessionId]);
|
||||
|
||||
const ready = !!sessionId;
|
||||
|
||||
return (
|
||||
<div className="fixed bottom-3 right-7 z-50">
|
||||
{open && (
|
||||
<div role="dialog" aria-label="华夏基金智能助手" className="mb-4 flex h-[520px] w-[380px] flex-col overflow-hidden rounded-2xl border border-slate-200 bg-white shadow-2xl">
|
||||
<div
|
||||
role="dialog"
|
||||
aria-label="华夏基金智能助手"
|
||||
className="mb-4 flex h-[520px] w-[380px] flex-col overflow-hidden rounded-2xl border border-slate-200 bg-white shadow-2xl"
|
||||
>
|
||||
<div className="flex items-center justify-between bg-[var(--brand-navy)] px-5 py-4 text-white">
|
||||
<div className="flex items-center gap-3">
|
||||
<div className="flex h-9 w-9 items-center justify-center rounded-xl bg-[var(--brand-primary)]"><Bot className="h-5 w-5" /></div>
|
||||
<div><p className="text-sm font-semibold">华夏基金智能助手</p><p className="mt-0.5 text-xs text-blue-100">在线为您提供基金信息服务</p></div>
|
||||
<div className="flex h-9 w-9 items-center justify-center rounded-xl bg-[var(--brand-primary)]">
|
||||
<Bot className="h-5 w-5" />
|
||||
</div>
|
||||
<div>
|
||||
<p className="text-sm font-semibold">华夏基金智能助手</p>
|
||||
<p className="mt-0.5 text-xs text-blue-100">
|
||||
{mode === "client" ? "已接入您的账户数据" : "在线为您提供基金信息服务"}
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
<button aria-label="关闭聊天框" onClick={() => void closeChat()} className="rounded-lg p-2 text-blue-100 transition hover:bg-white/10 hover:text-white"><X className="h-4 w-4" /></button>
|
||||
<button
|
||||
aria-label="关闭聊天框"
|
||||
onClick={() => void closeChat()}
|
||||
className="rounded-lg p-2 text-blue-100 transition hover:bg-white/10 hover:text-white"
|
||||
>
|
||||
<X className="h-4 w-4" />
|
||||
</button>
|
||||
</div>
|
||||
<div className="flex-1 space-y-4 overflow-y-auto bg-slate-50 p-4">
|
||||
{messages.map((message) => <div key={message.id} className={`flex ${message.role === "user" ? "justify-end" : "justify-start"}`}><div className={`max-w-[82%] rounded-2xl px-4 py-3 text-sm leading-6 ${message.role === "user" ? "rounded-br-md bg-[var(--brand-primary)] text-white" : "rounded-bl-md border border-slate-200 bg-white text-slate-700"}`}>{message.content}</div></div>)}
|
||||
{loading && <div className="text-xs text-slate-400">正在整理信息...</div>}
|
||||
|
||||
<div ref={listRef} className="flex-1 space-y-4 overflow-y-auto bg-slate-50 p-4">
|
||||
{messages.map((message) => (
|
||||
<div
|
||||
key={message.id}
|
||||
className={`flex ${message.role === "user" ? "justify-end" : "justify-start"}`}
|
||||
>
|
||||
<div
|
||||
className={`max-w-[82%] rounded-2xl px-4 py-3 text-sm leading-6 ${
|
||||
message.role === "user"
|
||||
? "rounded-br-md bg-[var(--brand-primary)] text-white"
|
||||
: "rounded-bl-md border border-slate-200 bg-white text-slate-700"
|
||||
}`}
|
||||
>
|
||||
{message.content}
|
||||
</div>
|
||||
</div>
|
||||
))}
|
||||
|
||||
{connecting && (
|
||||
<div className="flex items-center gap-2 text-xs text-slate-400">
|
||||
<RotateCw className="h-3 w-3 animate-spin" />
|
||||
正在建立会话...
|
||||
</div>
|
||||
)}
|
||||
{sending && !connecting && <div className="text-xs text-slate-400">正在整理信息...</div>}
|
||||
</div>
|
||||
<form onSubmit={(event) => { event.preventDefault(); void sendMessage(); }} className="flex gap-2 border-t border-slate-200 bg-white p-3">
|
||||
<input value={input} onChange={(event) => setInput(event.target.value)} placeholder="输入您的问题" className="min-w-0 flex-1 rounded-xl border border-slate-200 px-3 py-2.5 text-sm outline-none transition focus:border-[var(--brand-primary)]" />
|
||||
<button type="submit" aria-label="发送消息" className="flex h-10 w-10 shrink-0 items-center justify-center rounded-xl bg-[var(--brand-primary)] text-white transition hover:bg-[var(--brand-primary-dark)] disabled:cursor-not-allowed disabled:opacity-50" disabled={loading || !input.trim()}><Send className="h-4 w-4" /></button>
|
||||
|
||||
{error && (
|
||||
<div className="flex items-center justify-between gap-3 border-t border-rose-100 bg-rose-50 px-4 py-2.5">
|
||||
<span className="text-xs text-rose-600">{error}</span>
|
||||
{!ready && (
|
||||
<button
|
||||
onClick={() => void connect()}
|
||||
disabled={connecting}
|
||||
className="shrink-0 rounded-lg border border-rose-200 bg-white px-2.5 py-1 text-xs text-rose-600 transition hover:bg-rose-100 disabled:opacity-50"
|
||||
>
|
||||
重试
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<form
|
||||
onSubmit={(event) => {
|
||||
event.preventDefault();
|
||||
void sendMessage();
|
||||
}}
|
||||
className="flex gap-2 border-t border-slate-200 bg-white p-3"
|
||||
>
|
||||
<input
|
||||
value={input}
|
||||
onChange={(event) => setInput(event.target.value)}
|
||||
placeholder={ready ? "输入您的问题" : "正在准备会话..."}
|
||||
disabled={connecting}
|
||||
className="min-w-0 flex-1 rounded-xl border border-slate-200 px-3 py-2.5 text-sm outline-none transition focus:border-[var(--brand-primary)] disabled:bg-slate-50"
|
||||
/>
|
||||
<button
|
||||
type="submit"
|
||||
aria-label="发送消息"
|
||||
className="flex h-10 w-10 shrink-0 items-center justify-center rounded-xl bg-[var(--brand-primary)] text-white transition hover:bg-[var(--brand-primary-dark)] disabled:cursor-not-allowed disabled:opacity-50"
|
||||
disabled={sending || connecting || !input.trim()}
|
||||
>
|
||||
<Send className="h-4 w-4" />
|
||||
</button>
|
||||
</form>
|
||||
</div>
|
||||
)}
|
||||
<button aria-label={open ? "收起智能助手" : "打开智能助手"} onClick={() => setOpen((value) => !value)} className="flex h-14 w-14 items-center justify-center rounded-full bg-[var(--brand-primary)] text-white shadow-lg shadow-indigo-900/25 transition hover:-translate-y-0.5 hover:bg-[var(--brand-primary-dark)]">{open ? <X className="h-5 w-5" /> : <MessageCircle className="h-5 w-5" />}</button>
|
||||
|
||||
<button
|
||||
aria-label={open ? "收起智能助手" : "打开智能助手"}
|
||||
onClick={() => (open ? void closeChat() : openChat())}
|
||||
className="flex h-14 w-14 items-center justify-center rounded-full bg-[var(--brand-primary)] text-white shadow-lg shadow-indigo-900/25 transition hover:-translate-y-0.5 hover:bg-[var(--brand-primary-dark)]"
|
||||
>
|
||||
{open ? <X className="h-5 w-5" /> : <MessageCircle className="h-5 w-5" />}
|
||||
</button>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
+62
-22
@@ -29,6 +29,22 @@ function handleUnauthorized(status: number) {
|
||||
if (!window.location.pathname.startsWith("/login")) window.location.href = "/login";
|
||||
}
|
||||
|
||||
/**
|
||||
* 从错误响应体里取后端给的原因(如「会话不存在或已过期」)。
|
||||
* 非 2xx 时后端返回的是 JSON(utils/response.fail),部分场景也会包成 SSE。
|
||||
*/
|
||||
function extractErrorMessage(body: string): string | undefined {
|
||||
const text = body.trim();
|
||||
if (!text) return undefined;
|
||||
const jsonText = text.startsWith("data:") ? text.split("\n")[0].slice(5).trim() : text;
|
||||
try {
|
||||
const payload = JSON.parse(jsonText) as { message?: string; error?: string };
|
||||
return payload?.message ?? payload?.error;
|
||||
} catch {
|
||||
return undefined;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 后端存在两套成功码,必须都认:
|
||||
* - utils/response.success → code 200(工作台、风控、工单等)
|
||||
@@ -126,6 +142,8 @@ export interface SSESummary {
|
||||
traceId?: string;
|
||||
/** error 事件里的提示文案。 */
|
||||
error?: string;
|
||||
/** 客服 Agent 命中的知识来源(裸 JSON 响应里的 sources)。 */
|
||||
sources?: unknown[];
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -152,7 +170,13 @@ export async function readSSE(
|
||||
});
|
||||
|
||||
handleUnauthorized(response.status);
|
||||
if (!response.ok) throw new ApiError("请求失败,请稍后重试", response.status);
|
||||
if (!response.ok) {
|
||||
const body = await response.text();
|
||||
throw new ApiError(
|
||||
extractErrorMessage(body) ?? "请求失败,请稍后重试",
|
||||
response.status
|
||||
);
|
||||
}
|
||||
|
||||
const summary: SSESummary = { text: "" };
|
||||
const raw = await response.text();
|
||||
@@ -170,29 +194,45 @@ export async function readSSE(
|
||||
continue; /* 忽略心跳与非 JSON 块 */
|
||||
}
|
||||
|
||||
if (event.type === "text") {
|
||||
summary.text += event.content ?? "";
|
||||
} else if (event.type === "meta") {
|
||||
summary.intent = event.intent;
|
||||
summary.draftId = event.draft_id;
|
||||
summary.queryId = event.query_id;
|
||||
summary.traceId = event.trace_id;
|
||||
} else if (event.type === "done") {
|
||||
if (event.draft_id) summary.draftId = event.draft_id;
|
||||
if (event.query_id) summary.queryId = event.query_id;
|
||||
} else if (event.type === "error") {
|
||||
summary.error = event.message ?? "助手返回错误,请稍后重试";
|
||||
const record = event as Record<string, unknown>;
|
||||
const type = typeof record.type === "string" ? record.type : undefined;
|
||||
|
||||
if (type === "text") {
|
||||
summary.text += (record.content as string) ?? "";
|
||||
} else if (type === "meta") {
|
||||
summary.intent = record.intent as string | undefined;
|
||||
summary.draftId = record.draft_id as string | undefined;
|
||||
summary.queryId = record.query_id as string | undefined;
|
||||
summary.traceId = record.trace_id as string | undefined;
|
||||
} else if (type === "done") {
|
||||
if (record.draft_id) summary.draftId = record.draft_id as string;
|
||||
if (record.query_id) summary.queryId = record.query_id as string;
|
||||
} else if (type === "error") {
|
||||
summary.error = (record.message as string) ?? "助手返回错误,请稍后重试";
|
||||
} else if (!type) {
|
||||
/*
|
||||
* 无 type 的响应——客服 Agent / 客户 Agent 不用事件包装,直接 dump 业务对象:
|
||||
* service/customer_agent/chat.py → {answer, sources, intent, rewritten_query, trace_id}
|
||||
* 另兼容 success() 包装的 {code, message, data}(取值时自动下钻一层 data)。
|
||||
* 旧实现只认 type=text,导致这两个 Agent 的回复恒为空串(聊天框永远回「暂时没有找到合适的回答」)。
|
||||
*/
|
||||
const node =
|
||||
record.data && typeof record.data === "object"
|
||||
? (record.data as Record<string, unknown>)
|
||||
: record;
|
||||
const answer = node.answer ?? node.content ?? node.summary;
|
||||
if (typeof answer === "string" && answer) {
|
||||
summary.text += answer;
|
||||
} else if (typeof record.message === "string" && record.message) {
|
||||
summary.error = record.message;
|
||||
}
|
||||
if (typeof node.intent === "string") summary.intent = node.intent;
|
||||
if (typeof node.query_id === "string") summary.queryId = node.query_id;
|
||||
const trace = node.trace_id ?? record.trace_id;
|
||||
if (typeof trace === "string") summary.traceId = trace;
|
||||
if (Array.isArray(node.sources)) summary.sources = node.sources;
|
||||
}
|
||||
}
|
||||
|
||||
return summary;
|
||||
}
|
||||
|
||||
/**
|
||||
* 取 SSE 响应的正文文本(客服/客户 Agent 聊天用)。
|
||||
* 旧实现只取第一个 data 行,命中 meta 事件时会拿到空串,这里改为聚合全部 text 事件。
|
||||
*/
|
||||
export async function apiStream(path: string, body: unknown): Promise<string> {
|
||||
const summary = await readSSE(path, body);
|
||||
return summary.error || summary.text;
|
||||
}
|
||||
|
||||
@@ -72,10 +72,7 @@ export const assistantApi = {
|
||||
sessionHistory: (sessionId: string) =>
|
||||
apiFetch<AgentMessage[]>(`/advisor-agent/session/${encodeURIComponent(sessionId)}/history`),
|
||||
|
||||
/**
|
||||
* POST /chat/stream —— 响应是 text/event-stream,但后端整包输出,
|
||||
* 因此读完整条响应后一次性返回聚合结果(非流式渲染)。
|
||||
*/
|
||||
/** POST /chat/stream —— 读取 SSE,并支持调用方实时消费事件。 */
|
||||
chatOnce(
|
||||
body: {
|
||||
query: string;
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
import { apiFetch, readSSE, type SSESummary } from "@/lib/api";
|
||||
|
||||
/**
|
||||
* 客户端两个 Agent 的会话适配层。
|
||||
*
|
||||
* 两者都是「打开建会话 → 带 session_id 对话 → 关闭结束会话」的三段式,
|
||||
* 后端的会话归属校验会拒绝未创建或不属于自己的 session_id(SessionOwnershipError),
|
||||
* 所以 session_id 必须来自 session/create,不能在前端自己造。
|
||||
*
|
||||
* - customer:匿名客服 Agent(/api/agent/customer/*),游客与登录用户都可用,无需 token
|
||||
* - client :登录客户 Agent(/api/agent/client/*),需 require_customer,
|
||||
* 能读到本人持仓 / 交易(匿名版没有 customer_id,数据查询会被引导去登录)
|
||||
*/
|
||||
export type ChatAgentMode = "customer" | "client";
|
||||
|
||||
const PREFIX: Record<ChatAgentMode, string> = {
|
||||
customer: "/agent/customer",
|
||||
client: "/agent/client",
|
||||
};
|
||||
|
||||
export interface AgentSession {
|
||||
session_id: string;
|
||||
customer_id: number | null;
|
||||
}
|
||||
|
||||
export const chatAgentApi = {
|
||||
/** POST /session/create —— 建立会话,返回后端签发的 session_id。 */
|
||||
createSession: (mode: ChatAgentMode) =>
|
||||
apiFetch<AgentSession>(`${PREFIX[mode]}/session/create`, { method: "POST" }),
|
||||
|
||||
/** POST /chat —— 响应是 SSE 包装,但内容为 {answer, sources, intent...}。 */
|
||||
chat: (
|
||||
mode: ChatAgentMode,
|
||||
body: { session_id: string; query: string },
|
||||
signal?: AbortSignal
|
||||
): Promise<SSESummary> => readSSE(`${PREFIX[mode]}/chat`, body, signal),
|
||||
|
||||
/** POST /session/end —— 关闭会话并清理服务端上下文与限流键。 */
|
||||
endSession: (mode: ChatAgentMode, sessionId: string) =>
|
||||
apiFetch<{ session_id: string; archived?: boolean }>(`${PREFIX[mode]}/session/end`, {
|
||||
method: "POST",
|
||||
body: JSON.stringify({ session_id: sessionId }),
|
||||
}),
|
||||
};
|
||||
Vendored
+2
-2
@@ -1,7 +1,7 @@
|
||||
/// <reference types="next" />
|
||||
/// <reference types="next/image-types/global" />
|
||||
import "./.next/dev/types/routes.d.ts";
|
||||
import "./.next/dev/types/root-params.d.ts";
|
||||
import "./.next/types/routes.d.ts";
|
||||
import "./.next/types/root-params.d.ts";
|
||||
|
||||
// NOTE: This file should not be edited
|
||||
// see https://nextjs.org/docs/app/api-reference/config/typescript for more information.
|
||||
|
||||
@@ -90,6 +90,7 @@ async def lifespan(app: FastAPI):
|
||||
await app.state.advisor_event_consumer.stop()
|
||||
if app.state.advisor_scheduler is not None:
|
||||
app.state.advisor_scheduler.shutdown()
|
||||
await llm_client.aclose()
|
||||
await database.dispose()
|
||||
|
||||
|
||||
|
||||
@@ -49,6 +49,32 @@ def build_cache_key(
|
||||
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:
|
||||
"""读取缓存,Redis 异常或内容损坏时返回空结果。"""
|
||||
try:
|
||||
|
||||
@@ -16,6 +16,29 @@ _SYSTEM_PROMPT = """你是基金平台 NL2SQL 的问题改写器。
|
||||
不能扩大权限范围,不能臆造查询结果;上下文不足时保留原问题。
|
||||
"""
|
||||
|
||||
_REFERENCE_TERMS = (
|
||||
"他", "她", "它", "他们", "她们", "它们", "这个客户", "该客户", "那个客户",
|
||||
"这个基金", "该基金", "那只基金", "这只基金", "这个结果", "上一轮", "刚才", "前面",
|
||||
)
|
||||
|
||||
_CUSTOMER_RISK_LABELS = {
|
||||
"保守型": "C1",
|
||||
"稳健型": "C2",
|
||||
"平衡型": "C3",
|
||||
"进取型": "C4",
|
||||
"激进型": "C5",
|
||||
}
|
||||
|
||||
|
||||
def _normalize_customer_risk_label(question: str) -> str:
|
||||
"""将展示层客户风险标签规范化为数据库 C1-C5 编码。"""
|
||||
normalized = re.sub(r"\s+", "", question)
|
||||
if "客户" not in normalized and "风险" not in normalized:
|
||||
return normalized
|
||||
for label, code in _CUSTOMER_RISK_LABELS.items():
|
||||
normalized = normalized.replace(label, code)
|
||||
return normalized
|
||||
|
||||
|
||||
def _rewrite_explicit_customer_identity(question: str) -> str | None:
|
||||
normalized = re.sub(r"\s+", "", question)
|
||||
@@ -58,7 +81,7 @@ async def rewrite_query(
|
||||
llm_client=llm,
|
||||
) -> str:
|
||||
"""使用有限会话上下文补全问题;改写失败时安全返回原问题。"""
|
||||
original = (question or "").strip()
|
||||
original = _normalize_customer_risk_label((question or "").strip())
|
||||
context = (conversation_context or "").strip()
|
||||
if not original:
|
||||
return original
|
||||
@@ -68,7 +91,7 @@ async def rewrite_query(
|
||||
contextual_identity = _rewrite_customer_name_follow_up(original, context)
|
||||
if 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
|
||||
|
||||
prompt = (
|
||||
|
||||
@@ -67,6 +67,7 @@ async def render_answer(
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
temperature=0,
|
||||
max_tokens=256,
|
||||
)
|
||||
except Exception: # noqa: BLE001 摘要失败回退固定文本
|
||||
return fallback
|
||||
|
||||
+10
-3
@@ -1,6 +1,7 @@
|
||||
"""NL2SQL 元数据召回与授权过滤。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from typing import Any
|
||||
|
||||
@@ -70,15 +71,21 @@ async def retrieve_metadata(
|
||||
"is_valid",
|
||||
"is_deprecated",
|
||||
]
|
||||
results: list[dict[str, Any]] = []
|
||||
for chunk_type in ("table_meta", "field_meta"):
|
||||
hits = await milvus_client.search(
|
||||
async def search_chunk(chunk_type: str):
|
||||
return 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,
|
||||
)
|
||||
|
||||
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):
|
||||
entity = hit.get("entity") or hit
|
||||
if not entity.get("is_valid", True) or entity.get("is_deprecated", False):
|
||||
|
||||
+31
-2
@@ -1,12 +1,23 @@
|
||||
"""从 MySQL information_schema 加载并校验 NL2SQL 权威 Schema。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import bindparam, text
|
||||
|
||||
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(
|
||||
"""
|
||||
@@ -34,13 +45,25 @@ async def load_authoritative_schema(
|
||||
*,
|
||||
database: str,
|
||||
candidate_tables: set[str] | list[str],
|
||||
redis=None,
|
||||
cache_ttl: int = SCHEMA_CACHE_TTL,
|
||||
) -> 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)}
|
||||
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)
|
||||
column_result = await session.execute(COLUMNS_SQL, params)
|
||||
|
||||
@@ -56,4 +79,10 @@ async def load_authoritative_schema(
|
||||
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}
|
||||
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
|
||||
|
||||
@@ -18,12 +18,12 @@
|
||||
"value_hint": "fin_product.product_type 实际枚举值为:货币型/债券型/混合型/股票型/指数型/QDII;股票型基金必须用 product_type = '股票型',不要使用 Schema 注释里的'股票基金'"
|
||||
},
|
||||
{
|
||||
"term": "客户",
|
||||
"term": "客户风险等级",
|
||||
"enabled": true,
|
||||
"aliases": ["客户", "投资人", "持有人"],
|
||||
"aliases": ["客户", "投资人", "持有人", "风险评级", "风险等级", "风险类型", "激进型", "进取型", "平衡型", "稳健型", "保守型", "C1", "C2", "C3", "C4", "C5"],
|
||||
"tables": ["fin_customer_profile", "fin_holdings", "fin_risk_assessment"],
|
||||
"fields": ["customer_id", "risk_level", "customer_level"],
|
||||
"value_hint": "fin_customer_profile.risk_level / fin_risk_assessment.risk_level 实际枚举值为 R1-R5(R1 最保守,R5 最激进,与产品风险等级同一口径);查'保守型/稳健型客户'对应 R1/R2,'激进型客户'对应 R5"
|
||||
"value_hint": "客户风险字段实际枚举值为 C1-C5(C1 最保守,C5 最激进);展示层映射为保守型=C1、稳健型=C2、平衡型=C3、进取型=C4、激进型=C5。产品风险等级才使用 R1-R5"
|
||||
},
|
||||
{
|
||||
"term": "持仓",
|
||||
|
||||
+1
-1
@@ -74,7 +74,7 @@ async def generate_sql(
|
||||
},
|
||||
]
|
||||
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 统一收敛模型异常
|
||||
raise SqlGenerationError("SQL 模型调用失败") from exc
|
||||
sql = _clean_model_output(output)
|
||||
|
||||
@@ -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):
|
||||
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(
|
||||
|
||||
+36
-21
@@ -103,6 +103,21 @@ class LLMClient:
|
||||
def __init__(self, cfg: LLMCfg | None = None):
|
||||
self.cfg = cfg or settings.llm
|
||||
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 使用) -----------------------------
|
||||
@property
|
||||
@@ -185,10 +200,14 @@ class LLMClient:
|
||||
"max_tokens": request_max_tokens,
|
||||
"stream": False,
|
||||
}
|
||||
# SQL/分类等短输出任务不需要深度思考;DeepSeek-V4 若开启思考,
|
||||
# 可能耗尽 token 预算而返回空的 content。
|
||||
if model.lower().startswith("deepseek-v4"):
|
||||
payload["enable_thinking"] = False
|
||||
try:
|
||||
async with backend.client(self.cfg.timeout) as client:
|
||||
r = await client.post(url, headers=backend.headers, json=payload)
|
||||
r.raise_for_status()
|
||||
client = self._client_for(backend)
|
||||
r = await client.post(url, headers=backend.headers, json=payload)
|
||||
r.raise_for_status()
|
||||
choice = r.json()["choices"][0]
|
||||
message = choice["message"]
|
||||
content = message.get("content")
|
||||
@@ -250,20 +269,16 @@ class LLMClient:
|
||||
response = None
|
||||
max_retries = max(1, int(getattr(self.cfg, "max_retries", 1)))
|
||||
retry_backoff_sec = float(getattr(self.cfg, "retry_backoff_sec", 0))
|
||||
async with backend.client(self.cfg.timeout) as client:
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
response = await client.post(
|
||||
url,
|
||||
headers=backend.headers,
|
||||
json=payload,
|
||||
)
|
||||
response.raise_for_status()
|
||||
break
|
||||
except Exception:
|
||||
if attempt == max_retries - 1:
|
||||
raise
|
||||
await asyncio.sleep(retry_backoff_sec * (2**attempt))
|
||||
client = self._client_for(backend)
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
response = await client.post(url, headers=backend.headers, json=payload)
|
||||
response.raise_for_status()
|
||||
break
|
||||
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
|
||||
raise LLMFailError("Embedding 请求未返回响应")
|
||||
data = response.json()
|
||||
@@ -286,10 +301,10 @@ class LLMClient:
|
||||
url = backend.base_url.removesuffix("/v1") + "/api/tags"
|
||||
else:
|
||||
url = f"{backend.base_url}/models"
|
||||
async with backend.client(min(self.cfg.timeout, 10)) as client:
|
||||
r = await client.get(url, headers=backend.headers)
|
||||
if r.status_code >= 500:
|
||||
r.raise_for_status()
|
||||
client = self._client_for(backend)
|
||||
r = await client.get(url, headers=backend.headers)
|
||||
if r.status_code >= 500:
|
||||
r.raise_for_status()
|
||||
|
||||
|
||||
# 全局单例:Agent 统一引入
|
||||
|
||||
Reference in New Issue
Block a user