- app/tool/core_tools.py: 三只读 Tool(query_customer_profile/holdings/recent_trades, Core RO 仅 SELECT) + TOOL_REGISTRY 白名单(requires_customer); JSON 安全化(Decimal 两位/时间 isoformat) - app/service/tool_service.py: match_intent 关键词意图(仅 customer/advisor; risk/analyst 空转归 C1/C2) + assert_tool_access 归属断言(customer 本人/advisor assigned/risk_officer 全量/其余拒, 口径对齐 deps.assert_customer_access) + run_tool 编排(白名单→校验→执行→agent_tool_call 落库, 落库失败降级 warning) - agent_service: 图 START→tool→llm→guard; Tool 结果注入 LLM 上下文(SystemMessage); 降级回复带 Tool 摘要; chat() 可选 session 上下文(缺省空转, T-07 兼容) - session_repository: insert_tool_call(message_id 一期 NULL, session_id+trace_id 可还原) - 归属拒绝口径: Tool 层不抛 403 改 blocked 留痕(AUTH_403_* 同码), 对话内呈现——API 层 deny 铁律不变 - tests: test_chat_tools 20 例(三态+落库字段+降级+图注入) + test_chat 4 例(全链路/JWT 外 debug 通道), _ddl 补 agent_tool_call/core_holding, 297 绿 - 真库冒烟: uvicorn+真 MySQL/Redis chat 触发持仓查询 success 落痕 47ms, 现场已清理
176 lines
6.4 KiB
Python
176 lines
6.4 KiB
Python
"""对话接口(T-06 最小闭环 · FLOW §2):四 Agent 统一 chat 入口。
|
||
|
||
链路:X-Agent-Type 分流 + Agent 准入(deps.assert_agent_access,手册 §5.4)
|
||
→ customer 归属固定本人 / 代理人等指定客户走 G-01 归属校验(A-01 语义)
|
||
→ SessionGuard(会话存在、actor/agent_type 一致 AUTH_403_SESSION_AGENT、
|
||
active 状态,手册 §9)→ memory_service 窗口 → agent_service(T-07 图 +
|
||
T-04 Tool 节点:意图→归属校验→Core RO 只读查询)→ user/assistant 双消息
|
||
落 MySQL + Redis 窗口 → 响应 {session_id, reply, has_disclaimer, trace_id}。
|
||
落库:agent_session/agent_message(同 trace_id);agent_tool_call 由 Tool
|
||
节点落(T-04,success/blocked/error 全留痕)。审计:鉴权失败/越权经
|
||
deps.deny 双写留痕。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from uuid import uuid4
|
||
|
||
from fastapi import APIRouter, Depends, Request
|
||
from pydantic import BaseModel, Field
|
||
|
||
from app.api.deps import (
|
||
AGENT_TYPES,
|
||
AuthContext,
|
||
assert_agent_access,
|
||
assert_customer_access,
|
||
deny,
|
||
get_auth_context,
|
||
)
|
||
from app.repository.core_ro import CoreReadOnlyRepository
|
||
from app.repository.risk_repository import RiskRepository
|
||
from app.repository.session_repository import SessionRepository
|
||
from app.service import agent_service, memory_service
|
||
from app.utils.exceptions import ApiError, StateConflict
|
||
from app.utils.trace import current_trace
|
||
|
||
router = APIRouter(prefix="/api/chat", tags=["chat"])
|
||
|
||
MESSAGE_MAX_LENGTH = 4000
|
||
|
||
|
||
def _repo() -> RiskRepository:
|
||
"""审计仓储(deny 留痕用;测试 monkeypatch 点)。"""
|
||
return RiskRepository()
|
||
|
||
|
||
def _session_repo() -> SessionRepository:
|
||
"""会话仓储(测试 monkeypatch 点)。"""
|
||
return SessionRepository()
|
||
|
||
|
||
def _core_ro() -> CoreReadOnlyRepository:
|
||
"""归属校验仓储(测试 monkeypatch 点)。"""
|
||
return CoreReadOnlyRepository()
|
||
|
||
|
||
class ChatRequest(BaseModel):
|
||
session_id: str | None = Field(None, description="缺省新建会话;传入则续聊")
|
||
message: str = Field(..., min_length=1, max_length=MESSAGE_MAX_LENGTH)
|
||
customer_id: str | None = Field(
|
||
None, description="目标客户:customer 角色忽略(强制本人);advisor/risk/analyst 可指定(过归属校验)"
|
||
)
|
||
title: str | None = Field(None, max_length=256)
|
||
|
||
|
||
def _primary_role(auth: AuthContext) -> str:
|
||
"""会话主角色(agent_session.actor_role):按 Agent 边界优先序取(评审 P3-11,
|
||
多角色 token 落库稳定)。"""
|
||
for role in ("customer", "advisor", "analyst", "risk_officer", "compliance", "ops"):
|
||
if role in auth.roles:
|
||
return role
|
||
return auth.roles[0] if auth.roles else "unknown"
|
||
|
||
|
||
def _resolve_customer_id(
|
||
auth: AuthContext, agent_type: str, requested: str | None
|
||
) -> str | None:
|
||
"""会话关联客户:customer 强制本人;其余角色指定时过 G-01 归属校验。"""
|
||
if agent_type == "customer":
|
||
if requested and requested != auth.customer_id:
|
||
deny(auth, "AUTH_403_NOT_OWNER", _repo(), customer_id=requested, agent_type=agent_type)
|
||
return auth.customer_id
|
||
if requested:
|
||
assert_customer_access(auth, requested, core_ro=_core_ro(), risk_repo=_repo())
|
||
return requested
|
||
return None
|
||
|
||
|
||
@router.post("")
|
||
def chat_api(req: ChatRequest, request: Request, auth: AuthContext = Depends(get_auth_context)) -> dict:
|
||
agent_type = request.headers.get("X-Agent-Type", "").strip()
|
||
if not agent_type:
|
||
raise ApiError(401, "AUTH_401_MISSING_AGENT_TYPE", "missing X-Agent-Type header")
|
||
if agent_type not in AGENT_TYPES:
|
||
raise ApiError(400, "BAD_REQUEST", f"invalid X-Agent-Type: {agent_type}")
|
||
assert_agent_access(auth, agent_type, risk_repo=_repo())
|
||
|
||
message = req.message.strip()
|
||
if not message:
|
||
raise ApiError(400, "BAD_REQUEST", "message is blank")
|
||
|
||
customer_id = _resolve_customer_id(auth, agent_type, req.customer_id)
|
||
session_repo = _session_repo()
|
||
|
||
if req.session_id:
|
||
session = session_repo.get_session(req.session_id)
|
||
if session is None:
|
||
raise ApiError(404, "NOT_FOUND", f"session not found: {req.session_id}")
|
||
# SessionGuard(手册 §9):actor/agent_type 一致;他人会话 fail-closed
|
||
if session["actor_id"] != auth.actor_id or session["agent_type"] != agent_type:
|
||
deny(
|
||
auth,
|
||
"AUTH_403_SESSION_AGENT",
|
||
_repo(),
|
||
customer_id=session.get("customer_id"),
|
||
message="session belongs to another actor or agent",
|
||
agent_type=agent_type,
|
||
)
|
||
if session["status"] != "active":
|
||
raise ApiError(409, "STATE_CONFLICT", f"session is {session['status']}")
|
||
sid = session["session_id"]
|
||
else:
|
||
sid = f"sess-{uuid4().hex[:16]}"
|
||
session_repo.create_session(
|
||
session_id=sid,
|
||
trace_id=current_trace(),
|
||
agent_type=agent_type,
|
||
actor_id=auth.actor_id,
|
||
actor_role=_primary_role(auth),
|
||
customer_id=customer_id,
|
||
advisor_id=auth.actor_id if agent_type == "advisor" else None,
|
||
title=req.title,
|
||
)
|
||
|
||
history = memory_service.get_recent(agent_type, sid)
|
||
result = agent_service.chat(
|
||
agent_type,
|
||
history,
|
||
message,
|
||
session_id=sid,
|
||
trace_id=current_trace(),
|
||
actor={"actor_id": auth.actor_id, "roles": auth.roles, "token_type": auth.token_type},
|
||
customer_id=customer_id,
|
||
)
|
||
|
||
# 落盘:user + assistant 同步写(异步化归后续);同 trace_id 贯通
|
||
trace_id = current_trace()
|
||
seq = session_repo.next_seq_no(sid)
|
||
session_repo.insert_message(
|
||
session_id=sid, trace_id=trace_id, seq_no=seq, role="user", content=message
|
||
)
|
||
session_repo.insert_message(
|
||
session_id=sid,
|
||
trace_id=trace_id,
|
||
seq_no=seq + 1,
|
||
role="assistant",
|
||
content=result["reply"],
|
||
has_disclaimer=bool(result["has_disclaimer"]),
|
||
)
|
||
memory_service.append_window(
|
||
agent_type,
|
||
sid,
|
||
[
|
||
{"role": "user", "content": message},
|
||
{"role": "assistant", "content": result["reply"]},
|
||
],
|
||
)
|
||
|
||
return {
|
||
"session_id": sid,
|
||
"agent_type": agent_type,
|
||
"customer_id": customer_id,
|
||
"reply": result["reply"],
|
||
"has_disclaimer": result["has_disclaimer"],
|
||
"trace_id": trace_id,
|
||
}
|