Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
126 changes: 126 additions & 0 deletions app/agent/cost.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
"""模型调用成本与延迟日志(EVA-03)。

Agent 层不查库,这里用进程内列表 + 结构化日志记录每次模型调用的
路由、模型名、模式、耗时、token、错误与降级状态。后端可在请求结束后
从 get_call_logs() 读取并落库到 AIAnalysis 表。

session_id 通过 contextvar 注入:run_copilot / extract_promises 在入口设置,
所有嵌套调用自动带上当前会话,用于按会话统计成本。
"""

import contextvars
import json
import logging
import threading
from dataclasses import asdict, dataclass
from datetime import UTC, datetime
from typing import Any

from langchain_core.callbacks import BaseCallbackHandler

logger = logging.getLogger("app.agent.cost")

_current_session: contextvars.ContextVar[str | None] = contextvars.ContextVar("cost_session_id", default=None)


@dataclass
class ModelCallRecord:
session_id: str | None
route: str
model: str
mode: str
latency_ms: int
input_tokens: int
output_tokens: int
total_tokens: int
error: str | None
degraded: bool
started_at: str


class TokenCapture(BaseCallbackHandler):
"""在 LLM 调用结束时捕获 token 用量,供成本日志使用。"""

def __init__(self) -> None:
self.input_tokens = 0
self.output_tokens = 0
self.total_tokens = 0

def on_llm_end(self, response: Any, **kwargs: Any) -> None:
try:
first = response.generations[0][0]
message = getattr(first, "message", None)
if message is None:
return
usage = getattr(message, "usage_metadata", None)
if not usage:
usage = (getattr(message, "response_metadata", None) or {}).get("token_usage")
if not usage:
return
self.input_tokens = int(usage.get("input_tokens") or 0)
self.output_tokens = int(usage.get("output_tokens") or 0)
self.total_tokens = int(usage.get("total_tokens") or (self.input_tokens + self.output_tokens))
except Exception:
# token 统计失败不阻断主流程,成本日志只是观测
return


_records: list[ModelCallRecord] = []
_lock = threading.Lock()


def set_session(session_id: str | None) -> None:
_current_session.set(session_id)


def record_call(
*,
route: str,
model: str,
mode: str,
latency_ms: int,
input_tokens: int = 0,
output_tokens: int = 0,
total_tokens: int = 0,
error: str | None = None,
degraded: bool = False,
) -> None:
"""记录一次模型调用(含 Mock / 降级调用,token 为 0)。"""
rec = ModelCallRecord(
session_id=_current_session.get(),
route=route,
model=model,
mode=mode,
latency_ms=latency_ms,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
error=error,
degraded=degraded,
started_at=datetime.now(UTC).isoformat(),
)
with _lock:
_records.append(rec)
logger.info("model_call %s", json.dumps(asdict(rec), ensure_ascii=False))


def get_call_logs() -> list[ModelCallRecord]:
with _lock:
return list(_records)


def clear_call_logs() -> None:
with _lock:
_records.clear()


def session_totals(session_id: str) -> dict[str, int]:
"""按会话聚合:总调用数、总 token、总耗时。"""
calls = [r for r in get_call_logs() if r.session_id == session_id]
return {
"calls": len(calls),
"input_tokens": sum(r.input_tokens for r in calls),
"output_tokens": sum(r.output_tokens for r in calls),
"total_tokens": sum(r.total_tokens for r in calls),
"latency_ms": sum(r.latency_ms for r in calls),
}
12 changes: 10 additions & 2 deletions app/agent/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@
from app.agent.nodes.reply import reply_node
from app.agent.nodes.risk_engine import risk_engine_node
from app.agent.nodes.risk_extraction import risk_extraction_node
from app.agent.nodes.summary import summary_node
from app.agent.nodes.tool_query import tool_query_node
from app.agent.nodes.vision import vision_node
from app.agent.state import CustomerState


Expand All @@ -27,6 +29,8 @@ def route_by_risk(state: CustomerState) -> str:
builder.add_node("risk_engine", risk_engine_node)
builder.add_node("tool_query", tool_query_node)
builder.add_node("evidence", evidence_node)
builder.add_node("summary", summary_node)
builder.add_node("vision", vision_node)
builder.add_node("adverse", adverse_node)
builder.add_node("reply", reply_node)
builder.add_node("fact_check", fact_check_node)
Expand All @@ -39,11 +43,15 @@ def route_by_risk(state: CustomerState) -> str:
# 事实与证据在分支之前完成,保证两条路都拿得到证据
builder.add_edge("risk_engine", "tool_query")
builder.add_edge("tool_query", "evidence")
# 轨迹摘要在分支前完成,普通与高风险两条路都能拿到历史轨迹
builder.add_edge("evidence", "summary")
builder.add_conditional_edges(
"evidence",
"summary",
route_by_risk,
{"adverse": "adverse", "normal": "reply"},
{"adverse": "vision", "normal": "reply"},
)
# 高风险分支:先识别图片,再走不良反应专项处置
builder.add_edge("vision", "adverse")
builder.add_edge("adverse", "reply")
builder.add_edge("reply", "fact_check")
builder.add_edge("fact_check", END)
Expand Down
25 changes: 23 additions & 2 deletions app/agent/llm.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,12 @@
import logging
import time
from collections.abc import Callable
from typing import TypeVar

