Files
kefu/deploy/token-usage-20260917/model_gateway.py.diff
T
2026-09-21 10:34:06 +08:00

529 lines
24 KiB
Diff

--- server/model_gateway.py
+++ local/model_gateway.py
@@ -45,4 +45,5 @@
import model_protocol
import admin_backend
+from model_usage import UsageRecorder
from knowledge_store import KnowledgeStore
from knowledge_retriever import KnowledgeRetriever, inject_references
@@ -186,5 +187,5 @@
self.refresh()
- def plan(self) -> tuple[list[Outlet], Outlet | None, str, int]:
+ def plan(self, *, purpose: str = "chat") -> tuple[list[Outlet], Outlet | None, str, int]:
"""本轮问谁、谁当裁判、什么模式、编排版本号。"""
roles = self.roles or {}
@@ -202,4 +203,11 @@
if outlet is not None
]
+ if purpose == "guard":
+ # Internal screenshot classification uses the configured vision role.
+ # Older installations without that role keep their primary outlet;
+ # customer-answer fanout and reply scoring cannot judge UI state.
+ vision = pick(roles.get("vision_id") or "")
+ return ([vision] if vision is not None else answers[:1], None,
+ "shadow", int(roles.get("version") or 0))
judge = pick(roles.get("judge_id") or "")
mode = str(roles.get("judge_mode") or "shadow").lower()
@@ -210,7 +218,8 @@
class GatewayError(Exception):
- def __init__(self, message: str, *, retriable: bool = False):
+ def __init__(self, message: str, *, retriable: bool = False, response_data=None):
super().__init__(message)
self.retriable = retriable
+ self.response_data = response_data
@@ -229,16 +238,19 @@
except (httpx.TimeoutException, httpx.TransportError) as exc:
raise GatewayError(f"{type(exc).__name__}: {exc}", retriable=True) from exc
+ try:
+ data = response.json()
+ except ValueError:
+ data = None
if response.status_code in RETRIABLE_STATUS:
raise GatewayError(
- f"上游返回 {response.status_code}", retriable=True
+ f"上游返回 {response.status_code}", retriable=True, response_data=data
)
if response.status_code >= 400:
raise GatewayError(
- f"上游返回 {response.status_code}: {response.text[:200]}", retriable=False
+ f"上游返回 {response.status_code}: {response.text[:200]}", retriable=False, response_data=data
)
- try:
- return response.json()
- except ValueError as exc:
- raise GatewayError("上游响应不是合法 JSON", retriable=False) from exc
+ if data is None:
+ raise GatewayError("上游响应不是合法 JSON", retriable=False)
+ return data
@@ -253,4 +265,6 @@
tools: list | None = None,
deadline: float = SINGLE_CALL_TIMEOUT,
+ usage_recorder: UsageRecorder | None = None,
+ usage_role: str = "answer",
) -> dict:
"""向一个出口要一次回复。并发闸 + 熔断 + 有界重试都在这里。
@@ -284,6 +298,6 @@
payload_messages = list(messages)
dify_files = None
- if image_b64 and kind == "dify":
- dify_files = [await _dify_upload(client, outlet, image_b64, timeout)]
+ if kind == "dify":
+ dify_files = await _dify_message_files(client, outlet, messages, image_b64, timeout)
elif image_b64:
payload_messages = payload_messages + [
@@ -320,4 +334,7 @@
last: Exception | None = None
for attempt in range(1, MAX_ATTEMPTS + 1):
+ data = None
+ status = "error"
+ attempt_started = time.monotonic()
try:
data = await _post_once(
@@ -327,4 +344,5 @@
kind, data, config.get("base_url") or ""
)
+ status = "success"
outlet.breaker.record(True)
return {
@@ -336,11 +354,18 @@
}
except GatewayError as exc:
+ data = exc.response_data
last = exc
if not exc.retriable or attempt >= MAX_ATTEMPTS:
break
- await asyncio.sleep(0.35 * attempt + secrets.randbelow(200) / 1000.0)
except ValueError as exc: # 解析失败,确定性故障
last = exc
break
+ finally:
+ if usage_recorder is not None:
+ usage_recorder.capture(
+ outlet, data, attempt=attempt, role=usage_role, status=status,
+ latency_ms=int((time.monotonic() - attempt_started) * 1000),
+ )
+ await asyncio.sleep(0.35 * attempt + secrets.randbelow(200) / 1000.0)
outlet.breaker.record(False)
return {
@@ -359,10 +384,53 @@
}
+ except (GatewayError, httpx.HTTPError, ValueError) as exc:
+ # Attachment preparation is part of this outlet, not a failure of every
+ # concurrent model. Keep healthy candidates available to the caller.
+ outlet.breaker.record(False)
+ return {"provider": outlet.name, "text": "", "latency_ms": int((time.monotonic() - started) * 1000),
+ "error": str(exc)[:200] or type(exc).__name__}
+
+
+async def _dify_message_files(client, outlet, messages, image_b64, timeout):
+ """Translate current-turn attachments without leaking image bytes into query."""
+ urls = []
+ for message in reversed(messages or []):
+ if not isinstance(message, dict) or message.get("role") != "user":
+ continue
+ content = message.get("content")
+ if isinstance(content, list):
+ for part in content:
+ if isinstance(part, dict) and part.get("type") == "image_url":
+ image = part.get("image_url") or {}
+ url = image.get("url", "") if isinstance(image, dict) else image
+ if isinstance(url, str) and url:
+ urls.append(url)
+ break
+ if image_b64:
+ urls.append("data:image/png;base64," + image_b64)
+ files = []
+ seen = set()
+ for url in urls:
+ if url in seen:
+ continue
+ seen.add(url)
+ if url.startswith("data:"):
+ header, separator, data = url.partition(",")
+ mime = header[5:].split(";")[0].lower()
+ if not separator or ";base64" not in header or mime not in {"image/png", "image/jpeg", "image/webp", "image/gif"}:
+ raise GatewayError("Dify 图片数据格式不受支持", retriable=False)
+ files.append(await _dify_upload(client, outlet, data, timeout, mime=mime))
+ elif url.startswith(("https://", "http://")):
+ files.append({"type": "image", "transfer_method": "remote_url", "url": url})
+ else:
+ raise GatewayError("Dify 图片地址格式不受支持", retriable=False)
+ return files or None
+
async def _dify_upload(
- client: httpx.AsyncClient, outlet: Outlet, image_b64: str, timeout: float
+ client: httpx.AsyncClient, outlet: Outlet, image_b64: str, timeout: float, *, mime: str = "image/png"
) -> dict:
"""Dify 的图片要先 upload 拿 id 再在 chat-messages 里引用。"""
- root = model_protocol.dify_api_root(outlet.config.get("base_url") or "")
+ root = model_protocol.dify_api_root(outlet.config.get("base_url") or "", outlet.config.get("endpoint_mode") or "auto")
response = await client.post(
f"{root}/files/upload",
@@ -371,6 +439,6 @@
"Accept": "application/json",
},
- data={"user": "wechat-rpa-vision"},
- files={"file": ("chat.png", base64.b64decode(image_b64), "image/png")},
+ data={"user": "wechat-rpa"},
+ files={"file": ("chat." + {"image/png":"png", "image/jpeg":"jpg", "image/webp":"webp", "image/gif":"gif"}[mime], base64.b64decode(image_b64, validate=True), mime)},
timeout=timeout,
)
@@ -379,7 +447,14 @@
f"Dify 截图上传失败 {response.status_code}", retriable=response.status_code in RETRIABLE_STATUS
)
- upload_id = str((response.json() or {}).get("id") or "").strip()
- if not upload_id:
- raise GatewayError("Dify 上传成功但没返回文件 ID", retriable=False)
+ try:
+ uploaded = response.json()
+ except ValueError as exc:
+ raise GatewayError("Dify 上传响应不是合法 JSON", retriable=False) from exc
+ if not isinstance(uploaded, dict):
+ raise GatewayError("Dify 上传响应不是 JSON 对象", retriable=False)
+ upload_id = uploaded.get("id")
+ if not isinstance(upload_id, str) or not upload_id.strip():
+ raise GatewayError("Dify 上传成功但没返回有效文件 ID", retriable=False)
+ upload_id = upload_id.strip()
return {
"type": "image",
@@ -450,4 +525,5 @@
first: str,
second: str = "",
+ *, usage_recorder: UsageRecorder | None = None,
) -> dict:
if judge is None or not str(first or "").strip():
@@ -469,4 +545,5 @@
temperature=0.0,
deadline=JUDGE_DEADLINE,
+ usage_recorder=usage_recorder, usage_role="judge",
)
latency = int((time.monotonic() - started) * 1000)
@@ -607,5 +684,6 @@
body = await request.json()
- answers, judge, mode, version = catalog.plan()
+ purpose = str(body.get("purpose") or "chat")
+ answers, judge, mode, version = catalog.plan(purpose=purpose)
if not answers:
raise HTTPException(status_code=503, detail="没有可用的答题出口")
@@ -625,151 +703,162 @@
tools = body.get("tools") or None
- started = time.monotonic()
- knowledge_result = None
- knowledge_trace = {"enabled": False, "hits": [], "mode": "disabled", "elapsed_ms": 0}
- tenant = str(account["tenant_id"])
- if str(body.get("purpose") or "chat") == "chat":
- settings = await asyncio.to_thread(knowledge_store.settings, tenant)
- if settings["enabled"]:
- from knowledge_processing import latest_query
- query = str(body.get("customer_text") or "")
- if not query:
+ async with UsageRecorder(
+ database, tenant_id=str(account["tenant_id"]), desktop_account_id=int(account["id"]),
+ task_id=str(body.get("task_id") or ""), purpose=purpose,
+ ) as usage_recorder:
+ started = time.monotonic()
+ knowledge_result = None
+ knowledge_trace = {"enabled": False, "hits": [], "mode": "disabled", "elapsed_ms": 0}
+ tenant = str(account["tenant_id"])
+ if str(body.get("purpose") or "chat") == "chat":
+ settings = await asyncio.to_thread(knowledge_store.settings, tenant)
+ if settings["enabled"]:
+ from knowledge_processing import latest_query
+ query = ""
for turn in reversed(messages):
- if turn.get("role") == "user" and isinstance(turn.get("content"), str):
- query = turn["content"]
+ if isinstance(turn, dict) and turn.get("role") == "user":
+ query = model_protocol.message_text(turn.get("content"))
break
- query = latest_query(query) or query
+ if not query:
+ query = str(body.get("customer_text") or "")
+ query = latest_query(query) or query
+ try:
+ knowledge_result = await asyncio.wait_for(
+ asyncio.to_thread(knowledge_retriever.search, tenant, query), timeout=4.0)
+ knowledge_trace = {"enabled": True, "mode": knowledge_result["mode"],
+ "elapsed_ms": knowledge_result["elapsed_ms"], "hits": [
+ {k: h[k] for k in ("id", "revision", "score")} for h in knowledge_result["hits"]]}
+ messages = inject_references(messages, knowledge_result["hits"])
+ except Exception:
+ knowledge_trace = {"enabled": True, "hits": [], "mode": "unavailable", "elapsed_ms": 4000}
+ from knowledge_processing import redact
+ knowledge_result = {**knowledge_trace, "query": redact(query)[:1000]}
+ messages = inject_references(messages, [])
+
+ if tools:
+ # 工具轮:只问主出口,不并发也不评审。
+ #
+ # 并发问多个模型在这里是错的——它们会给出不同的 tool_calls,而工具
+ # 跑在桌面端本机、有副作用(挂号登记会真的建一条记录)。执行哪一份?
+ # 都执行就是重复下单,选一份就是把另一份的上下文丢了,接下来那一路
+ # 的对话直接对不上。
+ #
+ # 评审同理:一句"我要查客户资料"没什么好打分的。等工具跑完、模型给
+ # 出真正的答复(那一轮不带 tools),best-of-N 和裁判照常生效。
+ primary = answers[0]
+ single = await call_outlet(
+ request.app.state.client,
+ primary,
+ messages,
+ image_b64=image_b64,
+ tools=tools,
+ deadline=ANSWER_DEADLINE,
+ usage_recorder=usage_recorder,
+ )
+ payload = {
+ "reply": single.get("text") or "",
+ "tool_calls": single.get("tool_calls") or [],
+ "chosen": single.get("provider") or "",
+ "candidates": [single],
+ "judge": dict(_EMPTY_VERDICT),
+ "judge_mode": mode,
+ "roles_version": version,
+ "knowledge": knowledge_trace,
+ "total_ms": int((time.monotonic() - started) * 1000),
+ }
+ if single.get("error") and not payload["reply"] and not payload["tool_calls"]:
+ return JSONResponse(payload, status_code=502)
+ if idempotency_key:
+ idempotency[idempotency_key] = (now, payload)
try:
- knowledge_result = await asyncio.wait_for(
- asyncio.to_thread(knowledge_retriever.search, tenant, query), timeout=4.0)
- knowledge_trace = {"enabled": True, "mode": knowledge_result["mode"],
- "elapsed_ms": knowledge_result["elapsed_ms"], "hits": [
- {k: h[k] for k in ("id", "revision", "score")} for h in knowledge_result["hits"]]}
- messages = inject_references(messages, knowledge_result["hits"])
- except Exception:
- knowledge_trace = {"enabled": True, "hits": [], "mode": "unavailable", "elapsed_ms": 4000}
- from knowledge_processing import redact
- knowledge_result = {**knowledge_trace, "query": redact(query)[:1000]}
- messages = inject_references(messages, [])
-
- if tools:
- # 工具轮:只问主出口,不并发也不评审。
- #
- # 并发问多个模型在这里是错的——它们会给出不同的 tool_calls,而工具
- # 跑在桌面端本机、有副作用(挂号登记会真的建一条记录)。执行哪一份?
- # 都执行就是重复下单,选一份就是把另一份的上下文丢了,接下来那一路
- # 的对话直接对不上。
- #
- # 评审同理:一句"我要查客户资料"没什么好打分的。等工具跑完、模型给
- # 出真正的答复(那一轮不带 tools),best-of-N 和裁判照常生效。
- primary = answers[0]
- single = await call_outlet(
- request.app.state.client,
- primary,
- messages,
- image_b64=image_b64,
- tools=tools,
- deadline=ANSWER_DEADLINE,
+ log_queue.put_nowait({
+ "desktop_account_id": int(account["id"]), "tenant_id": tenant,
+ "device_id": x_device_id, "task_id": usage_recorder.context["task_id"],
+ "roles_version": version, "judge_mode": mode, "chosen": payload["chosen"],
+ "judge": payload["judge"], "candidates": [single], "total_ms": payload["total_ms"],
+ "customer_text": str(body.get("customer_text") or ""), "reply_text": payload["reply"],
+ "purpose": str(body.get("purpose") or "chat"), "knowledge_retrieval": knowledge_result,
+ })
+ except asyncio.QueueFull:
+ pass
+ return JSONResponse(payload)
+
+ results = await asyncio.gather(
+ *[
+ call_outlet(
+ request.app.state.client,
+ outlet,
+ messages,
+ image_b64=image_b64,
+ deadline=ANSWER_DEADLINE,
+ usage_recorder=usage_recorder,
+ )
+ for outlet in answers
+ ]
)
+ usable = [item for item in results if item.get("text")]
+ if not usable:
+ payload = {
+ "reply": "", "chosen": "", "candidates": results,
+ "judge": dict(_EMPTY_VERDICT), "judge_mode": mode,
+ "roles_version": version,
+ "total_ms": int((time.monotonic() - started) * 1000),
+ }
+ return JSONResponse(payload, status_code=502)
+
+ if purpose == "guard":
+ verdict = dict(_EMPTY_VERDICT)
+ else:
+ verdict = await run_judge(
+ request.app.state.client,
+ judge,
+ str(body.get("customer_text") or ""),
+ usable[0]["text"],
+ usable[1]["text"] if len(usable) > 1 else "",
+ usage_recorder=usage_recorder,
+ )
+
+ chosen = usable[0]
+ if (
+ mode == "arbitrate"
+ and verdict.get("participated")
+ and verdict.get("winner") == "B"
+ and len(usable) > 1
+ ):
+ chosen = usable[1]
+
payload = {
- "reply": single.get("text") or "",
- "tool_calls": single.get("tool_calls") or [],
- "chosen": single.get("provider") or "",
- "candidates": [single],
- "judge": dict(_EMPTY_VERDICT),
+ "reply": chosen["text"],
+ "knowledge": knowledge_trace,
+ "tool_calls": [],
+ "chosen": chosen["provider"],
+ "candidates": results,
+ "judge": verdict,
"judge_mode": mode,
"roles_version": version,
- "knowledge": knowledge_trace,
"total_ms": int((time.monotonic() - started) * 1000),
}
- if single.get("error") and not payload["reply"] and not payload["tool_calls"]:
- return JSONResponse(payload, status_code=502)
if idempotency_key:
idempotency[idempotency_key] = (now, payload)
try:
log_queue.put_nowait({
- "desktop_account_id": int(account["id"]), "tenant_id": tenant,
- "device_id": x_device_id, "task_id": str(body.get("task_id") or ""),
- "roles_version": version, "judge_mode": mode, "chosen": payload["chosen"],
- "judge": payload["judge"], "candidates": [single], "total_ms": payload["total_ms"],
- "customer_text": str(body.get("customer_text") or ""), "reply_text": payload["reply"],
- "purpose": str(body.get("purpose") or "chat"), "knowledge_retrieval": knowledge_result,
+ "desktop_account_id": int(account["id"]),
+ "tenant_id": str(account["tenant_id"]),
+ "device_id": x_device_id,
+ "task_id": usage_recorder.context["task_id"],
+ "roles_version": version,
+ "judge_mode": mode,
+ "chosen": payload["chosen"],
+ "judge": verdict,
+ "candidates": results,
+ "total_ms": payload["total_ms"],
+ "customer_text": str(body.get("customer_text") or ""),
+ "reply_text": payload["reply"],
+ "purpose": str(body.get("purpose") or "chat"),
+ "knowledge_retrieval": knowledge_result,
})
except asyncio.QueueFull:
- pass
+ pass # 观测数据丢一条,不能因此拖慢回复
return JSONResponse(payload)
-
- results = await asyncio.gather(
- *[
- call_outlet(
- request.app.state.client,
- outlet,
- messages,
- image_b64=image_b64,
- deadline=ANSWER_DEADLINE,
- )
- for outlet in answers
- ]
- )
- usable = [item for item in results if item.get("text")]
- if not usable:
- payload = {
- "reply": "", "chosen": "", "candidates": results,
- "judge": dict(_EMPTY_VERDICT), "judge_mode": mode,
- "roles_version": version,
- "total_ms": int((time.monotonic() - started) * 1000),
- }
- return JSONResponse(payload, status_code=502)
-
- verdict = await run_judge(
- request.app.state.client,
- judge,
- str(body.get("customer_text") or ""),
- usable[0]["text"],
- usable[1]["text"] if len(usable) > 1 else "",
- )
-
- chosen = usable[0]
- if (
- mode == "arbitrate"
- and verdict.get("participated")
- and verdict.get("winner") == "B"
- and len(usable) > 1
- ):
- chosen = usable[1]
-
- payload = {
- "reply": chosen["text"],
- "knowledge": knowledge_trace,
- "tool_calls": [],
- "chosen": chosen["provider"],
- "candidates": results,
- "judge": verdict,
- "judge_mode": mode,
- "roles_version": version,
- "total_ms": int((time.monotonic() - started) * 1000),
- }
- if idempotency_key:
- idempotency[idempotency_key] = (now, payload)
- try:
- log_queue.put_nowait({
- "desktop_account_id": int(account["id"]),
- "tenant_id": str(account["tenant_id"]),
- "device_id": x_device_id,
- "task_id": str(body.get("task_id") or ""),
- "roles_version": version,
- "judge_mode": mode,
- "chosen": payload["chosen"],
- "judge": verdict,
- "candidates": results,
- "total_ms": payload["total_ms"],
- "customer_text": str(body.get("customer_text") or ""),
- "reply_text": payload["reply"],
- "purpose": str(body.get("purpose") or "chat"),
- "knowledge_retrieval": knowledge_result,
- })
- except asyncio.QueueFull:
- pass # 观测数据丢一条,不能因此拖慢回复
- return JSONResponse(payload)
return app