from langchain_core.output_parsers import PydanticOutputParser
from pydantic import BaseModel

from app.agent import cost
from app.integrations.model.gateway import ModelUnavailableError, gateway

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -40,21 +42,40 @@ def run_structured[T: BaseModel](
auto —— 有配置就走模型,无配置或调用失败都降级到 fallback
real —— 强制走模型,缺配置或失败直接抛异常(仅开发/评测用)
"""
model = gateway.model_name_for(route)

if mode == "mock":
cost.record_call(route=route, model="", mode=mode, latency_ms=0, degraded=True)
return fallback(), True

if not gateway.is_available(route):
if mode == "real":
raise ModelUnavailableError(f"mode=real 但路由 {route} 未配置模型,无法执行")
logger.warning("模型路由 %s 不可用,降级到 Mock 结果", route)
cost.record_call(route=route, model=model, mode=mode, latency_ms=0, error="route unavailable", degraded=True)
return fallback(), True

parser = parser_for(schema)
chain = gateway.get(route).bind(response_format={"type": "json_object"}) | parser
token_capture = cost.TokenCapture()
started = time.monotonic()
try:
return chain.invoke(prompt), False
except Exception:
result = chain.invoke(prompt, config={"callbacks": [token_capture]})
latency_ms = int((time.monotonic() - started) * 1000)
cost.record_call(
route=route,
model=model,
mode=mode,
latency_ms=latency_ms,
input_tokens=token_capture.input_tokens,
output_tokens=token_capture.output_tokens,
total_tokens=token_capture.total_tokens,
)
return result, False
except Exception as exc:
latency_ms = int((time.monotonic() - started) * 1000)
if mode == "real":
raise
logger.exception("模型调用失败,降级到 Mock 结果(route=%s)", route)
cost.record_call(route=route, model=model, mode=mode, latency_ms=latency_ms, error=str(exc), degraded=True)
return fallback(), True
72 changes: 70 additions & 2 deletions app/agent/mock.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,8 @@
PromiseExtraction,
ReplyDraft,
RiskExtraction,
TrajectorySummary,
VisionExtraction,
)

# 不良反应症状词
Expand Down Expand Up @@ -55,12 +57,27 @@
_ORDER_NO_RE = re.compile(r"\d{6,}")


# 否定词:紧邻关键词之前的否定会取消命中(「没有红肿」「不痒」不算不良反应)
_NEGATION_WORDS = ("没有", "没", "不是", "并非", "不", "无", "未", "否认", "排除")


def _is_negated(text: str, pos: int) -> bool:
"""判断 text 中 pos 处的关键词是否被其前面紧邻的否定词修饰。"""
window = text[max(0, pos - 4) : pos]
return any(neg in window for neg in _NEGATION_WORDS)


def _word_hit(text: str, word: str) -> bool:
"""关键词命中且至少有一处未被否定。"""
return any(not _is_negated(text, m.start()) for m in re.finditer(re.escape(word), text))


def _hit(text: str, words: tuple[str, ...]) -> bool:
return any(word in text for word in words)
return any(_word_hit(text, word) for word in words)


def _hits(text: str, words: tuple[str, ...]) -> list[str]:
return [word for word in words if word in text]
return [word for word in words if _word_hit(text, word)]


def mock_intent(message: str) -> IntentAnalysis:
Expand Down Expand Up @@ -172,3 +189,54 @@ def mock_promise(message_text: str) -> PromiseExtraction:
owner_type="agent",
confidence=0.6 if has_promise else 0.0,
)


def mock_vision(images: list[dict]) -> VisionExtraction:
"""图片识别 Mock。

有真实图片数据(url / b64 / path)时返回确定性的「清晰 / 患处照片」;
只有占位(如 [图片])时返回「待识别」,交人工确认 —— 绝不编造图中内容。
"""
if not images:
return VisionExtraction()

first = images[0]
if any(first.get(key) for key in ("image_url", "image_b64", "image_path")):
return VisionExtraction(
image_type="患处照片",
clarity="清晰",
batch_no=first.get("batch_no_masked") or first.get("batch_no") or "",
visible_symptoms=["局部泛红"],
extra_fields={},
)
return VisionExtraction()


def mock_summary(
events: list[dict],
messages: list[dict],
orders: list[dict],
tickets: list[dict],
promises: list[dict],
) -> TrajectorySummary:
"""轨迹摘要 Mock:把最近的历史事件标题串成一句,事件 ID 一并保留用于溯源。"""
lines: list[str] = []
ids: list[str] = []

for event in (events or [])[-5:]:
eid = event.get("event_id") or event.get("id") or ""
title = event.get("title") or event.get("event_type") or ""
if title:
lines.append(str(title))
if eid:
ids.append(str(eid))

if not lines:
if promises:
lines.append("存在历史服务承诺")
elif orders:
lines.append("存在关联订单")
else:
lines.append("暂无跨会话历史轨迹")

return TrajectorySummary(summary=";".join(lines), relevant_event_ids=ids)
10 changes: 7 additions & 3 deletions app/agent/nodes/adverse.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,8 +41,9 @@ def _mentions(state: CustomerState, words: tuple[str, ...]) -> bool:
return any(any(word in text for word in words) for text in _texts(state))


def _has_image(state: CustomerState) -> bool:
return any((message.get("content_type") == "image") for message in (state.get("messages") or []))
def _image_known(state: CustomerState) -> bool:
"""图片是否已可识别(清晰)。没有图片或图片待识别都算缺失。"""
return (state.get("vision") or {}).get("clarity") == "清晰"


def _has_batch_no(state: CustomerState) -> bool:
Expand All @@ -56,7 +57,7 @@ def _has_batch_no(state: CustomerState) -> bool:
("症状出现时间", lambda state: _mentions(state, _OCCURRENCE_WORDS)),
("是否停止使用", lambda state: _mentions(state, _STOPPED_WORDS)),
("是否就医", lambda state: bool(state["risk"].get("medical_visit"))),
("图片或门诊资料", _has_image),
("图片或门诊资料", _image_known),
)


Expand Down Expand Up @@ -143,12 +144,15 @@ def adverse_node(state: CustomerState) -> dict:
missing = _missing_fields(state)
actions = _build_actions(grade)
symptom_summary = _symptom_summary(state)
vision = state.get("vision") or {}

assessment = AdverseAssessment(
grade=grade,
symptom_summary=symptom_summary,
medical_visit=bool(state["risk"].get("medical_visit")),
stopped_use=None, # 未知就留空,由客服确认,不猜
image_clarity=vision.get("clarity"),
image_type=vision.get("image_type"),
missing_fields=missing,
ticket_draft=None if grade == "L1" else _build_ticket_draft(state, grade),
safe_reply=_safe_reply(grade, symptom_summary, product_name),
Expand Down
1 change: 1 addition & 0 deletions app/agent/nodes/reply.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ def reply_node(state: CustomerState) -> dict:
evidence=_evidence_lines(state) or "暂无系统可核验的事实",
suggested_actions=state.get("suggested_actions") or [],
safe_reply=safe_reply,
timeline_summary=state.get("timeline_summary") or "无历史轨迹",
format_instructions=format_instructions(ReplyDraft),
)

Expand Down
40 changes: 40 additions & 0 deletions app/agent/nodes/summary.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
"""跨会话轨迹摘要节点(AG-05)。

消费 CopilotRequest.events(后端按消费者预组装的历史事件)产出
timeline_summary,供右侧「全轨迹」与回复草稿引用历史承诺
/ 进度。
"""

from app.agent.llm import format_instructions, run_structured
from app.agent.mock import mock_summary
from app.agent.schemas.analysis import TrajectorySummary
from app.agent.state import CustomerState
from app.prompt.summary import TRAJECTORY_SUMMARY_PROMPT


def summary_node(state: CustomerState) -> dict:
prompt = TRAJECTORY_SUMMARY_PROMPT.format(
message=state.get("current_message", ""),
risk_level=state.get("risk_level", "L0"),
events=state.get("events") or [],
orders=state.get("orders") or [],
tickets=state.get("tickets") or [],
promises=state.get("promises") or [],
format_instructions=format_instructions(TrajectorySummary),
)

result, degraded = run_structured(
prompt=prompt,
schema=TrajectorySummary,
route="reasoning",
mode=state.get("mode", "auto"),
fallback=lambda: mock_summary(
state.get("events") or [],
state.get("messages") or [],
state.get("orders") or [],
state.get("tickets") or [],
state.get("promises") or [],
),
)

return {"timeline_summary": result.summary, "degraded": degraded}
Loading
Loading