From 2dcdda2bb7aebc66915ce76dd307abbbb062d896 Mon Sep 17 00:00:00 2001 From: newany Date: Sat, 26 Sep 2026 17:20:49 +0800 Subject: [PATCH] =?UTF-8?q?fix(responses):=20=E6=94=AF=E6=8C=81=20Codex=20?= =?UTF-8?q?=E5=A4=9A=E6=99=BA=E8=83=BD=E4=BD=93=E5=8D=8F=E8=AE=AE=E4=B8=8E?= =?UTF-8?q?=E5=AD=90=E4=BB=BB=E5=8A=A1=E6=98=8E=E6=96=87=E6=8F=90=E5=8F=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 支持递归解包 `namespace` 命名空间工具,解决直连时 `spawn_agent` 等协作工具被静默丢弃的问题 - 支持提取 `agent_message` 中的 `encrypted_content`,解决子智能体任务 Payload 为空及结果回传丢失问题 - 清理工具参数中的 `"encrypted": true` 私有属性,保证转发给上游的 JSON Schema 格式纯净 - 在 `ResponsesStreamConverter` 中建立双向命名空间映射,回传工具调用时补齐 `namespace: "collaboration"` - 新增多智能体消息提取与工具解包的单元测试用例 --- app/adapters/responses_adapter.py | 1482 +++--- converter.py | 8293 +++++++++++++++-------------- tests/test_responses_adapter.py | 1336 ++--- 3 files changed, 5645 insertions(+), 5466 deletions(-) diff --git a/app/adapters/responses_adapter.py b/app/adapters/responses_adapter.py index 6dbd900..7e854d2 100644 --- a/app/adapters/responses_adapter.py +++ b/app/adapters/responses_adapter.py @@ -1,235 +1,260 @@ -"""Translate requests and SSE responses between OpenAI Responses and Chat Completions.""" - -from __future__ import annotations - -import json -import os -import time -from typing import Any - -from app.upstream_io import (StreamOutputBudget, merge_tool_call_delta, new_tool_state, - seal_tool_identities, tool_identity_complete) - -# --------------------------------------------------------------------------- -# ID generation -# --------------------------------------------------------------------------- - -def _rand_id(prefix: str = "resp_") -> str: - return prefix + os.urandom(12).hex() - -# --------------------------------------------------------------------------- -# Responses to Chat requests -# --------------------------------------------------------------------------- - -def _text_format_to_response_format(fmt) -> dict | None: - """Map Responses text.format to equivalent Chat formats, rejecting unsupported variants.""" - if fmt is None: - return None - if not isinstance(fmt, dict): - raise ValueError("text.format must be an object") - kind = fmt.get("type") - if kind in (None, "text"): - return None - if kind == "json_object": - return {"type": "json_object"} - if kind == "json_schema": - schema = fmt.get("schema") - if not isinstance(schema, dict): - raise ValueError("text.format json_schema requires a schema object") - js: dict[str, Any] = {"name": fmt.get("name") or "response", "schema": schema} - if "strict" in fmt: - js["strict"] = bool(fmt["strict"]) - if isinstance(fmt.get("description"), str): - js["description"] = fmt["description"] - return {"type": "json_schema", "json_schema": js} - raise ValueError(f"unsupported text.format type: {kind}") - - -def responses_request_to_chat(body: dict) -> dict: - """Convert Responses input, instructions and tools to a Chat request.""" - messages: list[dict] = [] - - # instructions → system message - instructions = body.get("instructions") - if instructions: - messages.append({"role": "system", "content": instructions}) - - # input → messages +"""Translate requests and SSE responses between OpenAI Responses and Chat Completions.""" + +from __future__ import annotations + +import json +import os +import time +from typing import Any + +from app.upstream_io import (StreamOutputBudget, merge_tool_call_delta, new_tool_state, + seal_tool_identities, tool_identity_complete) + +# --------------------------------------------------------------------------- +# ID generation +# --------------------------------------------------------------------------- + +def _rand_id(prefix: str = "resp_") -> str: + return prefix + os.urandom(12).hex() + +# --------------------------------------------------------------------------- +# Responses to Chat requests +# --------------------------------------------------------------------------- + +def _text_format_to_response_format(fmt) -> dict | None: + """Map Responses text.format to equivalent Chat formats, rejecting unsupported variants.""" + if fmt is None: + return None + if not isinstance(fmt, dict): + raise ValueError("text.format must be an object") + kind = fmt.get("type") + if kind in (None, "text"): + return None + if kind == "json_object": + return {"type": "json_object"} + if kind == "json_schema": + schema = fmt.get("schema") + if not isinstance(schema, dict): + raise ValueError("text.format json_schema requires a schema object") + js: dict[str, Any] = {"name": fmt.get("name") or "response", "schema": schema} + if "strict" in fmt: + js["strict"] = bool(fmt["strict"]) + if isinstance(fmt.get("description"), str): + js["description"] = fmt["description"] + return {"type": "json_schema", "json_schema": js} + raise ValueError(f"unsupported text.format type: {kind}") + + +def responses_request_to_chat(body: dict) -> dict: + """Convert Responses input, instructions and tools to a Chat request.""" + messages: list[dict] = [] + + # instructions → system message + instructions = body.get("instructions") + if instructions: + messages.append({"role": "system", "content": instructions}) + + # input → messages + inp = body.get("input", []) + if isinstance(inp, str): + messages.append({"role": "user", "content": inp}) + elif isinstance(inp, list): + messages.extend(_convert_input_items(inp)) + + # Build the Chat request body. + chat: dict[str, Any] = {"messages": messages, "stream": True} + + # model + if "model" in body: + chat["model"] = body["model"] + + # Normalize function tool definitions. + tools = list(body.get("tools") or []) inp = body.get("input", []) - if isinstance(inp, str): - messages.append({"role": "user", "content": inp}) - elif isinstance(inp, list): - messages.extend(_convert_input_items(inp)) - - # Build the Chat request body. - chat: dict[str, Any] = {"messages": messages, "stream": True} + if isinstance(inp, list): + for item in inp: + if isinstance(item, dict) and item.get("type") == "additional_tools" and isinstance(item.get("tools"), list): + tools.extend(item["tools"]) - # model - if "model" in body: - chat["model"] = body["model"] + ns_map = _build_tool_namespace_map(tools) - # Normalize function tool definitions. - tools = body.get("tools") if tools: chat["tools"] = _convert_tools_for_chat(tools) + if ns_map: + chat["_tool_namespaces"] = ns_map if "tool_choice" in body: - chat["tool_choice"] = body["tool_choice"] - - # Forward supported parameters. - for key in ("temperature", "top_p", "stop", "seed", - "presence_penalty", "frequency_penalty", - "response_format", "reasoning_effort", "parallel_tool_calls", "prompt_cache_key"): - if key in body: - chat[key] = body[key] - - # Explicit top-level values override equivalent nested fields. - reasoning = body.get("reasoning") - if isinstance(reasoning, dict) and "reasoning_effort" not in chat: - effort = reasoning.get("effort") - if isinstance(effort, str) and effort.strip(): - chat["reasoning_effort"] = effort - text = body.get("text") - if isinstance(text, dict) and "response_format" not in chat: - mapped = _text_format_to_response_format(text.get("format")) - if mapped is not None: - chat["response_format"] = mapped - # max_output_tokens → max_tokens - if "max_output_tokens" in body: - chat["max_tokens"] = body["max_output_tokens"] - elif "max_tokens" in body: - chat["max_tokens"] = body["max_tokens"] - - return chat - - -def _convert_input_items(items: list) -> list[dict]: - """Convert input items and merge adjacent assistant messages with tool calls.""" - messages: list[dict] = [] - # Buffer adjacent assistant text and function calls. - pending_assistant_content: str | list[dict] | None = None - pending_tool_calls: list[dict] = [] - - def _flush_assistant(): - nonlocal pending_assistant_content, pending_tool_calls - if pending_assistant_content is not None or pending_tool_calls: - msg: dict[str, Any] = {"role": "assistant", - "content": pending_assistant_content or ""} - if pending_tool_calls: - msg["tool_calls"] = pending_tool_calls[:] - messages.append(msg) - pending_assistant_content = None - pending_tool_calls.clear() - - for item in items: - if not isinstance(item, dict): - continue - + chat["tool_choice"] = body["tool_choice"] + + # Forward supported parameters. + for key in ("temperature", "top_p", "stop", "seed", + "presence_penalty", "frequency_penalty", + "response_format", "reasoning_effort", "parallel_tool_calls", "prompt_cache_key"): + if key in body: + chat[key] = body[key] + + # Explicit top-level values override equivalent nested fields. + reasoning = body.get("reasoning") + if isinstance(reasoning, dict) and "reasoning_effort" not in chat: + effort = reasoning.get("effort") + if isinstance(effort, str) and effort.strip(): + chat["reasoning_effort"] = effort + text = body.get("text") + if isinstance(text, dict) and "response_format" not in chat: + mapped = _text_format_to_response_format(text.get("format")) + if mapped is not None: + chat["response_format"] = mapped + # max_output_tokens → max_tokens + if "max_output_tokens" in body: + chat["max_tokens"] = body["max_output_tokens"] + elif "max_tokens" in body: + chat["max_tokens"] = body["max_tokens"] + + return chat + + +def _convert_input_items(items: list) -> list[dict]: + """Convert input items and merge adjacent assistant messages with tool calls.""" + messages: list[dict] = [] + # Buffer adjacent assistant text and function calls. + pending_assistant_content: str | list[dict] | None = None + pending_tool_calls: list[dict] = [] + + def _flush_assistant(): + nonlocal pending_assistant_content, pending_tool_calls + if pending_assistant_content is not None or pending_tool_calls: + msg: dict[str, Any] = {"role": "assistant", + "content": pending_assistant_content or ""} + if pending_tool_calls: + msg["tool_calls"] = pending_tool_calls[:] + messages.append(msg) + pending_assistant_content = None + pending_tool_calls.clear() + + for item in items: + if not isinstance(item, dict): + continue + item_type = item.get("type") role = item.get("role", "") - # Untyped role messages - if item_type is None and role in ("user", "system", "developer"): - _flush_assistant() - mapped_role = "system" if role == "developer" else role - content = _extract_content(item.get("content", "")) - messages.append({"role": mapped_role, "content": content}) + # Ignore additional_tools in message sequence + if item_type == "additional_tools": continue - # Typed message items - if item_type == "message" and role in ("user", "system", "developer"): + # Agent message from multi-agent collaboration + if item_type == "agent_message": _flush_assistant() - mapped_role = "system" if role == "developer" else role content = _extract_content(item.get("content", "")) - messages.append({"role": mapped_role, "content": content}) + messages.append({"role": "user", "content": content}) continue - # Assistant output from history - if item_type == "message" and role == "assistant": - _flush_assistant() - content_parts = item.get("content", []) - text = _extract_output_text(content_parts) if isinstance(content_parts, list) else str(content_parts) - pending_assistant_content = text - continue - - # Untyped assistant messages - if item_type is None and role == "assistant": - _flush_assistant() - content = _extract_content(item.get("content", "")) - pending_assistant_content = content - continue - - # Merge calls into the preceding assistant message. - if item_type == "function_call": - arguments = item.get("arguments", "{}") - if not isinstance(arguments, str): - raise ValueError("function_call.arguments must be a JSON string") - if pending_assistant_content is None: - pending_assistant_content = "" - pending_tool_calls.append({ - "id": item.get("call_id", item.get("id", _rand_id("call_"))), - "type": "function", - "function": { - "name": item.get("name", ""), - "arguments": arguments, - }, - }) - continue - - # Map function results to tool messages. - if item_type == "function_call_output": - _flush_assistant() - messages.append({ - "role": "tool", - "tool_call_id": item.get("call_id", ""), - "content": (_extract_content(item["output"]) - if isinstance(item.get("output"), list) else item.get("output", "")), - }) - continue - - # Retain compatible message content from unknown item types. - if role: - _flush_assistant() - content = _extract_content(item.get("content", "")) - messages.append({"role": role, "content": content}) - - _flush_assistant() - return messages - - -def _extract_content(content) -> str | list[dict]: - """Convert protocol blocks without stringifying image content.""" - if isinstance(content, str): - return content - if isinstance(content, list): - parts = [] - has_image = False + # Untyped role messages + if item_type is None and role in ("user", "system", "developer"): + _flush_assistant() + mapped_role = "system" if role == "developer" else role + content = _extract_content(item.get("content", "")) + messages.append({"role": mapped_role, "content": content}) + continue + + # Typed message items + if item_type == "message" and role in ("user", "system", "developer"): + _flush_assistant() + mapped_role = "system" if role == "developer" else role + content = _extract_content(item.get("content", "")) + messages.append({"role": mapped_role, "content": content}) + continue + + # Assistant output from history + if item_type == "message" and role == "assistant": + _flush_assistant() + content_parts = item.get("content", []) + text = _extract_output_text(content_parts) if isinstance(content_parts, list) else str(content_parts) + pending_assistant_content = text + continue + + # Untyped assistant messages + if item_type is None and role == "assistant": + _flush_assistant() + content = _extract_content(item.get("content", "")) + pending_assistant_content = content + continue + + # Merge calls into the preceding assistant message. + if item_type == "function_call": + arguments = item.get("arguments", "{}") + if not isinstance(arguments, str): + raise ValueError("function_call.arguments must be a JSON string") + if pending_assistant_content is None: + pending_assistant_content = "" + pending_tool_calls.append({ + "id": item.get("call_id", item.get("id", _rand_id("call_"))), + "type": "function", + "function": { + "name": item.get("name", ""), + "arguments": arguments, + }, + }) + continue + + # Map function results to tool messages. + if item_type == "function_call_output": + _flush_assistant() + messages.append({ + "role": "tool", + "tool_call_id": item.get("call_id", ""), + "content": (_extract_content(item["output"]) + if isinstance(item.get("output"), list) else item.get("output", "")), + }) + continue + + # Retain compatible message content from unknown item types. + if role: + _flush_assistant() + content = _extract_content(item.get("content", "")) + messages.append({"role": role, "content": content}) + + _flush_assistant() + return messages + + +def _extract_content(content) -> str | list[dict]: + """Convert protocol blocks without stringifying image content.""" + if isinstance(content, str): + return content + if isinstance(content, list): + parts = [] + has_image = False for p in content: if isinstance(p, dict): kind = p.get("type") if kind in ("input_text", "text", "output_text"): parts.append({"type": "text", "text": p.get("text", "")}) + elif kind == "encrypted_content": + text = p.get("encrypted_content") or p.get("text", "") + if isinstance(text, str) and text: + parts.append({"type": "text", "text": text}) elif kind == "input_image": - if p.get("file_id"): - raise ValueError("Responses input_image file_id is not supported; provide image_url instead") - url = p.get("image_url") - if not isinstance(url, str) or not url: - raise ValueError("Responses input_image requires a non-empty image_url") - image_url = {"url": url} - if "detail" in p: - image_url["detail"] = p["detail"] - parts.append({"type": "image_url", "image_url": image_url}) - has_image = True - elif kind == "image_url": - parts.append(p) - has_image = True - elif isinstance(p, str): - parts.append({"type": "text", "text": p}) - if has_image: - return parts - return "".join(part["text"] for part in parts) or str(content) - return str(content) - - + if p.get("file_id"): + raise ValueError("Responses input_image file_id is not supported; provide image_url instead") + url = p.get("image_url") + if not isinstance(url, str) or not url: + raise ValueError("Responses input_image requires a non-empty image_url") + image_url = {"url": url} + if "detail" in p: + image_url["detail"] = p["detail"] + parts.append({"type": "image_url", "image_url": image_url}) + has_image = True + elif kind == "image_url": + parts.append(p) + has_image = True + elif isinstance(p, str): + parts.append({"type": "text", "text": p}) + if has_image: + return parts + return "".join(part["text"] for part in parts) or str(content) + return str(content) + + def _extract_output_text(content_parts: list) -> str | list[dict]: """Preserve historical images and extract output_text for text-only messages.""" if any(isinstance(part, dict) and part.get("type") in ("input_image", "image_url") @@ -242,40 +267,94 @@ def _extract_output_text(content_parts: list) -> str | list[dict]: return "".join(texts) +def _build_tool_namespace_map(tools: list, current_ns: str = "") -> dict[str, str]: + """Map tool names to their declaring namespace for Responses event reconstruction.""" + ns_map: dict[str, str] = {} + for t in tools: + if not isinstance(t, dict): + continue + t_type = t.get("type") + if t_type == "namespace" and isinstance(t.get("tools"), list): + ns_name = t.get("name", "") + full_ns = f"{current_ns}.{ns_name}" if current_ns else ns_name + ns_map.update(_build_tool_namespace_map(t["tools"], full_ns)) + elif t_type == "function": + name = t.get("name") or (t.get("function", {}).get("name") if isinstance(t.get("function"), dict) else "") + if name and current_ns: + ns_map[name] = current_ns + return ns_map + + def _convert_tools_for_chat(tools: list) -> list: """Convert Responses tool definitions to Chat function objects.""" result = [] + seen_names = set() for t in tools: if not isinstance(t, dict): continue - if t.get("type") != "function": + tool_type = t.get("type") + if tool_type == "namespace" and isinstance(t.get("tools"), list): + for sub_tool in _convert_tools_for_chat(t["tools"]): + fn_name = sub_tool.get("function", {}).get("name") + if fn_name and fn_name not in seen_names: + seen_names.add(fn_name) + result.append(sub_tool) + continue + if tool_type != "function": continue # Already in Chat format. if "function" in t: - result.append(t) + tool_obj = json.loads(json.dumps(t)) + fn_dict = tool_obj.get("function") + if isinstance(fn_dict, dict) and "parameters" in fn_dict: + props = (fn_dict.get("parameters") or {}).get("properties") + if isinstance(props, dict): + for p_name, p_val in props.items(): + if isinstance(p_val, dict) and "encrypted" in p_val: + p_val.pop("encrypted", None) + fn_name = fn_dict.get("name") if isinstance(fn_dict, dict) else None + if fn_name and fn_name not in seen_names: + seen_names.add(fn_name) + result.append(tool_obj) + elif not fn_name: + result.append(tool_obj) continue # Nest the flat Responses function fields. - fn: dict[str, Any] = {"name": t.get("name", "")} + name = t.get("name", "") + if name and name in seen_names: + continue + if name: + seen_names.add(name) + fn: dict[str, Any] = {"name": name} if "description" in t: fn["description"] = t["description"] if "parameters" in t: - fn["parameters"] = t["parameters"] + params = t["parameters"] + if isinstance(params, dict): + params = json.loads(json.dumps(params)) + props = params.get("properties") + if isinstance(props, dict): + for p_name, p_val in props.items(): + if isinstance(p_val, dict) and "encrypted" in p_val: + p_val.pop("encrypted", None) + fn["parameters"] = params if "strict" in t: - fn["strict"] = t["strict"] - result.append({"type": "function", "function": fn}) - return result - - -# --------------------------------------------------------------------------- -# Chat SSE to Responses events -# --------------------------------------------------------------------------- - -class ResponsesStreamConverter: - """Convert Chat SSE increments into Responses events.""" - + fn["strict"] = t["strict"] + result.append({"type": "function", "function": fn}) + return result + + +# --------------------------------------------------------------------------- +# Chat SSE to Responses events +# --------------------------------------------------------------------------- + +class ResponsesStreamConverter: + """Convert Chat SSE increments into Responses events.""" + def __init__(self, model: str = "unknown", parallel_tool_calls: bool = True, *, realtime: bool = False, budget: StreamOutputBudget | None = None, - tool_states: dict | None = None, declared_names=None): + tool_states: dict | None = None, declared_names=None, + tool_namespaces: dict[str, str] | None = None): self.resp_id = _rand_id("resp_") self.msg_id = _rand_id("msg_") self.model = model @@ -284,476 +363,497 @@ def __init__(self, model: str = "unknown", parallel_tool_calls: bool = True, *, self._budget = budget if budget is not None else StreamOutputBudget(0) self._tool_states = tool_states self._local_tool_states: dict[int, dict] = {} + self._tool_namespaces = dict(tool_namespaces or {}) self._declared_names = frozenset( value for value in (declared_names or ()) if isinstance(value, str) and value) self.created_at = int(time.time()) - - # Stream state - self._emitted_created = False - self._emitted_msg_item = False - self._emitted_content_part = False - - # Collected content - self._content = "" - # Reasoning items precede message items. - self._reasoning = "" - self._reasoning_item_id = _rand_id("rs_") - self._emitted_reasoning_item = False - self._reasoning_output_idx: int | None = None - self._message_output_idx: int | None = None - self._tool_calls: dict[int, dict] = {} # index → stable output item state - self._output_order: list[tuple[str, int | None]] = [] - self._next_output_idx = 0 - self._finish_reason: str | None = None - self._usage: dict | None = None - self._content_filter = False - self._seq = 0 # Monotonic emitted-event sequence - # Public methods - - def feed_line(self, line: str) -> str: - """Convert one SSE line into Responses event text.""" - line = line.strip() - if not line or not line.startswith("data:"): - return "" - data = line[5:].strip() - if data == "[DONE]": - return "" - try: - chunk = json.loads(data) - except json.JSONDecodeError: - return "" - return self._process_chunk(chunk) - - def finish(self) -> str: - """Close output items and emit the terminal response status.""" - if self._realtime: - if (self._tool_calls and self._finish_reason not in - ("tool_calls", "length", "content_filter", "content-filter", "refusal") - and not self._content_filter): - raise ValueError("tool calls require a tool_calls finish reason") - seal_tool_identities(self._tool_states if self._tool_states is not None - else self._local_tool_states, self._declared_names) - pending = self._flush_ready_tools() - else: - pending = [] - status, reason = self._final_status() - if self._realtime: - return "".join(pending) + self._finish_realtime(status) - events: list[str] = [] - - # Close reasoning items. - if self._emitted_reasoning_item: - events.append(self._evt("response.reasoning_summary_text.done", { - "output_index": 0, "summary_index": 0, "text": self._reasoning, - "item_id": self._reasoning_item_id - })) - events.append(self._evt("response.output_item.done", { - "output_index": 0, "item": self._reasoning_item(status) - })) - - # Close text content. - if self._emitted_content_part: - events.append(self._evt("response.output_text.done", { - "output_index": self._msg_idx(), "content_index": 0, "text": self._content, - "item_id": self.msg_id - })) - events.append(self._evt("response.content_part.done", { - "output_index": self._msg_idx(), "content_index": 0, - "part": {"type": "output_text", "text": self._content, "annotations": []}, - "item_id": self.msg_id - })) - - if self._emitted_msg_item: - events.append(self._evt("response.output_item.done", { - "output_index": self._msg_idx(), - "item": self._msg_item(status) - })) - - # Close function calls. - for idx in sorted(self._tool_calls): - tc = self._tool_calls[idx] - if tc.get("emitted"): - oi = tc["output_idx"] - events.append(self._evt("response.function_call_arguments.done", { - "output_index": oi, "arguments": tc["args"], "item_id": tc["fc_id"] - })) - events.append(self._evt("response.output_item.done", { - "output_index": oi, "item": self._fc_item(tc, status) - })) - - # Truncated or filtered responses must not report completion. - events.append(self._evt(f"response.{status}", { - "response": self._response_obj(status, incomplete_reason=reason) - })) - return "".join(events) - - def mark_content_filter(self) -> None: - """Keep a detector-confirmed refusal from becoming a successful terminal.""" - self._content_filter = True - - def set_validated_tools(self, tool_calls) -> None: - """Accept terminal tool metadata already validated by the shared Chat accumulator.""" - if not self._realtime: - return - expected = [self._tool_calls[index] for index in sorted(self._tool_calls)] - if len(expected) != len(tool_calls or []): - raise ValueError("validated tool calls do not match stream items") - for slot, call in zip(expected, tool_calls or []): - function = call.get("function") or {} - state = slot.get("state") - values = ((slot.get("id"), call.get("id")), - (slot.get("name"), function.get("name"))) - if state is not None: - values += ((state.get("id"), call.get("id")), - (state.get("name"), function.get("name")), - (state.get("arguments"), function.get("arguments"))) - if any(left != right for left, right in values): - raise ValueError("validated tool metadata does not match stream items") - slot["validated"] = True - - def _finish_realtime(self, status: str) -> str: - events: list[str] = [] - for kind, index in self._output_order: - if kind == "reasoning" and self._emitted_reasoning_item: - output_index = self._reasoning_output_idx - events.append(self._evt("response.reasoning_summary_text.done", { - "output_index": output_index, "summary_index": 0, "text": self._reasoning, - "item_id": self._reasoning_item_id})) - events.append(self._evt("response.output_item.done", { - "output_index": output_index, "item": self._reasoning_item(status)})) - elif kind == "message" and self._emitted_msg_item: - output_index = self._message_output_idx - events.append(self._evt("response.output_text.done", { - "output_index": output_index, "content_index": 0, "text": self._content, - "item_id": self.msg_id})) - events.append(self._evt("response.content_part.done", { - "output_index": output_index, "content_index": 0, - "part": {"type": "output_text", "text": self._content, "annotations": []}, - "item_id": self.msg_id})) - events.append(self._evt("response.output_item.done", { - "output_index": output_index, "item": self._msg_item(status)})) - elif kind == "tool": - slot = self._tool_calls[index] - if slot.get("emitted"): - output_index = slot["output_idx"] - events.append(self._evt("response.function_call_arguments.done", { - "output_index": output_index, "arguments": self._tool_arguments(slot), - "item_id": slot["fc_id"]})) - events.append(self._evt("response.output_item.done", { - "output_index": output_index, "item": self._fc_item(slot, status)})) - events.append(self._evt(f"response.{status}", { - "response": self._response_obj(status, incomplete_reason=self._final_status()[1])})) - return "".join(events) - - def get_nonstream_response(self) -> dict: - """Return the complete non-streaming Response object.""" - status, reason = self._final_status() - return self._response_obj(status, incomplete_reason=reason) - - def _final_status(self) -> tuple[str, str | None]: - """Map finish reasons to response status without hiding truncation or filtering.""" - fr = self._finish_reason - if self._content_filter and fr in (None, "stop", "tool_calls"): - return "incomplete", "content_filter" - if fr in (None, "stop", "tool_calls"): - return "completed", None - if fr == "length": - return "incomplete", "max_output_tokens" - if fr in ("content_filter", "content-filter", "refusal"): - return "incomplete", "content_filter" - return "incomplete", None - # Internal helpers - - def _process_chunk(self, chunk: dict) -> str: - events: list[str] = [] - - # Model identity - if chunk.get("model"): - self.model = chunk["model"] - - # Emit created and in_progress once. - if not self._emitted_created: - resp = self._response_obj("in_progress") - events.append(self._evt("response.created", {"response": resp})) - events.append(self._evt("response.in_progress", {"response": resp})) - self._emitted_created = True - - # usage - if chunk.get("usage"): - self._usage = chunk["usage"] - - for choice in chunk.get("choices", []): - delta = choice.get("delta", {}) - finish = choice.get("finish_reason") or None - if self._finish_reason is not None: - if (any(delta.get(key) for key in ("content", "reasoning_content", "refusal")) - or bool(delta.get("tool_calls")) or bool(delta.get("function_call"))): - raise ValueError("output after finish_reason") - if finish is not None and finish != self._finish_reason: - raise ValueError("changed finish_reason") - - # Emit reasoning before message content. - reasoning = delta.get("reasoning_content") - if reasoning: - if not self._emitted_reasoning_item: - self._reasoning_output_idx = self._claim_output("reasoning") - events.append(self._evt("response.output_item.added", { - "output_index": self._reasoning_output_idx, - "item": {"type": "reasoning", "id": self._reasoning_item_id, - "summary": [], "status": "in_progress"} - })) - self._emitted_reasoning_item = True - self._budget.charge_text(reasoning) - self._reasoning += reasoning - events.append(self._evt("response.reasoning_summary_text.delta", { - "output_index": self._reasoning_output_idx, "summary_index": 0, "delta": reasoning, - "item_id": self._reasoning_item_id - })) - - # Preserve refusal content as valid output_text. - content = (delta.get("content") or "") + (delta.get("refusal") or "") - if content: - if not self._emitted_msg_item: - self._message_output_idx = self._claim_output("message") - events.append(self._evt("response.output_item.added", { - "output_index": self._message_output_idx, - "item": self._msg_item("in_progress", empty=True) - })) - self._emitted_msg_item = True - - if not self._emitted_content_part: - events.append(self._evt("response.content_part.added", { - "output_index": self._msg_idx(), "content_index": 0, - "part": {"type": "output_text", "text": "", "annotations": []}, - "item_id": self.msg_id - })) - self._emitted_content_part = True - - self._budget.charge_text(content) - self._content += content - events.append(self._evt("response.output_text.delta", { - "output_index": self._msg_idx(), "content_index": 0, "delta": content, - "item_id": self.msg_id - })) - - # ---- tool_calls delta ---- - for tc in delta.get("tool_calls", []): - idx = tc.get("index", 0) - if idx not in self._tool_calls: - output_idx = None if self._realtime else self._claim_output("tool", idx) - self._tool_calls[idx] = { - "id": tc.get("id", ""), "name": "", "args": "", - "fc_id": _rand_id("fc_"), "output_idx": output_idx, - "emitted": False, "emitted_args_length": 0, - } - slot = self._tool_calls[idx] - if self._realtime: - state = self._sync_tool_state(idx, tc) - slot["state"] = state - else: - if tc.get("id"): - slot["id"] = tc["id"] - fn = tc.get("function", {}) - if fn.get("name"): - slot["name"] = fn["name"] - slot["_pending_args"] = fn.get("arguments") or "" - - if self._realtime: - events.extend(self._flush_tool_slot(idx, slot)) - else: - if not slot["emitted"]: - events.append(self._evt("response.output_item.added", { - "output_index": slot["output_idx"], - "item": self._fc_item(slot, "in_progress") - })) - slot["emitted"] = True - arguments = self._new_tool_arguments(slot) - if arguments: - events.append(self._evt("response.function_call_arguments.delta", { - "output_index": slot["output_idx"], - "delta": arguments, "item_id": slot["fc_id"] - })) - - if finish: - self._finish_reason = finish - if self._realtime: - seal_tool_identities(self._tool_states if self._tool_states is not None - else self._local_tool_states, self._declared_names) - for pending_idx in sorted(self._tool_calls): - events.extend(self._flush_tool_slot(pending_idx, self._tool_calls[pending_idx])) - - return "".join(events) - - def _evt(self, event_type: str, data: dict) -> str: - """Format SSE events with monotonically increasing sequence numbers.""" - self._seq += 1 - payload = {"type": event_type, **data, "sequence_number": self._seq} - return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n" - - def _claim_output(self, kind: str, index: int | None = None) -> int: - if self._realtime: - output_index = self._next_output_idx - self._next_output_idx += 1 - self._output_order.append((kind, index)) - return output_index - if kind == "reasoning": - return 0 - if kind == "message": - return 1 if self._emitted_reasoning_item else 0 - return ((1 if self._emitted_reasoning_item else 0) - + (1 if self._emitted_msg_item else 0) + len(self._tool_calls)) - - def _sync_tool_state(self, index: int, tool: dict) -> dict: - state = ((self._tool_states or {}).get(index) - if self._tool_states is not None else self._local_tool_states.get(index)) - if state is None: - if self._tool_states is not None: - raise ValueError("tool state missing from realtime accumulator") - state = new_tool_state() - self._local_tool_states[index] = state - if self._tool_states is None: - merge_tool_call_delta(state, tool, declared_names=self._declared_names, - charge=self._budget.charge_text) - else: - state["identity_complete"] = tool_identity_complete( - state, self._declared_names, terminal=bool(state.get("_terminal"))) - slot = self._tool_calls[index] - slot["id"] = state.get("id") or "" - slot["name"] = state.get("name") or "" - return state - - def _flush_tool_slot(self, index: int, slot: dict) -> list[str]: - """Start one ready tool and flush its buffered arguments exactly once.""" - if not self._realtime: - return [] - state = slot.get("state") - if state is None or not state.get("identity_complete"): - return [] - if not slot.get("emitted"): - if slot.get("output_idx") is None: - slot["output_idx"] = self._claim_output("tool", index) - events = [self._evt("response.output_item.added", { - "output_index": slot["output_idx"], - "item": self._fc_item(slot, "in_progress", include_arguments=False) - })] - slot["emitted"] = True - state["identity_emitted"] = True - else: - events = [] - arguments = self._new_tool_arguments(slot) - if arguments: - events.append(self._evt("response.function_call_arguments.delta", { - "output_index": slot["output_idx"], - "delta": arguments, "item_id": slot["fc_id"] - })) - return events - - def _flush_ready_tools(self) -> list[str]: - events: list[str] = [] - for index in sorted(self._tool_calls): - events.extend(self._flush_tool_slot(index, self._tool_calls[index])) - return events - - def _tool_arguments(self, slot: dict) -> str: - if self._realtime and slot.get("state") is not None: - return slot["state"].get("arguments") or "" - return slot["args"] - - def _new_tool_arguments(self, slot: dict) -> str: - if not self._realtime: - piece = slot.get("_pending_args", "") - slot["args"] += piece - return piece - if not slot.get("emitted"): - return "" - arguments = self._tool_arguments(slot) - emitted = slot.get("emitted_args_length", 0) - if len(arguments) < emitted or not arguments.startswith(slot.get("_emitted_prefix", "")): - raise ValueError("non-append-only tool arguments") - piece = arguments[emitted:] - slot["emitted_args_length"] = len(arguments) - slot["_emitted_prefix"] = arguments - return piece - - def _msg_idx(self) -> int: - """Return the stable message index in realtime mode or the canonical placement.""" - return (self._message_output_idx if self._realtime else - 1 if self._emitted_reasoning_item else 0) - - def _reasoning_item(self, status: str) -> dict: - """Build a reasoning item with its text in the first summary block.""" - return {"type": "reasoning", "id": self._reasoning_item_id, "status": status, - "summary": [{"type": "summary_text", "text": self._reasoning}]} - - def _msg_item(self, status: str = "in_progress", empty: bool = False) -> dict: - content = [] if empty else [ - {"type": "output_text", "text": self._content, "annotations": []} - ] - return { - "type": "message", - "id": self.msg_id, - "status": status, - "role": "assistant", - "content": content, - } - + + # Stream state + self._emitted_created = False + self._emitted_msg_item = False + self._emitted_content_part = False + + # Collected content + self._content = "" + # Reasoning items precede message items. + self._reasoning = "" + self._reasoning_item_id = _rand_id("rs_") + self._emitted_reasoning_item = False + self._reasoning_output_idx: int | None = None + self._message_output_idx: int | None = None + self._tool_calls: dict[int, dict] = {} # index → stable output item state + self._output_order: list[tuple[str, int | None]] = [] + self._next_output_idx = 0 + self._finish_reason: str | None = None + self._usage: dict | None = None + self._content_filter = False + self._seq = 0 # Monotonic emitted-event sequence + # Public methods + + def feed_line(self, line: str) -> str: + """Convert one SSE line into Responses event text.""" + line = line.strip() + if not line or not line.startswith("data:"): + return "" + data = line[5:].strip() + if data == "[DONE]": + return "" + try: + chunk = json.loads(data) + except json.JSONDecodeError: + return "" + return self._process_chunk(chunk) + + def finish(self) -> str: + """Close output items and emit the terminal response status.""" + if self._realtime: + if (self._tool_calls and self._finish_reason not in + ("tool_calls", "length", "content_filter", "content-filter", "refusal") + and not self._content_filter): + raise ValueError("tool calls require a tool_calls finish reason") + seal_tool_identities(self._tool_states if self._tool_states is not None + else self._local_tool_states, self._declared_names) + pending = self._flush_ready_tools() + else: + pending = [] + status, reason = self._final_status() + if self._realtime: + return "".join(pending) + self._finish_realtime(status) + events: list[str] = [] + + # Close reasoning items. + if self._emitted_reasoning_item: + events.append(self._evt("response.reasoning_summary_text.done", { + "output_index": 0, "summary_index": 0, "text": self._reasoning, + "item_id": self._reasoning_item_id + })) + events.append(self._evt("response.output_item.done", { + "output_index": 0, "item": self._reasoning_item(status) + })) + + # Close text content. + if self._emitted_content_part: + events.append(self._evt("response.output_text.done", { + "output_index": self._msg_idx(), "content_index": 0, "text": self._content, + "item_id": self.msg_id + })) + events.append(self._evt("response.content_part.done", { + "output_index": self._msg_idx(), "content_index": 0, + "part": {"type": "output_text", "text": self._content, "annotations": []}, + "item_id": self.msg_id + })) + + if self._emitted_msg_item: + events.append(self._evt("response.output_item.done", { + "output_index": self._msg_idx(), + "item": self._msg_item(status) + })) + + # Close function calls. + for idx in sorted(self._tool_calls): + tc = self._tool_calls[idx] + if tc.get("emitted"): + oi = tc["output_idx"] + events.append(self._evt("response.function_call_arguments.done", { + "output_index": oi, "arguments": tc["args"], "item_id": tc["fc_id"] + })) + events.append(self._evt("response.output_item.done", { + "output_index": oi, "item": self._fc_item(tc, status) + })) + + # Truncated or filtered responses must not report completion. + events.append(self._evt(f"response.{status}", { + "response": self._response_obj(status, incomplete_reason=reason) + })) + return "".join(events) + + def mark_content_filter(self) -> None: + """Keep a detector-confirmed refusal from becoming a successful terminal.""" + self._content_filter = True + + def set_validated_tools(self, tool_calls) -> None: + """Accept terminal tool metadata already validated by the shared Chat accumulator.""" + if not self._realtime: + return + expected = [self._tool_calls[index] for index in sorted(self._tool_calls)] + if len(expected) != len(tool_calls or []): + raise ValueError("validated tool calls do not match stream items") + for slot, call in zip(expected, tool_calls or []): + function = call.get("function") or {} + state = slot.get("state") + values = ((slot.get("id"), call.get("id")), + (slot.get("name"), function.get("name"))) + if state is not None: + values += ((state.get("id"), call.get("id")), + (state.get("name"), function.get("name")), + (state.get("arguments"), function.get("arguments"))) + if any(left != right for left, right in values): + raise ValueError("validated tool metadata does not match stream items") + slot["validated"] = True + + def _finish_realtime(self, status: str) -> str: + events: list[str] = [] + for kind, index in self._output_order: + if kind == "reasoning" and self._emitted_reasoning_item: + output_index = self._reasoning_output_idx + events.append(self._evt("response.reasoning_summary_text.done", { + "output_index": output_index, "summary_index": 0, "text": self._reasoning, + "item_id": self._reasoning_item_id})) + events.append(self._evt("response.output_item.done", { + "output_index": output_index, "item": self._reasoning_item(status)})) + elif kind == "message" and self._emitted_msg_item: + output_index = self._message_output_idx + events.append(self._evt("response.output_text.done", { + "output_index": output_index, "content_index": 0, "text": self._content, + "item_id": self.msg_id})) + events.append(self._evt("response.content_part.done", { + "output_index": output_index, "content_index": 0, + "part": {"type": "output_text", "text": self._content, "annotations": []}, + "item_id": self.msg_id})) + events.append(self._evt("response.output_item.done", { + "output_index": output_index, "item": self._msg_item(status)})) + elif kind == "tool": + slot = self._tool_calls[index] + if slot.get("emitted"): + output_index = slot["output_idx"] + events.append(self._evt("response.function_call_arguments.done", { + "output_index": output_index, "arguments": self._tool_arguments(slot), + "item_id": slot["fc_id"]})) + events.append(self._evt("response.output_item.done", { + "output_index": output_index, "item": self._fc_item(slot, status)})) + events.append(self._evt(f"response.{status}", { + "response": self._response_obj(status, incomplete_reason=self._final_status()[1])})) + return "".join(events) + + def get_nonstream_response(self) -> dict: + """Return the complete non-streaming Response object.""" + status, reason = self._final_status() + return self._response_obj(status, incomplete_reason=reason) + + def _final_status(self) -> tuple[str, str | None]: + """Map finish reasons to response status without hiding truncation or filtering.""" + fr = self._finish_reason + if self._content_filter and fr in (None, "stop", "tool_calls"): + return "incomplete", "content_filter" + if fr in (None, "stop", "tool_calls"): + return "completed", None + if fr == "length": + return "incomplete", "max_output_tokens" + if fr in ("content_filter", "content-filter", "refusal"): + return "incomplete", "content_filter" + return "incomplete", None + # Internal helpers + + def _process_chunk(self, chunk: dict) -> str: + events: list[str] = [] + + # Model identity + if chunk.get("model"): + self.model = chunk["model"] + + # Emit created and in_progress once. + if not self._emitted_created: + resp = self._response_obj("in_progress") + events.append(self._evt("response.created", {"response": resp})) + events.append(self._evt("response.in_progress", {"response": resp})) + self._emitted_created = True + + # usage + if chunk.get("usage"): + self._usage = chunk["usage"] + + for choice in chunk.get("choices", []): + delta = choice.get("delta", {}) + finish = choice.get("finish_reason") or None + if self._finish_reason is not None: + if (any(delta.get(key) for key in ("content", "reasoning_content", "refusal")) + or bool(delta.get("tool_calls")) or bool(delta.get("function_call"))): + raise ValueError("output after finish_reason") + if finish is not None and finish != self._finish_reason: + raise ValueError("changed finish_reason") + + # Emit reasoning before message content. + reasoning = delta.get("reasoning_content") + if reasoning: + if not self._emitted_reasoning_item: + self._reasoning_output_idx = self._claim_output("reasoning") + events.append(self._evt("response.output_item.added", { + "output_index": self._reasoning_output_idx, + "item": {"type": "reasoning", "id": self._reasoning_item_id, + "summary": [], "status": "in_progress"} + })) + self._emitted_reasoning_item = True + self._budget.charge_text(reasoning) + self._reasoning += reasoning + events.append(self._evt("response.reasoning_summary_text.delta", { + "output_index": self._reasoning_output_idx, "summary_index": 0, "delta": reasoning, + "item_id": self._reasoning_item_id + })) + + # Preserve refusal content as valid output_text. + content = (delta.get("content") or "") + (delta.get("refusal") or "") + if content: + if not self._emitted_msg_item: + self._message_output_idx = self._claim_output("message") + events.append(self._evt("response.output_item.added", { + "output_index": self._message_output_idx, + "item": self._msg_item("in_progress", empty=True) + })) + self._emitted_msg_item = True + + if not self._emitted_content_part: + events.append(self._evt("response.content_part.added", { + "output_index": self._msg_idx(), "content_index": 0, + "part": {"type": "output_text", "text": "", "annotations": []}, + "item_id": self.msg_id + })) + self._emitted_content_part = True + + self._budget.charge_text(content) + self._content += content + events.append(self._evt("response.output_text.delta", { + "output_index": self._msg_idx(), "content_index": 0, "delta": content, + "item_id": self.msg_id + })) + + # ---- tool_calls delta ---- + for tc in delta.get("tool_calls", []): + idx = tc.get("index", 0) + if idx not in self._tool_calls: + output_idx = None if self._realtime else self._claim_output("tool", idx) + self._tool_calls[idx] = { + "id": tc.get("id", ""), "name": "", "args": "", + "fc_id": _rand_id("fc_"), "output_idx": output_idx, + "emitted": False, "emitted_args_length": 0, + } + slot = self._tool_calls[idx] + if self._realtime: + state = self._sync_tool_state(idx, tc) + slot["state"] = state + else: + if tc.get("id"): + slot["id"] = tc["id"] + fn = tc.get("function", {}) + if fn.get("name"): + slot["name"] = fn["name"] + slot["_pending_args"] = fn.get("arguments") or "" + + if self._realtime: + events.extend(self._flush_tool_slot(idx, slot)) + else: + if not slot["emitted"]: + events.append(self._evt("response.output_item.added", { + "output_index": slot["output_idx"], + "item": self._fc_item(slot, "in_progress") + })) + slot["emitted"] = True + arguments = self._new_tool_arguments(slot) + if arguments: + events.append(self._evt("response.function_call_arguments.delta", { + "output_index": slot["output_idx"], + "delta": arguments, "item_id": slot["fc_id"] + })) + + if finish: + self._finish_reason = finish + if self._realtime: + seal_tool_identities(self._tool_states if self._tool_states is not None + else self._local_tool_states, self._declared_names) + for pending_idx in sorted(self._tool_calls): + events.extend(self._flush_tool_slot(pending_idx, self._tool_calls[pending_idx])) + + return "".join(events) + + def _evt(self, event_type: str, data: dict) -> str: + """Format SSE events with monotonically increasing sequence numbers.""" + self._seq += 1 + payload = {"type": event_type, **data, "sequence_number": self._seq} + return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n" + + def _claim_output(self, kind: str, index: int | None = None) -> int: + if self._realtime: + output_index = self._next_output_idx + self._next_output_idx += 1 + self._output_order.append((kind, index)) + return output_index + if kind == "reasoning": + return 0 + if kind == "message": + return 1 if self._emitted_reasoning_item else 0 + return ((1 if self._emitted_reasoning_item else 0) + + (1 if self._emitted_msg_item else 0) + len(self._tool_calls)) + + def _sync_tool_state(self, index: int, tool: dict) -> dict: + state = ((self._tool_states or {}).get(index) + if self._tool_states is not None else self._local_tool_states.get(index)) + if state is None: + if self._tool_states is not None: + raise ValueError("tool state missing from realtime accumulator") + state = new_tool_state() + self._local_tool_states[index] = state + if self._tool_states is None: + merge_tool_call_delta(state, tool, declared_names=self._declared_names, + charge=self._budget.charge_text) + else: + state["identity_complete"] = tool_identity_complete( + state, self._declared_names, terminal=bool(state.get("_terminal"))) + slot = self._tool_calls[index] + slot["id"] = state.get("id") or "" + slot["name"] = state.get("name") or "" + return state + + def _flush_tool_slot(self, index: int, slot: dict) -> list[str]: + """Start one ready tool and flush its buffered arguments exactly once.""" + if not self._realtime: + return [] + state = slot.get("state") + if state is None or not state.get("identity_complete"): + return [] + if not slot.get("emitted"): + if slot.get("output_idx") is None: + slot["output_idx"] = self._claim_output("tool", index) + events = [self._evt("response.output_item.added", { + "output_index": slot["output_idx"], + "item": self._fc_item(slot, "in_progress", include_arguments=False) + })] + slot["emitted"] = True + state["identity_emitted"] = True + else: + events = [] + arguments = self._new_tool_arguments(slot) + if arguments: + events.append(self._evt("response.function_call_arguments.delta", { + "output_index": slot["output_idx"], + "delta": arguments, "item_id": slot["fc_id"] + })) + return events + + def _flush_ready_tools(self) -> list[str]: + events: list[str] = [] + for index in sorted(self._tool_calls): + events.extend(self._flush_tool_slot(index, self._tool_calls[index])) + return events + + def _tool_arguments(self, slot: dict) -> str: + if self._realtime and slot.get("state") is not None: + return slot["state"].get("arguments") or "" + return slot["args"] + + def _new_tool_arguments(self, slot: dict) -> str: + if not self._realtime: + piece = slot.get("_pending_args", "") + slot["args"] += piece + return piece + if not slot.get("emitted"): + return "" + arguments = self._tool_arguments(slot) + emitted = slot.get("emitted_args_length", 0) + if len(arguments) < emitted or not arguments.startswith(slot.get("_emitted_prefix", "")): + raise ValueError("non-append-only tool arguments") + piece = arguments[emitted:] + slot["emitted_args_length"] = len(arguments) + slot["_emitted_prefix"] = arguments + return piece + + def _msg_idx(self) -> int: + """Return the stable message index in realtime mode or the canonical placement.""" + return (self._message_output_idx if self._realtime else + 1 if self._emitted_reasoning_item else 0) + + def _reasoning_item(self, status: str) -> dict: + """Build a reasoning item with its text in the first summary block.""" + return {"type": "reasoning", "id": self._reasoning_item_id, "status": status, + "summary": [{"type": "summary_text", "text": self._reasoning}]} + + def _msg_item(self, status: str = "in_progress", empty: bool = False) -> dict: + content = [] if empty else [ + {"type": "output_text", "text": self._content, "annotations": []} + ] + return { + "type": "message", + "id": self.msg_id, + "status": status, + "role": "assistant", + "content": content, + } + def _fc_item(self, tc: dict, status: str, *, include_arguments: bool = True) -> dict: - return { + raw_name = tc["name"] or (tc.get("state", {}).get("name") or "") + name = raw_name + ns = tc.get("namespace") + if not ns: + if "." in raw_name: + ns_prefix, base_name = raw_name.split(".", 1) + if ns_prefix in ("collaboration", "collaboration-optimize") or (self._tool_namespaces and base_name in self._tool_namespaces): + ns = ns_prefix + name = base_name + elif "__" in raw_name: + ns_prefix, base_name = raw_name.split("__", 1) + if ns_prefix in ("collaboration", "collaboration-optimize") or (self._tool_namespaces and base_name in self._tool_namespaces): + ns = ns_prefix + name = base_name + elif self._tool_namespaces and raw_name in self._tool_namespaces: + ns = self._tool_namespaces[raw_name] + + item: dict[str, Any] = { "type": "function_call", "id": tc["fc_id"], "call_id": tc["id"] or (tc.get("state", {}).get("id") or ""), - "name": tc["name"] or (tc.get("state", {}).get("name") or ""), + "name": name, "arguments": self._tool_arguments(tc) if include_arguments else "", "status": status, } - - def _response_obj(self, status: str, incomplete_reason: str | None = None) -> dict: - output = [] - if self._realtime: - for kind, index in self._output_order: - if kind == "reasoning" and self._emitted_reasoning_item: - output.append(self._reasoning_item(status)) - elif kind == "message" and self._emitted_msg_item: - output.append(self._msg_item(status)) - elif kind == "tool" and self._tool_calls[index].get("emitted"): - output.append(self._fc_item(self._tool_calls[index], status)) - else: - if self._emitted_reasoning_item: - output.append(self._reasoning_item(status)) - if self._emitted_msg_item or self._content: - output.append(self._msg_item(status)) - for idx in sorted(self._tool_calls): - tc = self._tool_calls[idx] - if tc.get("emitted"): - output.append(self._fc_item(tc, status)) - - usage = None - if self._usage: - u = self._usage - reasoning_tokens = (u.get("completion_tokens_details") or {}).get("reasoning_tokens", 0) - # Omit unknown cache details instead of reporting a fabricated zero. - cached = (u.get("prompt_tokens_details") or {}).get("cached_tokens", - u.get("cache_read_input_tokens")) - usage = { - "input_tokens": u.get("prompt_tokens", u.get("input_tokens", 0)), - "output_tokens": u.get("completion_tokens", u.get("output_tokens", 0)), - "output_tokens_details": {"reasoning_tokens": reasoning_tokens}, - "total_tokens": u.get("total_tokens", 0), - } - if isinstance(cached, int) and not isinstance(cached, bool) and cached >= 0: - usage["input_tokens_details"] = {"cached_tokens": cached} - - obj = { - "id": self.resp_id, - "object": "response", - "created_at": self.created_at, - "status": status, - "model": self.model, - "output": output, - "parallel_tool_calls": self._parallel_tool_calls, - "usage": usage, - } - if status == "incomplete": - # Preserve a null incomplete reason when the upstream supplies none. - obj["incomplete_details"] = {"reason": incomplete_reason} - return obj + if ns: + item["namespace"] = ns + return item + + def _response_obj(self, status: str, incomplete_reason: str | None = None) -> dict: + output = [] + if self._realtime: + for kind, index in self._output_order: + if kind == "reasoning" and self._emitted_reasoning_item: + output.append(self._reasoning_item(status)) + elif kind == "message" and self._emitted_msg_item: + output.append(self._msg_item(status)) + elif kind == "tool" and self._tool_calls[index].get("emitted"): + output.append(self._fc_item(self._tool_calls[index], status)) + else: + if self._emitted_reasoning_item: + output.append(self._reasoning_item(status)) + if self._emitted_msg_item or self._content: + output.append(self._msg_item(status)) + for idx in sorted(self._tool_calls): + tc = self._tool_calls[idx] + if tc.get("emitted"): + output.append(self._fc_item(tc, status)) + + usage = None + if self._usage: + u = self._usage + reasoning_tokens = (u.get("completion_tokens_details") or {}).get("reasoning_tokens", 0) + # Omit unknown cache details instead of reporting a fabricated zero. + cached = (u.get("prompt_tokens_details") or {}).get("cached_tokens", + u.get("cache_read_input_tokens")) + usage = { + "input_tokens": u.get("prompt_tokens", u.get("input_tokens", 0)), + "output_tokens": u.get("completion_tokens", u.get("output_tokens", 0)), + "output_tokens_details": {"reasoning_tokens": reasoning_tokens}, + "total_tokens": u.get("total_tokens", 0), + } + if isinstance(cached, int) and not isinstance(cached, bool) and cached >= 0: + usage["input_tokens_details"] = {"cached_tokens": cached} + + obj = { + "id": self.resp_id, + "object": "response", + "created_at": self.created_at, + "status": status, + "model": self.model, + "output": output, + "parallel_tool_calls": self._parallel_tool_calls, + "usage": usage, + } + if status == "incomplete": + # Preserve a null incomplete reason when the upstream supplies none. + obj["incomplete_details"] = {"reason": incomplete_reason} + return obj diff --git a/converter.py b/converter.py index 543c27a..f9a3c23 100644 --- a/converter.py +++ b/converter.py @@ -1,3676 +1,3682 @@ -#!/usr/bin/env python3 -"""Expose CodeBuddy and WorkBuddy through compatible Chat, Responses and Messages APIs.""" - -from __future__ import annotations - -import argparse -import asyncio -import hashlib -import json -import math -import os -import re -import secrets -import shutil -import sys -import threading -import time -import uuid -from collections import OrderedDict -from contextlib import asynccontextmanager, nullcontext -from dataclasses import dataclass -from datetime import datetime, timedelta, timezone -from pathlib import Path -from typing import Optional - -import httpx -from fastapi import FastAPI, Header, HTTPException, Request -from fastapi.exception_handlers import http_exception_handler as _default_http_exception_handler -from fastapi.responses import JSONResponse, Response, StreamingResponse -from starlette.concurrency import run_in_threadpool -import uvicorn as uvicorn # Keep the existing embedding/test hook. - -try: - from app.desensitize import desensitize_body -except ImportError: # Disable desensitization when its module is unavailable. - def desensitize_body(body, roles=("system",), desensitize_harness_user=False, - desensitize_tools=False, compact_harness=False, - strip_tool_metadata=False): - return body - -from app.adapters.responses_adapter import ( - responses_request_to_chat, - ResponsesStreamConverter, -) -from app.adapters.responses_projection import PROJECTION_MODES, project_responses_chat_body -from app.adapters.anthropic_adapter import ( - anthropic_request_to_chat, - AnthropicStreamConverter, -) - -from app import auth_oauth -from app import trial_rewards -from app import buddy, checkin as checkin_service, model_policy, travel -from app.credential_cooldowns import CredentialCooldowns -from app.model_blocks import ModelBlocks -from app.usage_snapshots import UsageSnapshots -from app.client_hangup import ClientHungUp, await_or_hangup -from app.observability import (AuditMiddleware, observe_recovery, observe_route, observe_stream_mode, - observe_responses_projection, observe_usage, observe_attempt, - observe_failure, observe_failure_seq) -from app.credential_io import (CredentialFileError, read_import_file, atomic_write_credential, - credential_file_lock) -from app.upstream_io import (ChatSSEAccumulator, StreamOutputBudget, UpstreamHTTPError, - UpstreamResponseError, open_backend_stream, parse_retry_after, - read_bounded_error) -from app.inference_resources import (AccountCapacity, InferenceResourcesMiddleware, inference_lifespan, - request_resources, release_credential) -from app.request_context import SessionIdentifierError, current_context -from app import model_capabilities -from app.message_normalization import merge_intl_user_images -from app.adapters.chat_input import normalize_chat_messages -from app.model_catalog_view import INTERNATIONAL as SHARED_INTL_PROFILES, share_models -from app.inference_auth import require_api_key -from app.admin_auth import SessionStoreError -from app.content_filter import ContentFilterDetector, is_filter_error -from app.request_limits import ImageLimitError, apply_image_policy -from app.safe_logging import format_log_body, sanitize_log_text -from app.startup import allows_open_noauth, load_startup_env, run_server -from app.site_routing import (DOMESTIC, INTERNATIONAL, PROFILE_ENDPOINTS, site_for_auth, site_for_headers, - profile_for_auth, profile_for_headers, profile_region, profile_product, - profile_site, chat_url_for_headers, refresh_url_for_auth) -from app.client_profiles import CLI_VERSION, CLI_USER_AGENT, credential_headers, catalog_cache_key, account_key -try: - from app import credits as credits_mod -except ImportError: # Disable credit maintenance when its module is unavailable. - credits_mod = None - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - -APP_VERSION = Path(__file__).with_name("VERSION").read_text(encoding="utf-8").strip() -BACKEND = "https://copilot.tencent.com" -DEFAULT_DOMAIN = "www.codebuddy.cn" -CBC_VERSION = CLI_VERSION -USER_AGENT = CLI_USER_AGENT - -_STREAM_MODES = ("compatible", "realtime") -_REQUEST_POLICY_KEY = object() # Process-local object key; never serializable or client injectable. - - -@dataclass(frozen=True) -class _StreamRequestPolicy: - mode: str - aggregate: bool - max_collect_bytes: int - - @property - def realtime(self): - return self.mode == "realtime" - - -def _snapshot_stream_policy(protocol: str, body: dict) -> _StreamRequestPolicy: - """Freeze hot streaming choices and memory policy for one request and all failovers.""" - mode = CONFIG.get("stream_mode", "compatible") - if mode not in _STREAM_MODES: - raise ValueError("invalid stream mode") - compatible_aggregate = protocol == "responses" or bool(body.get("tools")) - return _StreamRequestPolicy(mode, mode == "compatible" and compatible_aggregate, - max(0, int(CONFIG.get("max_collect_bytes", 0) or 0))) - - -def _body_with_stream_policy(body: dict, policy: _StreamRequestPolicy) -> dict: - """Attach a process-local snapshot without exposing an upstream payload override.""" - routed = dict(body) - routed[_REQUEST_POLICY_KEY] = policy - return routed - - -# --------------------------------------------------------------------------- -# Platform-specific credential directories -# --------------------------------------------------------------------------- - -def managed_auth_dir() -> Path: - """Use CODEBUDDY_AUTH_DIR when set, otherwise the project's auth directory.""" - env_dir = os.environ.get("CODEBUDDY_AUTH_DIR") - return Path(env_dir) if env_dir else Path(__file__).resolve().parent / "auth" - - -def auth_dirs() -> list[Path]: - """Locate desktop credentials used only as seed files.""" - home = Path.home() - plat = sys.platform - if plat == "darwin": - return [home / "Library" / "Application Support" / "CodeBuddyExtension" / "Data" / "Public" / "auth"] - if plat == "win32": - local = Path(os.environ.get("LOCALAPPDATA", home / "AppData" / "Local")) - return [local / "CodeBuddyExtension" / "Data" / "Public" / "auth"] - xdg = Path(os.environ.get("XDG_DATA_HOME", home / ".local" / "share")) - return [xdg / "CodeBuddyExtension" / "Data" / "Public" / "auth"] - - -def seed_credentials(): - """Seed missing managed credentials without overwriting files; skip custom auth directories.""" - if os.environ.get("CODEBUDDY_AUTH_DIR"): - return - dst_dir = managed_auth_dir() - dst_dir.mkdir(parents=True, exist_ok=True) - try: - os.chmod(dst_dir, 0o700) - except OSError: - pass - have_uids = {u for u in (_cred_identity(f) for f in dst_dir.glob("*.info")) if u} - for src_dir in auth_dirs(): - if not src_dir.is_dir(): - continue - for f in sorted(src_dir.glob("*.info")): - try: - data = json.loads(f.read_text(encoding="utf-8")) - except Exception as e: - _log(f"[cred] 种子跳过(无法解析): {f.name}: {e}") - continue - uid, verr = auth_oauth.validate_cred_data(data) - if verr: - _log(f"[cred] 种子跳过(入库校验失败:{verr}): {f.name}") - continue - identity = _credential_identity(data) - if uid and identity in have_uids: - _log(f"[cred] 种子跳过(同账号已在自管目录): {f.name}") - continue - dst = dst_dir / f.name - if not dst.exists(): - try: - shutil.copyfile(f, dst) - os.chmod(dst, 0o600) - if uid: - have_uids.add(identity) - _log(f"[cred] 已复制桌面端凭据到自管目录: {f.name}") - except OSError as e: - _log(f"[cred] 复制凭据失败 {f.name}: {e}") - - -def find_auth_files() -> list[Path]: - """Find managed .info credential files.""" - d = managed_auth_dir() - return sorted(d.glob("*.info")) if d.is_dir() else [] - - -def _cred_uid(path) -> Optional[str]: - """Read the account UID for deduplication, or return None when unavailable.""" - try: - d = json.loads(Path(path).read_text(encoding="utf-8")) - acct = d.get("account") - if not isinstance(acct, dict): - arr = d.get("accounts") - acct = arr[0] if isinstance(arr, list) and arr and isinstance(arr[0], dict) else {} - return acct.get("uid") - except Exception: - return None - -def _credential_account(data: dict) -> dict: - account = data.get("account") - if not isinstance(account, dict): - accounts = data.get("accounts") or [] - account = accounts[0] if isinstance(accounts, list) and accounts and isinstance(accounts[0], dict) else {} - return account - - -def _credential_identity(data: dict) -> str: - account = _credential_account(data) - return account_key(profile_for_auth(data.get("auth") or {}), account.get("uid"), account.get("enterpriseId")) - - -def _cred_identity(path) -> str | None: - try: - return _credential_identity(json.loads(Path(path).read_text(encoding="utf-8"))) - except (OSError, ValueError, TypeError, AttributeError): - return None - - -def find_auth_file() -> Path | None: - files = find_auth_files() - return files[0] if files else None - - -# --------------------------------------------------------------------------- -# Credential loading, refresh and persistence -# --------------------------------------------------------------------------- - -class CredentialManager: - """Load credentials and refresh expiring tokens with persistence.""" - - def __init__(self, path: Path): - self.path = path - self._lock = threading.RLock() - self._cached: dict | None = None - self._mtime = None - self._generation = 0 - - def _read_raw(self) -> dict: - with open(self.path, "r", encoding="utf-8") as f: - return json.load(f) - - def _file_version(self): - st = self.path.stat() - return st.st_dev, st.st_ino, st.st_mtime_ns, st.st_size - - def _load_if_stale(self): - """Reload changed credentials and invalidate leases held by older requests.""" - mt = self._file_version() - if self._cached is None or mt != self._mtime: - self._cached = self._read_raw() - self._mtime = mt - self._generation += 1 - - def _session(self) -> dict: - self._load_if_stale() - if self._cached is None: - raise RuntimeError(f"无法读取 auth 文件:{self.path}") - return self._cached - - def _is_expired(self) -> bool: - s = self._session() - expires_at = (s.get("auth") or {}).get("expiresAt") or 0 - # Treat tokens as expired 60 seconds early. - return time.time() * 1000 >= (expires_at - 60_000) - - def _refresh_needed(self, margin_s, keepalive_s): - summary = self.summary() - now = time.time() - exp = (summary.get("token_expires_at") or 0) / 1000 - last = (summary.get("last_refresh_time") or 0) / 1000 - return bool(summary.get("token_expired") or (exp and exp - now < margin_s) - or (keepalive_s > 0 and (last <= 0 or now - last >= keepalive_s))) - - def _refresh(self, margin_s=60, keepalive_s=0): - with self._lock: - if not self._refresh_needed(margin_s, keepalive_s): - return False - with credential_file_lock(self.path.parent, self.path.name): - if not self._refresh_needed(margin_s, keepalive_s): - return False - self._refresh_locked() - return True - - def _refresh_locked(self): - """Share the import lock so token refresh cannot overwrite a newer login.""" - s = self._session() - auth = s.get("auth") or {} - headers = self._build_headers_from(auth, _credential_account(s)) - headers["X-Refresh-Token"] = auth.get("refreshToken", "") - headers["X-Auth-Refresh-Source"] = "plugin" - url = refresh_url_for_auth(auth) - try: - with httpx.Client(timeout=15) as c: - r = c.post(url, headers=headers, json={}) - data = r.json() - except Exception as e: - raise RuntimeError(f"刷新 token 网络失败:{e}") - if data.get("code") != 0 or not data.get("data"): - raise RuntimeError(f"刷新 token 失败:{data.get('msg', data)}") - new_auth = dict(data["data"]) - if not isinstance(new_auth.get("accessToken"), str) or not new_auth["accessToken"]: - raise RuntimeError("刷新接口未返回有效的 accessToken") - new_auth["domain"] = new_auth.get("domain") or auth.get("domain") - new_auth["refreshToken"] = new_auth.get("refreshToken") or auth.get("refreshToken") - new_auth["lastRefreshTime"] = int(time.time() * 1000) - if not new_auth.get("expiresAt") and new_auth.get("expiresIn"): - new_auth["expiresAt"] = int(time.time() * 1000) + new_auth["expiresIn"] * 1000 - if not new_auth.get("refreshExpiresAt") and new_auth.get("refreshExpiresIn"): - new_auth["refreshExpiresAt"] = int(time.time() * 1000) + new_auth["refreshExpiresIn"] * 1000 - updated = dict(s, auth=new_auth) - atomic_write_credential(self.path.parent, self.path.name, - json.dumps(updated, ensure_ascii=False, indent=2).encode("utf-8")) - self._cached = updated - self._mtime = self._file_version() - self._generation += 1 - - def _build_headers_from(self, auth: dict, account: dict) -> dict: - return credential_headers(auth, account) - - def get_headers(self) -> dict: - """Return upstream headers with a current token, refreshing when necessary.""" - with self._lock: - if self._is_expired(): - self._refresh() - s = self._session() - return self._build_headers_from(s.get("auth") or {}, _credential_account(s)) - - def refresh_if_due(self, margin_s: int, keepalive_s: int) -> bool: - """Refresh under the shared foreground/background lock after rechecking expiry.""" - return self._refresh(margin_s, keepalive_s) - - def invalidate(self): - """Reload imported credentials while retaining the manager and refresh lock.""" - with self._lock: - self._cached = None - self._mtime = None - self._generation += 1 - - - def summary(self) -> dict: - with self._lock: - s = self._session() - auth = s.get("auth") or {} - acct = _credential_account(s) - profile = profile_for_auth(auth) - return { - "uid": str(acct.get("uid") or "") or None, - "account_key": _credential_identity(s), - "site": profile_site(profile), - "profile": profile, "region": profile_region(profile), "product": profile_product(profile), - "nickname": acct.get("nickname"), - "enterpriseName": acct.get("enterpriseName"), - "token_expires_at": auth.get("expiresAt", 0), - "token_expired": self._is_expired(), - "last_refresh_time": auth.get("lastRefreshTime") or 0, - } - - -STORAGE_WARN_INTERVAL = 300 # Rate limit for persistence-failure warnings -USAGE_CACHE_MAX_AGE_S = 7 * 24 * 3600 # A cached usage snapshot older than this is not shown. -STICKY_TTL = 30 * 60 # Idle session binding lifetime in seconds -STICKY_MAX = 512 # Session binding capacity -CRED_COOLDOWN = 300 # Credential cooldown in seconds -MODEL_COOLDOWN = 600 # Model cooldown when a 429 omits reset time -MODEL_COOLDOWN_MAX = 86400 # Maximum model cooldown in seconds -MODEL_SITE_BLOCK_S = 6 * 3600 # Initial unsupported-model backoff -MODEL_SITE_BLOCK_MAX_S = 24 * 3600 # Maximum unsupported-model backoff -# An unsupported model must be routed to a different backend. -MODEL_NOT_SERVABLE_CODES = frozenset({"11102"}) -_NOT_SERVABLE_MSG = re.compile(r"service info not found|model .{0,80}not (?:found|supported)", re.I) -CRED_REFRESH_MARGIN = 600 # Proactive refresh margin in seconds -CRED_KEEPALIVE_S = 24 * 3600 # Maximum idle interval before refreshing -CRED_KEEPALIVE_RETRY_S = 3600 # Keepalive retry interval, independent of expiry retries - - -def _msg_text(m: dict) -> str: - c = m.get("content") - if isinstance(c, str): - return c - if isinstance(c, list): - return "\n".join(p.get("text", "") for p in c if isinstance(p, dict)) - return "" - - -def session_key(payload: dict) -> str | None: - """Derive a stable session key from system instructions and the first user message.""" - msgs = payload.get("messages") - if not msgs: - inp = payload.get("input") # Responses API - if isinstance(inp, str): - msgs = [{"role": "user", "content": inp}] - elif isinstance(inp, list): - msgs = inp - msgs = msgs or [] - if not msgs: - return None - system = "" - for m in msgs: - if m.get("role") in ("system", "developer"): - system += _msg_text(m) - else: - break - first_user = next((_msg_text(m) for m in msgs if m.get("role") == "user"), "") - if not system and not first_user: - return None - return hashlib.sha256((system + "\x00" + first_user).encode("utf-8", "replace")).hexdigest()[:32] - - -def _parse_reset_time(raw: bytes) -> float | None: - """Parse a quota reset timestamp from a 429 response and return epoch seconds.""" - try: - text = raw.decode("utf-8", "replace") - except Exception: - return None - m = re.search(r"(\d{4}-\d{2}-\d{2})\s+(\d{2}:\d{2}:\d{2})\s*UTC\s*([+-]?\d+)", text) - if not m: - return None - try: - dt = datetime.strptime(f"{m.group(1)} {m.group(2)}", "%Y-%m-%d %H:%M:%S") - tz = timezone(timedelta(hours=int(m.group(3)))) - return dt.replace(tzinfo=tz).timestamp() - except ValueError: - return None - - - - -def _parse_not_servable(raw: bytes, status: int): - """Recognize unsupported-model errors from code/message fields, excluding incidental IDs.""" - if status not in (400, 404) or not raw: - return None - try: - payload = json.loads(raw.decode("utf-8", "replace")) - except (ValueError, UnicodeError, RecursionError): - return None - if not isinstance(payload, dict): - return None - nodes = [payload] - inner = payload.get("error") - if isinstance(inner, dict): - nodes.append(inner) - code = msg = "" - for node in nodes: - for key in ("code", "errCode", "error_code"): - value = node.get(key) - if value is not None and str(value).strip(): - code = code or str(value).strip() - for key in ("msg", "message"): - value = node.get(key) - if isinstance(value, str) and value.strip(): - msg = msg or value.strip() - if not code and not msg: - return None - if code in MODEL_NOT_SERVABLE_CODES or _NOT_SERVABLE_MSG.search(msg): - return code or "11102", msg[:200] - return None - - -def _block_model(model: str | None) -> str | None: - """Track blocked models by public ID, normalizing the international auto alias.""" - return "auto" if model == "default-model" else model - - -def _dynamic_request_headers(skey: str | None) -> dict: - """Generate upstream request IDs while keeping session IDs stable.""" - rid = secrets.token_hex(16) # X-Request-ID == X-Conversation-Message-ID - crid = secrets.token_hex(16) # X-Conversation-Request-ID == X-Root-Request-ID == trace id - span, parent = secrets.token_hex(8), secrets.token_hex(8) - if skey: - conv = str(uuid.UUID(hex=hashlib.sha256(skey.encode()).hexdigest()[:32])) - else: - conv = str(uuid.uuid4()) - return { - "X-Conversation-ID": conv, - "X-Request-ID": rid, - "X-Conversation-Message-ID": rid, - "X-Conversation-Request-ID": crid, - "X-Root-Request-ID": crid, - "X-Trace-ID": crid, - "traceparent": f"00-{crid}-{span}-01", - "b3": f"{crid}-{span}-1-{parent}", - "X-B3-TraceId": crid, - "X-B3-SpanId": span, - "X-B3-ParentSpanId": parent, - "X-B3-Sampled": "1", - } - - -class CredentialPool: - """Manage credential discovery, reloads, sticky sessions, cooldowns and refresh.""" - - def __init__(self, paths: list[Path] | None = None, scan: bool = False, - blocks_path: Path | None = None, cooldowns_path: Path | None = None, *, state_store=None): - self._lock = threading.RLock() - self._entries: list[dict] = [] # {id, cm, fail_until} - self._sticky: "OrderedDict[str, tuple[str, float]]" = OrderedDict() - self._model_fail: dict[tuple[str, str], float] = {} # Per-credential/model 429 expiry - # Keep unsupported-model backoff isolated by backend and model. - self._blocks = ModelBlocks(blocks_path, ttl_s=MODEL_SITE_BLOCK_S, max_ttl_s=MODEL_SITE_BLOCK_MAX_S, store=state_store) - # Cooldowns outlive a restart so a backend that just refused is not retried immediately. - self._cooldowns = CredentialCooldowns(cooldowns_path, store=state_store) - self._storage_warned = 0.0 # Rate limit for persistence-failure warnings - self._rr = {None: 0, "cn": 0, "intl": 0} - self._ledger = None # Prefer credits expiring sooner. - self._capacity = AccountCapacity() - self._scan = scan # Rescan credentials before selection. - self._ignored_duplicates: set[str] = set() - self._sync_pending: set[str] = set() - self._syncing: set[str] = set() - self._sync_event = threading.Event() - self._sync_retry: dict[str, float] = {} - self._sync_attempts: dict[str, int] = {} - self.reload(paths or []) - if self._scan: - self._rescan() # Discover credentials at startup. - - def reload(self, paths: list[Path], *, reset: bool = True): - """Reset authentication only for changed or imported files and schedule catalog refresh.""" - with self._lock: - by_id = {entry["id"]: entry for entry in self._entries} - have_uids = {entry["account_key"]: entry["id"] - for entry in self._entries if entry.get("uid")} - for path in paths: - cid = str(Path(path).resolve()) - if not os.path.exists(cid): - continue - entry = by_id.get(cid) - if entry is not None: - if reset: - entry["cm"].invalidate() - try: - summary = entry["cm"].summary() - except Exception: - continue # A damaged file must not block other credentials. - generation = entry["cm"]._generation - identity = summary["account_key"] - changed = reset or generation != entry.get("generation") - if changed: - old_identity = entry.get("account_key") - replaced = old_identity != identity - if replaced: - self._model_fail = {key: until for key, until in self._model_fail.items() if key[0] != cid} - self._sticky = OrderedDict((key, value) for key, value in self._sticky.items() if value[0] != cid) - # The path now belongs to another account; its cooldowns must not carry over. - self.forget_credential_state(entry) - if entry.get("uid"): - have_uids.pop(old_identity, None) - entry.update(uid=summary.get("uid"), profile=summary["profile"], site=summary["site"], - account_key=identity, generation=generation, catalog_dirty=True) - self._bind_entry(entry) - if reset or replaced: - entry.update(fail_until=0.0, keepalive_after=0.0) - if reset and not replaced: - # An explicit reload re-evaluates auth for the same account, so the - # persisted breaker is lifted; a replacement must keep the incoming - # account's own breaker, which is hydrated just below. - self._forget_credential_cooldown(entry) - self._adopt_cooldowns(entry) - if entry.get("uid"): - have_uids[identity] = cid - self._queue_sync(cid) - continue - manager = CredentialManager(Path(cid)) - try: - summary = manager.summary() - except Exception: - summary = {} - uid = summary.get("uid") - profile = summary.get("profile", "cn-cli") - identity_key = summary.get("account_key") - if uid and identity_key in have_uids: - if cid not in self._ignored_duplicates: - _log(f"[cred] 忽略重复账号凭据: {Path(cid).name}(同产品账号与 {Path(have_uids[identity_key]).name} 重复)") - self._ignored_duplicates.add(cid) - continue - entry = {"id": cid, "cm": manager, "fail_until": 0.0 if summary else time.time() + CRED_COOLDOWN, "uid": uid, - "site": summary.get("site"), "profile": profile, "generation": manager._generation, - "account_key": identity_key, "catalog_dirty": True} - self._bind_entry(entry) - self._adopt_cooldowns(entry) - self._entries.append(entry) - by_id[cid] = entry - self._ignored_duplicates.discard(cid) - if uid: - have_uids[identity_key] = cid - self._queue_sync(cid) - - def _queue_sync(self, cid): - entry = next((entry for entry in self._entries if entry["id"] == cid), None) - if entry is not None and not model_policy.credential_enabled(CONFIG, entry): - return - self._sync_pending.add(cid) - self._sync_retry.pop(cid, None) - self._sync_attempts.pop(cid, None) - self._sync_event.set() - if CONFIG.get("cred_pool") is self: - _publish_model_cache() - else: - invalidate_model_table() - - def begin_sync(self, *, all_entries=False): - """Drain refresh work and clear its wake event under the same lock.""" - with self._lock: - active = {entry["id"] for entry in self._entries if model_policy.credential_enabled(CONFIG, entry)} - self._sync_pending.intersection_update(active) - self._sync_retry = {cid: deadline for cid, deadline in self._sync_retry.items() if cid in active} - due = {cid for cid, deadline in self._sync_retry.items() if deadline <= time.monotonic()} - self._sync_pending.update(due) - ids = active if all_entries else set(self._sync_pending) - self._sync_pending.difference_update(ids) - if not self._sync_pending: - self._sync_event.clear() - self._syncing.update(ids) - return ids - - def end_sync(self, ids, failed=()): - with self._lock: - self._syncing.difference_update(ids) - present = {entry["id"] for entry in self._entries if model_policy.credential_enabled(CONFIG, entry)} - for cid in ids: - if cid in failed and cid in present and cid not in self._sync_pending: - attempt = min(self._sync_attempts.get(cid, 0) + 1, 5) - self._sync_attempts[cid] = attempt - self._sync_retry[cid] = time.monotonic() + min(60 * 2 ** (attempt - 1), 900) - else: - self._sync_retry.pop(cid, None) - self._sync_attempts.pop(cid, None) - - def sync_pending(self, region=None): - with self._lock: - pending = self._sync_pending | self._syncing | self._sync_retry.keys() - return bool(pending) if region is None else any( - entry["id"] in pending and (profile := self._entry_profile(entry)) - and profile_region(profile) == region for entry in self._entries) - - def sync_wait(self, periodic_delay): - with self._lock: - retry_delay = min(self._sync_retry.values(), default=float("inf")) - time.monotonic() - return max(0, min(periodic_delay, retry_delay)) - - def apply_if_current(self, cm, generation, update): - """Run updates only for enabled accounts with the current credential lease.""" - with self._lock, cm._lock: - entry = next((entry for entry in self._entries if entry["cm"] is cm), None) - if entry is None or not model_policy.credential_enabled(CONFIG, entry): - return False - if not self._lease_matches(cm, generation): - self._queue_sync(entry["id"]) - return False - self.reload([cm.path], reset=False) - update() - return True - - def prune(self): - """Remove missing credential files and their session bindings.""" - with self._lock: - self._ignored_duplicates = {p for p in self._ignored_duplicates if os.path.exists(p)} - before = len(self._entries) - removed = [e for e in self._entries if not os.path.exists(e["id"])] - for entry in removed: - if self._ledger is not None: - self._ledger.remove(entry["id"]) - self.forget_credential_state(entry) - self._entries = [e for e in self._entries if e not in removed] - if len(self._entries) != before: - ids = {e["id"] for e in self._entries} - self._sync_pending.intersection_update(ids) - self._syncing.intersection_update(ids) - self._sync_retry = {cid: deadline for cid, deadline in self._sync_retry.items() if cid in ids} - self._sync_attempts = {cid: count for cid, count in self._sync_attempts.items() if cid in ids} - invalidate_model_table() - self._sticky = OrderedDict((k, v) for k, v in self._sticky.items() if v[0] in ids) - self._model_fail = {k: v for k, v in self._model_fail.items() if k[0] in ids} - if CONFIG.get("cred_pool") is self: - _publish_model_cache() - - - def find_by_uid(self, uid: str, identity: str | None = None) -> Optional[str]: - """Find a credential ID by account UID for import conflict checks.""" - with self._lock: - for e in self._entries: - if e.get("uid") == uid and (identity is None or e.get("account_key") == identity): - return e["id"] - return None - - def set_ledger(self, ledger): - """Attach the credit ledger used for expiry-aware credential selection.""" - with self._lock: - self._ledger = ledger - self.reload([Path(entry["id"]) for entry in self._entries], reset=False) - for entry in self._entries: - self._bind_entry(entry) - - def _bind_entry(self, entry): - if self._ledger is not None: - if entry.get("account_key"): - self._ledger.bind_identity(entry["id"], entry["account_key"]) - else: - self._ledger.remove(entry["id"]) - - def _adopt_cooldowns(self, entry): - """Hydrate persisted cooldowns once per identity, so the in-memory table stays authoritative.""" - identity, profile = entry.get("account_key"), entry.get("profile") - if not self._durable_identity(entry): - return # Adopt later, once this account's identity is validated. - if entry.get("cooldowns_adopted") == identity: - return - entry["cooldowns_adopted"] = identity - state = self._cooldowns.restore(identity, profile) - if not state: - return - # Deadlines are absolute and already bounded, so adopting one never extends a cooldown. - if state.get("fail_until"): - entry["fail_until"] = max(entry.get("fail_until") or 0.0, state["fail_until"]) - if state.get("reason"): - entry["last_error"] = state["reason"] - if state.get("failed_at"): - entry["last_failure_at"] = state["failed_at"] - for model, until in (state.get("models") or {}).items(): - key = (entry["id"], model) - self._model_fail[key] = max(self._model_fail.get(key, 0.0), until) - - def _durable_identity(self, entry) -> bool: - """Whether this account's identity is complete enough to key durable state. - - A hash is derived from the profile and UID, so an account with no UID still yields a - stable-looking hash shared by every other account in the same state. Durable rows - must therefore be keyed only when the identifying components are actually present. - """ - return bool(entry.get("account_key") and entry.get("profile") and entry.get("uid")) - - def _remember_credential(self, entry, reason): - """Mirror a credential circuit breaker to disk; returns whether it is durable.""" - if not self._durable_identity(entry): - return False - durable = self._cooldowns.note_credential(entry["account_key"], entry["profile"], - entry["fail_until"], reason=reason) - self._warn_storage("cooldown", durable) - return durable - - def _remember_model(self, entry, model, until): - """Mirror a model cooldown to disk; returns whether it is durable.""" - if not self._durable_identity(entry): - return False - durable = self._cooldowns.note_model(entry["account_key"], entry["profile"], model, until) - self._warn_storage("cooldown", durable) - return durable - - def _warn_storage(self, label, durable): - """Report a persistence failure operationally, rate limited so a hot path cannot flood.""" - if durable: - self._storage_warned = 0.0 - return - now = time.time() - if now - getattr(self, "_storage_warned", 0.0) < STORAGE_WARN_INTERVAL: - return - self._storage_warned = now - _log(f"[cred] {label}持久化失败({self._cooldowns.last_error});本次运行仍按内存态生效") - - def _forget_credential_cooldown(self, entry): - """Lift a persisted breaker while keeping this account's model cooldowns.""" - if self._durable_identity(entry): - outcome = self._cooldowns.clear_credential(entry["account_key"], entry["profile"]) - self._warn_storage("cooldown", outcome["durable"]) - - def cooldown_detail(self) -> list: - """Return persisted cooldown rows for diagnostics.""" - return self._cooldowns.detail() - - def cooldown_storage(self) -> dict: - """Report whether cooldown persistence is currently usable.""" - return {"available": self._cooldowns.path is not None, "path": self._cooldowns.path, - "degraded": self._cooldowns.last_error is not None, - "last_error": self._cooldowns.last_error, "rows": len(self._cooldowns.detail()), - "warning": "冷却持久化写入失败,本次运行仍按内存态生效。" if self._cooldowns.last_error else None} - - def clear_cooldowns(self, cm, model: str | None = None) -> dict: - """Lift a circuit breaker or model cooldown after a confirmed recovery or admin reset. - - Persisting cooldowns removes the old "restart the gateway to clear it" workaround, - so an explicit reset has to be able to lift one both in memory and on disk. The two - outcomes are reported separately: an in-memory reset that could not be written is - not a durable reset, and saying otherwise would hide a cooldown that comes back. - """ - with self._lock: - entry = next((e for e in self._entries if e["cm"] is cm), None) - if entry is None: - return {"changed_in_memory": False, "durable": False} - # The in-memory reset always happens, even when this account's identity is not - # complete enough to key durable state. - durable = not self._durable_identity(entry) - if model: - routed_model = _upstream_model(model, self._entry_profile(entry)) - changed = self._model_fail.pop((entry["id"], routed_model), None) is not None - if not durable: - outcome = self._cooldowns.clear_model(entry["account_key"], entry["profile"], routed_model) - durable = outcome["durable"] - self._warn_storage("cooldown", durable) - return {"changed_in_memory": changed, "durable": durable} - changed = entry["fail_until"] > time.time() - entry["fail_until"] = 0.0 - entry["last_error"] = None - if not durable: - durable = self._cooldowns.clear_credential(entry["account_key"], entry["profile"])["durable"] - self._warn_storage("cooldown", durable) - return {"changed_in_memory": changed, "durable": durable} - - def reset_cooldowns_for(self, identity: str) -> dict: - """Lift every cooldown held by one account, addressed by its public identity. - - Persisting cooldowns removed the old "restart the gateway to clear it" workaround, so an - operator needs a supported way back when a breaker or a 429 cooldown was recorded in - error, or when upstream has demonstrably recovered. The whole account is reset rather - than just its circuit breaker: a credential that is mid-429 is exactly the case an - operator is trying to unstick, and lifting only the breaker would leave it unusable. - - This is a local state change. It never refreshes a token, contacts upstream, or queues - synchronization, so a subsequent genuine failure is free to arm the cooldown again. - """ - with self._lock: - entry = next((e for e in self._entries if e.get("account_key") == identity), None) - if entry is None: - raise KeyError(identity) - now = time.time() - changed = entry["fail_until"] > now - entry["fail_until"] = 0.0 - entry["last_error"] = None - for key in [k for k in self._model_fail if k[0] == entry["id"]]: - if self._model_fail[key] > now: - changed = True - del self._model_fail[key] - # One durable write for the whole account, replacing any breaker and model rows. An - # account whose identity is incomplete still resets in memory, but owns no disk row. - durable = not self._durable_identity(entry) - if not durable: - durable = self._cooldowns.forget(entry["account_key"])["durable"] - self._warn_storage("cooldown", durable) - return {"changed_in_memory": changed, "durable": durable} - - def entries(self) -> list[dict]: - """Return credential snapshots for account maintenance.""" - with self._lock: - return [dict(e) for e in self._entries] - - def _expiry_rank(self, e: dict) -> tuple: - """Order by earliest credit expiry, placing unknown balances last.""" - exp = self._ledger.soonest_expiry_of(e["id"]) if self._ledger else None - return (exp is None, exp or 0.0) - def _rescan(self): - self.prune() - paths = find_auth_files() if self._scan else [Path(entry["id"]) for entry in self.entries()] - self.reload(paths, reset=False) - - def _healthy(self, e: dict) -> bool: - return model_policy.credential_enabled(CONFIG, e) and time.time() >= e["fail_until"] - - @staticmethod - def _entry_profile(entry): - try: - return entry["cm"].summary().get("profile", "cn-cli") - except Exception: - return None - - @classmethod - def _entry_site(cls, entry): - profile = cls._entry_profile(entry) - return profile_site(profile) if profile else None - - def _zero_balance(self, entry, profile) -> bool: - """Restrict a confirmed zero-balance account to its advertised zero-rate models.""" - balance = (self._ledger.entry(entry["id"]).get("credits") or {}) if self._ledger else {} - if not balance: - return False - try: - return (bool(balance.get("intl")) == (profile_region(profile) == "intl") - and float(balance.get("credits") or 0) <= 0) - except (TypeError, ValueError): - return False - - def _has_credit(self, entry, profile): - balance = (self._ledger.entry(entry["id"]).get("credits") or {}) if self._ledger else {} - if not balance: - return profile_region(profile) == "cn" - try: - return (bool(balance.get("intl")) == (profile_region(profile) == "intl") - and float(balance.get("credits") or 0) > 0) - except (TypeError, ValueError): - return False - - def _eligible(self, entry, model, *, region=None, profile=None, rule=None): - if not model_policy.route_allowed(CONFIG, entry, model, rule=rule): - return False - actual = self._entry_profile(entry) - profile = profile or actual - if not profile or profile != actual or not _in_region(profile, region): - return False - configured = {candidate for item in self._entries if (candidate := self._entry_profile(item)) - and _in_region(candidate, region)} - if profile not in _model_profiles(model, region, configured): - return False - if CONFIG.get("account_catalogs") is not None or CONFIG.get("model_cache") is not None: - try: - identity = entry["cm"].summary()["account_key"] - except Exception: - return False - if identity != entry.get("account_key"): - return False - account = (CONFIG.get("account_catalogs") or {}).get(identity) or {} - models = _effective_account_scope(account, "serves", model_id=_upstream_model(model, profile)) - if account.get("profile") != profile or models is None: - return False - usable = _usable_models(models) - supported = any(item["id"] == _upstream_model(model, profile) for item in usable) - cli_auto = model == "auto" and profile == "cn-cli" and bool(usable) - # Disabling the guard must not borrow another account's model capabilities. - declared = any(item["id"] == _upstream_model(model, profile) - for item in _models_for_profile(profile, configured, scope="serves")) - passthrough = (model != "auto" and not declared and not CONFIG.get("model_guard") - and len(configured) == 1) - if model and not (supported or cli_auto or passthrough): - return False - # Zero-balance accounts may only use their own advertised zero-rate models. - return (not model or self._has_credit(entry, profile) - or self._model_free(entry, model, profile=profile)) - - def _model_free(self, entry, model: str | None, *, profile=None) -> bool: - """Check whether this account advertises the model as zero-rate.""" - if not model or model == "auto": - return False - profile = profile or self._entry_profile(entry) - if not profile: - return False - accounts = CONFIG.get("account_catalogs") - if accounts is not None or CONFIG.get("model_cache") is not None: - account = (accounts or {}).get(entry.get("account_key")) or {} - if account.get("profile") != profile: - return False - return _model_free(_effective_account_scope(account, "serves", model_id=_upstream_model(model, profile)), model, profile) - return _model_free(_models_for_profile(profile, model_id=_upstream_model(model, profile)), model, profile) - - @classmethod - def _entry_endpoint(cls, e: dict) -> str | None: - """Return the credential's backend endpoint for isolated model availability checks.""" - profile = cls._entry_profile(e) - return PROFILE_ENDPOINTS.get(profile) if profile else None - - @classmethod - def _model_block_key(cls, entry): - endpoint = cls._entry_endpoint(entry) - if endpoint and cls._entry_profile(entry) in SHARED_INTL_PROFILES: - identity = entry.get("account_key") or hashlib.sha256(str(entry.get("id", "")).encode()).hexdigest() - return f"{endpoint}#account:{identity}" - return endpoint - - - def _model_servable(self, e: dict, model: str | None) -> bool: - """Check backend/model backoff, skipping the check when no model is supplied.""" - if not model: - return True - endpoint = self._model_block_key(e) - if not endpoint: - return True - return time.time() >= self._blocks.until(endpoint, _block_model(model)) - - def _model_healthy(self, e: dict, model: str | None) -> bool: - """Check this credential's model-specific 429 cooldown.""" - if not model: - return True - routed_model = _upstream_model(model, self._entry_profile(e)) - return time.time() >= self._model_fail.get((e["id"], routed_model), 0.0) - - def _evict_sticky(self): - now = time.time() - while self._sticky: - k, (_, ts) = next(iter(self._sticky.items())) - if now - ts > STICKY_TTL or len(self._sticky) > STICKY_MAX: - self._sticky.pop(k) - else: - break - - def _candidates(self, model: str | None, *, region=None, tried=()) -> list[dict]: - """Exclude tried credentials and rank candidates by zero rate and credit expiry.""" - tried = set(tried) - healthy = [entry for entry in self._entries if entry["cm"] not in tried - and self._healthy(entry) - and self._eligible(entry, model, region=region) and self._model_healthy(entry, model) - and self._model_servable(entry, model)] - if not healthy: - return [] - # Prefer zero-rate models, then earlier credit expiry; unknown balances sort last. - healthy.sort(key=lambda entry: (not self._model_free(entry, model), *self._expiry_rank(entry))) - return healthy - - @staticmethod - def _capacity_error(): - return HTTPException(status_code=503, headers={"Retry-After": "3"}, detail={"error": { - "message": "符合当前路由和免费优先策略的账号在途名额已满,请稍后重试", - "type": "service_unavailable", "code": "credential_concurrency_limit"}}) - - @staticmethod - def _capacity_key(entry): - return entry.get("account_key") or entry["id"] - - - def pick(self, skey: str | None, model: str | None = None, *, region=None, - tried=(), with_capacity=False, requirements=None) -> CredentialManager | None: - """Select a healthy sticky or round-robin credential, preferring eligible zero-rate accounts.""" - self._rescan() # Reload and prune acquire their own locks. - with self._lock: - self._evict_sticky() - candidates = self._candidates(model, region=region, tried=tried) - if not candidates: - if skey: - self._sticky.pop(skey, None) - return None - if requirements is not None: - free = self._model_free(candidates[0], model) - checked = [(entry, requirements.violations( - model_capabilities.entry_model(sys.modules[__name__], entry, model))) - for entry in candidates if self._model_free(entry, model) == free] - candidates = [entry for entry, failures in checked if not failures] - if not candidates: - raise model_capabilities.capability_error([failure for _, failures in checked for failure in failures]) - limit = CONFIG.get("max_inflight_per_account", 0) - if with_capacity and limit: - free = self._model_free(candidates[0], model) - candidates = [entry for entry in candidates if self._model_free(entry, model) == free - and self._capacity.count(self._capacity_key(entry)) < limit] - if not candidates: - raise self._capacity_error() - best = candidates[0] - free = self._model_free(best, model) - top = [e for e in candidates if self._model_free(e, model) == free - and self._expiry_rank(e) == self._expiry_rank(best)] - if skey and skey in self._sticky: - cid, _ = self._sticky[skey] - sticky = next((e for e in top if e["id"] == cid), None) - if sticky is not None: - self._sticky[skey] = (cid, time.time()) - self._sticky.move_to_end(skey) - return sticky["cm"] - e = top[self._rr[region] % len(top)] - self._rr[region] += 1 - if skey: - self._sticky[skey] = (e["id"], time.time()) - return e["cm"] - - def headers_for(self, skey: str | None, model: str | None = None, *, region=None, - with_generation=False, tried=(), with_capacity=False, requirements=None): - """Recheck identity and atomically reserve account capacity before sending.""" - capacity_race = False - capability_failures = [] - for _ in range(max(1, len(self._entries))): - cm = self.pick(skey, model, region=region, tried=tried, with_capacity=with_capacity, requirements=requirements) - if cm is None: - return None - reason = None - with cm._lock: - try: - headers = cm.get_headers() - profile = profile_for_headers(headers) - generation = cm._generation - except Exception as error: - generation, reason = cm._generation, str(error) - if reason is not None: - self.cooldown(cm, reason=reason, generation=generation) - continue - with self._lock: - self.reload([cm.path], reset=False) - entry = next((entry for entry in self._entries if entry["cm"] is cm), None) - if (entry is not None and cm._generation == generation and self._healthy(entry) - and self._eligible(entry, model, region=region, profile=profile) and self._model_healthy(entry, model)): - if requirements is not None: - failures = requirements.violations(model_capabilities.entry_model(sys.modules[__name__], entry, model)) - if failures: - capability_failures.extend(failures) - continue - if with_capacity: - lease = self._capacity.acquire(self._capacity_key(entry), - CONFIG.get("max_inflight_per_account", 0), cm, generation) - if lease is None: - capacity_race = True - continue - return lease, headers - return ((cm, generation) if with_generation else cm), headers - if capacity_race: - raise self._capacity_error() - if capability_failures: - raise model_capabilities.capability_error(capability_failures) - return None - - @staticmethod - def _lease_matches(cm, generation): - if generation is None: - return True - try: - cm._load_if_stale() - except (OSError, ValueError): - return generation == cm._generation - return generation == cm._generation - - def cooldown(self, cm: CredentialManager, reason: str = "", *, generation=None): - with self._lock, (cm._lock if generation is not None else nullcontext()): - if not self._lease_matches(cm, generation): - return - for e in self._entries: - if e["cm"] is cm: - e["fail_until"] = time.time() + CRED_COOLDOWN - e["last_error"] = sanitize_log_text(reason, 256) - e["last_failure_at"] = time.time() - self._remember_credential(e, e["last_error"]) - _log(f"[cred] 凭证熔断 {CRED_COOLDOWN}s: {Path(cm.path).name} {reason}") - - def note_status(self, cm: CredentialManager | None, status: int, - model: str | None = None, raw: bytes = b"", *, generation=None, retry_after=None): - """Apply credential-wide auth cooldowns, per-model 429 cooldowns and backend/model backoff.""" - if cm is None: - return - if status in (401, 403): - self.cooldown(cm, reason=f"backend HTTP {status}", generation=generation) - return - not_servable = _parse_not_servable(raw, status) if model else None - if not_servable: - self.note_not_servable(cm, model, code=not_servable[0], msg=not_servable[1], generation=generation) - return - if status != 429 or not model: - return - now = time.time() - if retry_after is not None: - until = now + retry_after - else: - reset = _parse_reset_time(raw) - until = reset if reset is not None and reset > now else now + MODEL_COOLDOWN - until = min(until, now + MODEL_COOLDOWN_MAX) - with self._lock, (cm._lock if generation is not None else nullcontext()): - if not self._lease_matches(cm, generation): - return - self._model_fail = {k: v for k, v in self._model_fail.items() if v > now} - for e in self._entries: - if e["cm"] is cm: - routed_model = _upstream_model(model, self._entry_profile(e)) - key = (e["id"], routed_model) - until = max(until, self._model_fail.get(key, 0.0)) - self._model_fail[key] = until - self._remember_model(e, routed_model, until) - _log(f"[cred] 模型冷却 {model} @ {Path(cm.path).name} 至 " - f"{time.strftime('%m-%d %H:%M:%S', time.localtime(until))} (HTTP 429)") - - def model_cooldown_until(self, model: str | None, *, region=None) -> float | None: - """Return the earliest reset only when all healthy credentials are cooling down.""" - if not model: - return None - with self._lock: - now = time.time() - pool = [entry for entry in self._entries if self._healthy(entry) - and self._eligible(entry, model, region=region)] - if not pool: - return None - untils = [self._model_fail.get((entry["id"], _upstream_model(model, self._entry_profile(entry))), 0.0) - for entry in pool] - if any(now >= u for u in untils): - return None - return min(untils) - - def note_not_servable(self, cm, model: str, code: str = "", msg: str = "", *, generation=None) -> float: - """Isolate international model rejection by account; retain domestic backend backoff.""" - if not model: - return 0.0 - with self._lock, (cm._lock if generation is not None else nullcontext()): - if not self._lease_matches(cm, generation): - return 0.0 - entry = next((e for e in self._entries if e["cm"] is cm), None) - endpoint = self._model_block_key(entry) if entry else None - if not endpoint: - return 0.0 - row = self._blocks.note(endpoint, _block_model(model), code=code, msg=msg) - until = float(row.get("until") or 0.0) - _log(f"[block] 模型 {model} @{endpoint} 官方回 {code}," - f"{time.strftime('%m-%d %H:%M', time.localtime(until))} 前不再派发 " - f"(第 {row.get('hits')} 次){' | ' + msg[:80] if msg else ''}") - return until - - def note_model_ok(self, cm, model: str) -> bool: - """Clear model backoff immediately after a successful backend response.""" - if not model: - return False - entry = next((e for e in self._entries if e["cm"] is cm), None) - endpoint = self._model_block_key(entry) if entry else None - return bool(endpoint) and self._blocks.clear(endpoint, _block_model(model)) - - def model_block_until(self, model: str | None, *, region=None) -> float | None: - """Return a retry time only when every potential backend has confirmed backoff.""" - if not model: - return None - now = time.time() - with self._lock: - candidates = [e for e in self._entries - if self._healthy(e) and (region is None - or _in_region(self._entry_profile(e), region)) - and model_policy.route_allowed(CONFIG, e, model)] - endpoints = {self._model_block_key(e) for e in candidates} - # Unknown catalogs remain potential sources, but cannot authorize dispatch. - capable = {self._model_block_key(e) for e in candidates - if (profile := self._entry_profile(e)) - and profile in _model_profiles(model, profile_region(profile))} - accounts = CONFIG.get("account_catalogs") - def catalog_unknown(entry): - profile = self._entry_profile(entry) - if not profile: - return False - if accounts is not None or CONFIG.get("model_cache") is not None: - account = (accounts or {}).get(entry.get("account_key")) or {} - return (account.get("profile") != profile - or _account_scope(account, "serves") is None) - return _catalog_for(profile, "serves") is None - unknown = {self._model_block_key(e) for e in candidates if catalog_unknown(e)} - endpoints.discard(None) - endpoints &= capable | unknown - if not endpoints: - return None - routed = _block_model(model) - untils = [self._blocks.until(endpoint, routed) for endpoint in endpoints] - if any(until <= now for until in untils): - return None - return max(untils) - - def model_blocks_detail(self) -> list: - """Return model backoff details for diagnostics.""" - return self._blocks.detail() - - def refresh_due(self, margin_s: int = CRED_REFRESH_MARGIN, keepalive_s: int = CRED_KEEPALIVE_S): - """Refresh expiring or idle tokens with generation-scoped failure backoff.""" - with self._lock: - entries = list(self._entries) - now = time.time() - for entry in entries: - if now < entry.get("fail_until", 0.0) or not model_policy.credential_enabled(CONFIG, entry): - continue - cm = entry["cm"] - failure = None - refreshed = keepalive_due = False - with cm._lock: - try: - summary = cm.summary() - exp = (summary.get("token_expires_at") or 0) / 1000 - last = (summary.get("last_refresh_time") or 0) / 1000 - expiry_due = bool(summary.get("token_expired") or (exp and exp - now < margin_s)) - keepalive_due = (not expiry_due and keepalive_s > 0 - and now >= entry.get("keepalive_after", 0.0) - and (last <= 0 or now - last >= keepalive_s)) - if not (expiry_due or keepalive_due): - continue - refreshed = cm.refresh_if_due(margin_s, keepalive_s if keepalive_due else 0) - entry["keepalive_after"] = 0.0 - except Exception as error: - failure = (str(error), cm._generation) - if keepalive_due: - entry["keepalive_after"] = now + CRED_KEEPALIVE_RETRY_S - if failure: - self.cooldown(cm, reason=failure[0], generation=failure[1]) - elif refreshed: - _log(f"[cred] {'每日保活刷新' if keepalive_due else '已主动刷新'}并回写: {Path(entry['id']).name}") - - def remove_file(self, name: str) -> bool: - """Share the refresh lock so an in-flight refresh cannot recreate a deleted file.""" - with self._lock: - entry = next((x for x in self._entries if os.path.basename(x["id"]) == name), None) - if entry is None: - return False - cm = entry["cm"] - try: - with cm._lock, credential_file_lock(cm.path.parent, cm.path.name): - os.unlink(entry["id"]) - cm.invalidate() - except FileNotFoundError: - pass - except OSError: - return False - self.prune() - return True - - def forget_cooldowns(self, entry): - """Drop persisted cooldowns for a credential that no longer exists.""" - identity = entry.get("account_key") - if identity: - self._warn_storage("cooldown", self._cooldowns.forget(identity)["durable"]) - - def forget_usage(self, entry): - """Drop a deleted credential's cached usage, in the store and the live aggregate.""" - snapshots = CONFIG.get("usage_snapshots") - if snapshots is not None: - snapshots.forget(entry["id"]) - accounts = CONFIG.get("usage_daily_accounts") - if isinstance(accounts, dict): - accounts.pop(entry["id"], None) - - def forget_credential_state(self, entry): - """Run both independent cleanups for a credential that is gone or replaced.""" - self.forget_cooldowns(entry) - self.forget_usage(entry) - - def first(self) -> CredentialManager | None: - with self._lock: - return self._entries[0]["cm"] if self._entries else None - - def snapshot(self) -> list[dict]: - with self._lock: - now = time.time() - out = [] - for e in self._entries: - s: dict = {"auth_file": e["id"], "healthy": self._healthy(e), - "in_flight": self._capacity.count(self._capacity_key(e)), - "max_in_flight": CONFIG.get("max_inflight_per_account", 0), - "model_cooldowns": {m: time.strftime("%m-%d %H:%M:%S", time.localtime(u)) - for (cid, m), u in self._model_fail.items() - if cid == e["id"] and u > now}, - "sticky_sessions": sum(1 for _, (cid, ts) in self._sticky.items() - if cid == e["id"] and now - ts <= STICKY_TTL)} - try: - s.update(e["cm"].summary()) - except Exception: - s["error"] = "凭据读取失败" - out.append(s) - return out - - -def _refresher_loop(pool: CredentialPool): - """Refresh idle credentials before expiry and persist renewed tokens.""" - while True: - time.sleep(60) - try: - pool.refresh_due() - except Exception as e: - _log(f"[cred] 刷新线程异常: {e}") - -CHECKIN_FIRST_DELAY = 30 # Initial check-in delay in seconds -HOUSEKEEP_INTERVAL = 3600 # Account maintenance interval in seconds - - -def _bearer_token(headers: dict) -> str: - return (headers.get("Authorization") or "").removeprefix("Bearer ").strip() - - -_HOUSEKEEP_LOCK = threading.Lock() - - -def _sync_error(pool, ledger, entry, generation, phase, error): - message = f"{phase}: {_network_error_text(error)}" - pool.apply_if_current(entry["cm"], generation, lambda: ledger.note_error(entry["id"], message)) - _log(f"[{phase}] {Path(entry['id']).name} 同步失败(保留旧数据): {message}") - - -def _buddy_context(entry, headers, consent_revision=None): - def select_model(requested=None): - from app.audit_store import safe_label - pool = CONFIG.get("cred_pool") - account = (CONFIG.get("account_catalogs") or {}).get(entry.get("account_key")) or {} - if pool is None or entry.get("profile") != "cn-work" or account.get("profile") != "cn-work": - return None - with pool._lock: - current = next((item for item in pool._entries if item["cm"] is entry["cm"] - and item.get("account_key") == entry.get("account_key")), None) - if current is None or not pool._healthy(current): - return None - candidates = [] - for item in _usable_models(_account_scope(account, "serves")): - model = item["id"] - rate = _multiplier_value(item.get("credits")) - if (not safe_label(model) or model in {".", ".."} or requested and model != requested - or rate is None or not 0 <= rate < float("inf") or "custom" in (item.get("tags") or [])): - continue - rule = model_policy.rule_for(CONFIG, model) - if (rule["upstream_id"] != model or not pool._eligible(current, model, profile="cn-work", rule=rule) - or not pool._model_healthy(current, model) or not pool._model_servable(current, model)): - continue - name = item.get("name") - candidates.append((rate, model, name if isinstance(name, str) and len(name) <= 160 else model)) - if not candidates: - return None - _, model, name = min(candidates) - return {"id": model, "name": name} - return buddy.context(CONFIG, entry, consent_revision, headers=headers, task_model=select_model) - - -def _sync_credits(pool, ledger, entry, *, checkin, failed, expected_identity=None): - if not model_policy.credential_enabled(CONFIG, entry): - return None - cm, cid = entry["cm"], entry["id"] - generation = None - try: - with cm._lock: - try: - if expected_identity is not None and cm.summary().get("account_key") != expected_identity: - failed.add(cid) - return None - headers = cm.get_headers() - finally: - generation = cm._generation - profile = profile_for_headers(headers) - identity = account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")) - if entry.get("account_key") and entry["account_key"] != identity: - failed.add(cid) - return None # A path now owned by another account must be rescheduled with its own preferences. - site = site_for_headers(headers) - token, uid, domain = _bearer_token(headers), headers.get("X-User-Id", ""), headers.get("X-Domain", "") - day = time.strftime("%Y-%m-%d") - if checkin and model_policy.credential_auto_checkin(CONFIG, entry) and not ledger.checkin_done(cid, day): - try: - def can_claim(): - return (model_policy.credential_auto_checkin(CONFIG, entry) - and pool.apply_if_current(cm, generation, lambda: None)) - result = checkin_service.perform(token, uid=uid, domain=domain, can_claim=can_claim) - if result["state"] == "cancelled": - if not pool.apply_if_current(cm, generation, lambda: None): - failed.add(cid) - return None - # A preference-only cancellation must not interrupt balance refresh. - elif not pool.apply_if_current(cm, generation, lambda: ledger.mark_checkin( - cid, day, result["ok"], result.get("code"), result["message"], state=result["state"])): - failed.add(cid) - return None - _log(f"[checkin] {Path(cid).name}: ok={result['ok']} already={result.get('already')} code={result.get('code')}") - except Exception as error: - _sync_error(pool, ledger, entry, generation, "checkin", error) - if not model_policy.credential_enabled(CONFIG, entry): - return None - if checkin and model_policy.credential_auto_travel(CONFIG, entry): - try: - def can_travel(): - return (model_policy.credential_auto_travel(CONFIG, entry) - and pool.apply_if_current(cm, generation, lambda: None)) - trip = travel.perform(token, profile_for_headers(headers), can_write=can_travel, - buddy_context=_buddy_context(entry, headers)) - if not pool.apply_if_current(cm, generation, lambda: travel.remember(ledger, cid, trip)): - failed.add(cid) - return None - buddy.daily_warning(CONFIG, entry.get("account_key"), entry.get("profile"), trip) - except Exception as error: - _sync_error(pool, ledger, entry, generation, "travel", error) - if not model_policy.credential_enabled(CONFIG, entry): - return None - balance = credits_mod.fetch_credits(token, uid=uid, domain=domain) - if bool(balance.get("intl")) != (site == INTERNATIONAL): - raise ValueError("积分响应与凭据站点不一致") - if not pool.apply_if_current(cm, generation, lambda: ledger.update_credits(cid, balance)): - failed.add(cid) - return None - _log(f"[credits] {Path(cid).name}: 站点 {site},余额 {balance['credits']}") - profile = profile_for_headers(headers) - return entry, generation, headers, profile - except Exception as error: - failed.add(cid) - _sync_error(pool, ledger, entry, generation, "credits", error) - return None - - -def _publish_model_cache(): - """Publish account-scoped versioned catalogs, excluding ownerless shared caches.""" - cache = CONFIG.get("model_cache") - if cache is not None: - pool = CONFIG.get("cred_pool") - with pool._lock if pool is not None else nullcontext(): - accounts = {} - for entry in pool.entries() if pool is not None else []: - identity, profile = entry.get("account_key"), entry.get("profile") - if not identity or not profile: - continue - key = catalog_cache_key(profile, identity) - known = cache.age(key) is not None - accounts[identity] = {"profile": profile, - "models": cache.models(key) if known else None, - "serves": cache.serves(key) if known else None} - CONFIG["account_catalogs"] = accounts - catalogs = {profile: None for profile in PROFILE_ENDPOINTS} - for account in accounts.values(): - if account["models"] is not None: - models = catalogs[account["profile"]] - if models is None: - models = catalogs[account["profile"]] = [] - models.extend(account["models"]) - CONFIG["model_catalogs"] = catalogs - CONFIG["models_remote"], CONFIG["models_intl"] = catalogs["cn-cli"], catalogs["intl-cli"] - invalidate_model_table() - - -def _sync_model_catalogs(pool, ledger, refs, failed): - cache = CONFIG.get("model_cache") - if cache is None: - return - for entry, generation, headers, profile in refs.values(): - identity = account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")) - if not model_policy.credential_enabled(CONFIG, entry): - continue - key = catalog_cache_key(profile, identity) - if cache.fresh(key) and not entry.get("catalog_dirty"): - continue - try: - scopes = credits_mod.fetch_model_scopes( - _bearer_token(headers), domain=headers.get("X-Domain", ""), - uid=headers.get("X-User-Id", ""), enterprise_id=headers.get("X-Enterprise-Id", "")) - models, serves = scopes["picker"], scopes["account"] - def publish(): - cache.put(key, models, serves=serves) - for current in pool._entries: - if current["cm"] is entry["cm"]: - current["catalog_dirty"] = False - if pool.apply_if_current(entry["cm"], generation, publish): - _log(f"[models] {profile} 模型表已刷新: 选择器 {len(models)} 个," - f"账号根表 {len(serves)} 个") - else: - failed.add(entry["id"]) - except Exception as error: - failed.add(entry["id"]) - _sync_error(pool, ledger, entry, generation, "models", error) - _publish_model_cache() - - -def _sync_usage(pool, entries=None, expected_identity=None): - """Refresh usage during maintenance with independent per-account snapshots.""" - accounts = CONFIG.get("usage_daily_accounts") - if not isinstance(accounts, dict): - accounts = CONFIG["usage_daily_accounts"] = {} - targets = pool.entries() if entries is None else entries - target_ids = {entry["id"] for entry in targets} - previous_stale = set((CONFIG.get("usage_daily") or {}).get("stale_accounts", [])) - stale = {entry["id"] for entry in pool.entries() - if entry["id"] not in target_ids and Path(entry["id"]).name in previous_stale} - for entry in targets: - if not model_policy.credential_enabled(CONFIG, entry): - continue - try: - cm = entry["cm"] - with cm._lock: - if expected_identity is not None and cm.summary().get("account_key") != expected_identity: - stale.add(entry["id"]) - continue - headers = cm.get_headers() - generation = cm._generation - site = site_for_headers(headers) - usage = credits_mod.fetch_request_usage(_bearer_token(headers), uid=headers.get("X-User-Id", ""), - domain=headers.get("X-Domain", "")) - def store(): - accounts[entry["id"]] = {"identity": entry.get("account_key"), "site": site, - "by_day": usage["by_day"], - "total_credits": round(usage["total_credits"], 2), - "requests": usage["requests"], - "partial": bool(usage.get("partial")), - "fetched_at": time.time()} - snapshots = CONFIG.get("usage_snapshots") - if snapshots is not None and entry.get("account_key") and entry.get("uid"): - # Only a validated identity may key durable state. An absent flag defaults to - # False, but a present value is passed through raw so the store's own strict - # validator governs: coercing it here with bool() would mask a malformed value - # and persist it as a legitimate flag. - snapshots.store(entry["id"], entry["account_key"], site, usage, - partial=usage.get("partial", False)) - if not pool.apply_if_current(cm, generation, store): - stale.add(entry["id"]) - except Exception as error: - stale.add(entry["id"]) - _log(f"[usage] {Path(entry['id']).name} 明细拉取失败(保留其上次成功快照): {_network_error_text(error)}") - _publish_usage_daily(pool, stale) - return stale & target_ids - - -def _usage_row_expired(snap) -> bool: - """Whether a cached row is too old to keep showing.""" - return time.time() - float(snap.get("fetched_at") or 0.0) > USAGE_CACHE_MAX_AGE_S - - -def _adopt_cached_usage(pool, snapshots, accounts): - """Seed the aggregate from the cache, keeping only rows that still belong to their account. - - Ownership is checked here, at the moment of use, rather than only when the cache was - written: a path can be reused between restarts, and a row hydrated on an earlier pass - would otherwise keep showing the previous account's usage indefinitely. - """ - # The live map is mutated by forget_usage under the pool lock, so the reads below need that - # same lock to stay consistent with a concurrent credential deletion or replacement. - with pool._lock: - cached = snapshots.accounts() - if not cached and not accounts: - return # Nothing cached and nothing live: leave the pool untouched. - # Only an account whose identity components are all present may own durable usage: an - # account_key is a hash of the profile and UID, so an account with no UID would otherwise - # share one identity with every other such account. - owners = {entry["id"]: entry.get("account_key") for entry in pool.entries() - if entry.get("account_key") and entry.get("profile") and entry.get("uid")} - # Drop any row whose recorded owner no longer matches the account at that path. This covers - # rows hydrated by an earlier pass as well as live rows, so a reused path cannot keep - # displaying the previous account's usage. Iterate over a copy of the items, so a row is - # never looked up in the live map after the keys were copied. - for path, row in list(accounts.items()): - if owners.get(path) != row.get("identity"): - accounts.pop(path, None) - for path, row in cached.items(): - identity = row.get("identity") - if owners.get(path) != identity: - snapshots.forget(path) # The path moved on; the cached row is not ours to show. - continue - if _usage_row_expired(row): - snapshots.forget(path) - continue - # A live snapshot from this run always wins; the cache only fills a gap. - if path in accounts: - continue - # A restored snapshot is stale until a refresh confirms it; clearing that is per account. - accounts[path] = dict(row) - accounts[path]["stale"] = True - CONFIG["usage_daily_accounts"] = accounts - - -def _publish_usage_daily(pool, stale=()): - """Aggregate enabled accounts' usage, retaining failed snapshots with explicit staleness.""" - # CredentialPool.forget_usage removes entries from the live map under the pool lock, so the - # snapshot taken here and the pruning inside the loop must hold that same lock. Copying the - # keys and then indexing the live dictionary would otherwise raise KeyError when a credential - # is deleted or replaced while a usage-maintenance pass is running. - with pool._lock: - accounts = CONFIG.get("usage_daily_accounts") - if not isinstance(accounts, dict): - accounts = {} - CONFIG["usage_daily_accounts"] = accounts - # Seed from the on-disk cache so the dashboard is not blank until the first refresh. - snapshots = CONFIG.get("usage_snapshots") - if snapshots is not None: - _adopt_cached_usage(pool, snapshots, accounts) - enabled = {e["id"] for e in pool.entries() if model_policy.credential_enabled(CONFIG, e)} - # Aggregate from a consistent local copy of the live map rather than indexing it per row. - rows = dict(accounts) - by_day, groups = {}, {} - used, count = 0.0, 0 - partial = False # Upstream paging hid additional usage for some account. - stale_out = set() # Accounts whose displayed figures are not from this run. - newest = 0.0 - for cred_id, snap in rows.items(): - if cred_id not in enabled: - continue - if _usage_row_expired(snap): - accounts.pop(cred_id, None) # Drop it from the live map, not just from this sum. - continue - site = snap.get("site") or "domestic" - group = groups.setdefault(site, {"by_day": {}, "total_credits": 0.0, "requests": 0}) - for day, models in (snap.get("by_day") or {}).items(): - total_day = by_day.setdefault(day, {}) - site_day = group["by_day"].setdefault(day, {}) - for model, credit in models.items(): - total_day[model] = round(total_day.get(model, 0.0) + credit, 6) - site_day[model] = round(site_day.get(model, 0.0) + credit, 6) - group["total_credits"] += float(snap.get("total_credits") or 0) - group["requests"] += int(snap.get("requests") or 0) - used += float(snap.get("total_credits") or 0) - count += int(snap.get("requests") or 0) - newest = max(newest, float(snap.get("fetched_at") or 0)) - # `partial` is the aggregate "this view is incomplete" flag the dashboard shows, so a - # failed or not-yet-refreshed account sets it; `stale_accounts` names which ones. - if snap.get("partial") or snap.get("stale"): - partial = True - if snap.get("stale"): - stale_out.add(Path(cred_id).name) - # Failed enabled accounts must remain visible even without a prior snapshot. - for cred_id in stale: - if cred_id in enabled: - partial = True - stale_out.add(Path(cred_id).name) - # A zero timestamp preserves quota-difference fallback when no usage snapshot exists. - for group in groups.values(): - group["total_credits"] = round(group["total_credits"], 2) - out = {"by_day": by_day, "groups": groups, "total_credits": round(used, 2), - "requests": count, "fetched_at": newest, "partial": partial} - if stale_out: - out["stale_accounts"] = sorted(stale_out) - CONFIG["usage_daily"] = out - _log(f"[usage] 明细已同步: {count} 请求 / {used:.2f} credits" - + (f" | {len(stale_out)} 账号同步失败" if stale_out else "")) - - -def _housekeep_once(pool: CredentialPool, ledger, *, pending_only=False): - """Serialize generation-scoped maintenance; new credentials only trigger balance/catalog reads.""" - if credits_mod is None: - return - with _HOUSEKEEP_LOCK: - pool._rescan() - # Drop expired deadlines and aged usage so neither table grows without bound; a failed - # prune is reported rather than silently leaving stale rows on disk. - pool._cooldowns.prune() - pool._warn_storage("cooldown", pool._cooldowns.last_error is None) - if CONFIG.get("usage_snapshots") is not None: - CONFIG["usage_snapshots"].prune() - ids = pool.begin_sync(all_entries=not pending_only) - failed = set() - try: - refs = {} - for entry in pool.entries(): - if entry["id"] not in ids: - continue - result = _sync_credits(pool, ledger, entry, checkin=not pending_only, failed=failed) - if result is not None: - refs[entry["id"]] = result - _sync_model_catalogs(pool, ledger, refs, failed) - if not pending_only: - _sync_usage(pool) - except Exception: - failed.update(ids) - raise - finally: - pool.end_sync(ids, failed) - - -def _housekeeper_loop(pool: CredentialPool, ledger) -> None: - """Wake for new credentials, retry failed work with backoff, and run hourly maintenance.""" - next_full = time.monotonic() + CHECKIN_FIRST_DELAY - while True: - pool._sync_event.wait(pool.sync_wait(next_full - time.monotonic())) - full_due = time.monotonic() >= next_full - try: - _housekeep_once(pool, ledger, pending_only=not full_due) - except Exception as error: - _log(f"[housekeeper] 循环异常: {_network_error_text(error)}") - if full_due: - next_full = time.monotonic() + HOUSEKEEP_INTERVAL - - -# --------------------------------------------------------------------------- -# Model inventory -# --------------------------------------------------------------------------- - -# Fallback models for legacy domestic deployments without a cloud catalog. -DEFAULT_MODELS = [ - "hy4-preview", "hy4-preview-x", - "hy3", "hy3-x", - "deepseek-v4-pro", "deepseek-v4-flash", "deepseek-v4.1-flash", "deepseek-v3-2-volc", - "glm-5.3", "glm-5.3-flash", "glm-5.2", "glm-5.1", "glm-5.0", "glm-5.0-turbo", - "glm-5v-turbo", "glm-4.7", "glm-4.6", "glm-4.6v", - "minimax-m3", "minimax-m2.7", "minimax-m2.5", - "kimi-k3-1", "kimi-k2.7", "kimi-k2.6", "kimi-k2.5", "kimi-k2-thinking", - "hunyuan-chat", "default", - "auto", # Backend-selected model alias -] - - -# Supported optional upstream request fields. -PASSTHROUGH_BODY_KEYS = { - "model", "messages", "tools", "tool_choice", "temperature", - "max_tokens", "max_completion_tokens", "top_p", "stream", - "stream_options", "stop", "presence_penalty", "frequency_penalty", - "n", "response_format", "seed", "user", "reasoning_effort", "prompt_cache_key", - "verbosity", "reasoning_summary", "parallel_tool_calls", -} - -# --------------------------------------------------------------------------- -# FastAPI application -# --------------------------------------------------------------------------- - -app = FastAPI(title="codebuddy2api", version=APP_VERSION, lifespan=inference_lifespan) -app.add_middleware(InferenceResourcesMiddleware, config=lambda: CONFIG) - -# Anthropic error types: https://platform.claude.com/docs/en/api/errors -_ANTHROPIC_ERROR_TYPES = { - "auth_error": "authentication_error", - "rate_limit_error": "rate_limit_error", - "invalid_request_error": "invalid_request_error", - "not_found_error": "not_found_error", - "upstream_error": "api_error", -} - - -@app.exception_handler(HTTPException) -async def _protocol_http_exception(request: Request, exc: HTTPException): - """Shape /v1 errors for the client protocol; retain FastAPI defaults elsewhere.""" - path = request.url.path - if not path.startswith("/v1/"): - return await _default_http_exception_handler(request, exc) - detail = exc.detail - err = detail.get("error") if isinstance(detail, dict) else None - if not isinstance(err, dict): - err = {"message": str(detail), "type": "error"} - message = str(err.get("message") or "") - if path.startswith("/v1/messages"): - # Preserve upstream business codes in Anthropic error envelopes. - etype = _ANTHROPIC_ERROR_TYPES.get(str(err.get("type") or "")) - if exc.status_code == 404: - etype = "not_found_error" # Anthropic's required type for HTTP 404. - elif etype is None: - etype = "api_error" if exc.status_code >= 500 else "invalid_request_error" - error_obj = {**err, "type": etype, "message": message} # Retain structured error fields. - return JSONResponse({"type": "error", "error": error_obj}, - status_code=exc.status_code, headers=exc.headers) - # OpenAI uses a top-level error object. - body = {"error": {**err, "message": message}} - return JSONResponse(body, status_code=exc.status_code, headers=exc.headers) -CONFIG: dict = {"api_key": "", "cred": None, "log_path": None, "ledger": None, - "admin_csrf": True, # Startup-only Origin/CSRF policy - "admin_allowed_origins": "", # Extra trusted management Origins (hot) - "models_remote": None, # Domestic cloud model inventory - "models_intl": None, # Eligible international model inventory - "model_cache": None, # Versioned catalog cache - "model_catalogs": {}, # Display-only merged product catalogs - "account_catalogs": None, # Account-scoped models/serves; None enables legacy embedding - "trial_ledger": None, - "model_guard": True, # Reject models absent from authorized catalogs - "max_images": 16, "image_policy": "truncate", - "max_request_bytes": 32 * 1024 * 1024, "log_body_limit": 65536, - "max_inbound_bytes": 64 * 1024 * 1024, - "max_collect_bytes": 8 * 1024 * 1024, "max_concurrent": 64, - "upstream_keepalive": False, "max_inflight_per_account": 0, - "request_context_mode": "legacy", - "stream_mode": "compatible", "model_capability_guard": True, - "responses_projection_mode": "balanced", "responses_projection_max_bytes": 40000, - "failover_max": 0, # Credential failovers allowed before the first response byte - "retry_write_timeout": False, # Opt-in replay after incomplete writes - "usage_daily": None, # Usage aggregated by date and model - "usage_daily_accounts": None, # Independent per-account usage snapshots - "usage_snapshots": None, # On-disk cache of the per-account snapshots - "credit_price_cny": None, "credit_price_usd": None, "usd_rate": None, - "desensitize": False, "no_compact": False, "keep_tool_metadata": False} # None prices use module defaults. - -# In-memory OAuth sessions do not survive restarts. -_OAUTH = auth_oauth.OAuthManager(user_agent=USER_AGENT) - - -# --------------------------------------------------------------------------- -# Runtime audit events -# --------------------------------------------------------------------------- - - - -def _log(msg: str): - """Persist allowlisted runtime events in SQLite, never free-form text or secrets.""" - audit = CONFIG.get("audit_store") - component = re.match(r"\[(cred|credits|models|usage|trial|checkin|housekeeper)\]", msg) - if audit is not None and component: - # Persist event codes, not free-form lines which may contain upstream data. - code = "cooldown" if "熔断" in msg or "冷却" in msg else "failure" if "失败" in msg or "异常" in msg else "updated" - audit.event("runtime", component.group(1), {"code": code}) - - -def _log_json(label: str, value): - if CONFIG.get("log_path") and CONFIG.get("log_body_limit", 65536): - _log(f"{label}\n{format_log_body(value, CONFIG.get('log_body_limit', 65536))}") - - -def _log_text_body(label: str, text: str): - if CONFIG.get("log_path") and CONFIG.get("log_body_limit", 65536): - _log(f"{label}\n{sanitize_log_text(text, CONFIG.get('log_body_limit', 65536))}") - - - - -def _truncate(s: str, n: int = 80) -> str: - s = str(s).replace("\n", " ").strip() - return s[:n] + ("…" if len(s) > n else "") - - -def _check_auth(authorization: Optional[str], x_api_key: Optional[str]): - require_api_key(CONFIG["api_key"], authorization, x_api_key) - - -def _check_admin_auth(authorization: Optional[str], x_api_key: Optional[str]): - if not CONFIG.get("api_key"): - raise HTTPException(status_code=503, detail={"error": {"message": "管理接口需要配置 API key", - "type": "management_locked"}}) - _check_auth(authorization, x_api_key) - - -def _cred_for(payload: dict, model: str | None = None, *, region=None, tried=(), requirements=None): - """Select a fresh credential lease and headers, excluding tried accounts; report unavailable capacity.""" - context = current_context() - raw_key = context.session_key if context is not None and context.scoped else session_key(payload) - skey = f"{region}:{raw_key}" if raw_key and region is not None else raw_key - skey = model_policy.sticky_scope(CONFIG, skey, model) - pool = CONFIG.get("cred_pool") - if pool is not None: - resources = request_resources.get() - picked = pool.headers_for(skey, model, region=region, with_generation=True, tried=tried, - with_capacity=resources is not None, requirements=requirements) - if picked is None: - until = pool.model_cooldown_until(model, region=region) - if until: - t = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(until)) - raise HTTPException(status_code=429, - headers={"Retry-After": str(max(1, math.ceil(until - time.time())))}, - detail={"error": { - "message": f"模型 {model} 额度冷却中(全部凭证),预计 {t} 重置后恢复", - "type": "rate_limit_error"}}) - blocked = pool.model_block_until(model, region=region) - if blocked: - # Report confirmed unsupported models as HTTP 404. - t = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(blocked)) - raise HTTPException(status_code=404, detail={"error": { - "message": f"模型 {model} 在当前所有已登录后端均不可用(官方回 service info not found)," - f"预计 {t} 后重试;请改用 /v1/models 列出的模型", - "type": "invalid_request_error", "code": "model_not_found", - "param": "model"}}) - raise HTTPException(status_code=503, headers={"Retry-After": "3" if _catalog_pending(region) else "30"}, - detail={"error": {"message": "无可用凭证(未登录、目录/额度未就绪或全部熔断)", - "type": "auth_error"}}) - cm, headers = picked - if resources is not None: - resources.add(cm) - else: - cm = CONFIG["cred"] - if cm is None or cm in {_cred_manager(item) for item in tried}: - raise HTTPException(status_code=503, detail={"error": {"message": "未找到登录凭据,请先在桌面端登录 CodeBuddy/WorkBuddy", "type": "auth_error"}}) - with cm._lock: - headers = cm.get_headers() - cm = (cm, cm._generation) - profile = profile_for_headers(headers) - if not _in_region(profile, region): - raise HTTPException(status_code=503, detail={"error": {"message": "未找到指定地域凭据", "type": "auth_error"}}) - if context is not None and context.scoped: - identity = account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")) - headers["X-Conversation-ID"] = context.conversation_id(profile, identity) - else: - headers.update(_dynamic_request_headers(f"{profile}:{skey}" if skey else None)) - return cm, headers - - -def _route_chat(payload, body, rid, *, tried=()): - """Validate account capabilities and derive each routed body from canonical input.""" - context = current_context() - enabled = context.capability_guard if context is not None else CONFIG.get("model_capability_guard", True) - cred = None - try: - requirements = (model_capabilities.Requirements.from_request( - body, payload, context.protocol if context is not None else "chat") if enabled else None) - cred, headers = _cred_for(payload, body.get("model"), tried=tried, requirements=requirements) - profile = profile_for_headers(headers) - routed_model = _upstream_model(body.get("model"), profile) - if requirements is not None and CONFIG.get("cred_pool") is None: - entry = {"profile": profile, "account_key": account_key( - profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id"))} - failures = requirements.violations(model_capabilities.entry_model(sys.modules[__name__], entry, body.get("model"))) - if failures: - raise model_capabilities.capability_error(failures) - canonical = body - if routed_model != body.get("model"): - body = {**body, "model": routed_model} - body, merged_runs, merged_messages = merge_intl_user_images(body, profile) - if body is not canonical: - _guard_request_size(body) - if merged_runs: - observe_attempt("intl_image_merge", merged_runs=merged_runs, merged_messages=merged_messages) - url = chat_url_for_headers(headers) - observe_route(public_model=payload.get("model", "auto"), upstream_model=routed_model, - profile=profile, credential=account_key(profile, headers.get("X-User-Id"), - headers.get("X-Enterprise-Id"))) - _log(f"[{rid}] ROUTE | region={profile_region(profile)} | profile={profile} | model={routed_model} | url={url}") - return body, cred, headers, url - except HTTPException as error: - release_credential(cred) - detail = error.detail.get("error", {}) if isinstance(error.detail, dict) else {} - observe_attempt("request_preflight_rejected", code=detail.get("code")) - raise - - -def _note_cred_model_ok(cred, model: str | None) -> None: - """Clear backend/model backoff after an upstream HTTP 200 response.""" - pool = CONFIG.get("cred_pool") - if pool is not None and cred is not None and model: - cm = cred[0] if isinstance(cred, tuple) else cred - pool.note_model_ok(cm, model) - - -def _note_cred_status(cred, status: int, model: str | None = None, raw: bytes = b"", *, retry_after=None): - """Record generation-scoped authentication, quota and unsupported-model failures.""" - pool = CONFIG.get("cred_pool") - if pool is not None and cred is not None: - cm, generation = cred if isinstance(cred, tuple) else (cred, None) - pool.note_status(cm, status, model=model, raw=raw, generation=generation, retry_after=retry_after) - -@app.get("/health") -def health(): - """Return public liveness without accessing or exposing credentials.""" - return {"status": "ok"} - - -@app.get("/admin/credentials") -def admin_list_credentials(authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Return account expiry, health and session-binding metadata.""" - _check_admin_auth(authorization, x_api_key) - pool = CONFIG.get("cred_pool") - if CONFIG.get("management") is not None: - return {"credentials": CONFIG["management"].admin_credential_inventory()} - return {"credentials": pool.snapshot() if pool else []} - - -class CredentialConflictError(CredentialFileError): - """Signal that another credential file already owns the account.""" - - -def _store_credential(directory: Path, name: str, content: bytes, uid: str, *, replace_identity=True, - replace_existing=True) -> Path: - """Serialize imports and logins with background refresh and standalone CLI writes.""" - pool = CONFIG.get("cred_pool") - identity = _credential_identity(json.loads(content)) - target = directory.resolve() / name - with pool._lock if pool is not None else nullcontext(): - cm = None - if pool is not None: - pool._rescan() - holder = pool.find_by_uid(uid, identity) - if holder and holder != str(target): - raise CredentialConflictError("该账号已在凭证池中") - entry = next((entry for entry in pool._entries if entry["id"] == str(target)), None) - cm = entry["cm"] if entry else None - elif CONFIG.get("cred") is not None and CONFIG["cred"].path.resolve() == target: - cm = CONFIG["cred"] - with cm._lock if cm is not None else nullcontext(): - with credential_file_lock(directory, name): - if not replace_existing and target.exists(): - raise CredentialConflictError("文件已存在,需明确允许替换") - if not replace_identity and target.exists() and _cred_identity(target) != identity: - raise CredentialConflictError("OAuth 不可覆盖其他产品或账号的凭据") - target = atomic_write_credential(directory, name, content) - if pool is not None: - pool.reload([target]) - elif cm is not None: - cm.invalidate() - return target - - -@app.post("/admin/credentials") -async def admin_add_credential(request: Request, - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Validate and atomically import credentials from the controlled directory.""" - _check_admin_auth(authorization, x_api_key) - try: - body = await request.json() - except (ValueError, UnicodeError): - raise HTTPException(status_code=400, detail={"error": {"message": "请求体必须是 JSON 对象", "type": "invalid_request_error"}}) from None - if not isinstance(body, dict): - raise HTTPException(status_code=400, detail={"error": {"message": "请求体必须是 JSON 对象", "type": "invalid_request_error"}}) - dst_dir = managed_auth_dir().resolve() - import_dir = Path(os.environ.get("CODEBUDDY_IMPORT_DIR") or dst_dir / "imports") - try: - name, content = read_import_file(import_dir, body.get("path")) - cred_data = auth_oauth.loads_strict(content.decode("utf-8")) - src_uid, verr = auth_oauth.validate_cred_data(cred_data) - if verr: - raise CredentialFileError("凭据格式或站点校验失败") - if (not isinstance(cred_data.get("account") or {}, dict) - or not isinstance(cred_data["auth"].get("expiresAt", 0), (int, float))): - raise CredentialFileError("凭据账号或过期时间格式无效") - # Normalize token aliases before persistence. - content = json.dumps(auth_oauth.normalize_cred_data(cred_data), ensure_ascii=False).encode("utf-8") - except CredentialFileError: - raise HTTPException(status_code=400, detail={"error": {"message": "凭据文件不符合导入要求", "type": "invalid_request_error"}}) from None - except (ValueError, UnicodeError, RecursionError): - raise HTTPException(status_code=400, detail={"error": {"message": "凭据必须是有效的 UTF-8 JSON 对象", "type": "invalid_request_error"}}) from None - except OSError: - raise HTTPException(status_code=400, detail={"error": {"message": "导入目录或文件不可读", "type": "invalid_request_error"}}) from None - try: - dst = _store_credential(dst_dir, name, content, src_uid) - except CredentialConflictError: - raise HTTPException(status_code=409, detail={"error": {"message": "该账号已在池中,请使用同文件名更新或先移除旧凭据", "type": "invalid_request_error"}}) from None - except CredentialFileError: - raise HTTPException(status_code=400, detail={"error": {"message": "凭据文件名或保存目标不符合要求", "type": "invalid_request_error"}}) from None - except OSError: - raise HTTPException(status_code=500, detail={"error": {"message": "凭据保存失败", "type": "server_error"}}) from None - return {"imported": str(dst)} - - -@app.delete("/admin/credentials/{name}") -def admin_del_credential(name: str, - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Delete the named .info file and remove its credential from the pool.""" - _check_admin_auth(authorization, x_api_key) - pool = CONFIG.get("cred_pool") - if CONFIG.get("management") is not None: - CONFIG["management"].admin_delete_guard(os.path.basename(name)) - if pool is None or not pool.remove_file(os.path.basename(name)): - raise HTTPException(status_code=404, detail={"error": {"message": f"凭据不在池中: {name}", "type": "invalid_request_error"}}) - return {"removed": os.path.basename(name)} - - - -def _save_oauth_credential(cred: dict) -> Path: - """Update credentials by product, account and tenant without crossing identities.""" - uid, error = auth_oauth.validate_cred_data(cred) - if error: - raise CredentialFileError("凭据格式或站点校验失败") - dst_dir = managed_auth_dir() - identity = _credential_identity(cred) - profile = profile_for_auth(cred["auth"]) - target = next((f for f in sorted(dst_dir.glob("*.info")) if _cred_identity(f) == identity), None) - existing = None - if target is not None: - try: - existing = json.loads(target.read_text(encoding="utf-8")) - except (OSError, ValueError): - pass - name = target.name if target is not None else f"{uid}.info" - if target is None and (dst_dir / name).exists(): - name = f"{uid}-{profile}.info" - if (dst_dir / name).exists(): - name = f"{uid}-{profile}-{identity}.info" - if (dst_dir / name).exists(): - raise CredentialConflictError("OAuth 保存目标已被其他身份占用") - cred = auth_oauth.merge_existing_accounts(cred, existing) - return _store_credential( - dst_dir, name, json.dumps(cred, ensure_ascii=False, indent=2).encode("utf-8"), uid, - replace_identity=False) - - -@app.post("/admin/oauth/start") -def admin_oauth_start(site: str = "cn", - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Start OAuth and return the browser authorization URL.""" - _check_admin_auth(authorization, x_api_key) - try: - return _OAUTH.start(site=site) - except ValueError as e: - raise HTTPException(status_code=400, detail={"error": {"message": str(e), "type": "invalid_request_error"}}) - except Exception as e: - raise HTTPException(status_code=502, detail={"error": {"message": f"发起失败: {e}", "type": "upstream_error"}}) - - -@app.get("/admin/oauth/poll") -def admin_oauth_poll(login_id: str = "", - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Poll OAuth and persist completed logins with pool reload.""" - _check_admin_auth(authorization, x_api_key) - try: - r = _OAUTH.poll(login_id) - except Exception as e: - raise HTTPException(status_code=502, detail={"error": {"message": f"轮询失败: {e}", "type": "upstream_error"}}) - if not r.get("done"): - return {"done": False} - cred = r.get("cred") - if r.get("error") or not cred: - return {"done": True, "error": r.get("error") or "登录失败"} - uid = r["uid"] - try: - target = _save_oauth_credential(cred) - except CredentialFileError: - return {"done": True, "error": "凭据格式、站点或保存目标不符合要求"} - except OSError: - raise HTTPException(status_code=500, detail={"error": {"message": "凭据保存失败", "type": "server_error"}}) from None - _log(f"[oauth] 无感登录已入库: {r.get('nickname') or uid} ({uid}) -> {target.name}") - return {"done": True, "uid": uid, "nickname": r.get("nickname") or "", "imported": str(target)} - -@app.get("/admin/credits") -def admin_credits(authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Return cached balances, credit expiry segments and check-in status.""" - _check_admin_auth(authorization, x_api_key) - ledger = CONFIG.get("ledger") - return {"credits": ledger.snapshot() if ledger else {}} - - -@app.get("/admin/model-blocks") -def admin_model_blocks(authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Return unsupported backend/model pairs and their retry deadlines.""" - _check_admin_auth(authorization, x_api_key) - pool = CONFIG.get("cred_pool") - return {"model_blocks": pool.model_blocks_detail() if pool is not None else []} - - -@app.post("/admin/checkin") -def admin_checkin(authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Run daily-idempotent manual check-in without implicitly syncing balances or usage.""" - _check_admin_auth(authorization, x_api_key) - return _admin_credential_action("checkin") - - -def _admin_credential_action(action, identity=None, *, consent_revision=None): - from app.credential_actions import run - return run(sys.modules[__name__], action, identity, consent_revision=consent_revision) - - -@app.post("/admin/sync") -def admin_sync(authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - _check_admin_auth(authorization, x_api_key) - return _admin_credential_action("sync") - - -@app.post("/admin/credentials/{identity}/{action}") -async def admin_credential_action(identity: str, action: str, request: Request, - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - _check_admin_auth(authorization, x_api_key) - from app.admin_api import _body - try: - body = await _body(request, 4096, allow_empty=True) - except ValueError: - raise HTTPException(400, "请求体必须是有效 JSON 对象") from None - if body: - if action != "travel": - raise HTTPException(400, "该操作不接受请求体") - if (set(body) != {"confirm_buddy", "agreement_revision"} - or body["confirm_buddy"] is not True or not isinstance(body["agreement_revision"], str)): - raise HTTPException(400, "首领确认参数无效") - return await run_in_threadpool(_admin_credential_action, action, identity, - consent_revision=body.get("agreement_revision")) - - -# --------------------------------------------------------------------------- -# OpenAI-compatible billing estimates -# --------------------------------------------------------------------------- - -def _billing_totals() -> dict: - """Convert regional balances independently, using official usage or a quota-difference fallback.""" - empty_grp = {"remaining": 0.0, "used_by_quota": 0.0, "soonest_expiry": None} - ledger = CONFIG.get("ledger") - snap = ledger.snapshot() if ledger else {} - price_cny = CONFIG.get("credit_price_cny") - price_usd = CONFIG.get("credit_price_usd") - rate = CONFIG.get("usd_rate") - if price_cny is None: - price_cny = credits_mod.CREDIT_PRICE_CNY if credits_mod else 0.014 - if price_usd is None: - price_usd = credits_mod.CREDIT_PRICE_USD if credits_mod else 0.03 - if rate is None: - rate = credits_mod.USD_RATE_CNY if credits_mod else 7.15 - agg = (credits_mod.aggregate_credits(snap) if credits_mod else - {"remaining": 0.0, "used_by_quota": 0.0, "soonest_expiry": None, - "groups": {"domestic": dict(empty_grp), "international": dict(empty_grp)}}) - cache = CONFIG.get("usage_daily") or {} - detail = bool(cache.get("fetched_at")) - detail_groups = cache.get("groups") or {} - per_usd = {"domestic": price_cny / rate, "international": price_usd} - remaining = used = rem_usd = used_usd = rem_cny = 0.0 - groups_out: dict = {} - for grp, g in (agg.get("groups") or {}).items(): - r = float(g.get("remaining") or 0) - gd = detail_groups.get(grp) - u = (float(gd.get("total_credits") or 0) if (detail and gd) - else float(g.get("used_by_quota") or 0)) # Fall back to the group's quota difference. - unit = per_usd.get(grp, 0.0) - remaining += r - used += u - rem_usd += r * unit - used_usd += u * unit - rem_cny += r * (unit * rate if grp == "international" else price_cny) - groups_out[grp] = {"credits_remaining": round(r, 2), "credits_used": round(u, 2), - "balance_usd": round(r * unit, 4), - "price_usd_per_credit": round(unit, 6)} - return {"remaining": round(remaining, 2), "used": round(used, 2), - "quota": round(remaining + used, 2), - "remaining_usd": round(rem_usd, 4), "used_usd": round(used_usd, 4), - "quota_usd": round(rem_usd + used_usd, 4), - "remaining_cny": round(rem_cny, 4), - "soonest_expiry": agg.get("soonest_expiry"), - "price_cny": price_cny, "price_usd": price_usd, "rate": rate, - "used_source": "official_usage_detail" if detail else "quota_delta", - # Expose incomplete pagination or failed account synchronization. - "partial": bool(agg.get("partial") or cache.get("partial")), - "groups": groups_out, "by_day": cache.get("by_day") or {}} - - -@app.get("/v1/dashboard/billing/subscription") -def billing_subscription(authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Expose subscription estimates with balance equal to the hard limit minus usage.""" - _check_auth(authorization, x_api_key) - t = _billing_totals() - limit = t["quota_usd"] - return { - "object": "billing_subscription", - "has_payment_method": True, "canceled": False, "canceled_at": None, "delinquent": None, - # Conservatively use the earliest credit expiry as the access deadline. - "access_until": int(t["soonest_expiry"] or (time.time() + 30 * 86400)), - "soft_limit": int(limit * 100), "hard_limit": int(limit * 100), - "soft_limit_usd": limit, "hard_limit_usd": limit, "system_hard_limit_usd": limit, - "plan": {"title": f"CodeBuddy Credits (CN {t['price_cny']:g} CNY/credit · " - f"INTL {t['price_usd']:g} USD/credit)"}, - # Include total and regional estimates as optional response fields. - "codebuddy_credits_remaining": t["remaining"], - "codebuddy_credits_used": t["used"], - "codebuddy_balance_usd": t["remaining_usd"], - "codebuddy_balance_cny": t["remaining_cny"], - "codebuddy_sites": t["groups"], - # Expose incomplete balance or usage data to callers. - "codebuddy_partial": t["partial"], - **({"codebuddy_stale_accounts": stale} if (stale := (CONFIG.get("usage_daily") or {}).get("stale_accounts")) else {}), - } - - -@app.get("/v1/dashboard/billing/usage") -def billing_usage(start_date: Optional[str] = None, end_date: Optional[str] = None, - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Return estimated usage in cents, with daily model costs for the last 30 days.""" - _check_auth(authorization, x_api_key) - t = _billing_totals() - # Convert each region at its own rate before combining daily and total usage. - cents = {"domestic": t["price_cny"] / t["rate"] * 100, "international": t["price_usd"] * 100} - detail = CONFIG.get("usage_daily") or {} - priced: dict = {} - for site, group in (detail.get("groups") or {}).items(): - unit = cents.get(site) - if unit is None: - continue - for day, models in (group.get("by_day") or {}).items(): - slot = priced.setdefault(day, {}) - for model, credit in models.items(): - slot[model] = slot.get(model, 0.0) + float(credit) * unit - daily = [] - for day in sorted(priced): - if start_date and day < start_date: - continue - if end_date and day > end_date: - continue - items = [{"name": m, "cost": round(c, 4)} - for m, c in sorted(priced[day].items()) if c > 0] - try: - ts = int(time.mktime(time.strptime(day, "%Y-%m-%d"))) - except ValueError: - ts = 0 - daily.append({"timestamp": ts, "line_items": items}) - if start_date or end_date: # Sum only the requested interval. - total_cents = round(sum(sum(i["cost"] for i in d["line_items"]) for d in daily), 2) - else: # Preserve the subscription balance identity. - total_cents = round(t["used_usd"] * 100, 2) - out = {"object": "list", "total_usage": total_cents, "daily_costs": daily} - if t.get("partial"): - out["partial"] = True - if detail.get("stale_accounts"): - out["stale_accounts"] = detail["stale_accounts"] - return out - - -# Prefer cloud catalogs over static fallback models. -_MODEL_TABLE_TTL = 60.0 # Model snapshot lifetime in seconds -_model_table_cache: dict = {} - - -def invalidate_model_table() -> None: - """Invalidate the public model snapshot after catalog changes.""" - global _model_table_cache - _model_table_cache = {} - - -def _catalog_for(profile: str, scope: str = "models", *, model_id=None): - """Return a profile catalog within the requested account scope.""" - accounts = CONFIG.get("account_catalogs") - if accounts is not None or CONFIG.get("model_cache") is not None: - pool = CONFIG.get("cred_pool") - models = None - for entry in pool.entries() if pool is not None else []: - if entry.get("profile") != profile: - continue - account = (accounts or {}).get(entry.get("account_key")) or {} - items = _effective_account_scope(account, scope, model_id=model_id) - if account.get("profile") == profile and items is not None: - if models is None: - models = [] - models.extend(items) - return models - catalogs = CONFIG.get("model_catalogs") or {} - if profile in catalogs: - return catalogs[profile] - legacy = {"cn-cli": "models_remote", "intl-cli": "models_intl"} - return CONFIG.get(legacy[profile]) if profile in legacy else None - - -def _in_region(profile: str, region: str | None) -> bool: - return region is None or profile_region(profile) == region - - -def _configured_profiles(region: str | None) -> set[str]: - pool = CONFIG.get("cred_pool") - if pool is not None: - return {profile for entry in pool.entries() if (profile := pool._entry_profile(entry)) - and model_policy.credential_enabled(CONFIG, entry) and _in_region(profile, region)} - cm = CONFIG.get("cred") - if cm is not None: - profile = cm.summary()["profile"] - return {profile} if _in_region(profile, region) else set() - known = {profile for profile in PROFILE_ENDPOINTS - if _in_region(profile, region) and _catalog_for(profile) is not None} - return known or ({"intl-cli"} if region == "intl" else {"cn-cli"}) - - -def _usable_models(models): - return [model for model in models or [] if model.get("id") and model.get("supportsToolCall") - and not model.get("disabled")] - - -def _account_scope(account: dict, scope: str = "models") -> list[dict] | None: - """Merge root candidates into selector models while retaining selector metadata.""" - picker = account.get("models") - if scope == "models" or picker is None: - return picker - seen = {item.get("id") for item in picker} - # Root entries supply missing names without replacing selector metadata. - return picker + [item for item in account.get("serves") or [] if item.get("id") not in seen] - - -def _shared_catalog_sources(scope="models"): - accounts, pool = CONFIG.get("account_catalogs"), CONFIG.get("cred_pool") - if accounts is not None or CONFIG.get("model_cache") is not None: - return [(entry["profile"], _account_scope(account, scope)) - for entry in (pool.entries() if pool is not None else []) - if entry.get("profile") in SHARED_INTL_PROFILES and model_policy.credential_enabled(CONFIG, entry) - and (account := (accounts or {}).get(entry.get("account_key"))) - and account.get("profile") == entry["profile"] and account.get("models") is not None] - configured = _configured_profiles(None) - return [(profile, _catalog_for(profile, scope)) for profile in sorted(configured & SHARED_INTL_PROFILES)] - - -def _effective_account_scope(account, scope="models", *, model_id=None): - native = _account_scope(account, scope) - profile = account.get("profile") - if profile not in SHARED_INTL_PROFILES or native is None: - return native - return share_models(native, profile, _shared_catalog_sources(scope), model_id=model_id) - - - -def _models_for_profile(profile: str, configured=None, *, scope: str = "models", model_id=None) -> list[dict]: - models = _catalog_for(profile, scope, model_id=model_id) - if models is None: - # Static fallback is limited to legacy domestic CLI deployments. - configured = _configured_profiles(profile_region(profile)) if configured is None else configured - return ([{"id": name, "supportsToolCall": True} for name in DEFAULT_MODELS] - if CONFIG.get("model_cache") is None and CONFIG.get("account_catalogs") is None - and profile == "cn-cli" and configured <= {"cn-cli"} else []) - if profile in SHARED_INTL_PROFILES and CONFIG.get("account_catalogs") is None and CONFIG.get("model_cache") is None: - models = share_models(models, profile, _shared_catalog_sources(scope), model_id=model_id) - return _usable_models(models) - - -def _upstream_model(model: str | None, profile: str) -> str | None: - return "default-model" if model == "auto" and profile_region(profile) == "intl" else model - - -def _free_multiplier(credits) -> bool: - """Check whether the account's catalog explicitly declares a zero credit rate.""" - if not isinstance(credits, str): - return False - match = re.fullmatch(r"x\s*0(?:\.0+)?\s*(?:credits?)?", credits.strip(), re.IGNORECASE) - return match is not None - - -def _multiplier_value(credits): - """Parse an official model rate, returning None for missing or unknown formats.""" - if not isinstance(credits, str): - return None - match = re.fullmatch(r"x\s*([0-9]+(?:\.[0-9]+)?)\s*(?:credits?)?", credits.strip(), re.IGNORECASE) - if match is None: - return None - try: - return float(match.group(1)) - except ValueError: - return None - - -def _model_free(models, model: str | None, profile: str) -> bool: - """Require an explicit zero-rate entry in this account's model catalog.""" - if not model: - return False - routed = _upstream_model(model, profile) - return any(item.get("id") == routed and _free_multiplier(item.get("credits")) - for item in models or []) - - -def _model_profiles(model: str | None, region: str | None = None, configured=None) -> set[str]: - configured = _configured_profiles(region) if configured is None else configured - profiles = {profile for profile in PROFILE_ENDPOINTS if _in_region(profile, region)} - if not model: - return profiles - supported = {profile for profile in profiles - if any(item["id"] == _upstream_model(model, profile) - for item in _models_for_profile(profile, configured, scope="serves", model_id=_upstream_model(model, profile)))} - if model == "auto" and region == "cn": - # WorkBuddy uses its advertised Auto; legacy CLI defaults remain separate. - if "cn-work" in configured and "cn-work" in supported: - return {"cn-work"} - if "cn-cli" in configured and _models_for_profile("cn-cli", configured): - return {"cn-cli"} - if model == "auto" and region is None and "cn-cli" in configured and _models_for_profile("cn-cli", configured): - supported.add("cn-cli") - if model != "auto" and not supported and not CONFIG.get("model_guard") and len(configured) == 1: - return configured - return supported - - -def _catalog_pending(region: str | None = None) -> bool: - pool = CONFIG.get("cred_pool") - if CONFIG.get("model_cache") is None and CONFIG.get("account_catalogs") is None: - return False - if pool is not None and pool.sync_pending(region): - return True - return not any(_catalog_for(profile) is not None for profile in _configured_profiles(region)) - - -def _profile_has_credits(profile: str) -> bool: - pool = CONFIG.get("cred_pool") - if pool is None: - if profile_region(profile) == "cn": - return True - ledger = CONFIG.get("ledger") - return bool(ledger and any((entry.get("credits") or {}).get("intl") - and float((entry.get("credits") or {}).get("credits") or 0) > 0 - for entry in ledger.snapshot().values())) - return any(pool._entry_profile(entry) == profile and pool._has_credit(entry, profile) - for entry in pool.entries() if model_policy.credential_enabled(CONFIG, entry)) - - -def current_models(region: str | None = None) -> list[str]: - """Merge models eligible for current accounts without region-specific client URLs.""" - pool = CONFIG.get("cred_pool") - if pool is not None: - pool._rescan() - with pool._lock if pool is not None else nullcontext(): - configured = _configured_profiles(region) - out, has_auto = [], False - auto_profiles = _model_profiles("auto", region, configured) - if pool is not None and (CONFIG.get("account_catalogs") is not None or CONFIG.get("model_cache") is not None): - accounts = CONFIG.get("account_catalogs") or {} - for entry in pool.entries(): - if not model_policy.credential_enabled(CONFIG, entry): - continue - profile = entry.get("profile") - if not profile or not _in_region(profile, region): - continue - # Publish only advertised zero-rate models for empty accounts. - zero = pool._zero_balance(entry, profile) - if not zero and not pool._has_credit(entry, profile): - continue - account = accounts.get(entry.get("account_key")) or {} - if account.get("profile") != profile: - continue - models = _usable_models(_effective_account_scope(account, "serves")) - if zero: - models = [model for model in models if _free_multiplier(model.get("credits"))] - out.extend(model["id"] for model in models) - if profile in auto_profiles and models: - has_auto |= profile == "cn-cli" or any(model["id"] == _upstream_model("auto", profile) for model in models) - else: - for profile in sorted(configured): - entries = ([entry for entry in pool.entries() if pool._entry_profile(entry) == profile] - if pool is not None else []) - # Empty products may publish only their advertised zero-rate models. - zero_only = bool(entries) and all(pool._zero_balance(entry, profile) for entry in entries) - if _profile_has_credits(profile) or zero_only: - models = _models_for_profile(profile, configured) - if zero_only: - models = [model for model in models if _free_multiplier(model.get("credits"))] - out.extend(model["id"] for model in models) - has_auto |= bool(models) and profile in auto_profiles - if has_auto: - out.append("auto") - return list(dict.fromkeys(out)) - - -def current_model_details(region: str | None = None) -> list[dict]: - """Return public model rates with per-profile values and the minimum eligible rate.""" - pool = CONFIG.get("cred_pool") - if pool is not None: - pool._rescan() - details: dict[str, dict] = {} - declarations = {} - for name in current_models(region): - details[name] = {"id": name, "credits": None, "credits_by_profile": {}} - if pool is None: - return [{**item, **model_capabilities.describe_models(())} for item in details.values()] - with pool._lock: - def record(profile: str, item: dict, *, zero: bool) -> None: - name = item.get("id") - if name not in details: - return - if zero and not _free_multiplier(item.get("credits")): - return # Empty accounts cannot supply paid model rates. - declarations.setdefault(name, []).append((profile, item)) - value = _multiplier_value(item.get("credits")) - if value is None: - return - details[name]["credits_by_profile"][profile] = value - best = details[name]["credits"] - details[name]["credits"] = value if best is None else min(best, value) - - if CONFIG.get("account_catalogs") is not None or CONFIG.get("model_cache") is not None: - accounts = CONFIG.get("account_catalogs") or {} - for entry in pool.entries(): - if not model_policy.credential_enabled(CONFIG, entry): - continue - profile = entry.get("profile") - if not profile or not _in_region(profile, region): - continue - zero = pool._zero_balance(entry, profile) - if not zero and not pool._has_credit(entry, profile): - continue - account = accounts.get(entry.get("account_key")) or {} - if account.get("profile") != profile: - continue - for item in _usable_models(_effective_account_scope(account, "serves")): - record(profile, item, zero=zero) - else: - configured = _configured_profiles(region) - for profile in sorted(configured): - entries = [entry for entry in pool.entries() if pool._entry_profile(entry) == profile] - zero_only = bool(entries) and all(pool._zero_balance(entry, profile) for entry in entries) - if not (_profile_has_credits(profile) or zero_only): - continue - for item in _models_for_profile(profile, configured): - record(profile, item, zero=zero_only) - return [{**item, **model_capabilities.describe_models(declarations.get(item["id"], []))} for item in details.values()] - - -def _client_wants_stream(payload: dict) -> bool: - """Default stream to false and reject non-Boolean values.""" - value = payload.get("stream", False) - if not isinstance(value, bool): - raise HTTPException(status_code=400, detail={"error": { - "message": "stream must be a boolean", "type": "invalid_request_error", "param": "stream"}}) - return value - - -def _prepare_payload(payload, field="messages") -> dict: - """Apply request-wide image limits before adaptation, logging and credential selection.""" - if not isinstance(payload, dict): - raise HTTPException(status_code=400, detail={"error": { - "message": "请求体必须是 JSON 对象", "type": "invalid_request_error"}}) - try: - prepared, stats = apply_image_policy( - payload, field=field, max_images=CONFIG["max_images"], policy=CONFIG["image_policy"]) - except ImageLimitError as error: - _log(f"[limit] 图片超限,拒绝请求 | count={error.count} | limit={error.limit}") - raise HTTPException(status_code=413, detail={"error": { - "message": str(error), "type": "invalid_request_error", "param": field, - "code": "too_many_images", "image_count": error.count, "max_images": error.limit}}) from None - if stats["dropped"]: - _log(f"[limit] 保留最新图片 | count={stats['count']} | retained={stats['retained']} | dropped={stats['dropped']}") - return prepared - - -def _normalize_tool_choice(body): - """Map named tool choice to a single required tool for string-only upstream selection.""" - choice = body.get("tool_choice") - if not isinstance(choice, dict): - return - function = choice.get("function", choice) - name = function.get("name") if isinstance(function, dict) else None - tools = body.get("tools") - matches = [tool for tool in tools if isinstance(tool, dict) and tool.get("type") == "function" - and isinstance(tool.get("function"), dict) and tool["function"].get("name") == name] if isinstance(tools, list) else [] - if choice.get("type") != "function" or not isinstance(name, str) or not name.strip() or len(matches) != 1: - raise HTTPException(status_code=400, detail={"error": {"message": "tool_choice must name exactly one declared function", - "type": "invalid_request_error", "param": "tool_choice"}}) - body["tools"], body["tool_choice"] = matches, "required" - - -def _bind_request_session(payload, body): - context = current_context() - if context is not None and context.scoped: - try: - context.bind_session(payload, body.get("messages")) - except SessionIdentifierError as error: - raise HTTPException(status_code=400, detail={"error": {"message": str(error), - "type": "invalid_request_error", "param": "session_id"}}) from None - except (ValueError, TypeError, UnicodeError, RecursionError): - raise HTTPException(status_code=400, detail={"error": {"message": "invalid session input", - "type": "invalid_request_error"}}) from None - - -def _request_id(): - context = current_context() - return context.request_id if context is not None else uuid.uuid4().hex - - -def _prepare_chat_body(body: dict, *, region=None, session_payload=None) -> dict: - """Normalize models, system messages, streaming, desensitization and payload budgets.""" - if session_payload is not None: - _bind_request_session(session_payload, body) - body = dict(body) - body["model"] = model_policy.resolve(CONFIG, body.get("model", "auto")) - guard_model(body["model"], region=region, resolved=True) - messages = body.get("messages") - if not isinstance(messages, list) or not messages or any(not isinstance(message, dict) for message in messages): - raise HTTPException(status_code=400, detail={"error": { - "message": "messages must be a non-empty array of objects", "type": "invalid_request_error"}}) - # Keep error paths tied to caller positions when the upstream system message moves. - message_indices = list(range(len(messages))) - # Upstreams reject developer roles; copy them as system messages without changing content. - messages = [ - dict(message, role="system") if message.get("role") == "developer" else message - for message in messages - ] - if messages[0].get("role") != "system": - system_index = next((index for index, message in enumerate(messages) if message.get("role") == "system"), None) - if system_index is None: - messages = [{"role": "system", "content": "You are a helpful assistant."}, *messages] - message_indices.insert(0, None) - else: - messages = [messages[system_index], *messages[:system_index], *messages[system_index + 1:]] - message_indices.insert(0, message_indices.pop(system_index)) - body["messages"] = normalize_chat_messages(messages, message_indices=message_indices) - _normalize_tool_choice(body) - body["stream"] = True - body.setdefault("stream_options", {"include_usage": True}) - body = _chat_body_desensitize(body) - _guard_request_size(body) - return body - - -def _guard_request_size(body: dict) -> int: - """Validate and measure upstream JSON bytes without truncating text or tool arguments.""" - size = 0 - limit = CONFIG["max_request_bytes"] - try: - for part in json.JSONEncoder(ensure_ascii=False, separators=(",", ":"), allow_nan=False).iterencode(body): - size += len(part.encode("utf-8")) - if size > limit: - _log(f"[limit] 请求体超限,拒绝请求 | limit_bytes={limit}") - raise HTTPException(status_code=413, detail={"error": { - "message": f"处理后的请求体超过网关上限 {limit} 字节,请缩短历史或压缩图片", - "type": "invalid_request_error", "code": "request_too_large", "max_bytes": limit}}) - except (ValueError, UnicodeError) as error: - raise HTTPException(status_code=400, detail={"error": { - "message": "请求体包含无法序列化的 JSON 值", "type": "invalid_request_error"}}) from None - return size - - -def guard_model(name: str, *, region=None, resolved=False) -> None: - """Reject unauthorized models and route only through accounts with confirmed support.""" - if not isinstance(name, str) or not name.strip(): - raise HTTPException(status_code=400, detail={"error": { - "message": "model must be a non-empty string", "type": "invalid_request_error", "param": "model"}}) - if not resolved: - name = model_policy.resolve(CONFIG, name) - model_policy.check_resolved(CONFIG, name) - if not CONFIG.get("model_guard"): - return - pool = CONFIG.get("cred_pool") - if pool is not None: - pool._rescan() - if _model_profiles(name, region): - return - if _catalog_pending(region): - raise HTTPException(status_code=503, headers={"Retry-After": "3"}, detail={"error": { - "message": "模型目录正在同步,请稍后重试", "type": "service_unavailable", "code": "catalog_syncing"}}) - raise HTTPException(status_code=404, detail={"error": { - "message": f"The model '{name}' is not supported by this gateway. See GET /v1/models.", - "type": "invalid_request_error", "param": "model", "code": "model_not_found"}}) - - - -@app.get("/v1/models") -def list_models(authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - _check_auth(authorization, x_api_key) - data = [{"id": item["id"], "object": "model", "created": 1700000000, "owned_by": "codebuddy", - "credits": item["credits"], "credits_by_profile": item["credits_by_profile"], - **{key: item[key] for key in ("capabilities", "limits", "metadata_by_profile") if key in item}} - for item in model_policy.public_details(sys.modules[__name__])] - return {"object": "list", "data": data} - - -@app.post("/v1/chat/completions") -async def chat_completions(request: Request, - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - _check_auth(authorization, x_api_key) - # Select sticky credentials while building upstream headers. - - try: - payload = await request.json() - except Exception as e: - raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}}) - - payload = _prepare_payload(payload) - # Aggregation supports exactly one completion, so reject other n values. - n_value = payload.get("n") - if n_value is not None and not (isinstance(n_value, int) and not isinstance(n_value, bool) and n_value == 1): - raise HTTPException(status_code=400, detail={"error": { - "message": "only n=1 is supported: multiple candidates would be merged into one answer", - "type": "invalid_request_error", "param": "n"}}) - messages = payload.get("messages") or [] - if not messages: - raise HTTPException(status_code=400, detail={"error": {"message": "messages is required", "type": "invalid_request_error"}}) - - # Forward only supported request fields. - client_wants_stream = _client_wants_stream(payload) - body = {k: payload[k] for k in PASSTHROUGH_BODY_KEYS if k in payload} - body = await run_in_threadpool(_prepare_chat_body, body, session_payload=payload) - stream_policy = _snapshot_stream_policy("chat", body) - observe_stream_mode(stream_policy.mode) - - # Record request metadata. - model_name = payload.get("model", "auto") - tool_names = [t.get("function", {}).get("name") for t in (payload.get("tools") or []) - if isinstance(t, dict)] - last_user = _last_user_text(messages) - rid = _request_id() - _log(f"[{rid}] ▶ REQUEST {model_name} | stream={client_wants_stream} | msgs={len(messages)}" - + (f" | tools={tool_names}" if tool_names else "") - + (f" | last_user={_truncate(last_user, 60)!r}" if last_user else "")) - # Credential selection and refresh perform blocking file and network I/O. - prepared = body # Keep canonical input for failover policy checks. - body, cred, headers, url = await run_in_threadpool(_route_chat, payload, body, rid) - _log_json(f"[{rid}] REQUEST BODY (发往后端,预览)", body) - t0 = time.time() - - if client_wants_stream: - def attempt(routed, cred, headers, url): - return _stream_upstream(url, headers, _body_with_stream_policy(routed, stream_policy), - model_name, t0, rid, cred=cred) - return _routed_stream(payload, prepared, model_name, rid, t0, attempt, - body, cred, headers, url) - - # Aggregate upstream SSE for non-streaming clients. - async def fetch(routed, cred, headers, url): - return await _fetch_checked_chat(url, headers, routed, model_name, rid, cred, - filter_retry=True, max_collect_bytes=stream_policy.max_collect_bytes) - # Watch for disconnects across the entire failover sequence. - try: - collected = await await_or_hangup( - _routed_fetch(payload, prepared, model_name, rid, t0, fetch, - body, cred, headers, url), request) - except ClientHungUp: - return _hungup_response(rid, model_name, t0) - _log_finish(model_name, t0, collected, rid) - if CONFIG.get("control_store") is not None: - collected = {**collected, "model": model_name} - return JSONResponse(content=collected) - - -def _last_user_text(messages: list) -> str: - """Extract the latest user text for bounded log previews.""" - for m in reversed(messages): - if m.get("role") != "user": - continue - content = m.get("content", "") - if isinstance(content, list): - for blk in content: - if isinstance(blk, dict) and blk.get("type") == "text": - return str(blk.get("text", "")) - return "" - return str(content) - return "" - - -def _log_finish(model_name: str, t0: float, result: dict, rid: str = ""): - """Log request timing, finish reason, usage, tools and bounded response previews.""" - elapsed = time.time() - t0 - prefix = f"[{rid}] " if rid else "" - choice = (result.get("choices") or [{}])[0] - finish = choice.get("finish_reason") - msg = choice.get("message") or {} - detector = ContentFilterDetector() - detector.feed(msg, finish) - if detector.detected: - return # Filtered responses must not expose echoed content in previews. - tcs = msg.get("tool_calls") or [] - usage = result.get("usage") or {} - tc_names = [t.get("function", {}).get("name") for t in tcs] - _log(f"{prefix}◀ RESPONSE {model_name} | {elapsed:.1f}s | finish={finish}" - + (f" | tool_calls={tc_names}" if tc_names else "") - + f" | tokens={usage.get('total_tokens', '?')}") - _log_json(f"{prefix}RESPONSE BODY (预览)", result) - - -def _chat_completion(merged: dict) -> dict: - message = {"role": "assistant", "content": merged["content"] or None} - for key in ("reasoning_content", "refusal", "tool_calls"): - if merged.get(key): - message[key] = merged[key] - return { - "id": "chatcmpl-" + os.urandom(12).hex(), "object": "chat.completion", - "created": int(time.time()), "model": merged.get("model") or "unknown", - "choices": [{"index": 0, "message": message, - "finish_reason": merged.get("finish_reason") or ("tool_calls" if merged.get("tool_calls") else "stop")}], - "usage": merged.get("usage") or {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}, - } - - -def _completion_to_merged(result: dict) -> dict: - choice = result["choices"][0] - return {**choice["message"], "finish_reason": choice["finish_reason"], - "model": result.get("model"), "usage": result.get("usage")} - - -async def _collect_stream(response: httpx.Response, *, accumulator=None) -> dict: - """Collect content, reasoning and tools while validating stream completion.""" - accumulator = accumulator if accumulator is not None else ChatSSEAccumulator() - async for line in response.aiter_lines(): - accumulator.feed_line(line) - if accumulator.done: - break - return _chat_completion(accumulator.result()) - - -_TOOL_CALL_MAX_RETRY = 3 - - -def _tool_choice_satisfied(tool_calls, body): - calls = tool_calls or [] - choice = body.get("tool_choice") - if choice == "none": - return not calls - if choice == "required": - names = _declared_tool_names(body) - return bool(calls) and all(call.get("function", {}).get("name") in names for call in calls) - if isinstance(choice, dict): - name = (choice.get("function") or {}).get("name") if isinstance(choice.get("function"), dict) else None - return bool(calls) and all(call.get("function", {}).get("name") == name for call in calls) - return True - - -def _declared_tool_names(body: dict) -> frozenset[str]: - """Extract declared function names for conservative realtime identity boundaries.""" - names = set() - for tool in body.get("tools") or []: - if not isinstance(tool, dict): - continue - function = tool.get("function") - if isinstance(function, dict) and isinstance(function.get("name"), str): - names.add(function["name"]) - elif isinstance(tool.get("name"), str): - names.add(tool["name"]) - return frozenset(name for name in names if name) - - -def _tool_calls_healthy(tool_calls, body: dict | None = None, *, require_declarations=False) -> bool: - """Validate unique calls, declared names and JSON-object arguments.""" - calls = tool_calls or [] - names = set(_declared_tool_names(body)) if body is not None else None - seen_ids = set() - for tc in calls: - if not isinstance(tc, dict): - return False - call_id = tc.get("id") - if (not isinstance(call_id, str) or not call_id.strip() - or (require_declarations and call_id in seen_ids) - or (require_declarations and tc.get("type") != "function")): - return False - if not require_declarations: - seen_ids.add(call_id) - seen_ids.add(call_id) - fn = tc.get("function") or {} - if not isinstance(fn, dict): - return False - name = fn.get("name") or "" - arguments = fn.get("arguments") or "" - if not isinstance(name, str) or not name.strip() or not isinstance(arguments, str) or not arguments.strip(): - return False - if names is not None and ((require_declarations or names) and name not in names): - return False - try: - decoded = json.loads(arguments) - except (TypeError, ValueError, UnicodeError, RecursionError): - return False - if not isinstance(decoded, dict): - return False - if (require_declarations and body is not None - and body.get("parallel_tool_calls") is False and len(calls) > 1): - return False - return True - - -def _validate_realtime_tools(tool_calls, body: dict, finish_reason=None, *, filtered=False): - """Validate terminal state without regenerating or replaying the request. - - Explicit refusal/filter and length results are incomplete upstream outcomes, - not malformed tool calls. Any tool bytes that are present still have to be - structurally healthy, but a required/named choice need not be satisfied when - the upstream legitimately stopped before producing a call. - """ - if not _tool_calls_healthy(tool_calls, body, require_declarations=True): - healthy = False - elif finish_reason is None: - # Realtime success terminals require an explicit upstream completion - # marker; [DONE] alone remains compatible only on the legacy path. - healthy = False - elif filtered or finish_reason in ("length", "content_filter", "content-filter", "refusal"): - healthy = True - else: - healthy = (_tool_choice_satisfied(tool_calls, body) - and (finish_reason == "tool_calls" if tool_calls else finish_reason != "tool_calls")) - if healthy: - return - raw = {"error": {"message": "Invalid upstream tool_calls", "type": "upstream_error", - "code": "invalid_tool_calls"}} - raise UpstreamResponseError(502, json.dumps(raw).encode("utf-8")) - - -def _merge_chat_sse_text(text: str) -> dict: - """Use the shared SSE accumulator for collected text responses.""" - accumulator = ChatSSEAccumulator(max_collect_bytes=CONFIG.get("max_collect_bytes", 0)) - for line in text.splitlines(): - accumulator.feed_line(line) - return accumulator.result() - - -def _chat_result_to_sse_lines(m: dict) -> list[str]: - """Replay collected Chat output as SSE, emitting reasoning before content.""" - content = m.get("content") or "" - reasoning = m.get("reasoning_content") or "" - tcs = m.get("tool_calls") or [] - finish = m.get("finish_reason") or "stop" - model = m.get("model") - # All chunks in a completion share one stable identifier. - completion_id = "chatcmpl-" + os.urandom(12).hex() - created = int(time.time()) - - def _line(delta: dict, fr=None) -> str: - payload = {"id": completion_id, "object": "chat.completion.chunk", "created": created, - "choices": [{"index": 0, "delta": delta, "finish_reason": fr}]} - if model: - payload["model"] = model - return "data: " + json.dumps(payload, ensure_ascii=False) - - lines = [_line({"role": "assistant", "content": ""})] - for i in range(0, len(reasoning), 48): - lines.append(_line({"reasoning_content": reasoning[i:i + 48]})) - for i in range(0, len(content), 48): - lines.append(_line({"content": content[i:i + 48]})) - refusal = m.get("refusal") or "" - for i in range(0, len(refusal), 48): - lines.append(_line({"refusal": refusal[i:i + 48]})) - for i, tc in enumerate(tcs): - lines.append(_line({"tool_calls": [dict(tc, index=i)]})) - lines.append(_line({}, finish)) - if m.get("usage"): - usage_chunk = {"id": completion_id, "object": "chat.completion.chunk", "created": created, - "choices": [], "usage": m["usage"]} - if model: - usage_chunk["model"] = model - lines.append("data: " + json.dumps(usage_chunk, ensure_ascii=False)) - lines.append("data: [DONE]") - return lines - - -def _network_error_text(error: Exception) -> str: - return sanitize_log_text(f"{type(error).__name__}: {str(error).strip() or 'upstream transport failed'}", 512) - -def _public_sse_line(line, model_name): - if CONFIG.get("control_store") is not None and line.startswith("data:"): - try: - event = json.loads(line[5:].strip()) - if isinstance(event, dict) and ("model" in event or "choices" in event): - event["model"] = model_name - return "data: " + json.dumps(event, ensure_ascii=False) - except (ValueError, TypeError): - pass - return line - - -@asynccontextmanager -async def _backend_stream(url, headers, body, *, timeout=300, rid="", model_name="?"): - started, opened = time.monotonic(), False - context = current_context() - if context is not None: - context.attempt = None - - def attempt_headers(): - if context is None: - return dict(headers) - attempt = context.start_attempt() - outgoing = context.attempt_headers(headers, attempt) - profile = profile_for_headers(headers) - observe_attempt("upstream_attempt", profile=profile, upstream_model=body.get("model"), - credential=account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")), - conversation_id=outgoing.get("X-Conversation-ID"), - upstream_request_id=outgoing.get("X-Request-ID")) - return outgoing - - def retry(error): - """Record connection retries and flag possible billing after write timeouts.""" - timeout_on_write = isinstance(error, WRITE_TIMEOUT_TRANSPORT) - observe_attempt("write_timeout_retry" if timeout_on_write else "connect_retry", - error_code=type(error).__name__, - duration_ms=(time.monotonic() - started) * 1000) - _log(f"[{rid}] {'写超时重放' if timeout_on_write else '建连失败'},重试 1/1 | {model_name}" - f" | {_network_error_text(error)}{_replay_cost_note(error)}") - try: - resources = request_resources.get() - clients = resources.clients if resources is not None and CONFIG.get("upstream_keepalive") else None - async with open_backend_stream(url, headers, body, read_timeout=timeout, on_retry=retry, - retry_write_timeout=bool(CONFIG.get("retry_write_timeout")), - clients=clients, headers_for_attempt=attempt_headers) as response: - opened = True - observe_attempt("upstream_http", status_code=response.status_code, - duration_ms=(time.monotonic() - started) * 1000) - yield response - except (httpx.HTTPError, UpstreamResponseError) as error: - if not opened: - observe_attempt("transport", error_code=type(error).__name__, - duration_ms=(time.monotonic() - started) * 1000) - raise - - -def _safe_err_raw(raw: bytes, status: int) -> dict: - try: - return json.loads(raw.decode("utf-8", "replace")) - except Exception: - return {"error": {"message": raw.decode("utf-8", "replace")[:500], "type": "upstream_error", "code": status}} - - -def _check_upstream_status(status, raw, cred, model, *, headers=None): - if status != 200: - retry_after = parse_retry_after((headers or {}).get("Retry-After")) - if not is_filter_error(raw): - _note_cred_status(cred, status, model=model, raw=raw, retry_after=retry_after) - raise UpstreamHTTPError(status, raw, retry_after=retry_after) - - -def _upstream_failure(error, model_name, t0, rid): - """Normalize failure logs and payloads before endpoint-specific error wrapping.""" - if isinstance(error, UpstreamResponseError): - status, raw = error.status, error.raw - category = f"HTTP {status}" - else: - status, raw = 502, _network_error_text(error).encode("utf-8") - category = "网络错误" - if isinstance(error, UpstreamResponseError) and is_filter_error(raw): - _note_content_filter(rid, model_name, final=True) - return status, raw - else: - observe_failure(f"upstream_{status}" if isinstance(error, UpstreamResponseError) else type(error).__name__) - elapsed = time.time() - t0 if t0 else 0 - _log(f"[{rid}] ✗ {category} | {model_name} | {elapsed:.1f}s | {sanitize_log_text(raw.decode('utf-8', 'replace'), 512)}") - _log_text_body(f"[{rid}] ERROR BODY", raw.decode("utf-8", "replace")) - return status, raw - - -def _hungup_response(rid, model_name, t0): - """Finish a disconnected ASGI request with an empty 204; auditing records cancellation.""" - elapsed = time.time() - t0 if t0 else 0 - _log(f"[{rid}] ✂ 下游已断连,取消这次聚合 | {model_name} | {elapsed:.1f}s") - return Response(status_code=204) - - -async def _fetch_checked_chat(url, headers, body, model_name, rid, cred=None, *, filter_retry=False, - max_collect_bytes=None): - """Collect with bounded repair and one eligible filter retry using a frozen request budget.""" - tool_attempt = 0 - filter_retried = False - collection_limit = (CONFIG.get("max_collect_bytes", 0) if max_collect_bytes is None - else max_collect_bytes) - while True: - accumulator = ChatSSEAccumulator(max_collect_bytes=collection_limit) - rejection = None - async with _backend_stream(url, headers, body, rid=rid, model_name=model_name) as response: - if response.status_code != 200: - _check_upstream_status(response.status_code, await read_bounded_error(response), cred, body.get("model"), - headers=response.headers) - else: - _note_cred_model_ok(cred, body.get("model")) - try: - result = await _collect_stream(response, accumulator=accumulator) - except UpstreamResponseError as error: - if (not accumulator.done or not accumulator.filter_detector.detected - or accumulator.saw_output): - raise - rejection = error - result = None - - detector = accumulator.filter_detector - if detector.detected: - retry_body = body - if (filter_retry and not filter_retried and detector.retry_safe - and CONFIG.get("desensitize") and CONFIG.get("no_compact")): - retry_body = _chat_body_desensitize(body, force_compact=True) - try: - if _guard_request_size(retry_body) >= _guard_request_size(body): - retry_body = body - except HTTPException: - retry_body = body - if retry_body != body: - _note_content_filter(rid, model_name, final=False) - body = retry_body - filter_retried = True - continue - if rejection is not None: - raise rejection - _note_content_filter(rid, model_name, final=True) - - calls = result["choices"][0]["message"].get("tool_calls") - if _tool_calls_healthy(calls, body) and (detector.detected or _tool_choice_satisfied(calls, body)): - observe_usage(result.get("usage") or {}) - return result - # Content filtering must not trigger tool-repair regeneration. - budget = CONFIG.get("tool_call_max_retry", _TOOL_CALL_MAX_RETRY) - if detector.detected or not body.get("tools") or tool_attempt >= budget: - if not detector.detected and body.get("tools"): - # Account for the final failed generation before returning an error. - exhausted = result.get("usage") or {} - observe_attempt("tool_args_exhausted", attempt=tool_attempt, max_attempts=budget, - total_tokens=exhausted.get("total_tokens")) - raise UpstreamResponseError(502, b"Invalid upstream tool_calls after retries") - tool_attempt += 1 - # Discarded generations still consume credits and belong in the audit trail. - discarded = result.get("usage") or {} - observe_attempt("tool_args_retry", attempt=tool_attempt, max_attempts=budget, - total_tokens=discarded.get("total_tokens")) - _log(f"[{rid}] tool_calls 损坏,重试 {tool_attempt}/{budget} | {model_name}") - -async def _chat_sse_lines(url, headers, body, model_name, t0, rid, cred=None, *, - policy=None, tracker=None, state=None): - """Yield validated Chat SSE with bounded filter detection and no streaming filter retries.""" - policy = policy or _snapshot_stream_policy("chat", body) - if policy.aggregate: - result = await _fetch_checked_chat( - url, headers, body, model_name, rid, cred, - max_collect_bytes=policy.max_collect_bytes) - for line in _chat_result_to_sse_lines(_completion_to_merged(result)): - yield line - yield "" - _log_finish(model_name, t0, result, rid) - return - if tracker is None: - tracker = ChatSSEAccumulator( - collect=False, retain_tools=policy.realtime, - budget=StreamOutputBudget(policy.max_collect_bytes)) - if state is not None: - state["tracker"] = tracker - - def completed(): - merged = tracker.result(allow_empty_filter=policy.realtime) - if policy.realtime: - _validate_realtime_tools( - merged.get("tool_calls"), body, merged.get("finish_reason"), - filtered=tracker.filter_detector.detected) - if state is not None: - state["merged"] = merged - return merged - - preview = bytearray() - # Realtime output is never copied into the retired raw-text preview log. - budget = 0 if policy.realtime else ( - CONFIG["log_body_limit"] if CONFIG.get("log_path") else 0) - try: - async with _backend_stream(url, headers, body, rid=rid, model_name=model_name) as response: - if response.status_code != 200: - _check_upstream_status(response.status_code, await read_bounded_error(response), cred, body.get("model"), - headers=response.headers) - else: - _note_cred_model_ok(cred, body.get("model")) - async for line in response.aiter_lines(): - tracker.feed_line(line) - if tracker.done or tracker.finish_reason: - completed() # Never expose a success marker before terminal validation. - remaining = budget - len(preview) - if remaining > 0: - preview.extend((line[:remaining] + "\n").encode("utf-8")[:remaining]) - yield line - if tracker.done: - yield "" - break - merged = completed() - finally: - if tracker.usage: - observe_usage(tracker.usage) - if tracker.filter_detector.detected: - _note_content_filter(rid, model_name, final=True) - return - _log(f"[{rid}] ◀ RESPONSE {model_name} | {time.time() - t0:.1f}s | stream finish={merged['finish_reason']}" - + f" | tokens={(merged['usage'] or {}).get('total_tokens', '?')}") - _log_text_body(f"[{rid}] RESPONSE SSE PREVIEW", preview.decode("utf-8", "replace")) - - -async def _stream_upstream(url: str, headers: dict, body: dict, - model_name: str = "?", t0: float = 0.0, rid: str = "", cred=None): - policy = body.pop(_REQUEST_POLICY_KEY, None) or _snapshot_stream_policy("chat", body) - sent = False - upstream = _chat_sse_lines(url, headers, body, model_name, t0, rid, cred, policy=policy) - try: - try: - async for line in upstream: - sent = True - yield (_public_sse_line(line, model_name) + "\n").encode("utf-8") - except (httpx.HTTPError, UpstreamResponseError) as error: - if not sent: - raise # Preserve the HTTP error while no response bytes have been sent. - status, raw = _upstream_failure(error, model_name, t0, rid) - yield _err_event(raw, status) - finally: - # The inner generator may be suspended at a yielded line when its - # consumer reports an adapter error or a downstream disconnect. - await _close_stream(upstream) - - - - -def _err_event(msg: bytes, status: int) -> bytes: - chunk = {"error": {"message": sanitize_log_text(msg.decode("utf-8", "replace"), 512), - "type": "upstream_error", "code": status}} - return f"data: {json.dumps(chunk, ensure_ascii=False)}\n\n".encode("utf-8") - - -def _cred_manager(cred): - """Unwrap a credential lease, retaining support for a standalone manager.""" - return cred[0] if isinstance(cred, tuple) else cred - - -# Retryable upstream auth, quota and gateway responses; deterministic request errors are excluded. -FAILOVER_CODES = frozenset({401, 403, 429, 502, 503, 504}) -# Connection failures occur before any request body is sent. -REPLAYABLE_TRANSPORT = (httpx.ConnectError, httpx.ConnectTimeout) -# Write-timeout replay requires explicit opt-in because partial requests may already be billed. -WRITE_TIMEOUT_TRANSPORT = (httpx.WriteTimeout,) -# Gateway timeouts may follow billable upstream work and need an explicit cost warning. -POSSIBLY_CHARGED_CODES = frozenset({502, 504}) - - -def _replay_cost_note(error) -> str: - """Label replays that may duplicate already billed work.""" - if isinstance(error, UpstreamHTTPError) and error.status in POSSIBLY_CHARGED_CODES: - return " | 上游可能已处理该请求" - if isinstance(error, WRITE_TIMEOUT_TRANSPORT): - return " | 上游可能已处理该请求(正文未写完)" - return "" - - -def _failover_safe(error, raw=b"") -> bool: - """Allow configured pre-response transport or HTTP failover, never filter or incomplete-stream replay.""" - if is_filter_error(raw): - return False - if isinstance(error, UpstreamHTTPError): - return error.status in FAILOVER_CODES - if isinstance(error, UpstreamResponseError): - return False - if isinstance(error, WRITE_TIMEOUT_TRANSPORT): - return bool(CONFIG.get("retry_write_timeout")) - return isinstance(error, REPLAYABLE_TRANSPORT) - - -class _StreamFailure(Exception): - """Carry the HTTP error and original exception from stream preflight.""" - - def __init__(self, status, raw, error=None): - self.status = status - self.raw = raw - self.error = error - self.headers = error.headers if isinstance(error, UpstreamHTTPError) else None - super().__init__(f"stream failed before first byte (HTTP {status})") - - -# Bound teardown waits by cycles so repeated cancellation cannot cause a busy loop. -TEARDOWN_GRACE_CYCLES = 100 -TEARDOWN_POLL_SECONDS = 0.01 - - -def _drain_teardown(future) -> None: - """Retrieve teardown exceptions without rethrowing them.""" - if not future.cancelled(): - future.exception() - - -async def _teardown_finished(task) -> None: - """Wait briefly for isolated cleanup, then drain it in the background without delaying cancellation.""" - for _ in range(TEARDOWN_GRACE_CYCLES): - if task.done(): - _drain_teardown(task) - return - try: - await asyncio.wait([task], timeout=TEARDOWN_POLL_SECONDS) - except asyncio.CancelledError: - pass - if not task.done(): - task.add_done_callback(_drain_teardown) - - -async def _first_segment(agen): - """Read one stream segment in a shielded task so cancellation cannot interrupt its cleanup.""" - task = asyncio.ensure_future(agen.__anext__()) - try: - return await asyncio.shield(task) - except BaseException: - task.cancel() - await _teardown_finished(task) - raise - - -async def _stream_segments(agen): - """Read cancellable segments while allowing the generator's cleanup to finish.""" - while True: - try: - yield await _first_segment(agen) - except StopAsyncIteration: - return - - -async def _preflight_stream(agen, model_name, t0, rid): - """Read the first segment before committing HTTP 200, preserving pre-response error status.""" - try: - return await _first_segment(agen) - except StopAsyncIteration: - empty = UpstreamResponseError(502, b'{"error":{"message":"upstream returned an empty stream",' - b'"type":"upstream_error","code":"empty_response"}}') - status, raw = _upstream_failure(empty, model_name, t0, rid) - raise _StreamFailure(status, raw, empty) from None - except (httpx.HTTPError, UpstreamResponseError) as error: - status, raw = _upstream_failure(error, model_name, t0, rid) - raise _StreamFailure(status, raw, error) from None - - -async def _close_stream(agen) -> None: - """Close the upstream generator in an isolated task that survives repeated cancellation.""" - if agen is None: - return - - async def close() -> None: - try: - await agen.aclose() - except Exception: - pass - - await _teardown_finished(asyncio.ensure_future(close())) - - -def _chunk_bytes(chunk, charset: str = "utf-8"): - return chunk if isinstance(chunk, (bytes, memoryview)) else chunk.encode(charset) - - -class _DeferredStreamResponse(StreamingResponse): - """Run stream preflight and failover inside ASGI disconnect monitoring before sending headers.""" - - def __init__(self, plan): - self._plan = plan # Async callable returning the upstream iterator and first segment. - super().__init__(content=(), media_type="text/event-stream", - headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}) - - async def stream_response(self, send) -> None: - agen, first = await self._plan() - try: - await send({"type": "http.response.start", "status": self.status_code, - "headers": self.raw_headers}) - await send({"type": "http.response.body", "body": _chunk_bytes(first, self.charset), - "more_body": True}) - async for chunk in _stream_segments(agen): - await send({"type": "http.response.body", "body": _chunk_bytes(chunk, self.charset), - "more_body": True}) - await send({"type": "http.response.body", "body": b"", "more_body": False}) - finally: - await _close_stream(agen) - - -def _failover_limit() -> int: - return int(CONFIG.get("failover_max") or 0) - - -async def _stream_plan(payload, canonical, model_name, rid, t0, make, routed, cred, headers, url): - """Prefetch with bounded credential failover, using canonical input to preserve routing restrictions.""" - tried = [] - recovered = None - while True: - stream = make(routed, cred, headers, url) - try: - first = await _preflight_stream(stream, model_name, t0, rid) - except _StreamFailure as failure: - await _close_stream(stream) # Release the failed upstream connection. - release_credential(cred) - recovered = observe_failure_seq() # Recover only this failure sequence. - tried.append(cred) - limit = _failover_limit() - surface = HTTPException(status_code=failure.status, headers=failure.headers, - detail=_safe_err_raw(failure.raw, failure.status)) - if limit <= 0 or len(tried) > limit or not _failover_safe(failure.error, failure.raw): - raise surface from None - try: - attempt = await run_in_threadpool(_route_chat, payload, canonical, rid, - tried={_cred_manager(item) for item in tried}) - except HTTPException: - raise surface from None # Preserve the failure when no alternative account exists. - if _cred_manager(attempt[1]) in {_cred_manager(item) for item in tried}: - raise surface from None - routed, cred, headers, url = attempt - _log(f"[{rid}] ↻ 换凭证重放 {len(tried)}/{limit} | {model_name} | 上游 HTTP " - f"{failure.status} → {profile_for_headers(headers)}" - f"{_replay_cost_note(failure.error)}") - continue # Prefetch from the replacement credential. - except BaseException: - # Release the current upstream before propagating cancellation or unexpected errors. - await _close_stream(stream) - raise - if tried: - # Preserve newer failures such as content filtering on the replacement account. - observe_recovery(recovered) - return stream, first - - -def _routed_stream(payload, canonical, model_name, rid, t0, make, routed, cred, headers, url): - """Defer stream preflight and failover until ASGI disconnect monitoring is active.""" - return _DeferredStreamResponse( - lambda: _stream_plan(payload, canonical, model_name, rid, t0, make, - routed, cred, headers, url)) - - -async def _routed_fetch(payload, canonical, model_name, rid, t0, fetch, routed, cred, headers, url): - """Apply bounded non-streaming failover while preserving canonical routing restrictions.""" - tried = [] - recovered = None - while True: - try: - collected = await fetch(routed, cred, headers, url) - if tried: - # Recover the earlier attempt without erasing a newer failure. - observe_recovery(recovered) - return collected - except (httpx.HTTPError, UpstreamResponseError) as error: - release_credential(cred) - status, raw = _upstream_failure(error, model_name, t0, rid) - recovered = observe_failure_seq() - tried.append(cred) - limit = _failover_limit() - surface = HTTPException(status_code=status, detail=_safe_err_raw(raw, status), - headers=error.headers if isinstance(error, UpstreamHTTPError) else None) - if limit <= 0 or len(tried) > limit or not _failover_safe(error, raw): - raise surface from None - try: - attempt = await run_in_threadpool(_route_chat, payload, canonical, rid, - tried={_cred_manager(item) for item in tried}) - except HTTPException: - raise surface from None - if _cred_manager(attempt[1]) in {_cred_manager(item) for item in tried}: - raise surface from None - routed, cred, headers, url = attempt - _log(f"[{rid}] ↻ 换凭证重放 {len(tried)}/{limit} | {model_name} | 上游 HTTP " - f"{status} → {profile_for_headers(headers)}" - f"{_replay_cost_note(error)}") - - -def _note_content_filter(rid, model_name, *, final): - stage = "content_filter" if final else "content_filter_retry" - observe_attempt(stage, error_code="content_filter") - if final: - observe_failure("content_filter") - action = "保留上游拒绝,不切换账号" if final else "纯审核拒绝,压缩模板重试 1/1" - _log(f"[{rid}] 内容审核拦截 | {model_name} | {action}") - - -def _chat_body_desensitize(body: dict, *, force_compact: bool = False) -> dict: - if not CONFIG.get("desensitize"): - return body - return desensitize_body( - body, - roles=("system", "developer"), - desensitize_harness_user=True, - desensitize_tools=True, - compact_harness=(force_compact or not CONFIG.get("no_compact")), - strip_tool_metadata=not CONFIG.get("keep_tool_metadata", False), - ) - - -# --------------------------------------------------------------------------- -# OpenAI Responses endpoint -# --------------------------------------------------------------------------- - -@app.post("/v1/responses") -async def create_response(request: Request, - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Serve Responses requests through the shared Chat upstream and event adapter.""" - _check_auth(authorization, x_api_key) - - try: - payload = await request.json() - except Exception as e: - raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}}) - - payload = _prepare_payload(payload, field="input") - # Reject server-side conversation references because this gateway is stateless. - for stateful in ("previous_response_id", "conversation"): - if payload.get(stateful): - raise HTTPException(status_code=400, detail={"error": { - "message": f"{stateful} is not supported: this gateway keeps no server-side response state; resubmit the full input instead", - "type": "invalid_request_error", "param": stateful}}) - # Convert Responses input to Chat format. - try: - chat_body = responses_request_to_chat(payload) - except Exception as e: - raise HTTPException(status_code=400, detail={"error": {"message": f"request conversion error: {e}", "type": "invalid_request_error"}}) - - await run_in_threadpool(_bind_request_session, payload, chat_body) - projection_mode = CONFIG.get("responses_projection_mode", "balanced") - projection_max_bytes = int(CONFIG.get("responses_projection_max_bytes", 40000)) - try: - chat_body, projection_stats = await run_in_threadpool( - project_responses_chat_body, chat_body, mode=projection_mode, - max_item_bytes=projection_max_bytes) - except UnicodeError: - raise HTTPException(status_code=400, detail={"error": { - "message": "request contains text that cannot be encoded as UTF-8", - "type": "invalid_request_error", "param": "input", - "code": "invalid_unicode"}}) from None - observe_responses_projection(projection_stats) - chat_body = await run_in_threadpool(_prepare_chat_body, chat_body) - - client_wants_stream = _client_wants_stream(payload) - stream_policy = _snapshot_stream_policy("responses", chat_body) - observe_stream_mode(stream_policy.mode) - model_name = payload.get("model", "auto") - rid = _request_id() - _log(f"[{rid}] ▶ RESPONSES {model_name} | stream={client_wants_stream} | input_items={len(payload.get('input', []))}") - _log( - f"[{rid}] ── RESPONSES PROJECTION ── " - f"mode={projection_stats.get('mode')} " - f"| msgs {projection_stats.get('original_messages')}→{projection_stats.get('projected_messages')} " - f"| chars {projection_stats.get('original_message_chars')}→{projection_stats.get('projected_message_chars')} " - f"| tools {projection_stats.get('original_tools')}→{projection_stats.get('projected_tools')} " - f"| tool_chars {projection_stats.get('original_tool_chars')}→{projection_stats.get('projected_tool_chars')} " - f"| harness_messages={projection_stats.get('harness_messages_projected', 0)} " - f"| truncated_items={projection_stats.get('truncated_items', 0)} " - f"| truncated_bytes={projection_stats.get('truncated_original_bytes', 0)}→" - f"{projection_stats.get('truncated_projected_bytes', 0)}" - ) - # Keep blocking credential selection and refresh off the event loop. - prepared = chat_body # Preserve canonical input for routing policy checks. - chat_body, cred, headers, url = await run_in_threadpool(_route_chat, payload, chat_body, rid) - _log_json(f"[{rid}] RESPONSES → CHAT BODY (预览)", chat_body) - t0 = time.time() - - if client_wants_stream: - def attempt(routed, cred, headers, url): - return _stream_responses(url, headers, _body_with_stream_policy(routed, stream_policy), - model_name, t0, rid, cred=cred) - response = _routed_stream(payload, prepared, model_name, rid, t0, attempt, - chat_body, cred, headers, url) - response.headers["X-CodeBuddy-Responses-Projection"] = projection_mode - return response - - response = await _nonstream_adapted( - url, headers, chat_body, model_name, t0, rid, cred, payload=payload, canonical=prepared, - request=request, policy=stream_policy) - response.headers["X-CodeBuddy-Responses-Projection"] = projection_mode - return response - - +#!/usr/bin/env python3 +"""Expose CodeBuddy and WorkBuddy through compatible Chat, Responses and Messages APIs.""" + +from __future__ import annotations + +import argparse +import asyncio +import hashlib +import json +import math +import os +import re +import secrets +import shutil +import sys +import threading +import time +import uuid +from collections import OrderedDict +from contextlib import asynccontextmanager, nullcontext +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Optional + +import httpx +from fastapi import FastAPI, Header, HTTPException, Request +from fastapi.exception_handlers import http_exception_handler as _default_http_exception_handler +from fastapi.responses import JSONResponse, Response, StreamingResponse +from starlette.concurrency import run_in_threadpool +import uvicorn as uvicorn # Keep the existing embedding/test hook. + +try: + from app.desensitize import desensitize_body +except ImportError: # Disable desensitization when its module is unavailable. + def desensitize_body(body, roles=("system",), desensitize_harness_user=False, + desensitize_tools=False, compact_harness=False, + strip_tool_metadata=False): + return body + +from app.adapters.responses_adapter import ( + responses_request_to_chat, + ResponsesStreamConverter, +) +from app.adapters.responses_projection import PROJECTION_MODES, project_responses_chat_body +from app.adapters.anthropic_adapter import ( + anthropic_request_to_chat, + AnthropicStreamConverter, +) + +from app import auth_oauth +from app import trial_rewards +from app import buddy, checkin as checkin_service, model_policy, travel +from app.credential_cooldowns import CredentialCooldowns +from app.model_blocks import ModelBlocks +from app.usage_snapshots import UsageSnapshots +from app.client_hangup import ClientHungUp, await_or_hangup +from app.observability import (AuditMiddleware, observe_recovery, observe_route, observe_stream_mode, + observe_responses_projection, observe_usage, observe_attempt, + observe_failure, observe_failure_seq) +from app.credential_io import (CredentialFileError, read_import_file, atomic_write_credential, + credential_file_lock) +from app.upstream_io import (ChatSSEAccumulator, StreamOutputBudget, UpstreamHTTPError, + UpstreamResponseError, open_backend_stream, parse_retry_after, + read_bounded_error) +from app.inference_resources import (AccountCapacity, InferenceResourcesMiddleware, inference_lifespan, + request_resources, release_credential) +from app.request_context import SessionIdentifierError, current_context +from app import model_capabilities +from app.message_normalization import merge_intl_user_images +from app.adapters.chat_input import normalize_chat_messages +from app.model_catalog_view import INTERNATIONAL as SHARED_INTL_PROFILES, share_models +from app.inference_auth import require_api_key +from app.admin_auth import SessionStoreError +from app.content_filter import ContentFilterDetector, is_filter_error +from app.request_limits import ImageLimitError, apply_image_policy +from app.safe_logging import format_log_body, sanitize_log_text +from app.startup import allows_open_noauth, load_startup_env, run_server +from app.site_routing import (DOMESTIC, INTERNATIONAL, PROFILE_ENDPOINTS, site_for_auth, site_for_headers, + profile_for_auth, profile_for_headers, profile_region, profile_product, + profile_site, chat_url_for_headers, refresh_url_for_auth) +from app.client_profiles import CLI_VERSION, CLI_USER_AGENT, credential_headers, catalog_cache_key, account_key +try: + from app import credits as credits_mod +except ImportError: # Disable credit maintenance when its module is unavailable. + credits_mod = None + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + +APP_VERSION = Path(__file__).with_name("VERSION").read_text(encoding="utf-8").strip() +BACKEND = "https://copilot.tencent.com" +DEFAULT_DOMAIN = "www.codebuddy.cn" +CBC_VERSION = CLI_VERSION +USER_AGENT = CLI_USER_AGENT + +_STREAM_MODES = ("compatible", "realtime") +_REQUEST_POLICY_KEY = object() # Process-local object key; never serializable or client injectable. + + +@dataclass(frozen=True) +class _StreamRequestPolicy: + mode: str + aggregate: bool + max_collect_bytes: int + + @property + def realtime(self): + return self.mode == "realtime" + + +def _snapshot_stream_policy(protocol: str, body: dict) -> _StreamRequestPolicy: + """Freeze hot streaming choices and memory policy for one request and all failovers.""" + mode = CONFIG.get("stream_mode", "compatible") + if mode not in _STREAM_MODES: + raise ValueError("invalid stream mode") + compatible_aggregate = protocol == "responses" or bool(body.get("tools")) + return _StreamRequestPolicy(mode, mode == "compatible" and compatible_aggregate, + max(0, int(CONFIG.get("max_collect_bytes", 0) or 0))) + + +def _body_with_stream_policy(body: dict, policy: _StreamRequestPolicy) -> dict: + """Attach a process-local snapshot without exposing an upstream payload override.""" + routed = dict(body) + routed[_REQUEST_POLICY_KEY] = policy + return routed + + +# --------------------------------------------------------------------------- +# Platform-specific credential directories +# --------------------------------------------------------------------------- + +def managed_auth_dir() -> Path: + """Use CODEBUDDY_AUTH_DIR when set, otherwise the project's auth directory.""" + env_dir = os.environ.get("CODEBUDDY_AUTH_DIR") + return Path(env_dir) if env_dir else Path(__file__).resolve().parent / "auth" + + +def auth_dirs() -> list[Path]: + """Locate desktop credentials used only as seed files.""" + home = Path.home() + plat = sys.platform + if plat == "darwin": + return [home / "Library" / "Application Support" / "CodeBuddyExtension" / "Data" / "Public" / "auth"] + if plat == "win32": + local = Path(os.environ.get("LOCALAPPDATA", home / "AppData" / "Local")) + return [local / "CodeBuddyExtension" / "Data" / "Public" / "auth"] + xdg = Path(os.environ.get("XDG_DATA_HOME", home / ".local" / "share")) + return [xdg / "CodeBuddyExtension" / "Data" / "Public" / "auth"] + + +def seed_credentials(): + """Seed missing managed credentials without overwriting files; skip custom auth directories.""" + if os.environ.get("CODEBUDDY_AUTH_DIR"): + return + dst_dir = managed_auth_dir() + dst_dir.mkdir(parents=True, exist_ok=True) + try: + os.chmod(dst_dir, 0o700) + except OSError: + pass + have_uids = {u for u in (_cred_identity(f) for f in dst_dir.glob("*.info")) if u} + for src_dir in auth_dirs(): + if not src_dir.is_dir(): + continue + for f in sorted(src_dir.glob("*.info")): + try: + data = json.loads(f.read_text(encoding="utf-8")) + except Exception as e: + _log(f"[cred] 种子跳过(无法解析): {f.name}: {e}") + continue + uid, verr = auth_oauth.validate_cred_data(data) + if verr: + _log(f"[cred] 种子跳过(入库校验失败:{verr}): {f.name}") + continue + identity = _credential_identity(data) + if uid and identity in have_uids: + _log(f"[cred] 种子跳过(同账号已在自管目录): {f.name}") + continue + dst = dst_dir / f.name + if not dst.exists(): + try: + shutil.copyfile(f, dst) + os.chmod(dst, 0o600) + if uid: + have_uids.add(identity) + _log(f"[cred] 已复制桌面端凭据到自管目录: {f.name}") + except OSError as e: + _log(f"[cred] 复制凭据失败 {f.name}: {e}") + + +def find_auth_files() -> list[Path]: + """Find managed .info credential files.""" + d = managed_auth_dir() + return sorted(d.glob("*.info")) if d.is_dir() else [] + + +def _cred_uid(path) -> Optional[str]: + """Read the account UID for deduplication, or return None when unavailable.""" + try: + d = json.loads(Path(path).read_text(encoding="utf-8")) + acct = d.get("account") + if not isinstance(acct, dict): + arr = d.get("accounts") + acct = arr[0] if isinstance(arr, list) and arr and isinstance(arr[0], dict) else {} + return acct.get("uid") + except Exception: + return None + +def _credential_account(data: dict) -> dict: + account = data.get("account") + if not isinstance(account, dict): + accounts = data.get("accounts") or [] + account = accounts[0] if isinstance(accounts, list) and accounts and isinstance(accounts[0], dict) else {} + return account + + +def _credential_identity(data: dict) -> str: + account = _credential_account(data) + return account_key(profile_for_auth(data.get("auth") or {}), account.get("uid"), account.get("enterpriseId")) + + +def _cred_identity(path) -> str | None: + try: + return _credential_identity(json.loads(Path(path).read_text(encoding="utf-8"))) + except (OSError, ValueError, TypeError, AttributeError): + return None + + +def find_auth_file() -> Path | None: + files = find_auth_files() + return files[0] if files else None + + +# --------------------------------------------------------------------------- +# Credential loading, refresh and persistence +# --------------------------------------------------------------------------- + +class CredentialManager: + """Load credentials and refresh expiring tokens with persistence.""" + + def __init__(self, path: Path): + self.path = path + self._lock = threading.RLock() + self._cached: dict | None = None + self._mtime = None + self._generation = 0 + + def _read_raw(self) -> dict: + with open(self.path, "r", encoding="utf-8") as f: + return json.load(f) + + def _file_version(self): + st = self.path.stat() + return st.st_dev, st.st_ino, st.st_mtime_ns, st.st_size + + def _load_if_stale(self): + """Reload changed credentials and invalidate leases held by older requests.""" + mt = self._file_version() + if self._cached is None or mt != self._mtime: + self._cached = self._read_raw() + self._mtime = mt + self._generation += 1 + + def _session(self) -> dict: + self._load_if_stale() + if self._cached is None: + raise RuntimeError(f"无法读取 auth 文件:{self.path}") + return self._cached + + def _is_expired(self) -> bool: + s = self._session() + expires_at = (s.get("auth") or {}).get("expiresAt") or 0 + # Treat tokens as expired 60 seconds early. + return time.time() * 1000 >= (expires_at - 60_000) + + def _refresh_needed(self, margin_s, keepalive_s): + summary = self.summary() + now = time.time() + exp = (summary.get("token_expires_at") or 0) / 1000 + last = (summary.get("last_refresh_time") or 0) / 1000 + return bool(summary.get("token_expired") or (exp and exp - now < margin_s) + or (keepalive_s > 0 and (last <= 0 or now - last >= keepalive_s))) + + def _refresh(self, margin_s=60, keepalive_s=0): + with self._lock: + if not self._refresh_needed(margin_s, keepalive_s): + return False + with credential_file_lock(self.path.parent, self.path.name): + if not self._refresh_needed(margin_s, keepalive_s): + return False + self._refresh_locked() + return True + + def _refresh_locked(self): + """Share the import lock so token refresh cannot overwrite a newer login.""" + s = self._session() + auth = s.get("auth") or {} + headers = self._build_headers_from(auth, _credential_account(s)) + headers["X-Refresh-Token"] = auth.get("refreshToken", "") + headers["X-Auth-Refresh-Source"] = "plugin" + url = refresh_url_for_auth(auth) + try: + with httpx.Client(timeout=15) as c: + r = c.post(url, headers=headers, json={}) + data = r.json() + except Exception as e: + raise RuntimeError(f"刷新 token 网络失败:{e}") + if data.get("code") != 0 or not data.get("data"): + raise RuntimeError(f"刷新 token 失败:{data.get('msg', data)}") + new_auth = dict(data["data"]) + if not isinstance(new_auth.get("accessToken"), str) or not new_auth["accessToken"]: + raise RuntimeError("刷新接口未返回有效的 accessToken") + new_auth["domain"] = new_auth.get("domain") or auth.get("domain") + new_auth["refreshToken"] = new_auth.get("refreshToken") or auth.get("refreshToken") + new_auth["lastRefreshTime"] = int(time.time() * 1000) + if not new_auth.get("expiresAt") and new_auth.get("expiresIn"): + new_auth["expiresAt"] = int(time.time() * 1000) + new_auth["expiresIn"] * 1000 + if not new_auth.get("refreshExpiresAt") and new_auth.get("refreshExpiresIn"): + new_auth["refreshExpiresAt"] = int(time.time() * 1000) + new_auth["refreshExpiresIn"] * 1000 + updated = dict(s, auth=new_auth) + atomic_write_credential(self.path.parent, self.path.name, + json.dumps(updated, ensure_ascii=False, indent=2).encode("utf-8")) + self._cached = updated + self._mtime = self._file_version() + self._generation += 1 + + def _build_headers_from(self, auth: dict, account: dict) -> dict: + return credential_headers(auth, account) + + def get_headers(self) -> dict: + """Return upstream headers with a current token, refreshing when necessary.""" + with self._lock: + if self._is_expired(): + self._refresh() + s = self._session() + return self._build_headers_from(s.get("auth") or {}, _credential_account(s)) + + def refresh_if_due(self, margin_s: int, keepalive_s: int) -> bool: + """Refresh under the shared foreground/background lock after rechecking expiry.""" + return self._refresh(margin_s, keepalive_s) + + def invalidate(self): + """Reload imported credentials while retaining the manager and refresh lock.""" + with self._lock: + self._cached = None + self._mtime = None + self._generation += 1 + + + def summary(self) -> dict: + with self._lock: + s = self._session() + auth = s.get("auth") or {} + acct = _credential_account(s) + profile = profile_for_auth(auth) + return { + "uid": str(acct.get("uid") or "") or None, + "account_key": _credential_identity(s), + "site": profile_site(profile), + "profile": profile, "region": profile_region(profile), "product": profile_product(profile), + "nickname": acct.get("nickname"), + "enterpriseName": acct.get("enterpriseName"), + "token_expires_at": auth.get("expiresAt", 0), + "token_expired": self._is_expired(), + "last_refresh_time": auth.get("lastRefreshTime") or 0, + } + + +STORAGE_WARN_INTERVAL = 300 # Rate limit for persistence-failure warnings +USAGE_CACHE_MAX_AGE_S = 7 * 24 * 3600 # A cached usage snapshot older than this is not shown. +STICKY_TTL = 30 * 60 # Idle session binding lifetime in seconds +STICKY_MAX = 512 # Session binding capacity +CRED_COOLDOWN = 300 # Credential cooldown in seconds +MODEL_COOLDOWN = 600 # Model cooldown when a 429 omits reset time +MODEL_COOLDOWN_MAX = 86400 # Maximum model cooldown in seconds +MODEL_SITE_BLOCK_S = 6 * 3600 # Initial unsupported-model backoff +MODEL_SITE_BLOCK_MAX_S = 24 * 3600 # Maximum unsupported-model backoff +# An unsupported model must be routed to a different backend. +MODEL_NOT_SERVABLE_CODES = frozenset({"11102"}) +_NOT_SERVABLE_MSG = re.compile(r"service info not found|model .{0,80}not (?:found|supported)", re.I) +CRED_REFRESH_MARGIN = 600 # Proactive refresh margin in seconds +CRED_KEEPALIVE_S = 24 * 3600 # Maximum idle interval before refreshing +CRED_KEEPALIVE_RETRY_S = 3600 # Keepalive retry interval, independent of expiry retries + + +def _msg_text(m: dict) -> str: + c = m.get("content") + if isinstance(c, str): + return c + if isinstance(c, list): + return "\n".join(p.get("text", "") for p in c if isinstance(p, dict)) + return "" + + +def session_key(payload: dict) -> str | None: + """Derive a stable session key from system instructions and the first user message.""" + msgs = payload.get("messages") + if not msgs: + inp = payload.get("input") # Responses API + if isinstance(inp, str): + msgs = [{"role": "user", "content": inp}] + elif isinstance(inp, list): + msgs = inp + msgs = msgs or [] + if not msgs: + return None + system = "" + for m in msgs: + if m.get("role") in ("system", "developer"): + system += _msg_text(m) + else: + break + first_user = next((_msg_text(m) for m in msgs if m.get("role") == "user"), "") + if not system and not first_user: + return None + return hashlib.sha256((system + "\x00" + first_user).encode("utf-8", "replace")).hexdigest()[:32] + + +def _parse_reset_time(raw: bytes) -> float | None: + """Parse a quota reset timestamp from a 429 response and return epoch seconds.""" + try: + text = raw.decode("utf-8", "replace") + except Exception: + return None + m = re.search(r"(\d{4}-\d{2}-\d{2})\s+(\d{2}:\d{2}:\d{2})\s*UTC\s*([+-]?\d+)", text) + if not m: + return None + try: + dt = datetime.strptime(f"{m.group(1)} {m.group(2)}", "%Y-%m-%d %H:%M:%S") + tz = timezone(timedelta(hours=int(m.group(3)))) + return dt.replace(tzinfo=tz).timestamp() + except ValueError: + return None + + + + +def _parse_not_servable(raw: bytes, status: int): + """Recognize unsupported-model errors from code/message fields, excluding incidental IDs.""" + if status not in (400, 404) or not raw: + return None + try: + payload = json.loads(raw.decode("utf-8", "replace")) + except (ValueError, UnicodeError, RecursionError): + return None + if not isinstance(payload, dict): + return None + nodes = [payload] + inner = payload.get("error") + if isinstance(inner, dict): + nodes.append(inner) + code = msg = "" + for node in nodes: + for key in ("code", "errCode", "error_code"): + value = node.get(key) + if value is not None and str(value).strip(): + code = code or str(value).strip() + for key in ("msg", "message"): + value = node.get(key) + if isinstance(value, str) and value.strip(): + msg = msg or value.strip() + if not code and not msg: + return None + if code in MODEL_NOT_SERVABLE_CODES or _NOT_SERVABLE_MSG.search(msg): + return code or "11102", msg[:200] + return None + + +def _block_model(model: str | None) -> str | None: + """Track blocked models by public ID, normalizing the international auto alias.""" + return "auto" if model == "default-model" else model + + +def _dynamic_request_headers(skey: str | None) -> dict: + """Generate upstream request IDs while keeping session IDs stable.""" + rid = secrets.token_hex(16) # X-Request-ID == X-Conversation-Message-ID + crid = secrets.token_hex(16) # X-Conversation-Request-ID == X-Root-Request-ID == trace id + span, parent = secrets.token_hex(8), secrets.token_hex(8) + if skey: + conv = str(uuid.UUID(hex=hashlib.sha256(skey.encode()).hexdigest()[:32])) + else: + conv = str(uuid.uuid4()) + return { + "X-Conversation-ID": conv, + "X-Request-ID": rid, + "X-Conversation-Message-ID": rid, + "X-Conversation-Request-ID": crid, + "X-Root-Request-ID": crid, + "X-Trace-ID": crid, + "traceparent": f"00-{crid}-{span}-01", + "b3": f"{crid}-{span}-1-{parent}", + "X-B3-TraceId": crid, + "X-B3-SpanId": span, + "X-B3-ParentSpanId": parent, + "X-B3-Sampled": "1", + } + + +class CredentialPool: + """Manage credential discovery, reloads, sticky sessions, cooldowns and refresh.""" + + def __init__(self, paths: list[Path] | None = None, scan: bool = False, + blocks_path: Path | None = None, cooldowns_path: Path | None = None, *, state_store=None): + self._lock = threading.RLock() + self._entries: list[dict] = [] # {id, cm, fail_until} + self._sticky: "OrderedDict[str, tuple[str, float]]" = OrderedDict() + self._model_fail: dict[tuple[str, str], float] = {} # Per-credential/model 429 expiry + # Keep unsupported-model backoff isolated by backend and model. + self._blocks = ModelBlocks(blocks_path, ttl_s=MODEL_SITE_BLOCK_S, max_ttl_s=MODEL_SITE_BLOCK_MAX_S, store=state_store) + # Cooldowns outlive a restart so a backend that just refused is not retried immediately. + self._cooldowns = CredentialCooldowns(cooldowns_path, store=state_store) + self._storage_warned = 0.0 # Rate limit for persistence-failure warnings + self._rr = {None: 0, "cn": 0, "intl": 0} + self._ledger = None # Prefer credits expiring sooner. + self._capacity = AccountCapacity() + self._scan = scan # Rescan credentials before selection. + self._ignored_duplicates: set[str] = set() + self._sync_pending: set[str] = set() + self._syncing: set[str] = set() + self._sync_event = threading.Event() + self._sync_retry: dict[str, float] = {} + self._sync_attempts: dict[str, int] = {} + self.reload(paths or []) + if self._scan: + self._rescan() # Discover credentials at startup. + + def reload(self, paths: list[Path], *, reset: bool = True): + """Reset authentication only for changed or imported files and schedule catalog refresh.""" + with self._lock: + by_id = {entry["id"]: entry for entry in self._entries} + have_uids = {entry["account_key"]: entry["id"] + for entry in self._entries if entry.get("uid")} + for path in paths: + cid = str(Path(path).resolve()) + if not os.path.exists(cid): + continue + entry = by_id.get(cid) + if entry is not None: + if reset: + entry["cm"].invalidate() + try: + summary = entry["cm"].summary() + except Exception: + continue # A damaged file must not block other credentials. + generation = entry["cm"]._generation + identity = summary["account_key"] + changed = reset or generation != entry.get("generation") + if changed: + old_identity = entry.get("account_key") + replaced = old_identity != identity + if replaced: + self._model_fail = {key: until for key, until in self._model_fail.items() if key[0] != cid} + self._sticky = OrderedDict((key, value) for key, value in self._sticky.items() if value[0] != cid) + # The path now belongs to another account; its cooldowns must not carry over. + self.forget_credential_state(entry) + if entry.get("uid"): + have_uids.pop(old_identity, None) + entry.update(uid=summary.get("uid"), profile=summary["profile"], site=summary["site"], + account_key=identity, generation=generation, catalog_dirty=True) + self._bind_entry(entry) + if reset or replaced: + entry.update(fail_until=0.0, keepalive_after=0.0) + if reset and not replaced: + # An explicit reload re-evaluates auth for the same account, so the + # persisted breaker is lifted; a replacement must keep the incoming + # account's own breaker, which is hydrated just below. + self._forget_credential_cooldown(entry) + self._adopt_cooldowns(entry) + if entry.get("uid"): + have_uids[identity] = cid + self._queue_sync(cid) + continue + manager = CredentialManager(Path(cid)) + try: + summary = manager.summary() + except Exception: + summary = {} + uid = summary.get("uid") + profile = summary.get("profile", "cn-cli") + identity_key = summary.get("account_key") + if uid and identity_key in have_uids: + if cid not in self._ignored_duplicates: + _log(f"[cred] 忽略重复账号凭据: {Path(cid).name}(同产品账号与 {Path(have_uids[identity_key]).name} 重复)") + self._ignored_duplicates.add(cid) + continue + entry = {"id": cid, "cm": manager, "fail_until": 0.0 if summary else time.time() + CRED_COOLDOWN, "uid": uid, + "site": summary.get("site"), "profile": profile, "generation": manager._generation, + "account_key": identity_key, "catalog_dirty": True} + self._bind_entry(entry) + self._adopt_cooldowns(entry) + self._entries.append(entry) + by_id[cid] = entry + self._ignored_duplicates.discard(cid) + if uid: + have_uids[identity_key] = cid + self._queue_sync(cid) + + def _queue_sync(self, cid): + entry = next((entry for entry in self._entries if entry["id"] == cid), None) + if entry is not None and not model_policy.credential_enabled(CONFIG, entry): + return + self._sync_pending.add(cid) + self._sync_retry.pop(cid, None) + self._sync_attempts.pop(cid, None) + self._sync_event.set() + if CONFIG.get("cred_pool") is self: + _publish_model_cache() + else: + invalidate_model_table() + + def begin_sync(self, *, all_entries=False): + """Drain refresh work and clear its wake event under the same lock.""" + with self._lock: + active = {entry["id"] for entry in self._entries if model_policy.credential_enabled(CONFIG, entry)} + self._sync_pending.intersection_update(active) + self._sync_retry = {cid: deadline for cid, deadline in self._sync_retry.items() if cid in active} + due = {cid for cid, deadline in self._sync_retry.items() if deadline <= time.monotonic()} + self._sync_pending.update(due) + ids = active if all_entries else set(self._sync_pending) + self._sync_pending.difference_update(ids) + if not self._sync_pending: + self._sync_event.clear() + self._syncing.update(ids) + return ids + + def end_sync(self, ids, failed=()): + with self._lock: + self._syncing.difference_update(ids) + present = {entry["id"] for entry in self._entries if model_policy.credential_enabled(CONFIG, entry)} + for cid in ids: + if cid in failed and cid in present and cid not in self._sync_pending: + attempt = min(self._sync_attempts.get(cid, 0) + 1, 5) + self._sync_attempts[cid] = attempt + self._sync_retry[cid] = time.monotonic() + min(60 * 2 ** (attempt - 1), 900) + else: + self._sync_retry.pop(cid, None) + self._sync_attempts.pop(cid, None) + + def sync_pending(self, region=None): + with self._lock: + pending = self._sync_pending | self._syncing | self._sync_retry.keys() + return bool(pending) if region is None else any( + entry["id"] in pending and (profile := self._entry_profile(entry)) + and profile_region(profile) == region for entry in self._entries) + + def sync_wait(self, periodic_delay): + with self._lock: + retry_delay = min(self._sync_retry.values(), default=float("inf")) - time.monotonic() + return max(0, min(periodic_delay, retry_delay)) + + def apply_if_current(self, cm, generation, update): + """Run updates only for enabled accounts with the current credential lease.""" + with self._lock, cm._lock: + entry = next((entry for entry in self._entries if entry["cm"] is cm), None) + if entry is None or not model_policy.credential_enabled(CONFIG, entry): + return False + if not self._lease_matches(cm, generation): + self._queue_sync(entry["id"]) + return False + self.reload([cm.path], reset=False) + update() + return True + + def prune(self): + """Remove missing credential files and their session bindings.""" + with self._lock: + self._ignored_duplicates = {p for p in self._ignored_duplicates if os.path.exists(p)} + before = len(self._entries) + removed = [e for e in self._entries if not os.path.exists(e["id"])] + for entry in removed: + if self._ledger is not None: + self._ledger.remove(entry["id"]) + self.forget_credential_state(entry) + self._entries = [e for e in self._entries if e not in removed] + if len(self._entries) != before: + ids = {e["id"] for e in self._entries} + self._sync_pending.intersection_update(ids) + self._syncing.intersection_update(ids) + self._sync_retry = {cid: deadline for cid, deadline in self._sync_retry.items() if cid in ids} + self._sync_attempts = {cid: count for cid, count in self._sync_attempts.items() if cid in ids} + invalidate_model_table() + self._sticky = OrderedDict((k, v) for k, v in self._sticky.items() if v[0] in ids) + self._model_fail = {k: v for k, v in self._model_fail.items() if k[0] in ids} + if CONFIG.get("cred_pool") is self: + _publish_model_cache() + + + def find_by_uid(self, uid: str, identity: str | None = None) -> Optional[str]: + """Find a credential ID by account UID for import conflict checks.""" + with self._lock: + for e in self._entries: + if e.get("uid") == uid and (identity is None or e.get("account_key") == identity): + return e["id"] + return None + + def set_ledger(self, ledger): + """Attach the credit ledger used for expiry-aware credential selection.""" + with self._lock: + self._ledger = ledger + self.reload([Path(entry["id"]) for entry in self._entries], reset=False) + for entry in self._entries: + self._bind_entry(entry) + + def _bind_entry(self, entry): + if self._ledger is not None: + if entry.get("account_key"): + self._ledger.bind_identity(entry["id"], entry["account_key"]) + else: + self._ledger.remove(entry["id"]) + + def _adopt_cooldowns(self, entry): + """Hydrate persisted cooldowns once per identity, so the in-memory table stays authoritative.""" + identity, profile = entry.get("account_key"), entry.get("profile") + if not self._durable_identity(entry): + return # Adopt later, once this account's identity is validated. + if entry.get("cooldowns_adopted") == identity: + return + entry["cooldowns_adopted"] = identity + state = self._cooldowns.restore(identity, profile) + if not state: + return + # Deadlines are absolute and already bounded, so adopting one never extends a cooldown. + if state.get("fail_until"): + entry["fail_until"] = max(entry.get("fail_until") or 0.0, state["fail_until"]) + if state.get("reason"): + entry["last_error"] = state["reason"] + if state.get("failed_at"): + entry["last_failure_at"] = state["failed_at"] + for model, until in (state.get("models") or {}).items(): + key = (entry["id"], model) + self._model_fail[key] = max(self._model_fail.get(key, 0.0), until) + + def _durable_identity(self, entry) -> bool: + """Whether this account's identity is complete enough to key durable state. + + A hash is derived from the profile and UID, so an account with no UID still yields a + stable-looking hash shared by every other account in the same state. Durable rows + must therefore be keyed only when the identifying components are actually present. + """ + return bool(entry.get("account_key") and entry.get("profile") and entry.get("uid")) + + def _remember_credential(self, entry, reason): + """Mirror a credential circuit breaker to disk; returns whether it is durable.""" + if not self._durable_identity(entry): + return False + durable = self._cooldowns.note_credential(entry["account_key"], entry["profile"], + entry["fail_until"], reason=reason) + self._warn_storage("cooldown", durable) + return durable + + def _remember_model(self, entry, model, until): + """Mirror a model cooldown to disk; returns whether it is durable.""" + if not self._durable_identity(entry): + return False + durable = self._cooldowns.note_model(entry["account_key"], entry["profile"], model, until) + self._warn_storage("cooldown", durable) + return durable + + def _warn_storage(self, label, durable): + """Report a persistence failure operationally, rate limited so a hot path cannot flood.""" + if durable: + self._storage_warned = 0.0 + return + now = time.time() + if now - getattr(self, "_storage_warned", 0.0) < STORAGE_WARN_INTERVAL: + return + self._storage_warned = now + _log(f"[cred] {label}持久化失败({self._cooldowns.last_error});本次运行仍按内存态生效") + + def _forget_credential_cooldown(self, entry): + """Lift a persisted breaker while keeping this account's model cooldowns.""" + if self._durable_identity(entry): + outcome = self._cooldowns.clear_credential(entry["account_key"], entry["profile"]) + self._warn_storage("cooldown", outcome["durable"]) + + def cooldown_detail(self) -> list: + """Return persisted cooldown rows for diagnostics.""" + return self._cooldowns.detail() + + def cooldown_storage(self) -> dict: + """Report whether cooldown persistence is currently usable.""" + return {"available": self._cooldowns.path is not None, "path": self._cooldowns.path, + "degraded": self._cooldowns.last_error is not None, + "last_error": self._cooldowns.last_error, "rows": len(self._cooldowns.detail()), + "warning": "冷却持久化写入失败,本次运行仍按内存态生效。" if self._cooldowns.last_error else None} + + def clear_cooldowns(self, cm, model: str | None = None) -> dict: + """Lift a circuit breaker or model cooldown after a confirmed recovery or admin reset. + + Persisting cooldowns removes the old "restart the gateway to clear it" workaround, + so an explicit reset has to be able to lift one both in memory and on disk. The two + outcomes are reported separately: an in-memory reset that could not be written is + not a durable reset, and saying otherwise would hide a cooldown that comes back. + """ + with self._lock: + entry = next((e for e in self._entries if e["cm"] is cm), None) + if entry is None: + return {"changed_in_memory": False, "durable": False} + # The in-memory reset always happens, even when this account's identity is not + # complete enough to key durable state. + durable = not self._durable_identity(entry) + if model: + routed_model = _upstream_model(model, self._entry_profile(entry)) + changed = self._model_fail.pop((entry["id"], routed_model), None) is not None + if not durable: + outcome = self._cooldowns.clear_model(entry["account_key"], entry["profile"], routed_model) + durable = outcome["durable"] + self._warn_storage("cooldown", durable) + return {"changed_in_memory": changed, "durable": durable} + changed = entry["fail_until"] > time.time() + entry["fail_until"] = 0.0 + entry["last_error"] = None + if not durable: + durable = self._cooldowns.clear_credential(entry["account_key"], entry["profile"])["durable"] + self._warn_storage("cooldown", durable) + return {"changed_in_memory": changed, "durable": durable} + + def reset_cooldowns_for(self, identity: str) -> dict: + """Lift every cooldown held by one account, addressed by its public identity. + + Persisting cooldowns removed the old "restart the gateway to clear it" workaround, so an + operator needs a supported way back when a breaker or a 429 cooldown was recorded in + error, or when upstream has demonstrably recovered. The whole account is reset rather + than just its circuit breaker: a credential that is mid-429 is exactly the case an + operator is trying to unstick, and lifting only the breaker would leave it unusable. + + This is a local state change. It never refreshes a token, contacts upstream, or queues + synchronization, so a subsequent genuine failure is free to arm the cooldown again. + """ + with self._lock: + entry = next((e for e in self._entries if e.get("account_key") == identity), None) + if entry is None: + raise KeyError(identity) + now = time.time() + changed = entry["fail_until"] > now + entry["fail_until"] = 0.0 + entry["last_error"] = None + for key in [k for k in self._model_fail if k[0] == entry["id"]]: + if self._model_fail[key] > now: + changed = True + del self._model_fail[key] + # One durable write for the whole account, replacing any breaker and model rows. An + # account whose identity is incomplete still resets in memory, but owns no disk row. + durable = not self._durable_identity(entry) + if not durable: + durable = self._cooldowns.forget(entry["account_key"])["durable"] + self._warn_storage("cooldown", durable) + return {"changed_in_memory": changed, "durable": durable} + + def entries(self) -> list[dict]: + """Return credential snapshots for account maintenance.""" + with self._lock: + return [dict(e) for e in self._entries] + + def _expiry_rank(self, e: dict) -> tuple: + """Order by earliest credit expiry, placing unknown balances last.""" + exp = self._ledger.soonest_expiry_of(e["id"]) if self._ledger else None + return (exp is None, exp or 0.0) + def _rescan(self): + self.prune() + paths = find_auth_files() if self._scan else [Path(entry["id"]) for entry in self.entries()] + self.reload(paths, reset=False) + + def _healthy(self, e: dict) -> bool: + return model_policy.credential_enabled(CONFIG, e) and time.time() >= e["fail_until"] + + @staticmethod + def _entry_profile(entry): + try: + return entry["cm"].summary().get("profile", "cn-cli") + except Exception: + return None + + @classmethod + def _entry_site(cls, entry): + profile = cls._entry_profile(entry) + return profile_site(profile) if profile else None + + def _zero_balance(self, entry, profile) -> bool: + """Restrict a confirmed zero-balance account to its advertised zero-rate models.""" + balance = (self._ledger.entry(entry["id"]).get("credits") or {}) if self._ledger else {} + if not balance: + return False + try: + return (bool(balance.get("intl")) == (profile_region(profile) == "intl") + and float(balance.get("credits") or 0) <= 0) + except (TypeError, ValueError): + return False + + def _has_credit(self, entry, profile): + balance = (self._ledger.entry(entry["id"]).get("credits") or {}) if self._ledger else {} + if not balance: + return profile_region(profile) == "cn" + try: + return (bool(balance.get("intl")) == (profile_region(profile) == "intl") + and float(balance.get("credits") or 0) > 0) + except (TypeError, ValueError): + return False + + def _eligible(self, entry, model, *, region=None, profile=None, rule=None): + if not model_policy.route_allowed(CONFIG, entry, model, rule=rule): + return False + actual = self._entry_profile(entry) + profile = profile or actual + if not profile or profile != actual or not _in_region(profile, region): + return False + configured = {candidate for item in self._entries if (candidate := self._entry_profile(item)) + and _in_region(candidate, region)} + if profile not in _model_profiles(model, region, configured): + return False + if CONFIG.get("account_catalogs") is not None or CONFIG.get("model_cache") is not None: + try: + identity = entry["cm"].summary()["account_key"] + except Exception: + return False + if identity != entry.get("account_key"): + return False + account = (CONFIG.get("account_catalogs") or {}).get(identity) or {} + models = _effective_account_scope(account, "serves", model_id=_upstream_model(model, profile)) + if account.get("profile") != profile or models is None: + return False + usable = _usable_models(models) + supported = any(item["id"] == _upstream_model(model, profile) for item in usable) + cli_auto = model == "auto" and profile == "cn-cli" and bool(usable) + # Disabling the guard must not borrow another account's model capabilities. + declared = any(item["id"] == _upstream_model(model, profile) + for item in _models_for_profile(profile, configured, scope="serves")) + passthrough = (model != "auto" and not declared and not CONFIG.get("model_guard") + and len(configured) == 1) + if model and not (supported or cli_auto or passthrough): + return False + # Zero-balance accounts may only use their own advertised zero-rate models. + return (not model or self._has_credit(entry, profile) + or self._model_free(entry, model, profile=profile)) + + def _model_free(self, entry, model: str | None, *, profile=None) -> bool: + """Check whether this account advertises the model as zero-rate.""" + if not model or model == "auto": + return False + profile = profile or self._entry_profile(entry) + if not profile: + return False + accounts = CONFIG.get("account_catalogs") + if accounts is not None or CONFIG.get("model_cache") is not None: + account = (accounts or {}).get(entry.get("account_key")) or {} + if account.get("profile") != profile: + return False + return _model_free(_effective_account_scope(account, "serves", model_id=_upstream_model(model, profile)), model, profile) + return _model_free(_models_for_profile(profile, model_id=_upstream_model(model, profile)), model, profile) + + @classmethod + def _entry_endpoint(cls, e: dict) -> str | None: + """Return the credential's backend endpoint for isolated model availability checks.""" + profile = cls._entry_profile(e) + return PROFILE_ENDPOINTS.get(profile) if profile else None + + @classmethod + def _model_block_key(cls, entry): + endpoint = cls._entry_endpoint(entry) + if endpoint and cls._entry_profile(entry) in SHARED_INTL_PROFILES: + identity = entry.get("account_key") or hashlib.sha256(str(entry.get("id", "")).encode()).hexdigest() + return f"{endpoint}#account:{identity}" + return endpoint + + + def _model_servable(self, e: dict, model: str | None) -> bool: + """Check backend/model backoff, skipping the check when no model is supplied.""" + if not model: + return True + endpoint = self._model_block_key(e) + if not endpoint: + return True + return time.time() >= self._blocks.until(endpoint, _block_model(model)) + + def _model_healthy(self, e: dict, model: str | None) -> bool: + """Check this credential's model-specific 429 cooldown.""" + if not model: + return True + routed_model = _upstream_model(model, self._entry_profile(e)) + return time.time() >= self._model_fail.get((e["id"], routed_model), 0.0) + + def _evict_sticky(self): + now = time.time() + while self._sticky: + k, (_, ts) = next(iter(self._sticky.items())) + if now - ts > STICKY_TTL or len(self._sticky) > STICKY_MAX: + self._sticky.pop(k) + else: + break + + def _candidates(self, model: str | None, *, region=None, tried=()) -> list[dict]: + """Exclude tried credentials and rank candidates by zero rate and credit expiry.""" + tried = set(tried) + healthy = [entry for entry in self._entries if entry["cm"] not in tried + and self._healthy(entry) + and self._eligible(entry, model, region=region) and self._model_healthy(entry, model) + and self._model_servable(entry, model)] + if not healthy: + return [] + # Prefer zero-rate models, then earlier credit expiry; unknown balances sort last. + healthy.sort(key=lambda entry: (not self._model_free(entry, model), *self._expiry_rank(entry))) + return healthy + + @staticmethod + def _capacity_error(): + return HTTPException(status_code=503, headers={"Retry-After": "3"}, detail={"error": { + "message": "符合当前路由和免费优先策略的账号在途名额已满,请稍后重试", + "type": "service_unavailable", "code": "credential_concurrency_limit"}}) + + @staticmethod + def _capacity_key(entry): + return entry.get("account_key") or entry["id"] + + + def pick(self, skey: str | None, model: str | None = None, *, region=None, + tried=(), with_capacity=False, requirements=None) -> CredentialManager | None: + """Select a healthy sticky or round-robin credential, preferring eligible zero-rate accounts.""" + self._rescan() # Reload and prune acquire their own locks. + with self._lock: + self._evict_sticky() + candidates = self._candidates(model, region=region, tried=tried) + if not candidates: + if skey: + self._sticky.pop(skey, None) + return None + if requirements is not None: + free = self._model_free(candidates[0], model) + checked = [(entry, requirements.violations( + model_capabilities.entry_model(sys.modules[__name__], entry, model))) + for entry in candidates if self._model_free(entry, model) == free] + candidates = [entry for entry, failures in checked if not failures] + if not candidates: + raise model_capabilities.capability_error([failure for _, failures in checked for failure in failures]) + limit = CONFIG.get("max_inflight_per_account", 0) + if with_capacity and limit: + free = self._model_free(candidates[0], model) + candidates = [entry for entry in candidates if self._model_free(entry, model) == free + and self._capacity.count(self._capacity_key(entry)) < limit] + if not candidates: + raise self._capacity_error() + best = candidates[0] + free = self._model_free(best, model) + top = [e for e in candidates if self._model_free(e, model) == free + and self._expiry_rank(e) == self._expiry_rank(best)] + if skey and skey in self._sticky: + cid, _ = self._sticky[skey] + sticky = next((e for e in top if e["id"] == cid), None) + if sticky is not None: + self._sticky[skey] = (cid, time.time()) + self._sticky.move_to_end(skey) + return sticky["cm"] + e = top[self._rr[region] % len(top)] + self._rr[region] += 1 + if skey: + self._sticky[skey] = (e["id"], time.time()) + return e["cm"] + + def headers_for(self, skey: str | None, model: str | None = None, *, region=None, + with_generation=False, tried=(), with_capacity=False, requirements=None): + """Recheck identity and atomically reserve account capacity before sending.""" + capacity_race = False + capability_failures = [] + for _ in range(max(1, len(self._entries))): + cm = self.pick(skey, model, region=region, tried=tried, with_capacity=with_capacity, requirements=requirements) + if cm is None: + return None + reason = None + with cm._lock: + try: + headers = cm.get_headers() + profile = profile_for_headers(headers) + generation = cm._generation + except Exception as error: + generation, reason = cm._generation, str(error) + if reason is not None: + self.cooldown(cm, reason=reason, generation=generation) + continue + with self._lock: + self.reload([cm.path], reset=False) + entry = next((entry for entry in self._entries if entry["cm"] is cm), None) + if (entry is not None and cm._generation == generation and self._healthy(entry) + and self._eligible(entry, model, region=region, profile=profile) and self._model_healthy(entry, model)): + if requirements is not None: + failures = requirements.violations(model_capabilities.entry_model(sys.modules[__name__], entry, model)) + if failures: + capability_failures.extend(failures) + continue + if with_capacity: + lease = self._capacity.acquire(self._capacity_key(entry), + CONFIG.get("max_inflight_per_account", 0), cm, generation) + if lease is None: + capacity_race = True + continue + return lease, headers + return ((cm, generation) if with_generation else cm), headers + if capacity_race: + raise self._capacity_error() + if capability_failures: + raise model_capabilities.capability_error(capability_failures) + return None + + @staticmethod + def _lease_matches(cm, generation): + if generation is None: + return True + try: + cm._load_if_stale() + except (OSError, ValueError): + return generation == cm._generation + return generation == cm._generation + + def cooldown(self, cm: CredentialManager, reason: str = "", *, generation=None): + with self._lock, (cm._lock if generation is not None else nullcontext()): + if not self._lease_matches(cm, generation): + return + for e in self._entries: + if e["cm"] is cm: + e["fail_until"] = time.time() + CRED_COOLDOWN + e["last_error"] = sanitize_log_text(reason, 256) + e["last_failure_at"] = time.time() + self._remember_credential(e, e["last_error"]) + _log(f"[cred] 凭证熔断 {CRED_COOLDOWN}s: {Path(cm.path).name} {reason}") + + def note_status(self, cm: CredentialManager | None, status: int, + model: str | None = None, raw: bytes = b"", *, generation=None, retry_after=None): + """Apply credential-wide auth cooldowns, per-model 429 cooldowns and backend/model backoff.""" + if cm is None: + return + if status in (401, 403): + self.cooldown(cm, reason=f"backend HTTP {status}", generation=generation) + return + not_servable = _parse_not_servable(raw, status) if model else None + if not_servable: + self.note_not_servable(cm, model, code=not_servable[0], msg=not_servable[1], generation=generation) + return + if status != 429 or not model: + return + now = time.time() + if retry_after is not None: + until = now + retry_after + else: + reset = _parse_reset_time(raw) + until = reset if reset is not None and reset > now else now + MODEL_COOLDOWN + until = min(until, now + MODEL_COOLDOWN_MAX) + with self._lock, (cm._lock if generation is not None else nullcontext()): + if not self._lease_matches(cm, generation): + return + self._model_fail = {k: v for k, v in self._model_fail.items() if v > now} + for e in self._entries: + if e["cm"] is cm: + routed_model = _upstream_model(model, self._entry_profile(e)) + key = (e["id"], routed_model) + until = max(until, self._model_fail.get(key, 0.0)) + self._model_fail[key] = until + self._remember_model(e, routed_model, until) + _log(f"[cred] 模型冷却 {model} @ {Path(cm.path).name} 至 " + f"{time.strftime('%m-%d %H:%M:%S', time.localtime(until))} (HTTP 429)") + + def model_cooldown_until(self, model: str | None, *, region=None) -> float | None: + """Return the earliest reset only when all healthy credentials are cooling down.""" + if not model: + return None + with self._lock: + now = time.time() + pool = [entry for entry in self._entries if self._healthy(entry) + and self._eligible(entry, model, region=region)] + if not pool: + return None + untils = [self._model_fail.get((entry["id"], _upstream_model(model, self._entry_profile(entry))), 0.0) + for entry in pool] + if any(now >= u for u in untils): + return None + return min(untils) + + def note_not_servable(self, cm, model: str, code: str = "", msg: str = "", *, generation=None) -> float: + """Isolate international model rejection by account; retain domestic backend backoff.""" + if not model: + return 0.0 + with self._lock, (cm._lock if generation is not None else nullcontext()): + if not self._lease_matches(cm, generation): + return 0.0 + entry = next((e for e in self._entries if e["cm"] is cm), None) + endpoint = self._model_block_key(entry) if entry else None + if not endpoint: + return 0.0 + row = self._blocks.note(endpoint, _block_model(model), code=code, msg=msg) + until = float(row.get("until") or 0.0) + _log(f"[block] 模型 {model} @{endpoint} 官方回 {code}," + f"{time.strftime('%m-%d %H:%M', time.localtime(until))} 前不再派发 " + f"(第 {row.get('hits')} 次){' | ' + msg[:80] if msg else ''}") + return until + + def note_model_ok(self, cm, model: str) -> bool: + """Clear model backoff immediately after a successful backend response.""" + if not model: + return False + entry = next((e for e in self._entries if e["cm"] is cm), None) + endpoint = self._model_block_key(entry) if entry else None + return bool(endpoint) and self._blocks.clear(endpoint, _block_model(model)) + + def model_block_until(self, model: str | None, *, region=None) -> float | None: + """Return a retry time only when every potential backend has confirmed backoff.""" + if not model: + return None + now = time.time() + with self._lock: + candidates = [e for e in self._entries + if self._healthy(e) and (region is None + or _in_region(self._entry_profile(e), region)) + and model_policy.route_allowed(CONFIG, e, model)] + endpoints = {self._model_block_key(e) for e in candidates} + # Unknown catalogs remain potential sources, but cannot authorize dispatch. + capable = {self._model_block_key(e) for e in candidates + if (profile := self._entry_profile(e)) + and profile in _model_profiles(model, profile_region(profile))} + accounts = CONFIG.get("account_catalogs") + def catalog_unknown(entry): + profile = self._entry_profile(entry) + if not profile: + return False + if accounts is not None or CONFIG.get("model_cache") is not None: + account = (accounts or {}).get(entry.get("account_key")) or {} + return (account.get("profile") != profile + or _account_scope(account, "serves") is None) + return _catalog_for(profile, "serves") is None + unknown = {self._model_block_key(e) for e in candidates if catalog_unknown(e)} + endpoints.discard(None) + endpoints &= capable | unknown + if not endpoints: + return None + routed = _block_model(model) + untils = [self._blocks.until(endpoint, routed) for endpoint in endpoints] + if any(until <= now for until in untils): + return None + return max(untils) + + def model_blocks_detail(self) -> list: + """Return model backoff details for diagnostics.""" + return self._blocks.detail() + + def refresh_due(self, margin_s: int = CRED_REFRESH_MARGIN, keepalive_s: int = CRED_KEEPALIVE_S): + """Refresh expiring or idle tokens with generation-scoped failure backoff.""" + with self._lock: + entries = list(self._entries) + now = time.time() + for entry in entries: + if now < entry.get("fail_until", 0.0) or not model_policy.credential_enabled(CONFIG, entry): + continue + cm = entry["cm"] + failure = None + refreshed = keepalive_due = False + with cm._lock: + try: + summary = cm.summary() + exp = (summary.get("token_expires_at") or 0) / 1000 + last = (summary.get("last_refresh_time") or 0) / 1000 + expiry_due = bool(summary.get("token_expired") or (exp and exp - now < margin_s)) + keepalive_due = (not expiry_due and keepalive_s > 0 + and now >= entry.get("keepalive_after", 0.0) + and (last <= 0 or now - last >= keepalive_s)) + if not (expiry_due or keepalive_due): + continue + refreshed = cm.refresh_if_due(margin_s, keepalive_s if keepalive_due else 0) + entry["keepalive_after"] = 0.0 + except Exception as error: + failure = (str(error), cm._generation) + if keepalive_due: + entry["keepalive_after"] = now + CRED_KEEPALIVE_RETRY_S + if failure: + self.cooldown(cm, reason=failure[0], generation=failure[1]) + elif refreshed: + _log(f"[cred] {'每日保活刷新' if keepalive_due else '已主动刷新'}并回写: {Path(entry['id']).name}") + + def remove_file(self, name: str) -> bool: + """Share the refresh lock so an in-flight refresh cannot recreate a deleted file.""" + with self._lock: + entry = next((x for x in self._entries if os.path.basename(x["id"]) == name), None) + if entry is None: + return False + cm = entry["cm"] + try: + with cm._lock, credential_file_lock(cm.path.parent, cm.path.name): + os.unlink(entry["id"]) + cm.invalidate() + except FileNotFoundError: + pass + except OSError: + return False + self.prune() + return True + + def forget_cooldowns(self, entry): + """Drop persisted cooldowns for a credential that no longer exists.""" + identity = entry.get("account_key") + if identity: + self._warn_storage("cooldown", self._cooldowns.forget(identity)["durable"]) + + def forget_usage(self, entry): + """Drop a deleted credential's cached usage, in the store and the live aggregate.""" + snapshots = CONFIG.get("usage_snapshots") + if snapshots is not None: + snapshots.forget(entry["id"]) + accounts = CONFIG.get("usage_daily_accounts") + if isinstance(accounts, dict): + accounts.pop(entry["id"], None) + + def forget_credential_state(self, entry): + """Run both independent cleanups for a credential that is gone or replaced.""" + self.forget_cooldowns(entry) + self.forget_usage(entry) + + def first(self) -> CredentialManager | None: + with self._lock: + return self._entries[0]["cm"] if self._entries else None + + def snapshot(self) -> list[dict]: + with self._lock: + now = time.time() + out = [] + for e in self._entries: + s: dict = {"auth_file": e["id"], "healthy": self._healthy(e), + "in_flight": self._capacity.count(self._capacity_key(e)), + "max_in_flight": CONFIG.get("max_inflight_per_account", 0), + "model_cooldowns": {m: time.strftime("%m-%d %H:%M:%S", time.localtime(u)) + for (cid, m), u in self._model_fail.items() + if cid == e["id"] and u > now}, + "sticky_sessions": sum(1 for _, (cid, ts) in self._sticky.items() + if cid == e["id"] and now - ts <= STICKY_TTL)} + try: + s.update(e["cm"].summary()) + except Exception: + s["error"] = "凭据读取失败" + out.append(s) + return out + + +def _refresher_loop(pool: CredentialPool): + """Refresh idle credentials before expiry and persist renewed tokens.""" + while True: + time.sleep(60) + try: + pool.refresh_due() + except Exception as e: + _log(f"[cred] 刷新线程异常: {e}") + +CHECKIN_FIRST_DELAY = 30 # Initial check-in delay in seconds +HOUSEKEEP_INTERVAL = 3600 # Account maintenance interval in seconds + + +def _bearer_token(headers: dict) -> str: + return (headers.get("Authorization") or "").removeprefix("Bearer ").strip() + + +_HOUSEKEEP_LOCK = threading.Lock() + + +def _sync_error(pool, ledger, entry, generation, phase, error): + message = f"{phase}: {_network_error_text(error)}" + pool.apply_if_current(entry["cm"], generation, lambda: ledger.note_error(entry["id"], message)) + _log(f"[{phase}] {Path(entry['id']).name} 同步失败(保留旧数据): {message}") + + +def _buddy_context(entry, headers, consent_revision=None): + def select_model(requested=None): + from app.audit_store import safe_label + pool = CONFIG.get("cred_pool") + account = (CONFIG.get("account_catalogs") or {}).get(entry.get("account_key")) or {} + if pool is None or entry.get("profile") != "cn-work" or account.get("profile") != "cn-work": + return None + with pool._lock: + current = next((item for item in pool._entries if item["cm"] is entry["cm"] + and item.get("account_key") == entry.get("account_key")), None) + if current is None or not pool._healthy(current): + return None + candidates = [] + for item in _usable_models(_account_scope(account, "serves")): + model = item["id"] + rate = _multiplier_value(item.get("credits")) + if (not safe_label(model) or model in {".", ".."} or requested and model != requested + or rate is None or not 0 <= rate < float("inf") or "custom" in (item.get("tags") or [])): + continue + rule = model_policy.rule_for(CONFIG, model) + if (rule["upstream_id"] != model or not pool._eligible(current, model, profile="cn-work", rule=rule) + or not pool._model_healthy(current, model) or not pool._model_servable(current, model)): + continue + name = item.get("name") + candidates.append((rate, model, name if isinstance(name, str) and len(name) <= 160 else model)) + if not candidates: + return None + _, model, name = min(candidates) + return {"id": model, "name": name} + return buddy.context(CONFIG, entry, consent_revision, headers=headers, task_model=select_model) + + +def _sync_credits(pool, ledger, entry, *, checkin, failed, expected_identity=None): + if not model_policy.credential_enabled(CONFIG, entry): + return None + cm, cid = entry["cm"], entry["id"] + generation = None + try: + with cm._lock: + try: + if expected_identity is not None and cm.summary().get("account_key") != expected_identity: + failed.add(cid) + return None + headers = cm.get_headers() + finally: + generation = cm._generation + profile = profile_for_headers(headers) + identity = account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")) + if entry.get("account_key") and entry["account_key"] != identity: + failed.add(cid) + return None # A path now owned by another account must be rescheduled with its own preferences. + site = site_for_headers(headers) + token, uid, domain = _bearer_token(headers), headers.get("X-User-Id", ""), headers.get("X-Domain", "") + day = time.strftime("%Y-%m-%d") + if checkin and model_policy.credential_auto_checkin(CONFIG, entry) and not ledger.checkin_done(cid, day): + try: + def can_claim(): + return (model_policy.credential_auto_checkin(CONFIG, entry) + and pool.apply_if_current(cm, generation, lambda: None)) + result = checkin_service.perform(token, uid=uid, domain=domain, can_claim=can_claim) + if result["state"] == "cancelled": + if not pool.apply_if_current(cm, generation, lambda: None): + failed.add(cid) + return None + # A preference-only cancellation must not interrupt balance refresh. + elif not pool.apply_if_current(cm, generation, lambda: ledger.mark_checkin( + cid, day, result["ok"], result.get("code"), result["message"], state=result["state"])): + failed.add(cid) + return None + _log(f"[checkin] {Path(cid).name}: ok={result['ok']} already={result.get('already')} code={result.get('code')}") + except Exception as error: + _sync_error(pool, ledger, entry, generation, "checkin", error) + if not model_policy.credential_enabled(CONFIG, entry): + return None + if checkin and model_policy.credential_auto_travel(CONFIG, entry): + try: + def can_travel(): + return (model_policy.credential_auto_travel(CONFIG, entry) + and pool.apply_if_current(cm, generation, lambda: None)) + trip = travel.perform(token, profile_for_headers(headers), can_write=can_travel, + buddy_context=_buddy_context(entry, headers)) + if not pool.apply_if_current(cm, generation, lambda: travel.remember(ledger, cid, trip)): + failed.add(cid) + return None + buddy.daily_warning(CONFIG, entry.get("account_key"), entry.get("profile"), trip) + except Exception as error: + _sync_error(pool, ledger, entry, generation, "travel", error) + if not model_policy.credential_enabled(CONFIG, entry): + return None + balance = credits_mod.fetch_credits(token, uid=uid, domain=domain) + if bool(balance.get("intl")) != (site == INTERNATIONAL): + raise ValueError("积分响应与凭据站点不一致") + if not pool.apply_if_current(cm, generation, lambda: ledger.update_credits(cid, balance)): + failed.add(cid) + return None + _log(f"[credits] {Path(cid).name}: 站点 {site},余额 {balance['credits']}") + profile = profile_for_headers(headers) + return entry, generation, headers, profile + except Exception as error: + failed.add(cid) + _sync_error(pool, ledger, entry, generation, "credits", error) + return None + + +def _publish_model_cache(): + """Publish account-scoped versioned catalogs, excluding ownerless shared caches.""" + cache = CONFIG.get("model_cache") + if cache is not None: + pool = CONFIG.get("cred_pool") + with pool._lock if pool is not None else nullcontext(): + accounts = {} + for entry in pool.entries() if pool is not None else []: + identity, profile = entry.get("account_key"), entry.get("profile") + if not identity or not profile: + continue + key = catalog_cache_key(profile, identity) + known = cache.age(key) is not None + accounts[identity] = {"profile": profile, + "models": cache.models(key) if known else None, + "serves": cache.serves(key) if known else None} + CONFIG["account_catalogs"] = accounts + catalogs = {profile: None for profile in PROFILE_ENDPOINTS} + for account in accounts.values(): + if account["models"] is not None: + models = catalogs[account["profile"]] + if models is None: + models = catalogs[account["profile"]] = [] + models.extend(account["models"]) + CONFIG["model_catalogs"] = catalogs + CONFIG["models_remote"], CONFIG["models_intl"] = catalogs["cn-cli"], catalogs["intl-cli"] + invalidate_model_table() + + +def _sync_model_catalogs(pool, ledger, refs, failed): + cache = CONFIG.get("model_cache") + if cache is None: + return + for entry, generation, headers, profile in refs.values(): + identity = account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")) + if not model_policy.credential_enabled(CONFIG, entry): + continue + key = catalog_cache_key(profile, identity) + if cache.fresh(key) and not entry.get("catalog_dirty"): + continue + try: + scopes = credits_mod.fetch_model_scopes( + _bearer_token(headers), domain=headers.get("X-Domain", ""), + uid=headers.get("X-User-Id", ""), enterprise_id=headers.get("X-Enterprise-Id", "")) + models, serves = scopes["picker"], scopes["account"] + def publish(): + cache.put(key, models, serves=serves) + for current in pool._entries: + if current["cm"] is entry["cm"]: + current["catalog_dirty"] = False + if pool.apply_if_current(entry["cm"], generation, publish): + _log(f"[models] {profile} 模型表已刷新: 选择器 {len(models)} 个," + f"账号根表 {len(serves)} 个") + else: + failed.add(entry["id"]) + except Exception as error: + failed.add(entry["id"]) + _sync_error(pool, ledger, entry, generation, "models", error) + _publish_model_cache() + + +def _sync_usage(pool, entries=None, expected_identity=None): + """Refresh usage during maintenance with independent per-account snapshots.""" + accounts = CONFIG.get("usage_daily_accounts") + if not isinstance(accounts, dict): + accounts = CONFIG["usage_daily_accounts"] = {} + targets = pool.entries() if entries is None else entries + target_ids = {entry["id"] for entry in targets} + previous_stale = set((CONFIG.get("usage_daily") or {}).get("stale_accounts", [])) + stale = {entry["id"] for entry in pool.entries() + if entry["id"] not in target_ids and Path(entry["id"]).name in previous_stale} + for entry in targets: + if not model_policy.credential_enabled(CONFIG, entry): + continue + try: + cm = entry["cm"] + with cm._lock: + if expected_identity is not None and cm.summary().get("account_key") != expected_identity: + stale.add(entry["id"]) + continue + headers = cm.get_headers() + generation = cm._generation + site = site_for_headers(headers) + usage = credits_mod.fetch_request_usage(_bearer_token(headers), uid=headers.get("X-User-Id", ""), + domain=headers.get("X-Domain", "")) + def store(): + accounts[entry["id"]] = {"identity": entry.get("account_key"), "site": site, + "by_day": usage["by_day"], + "total_credits": round(usage["total_credits"], 2), + "requests": usage["requests"], + "partial": bool(usage.get("partial")), + "fetched_at": time.time()} + snapshots = CONFIG.get("usage_snapshots") + if snapshots is not None and entry.get("account_key") and entry.get("uid"): + # Only a validated identity may key durable state. An absent flag defaults to + # False, but a present value is passed through raw so the store's own strict + # validator governs: coercing it here with bool() would mask a malformed value + # and persist it as a legitimate flag. + snapshots.store(entry["id"], entry["account_key"], site, usage, + partial=usage.get("partial", False)) + if not pool.apply_if_current(cm, generation, store): + stale.add(entry["id"]) + except Exception as error: + stale.add(entry["id"]) + _log(f"[usage] {Path(entry['id']).name} 明细拉取失败(保留其上次成功快照): {_network_error_text(error)}") + _publish_usage_daily(pool, stale) + return stale & target_ids + + +def _usage_row_expired(snap) -> bool: + """Whether a cached row is too old to keep showing.""" + return time.time() - float(snap.get("fetched_at") or 0.0) > USAGE_CACHE_MAX_AGE_S + + +def _adopt_cached_usage(pool, snapshots, accounts): + """Seed the aggregate from the cache, keeping only rows that still belong to their account. + + Ownership is checked here, at the moment of use, rather than only when the cache was + written: a path can be reused between restarts, and a row hydrated on an earlier pass + would otherwise keep showing the previous account's usage indefinitely. + """ + # The live map is mutated by forget_usage under the pool lock, so the reads below need that + # same lock to stay consistent with a concurrent credential deletion or replacement. + with pool._lock: + cached = snapshots.accounts() + if not cached and not accounts: + return # Nothing cached and nothing live: leave the pool untouched. + # Only an account whose identity components are all present may own durable usage: an + # account_key is a hash of the profile and UID, so an account with no UID would otherwise + # share one identity with every other such account. + owners = {entry["id"]: entry.get("account_key") for entry in pool.entries() + if entry.get("account_key") and entry.get("profile") and entry.get("uid")} + # Drop any row whose recorded owner no longer matches the account at that path. This covers + # rows hydrated by an earlier pass as well as live rows, so a reused path cannot keep + # displaying the previous account's usage. Iterate over a copy of the items, so a row is + # never looked up in the live map after the keys were copied. + for path, row in list(accounts.items()): + if owners.get(path) != row.get("identity"): + accounts.pop(path, None) + for path, row in cached.items(): + identity = row.get("identity") + if owners.get(path) != identity: + snapshots.forget(path) # The path moved on; the cached row is not ours to show. + continue + if _usage_row_expired(row): + snapshots.forget(path) + continue + # A live snapshot from this run always wins; the cache only fills a gap. + if path in accounts: + continue + # A restored snapshot is stale until a refresh confirms it; clearing that is per account. + accounts[path] = dict(row) + accounts[path]["stale"] = True + CONFIG["usage_daily_accounts"] = accounts + + +def _publish_usage_daily(pool, stale=()): + """Aggregate enabled accounts' usage, retaining failed snapshots with explicit staleness.""" + # CredentialPool.forget_usage removes entries from the live map under the pool lock, so the + # snapshot taken here and the pruning inside the loop must hold that same lock. Copying the + # keys and then indexing the live dictionary would otherwise raise KeyError when a credential + # is deleted or replaced while a usage-maintenance pass is running. + with pool._lock: + accounts = CONFIG.get("usage_daily_accounts") + if not isinstance(accounts, dict): + accounts = {} + CONFIG["usage_daily_accounts"] = accounts + # Seed from the on-disk cache so the dashboard is not blank until the first refresh. + snapshots = CONFIG.get("usage_snapshots") + if snapshots is not None: + _adopt_cached_usage(pool, snapshots, accounts) + enabled = {e["id"] for e in pool.entries() if model_policy.credential_enabled(CONFIG, e)} + # Aggregate from a consistent local copy of the live map rather than indexing it per row. + rows = dict(accounts) + by_day, groups = {}, {} + used, count = 0.0, 0 + partial = False # Upstream paging hid additional usage for some account. + stale_out = set() # Accounts whose displayed figures are not from this run. + newest = 0.0 + for cred_id, snap in rows.items(): + if cred_id not in enabled: + continue + if _usage_row_expired(snap): + accounts.pop(cred_id, None) # Drop it from the live map, not just from this sum. + continue + site = snap.get("site") or "domestic" + group = groups.setdefault(site, {"by_day": {}, "total_credits": 0.0, "requests": 0}) + for day, models in (snap.get("by_day") or {}).items(): + total_day = by_day.setdefault(day, {}) + site_day = group["by_day"].setdefault(day, {}) + for model, credit in models.items(): + total_day[model] = round(total_day.get(model, 0.0) + credit, 6) + site_day[model] = round(site_day.get(model, 0.0) + credit, 6) + group["total_credits"] += float(snap.get("total_credits") or 0) + group["requests"] += int(snap.get("requests") or 0) + used += float(snap.get("total_credits") or 0) + count += int(snap.get("requests") or 0) + newest = max(newest, float(snap.get("fetched_at") or 0)) + # `partial` is the aggregate "this view is incomplete" flag the dashboard shows, so a + # failed or not-yet-refreshed account sets it; `stale_accounts` names which ones. + if snap.get("partial") or snap.get("stale"): + partial = True + if snap.get("stale"): + stale_out.add(Path(cred_id).name) + # Failed enabled accounts must remain visible even without a prior snapshot. + for cred_id in stale: + if cred_id in enabled: + partial = True + stale_out.add(Path(cred_id).name) + # A zero timestamp preserves quota-difference fallback when no usage snapshot exists. + for group in groups.values(): + group["total_credits"] = round(group["total_credits"], 2) + out = {"by_day": by_day, "groups": groups, "total_credits": round(used, 2), + "requests": count, "fetched_at": newest, "partial": partial} + if stale_out: + out["stale_accounts"] = sorted(stale_out) + CONFIG["usage_daily"] = out + _log(f"[usage] 明细已同步: {count} 请求 / {used:.2f} credits" + + (f" | {len(stale_out)} 账号同步失败" if stale_out else "")) + + +def _housekeep_once(pool: CredentialPool, ledger, *, pending_only=False): + """Serialize generation-scoped maintenance; new credentials only trigger balance/catalog reads.""" + if credits_mod is None: + return + with _HOUSEKEEP_LOCK: + pool._rescan() + # Drop expired deadlines and aged usage so neither table grows without bound; a failed + # prune is reported rather than silently leaving stale rows on disk. + pool._cooldowns.prune() + pool._warn_storage("cooldown", pool._cooldowns.last_error is None) + if CONFIG.get("usage_snapshots") is not None: + CONFIG["usage_snapshots"].prune() + ids = pool.begin_sync(all_entries=not pending_only) + failed = set() + try: + refs = {} + for entry in pool.entries(): + if entry["id"] not in ids: + continue + result = _sync_credits(pool, ledger, entry, checkin=not pending_only, failed=failed) + if result is not None: + refs[entry["id"]] = result + _sync_model_catalogs(pool, ledger, refs, failed) + if not pending_only: + _sync_usage(pool) + except Exception: + failed.update(ids) + raise + finally: + pool.end_sync(ids, failed) + + +def _housekeeper_loop(pool: CredentialPool, ledger) -> None: + """Wake for new credentials, retry failed work with backoff, and run hourly maintenance.""" + next_full = time.monotonic() + CHECKIN_FIRST_DELAY + while True: + pool._sync_event.wait(pool.sync_wait(next_full - time.monotonic())) + full_due = time.monotonic() >= next_full + try: + _housekeep_once(pool, ledger, pending_only=not full_due) + except Exception as error: + _log(f"[housekeeper] 循环异常: {_network_error_text(error)}") + if full_due: + next_full = time.monotonic() + HOUSEKEEP_INTERVAL + + +# --------------------------------------------------------------------------- +# Model inventory +# --------------------------------------------------------------------------- + +# Fallback models for legacy domestic deployments without a cloud catalog. +DEFAULT_MODELS = [ + "hy4-preview", "hy4-preview-x", + "hy3", "hy3-x", + "deepseek-v4-pro", "deepseek-v4-flash", "deepseek-v4.1-flash", "deepseek-v3-2-volc", + "glm-5.3", "glm-5.3-flash", "glm-5.2", "glm-5.1", "glm-5.0", "glm-5.0-turbo", + "glm-5v-turbo", "glm-4.7", "glm-4.6", "glm-4.6v", + "minimax-m3", "minimax-m2.7", "minimax-m2.5", + "kimi-k3-1", "kimi-k2.7", "kimi-k2.6", "kimi-k2.5", "kimi-k2-thinking", + "hunyuan-chat", "default", + "auto", # Backend-selected model alias +] + + +# Supported optional upstream request fields. +PASSTHROUGH_BODY_KEYS = { + "model", "messages", "tools", "tool_choice", "temperature", + "max_tokens", "max_completion_tokens", "top_p", "stream", + "stream_options", "stop", "presence_penalty", "frequency_penalty", + "n", "response_format", "seed", "user", "reasoning_effort", "prompt_cache_key", + "verbosity", "reasoning_summary", "parallel_tool_calls", +} + +# --------------------------------------------------------------------------- +# FastAPI application +# --------------------------------------------------------------------------- + +app = FastAPI(title="codebuddy2api", version=APP_VERSION, lifespan=inference_lifespan) +app.add_middleware(InferenceResourcesMiddleware, config=lambda: CONFIG) + +# Anthropic error types: https://platform.claude.com/docs/en/api/errors +_ANTHROPIC_ERROR_TYPES = { + "auth_error": "authentication_error", + "rate_limit_error": "rate_limit_error", + "invalid_request_error": "invalid_request_error", + "not_found_error": "not_found_error", + "upstream_error": "api_error", +} + + +@app.exception_handler(HTTPException) +async def _protocol_http_exception(request: Request, exc: HTTPException): + """Shape /v1 errors for the client protocol; retain FastAPI defaults elsewhere.""" + path = request.url.path + if not path.startswith("/v1/"): + return await _default_http_exception_handler(request, exc) + detail = exc.detail + err = detail.get("error") if isinstance(detail, dict) else None + if not isinstance(err, dict): + err = {"message": str(detail), "type": "error"} + message = str(err.get("message") or "") + if path.startswith("/v1/messages"): + # Preserve upstream business codes in Anthropic error envelopes. + etype = _ANTHROPIC_ERROR_TYPES.get(str(err.get("type") or "")) + if exc.status_code == 404: + etype = "not_found_error" # Anthropic's required type for HTTP 404. + elif etype is None: + etype = "api_error" if exc.status_code >= 500 else "invalid_request_error" + error_obj = {**err, "type": etype, "message": message} # Retain structured error fields. + return JSONResponse({"type": "error", "error": error_obj}, + status_code=exc.status_code, headers=exc.headers) + # OpenAI uses a top-level error object. + body = {"error": {**err, "message": message}} + return JSONResponse(body, status_code=exc.status_code, headers=exc.headers) +CONFIG: dict = {"api_key": "", "cred": None, "log_path": None, "ledger": None, + "admin_csrf": True, # Startup-only Origin/CSRF policy + "admin_allowed_origins": "", # Extra trusted management Origins (hot) + "models_remote": None, # Domestic cloud model inventory + "models_intl": None, # Eligible international model inventory + "model_cache": None, # Versioned catalog cache + "model_catalogs": {}, # Display-only merged product catalogs + "account_catalogs": None, # Account-scoped models/serves; None enables legacy embedding + "trial_ledger": None, + "model_guard": True, # Reject models absent from authorized catalogs + "max_images": 16, "image_policy": "truncate", + "max_request_bytes": 32 * 1024 * 1024, "log_body_limit": 65536, + "max_inbound_bytes": 64 * 1024 * 1024, + "max_collect_bytes": 8 * 1024 * 1024, "max_concurrent": 64, + "upstream_keepalive": False, "max_inflight_per_account": 0, + "request_context_mode": "legacy", + "stream_mode": "compatible", "model_capability_guard": True, + "responses_projection_mode": "balanced", "responses_projection_max_bytes": 40000, + "failover_max": 0, # Credential failovers allowed before the first response byte + "retry_write_timeout": False, # Opt-in replay after incomplete writes + "usage_daily": None, # Usage aggregated by date and model + "usage_daily_accounts": None, # Independent per-account usage snapshots + "usage_snapshots": None, # On-disk cache of the per-account snapshots + "credit_price_cny": None, "credit_price_usd": None, "usd_rate": None, + "desensitize": False, "no_compact": False, "keep_tool_metadata": False} # None prices use module defaults. + +# In-memory OAuth sessions do not survive restarts. +_OAUTH = auth_oauth.OAuthManager(user_agent=USER_AGENT) + + +# --------------------------------------------------------------------------- +# Runtime audit events +# --------------------------------------------------------------------------- + + + +def _log(msg: str): + """Persist allowlisted runtime events in SQLite, never free-form text or secrets.""" + audit = CONFIG.get("audit_store") + component = re.match(r"\[(cred|credits|models|usage|trial|checkin|housekeeper)\]", msg) + if audit is not None and component: + # Persist event codes, not free-form lines which may contain upstream data. + code = "cooldown" if "熔断" in msg or "冷却" in msg else "failure" if "失败" in msg or "异常" in msg else "updated" + audit.event("runtime", component.group(1), {"code": code}) + + +def _log_json(label: str, value): + if CONFIG.get("log_path") and CONFIG.get("log_body_limit", 65536): + _log(f"{label}\n{format_log_body(value, CONFIG.get('log_body_limit', 65536))}") + + +def _log_text_body(label: str, text: str): + if CONFIG.get("log_path") and CONFIG.get("log_body_limit", 65536): + _log(f"{label}\n{sanitize_log_text(text, CONFIG.get('log_body_limit', 65536))}") + + + + +def _truncate(s: str, n: int = 80) -> str: + s = str(s).replace("\n", " ").strip() + return s[:n] + ("…" if len(s) > n else "") + + +def _check_auth(authorization: Optional[str], x_api_key: Optional[str]): + require_api_key(CONFIG["api_key"], authorization, x_api_key) + + +def _check_admin_auth(authorization: Optional[str], x_api_key: Optional[str]): + if not CONFIG.get("api_key"): + raise HTTPException(status_code=503, detail={"error": {"message": "管理接口需要配置 API key", + "type": "management_locked"}}) + _check_auth(authorization, x_api_key) + + +def _cred_for(payload: dict, model: str | None = None, *, region=None, tried=(), requirements=None): + """Select a fresh credential lease and headers, excluding tried accounts; report unavailable capacity.""" + context = current_context() + raw_key = context.session_key if context is not None and context.scoped else session_key(payload) + skey = f"{region}:{raw_key}" if raw_key and region is not None else raw_key + skey = model_policy.sticky_scope(CONFIG, skey, model) + pool = CONFIG.get("cred_pool") + if pool is not None: + resources = request_resources.get() + picked = pool.headers_for(skey, model, region=region, with_generation=True, tried=tried, + with_capacity=resources is not None, requirements=requirements) + if picked is None: + until = pool.model_cooldown_until(model, region=region) + if until: + t = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(until)) + raise HTTPException(status_code=429, + headers={"Retry-After": str(max(1, math.ceil(until - time.time())))}, + detail={"error": { + "message": f"模型 {model} 额度冷却中(全部凭证),预计 {t} 重置后恢复", + "type": "rate_limit_error"}}) + blocked = pool.model_block_until(model, region=region) + if blocked: + # Report confirmed unsupported models as HTTP 404. + t = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(blocked)) + raise HTTPException(status_code=404, detail={"error": { + "message": f"模型 {model} 在当前所有已登录后端均不可用(官方回 service info not found)," + f"预计 {t} 后重试;请改用 /v1/models 列出的模型", + "type": "invalid_request_error", "code": "model_not_found", + "param": "model"}}) + raise HTTPException(status_code=503, headers={"Retry-After": "3" if _catalog_pending(region) else "30"}, + detail={"error": {"message": "无可用凭证(未登录、目录/额度未就绪或全部熔断)", + "type": "auth_error"}}) + cm, headers = picked + if resources is not None: + resources.add(cm) + else: + cm = CONFIG["cred"] + if cm is None or cm in {_cred_manager(item) for item in tried}: + raise HTTPException(status_code=503, detail={"error": {"message": "未找到登录凭据,请先在桌面端登录 CodeBuddy/WorkBuddy", "type": "auth_error"}}) + with cm._lock: + headers = cm.get_headers() + cm = (cm, cm._generation) + profile = profile_for_headers(headers) + if not _in_region(profile, region): + raise HTTPException(status_code=503, detail={"error": {"message": "未找到指定地域凭据", "type": "auth_error"}}) + if context is not None and context.scoped: + identity = account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")) + headers["X-Conversation-ID"] = context.conversation_id(profile, identity) + else: + headers.update(_dynamic_request_headers(f"{profile}:{skey}" if skey else None)) + return cm, headers + + +def _route_chat(payload, body, rid, *, tried=()): + """Validate account capabilities and derive each routed body from canonical input.""" + context = current_context() + enabled = context.capability_guard if context is not None else CONFIG.get("model_capability_guard", True) + cred = None + try: + requirements = (model_capabilities.Requirements.from_request( + body, payload, context.protocol if context is not None else "chat") if enabled else None) + cred, headers = _cred_for(payload, body.get("model"), tried=tried, requirements=requirements) + profile = profile_for_headers(headers) + routed_model = _upstream_model(body.get("model"), profile) + if requirements is not None and CONFIG.get("cred_pool") is None: + entry = {"profile": profile, "account_key": account_key( + profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id"))} + failures = requirements.violations(model_capabilities.entry_model(sys.modules[__name__], entry, body.get("model"))) + if failures: + raise model_capabilities.capability_error(failures) + canonical = body + if routed_model != body.get("model"): + body = {**body, "model": routed_model} + body, merged_runs, merged_messages = merge_intl_user_images(body, profile) + if body is not canonical: + _guard_request_size(body) + if merged_runs: + observe_attempt("intl_image_merge", merged_runs=merged_runs, merged_messages=merged_messages) + url = chat_url_for_headers(headers) + observe_route(public_model=payload.get("model", "auto"), upstream_model=routed_model, + profile=profile, credential=account_key(profile, headers.get("X-User-Id"), + headers.get("X-Enterprise-Id"))) + _log(f"[{rid}] ROUTE | region={profile_region(profile)} | profile={profile} | model={routed_model} | url={url}") + return body, cred, headers, url + except HTTPException as error: + release_credential(cred) + detail = error.detail.get("error", {}) if isinstance(error.detail, dict) else {} + observe_attempt("request_preflight_rejected", code=detail.get("code")) + raise + + +def _note_cred_model_ok(cred, model: str | None) -> None: + """Clear backend/model backoff after an upstream HTTP 200 response.""" + pool = CONFIG.get("cred_pool") + if pool is not None and cred is not None and model: + cm = cred[0] if isinstance(cred, tuple) else cred + pool.note_model_ok(cm, model) + + +def _note_cred_status(cred, status: int, model: str | None = None, raw: bytes = b"", *, retry_after=None): + """Record generation-scoped authentication, quota and unsupported-model failures.""" + pool = CONFIG.get("cred_pool") + if pool is not None and cred is not None: + cm, generation = cred if isinstance(cred, tuple) else (cred, None) + pool.note_status(cm, status, model=model, raw=raw, generation=generation, retry_after=retry_after) + +@app.get("/health") +def health(): + """Return public liveness without accessing or exposing credentials.""" + return {"status": "ok"} + + +@app.get("/admin/credentials") +def admin_list_credentials(authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Return account expiry, health and session-binding metadata.""" + _check_admin_auth(authorization, x_api_key) + pool = CONFIG.get("cred_pool") + if CONFIG.get("management") is not None: + return {"credentials": CONFIG["management"].admin_credential_inventory()} + return {"credentials": pool.snapshot() if pool else []} + + +class CredentialConflictError(CredentialFileError): + """Signal that another credential file already owns the account.""" + + +def _store_credential(directory: Path, name: str, content: bytes, uid: str, *, replace_identity=True, + replace_existing=True) -> Path: + """Serialize imports and logins with background refresh and standalone CLI writes.""" + pool = CONFIG.get("cred_pool") + identity = _credential_identity(json.loads(content)) + target = directory.resolve() / name + with pool._lock if pool is not None else nullcontext(): + cm = None + if pool is not None: + pool._rescan() + holder = pool.find_by_uid(uid, identity) + if holder and holder != str(target): + raise CredentialConflictError("该账号已在凭证池中") + entry = next((entry for entry in pool._entries if entry["id"] == str(target)), None) + cm = entry["cm"] if entry else None + elif CONFIG.get("cred") is not None and CONFIG["cred"].path.resolve() == target: + cm = CONFIG["cred"] + with cm._lock if cm is not None else nullcontext(): + with credential_file_lock(directory, name): + if not replace_existing and target.exists(): + raise CredentialConflictError("文件已存在,需明确允许替换") + if not replace_identity and target.exists() and _cred_identity(target) != identity: + raise CredentialConflictError("OAuth 不可覆盖其他产品或账号的凭据") + target = atomic_write_credential(directory, name, content) + if pool is not None: + pool.reload([target]) + elif cm is not None: + cm.invalidate() + return target + + +@app.post("/admin/credentials") +async def admin_add_credential(request: Request, + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Validate and atomically import credentials from the controlled directory.""" + _check_admin_auth(authorization, x_api_key) + try: + body = await request.json() + except (ValueError, UnicodeError): + raise HTTPException(status_code=400, detail={"error": {"message": "请求体必须是 JSON 对象", "type": "invalid_request_error"}}) from None + if not isinstance(body, dict): + raise HTTPException(status_code=400, detail={"error": {"message": "请求体必须是 JSON 对象", "type": "invalid_request_error"}}) + dst_dir = managed_auth_dir().resolve() + import_dir = Path(os.environ.get("CODEBUDDY_IMPORT_DIR") or dst_dir / "imports") + try: + name, content = read_import_file(import_dir, body.get("path")) + cred_data = auth_oauth.loads_strict(content.decode("utf-8")) + src_uid, verr = auth_oauth.validate_cred_data(cred_data) + if verr: + raise CredentialFileError("凭据格式或站点校验失败") + if (not isinstance(cred_data.get("account") or {}, dict) + or not isinstance(cred_data["auth"].get("expiresAt", 0), (int, float))): + raise CredentialFileError("凭据账号或过期时间格式无效") + # Normalize token aliases before persistence. + content = json.dumps(auth_oauth.normalize_cred_data(cred_data), ensure_ascii=False).encode("utf-8") + except CredentialFileError: + raise HTTPException(status_code=400, detail={"error": {"message": "凭据文件不符合导入要求", "type": "invalid_request_error"}}) from None + except (ValueError, UnicodeError, RecursionError): + raise HTTPException(status_code=400, detail={"error": {"message": "凭据必须是有效的 UTF-8 JSON 对象", "type": "invalid_request_error"}}) from None + except OSError: + raise HTTPException(status_code=400, detail={"error": {"message": "导入目录或文件不可读", "type": "invalid_request_error"}}) from None + try: + dst = _store_credential(dst_dir, name, content, src_uid) + except CredentialConflictError: + raise HTTPException(status_code=409, detail={"error": {"message": "该账号已在池中,请使用同文件名更新或先移除旧凭据", "type": "invalid_request_error"}}) from None + except CredentialFileError: + raise HTTPException(status_code=400, detail={"error": {"message": "凭据文件名或保存目标不符合要求", "type": "invalid_request_error"}}) from None + except OSError: + raise HTTPException(status_code=500, detail={"error": {"message": "凭据保存失败", "type": "server_error"}}) from None + return {"imported": str(dst)} + + +@app.delete("/admin/credentials/{name}") +def admin_del_credential(name: str, + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Delete the named .info file and remove its credential from the pool.""" + _check_admin_auth(authorization, x_api_key) + pool = CONFIG.get("cred_pool") + if CONFIG.get("management") is not None: + CONFIG["management"].admin_delete_guard(os.path.basename(name)) + if pool is None or not pool.remove_file(os.path.basename(name)): + raise HTTPException(status_code=404, detail={"error": {"message": f"凭据不在池中: {name}", "type": "invalid_request_error"}}) + return {"removed": os.path.basename(name)} + + + +def _save_oauth_credential(cred: dict) -> Path: + """Update credentials by product, account and tenant without crossing identities.""" + uid, error = auth_oauth.validate_cred_data(cred) + if error: + raise CredentialFileError("凭据格式或站点校验失败") + dst_dir = managed_auth_dir() + identity = _credential_identity(cred) + profile = profile_for_auth(cred["auth"]) + target = next((f for f in sorted(dst_dir.glob("*.info")) if _cred_identity(f) == identity), None) + existing = None + if target is not None: + try: + existing = json.loads(target.read_text(encoding="utf-8")) + except (OSError, ValueError): + pass + name = target.name if target is not None else f"{uid}.info" + if target is None and (dst_dir / name).exists(): + name = f"{uid}-{profile}.info" + if (dst_dir / name).exists(): + name = f"{uid}-{profile}-{identity}.info" + if (dst_dir / name).exists(): + raise CredentialConflictError("OAuth 保存目标已被其他身份占用") + cred = auth_oauth.merge_existing_accounts(cred, existing) + return _store_credential( + dst_dir, name, json.dumps(cred, ensure_ascii=False, indent=2).encode("utf-8"), uid, + replace_identity=False) + + +@app.post("/admin/oauth/start") +def admin_oauth_start(site: str = "cn", + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Start OAuth and return the browser authorization URL.""" + _check_admin_auth(authorization, x_api_key) + try: + return _OAUTH.start(site=site) + except ValueError as e: + raise HTTPException(status_code=400, detail={"error": {"message": str(e), "type": "invalid_request_error"}}) + except Exception as e: + raise HTTPException(status_code=502, detail={"error": {"message": f"发起失败: {e}", "type": "upstream_error"}}) + + +@app.get("/admin/oauth/poll") +def admin_oauth_poll(login_id: str = "", + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Poll OAuth and persist completed logins with pool reload.""" + _check_admin_auth(authorization, x_api_key) + try: + r = _OAUTH.poll(login_id) + except Exception as e: + raise HTTPException(status_code=502, detail={"error": {"message": f"轮询失败: {e}", "type": "upstream_error"}}) + if not r.get("done"): + return {"done": False} + cred = r.get("cred") + if r.get("error") or not cred: + return {"done": True, "error": r.get("error") or "登录失败"} + uid = r["uid"] + try: + target = _save_oauth_credential(cred) + except CredentialFileError: + return {"done": True, "error": "凭据格式、站点或保存目标不符合要求"} + except OSError: + raise HTTPException(status_code=500, detail={"error": {"message": "凭据保存失败", "type": "server_error"}}) from None + _log(f"[oauth] 无感登录已入库: {r.get('nickname') or uid} ({uid}) -> {target.name}") + return {"done": True, "uid": uid, "nickname": r.get("nickname") or "", "imported": str(target)} + +@app.get("/admin/credits") +def admin_credits(authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Return cached balances, credit expiry segments and check-in status.""" + _check_admin_auth(authorization, x_api_key) + ledger = CONFIG.get("ledger") + return {"credits": ledger.snapshot() if ledger else {}} + + +@app.get("/admin/model-blocks") +def admin_model_blocks(authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Return unsupported backend/model pairs and their retry deadlines.""" + _check_admin_auth(authorization, x_api_key) + pool = CONFIG.get("cred_pool") + return {"model_blocks": pool.model_blocks_detail() if pool is not None else []} + + +@app.post("/admin/checkin") +def admin_checkin(authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Run daily-idempotent manual check-in without implicitly syncing balances or usage.""" + _check_admin_auth(authorization, x_api_key) + return _admin_credential_action("checkin") + + +def _admin_credential_action(action, identity=None, *, consent_revision=None): + from app.credential_actions import run + return run(sys.modules[__name__], action, identity, consent_revision=consent_revision) + + +@app.post("/admin/sync") +def admin_sync(authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + _check_admin_auth(authorization, x_api_key) + return _admin_credential_action("sync") + + +@app.post("/admin/credentials/{identity}/{action}") +async def admin_credential_action(identity: str, action: str, request: Request, + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + _check_admin_auth(authorization, x_api_key) + from app.admin_api import _body + try: + body = await _body(request, 4096, allow_empty=True) + except ValueError: + raise HTTPException(400, "请求体必须是有效 JSON 对象") from None + if body: + if action != "travel": + raise HTTPException(400, "该操作不接受请求体") + if (set(body) != {"confirm_buddy", "agreement_revision"} + or body["confirm_buddy"] is not True or not isinstance(body["agreement_revision"], str)): + raise HTTPException(400, "首领确认参数无效") + return await run_in_threadpool(_admin_credential_action, action, identity, + consent_revision=body.get("agreement_revision")) + + +# --------------------------------------------------------------------------- +# OpenAI-compatible billing estimates +# --------------------------------------------------------------------------- + +def _billing_totals() -> dict: + """Convert regional balances independently, using official usage or a quota-difference fallback.""" + empty_grp = {"remaining": 0.0, "used_by_quota": 0.0, "soonest_expiry": None} + ledger = CONFIG.get("ledger") + snap = ledger.snapshot() if ledger else {} + price_cny = CONFIG.get("credit_price_cny") + price_usd = CONFIG.get("credit_price_usd") + rate = CONFIG.get("usd_rate") + if price_cny is None: + price_cny = credits_mod.CREDIT_PRICE_CNY if credits_mod else 0.014 + if price_usd is None: + price_usd = credits_mod.CREDIT_PRICE_USD if credits_mod else 0.03 + if rate is None: + rate = credits_mod.USD_RATE_CNY if credits_mod else 7.15 + agg = (credits_mod.aggregate_credits(snap) if credits_mod else + {"remaining": 0.0, "used_by_quota": 0.0, "soonest_expiry": None, + "groups": {"domestic": dict(empty_grp), "international": dict(empty_grp)}}) + cache = CONFIG.get("usage_daily") or {} + detail = bool(cache.get("fetched_at")) + detail_groups = cache.get("groups") or {} + per_usd = {"domestic": price_cny / rate, "international": price_usd} + remaining = used = rem_usd = used_usd = rem_cny = 0.0 + groups_out: dict = {} + for grp, g in (agg.get("groups") or {}).items(): + r = float(g.get("remaining") or 0) + gd = detail_groups.get(grp) + u = (float(gd.get("total_credits") or 0) if (detail and gd) + else float(g.get("used_by_quota") or 0)) # Fall back to the group's quota difference. + unit = per_usd.get(grp, 0.0) + remaining += r + used += u + rem_usd += r * unit + used_usd += u * unit + rem_cny += r * (unit * rate if grp == "international" else price_cny) + groups_out[grp] = {"credits_remaining": round(r, 2), "credits_used": round(u, 2), + "balance_usd": round(r * unit, 4), + "price_usd_per_credit": round(unit, 6)} + return {"remaining": round(remaining, 2), "used": round(used, 2), + "quota": round(remaining + used, 2), + "remaining_usd": round(rem_usd, 4), "used_usd": round(used_usd, 4), + "quota_usd": round(rem_usd + used_usd, 4), + "remaining_cny": round(rem_cny, 4), + "soonest_expiry": agg.get("soonest_expiry"), + "price_cny": price_cny, "price_usd": price_usd, "rate": rate, + "used_source": "official_usage_detail" if detail else "quota_delta", + # Expose incomplete pagination or failed account synchronization. + "partial": bool(agg.get("partial") or cache.get("partial")), + "groups": groups_out, "by_day": cache.get("by_day") or {}} + + +@app.get("/v1/dashboard/billing/subscription") +def billing_subscription(authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Expose subscription estimates with balance equal to the hard limit minus usage.""" + _check_auth(authorization, x_api_key) + t = _billing_totals() + limit = t["quota_usd"] + return { + "object": "billing_subscription", + "has_payment_method": True, "canceled": False, "canceled_at": None, "delinquent": None, + # Conservatively use the earliest credit expiry as the access deadline. + "access_until": int(t["soonest_expiry"] or (time.time() + 30 * 86400)), + "soft_limit": int(limit * 100), "hard_limit": int(limit * 100), + "soft_limit_usd": limit, "hard_limit_usd": limit, "system_hard_limit_usd": limit, + "plan": {"title": f"CodeBuddy Credits (CN {t['price_cny']:g} CNY/credit · " + f"INTL {t['price_usd']:g} USD/credit)"}, + # Include total and regional estimates as optional response fields. + "codebuddy_credits_remaining": t["remaining"], + "codebuddy_credits_used": t["used"], + "codebuddy_balance_usd": t["remaining_usd"], + "codebuddy_balance_cny": t["remaining_cny"], + "codebuddy_sites": t["groups"], + # Expose incomplete balance or usage data to callers. + "codebuddy_partial": t["partial"], + **({"codebuddy_stale_accounts": stale} if (stale := (CONFIG.get("usage_daily") or {}).get("stale_accounts")) else {}), + } + + +@app.get("/v1/dashboard/billing/usage") +def billing_usage(start_date: Optional[str] = None, end_date: Optional[str] = None, + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Return estimated usage in cents, with daily model costs for the last 30 days.""" + _check_auth(authorization, x_api_key) + t = _billing_totals() + # Convert each region at its own rate before combining daily and total usage. + cents = {"domestic": t["price_cny"] / t["rate"] * 100, "international": t["price_usd"] * 100} + detail = CONFIG.get("usage_daily") or {} + priced: dict = {} + for site, group in (detail.get("groups") or {}).items(): + unit = cents.get(site) + if unit is None: + continue + for day, models in (group.get("by_day") or {}).items(): + slot = priced.setdefault(day, {}) + for model, credit in models.items(): + slot[model] = slot.get(model, 0.0) + float(credit) * unit + daily = [] + for day in sorted(priced): + if start_date and day < start_date: + continue + if end_date and day > end_date: + continue + items = [{"name": m, "cost": round(c, 4)} + for m, c in sorted(priced[day].items()) if c > 0] + try: + ts = int(time.mktime(time.strptime(day, "%Y-%m-%d"))) + except ValueError: + ts = 0 + daily.append({"timestamp": ts, "line_items": items}) + if start_date or end_date: # Sum only the requested interval. + total_cents = round(sum(sum(i["cost"] for i in d["line_items"]) for d in daily), 2) + else: # Preserve the subscription balance identity. + total_cents = round(t["used_usd"] * 100, 2) + out = {"object": "list", "total_usage": total_cents, "daily_costs": daily} + if t.get("partial"): + out["partial"] = True + if detail.get("stale_accounts"): + out["stale_accounts"] = detail["stale_accounts"] + return out + + +# Prefer cloud catalogs over static fallback models. +_MODEL_TABLE_TTL = 60.0 # Model snapshot lifetime in seconds +_model_table_cache: dict = {} + + +def invalidate_model_table() -> None: + """Invalidate the public model snapshot after catalog changes.""" + global _model_table_cache + _model_table_cache = {} + + +def _catalog_for(profile: str, scope: str = "models", *, model_id=None): + """Return a profile catalog within the requested account scope.""" + accounts = CONFIG.get("account_catalogs") + if accounts is not None or CONFIG.get("model_cache") is not None: + pool = CONFIG.get("cred_pool") + models = None + for entry in pool.entries() if pool is not None else []: + if entry.get("profile") != profile: + continue + account = (accounts or {}).get(entry.get("account_key")) or {} + items = _effective_account_scope(account, scope, model_id=model_id) + if account.get("profile") == profile and items is not None: + if models is None: + models = [] + models.extend(items) + return models + catalogs = CONFIG.get("model_catalogs") or {} + if profile in catalogs: + return catalogs[profile] + legacy = {"cn-cli": "models_remote", "intl-cli": "models_intl"} + return CONFIG.get(legacy[profile]) if profile in legacy else None + + +def _in_region(profile: str, region: str | None) -> bool: + return region is None or profile_region(profile) == region + + +def _configured_profiles(region: str | None) -> set[str]: + pool = CONFIG.get("cred_pool") + if pool is not None: + return {profile for entry in pool.entries() if (profile := pool._entry_profile(entry)) + and model_policy.credential_enabled(CONFIG, entry) and _in_region(profile, region)} + cm = CONFIG.get("cred") + if cm is not None: + profile = cm.summary()["profile"] + return {profile} if _in_region(profile, region) else set() + known = {profile for profile in PROFILE_ENDPOINTS + if _in_region(profile, region) and _catalog_for(profile) is not None} + return known or ({"intl-cli"} if region == "intl" else {"cn-cli"}) + + +def _usable_models(models): + return [model for model in models or [] if model.get("id") and model.get("supportsToolCall") + and not model.get("disabled")] + + +def _account_scope(account: dict, scope: str = "models") -> list[dict] | None: + """Merge root candidates into selector models while retaining selector metadata.""" + picker = account.get("models") + if scope == "models" or picker is None: + return picker + seen = {item.get("id") for item in picker} + # Root entries supply missing names without replacing selector metadata. + return picker + [item for item in account.get("serves") or [] if item.get("id") not in seen] + + +def _shared_catalog_sources(scope="models"): + accounts, pool = CONFIG.get("account_catalogs"), CONFIG.get("cred_pool") + if accounts is not None or CONFIG.get("model_cache") is not None: + return [(entry["profile"], _account_scope(account, scope)) + for entry in (pool.entries() if pool is not None else []) + if entry.get("profile") in SHARED_INTL_PROFILES and model_policy.credential_enabled(CONFIG, entry) + and (account := (accounts or {}).get(entry.get("account_key"))) + and account.get("profile") == entry["profile"] and account.get("models") is not None] + configured = _configured_profiles(None) + return [(profile, _catalog_for(profile, scope)) for profile in sorted(configured & SHARED_INTL_PROFILES)] + + +def _effective_account_scope(account, scope="models", *, model_id=None): + native = _account_scope(account, scope) + profile = account.get("profile") + if profile not in SHARED_INTL_PROFILES or native is None: + return native + return share_models(native, profile, _shared_catalog_sources(scope), model_id=model_id) + + + +def _models_for_profile(profile: str, configured=None, *, scope: str = "models", model_id=None) -> list[dict]: + models = _catalog_for(profile, scope, model_id=model_id) + if models is None: + # Static fallback is limited to legacy domestic CLI deployments. + configured = _configured_profiles(profile_region(profile)) if configured is None else configured + return ([{"id": name, "supportsToolCall": True} for name in DEFAULT_MODELS] + if CONFIG.get("model_cache") is None and CONFIG.get("account_catalogs") is None + and profile == "cn-cli" and configured <= {"cn-cli"} else []) + if profile in SHARED_INTL_PROFILES and CONFIG.get("account_catalogs") is None and CONFIG.get("model_cache") is None: + models = share_models(models, profile, _shared_catalog_sources(scope), model_id=model_id) + return _usable_models(models) + + +def _upstream_model(model: str | None, profile: str) -> str | None: + return "default-model" if model == "auto" and profile_region(profile) == "intl" else model + + +def _free_multiplier(credits) -> bool: + """Check whether the account's catalog explicitly declares a zero credit rate.""" + if not isinstance(credits, str): + return False + match = re.fullmatch(r"x\s*0(?:\.0+)?\s*(?:credits?)?", credits.strip(), re.IGNORECASE) + return match is not None + + +def _multiplier_value(credits): + """Parse an official model rate, returning None for missing or unknown formats.""" + if not isinstance(credits, str): + return None + match = re.fullmatch(r"x\s*([0-9]+(?:\.[0-9]+)?)\s*(?:credits?)?", credits.strip(), re.IGNORECASE) + if match is None: + return None + try: + return float(match.group(1)) + except ValueError: + return None + + +def _model_free(models, model: str | None, profile: str) -> bool: + """Require an explicit zero-rate entry in this account's model catalog.""" + if not model: + return False + routed = _upstream_model(model, profile) + return any(item.get("id") == routed and _free_multiplier(item.get("credits")) + for item in models or []) + + +def _model_profiles(model: str | None, region: str | None = None, configured=None) -> set[str]: + configured = _configured_profiles(region) if configured is None else configured + profiles = {profile for profile in PROFILE_ENDPOINTS if _in_region(profile, region)} + if not model: + return profiles + supported = {profile for profile in profiles + if any(item["id"] == _upstream_model(model, profile) + for item in _models_for_profile(profile, configured, scope="serves", model_id=_upstream_model(model, profile)))} + if model == "auto" and region == "cn": + # WorkBuddy uses its advertised Auto; legacy CLI defaults remain separate. + if "cn-work" in configured and "cn-work" in supported: + return {"cn-work"} + if "cn-cli" in configured and _models_for_profile("cn-cli", configured): + return {"cn-cli"} + if model == "auto" and region is None and "cn-cli" in configured and _models_for_profile("cn-cli", configured): + supported.add("cn-cli") + if model != "auto" and not supported and not CONFIG.get("model_guard") and len(configured) == 1: + return configured + return supported + + +def _catalog_pending(region: str | None = None) -> bool: + pool = CONFIG.get("cred_pool") + if CONFIG.get("model_cache") is None and CONFIG.get("account_catalogs") is None: + return False + if pool is not None and pool.sync_pending(region): + return True + return not any(_catalog_for(profile) is not None for profile in _configured_profiles(region)) + + +def _profile_has_credits(profile: str) -> bool: + pool = CONFIG.get("cred_pool") + if pool is None: + if profile_region(profile) == "cn": + return True + ledger = CONFIG.get("ledger") + return bool(ledger and any((entry.get("credits") or {}).get("intl") + and float((entry.get("credits") or {}).get("credits") or 0) > 0 + for entry in ledger.snapshot().values())) + return any(pool._entry_profile(entry) == profile and pool._has_credit(entry, profile) + for entry in pool.entries() if model_policy.credential_enabled(CONFIG, entry)) + + +def current_models(region: str | None = None) -> list[str]: + """Merge models eligible for current accounts without region-specific client URLs.""" + pool = CONFIG.get("cred_pool") + if pool is not None: + pool._rescan() + with pool._lock if pool is not None else nullcontext(): + configured = _configured_profiles(region) + out, has_auto = [], False + auto_profiles = _model_profiles("auto", region, configured) + if pool is not None and (CONFIG.get("account_catalogs") is not None or CONFIG.get("model_cache") is not None): + accounts = CONFIG.get("account_catalogs") or {} + for entry in pool.entries(): + if not model_policy.credential_enabled(CONFIG, entry): + continue + profile = entry.get("profile") + if not profile or not _in_region(profile, region): + continue + # Publish only advertised zero-rate models for empty accounts. + zero = pool._zero_balance(entry, profile) + if not zero and not pool._has_credit(entry, profile): + continue + account = accounts.get(entry.get("account_key")) or {} + if account.get("profile") != profile: + continue + models = _usable_models(_effective_account_scope(account, "serves")) + if zero: + models = [model for model in models if _free_multiplier(model.get("credits"))] + out.extend(model["id"] for model in models) + if profile in auto_profiles and models: + has_auto |= profile == "cn-cli" or any(model["id"] == _upstream_model("auto", profile) for model in models) + else: + for profile in sorted(configured): + entries = ([entry for entry in pool.entries() if pool._entry_profile(entry) == profile] + if pool is not None else []) + # Empty products may publish only their advertised zero-rate models. + zero_only = bool(entries) and all(pool._zero_balance(entry, profile) for entry in entries) + if _profile_has_credits(profile) or zero_only: + models = _models_for_profile(profile, configured) + if zero_only: + models = [model for model in models if _free_multiplier(model.get("credits"))] + out.extend(model["id"] for model in models) + has_auto |= bool(models) and profile in auto_profiles + if has_auto: + out.append("auto") + return list(dict.fromkeys(out)) + + +def current_model_details(region: str | None = None) -> list[dict]: + """Return public model rates with per-profile values and the minimum eligible rate.""" + pool = CONFIG.get("cred_pool") + if pool is not None: + pool._rescan() + details: dict[str, dict] = {} + declarations = {} + for name in current_models(region): + details[name] = {"id": name, "credits": None, "credits_by_profile": {}} + if pool is None: + return [{**item, **model_capabilities.describe_models(())} for item in details.values()] + with pool._lock: + def record(profile: str, item: dict, *, zero: bool) -> None: + name = item.get("id") + if name not in details: + return + if zero and not _free_multiplier(item.get("credits")): + return # Empty accounts cannot supply paid model rates. + declarations.setdefault(name, []).append((profile, item)) + value = _multiplier_value(item.get("credits")) + if value is None: + return + details[name]["credits_by_profile"][profile] = value + best = details[name]["credits"] + details[name]["credits"] = value if best is None else min(best, value) + + if CONFIG.get("account_catalogs") is not None or CONFIG.get("model_cache") is not None: + accounts = CONFIG.get("account_catalogs") or {} + for entry in pool.entries(): + if not model_policy.credential_enabled(CONFIG, entry): + continue + profile = entry.get("profile") + if not profile or not _in_region(profile, region): + continue + zero = pool._zero_balance(entry, profile) + if not zero and not pool._has_credit(entry, profile): + continue + account = accounts.get(entry.get("account_key")) or {} + if account.get("profile") != profile: + continue + for item in _usable_models(_effective_account_scope(account, "serves")): + record(profile, item, zero=zero) + else: + configured = _configured_profiles(region) + for profile in sorted(configured): + entries = [entry for entry in pool.entries() if pool._entry_profile(entry) == profile] + zero_only = bool(entries) and all(pool._zero_balance(entry, profile) for entry in entries) + if not (_profile_has_credits(profile) or zero_only): + continue + for item in _models_for_profile(profile, configured): + record(profile, item, zero=zero_only) + return [{**item, **model_capabilities.describe_models(declarations.get(item["id"], []))} for item in details.values()] + + +def _client_wants_stream(payload: dict) -> bool: + """Default stream to false and reject non-Boolean values.""" + value = payload.get("stream", False) + if not isinstance(value, bool): + raise HTTPException(status_code=400, detail={"error": { + "message": "stream must be a boolean", "type": "invalid_request_error", "param": "stream"}}) + return value + + +def _prepare_payload(payload, field="messages") -> dict: + """Apply request-wide image limits before adaptation, logging and credential selection.""" + if not isinstance(payload, dict): + raise HTTPException(status_code=400, detail={"error": { + "message": "请求体必须是 JSON 对象", "type": "invalid_request_error"}}) + try: + prepared, stats = apply_image_policy( + payload, field=field, max_images=CONFIG["max_images"], policy=CONFIG["image_policy"]) + except ImageLimitError as error: + _log(f"[limit] 图片超限,拒绝请求 | count={error.count} | limit={error.limit}") + raise HTTPException(status_code=413, detail={"error": { + "message": str(error), "type": "invalid_request_error", "param": field, + "code": "too_many_images", "image_count": error.count, "max_images": error.limit}}) from None + if stats["dropped"]: + _log(f"[limit] 保留最新图片 | count={stats['count']} | retained={stats['retained']} | dropped={stats['dropped']}") + return prepared + + +def _normalize_tool_choice(body): + """Map named tool choice to a single required tool for string-only upstream selection.""" + choice = body.get("tool_choice") + if not isinstance(choice, dict): + return + function = choice.get("function", choice) + name = function.get("name") if isinstance(function, dict) else None + tools = body.get("tools") + matches = [tool for tool in tools if isinstance(tool, dict) and tool.get("type") == "function" + and isinstance(tool.get("function"), dict) and tool["function"].get("name") == name] if isinstance(tools, list) else [] + if choice.get("type") != "function" or not isinstance(name, str) or not name.strip() or len(matches) != 1: + raise HTTPException(status_code=400, detail={"error": {"message": "tool_choice must name exactly one declared function", + "type": "invalid_request_error", "param": "tool_choice"}}) + body["tools"], body["tool_choice"] = matches, "required" + + +def _bind_request_session(payload, body): + context = current_context() + if context is not None and context.scoped: + try: + context.bind_session(payload, body.get("messages")) + except SessionIdentifierError as error: + raise HTTPException(status_code=400, detail={"error": {"message": str(error), + "type": "invalid_request_error", "param": "session_id"}}) from None + except (ValueError, TypeError, UnicodeError, RecursionError): + raise HTTPException(status_code=400, detail={"error": {"message": "invalid session input", + "type": "invalid_request_error"}}) from None + + +def _request_id(): + context = current_context() + return context.request_id if context is not None else uuid.uuid4().hex + + +def _prepare_chat_body(body: dict, *, region=None, session_payload=None) -> dict: + """Normalize models, system messages, streaming, desensitization and payload budgets.""" + if session_payload is not None: + _bind_request_session(session_payload, body) + body = dict(body) + body["model"] = model_policy.resolve(CONFIG, body.get("model", "auto")) + guard_model(body["model"], region=region, resolved=True) + messages = body.get("messages") + if not isinstance(messages, list) or not messages or any(not isinstance(message, dict) for message in messages): + raise HTTPException(status_code=400, detail={"error": { + "message": "messages must be a non-empty array of objects", "type": "invalid_request_error"}}) + # Keep error paths tied to caller positions when the upstream system message moves. + message_indices = list(range(len(messages))) + # Upstreams reject developer roles; copy them as system messages without changing content. + messages = [ + dict(message, role="system") if message.get("role") == "developer" else message + for message in messages + ] + if messages[0].get("role") != "system": + system_index = next((index for index, message in enumerate(messages) if message.get("role") == "system"), None) + if system_index is None: + messages = [{"role": "system", "content": "You are a helpful assistant."}, *messages] + message_indices.insert(0, None) + else: + messages = [messages[system_index], *messages[:system_index], *messages[system_index + 1:]] + message_indices.insert(0, message_indices.pop(system_index)) + body["messages"] = normalize_chat_messages(messages, message_indices=message_indices) + _normalize_tool_choice(body) + body["stream"] = True + body.setdefault("stream_options", {"include_usage": True}) + body = _chat_body_desensitize(body) + _guard_request_size(body) + return body + + +def _guard_request_size(body: dict) -> int: + """Validate and measure upstream JSON bytes without truncating text or tool arguments.""" + size = 0 + limit = CONFIG["max_request_bytes"] + try: + for part in json.JSONEncoder(ensure_ascii=False, separators=(",", ":"), allow_nan=False).iterencode(body): + size += len(part.encode("utf-8")) + if size > limit: + _log(f"[limit] 请求体超限,拒绝请求 | limit_bytes={limit}") + raise HTTPException(status_code=413, detail={"error": { + "message": f"处理后的请求体超过网关上限 {limit} 字节,请缩短历史或压缩图片", + "type": "invalid_request_error", "code": "request_too_large", "max_bytes": limit}}) + except (ValueError, UnicodeError) as error: + raise HTTPException(status_code=400, detail={"error": { + "message": "请求体包含无法序列化的 JSON 值", "type": "invalid_request_error"}}) from None + return size + + +def guard_model(name: str, *, region=None, resolved=False) -> None: + """Reject unauthorized models and route only through accounts with confirmed support.""" + if not isinstance(name, str) or not name.strip(): + raise HTTPException(status_code=400, detail={"error": { + "message": "model must be a non-empty string", "type": "invalid_request_error", "param": "model"}}) + if not resolved: + name = model_policy.resolve(CONFIG, name) + model_policy.check_resolved(CONFIG, name) + if not CONFIG.get("model_guard"): + return + pool = CONFIG.get("cred_pool") + if pool is not None: + pool._rescan() + if _model_profiles(name, region): + return + if _catalog_pending(region): + raise HTTPException(status_code=503, headers={"Retry-After": "3"}, detail={"error": { + "message": "模型目录正在同步,请稍后重试", "type": "service_unavailable", "code": "catalog_syncing"}}) + raise HTTPException(status_code=404, detail={"error": { + "message": f"The model '{name}' is not supported by this gateway. See GET /v1/models.", + "type": "invalid_request_error", "param": "model", "code": "model_not_found"}}) + + + +@app.get("/v1/models") +def list_models(authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + _check_auth(authorization, x_api_key) + data = [{"id": item["id"], "object": "model", "created": 1700000000, "owned_by": "codebuddy", + "credits": item["credits"], "credits_by_profile": item["credits_by_profile"], + **{key: item[key] for key in ("capabilities", "limits", "metadata_by_profile") if key in item}} + for item in model_policy.public_details(sys.modules[__name__])] + return {"object": "list", "data": data} + + +@app.post("/v1/chat/completions") +async def chat_completions(request: Request, + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + _check_auth(authorization, x_api_key) + # Select sticky credentials while building upstream headers. + + try: + payload = await request.json() + except Exception as e: + raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}}) + + payload = _prepare_payload(payload) + # Aggregation supports exactly one completion, so reject other n values. + n_value = payload.get("n") + if n_value is not None and not (isinstance(n_value, int) and not isinstance(n_value, bool) and n_value == 1): + raise HTTPException(status_code=400, detail={"error": { + "message": "only n=1 is supported: multiple candidates would be merged into one answer", + "type": "invalid_request_error", "param": "n"}}) + messages = payload.get("messages") or [] + if not messages: + raise HTTPException(status_code=400, detail={"error": {"message": "messages is required", "type": "invalid_request_error"}}) + + # Forward only supported request fields. + client_wants_stream = _client_wants_stream(payload) + body = {k: payload[k] for k in PASSTHROUGH_BODY_KEYS if k in payload} + body = await run_in_threadpool(_prepare_chat_body, body, session_payload=payload) + stream_policy = _snapshot_stream_policy("chat", body) + observe_stream_mode(stream_policy.mode) + + # Record request metadata. + model_name = payload.get("model", "auto") + tool_names = [t.get("function", {}).get("name") for t in (payload.get("tools") or []) + if isinstance(t, dict)] + last_user = _last_user_text(messages) + rid = _request_id() + _log(f"[{rid}] ▶ REQUEST {model_name} | stream={client_wants_stream} | msgs={len(messages)}" + + (f" | tools={tool_names}" if tool_names else "") + + (f" | last_user={_truncate(last_user, 60)!r}" if last_user else "")) + # Credential selection and refresh perform blocking file and network I/O. + prepared = body # Keep canonical input for failover policy checks. + body, cred, headers, url = await run_in_threadpool(_route_chat, payload, body, rid) + _log_json(f"[{rid}] REQUEST BODY (发往后端,预览)", body) + t0 = time.time() + + if client_wants_stream: + def attempt(routed, cred, headers, url): + return _stream_upstream(url, headers, _body_with_stream_policy(routed, stream_policy), + model_name, t0, rid, cred=cred) + return _routed_stream(payload, prepared, model_name, rid, t0, attempt, + body, cred, headers, url) + + # Aggregate upstream SSE for non-streaming clients. + async def fetch(routed, cred, headers, url): + return await _fetch_checked_chat(url, headers, routed, model_name, rid, cred, + filter_retry=True, max_collect_bytes=stream_policy.max_collect_bytes) + # Watch for disconnects across the entire failover sequence. + try: + collected = await await_or_hangup( + _routed_fetch(payload, prepared, model_name, rid, t0, fetch, + body, cred, headers, url), request) + except ClientHungUp: + return _hungup_response(rid, model_name, t0) + _log_finish(model_name, t0, collected, rid) + if CONFIG.get("control_store") is not None: + collected = {**collected, "model": model_name} + return JSONResponse(content=collected) + + +def _last_user_text(messages: list) -> str: + """Extract the latest user text for bounded log previews.""" + for m in reversed(messages): + if m.get("role") != "user": + continue + content = m.get("content", "") + if isinstance(content, list): + for blk in content: + if isinstance(blk, dict) and blk.get("type") == "text": + return str(blk.get("text", "")) + return "" + return str(content) + return "" + + +def _log_finish(model_name: str, t0: float, result: dict, rid: str = ""): + """Log request timing, finish reason, usage, tools and bounded response previews.""" + elapsed = time.time() - t0 + prefix = f"[{rid}] " if rid else "" + choice = (result.get("choices") or [{}])[0] + finish = choice.get("finish_reason") + msg = choice.get("message") or {} + detector = ContentFilterDetector() + detector.feed(msg, finish) + if detector.detected: + return # Filtered responses must not expose echoed content in previews. + tcs = msg.get("tool_calls") or [] + usage = result.get("usage") or {} + tc_names = [t.get("function", {}).get("name") for t in tcs] + _log(f"{prefix}◀ RESPONSE {model_name} | {elapsed:.1f}s | finish={finish}" + + (f" | tool_calls={tc_names}" if tc_names else "") + + f" | tokens={usage.get('total_tokens', '?')}") + _log_json(f"{prefix}RESPONSE BODY (预览)", result) + + +def _chat_completion(merged: dict) -> dict: + message = {"role": "assistant", "content": merged["content"] or None} + for key in ("reasoning_content", "refusal", "tool_calls"): + if merged.get(key): + message[key] = merged[key] + return { + "id": "chatcmpl-" + os.urandom(12).hex(), "object": "chat.completion", + "created": int(time.time()), "model": merged.get("model") or "unknown", + "choices": [{"index": 0, "message": message, + "finish_reason": merged.get("finish_reason") or ("tool_calls" if merged.get("tool_calls") else "stop")}], + "usage": merged.get("usage") or {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}, + } + + +def _completion_to_merged(result: dict) -> dict: + choice = result["choices"][0] + return {**choice["message"], "finish_reason": choice["finish_reason"], + "model": result.get("model"), "usage": result.get("usage")} + + +async def _collect_stream(response: httpx.Response, *, accumulator=None) -> dict: + """Collect content, reasoning and tools while validating stream completion.""" + accumulator = accumulator if accumulator is not None else ChatSSEAccumulator() + async for line in response.aiter_lines(): + accumulator.feed_line(line) + if accumulator.done: + break + return _chat_completion(accumulator.result()) + + +_TOOL_CALL_MAX_RETRY = 3 + + +def _tool_choice_satisfied(tool_calls, body): + calls = tool_calls or [] + choice = body.get("tool_choice") + if choice == "none": + return not calls + if choice == "required": + names = _declared_tool_names(body) + return bool(calls) and all(call.get("function", {}).get("name") in names for call in calls) + if isinstance(choice, dict): + name = (choice.get("function") or {}).get("name") if isinstance(choice.get("function"), dict) else None + return bool(calls) and all(call.get("function", {}).get("name") == name for call in calls) + return True + + +def _declared_tool_names(body: dict) -> frozenset[str]: + """Extract declared function names for conservative realtime identity boundaries.""" + names = set() + for tool in body.get("tools") or []: + if not isinstance(tool, dict): + continue + function = tool.get("function") + if isinstance(function, dict) and isinstance(function.get("name"), str): + names.add(function["name"]) + elif isinstance(tool.get("name"), str): + names.add(tool["name"]) + return frozenset(name for name in names if name) + + +def _tool_calls_healthy(tool_calls, body: dict | None = None, *, require_declarations=False) -> bool: + """Validate unique calls, declared names and JSON-object arguments.""" + calls = tool_calls or [] + names = set(_declared_tool_names(body)) if body is not None else None + seen_ids = set() + for tc in calls: + if not isinstance(tc, dict): + return False + call_id = tc.get("id") + if (not isinstance(call_id, str) or not call_id.strip() + or (require_declarations and call_id in seen_ids) + or (require_declarations and tc.get("type") != "function")): + return False + if not require_declarations: + seen_ids.add(call_id) + seen_ids.add(call_id) + fn = tc.get("function") or {} + if not isinstance(fn, dict): + return False + name = fn.get("name") or "" + arguments = fn.get("arguments") or "" + if not isinstance(name, str) or not name.strip() or not isinstance(arguments, str) or not arguments.strip(): + return False + if names is not None and ((require_declarations or names) and name not in names): + return False + try: + decoded = json.loads(arguments) + except (TypeError, ValueError, UnicodeError, RecursionError): + return False + if not isinstance(decoded, dict): + return False + if (require_declarations and body is not None + and body.get("parallel_tool_calls") is False and len(calls) > 1): + return False + return True + + +def _validate_realtime_tools(tool_calls, body: dict, finish_reason=None, *, filtered=False): + """Validate terminal state without regenerating or replaying the request. + + Explicit refusal/filter and length results are incomplete upstream outcomes, + not malformed tool calls. Any tool bytes that are present still have to be + structurally healthy, but a required/named choice need not be satisfied when + the upstream legitimately stopped before producing a call. + """ + if not _tool_calls_healthy(tool_calls, body, require_declarations=True): + healthy = False + elif finish_reason is None: + # Realtime success terminals require an explicit upstream completion + # marker; [DONE] alone remains compatible only on the legacy path. + healthy = False + elif filtered or finish_reason in ("length", "content_filter", "content-filter", "refusal"): + healthy = True + else: + healthy = (_tool_choice_satisfied(tool_calls, body) + and (finish_reason == "tool_calls" if tool_calls else finish_reason != "tool_calls")) + if healthy: + return + raw = {"error": {"message": "Invalid upstream tool_calls", "type": "upstream_error", + "code": "invalid_tool_calls"}} + raise UpstreamResponseError(502, json.dumps(raw).encode("utf-8")) + + +def _merge_chat_sse_text(text: str) -> dict: + """Use the shared SSE accumulator for collected text responses.""" + accumulator = ChatSSEAccumulator(max_collect_bytes=CONFIG.get("max_collect_bytes", 0)) + for line in text.splitlines(): + accumulator.feed_line(line) + return accumulator.result() + + +def _chat_result_to_sse_lines(m: dict) -> list[str]: + """Replay collected Chat output as SSE, emitting reasoning before content.""" + content = m.get("content") or "" + reasoning = m.get("reasoning_content") or "" + tcs = m.get("tool_calls") or [] + finish = m.get("finish_reason") or "stop" + model = m.get("model") + # All chunks in a completion share one stable identifier. + completion_id = "chatcmpl-" + os.urandom(12).hex() + created = int(time.time()) + + def _line(delta: dict, fr=None) -> str: + payload = {"id": completion_id, "object": "chat.completion.chunk", "created": created, + "choices": [{"index": 0, "delta": delta, "finish_reason": fr}]} + if model: + payload["model"] = model + return "data: " + json.dumps(payload, ensure_ascii=False) + + lines = [_line({"role": "assistant", "content": ""})] + for i in range(0, len(reasoning), 48): + lines.append(_line({"reasoning_content": reasoning[i:i + 48]})) + for i in range(0, len(content), 48): + lines.append(_line({"content": content[i:i + 48]})) + refusal = m.get("refusal") or "" + for i in range(0, len(refusal), 48): + lines.append(_line({"refusal": refusal[i:i + 48]})) + for i, tc in enumerate(tcs): + lines.append(_line({"tool_calls": [dict(tc, index=i)]})) + lines.append(_line({}, finish)) + if m.get("usage"): + usage_chunk = {"id": completion_id, "object": "chat.completion.chunk", "created": created, + "choices": [], "usage": m["usage"]} + if model: + usage_chunk["model"] = model + lines.append("data: " + json.dumps(usage_chunk, ensure_ascii=False)) + lines.append("data: [DONE]") + return lines + + +def _network_error_text(error: Exception) -> str: + return sanitize_log_text(f"{type(error).__name__}: {str(error).strip() or 'upstream transport failed'}", 512) + +def _public_sse_line(line, model_name): + if CONFIG.get("control_store") is not None and line.startswith("data:"): + try: + event = json.loads(line[5:].strip()) + if isinstance(event, dict) and ("model" in event or "choices" in event): + event["model"] = model_name + return "data: " + json.dumps(event, ensure_ascii=False) + except (ValueError, TypeError): + pass + return line + + +@asynccontextmanager +async def _backend_stream(url, headers, body, *, timeout=300, rid="", model_name="?"): + started, opened = time.monotonic(), False + context = current_context() + if context is not None: + context.attempt = None + + def attempt_headers(): + if context is None: + return dict(headers) + attempt = context.start_attempt() + outgoing = context.attempt_headers(headers, attempt) + profile = profile_for_headers(headers) + observe_attempt("upstream_attempt", profile=profile, upstream_model=body.get("model"), + credential=account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")), + conversation_id=outgoing.get("X-Conversation-ID"), + upstream_request_id=outgoing.get("X-Request-ID")) + return outgoing + + def retry(error): + """Record connection retries and flag possible billing after write timeouts.""" + timeout_on_write = isinstance(error, WRITE_TIMEOUT_TRANSPORT) + observe_attempt("write_timeout_retry" if timeout_on_write else "connect_retry", + error_code=type(error).__name__, + duration_ms=(time.monotonic() - started) * 1000) + _log(f"[{rid}] {'写超时重放' if timeout_on_write else '建连失败'},重试 1/1 | {model_name}" + f" | {_network_error_text(error)}{_replay_cost_note(error)}") + try: + resources = request_resources.get() + clients = resources.clients if resources is not None and CONFIG.get("upstream_keepalive") else None + async with open_backend_stream(url, headers, body, read_timeout=timeout, on_retry=retry, + retry_write_timeout=bool(CONFIG.get("retry_write_timeout")), + clients=clients, headers_for_attempt=attempt_headers) as response: + opened = True + observe_attempt("upstream_http", status_code=response.status_code, + duration_ms=(time.monotonic() - started) * 1000) + yield response + except (httpx.HTTPError, UpstreamResponseError) as error: + if not opened: + observe_attempt("transport", error_code=type(error).__name__, + duration_ms=(time.monotonic() - started) * 1000) + raise + + +def _safe_err_raw(raw: bytes, status: int) -> dict: + try: + return json.loads(raw.decode("utf-8", "replace")) + except Exception: + return {"error": {"message": raw.decode("utf-8", "replace")[:500], "type": "upstream_error", "code": status}} + + +def _check_upstream_status(status, raw, cred, model, *, headers=None): + if status != 200: + retry_after = parse_retry_after((headers or {}).get("Retry-After")) + if not is_filter_error(raw): + _note_cred_status(cred, status, model=model, raw=raw, retry_after=retry_after) + raise UpstreamHTTPError(status, raw, retry_after=retry_after) + + +def _upstream_failure(error, model_name, t0, rid): + """Normalize failure logs and payloads before endpoint-specific error wrapping.""" + if isinstance(error, UpstreamResponseError): + status, raw = error.status, error.raw + category = f"HTTP {status}" + else: + status, raw = 502, _network_error_text(error).encode("utf-8") + category = "网络错误" + if isinstance(error, UpstreamResponseError) and is_filter_error(raw): + _note_content_filter(rid, model_name, final=True) + return status, raw + else: + observe_failure(f"upstream_{status}" if isinstance(error, UpstreamResponseError) else type(error).__name__) + elapsed = time.time() - t0 if t0 else 0 + _log(f"[{rid}] ✗ {category} | {model_name} | {elapsed:.1f}s | {sanitize_log_text(raw.decode('utf-8', 'replace'), 512)}") + _log_text_body(f"[{rid}] ERROR BODY", raw.decode("utf-8", "replace")) + return status, raw + + +def _hungup_response(rid, model_name, t0): + """Finish a disconnected ASGI request with an empty 204; auditing records cancellation.""" + elapsed = time.time() - t0 if t0 else 0 + _log(f"[{rid}] ✂ 下游已断连,取消这次聚合 | {model_name} | {elapsed:.1f}s") + return Response(status_code=204) + + +async def _fetch_checked_chat(url, headers, body, model_name, rid, cred=None, *, filter_retry=False, + max_collect_bytes=None): + """Collect with bounded repair and one eligible filter retry using a frozen request budget.""" + tool_attempt = 0 + filter_retried = False + collection_limit = (CONFIG.get("max_collect_bytes", 0) if max_collect_bytes is None + else max_collect_bytes) + while True: + accumulator = ChatSSEAccumulator(max_collect_bytes=collection_limit) + rejection = None + async with _backend_stream(url, headers, body, rid=rid, model_name=model_name) as response: + if response.status_code != 200: + _check_upstream_status(response.status_code, await read_bounded_error(response), cred, body.get("model"), + headers=response.headers) + else: + _note_cred_model_ok(cred, body.get("model")) + try: + result = await _collect_stream(response, accumulator=accumulator) + except UpstreamResponseError as error: + if (not accumulator.done or not accumulator.filter_detector.detected + or accumulator.saw_output): + raise + rejection = error + result = None + + detector = accumulator.filter_detector + if detector.detected: + retry_body = body + if (filter_retry and not filter_retried and detector.retry_safe + and CONFIG.get("desensitize") and CONFIG.get("no_compact")): + retry_body = _chat_body_desensitize(body, force_compact=True) + try: + if _guard_request_size(retry_body) >= _guard_request_size(body): + retry_body = body + except HTTPException: + retry_body = body + if retry_body != body: + _note_content_filter(rid, model_name, final=False) + body = retry_body + filter_retried = True + continue + if rejection is not None: + raise rejection + _note_content_filter(rid, model_name, final=True) + + calls = result["choices"][0]["message"].get("tool_calls") + if _tool_calls_healthy(calls, body) and (detector.detected or _tool_choice_satisfied(calls, body)): + observe_usage(result.get("usage") or {}) + return result + # Content filtering must not trigger tool-repair regeneration. + budget = CONFIG.get("tool_call_max_retry", _TOOL_CALL_MAX_RETRY) + if detector.detected or not body.get("tools") or tool_attempt >= budget: + if not detector.detected and body.get("tools"): + # Account for the final failed generation before returning an error. + exhausted = result.get("usage") or {} + observe_attempt("tool_args_exhausted", attempt=tool_attempt, max_attempts=budget, + total_tokens=exhausted.get("total_tokens")) + raise UpstreamResponseError(502, b"Invalid upstream tool_calls after retries") + tool_attempt += 1 + # Discarded generations still consume credits and belong in the audit trail. + discarded = result.get("usage") or {} + observe_attempt("tool_args_retry", attempt=tool_attempt, max_attempts=budget, + total_tokens=discarded.get("total_tokens")) + _log(f"[{rid}] tool_calls 损坏,重试 {tool_attempt}/{budget} | {model_name}") + +async def _chat_sse_lines(url, headers, body, model_name, t0, rid, cred=None, *, + policy=None, tracker=None, state=None): + """Yield validated Chat SSE with bounded filter detection and no streaming filter retries.""" + policy = policy or _snapshot_stream_policy("chat", body) + if policy.aggregate: + result = await _fetch_checked_chat( + url, headers, body, model_name, rid, cred, + max_collect_bytes=policy.max_collect_bytes) + for line in _chat_result_to_sse_lines(_completion_to_merged(result)): + yield line + yield "" + _log_finish(model_name, t0, result, rid) + return + if tracker is None: + tracker = ChatSSEAccumulator( + collect=False, retain_tools=policy.realtime, + budget=StreamOutputBudget(policy.max_collect_bytes)) + if state is not None: + state["tracker"] = tracker + + def completed(): + merged = tracker.result(allow_empty_filter=policy.realtime) + if policy.realtime: + _validate_realtime_tools( + merged.get("tool_calls"), body, merged.get("finish_reason"), + filtered=tracker.filter_detector.detected) + if state is not None: + state["merged"] = merged + return merged + + preview = bytearray() + # Realtime output is never copied into the retired raw-text preview log. + budget = 0 if policy.realtime else ( + CONFIG["log_body_limit"] if CONFIG.get("log_path") else 0) + try: + async with _backend_stream(url, headers, body, rid=rid, model_name=model_name) as response: + if response.status_code != 200: + _check_upstream_status(response.status_code, await read_bounded_error(response), cred, body.get("model"), + headers=response.headers) + else: + _note_cred_model_ok(cred, body.get("model")) + async for line in response.aiter_lines(): + tracker.feed_line(line) + if tracker.done or tracker.finish_reason: + completed() # Never expose a success marker before terminal validation. + remaining = budget - len(preview) + if remaining > 0: + preview.extend((line[:remaining] + "\n").encode("utf-8")[:remaining]) + yield line + if tracker.done: + yield "" + break + merged = completed() + finally: + if tracker.usage: + observe_usage(tracker.usage) + if tracker.filter_detector.detected: + _note_content_filter(rid, model_name, final=True) + return + _log(f"[{rid}] ◀ RESPONSE {model_name} | {time.time() - t0:.1f}s | stream finish={merged['finish_reason']}" + + f" | tokens={(merged['usage'] or {}).get('total_tokens', '?')}") + _log_text_body(f"[{rid}] RESPONSE SSE PREVIEW", preview.decode("utf-8", "replace")) + + +async def _stream_upstream(url: str, headers: dict, body: dict, + model_name: str = "?", t0: float = 0.0, rid: str = "", cred=None): + policy = body.pop(_REQUEST_POLICY_KEY, None) or _snapshot_stream_policy("chat", body) + sent = False + upstream = _chat_sse_lines(url, headers, body, model_name, t0, rid, cred, policy=policy) + try: + try: + async for line in upstream: + sent = True + yield (_public_sse_line(line, model_name) + "\n").encode("utf-8") + except (httpx.HTTPError, UpstreamResponseError) as error: + if not sent: + raise # Preserve the HTTP error while no response bytes have been sent. + status, raw = _upstream_failure(error, model_name, t0, rid) + yield _err_event(raw, status) + finally: + # The inner generator may be suspended at a yielded line when its + # consumer reports an adapter error or a downstream disconnect. + await _close_stream(upstream) + + + + +def _err_event(msg: bytes, status: int) -> bytes: + chunk = {"error": {"message": sanitize_log_text(msg.decode("utf-8", "replace"), 512), + "type": "upstream_error", "code": status}} + return f"data: {json.dumps(chunk, ensure_ascii=False)}\n\n".encode("utf-8") + + +def _cred_manager(cred): + """Unwrap a credential lease, retaining support for a standalone manager.""" + return cred[0] if isinstance(cred, tuple) else cred + + +# Retryable upstream auth, quota and gateway responses; deterministic request errors are excluded. +FAILOVER_CODES = frozenset({401, 403, 429, 502, 503, 504}) +# Connection failures occur before any request body is sent. +REPLAYABLE_TRANSPORT = (httpx.ConnectError, httpx.ConnectTimeout) +# Write-timeout replay requires explicit opt-in because partial requests may already be billed. +WRITE_TIMEOUT_TRANSPORT = (httpx.WriteTimeout,) +# Gateway timeouts may follow billable upstream work and need an explicit cost warning. +POSSIBLY_CHARGED_CODES = frozenset({502, 504}) + + +def _replay_cost_note(error) -> str: + """Label replays that may duplicate already billed work.""" + if isinstance(error, UpstreamHTTPError) and error.status in POSSIBLY_CHARGED_CODES: + return " | 上游可能已处理该请求" + if isinstance(error, WRITE_TIMEOUT_TRANSPORT): + return " | 上游可能已处理该请求(正文未写完)" + return "" + + +def _failover_safe(error, raw=b"") -> bool: + """Allow configured pre-response transport or HTTP failover, never filter or incomplete-stream replay.""" + if is_filter_error(raw): + return False + if isinstance(error, UpstreamHTTPError): + return error.status in FAILOVER_CODES + if isinstance(error, UpstreamResponseError): + return False + if isinstance(error, WRITE_TIMEOUT_TRANSPORT): + return bool(CONFIG.get("retry_write_timeout")) + return isinstance(error, REPLAYABLE_TRANSPORT) + + +class _StreamFailure(Exception): + """Carry the HTTP error and original exception from stream preflight.""" + + def __init__(self, status, raw, error=None): + self.status = status + self.raw = raw + self.error = error + self.headers = error.headers if isinstance(error, UpstreamHTTPError) else None + super().__init__(f"stream failed before first byte (HTTP {status})") + + +# Bound teardown waits by cycles so repeated cancellation cannot cause a busy loop. +TEARDOWN_GRACE_CYCLES = 100 +TEARDOWN_POLL_SECONDS = 0.01 + + +def _drain_teardown(future) -> None: + """Retrieve teardown exceptions without rethrowing them.""" + if not future.cancelled(): + future.exception() + + +async def _teardown_finished(task) -> None: + """Wait briefly for isolated cleanup, then drain it in the background without delaying cancellation.""" + for _ in range(TEARDOWN_GRACE_CYCLES): + if task.done(): + _drain_teardown(task) + return + try: + await asyncio.wait([task], timeout=TEARDOWN_POLL_SECONDS) + except asyncio.CancelledError: + pass + if not task.done(): + task.add_done_callback(_drain_teardown) + + +async def _first_segment(agen): + """Read one stream segment in a shielded task so cancellation cannot interrupt its cleanup.""" + task = asyncio.ensure_future(agen.__anext__()) + try: + return await asyncio.shield(task) + except BaseException: + task.cancel() + await _teardown_finished(task) + raise + + +async def _stream_segments(agen): + """Read cancellable segments while allowing the generator's cleanup to finish.""" + while True: + try: + yield await _first_segment(agen) + except StopAsyncIteration: + return + + +async def _preflight_stream(agen, model_name, t0, rid): + """Read the first segment before committing HTTP 200, preserving pre-response error status.""" + try: + return await _first_segment(agen) + except StopAsyncIteration: + empty = UpstreamResponseError(502, b'{"error":{"message":"upstream returned an empty stream",' + b'"type":"upstream_error","code":"empty_response"}}') + status, raw = _upstream_failure(empty, model_name, t0, rid) + raise _StreamFailure(status, raw, empty) from None + except (httpx.HTTPError, UpstreamResponseError) as error: + status, raw = _upstream_failure(error, model_name, t0, rid) + raise _StreamFailure(status, raw, error) from None + + +async def _close_stream(agen) -> None: + """Close the upstream generator in an isolated task that survives repeated cancellation.""" + if agen is None: + return + + async def close() -> None: + try: + await agen.aclose() + except Exception: + pass + + await _teardown_finished(asyncio.ensure_future(close())) + + +def _chunk_bytes(chunk, charset: str = "utf-8"): + return chunk if isinstance(chunk, (bytes, memoryview)) else chunk.encode(charset) + + +class _DeferredStreamResponse(StreamingResponse): + """Run stream preflight and failover inside ASGI disconnect monitoring before sending headers.""" + + def __init__(self, plan): + self._plan = plan # Async callable returning the upstream iterator and first segment. + super().__init__(content=(), media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}) + + async def stream_response(self, send) -> None: + agen, first = await self._plan() + try: + await send({"type": "http.response.start", "status": self.status_code, + "headers": self.raw_headers}) + await send({"type": "http.response.body", "body": _chunk_bytes(first, self.charset), + "more_body": True}) + async for chunk in _stream_segments(agen): + await send({"type": "http.response.body", "body": _chunk_bytes(chunk, self.charset), + "more_body": True}) + await send({"type": "http.response.body", "body": b"", "more_body": False}) + finally: + await _close_stream(agen) + + +def _failover_limit() -> int: + return int(CONFIG.get("failover_max") or 0) + + +async def _stream_plan(payload, canonical, model_name, rid, t0, make, routed, cred, headers, url): + """Prefetch with bounded credential failover, using canonical input to preserve routing restrictions.""" + tried = [] + recovered = None + while True: + stream = make(routed, cred, headers, url) + try: + first = await _preflight_stream(stream, model_name, t0, rid) + except _StreamFailure as failure: + await _close_stream(stream) # Release the failed upstream connection. + release_credential(cred) + recovered = observe_failure_seq() # Recover only this failure sequence. + tried.append(cred) + limit = _failover_limit() + surface = HTTPException(status_code=failure.status, headers=failure.headers, + detail=_safe_err_raw(failure.raw, failure.status)) + if limit <= 0 or len(tried) > limit or not _failover_safe(failure.error, failure.raw): + raise surface from None + try: + attempt = await run_in_threadpool(_route_chat, payload, canonical, rid, + tried={_cred_manager(item) for item in tried}) + except HTTPException: + raise surface from None # Preserve the failure when no alternative account exists. + if _cred_manager(attempt[1]) in {_cred_manager(item) for item in tried}: + raise surface from None + routed, cred, headers, url = attempt + _log(f"[{rid}] ↻ 换凭证重放 {len(tried)}/{limit} | {model_name} | 上游 HTTP " + f"{failure.status} → {profile_for_headers(headers)}" + f"{_replay_cost_note(failure.error)}") + continue # Prefetch from the replacement credential. + except BaseException: + # Release the current upstream before propagating cancellation or unexpected errors. + await _close_stream(stream) + raise + if tried: + # Preserve newer failures such as content filtering on the replacement account. + observe_recovery(recovered) + return stream, first + + +def _routed_stream(payload, canonical, model_name, rid, t0, make, routed, cred, headers, url): + """Defer stream preflight and failover until ASGI disconnect monitoring is active.""" + return _DeferredStreamResponse( + lambda: _stream_plan(payload, canonical, model_name, rid, t0, make, + routed, cred, headers, url)) + + +async def _routed_fetch(payload, canonical, model_name, rid, t0, fetch, routed, cred, headers, url): + """Apply bounded non-streaming failover while preserving canonical routing restrictions.""" + tried = [] + recovered = None + while True: + try: + collected = await fetch(routed, cred, headers, url) + if tried: + # Recover the earlier attempt without erasing a newer failure. + observe_recovery(recovered) + return collected + except (httpx.HTTPError, UpstreamResponseError) as error: + release_credential(cred) + status, raw = _upstream_failure(error, model_name, t0, rid) + recovered = observe_failure_seq() + tried.append(cred) + limit = _failover_limit() + surface = HTTPException(status_code=status, detail=_safe_err_raw(raw, status), + headers=error.headers if isinstance(error, UpstreamHTTPError) else None) + if limit <= 0 or len(tried) > limit or not _failover_safe(error, raw): + raise surface from None + try: + attempt = await run_in_threadpool(_route_chat, payload, canonical, rid, + tried={_cred_manager(item) for item in tried}) + except HTTPException: + raise surface from None + if _cred_manager(attempt[1]) in {_cred_manager(item) for item in tried}: + raise surface from None + routed, cred, headers, url = attempt + _log(f"[{rid}] ↻ 换凭证重放 {len(tried)}/{limit} | {model_name} | 上游 HTTP " + f"{status} → {profile_for_headers(headers)}" + f"{_replay_cost_note(error)}") + + +def _note_content_filter(rid, model_name, *, final): + stage = "content_filter" if final else "content_filter_retry" + observe_attempt(stage, error_code="content_filter") + if final: + observe_failure("content_filter") + action = "保留上游拒绝,不切换账号" if final else "纯审核拒绝,压缩模板重试 1/1" + _log(f"[{rid}] 内容审核拦截 | {model_name} | {action}") + + +def _chat_body_desensitize(body: dict, *, force_compact: bool = False) -> dict: + if not CONFIG.get("desensitize"): + return body + return desensitize_body( + body, + roles=("system", "developer"), + desensitize_harness_user=True, + desensitize_tools=True, + compact_harness=(force_compact or not CONFIG.get("no_compact")), + strip_tool_metadata=not CONFIG.get("keep_tool_metadata", False), + ) + + +# --------------------------------------------------------------------------- +# OpenAI Responses endpoint +# --------------------------------------------------------------------------- + +@app.post("/v1/responses") +async def create_response(request: Request, + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Serve Responses requests through the shared Chat upstream and event adapter.""" + _check_auth(authorization, x_api_key) + + try: + payload = await request.json() + except Exception as e: + raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}}) + + payload = _prepare_payload(payload, field="input") + # Reject server-side conversation references because this gateway is stateless. + for stateful in ("previous_response_id", "conversation"): + if payload.get(stateful): + raise HTTPException(status_code=400, detail={"error": { + "message": f"{stateful} is not supported: this gateway keeps no server-side response state; resubmit the full input instead", + "type": "invalid_request_error", "param": stateful}}) + # Convert Responses input to Chat format. + try: + chat_body = responses_request_to_chat(payload) + except Exception as e: + raise HTTPException(status_code=400, detail={"error": {"message": f"request conversion error: {e}", "type": "invalid_request_error"}}) + + await run_in_threadpool(_bind_request_session, payload, chat_body) + projection_mode = CONFIG.get("responses_projection_mode", "balanced") + projection_max_bytes = int(CONFIG.get("responses_projection_max_bytes", 40000)) + try: + chat_body, projection_stats = await run_in_threadpool( + project_responses_chat_body, chat_body, mode=projection_mode, + max_item_bytes=projection_max_bytes) + except UnicodeError: + raise HTTPException(status_code=400, detail={"error": { + "message": "request contains text that cannot be encoded as UTF-8", + "type": "invalid_request_error", "param": "input", + "code": "invalid_unicode"}}) from None + observe_responses_projection(projection_stats) + chat_body = await run_in_threadpool(_prepare_chat_body, chat_body) + + client_wants_stream = _client_wants_stream(payload) + stream_policy = _snapshot_stream_policy("responses", chat_body) + observe_stream_mode(stream_policy.mode) + model_name = payload.get("model", "auto") + rid = _request_id() + _log(f"[{rid}] ▶ RESPONSES {model_name} | stream={client_wants_stream} | input_items={len(payload.get('input', []))}") + _log( + f"[{rid}] ── RESPONSES PROJECTION ── " + f"mode={projection_stats.get('mode')} " + f"| msgs {projection_stats.get('original_messages')}→{projection_stats.get('projected_messages')} " + f"| chars {projection_stats.get('original_message_chars')}→{projection_stats.get('projected_message_chars')} " + f"| tools {projection_stats.get('original_tools')}→{projection_stats.get('projected_tools')} " + f"| tool_chars {projection_stats.get('original_tool_chars')}→{projection_stats.get('projected_tool_chars')} " + f"| harness_messages={projection_stats.get('harness_messages_projected', 0)} " + f"| truncated_items={projection_stats.get('truncated_items', 0)} " + f"| truncated_bytes={projection_stats.get('truncated_original_bytes', 0)}→" + f"{projection_stats.get('truncated_projected_bytes', 0)}" + ) + # Keep blocking credential selection and refresh off the event loop. + prepared = chat_body # Preserve canonical input for routing policy checks. + chat_body, cred, headers, url = await run_in_threadpool(_route_chat, payload, chat_body, rid) + _log_json(f"[{rid}] RESPONSES → CHAT BODY (预览)", chat_body) + t0 = time.time() + + if client_wants_stream: + def attempt(routed, cred, headers, url): + return _stream_responses(url, headers, _body_with_stream_policy(routed, stream_policy), + model_name, t0, rid, cred=cred) + response = _routed_stream(payload, prepared, model_name, rid, t0, attempt, + chat_body, cred, headers, url) + response.headers["X-CodeBuddy-Responses-Projection"] = projection_mode + return response + + response = await _nonstream_adapted( + url, headers, chat_body, model_name, t0, rid, cred, payload=payload, canonical=prepared, + request=request, policy=stream_policy) + response.headers["X-CodeBuddy-Responses-Projection"] = projection_mode + return response + + async def _nonstream_adapted(url, headers, body, model_name, t0, rid, cred, *, anthropic=False, payload=None, canonical=None, request=None, policy=None): policy = policy or _snapshot_stream_policy("messages" if anthropic else "responses", body) - converter = (AnthropicStreamConverter(model=model_name) if anthropic else ResponsesStreamConverter(model=model_name, parallel_tool_calls=body.get("parallel_tool_calls", True))) + tool_namespaces = body.get("_tool_namespaces") + converter = (AnthropicStreamConverter(model=model_name) if anthropic else + ResponsesStreamConverter(model=model_name, + parallel_tool_calls=body.get("parallel_tool_calls", True), + tool_namespaces=tool_namespaces)) async def fetch(routed, cred, headers, url): - return await _fetch_checked_chat(url, headers, routed, model_name, rid, cred, + routed_body = {k: v for k, v in routed.items() if k != "_tool_namespaces"} if "_tool_namespaces" in routed else routed + return await _fetch_checked_chat(url, headers, routed_body, model_name, rid, cred, filter_retry=True, max_collect_bytes=policy.max_collect_bytes) - try: - collected = await await_or_hangup( - _routed_fetch(payload, body if canonical is None else canonical, - model_name, rid, t0, fetch, body, cred, headers, url), request) - for line in _chat_result_to_sse_lines(_completion_to_merged(collected)): - converter.feed_line(_public_sse_line(line, model_name)) - converter.finish() - except (httpx.HTTPError, UpstreamResponseError) as error: - status, raw = _upstream_failure(error, model_name, t0, rid) - raise HTTPException(status_code=status, detail=_safe_err_raw(raw, status), - headers=error.headers if isinstance(error, UpstreamHTTPError) else None) from None - except ClientHungUp: - return _hungup_response(rid, model_name, t0) - result = converter.get_nonstream_response() - _log_finish(model_name, t0, collected, rid) - return JSONResponse(content=result) - - + try: + collected = await await_or_hangup( + _routed_fetch(payload, body if canonical is None else canonical, + model_name, rid, t0, fetch, body, cred, headers, url), request) + for line in _chat_result_to_sse_lines(_completion_to_merged(collected)): + converter.feed_line(_public_sse_line(line, model_name)) + converter.finish() + except (httpx.HTTPError, UpstreamResponseError) as error: + status, raw = _upstream_failure(error, model_name, t0, rid) + raise HTTPException(status_code=status, detail=_safe_err_raw(raw, status), + headers=error.headers if isinstance(error, UpstreamHTTPError) else None) from None + except ClientHungUp: + return _hungup_response(rid, model_name, t0) + result = converter.get_nonstream_response() + _log_finish(model_name, t0, collected, rid) + return JSONResponse(content=result) + + async def _stream_adapted(url, headers, body, model_name, t0, rid, cred=None, *, anthropic=False): """Map protocol events while sharing connection, aggregation and failure handling.""" protocol = "messages" if anthropic else "responses" policy = body.pop(_REQUEST_POLICY_KEY, None) or _snapshot_stream_policy(protocol, body) + tool_namespaces = body.get("_tool_namespaces") state = {} tracker = None declared_names = _declared_tool_names(body) @@ -3684,488 +3690,491 @@ async def _stream_adapted(url, headers, body, model_name, t0, rid, cred=None, *, ResponsesStreamConverter(model=model_name, parallel_tool_calls=body.get("parallel_tool_calls", True), realtime=True, budget=budget, tool_states=tracker.tools, - declared_names=declared_names)) + declared_names=declared_names, + tool_namespaces=tool_namespaces)) else: converter = (AnthropicStreamConverter(model=model_name) if anthropic else ResponsesStreamConverter(model=model_name, - parallel_tool_calls=body.get("parallel_tool_calls", True))) + parallel_tool_calls=body.get("parallel_tool_calls", True), + tool_namespaces=tool_namespaces)) sent = False + upstream_body = {k: v for k, v in body.items() if k != "_tool_namespaces"} if "_tool_namespaces" in body else body upstream = _chat_sse_lines( - url, headers, body, model_name, t0, rid, cred, + url, headers, upstream_body, model_name, t0, rid, cred, policy=policy, tracker=tracker, state=state) - try: - try: - async for line in upstream: - events = converter.feed_line(_public_sse_line(line, model_name)) - if events: - sent = True - yield events.encode("utf-8") - if policy.realtime: - converter.set_validated_tools((state.get("merged") or {}).get("tool_calls")) - if tracker.filter_detector.detected: - converter.mark_content_filter() - events = converter.finish() - if events: - sent = True - yield events.encode("utf-8") - except (httpx.HTTPError, UpstreamResponseError) as error: - if not sent: - raise # Preserve the HTTP error before any response bytes have been sent. - status, raw = _upstream_failure(error, model_name, t0, rid) - event = {"type": "error", "error": { - "message": sanitize_log_text(raw.decode("utf-8", "replace"), 512), - "type": "api_error" if anthropic else "upstream_error", "code": status}} - prefix = "event: error\n" if anthropic else "" - yield (prefix + f"data: {json.dumps(event, ensure_ascii=False)}\n\n").encode("utf-8") - except ValueError: - # Adapter metadata/append-only checks are protocol failures, not - # uncaught application errors after a response has opened. - failure = UpstreamResponseError(502, b'{"error":{"message":"Invalid upstream tool_calls",' - b'"type":"upstream_error","code":"invalid_tool_calls"}}') - if not sent: - raise failure from None - status, raw = _upstream_failure(failure, model_name, t0, rid) - event = {"type": "error", "error": { - "message": sanitize_log_text(raw.decode("utf-8", "replace"), 512), - "type": "api_error" if anthropic else "upstream_error", "code": status}} - prefix = "event: error\n" if anthropic else "" - yield (prefix + f"data: {json.dumps(event, ensure_ascii=False)}\n\n").encode("utf-8") - finally: - # Close the nested Chat generator explicitly; do not wait for GC after - # an adapter error, disconnect or cancellation. - await _close_stream(upstream) - - -async def _stream_responses(url: str, headers: dict, body: dict, - model_name: str = "?", t0: float = 0.0, rid: str = "", cred=None): - stream = _stream_adapted(url, headers, body, model_name, t0, rid, cred) - try: - async for chunk in stream: - yield chunk - finally: - await _close_stream(stream) - - -# --------------------------------------------------------------------------- -# Anthropic Messages endpoint -# --------------------------------------------------------------------------- - -@app.post("/v1/messages") -async def create_message(request: Request, - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Serve Anthropic Messages through the shared Chat upstream and event adapter.""" - _check_auth(authorization, x_api_key) - - try: - payload = await request.json() - except Exception as e: - raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}}) - - payload = _prepare_payload(payload) - # Convert Anthropic messages and tools to Chat format. - messages = payload.get("messages") or [] - if not messages: - raise HTTPException(status_code=400, detail={"error": {"message": "messages is required", "type": "invalid_request_error"}}) - - try: - chat_body = anthropic_request_to_chat(payload) - except Exception as e: - raise HTTPException(status_code=400, detail={"error": {"message": f"request conversion error: {e}", "type": "invalid_request_error"}}) - - chat_body = await run_in_threadpool(_prepare_chat_body, chat_body, session_payload=payload) - client_wants_stream = _client_wants_stream(payload) - stream_policy = _snapshot_stream_policy("messages", chat_body) - observe_stream_mode(stream_policy.mode) - model_name = payload.get("model", "auto") - chat_messages = chat_body.get("messages", []) - rid = _request_id() - _log(f"[{rid}] ▶ ANTHROPIC {model_name} | msgs={len(chat_messages)} | anthropic_msgs={len(messages)}") - # Keep blocking credential selection and refresh off the event loop. - prepared = chat_body # Preserve canonical input for routing policy checks. - chat_body, cred, headers, url = await run_in_threadpool(_route_chat, payload, chat_body, rid) - _log_json(f"[{rid}] ANTHROPIC → CHAT BODY (预览)", chat_body) - t0 = time.time() - - if not client_wants_stream: - return await _nonstream_adapted(url, headers, chat_body, model_name, t0, rid, cred, - anthropic=True, payload=payload, canonical=prepared, - request=request, policy=stream_policy) - - def attempt(routed, cred, headers, url): - return _stream_anthropic(url, headers, _body_with_stream_policy(routed, stream_policy), - model_name, t0, rid, cred=cred) - return _routed_stream(payload, prepared, model_name, rid, t0, attempt, - chat_body, cred, headers, url) - - -async def _stream_anthropic(url: str, headers: dict, body: dict, - model_name: str = "?", t0: float = 0.0, rid: str = "", cred=None): - stream = _stream_adapted(url, headers, body, model_name, t0, rid, cred, anthropic=True) - try: - async for chunk in stream: - yield chunk - finally: - await _close_stream(stream) - - -@app.post("/v1/messages/count_tokens") -async def count_tokens(request: Request, - authorization: Optional[str] = Header(default=None), - x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): - """Return heuristic token estimates for Anthropic request budgeting.""" - _check_auth(authorization, x_api_key) - try: - payload = await request.json() - except Exception as e: - raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}}) - if not isinstance(payload, dict): - raise HTTPException(status_code=400, detail={"error": {"message": "请求体必须是 JSON 对象", "type": "invalid_request_error"}}) - return {"input_tokens": _estimate_input_tokens(payload)} - - -def _estimate_input_tokens(payload: dict) -> int: - """Estimate tokens from character counts and message overhead, not upstream billing.""" - - def measure(value) -> int: - if isinstance(value, str): - ascii_chars = sum(1 for ch in value if ord(ch) < 128) - return (ascii_chars + 3) // 4 + (len(value) - ascii_chars) - if isinstance(value, list): - return sum(measure(item) for item in value) - if isinstance(value, dict): - return sum(measure(item) for item in value.values()) - return 0 - - total = measure(payload.get("system")) + measure(payload.get("tools")) - messages = payload.get("messages") - if isinstance(messages, list): - for message in messages: - if isinstance(message, dict): - total += measure(message.get("content")) + 4 # Message structure overhead. - return total - - -# --------------------------------------------------------------------------- -# Startup -# --------------------------------------------------------------------------- - -def preflight() -> bool: - files = find_auth_files() - sys.stderr.write("==== 预检 ====\n") - sys.stderr.write(f"平台 : {sys.platform}\n") - sys.stderr.write(f"Python : {sys.version.split()[0]}\n") - sys.stderr.write(f"后端 : {BACKEND} (直连,原生 function calling)\n") - sys.stderr.write(f"自管目录 : {managed_auth_dir()}\n") - sys.stderr.write(f"登录文件 : {len(files)} 个\n") - if not os.environ.get("CODEBUDDY_AUTH_DIR"): - sys.stderr.write(f"种子来源 : {', '.join(str(d) for d in auth_dirs())}\n") - ok = True - if not files: - sys.stderr.write("\n[警告] 未找到登录文件。请运行 python3 converter.py login 扫码添加账号,或用 --auth-file 指定。\n") - ok = False - for af in files: - try: - cm = CredentialManager(af) - info = cm.summary() - sys.stderr.write(f"账号 : {info.get('nickname')} / {info.get('enterpriseName')} ({af.name})\n") - sys.stderr.write(f"token过期 : {'是(将自动刷新)' if info['token_expired'] else '否'}\n") - except Exception as e: - sys.stderr.write(f"[警告] 读取凭据失败 {af.name}:{e}\n") - ok = False - sys.stderr.write("================\n") - return ok - - -def login(site: str = "cn", open_browser: bool = True) -> int: - """Complete browser login and persist managed credentials without local HTTP calls.""" - import webbrowser - - try: - started = _OAUTH.start(site=site) - uri = started["verification_uri"] - print(f"请打开以下链接扫码登录:\n{uri}", flush=True) - if open_browser: - try: - opened = webbrowser.open(uri) - except webbrowser.Error: - opened = False - if not opened: - print("无法自动打开浏览器,请手动打开上面的链接。", flush=True) - print("正在等待扫码授权;网页显示登录成功后,请继续等待终端确认入库。\n" - "按 Ctrl+C 取消。", flush=True) - while True: - result = _OAUTH.poll(started["login_id"]) - if result.get("done"): - if result.get("error") or not result.get("cred"): - print(f"登录失败:{result.get('error') or '未获取到凭据'}", file=sys.stderr) - return 1 - target = _save_oauth_credential(result["cred"]) - print(f"登录成功,账号已保存至:{target}\n" - "使用同一凭据目录的服务会在下次请求时自动加载(默认目录扫描模式)。", - flush=True) - return 0 - time.sleep(1.5) - except KeyboardInterrupt: - print("\n已取消登录。", file=sys.stderr) - return 130 - except CredentialFileError as e: - print(f"登录失败:{e}", file=sys.stderr) - return 1 - except OSError: - print("登录失败:无法保存凭据,请检查凭据目录的写入权限。", file=sys.stderr) - return 1 - except (httpx.HTTPError, ValueError, RuntimeError): - # Upstream errors may contain authorization URLs or sensitive response data. - print("登录失败:登录接口请求失败或响应无效,请检查网络后重试。", file=sys.stderr) - return 1 - - -def _nonnegative_int(value): - number = int(value) - if number < 0: - raise argparse.ArgumentTypeError("必须为非负整数") - return number - - -def _positive_int(value): - number = _nonnegative_int(value) - if number == 0: - raise argparse.ArgumentTypeError("必须为正整数") - return number - -def _projection_bytes_arg(value): - number = _nonnegative_int(value) - if number != 0 and number < 256: - raise argparse.ArgumentTypeError("必须为 0 或至少 256") - if number > 33554432: - raise argparse.ArgumentTypeError("不能超过 33554432") - return number - - -def _origins_arg(value): - from app.settings import normalize_allowed_origins - try: - return normalize_allowed_origins(value) - except ValueError: - raise argparse.ArgumentTypeError("必须为逗号分隔的 http/https 来源或域名") from None - - -def _boolean_arg(value): - if isinstance(value, bool): - return value - normalized = str(value).strip().lower() - if normalized in ("true", "1", "yes", "on"): - return True - if normalized in ("false", "0", "no", "off"): - return False - raise argparse.ArgumentTypeError("必须为 true 或 false") - - -def main(): - dotenv_keys = set() if any(arg in ("-h", "--help") for arg in sys.argv[1:]) else load_startup_env() - ap = argparse.ArgumentParser(description="CodeBuddy -> OpenAI 兼容转换器(直连后端)") - ap.add_argument("command", nargs="?", choices=("serve", "login"), default="serve", - help="serve 启动服务(默认);login 扫码登录、自动轮询并保存账号") - ap.add_argument("--site", choices=tuple(auth_oauth.SITE_HOSTS), default="cn", - help="login 站点:cn 国内站(默认);intl 国际 WorkBuddy;intl-codebuddy 国际 CodeBuddy") - ap.add_argument("--no-browser", action="store_true", - help="login 仅显示授权链接,不自动打开浏览器(服务器/容器环境)") - ap.add_argument("--host", default="127.0.0.1", help="监听地址;覆盖 CODEBUDDY2API_BIND") - ap.add_argument("--port", type=int, default=8787, help="监听端口;覆盖 CODEBUDDY2API_PORT") - ap.add_argument("--api-key", default=os.environ.get("CODEBUDDY2API_KEY", ""), - help="管理与推理密钥;未配置时首次本地交互启动生成并保存") - ap.add_argument("--admin-csrf", type=_boolean_arg, nargs="?", const=True, - default=os.environ.get("CODEBUDDY2API_ADMIN_CSRF", "true"), - help="管理 Origin/CSRF 校验,默认 true;仅在受信任本地环境设为 false,鉴权仍启用") - ap.add_argument("--admin-allowed-origins", type=_origins_arg, metavar="ORIGINS", - default=os.environ.get("CODEBUDDY2API_ADMIN_ORIGINS"), - help="额外信任的管理页来源(逗号分隔,支持域名或完整来源,裸域名按 https);" - "反代 HTTPS 域名登录报 Origin 校验失败时设置,也可在 WebUI 配置") - ap.add_argument("--log", default=None, metavar="PATH", - help="已停用:日志统一保存到数据目录中的 logs.sqlite3") - ap.add_argument("--desensitize", action="store_true", - help="适配固定 CLI 模板、压缩运行时提示并零宽脱敏关键词。默认关闭。") - ap.add_argument("--no-compact", action="store_true", - help="配合 --desensitize 保留主要行为指令,仍适配固定模板并裁剪运行时元数据;" - "非流式纯审核拒绝最多压缩兜底一次。") - ap.add_argument("--keep-tool-metadata", type=_boolean_arg, nargs="?", const=True, - default=os.environ.get("CODEBUDDY2API_KEEP_TOOL_METADATA", "false"), - help="保留工具描述及参数 description/title;启用脱敏时仍处理描述文本,默认 false") - ap.add_argument("--responses-projection-mode", choices=PROJECTION_MODES, - default=os.environ.get("CODEBUDDY2API_RESPONSES_PROJECTION_MODE", "balanced"), - help="Responses 上下文:balanced 仅改写有固定摘要的已识别 harness,其余文本原样保留;passthrough 完全关闭投影;默认 balanced") - ap.add_argument("--responses-projection-max-bytes", type=_projection_bytes_arg, metavar="BYTES", - default=os.environ.get("CODEBUDDY2API_RESPONSES_PROJECTION_MAX_BYTES", "40000"), - help="Responses 单项 UTF-8 字节上限;0 禁用,非零范围 256..33554432,默认 40000") - ap.add_argument("--skip-check", action="store_true", help="跳过启动预检") - ap.add_argument("--auth-file", action="append", default=[], metavar="PATH", - help="凭据文件(可重复传入组成凭证池;默认自动扫描 auth 目录全部 *.info)") - ap.add_argument("--credit-price-cny", type=float, default=None, metavar="PRICE", - help="积分折算单价(元/Credit),默认 0.014(旗舰版连续包月 700元/5万积分摊算)") - ap.add_argument("--usd-rate", type=float, default=None, metavar="RATE", - help="人民币→美元汇率,影响 /v1/dashboard/billing 端点金额") - ap.add_argument("--credit-price-usd", type=float, default=None, metavar="PRICE", - help="国际站积分折算单价(美元/Credit),默认 0.03(Pro 加量包 $15/500 积分)") - ap.add_argument("--model-catalog-ttl", type=int, default=6 * 3600, metavar="SECONDS", - help="云端模型表缓存有效期,默认 21600 秒(6 小时);TTL 内不再打 /v3/config") - ap.add_argument("--no-model-guard", action="store_true", - help="关闭表外模型本地拦截;默认拦截,避免无效请求打到上游并触发扣费") - ap.add_argument("--max-images", type=_nonnegative_int, metavar="N", - default=os.environ.get("CODEBUDDY2API_MAX_IMAGES", "16"), - help="单请求图片上限,默认 16;0 表示不允许图片") - ap.add_argument("--image-policy", choices=("truncate", "error"), - default=os.environ.get("CODEBUDDY2API_IMAGE_POLICY", "truncate"), - help="超额图片策略:truncate 保留最新图片(默认),error 返回 413") - ap.add_argument("--max-request-bytes", type=_positive_int, metavar="BYTES", - default=os.environ.get("CODEBUDDY2API_MAX_REQUEST_BYTES", str(32 * 1024 * 1024)), - help="图片处理与适配后请求体的字节上限,默认 32 MiB") - ap.add_argument("--max-inbound-bytes", type=_positive_int, metavar="BYTES", - default=os.environ.get("CODEBUDDY2API_MAX_INBOUND_BYTES", str(64 * 1024 * 1024)), - help="入站原始请求体字节上限(解析前生效,含 chunked),默认 64 MiB") - ap.add_argument("--max-collect-bytes", type=_nonnegative_int, metavar="BYTES", - default=os.environ.get("CODEBUDDY2API_MAX_COLLECT_BYTES", str(8 * 1024 * 1024)), - help="聚合路径输出收集总字节上限(正文+思考+工具参数),默认 8 MiB;0 不限制") - ap.add_argument("--max-concurrent", type=_nonnegative_int, metavar="N", - default=os.environ.get("CODEBUDDY2API_MAX_CONCURRENT", "64"), - help="推理端点并发上限(超出立即 503),默认 64;0 不限制") - ap.add_argument("--max-inflight-per-account", type=_nonnegative_int, metavar="N", - default=os.environ.get("CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT", "0"), - help="单账号在途上限,默认 0(不限制);满载返回 503,不借容量切换到收费账号") - ap.add_argument("--upstream-keepalive", type=_boolean_arg, nargs="?", const=True, - default=os.environ.get("CODEBUDDY2API_UPSTREAM_KEEPALIVE", "false"), - help="按上游入口复用有界连接池,默认 false;重启生效,不改变超时或重放规则") - ap.add_argument("--request-context-mode", choices=("legacy", "scoped"), - default=os.environ.get("CODEBUDDY2API_REQUEST_CONTEXT_MODE", "legacy"), - help="请求上下文:legacy 保持旧会话头,scoped 启用显式会话与逐尝试追踪;默认 legacy") - ap.add_argument("--stream-mode", choices=_STREAM_MODES, - default=os.environ.get("CODEBUDDY2API_STREAM_MODE", "compatible"), - help="流式传输:compatible 保持兼容聚合策略,realtime 增量发送且不重生成工具参数;默认 compatible") - ap.add_argument("--model-capability-guard", type=_boolean_arg, nargs="?", const=True, - default=os.environ.get("CODEBUDDY2API_MODEL_CAPABILITY_GUARD", "true"), - help="按账号模型声明预检图片、工具、思考和输出上限;false 仅关闭新增能力预检") - ap.add_argument("--log-body-limit", type=_nonnegative_int, metavar="BYTES", - default=os.environ.get("CODEBUDDY2API_LOG_BODY_LIMIT", "65536"), - help="旧文本预览兼容项;文本输出已停用,SQLite 诊断使用独立预算") - ap.add_argument("--tool-call-max-retry", type=_nonnegative_int, metavar="N", - default=os.environ.get("CODEBUDDY2API_TOOL_CALL_MAX_RETRY", "3"), - help="工具参数损坏时的额外生成上限,默认 3;0 表示不重试(每次额外生成都消耗额度)") - ap.add_argument("--failover-max", type=_nonnegative_int, metavar="N", - default=os.environ.get("CODEBUDDY2API_FAILOVER_MAX", "0"), - help="失败发生在向下游落第一个字节之前时,最多换几个凭证就地重放,默认 0(关闭);" - "只重放上游没收下请求体或用 401/403/429/502/503/504 拒绝的失败;" - "写请求体超时需另开 --retry-write-timeout 才参与") - ap.add_argument("--retry-write-timeout", type=_boolean_arg, nargs="?", const=True, - default=os.environ.get("CODEBUDDY2API_RETRY_WRITE_TIMEOUT", "false"), - help="把「写请求体超时」也算作上游没收下请求体从而参与重放,默认 false。写超时只能" - "证明正文没写完,上游是否已按半截正文计费看不到,因此要显式开启(同时作用于连接" - "重试与 --failover-max 换凭证重放)") - ap.add_argument("--auto-trial", type=_boolean_arg, nargs="?", const=True, - default=None, help=argparse.SUPPRESS) - args = ap.parse_args() - if args.auto_trial is not None or "CODEBUDDY2API_AUTO_TRIAL" in os.environ: - sys.stderr.write("[trial] AUTO_TRIAL / --auto-trial 已停用;请在 WebUI 凭证页手动领取体验积分。\n") - del args.auto_trial - if args.image_policy not in ("truncate", "error"): - ap.error("CODEBUDDY2API_IMAGE_POLICY 必须为 truncate 或 error") - if args.command == "login": - return login(site=args.site, open_browser=not args.no_browser) - - for key in ("max_images", "image_policy", "max_request_bytes", "log_body_limit", - "tool_call_max_retry", "max_inbound_bytes", "max_collect_bytes", "max_concurrent", - "failover_max", "retry_write_timeout", "upstream_keepalive", "max_inflight_per_account", - "request_context_mode", "stream_mode", "model_capability_guard", "admin_allowed_origins", - "responses_projection_mode", "responses_projection_max_bytes"): - CONFIG[key] = getattr(args, key) - CONFIG["api_key"] = args.api_key - CONFIG["desensitize"] = args.desensitize - CONFIG["no_compact"] = args.no_compact - CONFIG["credit_price_cny"] = args.credit_price_cny or None - CONFIG["usd_rate"] = args.usd_rate or None - CONFIG["credit_price_usd"] = args.credit_price_usd or None - CONFIG["model_guard"] = not args.no_model_guard - if args.log is not None or os.environ.get("CODEBUDDY2API_LOG"): - sys.stderr.write("[log] --log / CODEBUDDY2API_LOG 已停用;请在 WebUI 查看 SQLite 日志。\n") - CONFIG["log_path"] = None - from app import runtime_management - runtime_management.initialize(sys.modules[__name__], args, parser=ap, dotenv_keys=dotenv_keys) - # Validate effective binding and authentication before credential scans or background work. - if (args.host not in ("127.0.0.1", "::1", "localhost") and not CONFIG.get("api_key") - and not allows_open_noauth()): - runtime_management.close(CONFIG) - ap.error("非回环绑定且未设置 API key 会匿名开放推理额度;" - "请设置 CODEBUDDY2API_KEY,或确知风险后以 CODEBUDDY2API_ALLOW_OPEN_NOAUTH=true 显式放行") - files = [Path(p) for p in args.auth_file] - if not files: - seed_credentials() # Seed missing desktop credentials into managed storage. - CONFIG["cred_pool"] = CredentialPool(files, scan=not files, state_store=CONFIG["state_store"]) - CONFIG["usage_snapshots"] = UsageSnapshots(store=CONFIG["state_store"]) - CONFIG["cred"] = CONFIG["cred_pool"].first() - CONFIG["account_catalogs"] = {} # Disable static fallback before maintenance starts. - if credits_mod is not None: - ledger = credits_mod.CreditLedger(store=CONFIG["state_store"]) - CONFIG["ledger"] = ledger - CONFIG["model_cache"] = credits_mod.ModelCatalogCache(ttl=args.model_catalog_ttl, store=CONFIG["state_store"]) - CONFIG["cred_pool"].set_ledger(ledger) # Verify balance ownership before publishing catalogs. - _publish_model_cache() - # Publish cached usage before maintenance threads start, so the dashboard is populated from - # the first request. It stays a no-op when nothing is cached, so an empty deployment and a - # mocked pool are both unaffected. - if CONFIG["usage_snapshots"].detail(): - _publish_usage_daily(CONFIG["cred_pool"]) - - try: - runtime_management.install(sys.modules[__name__]) - except SessionStoreError as error: - # An obsolete session snapshot survived, so the new key epoch must not activate: - # a later start under the superseded key could adopt it and revive admin cookies. - runtime_management.close(CONFIG) - ap.error(str(error)) - threading.Thread(target=_refresher_loop, args=(CONFIG["cred_pool"],), - daemon=True, name="cred-refresher").start() - if credits_mod is not None: - threading.Thread(target=_housekeeper_loop, args=(CONFIG["cred_pool"], ledger), - daemon=True, name="cred-housekeeper").start() - - if not args.skip_check: - preflight() - - sys.stderr.write(f"\n✅ 监听 http://{args.host}:{args.port}(直连后端,原生 function calling)\n") - sys.stderr.write(f" WebUI : http://{args.host}:{args.port}/dashboard\n") - sys.stderr.write(" SQLite 审计默认开启,凭证仍以 .info 文件保存\n") - sys.stderr.write(" GET /v1/models\n") - sys.stderr.write(" POST /v1/chat/completions (原生 tools/tool_calls,支持流式)\n") - sys.stderr.write(" POST /v1/responses (Responses API,Codex CLI 兼容)\n") - sys.stderr.write(" POST /v1/messages (Anthropic API,Claude Code / CC Switch 兼容)\n") - sys.stderr.write(" GET /health\n") - sys.stderr.write(" GET/POST/DELETE /admin/credentials (凭证池管理)\n") - sys.stderr.write(" 添加账号:python3 converter.py login(自动等待扫码并保存)\n") - if credits_mod is not None: - sys.stderr.write(" GET /admin/credits (积分/签到状态)\n") - sys.stderr.write(" POST /admin/checkin (仅签到,按日幂等)\n") - sys.stderr.write(" POST /admin/sync (同步余额、目录与用量,不签到)\n") - sys.stderr.write(" POST /admin/credentials/{id}/reset-cooldown (清除该账号冷却,仅本地)\n") - sys.stderr.write(" 每日签到 + 快过期积分优先调度已启用\n") - if args.api_key: - sys.stderr.write(" 鉴权已启用(API key 已设置)\n") - if not CONFIG["admin_csrf"]: - sys.stderr.write(" 警告:管理 Origin/CSRF 校验已关闭,仅限受信任本地环境;API key 与会话校验仍启用。\n") - sys.stderr.write(f" 图片限制 : {CONFIG['max_images']} 张/请求,策略 {CONFIG['image_policy']}\n") - if CONFIG["log_path"]: - sys.stderr.write(f" 日志 : {CONFIG['log_path']}\n") - if args.desensitize: - mode = "零宽脱敏 + 保留全文" if args.no_compact else "零宽脱敏 + 压缩摘要" - sys.stderr.write(f" 脱敏 : 已启用({mode})\n") - sys.stderr.write("按 Ctrl+C 退出。\n\n") - - # Record service startup. - _log(f"==== converter 启动 ====") - - try: - run_server(app, CONFIG, host=args.host, port=args.port) - finally: - runtime_management.close(CONFIG) - - -if __name__ == "__main__": - sys.exit(main()) + try: + try: + async for line in upstream: + events = converter.feed_line(_public_sse_line(line, model_name)) + if events: + sent = True + yield events.encode("utf-8") + if policy.realtime: + converter.set_validated_tools((state.get("merged") or {}).get("tool_calls")) + if tracker.filter_detector.detected: + converter.mark_content_filter() + events = converter.finish() + if events: + sent = True + yield events.encode("utf-8") + except (httpx.HTTPError, UpstreamResponseError) as error: + if not sent: + raise # Preserve the HTTP error before any response bytes have been sent. + status, raw = _upstream_failure(error, model_name, t0, rid) + event = {"type": "error", "error": { + "message": sanitize_log_text(raw.decode("utf-8", "replace"), 512), + "type": "api_error" if anthropic else "upstream_error", "code": status}} + prefix = "event: error\n" if anthropic else "" + yield (prefix + f"data: {json.dumps(event, ensure_ascii=False)}\n\n").encode("utf-8") + except ValueError: + # Adapter metadata/append-only checks are protocol failures, not + # uncaught application errors after a response has opened. + failure = UpstreamResponseError(502, b'{"error":{"message":"Invalid upstream tool_calls",' + b'"type":"upstream_error","code":"invalid_tool_calls"}}') + if not sent: + raise failure from None + status, raw = _upstream_failure(failure, model_name, t0, rid) + event = {"type": "error", "error": { + "message": sanitize_log_text(raw.decode("utf-8", "replace"), 512), + "type": "api_error" if anthropic else "upstream_error", "code": status}} + prefix = "event: error\n" if anthropic else "" + yield (prefix + f"data: {json.dumps(event, ensure_ascii=False)}\n\n").encode("utf-8") + finally: + # Close the nested Chat generator explicitly; do not wait for GC after + # an adapter error, disconnect or cancellation. + await _close_stream(upstream) + + +async def _stream_responses(url: str, headers: dict, body: dict, + model_name: str = "?", t0: float = 0.0, rid: str = "", cred=None): + stream = _stream_adapted(url, headers, body, model_name, t0, rid, cred) + try: + async for chunk in stream: + yield chunk + finally: + await _close_stream(stream) + + +# --------------------------------------------------------------------------- +# Anthropic Messages endpoint +# --------------------------------------------------------------------------- + +@app.post("/v1/messages") +async def create_message(request: Request, + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Serve Anthropic Messages through the shared Chat upstream and event adapter.""" + _check_auth(authorization, x_api_key) + + try: + payload = await request.json() + except Exception as e: + raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}}) + + payload = _prepare_payload(payload) + # Convert Anthropic messages and tools to Chat format. + messages = payload.get("messages") or [] + if not messages: + raise HTTPException(status_code=400, detail={"error": {"message": "messages is required", "type": "invalid_request_error"}}) + + try: + chat_body = anthropic_request_to_chat(payload) + except Exception as e: + raise HTTPException(status_code=400, detail={"error": {"message": f"request conversion error: {e}", "type": "invalid_request_error"}}) + + chat_body = await run_in_threadpool(_prepare_chat_body, chat_body, session_payload=payload) + client_wants_stream = _client_wants_stream(payload) + stream_policy = _snapshot_stream_policy("messages", chat_body) + observe_stream_mode(stream_policy.mode) + model_name = payload.get("model", "auto") + chat_messages = chat_body.get("messages", []) + rid = _request_id() + _log(f"[{rid}] ▶ ANTHROPIC {model_name} | msgs={len(chat_messages)} | anthropic_msgs={len(messages)}") + # Keep blocking credential selection and refresh off the event loop. + prepared = chat_body # Preserve canonical input for routing policy checks. + chat_body, cred, headers, url = await run_in_threadpool(_route_chat, payload, chat_body, rid) + _log_json(f"[{rid}] ANTHROPIC → CHAT BODY (预览)", chat_body) + t0 = time.time() + + if not client_wants_stream: + return await _nonstream_adapted(url, headers, chat_body, model_name, t0, rid, cred, + anthropic=True, payload=payload, canonical=prepared, + request=request, policy=stream_policy) + + def attempt(routed, cred, headers, url): + return _stream_anthropic(url, headers, _body_with_stream_policy(routed, stream_policy), + model_name, t0, rid, cred=cred) + return _routed_stream(payload, prepared, model_name, rid, t0, attempt, + chat_body, cred, headers, url) + + +async def _stream_anthropic(url: str, headers: dict, body: dict, + model_name: str = "?", t0: float = 0.0, rid: str = "", cred=None): + stream = _stream_adapted(url, headers, body, model_name, t0, rid, cred, anthropic=True) + try: + async for chunk in stream: + yield chunk + finally: + await _close_stream(stream) + + +@app.post("/v1/messages/count_tokens") +async def count_tokens(request: Request, + authorization: Optional[str] = Header(default=None), + x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")): + """Return heuristic token estimates for Anthropic request budgeting.""" + _check_auth(authorization, x_api_key) + try: + payload = await request.json() + except Exception as e: + raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}}) + if not isinstance(payload, dict): + raise HTTPException(status_code=400, detail={"error": {"message": "请求体必须是 JSON 对象", "type": "invalid_request_error"}}) + return {"input_tokens": _estimate_input_tokens(payload)} + + +def _estimate_input_tokens(payload: dict) -> int: + """Estimate tokens from character counts and message overhead, not upstream billing.""" + + def measure(value) -> int: + if isinstance(value, str): + ascii_chars = sum(1 for ch in value if ord(ch) < 128) + return (ascii_chars + 3) // 4 + (len(value) - ascii_chars) + if isinstance(value, list): + return sum(measure(item) for item in value) + if isinstance(value, dict): + return sum(measure(item) for item in value.values()) + return 0 + + total = measure(payload.get("system")) + measure(payload.get("tools")) + messages = payload.get("messages") + if isinstance(messages, list): + for message in messages: + if isinstance(message, dict): + total += measure(message.get("content")) + 4 # Message structure overhead. + return total + + +# --------------------------------------------------------------------------- +# Startup +# --------------------------------------------------------------------------- + +def preflight() -> bool: + files = find_auth_files() + sys.stderr.write("==== 预检 ====\n") + sys.stderr.write(f"平台 : {sys.platform}\n") + sys.stderr.write(f"Python : {sys.version.split()[0]}\n") + sys.stderr.write(f"后端 : {BACKEND} (直连,原生 function calling)\n") + sys.stderr.write(f"自管目录 : {managed_auth_dir()}\n") + sys.stderr.write(f"登录文件 : {len(files)} 个\n") + if not os.environ.get("CODEBUDDY_AUTH_DIR"): + sys.stderr.write(f"种子来源 : {', '.join(str(d) for d in auth_dirs())}\n") + ok = True + if not files: + sys.stderr.write("\n[警告] 未找到登录文件。请运行 python3 converter.py login 扫码添加账号,或用 --auth-file 指定。\n") + ok = False + for af in files: + try: + cm = CredentialManager(af) + info = cm.summary() + sys.stderr.write(f"账号 : {info.get('nickname')} / {info.get('enterpriseName')} ({af.name})\n") + sys.stderr.write(f"token过期 : {'是(将自动刷新)' if info['token_expired'] else '否'}\n") + except Exception as e: + sys.stderr.write(f"[警告] 读取凭据失败 {af.name}:{e}\n") + ok = False + sys.stderr.write("================\n") + return ok + + +def login(site: str = "cn", open_browser: bool = True) -> int: + """Complete browser login and persist managed credentials without local HTTP calls.""" + import webbrowser + + try: + started = _OAUTH.start(site=site) + uri = started["verification_uri"] + print(f"请打开以下链接扫码登录:\n{uri}", flush=True) + if open_browser: + try: + opened = webbrowser.open(uri) + except webbrowser.Error: + opened = False + if not opened: + print("无法自动打开浏览器,请手动打开上面的链接。", flush=True) + print("正在等待扫码授权;网页显示登录成功后,请继续等待终端确认入库。\n" + "按 Ctrl+C 取消。", flush=True) + while True: + result = _OAUTH.poll(started["login_id"]) + if result.get("done"): + if result.get("error") or not result.get("cred"): + print(f"登录失败:{result.get('error') or '未获取到凭据'}", file=sys.stderr) + return 1 + target = _save_oauth_credential(result["cred"]) + print(f"登录成功,账号已保存至:{target}\n" + "使用同一凭据目录的服务会在下次请求时自动加载(默认目录扫描模式)。", + flush=True) + return 0 + time.sleep(1.5) + except KeyboardInterrupt: + print("\n已取消登录。", file=sys.stderr) + return 130 + except CredentialFileError as e: + print(f"登录失败:{e}", file=sys.stderr) + return 1 + except OSError: + print("登录失败:无法保存凭据,请检查凭据目录的写入权限。", file=sys.stderr) + return 1 + except (httpx.HTTPError, ValueError, RuntimeError): + # Upstream errors may contain authorization URLs or sensitive response data. + print("登录失败:登录接口请求失败或响应无效,请检查网络后重试。", file=sys.stderr) + return 1 + + +def _nonnegative_int(value): + number = int(value) + if number < 0: + raise argparse.ArgumentTypeError("必须为非负整数") + return number + + +def _positive_int(value): + number = _nonnegative_int(value) + if number == 0: + raise argparse.ArgumentTypeError("必须为正整数") + return number + +def _projection_bytes_arg(value): + number = _nonnegative_int(value) + if number != 0 and number < 256: + raise argparse.ArgumentTypeError("必须为 0 或至少 256") + if number > 33554432: + raise argparse.ArgumentTypeError("不能超过 33554432") + return number + + +def _origins_arg(value): + from app.settings import normalize_allowed_origins + try: + return normalize_allowed_origins(value) + except ValueError: + raise argparse.ArgumentTypeError("必须为逗号分隔的 http/https 来源或域名") from None + + +def _boolean_arg(value): + if isinstance(value, bool): + return value + normalized = str(value).strip().lower() + if normalized in ("true", "1", "yes", "on"): + return True + if normalized in ("false", "0", "no", "off"): + return False + raise argparse.ArgumentTypeError("必须为 true 或 false") + + +def main(): + dotenv_keys = set() if any(arg in ("-h", "--help") for arg in sys.argv[1:]) else load_startup_env() + ap = argparse.ArgumentParser(description="CodeBuddy -> OpenAI 兼容转换器(直连后端)") + ap.add_argument("command", nargs="?", choices=("serve", "login"), default="serve", + help="serve 启动服务(默认);login 扫码登录、自动轮询并保存账号") + ap.add_argument("--site", choices=tuple(auth_oauth.SITE_HOSTS), default="cn", + help="login 站点:cn 国内站(默认);intl 国际 WorkBuddy;intl-codebuddy 国际 CodeBuddy") + ap.add_argument("--no-browser", action="store_true", + help="login 仅显示授权链接,不自动打开浏览器(服务器/容器环境)") + ap.add_argument("--host", default="127.0.0.1", help="监听地址;覆盖 CODEBUDDY2API_BIND") + ap.add_argument("--port", type=int, default=8787, help="监听端口;覆盖 CODEBUDDY2API_PORT") + ap.add_argument("--api-key", default=os.environ.get("CODEBUDDY2API_KEY", ""), + help="管理与推理密钥;未配置时首次本地交互启动生成并保存") + ap.add_argument("--admin-csrf", type=_boolean_arg, nargs="?", const=True, + default=os.environ.get("CODEBUDDY2API_ADMIN_CSRF", "true"), + help="管理 Origin/CSRF 校验,默认 true;仅在受信任本地环境设为 false,鉴权仍启用") + ap.add_argument("--admin-allowed-origins", type=_origins_arg, metavar="ORIGINS", + default=os.environ.get("CODEBUDDY2API_ADMIN_ORIGINS"), + help="额外信任的管理页来源(逗号分隔,支持域名或完整来源,裸域名按 https);" + "反代 HTTPS 域名登录报 Origin 校验失败时设置,也可在 WebUI 配置") + ap.add_argument("--log", default=None, metavar="PATH", + help="已停用:日志统一保存到数据目录中的 logs.sqlite3") + ap.add_argument("--desensitize", action="store_true", + help="适配固定 CLI 模板、压缩运行时提示并零宽脱敏关键词。默认关闭。") + ap.add_argument("--no-compact", action="store_true", + help="配合 --desensitize 保留主要行为指令,仍适配固定模板并裁剪运行时元数据;" + "非流式纯审核拒绝最多压缩兜底一次。") + ap.add_argument("--keep-tool-metadata", type=_boolean_arg, nargs="?", const=True, + default=os.environ.get("CODEBUDDY2API_KEEP_TOOL_METADATA", "false"), + help="保留工具描述及参数 description/title;启用脱敏时仍处理描述文本,默认 false") + ap.add_argument("--responses-projection-mode", choices=PROJECTION_MODES, + default=os.environ.get("CODEBUDDY2API_RESPONSES_PROJECTION_MODE", "balanced"), + help="Responses 上下文:balanced 仅改写有固定摘要的已识别 harness,其余文本原样保留;passthrough 完全关闭投影;默认 balanced") + ap.add_argument("--responses-projection-max-bytes", type=_projection_bytes_arg, metavar="BYTES", + default=os.environ.get("CODEBUDDY2API_RESPONSES_PROJECTION_MAX_BYTES", "40000"), + help="Responses 单项 UTF-8 字节上限;0 禁用,非零范围 256..33554432,默认 40000") + ap.add_argument("--skip-check", action="store_true", help="跳过启动预检") + ap.add_argument("--auth-file", action="append", default=[], metavar="PATH", + help="凭据文件(可重复传入组成凭证池;默认自动扫描 auth 目录全部 *.info)") + ap.add_argument("--credit-price-cny", type=float, default=None, metavar="PRICE", + help="积分折算单价(元/Credit),默认 0.014(旗舰版连续包月 700元/5万积分摊算)") + ap.add_argument("--usd-rate", type=float, default=None, metavar="RATE", + help="人民币→美元汇率,影响 /v1/dashboard/billing 端点金额") + ap.add_argument("--credit-price-usd", type=float, default=None, metavar="PRICE", + help="国际站积分折算单价(美元/Credit),默认 0.03(Pro 加量包 $15/500 积分)") + ap.add_argument("--model-catalog-ttl", type=int, default=6 * 3600, metavar="SECONDS", + help="云端模型表缓存有效期,默认 21600 秒(6 小时);TTL 内不再打 /v3/config") + ap.add_argument("--no-model-guard", action="store_true", + help="关闭表外模型本地拦截;默认拦截,避免无效请求打到上游并触发扣费") + ap.add_argument("--max-images", type=_nonnegative_int, metavar="N", + default=os.environ.get("CODEBUDDY2API_MAX_IMAGES", "16"), + help="单请求图片上限,默认 16;0 表示不允许图片") + ap.add_argument("--image-policy", choices=("truncate", "error"), + default=os.environ.get("CODEBUDDY2API_IMAGE_POLICY", "truncate"), + help="超额图片策略:truncate 保留最新图片(默认),error 返回 413") + ap.add_argument("--max-request-bytes", type=_positive_int, metavar="BYTES", + default=os.environ.get("CODEBUDDY2API_MAX_REQUEST_BYTES", str(32 * 1024 * 1024)), + help="图片处理与适配后请求体的字节上限,默认 32 MiB") + ap.add_argument("--max-inbound-bytes", type=_positive_int, metavar="BYTES", + default=os.environ.get("CODEBUDDY2API_MAX_INBOUND_BYTES", str(64 * 1024 * 1024)), + help="入站原始请求体字节上限(解析前生效,含 chunked),默认 64 MiB") + ap.add_argument("--max-collect-bytes", type=_nonnegative_int, metavar="BYTES", + default=os.environ.get("CODEBUDDY2API_MAX_COLLECT_BYTES", str(8 * 1024 * 1024)), + help="聚合路径输出收集总字节上限(正文+思考+工具参数),默认 8 MiB;0 不限制") + ap.add_argument("--max-concurrent", type=_nonnegative_int, metavar="N", + default=os.environ.get("CODEBUDDY2API_MAX_CONCURRENT", "64"), + help="推理端点并发上限(超出立即 503),默认 64;0 不限制") + ap.add_argument("--max-inflight-per-account", type=_nonnegative_int, metavar="N", + default=os.environ.get("CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT", "0"), + help="单账号在途上限,默认 0(不限制);满载返回 503,不借容量切换到收费账号") + ap.add_argument("--upstream-keepalive", type=_boolean_arg, nargs="?", const=True, + default=os.environ.get("CODEBUDDY2API_UPSTREAM_KEEPALIVE", "false"), + help="按上游入口复用有界连接池,默认 false;重启生效,不改变超时或重放规则") + ap.add_argument("--request-context-mode", choices=("legacy", "scoped"), + default=os.environ.get("CODEBUDDY2API_REQUEST_CONTEXT_MODE", "legacy"), + help="请求上下文:legacy 保持旧会话头,scoped 启用显式会话与逐尝试追踪;默认 legacy") + ap.add_argument("--stream-mode", choices=_STREAM_MODES, + default=os.environ.get("CODEBUDDY2API_STREAM_MODE", "compatible"), + help="流式传输:compatible 保持兼容聚合策略,realtime 增量发送且不重生成工具参数;默认 compatible") + ap.add_argument("--model-capability-guard", type=_boolean_arg, nargs="?", const=True, + default=os.environ.get("CODEBUDDY2API_MODEL_CAPABILITY_GUARD", "true"), + help="按账号模型声明预检图片、工具、思考和输出上限;false 仅关闭新增能力预检") + ap.add_argument("--log-body-limit", type=_nonnegative_int, metavar="BYTES", + default=os.environ.get("CODEBUDDY2API_LOG_BODY_LIMIT", "65536"), + help="旧文本预览兼容项;文本输出已停用,SQLite 诊断使用独立预算") + ap.add_argument("--tool-call-max-retry", type=_nonnegative_int, metavar="N", + default=os.environ.get("CODEBUDDY2API_TOOL_CALL_MAX_RETRY", "3"), + help="工具参数损坏时的额外生成上限,默认 3;0 表示不重试(每次额外生成都消耗额度)") + ap.add_argument("--failover-max", type=_nonnegative_int, metavar="N", + default=os.environ.get("CODEBUDDY2API_FAILOVER_MAX", "0"), + help="失败发生在向下游落第一个字节之前时,最多换几个凭证就地重放,默认 0(关闭);" + "只重放上游没收下请求体或用 401/403/429/502/503/504 拒绝的失败;" + "写请求体超时需另开 --retry-write-timeout 才参与") + ap.add_argument("--retry-write-timeout", type=_boolean_arg, nargs="?", const=True, + default=os.environ.get("CODEBUDDY2API_RETRY_WRITE_TIMEOUT", "false"), + help="把「写请求体超时」也算作上游没收下请求体从而参与重放,默认 false。写超时只能" + "证明正文没写完,上游是否已按半截正文计费看不到,因此要显式开启(同时作用于连接" + "重试与 --failover-max 换凭证重放)") + ap.add_argument("--auto-trial", type=_boolean_arg, nargs="?", const=True, + default=None, help=argparse.SUPPRESS) + args = ap.parse_args() + if args.auto_trial is not None or "CODEBUDDY2API_AUTO_TRIAL" in os.environ: + sys.stderr.write("[trial] AUTO_TRIAL / --auto-trial 已停用;请在 WebUI 凭证页手动领取体验积分。\n") + del args.auto_trial + if args.image_policy not in ("truncate", "error"): + ap.error("CODEBUDDY2API_IMAGE_POLICY 必须为 truncate 或 error") + if args.command == "login": + return login(site=args.site, open_browser=not args.no_browser) + + for key in ("max_images", "image_policy", "max_request_bytes", "log_body_limit", + "tool_call_max_retry", "max_inbound_bytes", "max_collect_bytes", "max_concurrent", + "failover_max", "retry_write_timeout", "upstream_keepalive", "max_inflight_per_account", + "request_context_mode", "stream_mode", "model_capability_guard", "admin_allowed_origins", + "responses_projection_mode", "responses_projection_max_bytes"): + CONFIG[key] = getattr(args, key) + CONFIG["api_key"] = args.api_key + CONFIG["desensitize"] = args.desensitize + CONFIG["no_compact"] = args.no_compact + CONFIG["credit_price_cny"] = args.credit_price_cny or None + CONFIG["usd_rate"] = args.usd_rate or None + CONFIG["credit_price_usd"] = args.credit_price_usd or None + CONFIG["model_guard"] = not args.no_model_guard + if args.log is not None or os.environ.get("CODEBUDDY2API_LOG"): + sys.stderr.write("[log] --log / CODEBUDDY2API_LOG 已停用;请在 WebUI 查看 SQLite 日志。\n") + CONFIG["log_path"] = None + from app import runtime_management + runtime_management.initialize(sys.modules[__name__], args, parser=ap, dotenv_keys=dotenv_keys) + # Validate effective binding and authentication before credential scans or background work. + if (args.host not in ("127.0.0.1", "::1", "localhost") and not CONFIG.get("api_key") + and not allows_open_noauth()): + runtime_management.close(CONFIG) + ap.error("非回环绑定且未设置 API key 会匿名开放推理额度;" + "请设置 CODEBUDDY2API_KEY,或确知风险后以 CODEBUDDY2API_ALLOW_OPEN_NOAUTH=true 显式放行") + files = [Path(p) for p in args.auth_file] + if not files: + seed_credentials() # Seed missing desktop credentials into managed storage. + CONFIG["cred_pool"] = CredentialPool(files, scan=not files, state_store=CONFIG["state_store"]) + CONFIG["usage_snapshots"] = UsageSnapshots(store=CONFIG["state_store"]) + CONFIG["cred"] = CONFIG["cred_pool"].first() + CONFIG["account_catalogs"] = {} # Disable static fallback before maintenance starts. + if credits_mod is not None: + ledger = credits_mod.CreditLedger(store=CONFIG["state_store"]) + CONFIG["ledger"] = ledger + CONFIG["model_cache"] = credits_mod.ModelCatalogCache(ttl=args.model_catalog_ttl, store=CONFIG["state_store"]) + CONFIG["cred_pool"].set_ledger(ledger) # Verify balance ownership before publishing catalogs. + _publish_model_cache() + # Publish cached usage before maintenance threads start, so the dashboard is populated from + # the first request. It stays a no-op when nothing is cached, so an empty deployment and a + # mocked pool are both unaffected. + if CONFIG["usage_snapshots"].detail(): + _publish_usage_daily(CONFIG["cred_pool"]) + + try: + runtime_management.install(sys.modules[__name__]) + except SessionStoreError as error: + # An obsolete session snapshot survived, so the new key epoch must not activate: + # a later start under the superseded key could adopt it and revive admin cookies. + runtime_management.close(CONFIG) + ap.error(str(error)) + threading.Thread(target=_refresher_loop, args=(CONFIG["cred_pool"],), + daemon=True, name="cred-refresher").start() + if credits_mod is not None: + threading.Thread(target=_housekeeper_loop, args=(CONFIG["cred_pool"], ledger), + daemon=True, name="cred-housekeeper").start() + + if not args.skip_check: + preflight() + + sys.stderr.write(f"\n✅ 监听 http://{args.host}:{args.port}(直连后端,原生 function calling)\n") + sys.stderr.write(f" WebUI : http://{args.host}:{args.port}/dashboard\n") + sys.stderr.write(" SQLite 审计默认开启,凭证仍以 .info 文件保存\n") + sys.stderr.write(" GET /v1/models\n") + sys.stderr.write(" POST /v1/chat/completions (原生 tools/tool_calls,支持流式)\n") + sys.stderr.write(" POST /v1/responses (Responses API,Codex CLI 兼容)\n") + sys.stderr.write(" POST /v1/messages (Anthropic API,Claude Code / CC Switch 兼容)\n") + sys.stderr.write(" GET /health\n") + sys.stderr.write(" GET/POST/DELETE /admin/credentials (凭证池管理)\n") + sys.stderr.write(" 添加账号:python3 converter.py login(自动等待扫码并保存)\n") + if credits_mod is not None: + sys.stderr.write(" GET /admin/credits (积分/签到状态)\n") + sys.stderr.write(" POST /admin/checkin (仅签到,按日幂等)\n") + sys.stderr.write(" POST /admin/sync (同步余额、目录与用量,不签到)\n") + sys.stderr.write(" POST /admin/credentials/{id}/reset-cooldown (清除该账号冷却,仅本地)\n") + sys.stderr.write(" 每日签到 + 快过期积分优先调度已启用\n") + if args.api_key: + sys.stderr.write(" 鉴权已启用(API key 已设置)\n") + if not CONFIG["admin_csrf"]: + sys.stderr.write(" 警告:管理 Origin/CSRF 校验已关闭,仅限受信任本地环境;API key 与会话校验仍启用。\n") + sys.stderr.write(f" 图片限制 : {CONFIG['max_images']} 张/请求,策略 {CONFIG['image_policy']}\n") + if CONFIG["log_path"]: + sys.stderr.write(f" 日志 : {CONFIG['log_path']}\n") + if args.desensitize: + mode = "零宽脱敏 + 保留全文" if args.no_compact else "零宽脱敏 + 压缩摘要" + sys.stderr.write(f" 脱敏 : 已启用({mode})\n") + sys.stderr.write("按 Ctrl+C 退出。\n\n") + + # Record service startup. + _log(f"==== converter 启动 ====") + + try: + run_server(app, CONFIG, host=args.host, port=args.port) + finally: + runtime_management.close(CONFIG) + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/test_responses_adapter.py b/tests/test_responses_adapter.py index eab1ff2..cf663aa 100644 --- a/tests/test_responses_adapter.py +++ b/tests/test_responses_adapter.py @@ -1,650 +1,720 @@ -#!/usr/bin/env python3 -"""Test Responses request and response adaptation.""" - -import json -import sys -from pathlib import Path -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) # Allow direct execution. - -from app.adapters.responses_adapter import ( - responses_request_to_chat, - ResponsesStreamConverter, -) -from app.desensitize import desensitize_body -from app.adapters.responses_projection import project_responses_chat_body - - -def test_simple_text_request(): - """Convert plain text input into Chat messages.""" - req = { - "model": "glm-5.2", - "input": "Hello, how are you?", - "instructions": "You are a helpful assistant.", - "stream": True, - } - chat = responses_request_to_chat(req) - assert chat["messages"][0] == {"role": "system", "content": "You are a helpful assistant."} - assert chat["messages"][1] == {"role": "user", "content": "Hello, how are you?"} - assert chat["model"] == "glm-5.2" - print("✅ test_simple_text_request") - - -def test_array_input_request(): - """Convert mixed message, function-call and function-output items.""" - req = { - "model": "glm-5.2", - "input": [ - {"role": "user", "content": "Fix the bug"}, - {"type": "message", "id": "msg_1", "role": "assistant", - "content": [{"type": "output_text", "text": "I'll check the file."}]}, - {"type": "function_call", "id": "fc_1", "call_id": "call_123", - "name": "shell", "arguments": '{"cmd":"cat main.py"}'}, - {"type": "function_call_output", "call_id": "call_123", - "output": "print('hello')"}, - {"role": "user", "content": "Now fix it"}, - ], - "instructions": "You are a coding assistant.", - } - chat = responses_request_to_chat(req) - msgs = chat["messages"] - - assert msgs[0] == {"role": "system", "content": "You are a coding assistant."} - assert msgs[1] == {"role": "user", "content": "Fix the bug"} - assert msgs[2]["role"] == "assistant" - assert msgs[2]["content"] == "I'll check the file." - assert len(msgs[2]["tool_calls"]) == 1 - assert msgs[2]["tool_calls"][0]["function"]["name"] == "shell" - assert msgs[3]["role"] == "tool" - assert msgs[3]["tool_call_id"] == "call_123" - assert msgs[4] == {"role": "user", "content": "Now fix it"} - print("✅ test_array_input_request") - -def test_function_call_arguments_must_be_string(): - """Reject non-standard object arguments before Chat adaptation.""" - try: - responses_request_to_chat({"model": "auto", "input": [ - {"type": "function_call", "name": "tool", "arguments": {"x": 1}}, - ]}) - except ValueError as error: - assert "JSON string" in str(error) - else: - raise AssertionError("non-string function call arguments were accepted") - - -def test_tools_conversion(): - """Convert flat Responses tools to nested Chat definitions.""" - req = { - "model": "glm-5.2", - "input": "test", - "tools": [ - {"type": "function", "name": "shell", - "description": "Run a shell command", - "parameters": {"type": "object", "properties": {"cmd": {"type": "string"}}}}, - ], - } - chat = responses_request_to_chat(req) - tool = chat["tools"][0] - assert tool["type"] == "function" - assert "function" in tool - assert tool["function"]["name"] == "shell" - print("✅ test_tools_conversion") - - -def test_max_output_tokens(): - """Map max_output_tokens to max_tokens.""" - req = {"model": "glm-5.2", "input": "test", "max_output_tokens": 4096} - chat = responses_request_to_chat(req) - assert chat["max_tokens"] == 4096 - print("✅ test_max_output_tokens") - - -def test_developer_role(): - """Normalize developer roles to system roles.""" - req = {"model": "glm-5.2", "input": [ - {"role": "developer", "content": "Be concise."}, - {"role": "user", "content": "Hi"}, - ]} - chat = responses_request_to_chat(req) - assert chat["messages"][0] == {"role": "system", "content": "Be concise."} - assert chat["messages"][1] == {"role": "user", "content": "Hi"} - print("✅ test_developer_role") +#!/usr/bin/env python3 +"""Test Responses request and response adaptation.""" + +import json +import sys +from pathlib import Path +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) # Allow direct execution. + +from app.adapters.responses_adapter import ( + responses_request_to_chat, + ResponsesStreamConverter, +) +from app.desensitize import desensitize_body +from app.adapters.responses_projection import project_responses_chat_body + + +def test_simple_text_request(): + """Convert plain text input into Chat messages.""" + req = { + "model": "glm-5.2", + "input": "Hello, how are you?", + "instructions": "You are a helpful assistant.", + "stream": True, + } + chat = responses_request_to_chat(req) + assert chat["messages"][0] == {"role": "system", "content": "You are a helpful assistant."} + assert chat["messages"][1] == {"role": "user", "content": "Hello, how are you?"} + assert chat["model"] == "glm-5.2" + print("✅ test_simple_text_request") + + +def test_array_input_request(): + """Convert mixed message, function-call and function-output items.""" + req = { + "model": "glm-5.2", + "input": [ + {"role": "user", "content": "Fix the bug"}, + {"type": "message", "id": "msg_1", "role": "assistant", + "content": [{"type": "output_text", "text": "I'll check the file."}]}, + {"type": "function_call", "id": "fc_1", "call_id": "call_123", + "name": "shell", "arguments": '{"cmd":"cat main.py"}'}, + {"type": "function_call_output", "call_id": "call_123", + "output": "print('hello')"}, + {"role": "user", "content": "Now fix it"}, + ], + "instructions": "You are a coding assistant.", + } + chat = responses_request_to_chat(req) + msgs = chat["messages"] + + assert msgs[0] == {"role": "system", "content": "You are a coding assistant."} + assert msgs[1] == {"role": "user", "content": "Fix the bug"} + assert msgs[2]["role"] == "assistant" + assert msgs[2]["content"] == "I'll check the file." + assert len(msgs[2]["tool_calls"]) == 1 + assert msgs[2]["tool_calls"][0]["function"]["name"] == "shell" + assert msgs[3]["role"] == "tool" + assert msgs[3]["tool_call_id"] == "call_123" + assert msgs[4] == {"role": "user", "content": "Now fix it"} + print("✅ test_array_input_request") + +def test_function_call_arguments_must_be_string(): + """Reject non-standard object arguments before Chat adaptation.""" + try: + responses_request_to_chat({"model": "auto", "input": [ + {"type": "function_call", "name": "tool", "arguments": {"x": 1}}, + ]}) + except ValueError as error: + assert "JSON string" in str(error) + else: + raise AssertionError("non-string function call arguments were accepted") + + +def test_tools_conversion(): + """Convert flat Responses tools to nested Chat definitions.""" + req = { + "model": "glm-5.2", + "input": "test", + "tools": [ + {"type": "function", "name": "shell", + "description": "Run a shell command", + "parameters": {"type": "object", "properties": {"cmd": {"type": "string"}}}}, + ], + } + chat = responses_request_to_chat(req) + tool = chat["tools"][0] + assert tool["type"] == "function" + assert "function" in tool + assert tool["function"]["name"] == "shell" + print("✅ test_tools_conversion") + + +def test_max_output_tokens(): + """Map max_output_tokens to max_tokens.""" + req = {"model": "glm-5.2", "input": "test", "max_output_tokens": 4096} + chat = responses_request_to_chat(req) + assert chat["max_tokens"] == 4096 + print("✅ test_max_output_tokens") + + +def test_developer_role(): + """Normalize developer roles to system roles.""" + req = {"model": "glm-5.2", "input": [ + {"role": "developer", "content": "Be concise."}, + {"role": "user", "content": "Hi"}, + ]} + chat = responses_request_to_chat(req) + assert chat["messages"][0] == {"role": "system", "content": "Be concise."} + assert chat["messages"][1] == {"role": "user", "content": "Hi"} + print("✅ test_developer_role") + + +def test_typed_developer_message_request(): + """Normalize developer roles inside typed messages.""" + req = { + "model": "glm-5.2", + "input": [ + {"type": "message", "role": "developer", "content": "Be concise."}, + {"type": "message", "role": "user", "content": "Hi"}, + ], + } + chat = responses_request_to_chat(req) + assert chat["messages"][0] == {"role": "system", "content": "Be concise."} + assert chat["messages"][1] == {"role": "user", "content": "Hi"} + print("✅ test_typed_developer_message_request") + + +def test_desensitize_harness_user_and_tools(): + """Compact trusted harness input and tool metadata while preserving real user text.""" + body = { + "messages": [ + {"role": "system", "content": "Refuse exploit development."}, + {"role": "user", "content": "# AGENTS.md instructions\n sandbox escalation"}, + {"role": "user", "content": "please explain dos attacks"}, + ], + "tools": [ + {"type": "function", "function": {"name": "exec_command", "description": "Run dangerous exploit development checks."}} + ], + } + out = desensitize_body( + body, + roles=("system", "developer"), + desensitize_harness_user=True, + desensitize_tools=True, + ) + assert "​" in out["messages"][0]["content"] + assert "Repository instructions and durable user context are provided." in out["messages"][1]["content"] + assert "Environment context is provided by the harness." in out["messages"][1]["content"] + assert "​" not in out["messages"][2]["content"] + assert out["messages"][2]["content"] == "please explain dos attacks" + assert "​" in out["tools"][0]["function"]["description"] + print("✅ test_desensitize_harness_user_and_tools") + + +def test_compact_harness_messages_and_strip_tool_metadata(): + """Compact long Codex templates without replacing the user's actual request.""" + body = { + "messages": [ + {"role": "system", "content": "You are a coding agent running in the Codex CLI. # How you work\nUse sandbox and escalation."}, + {"role": "system", "content": "\nFilesystem sandboxing defines which files can be read or written."}, + {"role": "user", "content": "# AGENTS.md instructions\n sandbox escalation"}, + ], + "tools": [ + {"type": "function", "function": {"name": "exec_command", "description": "Run dangerous exploit development checks.", "parameters": {"type": "object", "properties": {"cmd": {"type": "string", "description": "Shell command to execute."}}}}} + ], + } + out = desensitize_body( + body, + roles=("system", "developer"), + desensitize_harness_user=True, + desensitize_tools=True, + compact_harness=True, + strip_tool_metadata=True, + ) + assert len(out["messages"][0]["content"]) < 220 + assert "Codex CLI" in out["messages"][0]["content"] + assert "sandboxing defines" not in out["messages"][1]["content"] + assert "Repository instructions and durable user context are provided." in out["messages"][2]["content"] + assert "Environment context is provided by the harness." in out["messages"][2]["content"] + assert "description" not in out["tools"][0]["function"] + assert "description" not in out["tools"][0]["function"]["parameters"]["properties"]["cmd"] + print("✅ test_compact_harness_messages_and_strip_tool_metadata") + + +def test_no_compact_still_prunes_codex_runtime_metadata(): + """Remove trusted runtime metadata even when full conversation text is preserved.""" + body = { + "messages": [ + { + "role": "system", + "content": ( + "You are a coding agent running in the Codex CLI.\n\n" + "# How you work\nUse sandbox and escalation carefully.\n\n" + "\nFilesystem sandboxing defines which files can be read or written.\n" + "## How to request escalation\n...\n\n\n" + "The following deferred tools are now available via ToolSearch.\n..." + ), + }, + { + "role": "user", + "content": ( + "# AGENTS.md instructions\n\nproject guidance\n" + "\nvery long runtime context\n\n" + "\nvery long skills metadata\n\n" + "test" + ), + }, + {"role": "user", "content": "test"}, + ], + } + out = desensitize_body( + body, + roles=("system", "developer"), + desensitize_harness_user=True, + compact_harness=False, + ) + system_text = out["messages"][0]["content"] + harness_text = out["messages"][1]["content"] + assert "You are a coding agent running in the Codex CLI" in system_text + assert "## Planning" not in system_text + assert "## Task execution" not in system_text + assert "### Final answer structure and style guidelines" not in system_text + assert "# How you work" in system_text + assert "Filesystem sandboxing defines" not in system_text + # Without a closed wrapper the deferred-tool paragraph is real text, not metadata. + assert "The following deferred tools are now available via ToolSearch.\n..." in system_text + assert "Runtime permissions apply" in system_text + assert "Runtime tool, agent, sk" not in system_text + assert "very long runtime context" not in harness_text + assert "very long skills metadata" not in harness_text + assert "# AGENTS.md instructions" not in harness_text + assert "Repository instructions and durable user context are provided." in harness_text + assert "Environment context is provided by the harness." in harness_text + # Remove inserted separators to isolate harness compaction from term adaptation. + assert "Runtime skill metadata is available" in harness_text.replace("​", "") + assert harness_text.strip().replace("​", "").endswith("test") + assert out["messages"][2]["content"] == "test" + print("✅ test_no_compact_still_prunes_codex_runtime_metadata") + + +def test_responses_projection_balanced_preserves_real_text_tools_and_structure(): + """Balanced mode replaces recognized harness blocks without changing tools or real text.""" + tool = { + "type": "function", + "function": { + "name": "exec_command", + "description": "Run a command", + "parameters": { + "type": "object", + "properties": {"cmd": {"type": "string", "description": "Command"}}, + "required": ["cmd"], + "additionalProperties": False, + "x-vendor-detail": {"deep": {"schema": "kept"}}, + }, + "strict": False, + }, + } + body = { + "model": "auto", + "messages": [ + {"role": "system", "content": "Repository policy: run tests."}, + {"role": "user", "content": "# AGENTS.md instructions\n\nvolatile context\n"}, + {"role": "user", "content": "实现该方案"}, + ], + "tools": [tool], + } + before = json.loads(json.dumps(body, ensure_ascii=False)) + out, stats = project_responses_chat_body(body) + assert body == before + assert out["tools"] == body["tools"] + assert out["messages"][0] == body["messages"][0] + assert out["messages"][1]["content"] != body["messages"][1]["content"] + assert "# AGENTS.md instructions" not in out["messages"][1]["content"] + assert "Environment context is provided by the harness." in out["messages"][1]["content"] + assert out["messages"][2] == body["messages"][2] + assert stats["mode"] == "balanced" + assert stats["original_tools"] == stats["projected_tools"] == 1 + assert stats["original_tool_chars"] == stats["projected_tool_chars"] + assert stats["harness_messages_projected"] == 1 + print("✅ test_responses_projection_balanced_preserves_real_text_tools_and_structure") + + +def test_responses_projection_preserves_history_and_tool_chain(): + """Balanced mode does not summarize history or drop tool-call relationships.""" + body = { + "messages": [ + {"role": "user", "content": "old task"}, + {"role": "assistant", "content": "old answer", "tool_calls": [{ + "id": "call_old", "type": "function", + "function": {"name": "exec_command", "arguments": '{"cmd":"ls"}'}, + }]}, + {"role": "tool", "tool_call_id": "call_old", "content": "old output"}, + {"role": "user", "content": "new task"}, + ], + "tools": [{"type": "function", "function": {"name": "exec_command", "parameters": {"type": "object"}}}], + } + before = json.loads(json.dumps(body)) + out, stats = project_responses_chat_body(body, max_item_bytes=40000) + assert body == before + assert out["messages"] == body["messages"] + assert stats["original_messages"] == stats["projected_messages"] == 4 + assert "anchor_user_preserved" not in stats + print("✅ test_responses_projection_preserves_history_and_tool_chain") + + +def test_responses_projection_truncates_generated_content_and_json_arguments(): + """Oversized assistant, tool output, and JSON values retain UTF-8 head and tail.""" + long_text = "HEAD\n" + ("中" * 180) + "\nTAIL" + long_output = "OUTPUT\n" + ("输出" * 180) + "\nEND" + arguments = json.dumps({"cmd": "echo " + ("x" * 500), "workdir": "/tmp"}) + apply_patch = json.dumps({"patch": "*** Begin Patch\n" + ("+" * 500) + "*** End Patch"}) + body = {"messages": [ + {"role": "assistant", "content": long_text, "tool_calls": [{ + "id": "call_1", "type": "function", + "function": {"name": "exec_command", "arguments": arguments}, + }]}, + {"role": "tool", "tool_call_id": "call_1", "content": long_output}, + {"role": "assistant", "content": "", "tool_calls": [{ + "id": "call_patch", "type": "function", + "function": {"name": "apply_patch", "arguments": apply_patch}, + }]}, + ]} + before = json.loads(json.dumps(body, ensure_ascii=False)) + out, stats = project_responses_chat_body(body, max_item_bytes=256) + assert body == before + assistant = out["messages"][0]["content"] + assert assistant.startswith("HEAD") + assert assistant.endswith("TAIL") + assert "middle omitted" in assistant + assert "original bytes:" in assistant + assert "estimated tokens:" in assistant + assert "total lines:" in assistant + tool_output = out["messages"][1]["content"] + assert tool_output.startswith("OUTPUT") + assert tool_output.endswith("END") + args_wire = out["messages"][0]["tool_calls"][0]["function"]["arguments"] + args = json.loads(args_wire) + assert len(args_wire.encode("utf-8")) <= 256 + assert "middle omitted" in args["_truncated"]["warning"] + assert '"cmd"' in args["head"] and "echo" in args["head"] + assert "workdir" in args["tail"] and args["tail"].endswith("}") + patch_wire = out["messages"][2]["tool_calls"][0]["function"]["arguments"] + patch_args = json.loads(patch_wire) + assert len(patch_wire.encode("utf-8")) <= 256 + assert '"patch"' in patch_args["head"] and "*** Begin Patch" in patch_args["head"] + assert patch_args["tail"].endswith('*** End Patch"}') + assert stats["truncated_items"] == 4 + assert stats["truncated_original_bytes"] > stats["truncated_projected_bytes"] + print("✅ test_responses_projection_truncates_generated_content_and_json_arguments") + + +def test_responses_projection_passthrough_and_zero_limit_are_lossless(): + """Passthrough and max_item_bytes=0 preserve the request payload.""" + body = {"messages": [ + {"role": "user", "content": "# AGENTS.md instructions\nvolatile\nreal task"}, + {"role": "assistant", "content": "assistant " + "x" * 500}, + {"role": "tool", "tool_call_id": "c", "content": "output " + "y" * 500}, + ], "tools": [{"type": "function", "function": {"name": "tool", "parameters": {"type": "object", "x": 1}}}]} + before = json.loads(json.dumps(body, ensure_ascii=False)) + for kwargs in ({"mode": "passthrough"}, {"max_item_bytes": 0}): + out, stats = project_responses_chat_body(body, **kwargs) + if "mode" in kwargs: + assert out == before + else: + assert out["messages"][1:] == before["messages"][1:] + assert out["tools"] == before["tools"] + assert body == before + assert stats["mode"] == ("passthrough" if "mode" in kwargs else "balanced") + assert stats["truncated_items"] == 0 + print("✅ test_responses_projection_passthrough_and_zero_limit_are_lossless") + + +def test_responses_projection_rejects_invalid_mode_and_limit(): + """Projection validates its public mode and byte-limit arguments.""" + for mode in ("aggressive", "conservative", "", None, True): + try: + project_responses_chat_body({"messages": []}, mode=mode) + except ValueError: + pass + else: + raise AssertionError(f"invalid mode accepted: {mode!r}") + for limit in (-1, 1, 128, 255, True, 1.5, "256", None, [], {}): + try: + project_responses_chat_body({"messages": []}, max_item_bytes=limit) + except ValueError: + pass + else: + raise AssertionError(f"invalid max_item_bytes accepted: {limit!r}") + assert project_responses_chat_body({"messages": []}, mode="balanced", max_item_bytes=256)[1]["max_item_bytes"] == 256 + print("✅ test_responses_projection_rejects_invalid_mode_and_limit") + + +def test_responses_projection_stats_have_official_shape(): + """Stats expose current projection counters without legacy mode fields.""" + body = {"messages": [{"role": "user", "content": "hello"}], "tools": [{"type": "function", "function": {"name": "t"}}]} + _, stats = project_responses_chat_body(body, mode="balanced", max_item_bytes=0) + assert stats == { + "mode": "balanced", "max_item_bytes": 0, "original_messages": 1, + "projected_messages": 1, "original_message_chars": stats["original_message_chars"], + "projected_message_chars": stats["projected_message_chars"], "original_tools": 1, + "projected_tools": 1, "original_tool_chars": stats["original_tool_chars"], + "projected_tool_chars": stats["projected_tool_chars"], "harness_messages_projected": 0, + "truncated_items": 0, "truncated_original_bytes": 0, "truncated_projected_bytes": 0, + } + print("✅ test_responses_projection_stats_have_official_shape") + + +def test_stream_converter_text(): + """Convert Chat text SSE into Responses events.""" + conv = ResponsesStreamConverter(model="glm-5.2") + + # Synthetic Chat SSE chunks + chunks = [ + 'data: {"id":"chatcmpl-1","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}', + 'data: {"id":"chatcmpl-1","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}', + 'data: {"id":"chatcmpl-1","choices":[{"index":0,"delta":{"content":" world"},"finish_reason":null}]}', + 'data: {"id":"chatcmpl-1","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":2,"total_tokens":12}}', + 'data: [DONE]', + ] + + all_events = [] + for line in chunks: + result = conv.feed_line(line) + if result: + for evt_line in result.strip().split("\n\n"): + if evt_line.startswith("data: "): + all_events.append(json.loads(evt_line[6:])) + + finish = conv.finish() + for evt_line in finish.strip().split("\n\n"): + if evt_line.startswith("data: "): + all_events.append(json.loads(evt_line[6:])) + + types = [e["type"] for e in all_events] + assert "response.created" in types + assert "response.in_progress" in types + assert "response.output_item.added" in types + assert "response.content_part.added" in types + assert "response.output_text.delta" in types + assert "response.output_text.done" in types + assert "response.content_part.done" in types + assert "response.output_item.done" in types + assert "response.completed" in types + + text_done = [e for e in all_events if e["type"] == "response.output_text.done"][0] + assert text_done["text"] == "Hello world" + + completed = [e for e in all_events if e["type"] == "response.completed"][0] + resp = completed["response"] + assert resp["status"] == "completed" + assert resp["output"][0]["type"] == "message" + assert resp["output"][0]["content"][0]["text"] == "Hello world" + assert resp["usage"]["input_tokens"] == 10 + + print("✅ test_stream_converter_text") + + +def test_stream_converter_function_call(): + """Convert Chat tool-call SSE into Responses function-call events.""" + conv = ResponsesStreamConverter(model="glm-5.2") + + chunks = [ + 'data: {"id":"chatcmpl-2","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"id":"call_abc","type":"function","function":{"name":"shell","arguments":""}}]},"finish_reason":null}]}', + 'data: {"id":"chatcmpl-2","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\\"cmd"}}]},"finish_reason":null}]}', + 'data: {"id":"chatcmpl-2","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\\": \\"ls\\"}"}}]},"finish_reason":null}]}', + 'data: {"id":"chatcmpl-2","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}', + 'data: [DONE]', + ] + + all_events = [] + for line in chunks: + result = conv.feed_line(line) + if result: + for evt_line in result.strip().split("\n\n"): + if evt_line.startswith("data: "): + all_events.append(json.loads(evt_line[6:])) + + finish = conv.finish() + for evt_line in finish.strip().split("\n\n"): + if evt_line.startswith("data: "): + all_events.append(json.loads(evt_line[6:])) + + types = [e["type"] for e in all_events] + assert "response.output_item.added" in types + assert "response.function_call_arguments.delta" in types + assert "response.function_call_arguments.done" in types + assert "response.completed" in types + + args_done = [e for e in all_events if e["type"] == "response.function_call_arguments.done"][0] + assert args_done["arguments"] == '{"cmd": "ls"}' + + print("✅ test_stream_converter_function_call") + + +def test_nonstream_response(): + """Build a non-streaming Response object.""" + conv = ResponsesStreamConverter(model="glm-5.2") + conv.feed_line('data: {"id":"c1","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":null}]}') + conv.feed_line('data: {"id":"c1","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}}') + + resp = conv.get_nonstream_response() + assert resp["object"] == "response" + assert resp["status"] == "completed" + assert resp["output"][0]["type"] == "message" + assert resp["output"][0]["content"][0]["text"] == "Hi" + assert resp["usage"]["input_tokens"] == 5 + + print("✅ test_nonstream_response") + + +def test_finish_reason_maps_to_terminal_status(): + """Report truncated or filtered output as incomplete in streaming and aggregated responses.""" + conv = ResponsesStreamConverter(model="glm-5.2") + conv.feed_line('data: {"id":"c2","choices":[{"index":0,"delta":{"content":"partial"},"finish_reason":null}]}') + conv.feed_line('data: {"id":"c2","choices":[{"index":0,"delta":{},"finish_reason":"length"}]}') + tail = conv.finish() + assert '"type": "response.incomplete"' in tail and '"type": "response.completed"' not in tail + resp = conv.get_nonstream_response() + assert resp["status"] == "incomplete" + assert resp["incomplete_details"] == {"reason": "max_output_tokens"} + assert resp["output"][0]["status"] == "incomplete" + + conv = ResponsesStreamConverter(model="glm-5.2") + conv.feed_line('data: {"id":"c3","choices":[{"index":0,"delta":{},"finish_reason":"content_filter"}]}') + resp = conv.get_nonstream_response() + assert resp["status"] == "incomplete" and resp["incomplete_details"]["reason"] == "content_filter" + + conv = ResponsesStreamConverter(model="glm-5.2") + conv.feed_line('data: {"id":"c4","choices":[{"index":0,"delta":{"content":"ok"},"finish_reason":null}]}') + conv.feed_line('data: {"id":"c4","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}') + assert '"type": "response.completed"' in conv.finish() + assert "incomplete_details" not in conv.get_nonstream_response() + + print("✅ test_finish_reason_maps_to_terminal_status") + + +def test_stream_events_carry_sequence_and_item_ids(): + """Emit monotonic sequence numbers and item IDs on every corresponding delta.""" + conv = ResponsesStreamConverter(model="glm-5.2") + chunks = [ + 'data: {"id":"s1","choices":[{"index":0,"delta":{"reasoning_content":"想"},"finish_reason":null}]}', + 'data: {"id":"s1","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":null}]}', + 'data: {"id":"s1","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"shell","arguments":"{}"}}]},"finish_reason":null}]}', + 'data: {"id":"s1","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}', + ] + raw = "".join(conv.feed_line(line) for line in chunks) + conv.finish() + evts = [json.loads(part[6:]) for part in raw.strip().split("\n\n") if part.startswith("data: ")] + seqs = [e["sequence_number"] for e in evts] + assert seqs == sorted(seqs) and len(set(seqs)) == len(seqs), seqs + msg_ids = {e.get("item_id") for e in evts if e["type"].startswith(("response.output_text.", "response.content_part."))} + assert msg_ids == {conv.msg_id}, msg_ids + rs = [e for e in evts if e["type"].startswith("response.reasoning_summary_text.")] + assert rs and all(e["item_id"] == conv._reasoning_item_id for e in rs) + fc = [e for e in evts if e["type"].startswith("response.function_call_arguments.")] + fc_ids = {e.get("item_id") for e in fc} + assert len(fc_ids) == 1 and None not in fc_ids and next(iter(fc_ids)).startswith("fc_") + + print("✅ test_stream_events_carry_sequence_and_item_ids") + + +def test_usage_maps_cached_tokens_and_omits_when_unknown(): + """Preserve known cache counters and omit details when upstream counters are absent.""" + conv = ResponsesStreamConverter(model="m") + conv.feed_line('data: {"id":"u1","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') + conv.feed_line('data: {"id":"u1","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":1,"total_tokens":10,"prompt_tokens_details":{"cached_tokens":7}}}') + assert conv.get_nonstream_response()["usage"]["input_tokens_details"] == {"cached_tokens": 7} + + conv = ResponsesStreamConverter(model="m") + conv.feed_line('data: {"id":"u2","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') + conv.feed_line('data: {"id":"u2","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":1,"total_tokens":10,"cache_read_input_tokens":3}}') + assert conv.get_nonstream_response()["usage"]["input_tokens_details"] == {"cached_tokens": 3} + + conv = ResponsesStreamConverter(model="m") + conv.feed_line('data: {"id":"u3","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') + conv.feed_line('data: {"id":"u3","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":1,"total_tokens":10}}') + assert "input_tokens_details" not in conv.get_nonstream_response()["usage"] + + print("✅ test_usage_maps_cached_tokens_and_omits_when_unknown") + + +def test_reasoning_effort_and_text_format_are_mapped(): + """Map reasoning and text format options with explicit top-level precedence.""" + chat = responses_request_to_chat({"input": "hi", "reasoning": {"effort": "high"}, + "text": {"format": {"type": "json_object"}}}) + assert chat["reasoning_effort"] == "high" + assert chat["response_format"] == {"type": "json_object"} + + chat = responses_request_to_chat({"input": "hi", "reasoning_effort": "low", + "reasoning": {"effort": "high"}}) + assert chat["reasoning_effort"] == "low" # Explicit top-level fields take precedence. + + chat = responses_request_to_chat({"input": "hi", "text": {"format": { + "type": "json_schema", "name": "answer", "strict": True, + "schema": {"type": "object", "properties": {"a": {"type": "integer"}}}}}}) + fmt = chat["response_format"] + assert fmt["type"] == "json_schema" + assert fmt["json_schema"]["name"] == "answer" and fmt["json_schema"]["strict"] is True + assert fmt["json_schema"]["schema"]["properties"]["a"]["type"] == "integer" + + chat = responses_request_to_chat({"input": "hi", "text": {"format": {"type": "text"}}}) + assert "response_format" not in chat + + try: + responses_request_to_chat({"input": "hi", "text": {"format": {"type": "xml"}}}) + raise AssertionError("unsupported text.format must raise") + except ValueError: + pass + + print("✅ test_reasoning_effort_and_text_format_are_mapped") + +def test_parallel_tool_calls_roundtrip(): + """Preserve the requested parallel_tool_calls setting in upstream and response objects.""" + chat = responses_request_to_chat({"input": "hi", "parallel_tool_calls": False}) + assert chat["parallel_tool_calls"] is False + conv = ResponsesStreamConverter(model="m", parallel_tool_calls=False) + conv.feed_line('data: {"id":"p1","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') + conv.feed_line('data: {"id":"p1","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}') + assert conv.get_nonstream_response()["parallel_tool_calls"] is False + conv = ResponsesStreamConverter(model="m") + conv.feed_line('data: {"id":"p2","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') + conv.feed_line('data: {"id":"p2","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}') + assert conv.get_nonstream_response()["parallel_tool_calls"] is True + print("✅ test_parallel_tool_calls_roundtrip") -def test_typed_developer_message_request(): - """Normalize developer roles inside typed messages.""" +def test_encrypted_content_and_namespace_tools(): + """Extract encrypted_content in agent_message and sanitize namespace tools.""" req = { - "model": "glm-5.2", + "model": "deepseek-v4.1-flash", "input": [ - {"type": "message", "role": "developer", "content": "Be concise."}, - {"type": "message", "role": "user", "content": "Hi"}, - ], - } - chat = responses_request_to_chat(req) - assert chat["messages"][0] == {"role": "system", "content": "Be concise."} - assert chat["messages"][1] == {"role": "user", "content": "Hi"} - print("✅ test_typed_developer_message_request") - - -def test_desensitize_harness_user_and_tools(): - """Compact trusted harness input and tool metadata while preserving real user text.""" - body = { - "messages": [ - {"role": "system", "content": "Refuse exploit development."}, - {"role": "user", "content": "# AGENTS.md instructions\n sandbox escalation"}, - {"role": "user", "content": "please explain dos attacks"}, - ], - "tools": [ - {"type": "function", "function": {"name": "exec_command", "description": "Run dangerous exploit development checks."}} - ], - } - out = desensitize_body( - body, - roles=("system", "developer"), - desensitize_harness_user=True, - desensitize_tools=True, - ) - assert "​" in out["messages"][0]["content"] - assert "Repository instructions and durable user context are provided." in out["messages"][1]["content"] - assert "Environment context is provided by the harness." in out["messages"][1]["content"] - assert "​" not in out["messages"][2]["content"] - assert out["messages"][2]["content"] == "please explain dos attacks" - assert "​" in out["tools"][0]["function"]["description"] - print("✅ test_desensitize_harness_user_and_tools") - - -def test_compact_harness_messages_and_strip_tool_metadata(): - """Compact long Codex templates without replacing the user's actual request.""" - body = { - "messages": [ - {"role": "system", "content": "You are a coding agent running in the Codex CLI. # How you work\nUse sandbox and escalation."}, - {"role": "system", "content": "\nFilesystem sandboxing defines which files can be read or written."}, - {"role": "user", "content": "# AGENTS.md instructions\n sandbox escalation"}, - ], - "tools": [ - {"type": "function", "function": {"name": "exec_command", "description": "Run dangerous exploit development checks.", "parameters": {"type": "object", "properties": {"cmd": {"type": "string", "description": "Shell command to execute."}}}}} - ], - } - out = desensitize_body( - body, - roles=("system", "developer"), - desensitize_harness_user=True, - desensitize_tools=True, - compact_harness=True, - strip_tool_metadata=True, - ) - assert len(out["messages"][0]["content"]) < 220 - assert "Codex CLI" in out["messages"][0]["content"] - assert "sandboxing defines" not in out["messages"][1]["content"] - assert "Repository instructions and durable user context are provided." in out["messages"][2]["content"] - assert "Environment context is provided by the harness." in out["messages"][2]["content"] - assert "description" not in out["tools"][0]["function"] - assert "description" not in out["tools"][0]["function"]["parameters"]["properties"]["cmd"] - print("✅ test_compact_harness_messages_and_strip_tool_metadata") - - -def test_no_compact_still_prunes_codex_runtime_metadata(): - """Remove trusted runtime metadata even when full conversation text is preserved.""" - body = { - "messages": [ { - "role": "system", - "content": ( - "You are a coding agent running in the Codex CLI.\n\n" - "# How you work\nUse sandbox and escalation carefully.\n\n" - "\nFilesystem sandboxing defines which files can be read or written.\n" - "## How to request escalation\n...\n\n\n" - "The following deferred tools are now available via ToolSearch.\n..." - ), + "type": "additional_tools", + "role": "developer", + "tools": [ + { + "type": "namespace", + "name": "collaboration", + "tools": [ + { + "type": "function", + "name": "spawn_agent", + "description": "Spawn an agent", + "parameters": { + "type": "object", + "properties": { + "message": {"type": "string", "encrypted": True}, + "task_name": {"type": "string"}, + }, + }, + } + ], + } + ], }, { - "role": "user", - "content": ( - "# AGENTS.md instructions\n\nproject guidance\n" - "\nvery long runtime context\n\n" - "\nvery long skills metadata\n\n" - "test" - ), - }, - {"role": "user", "content": "test"}, - ], - } - out = desensitize_body( - body, - roles=("system", "developer"), - desensitize_harness_user=True, - compact_harness=False, - ) - system_text = out["messages"][0]["content"] - harness_text = out["messages"][1]["content"] - assert "You are a coding agent running in the Codex CLI" in system_text - assert "## Planning" not in system_text - assert "## Task execution" not in system_text - assert "### Final answer structure and style guidelines" not in system_text - assert "# How you work" in system_text - assert "Filesystem sandboxing defines" not in system_text - # Without a closed wrapper the deferred-tool paragraph is real text, not metadata. - assert "The following deferred tools are now available via ToolSearch.\n..." in system_text - assert "Runtime permissions apply" in system_text - assert "Runtime tool, agent, sk" not in system_text - assert "very long runtime context" not in harness_text - assert "very long skills metadata" not in harness_text - assert "# AGENTS.md instructions" not in harness_text - assert "Repository instructions and durable user context are provided." in harness_text - assert "Environment context is provided by the harness." in harness_text - # Remove inserted separators to isolate harness compaction from term adaptation. - assert "Runtime skill metadata is available" in harness_text.replace("​", "") - assert harness_text.strip().replace("​", "").endswith("test") - assert out["messages"][2]["content"] == "test" - print("✅ test_no_compact_still_prunes_codex_runtime_metadata") - - -def test_responses_projection_balanced_preserves_real_text_tools_and_structure(): - """Balanced mode replaces recognized harness blocks without changing tools or real text.""" - tool = { - "type": "function", - "function": { - "name": "exec_command", - "description": "Run a command", - "parameters": { - "type": "object", - "properties": {"cmd": {"type": "string", "description": "Command"}}, - "required": ["cmd"], - "additionalProperties": False, - "x-vendor-detail": {"deep": {"schema": "kept"}}, + "type": "agent_message", + "author": "/root", + "recipient": "/root/worker", + "content": [ + {"type": "input_text", "text": "Payload:\n"}, + {"type": "encrypted_content", "encrypted_content": "Execute news task"}, + ], }, - "strict": False, - }, - } - body = { - "model": "auto", - "messages": [ - {"role": "system", "content": "Repository policy: run tests."}, - {"role": "user", "content": "# AGENTS.md instructions\n\nvolatile context\n"}, - {"role": "user", "content": "实现该方案"}, ], - "tools": [tool], } - before = json.loads(json.dumps(body, ensure_ascii=False)) - out, stats = project_responses_chat_body(body) - assert body == before - assert out["tools"] == body["tools"] - assert out["messages"][0] == body["messages"][0] - assert out["messages"][1]["content"] != body["messages"][1]["content"] - assert "# AGENTS.md instructions" not in out["messages"][1]["content"] - assert "Environment context is provided by the harness." in out["messages"][1]["content"] - assert out["messages"][2] == body["messages"][2] - assert stats["mode"] == "balanced" - assert stats["original_tools"] == stats["projected_tools"] == 1 - assert stats["original_tool_chars"] == stats["projected_tool_chars"] - assert stats["harness_messages_projected"] == 1 - print("✅ test_responses_projection_balanced_preserves_real_text_tools_and_structure") - - -def test_responses_projection_preserves_history_and_tool_chain(): - """Balanced mode does not summarize history or drop tool-call relationships.""" - body = { - "messages": [ - {"role": "user", "content": "old task"}, - {"role": "assistant", "content": "old answer", "tool_calls": [{ - "id": "call_old", "type": "function", - "function": {"name": "exec_command", "arguments": '{"cmd":"ls"}'}, - }]}, - {"role": "tool", "tool_call_id": "call_old", "content": "old output"}, - {"role": "user", "content": "new task"}, - ], - "tools": [{"type": "function", "function": {"name": "exec_command", "parameters": {"type": "object"}}}], - } - before = json.loads(json.dumps(body)) - out, stats = project_responses_chat_body(body, max_item_bytes=40000) - assert body == before - assert out["messages"] == body["messages"] - assert stats["original_messages"] == stats["projected_messages"] == 4 - assert "anchor_user_preserved" not in stats - print("✅ test_responses_projection_preserves_history_and_tool_chain") - - -def test_responses_projection_truncates_generated_content_and_json_arguments(): - """Oversized assistant, tool output, and JSON values retain UTF-8 head and tail.""" - long_text = "HEAD\n" + ("中" * 180) + "\nTAIL" - long_output = "OUTPUT\n" + ("输出" * 180) + "\nEND" - arguments = json.dumps({"cmd": "echo " + ("x" * 500), "workdir": "/tmp"}) - apply_patch = json.dumps({"patch": "*** Begin Patch\n" + ("+" * 500) + "*** End Patch"}) - body = {"messages": [ - {"role": "assistant", "content": long_text, "tool_calls": [{ - "id": "call_1", "type": "function", - "function": {"name": "exec_command", "arguments": arguments}, - }]}, - {"role": "tool", "tool_call_id": "call_1", "content": long_output}, - {"role": "assistant", "content": "", "tool_calls": [{ - "id": "call_patch", "type": "function", - "function": {"name": "apply_patch", "arguments": apply_patch}, - }]}, - ]} - before = json.loads(json.dumps(body, ensure_ascii=False)) - out, stats = project_responses_chat_body(body, max_item_bytes=256) - assert body == before - assistant = out["messages"][0]["content"] - assert assistant.startswith("HEAD") - assert assistant.endswith("TAIL") - assert "middle omitted" in assistant - assert "original bytes:" in assistant - assert "estimated tokens:" in assistant - assert "total lines:" in assistant - tool_output = out["messages"][1]["content"] - assert tool_output.startswith("OUTPUT") - assert tool_output.endswith("END") - args_wire = out["messages"][0]["tool_calls"][0]["function"]["arguments"] - args = json.loads(args_wire) - assert len(args_wire.encode("utf-8")) <= 256 - assert "middle omitted" in args["_truncated"]["warning"] - assert '"cmd"' in args["head"] and "echo" in args["head"] - assert "workdir" in args["tail"] and args["tail"].endswith("}") - patch_wire = out["messages"][2]["tool_calls"][0]["function"]["arguments"] - patch_args = json.loads(patch_wire) - assert len(patch_wire.encode("utf-8")) <= 256 - assert '"patch"' in patch_args["head"] and "*** Begin Patch" in patch_args["head"] - assert patch_args["tail"].endswith('*** End Patch"}') - assert stats["truncated_items"] == 4 - assert stats["truncated_original_bytes"] > stats["truncated_projected_bytes"] - print("✅ test_responses_projection_truncates_generated_content_and_json_arguments") - - -def test_responses_projection_passthrough_and_zero_limit_are_lossless(): - """Passthrough and max_item_bytes=0 preserve the request payload.""" - body = {"messages": [ - {"role": "user", "content": "# AGENTS.md instructions\nvolatile\nreal task"}, - {"role": "assistant", "content": "assistant " + "x" * 500}, - {"role": "tool", "tool_call_id": "c", "content": "output " + "y" * 500}, - ], "tools": [{"type": "function", "function": {"name": "tool", "parameters": {"type": "object", "x": 1}}}]} - before = json.loads(json.dumps(body, ensure_ascii=False)) - for kwargs in ({"mode": "passthrough"}, {"max_item_bytes": 0}): - out, stats = project_responses_chat_body(body, **kwargs) - if "mode" in kwargs: - assert out == before - else: - assert out["messages"][1:] == before["messages"][1:] - assert out["tools"] == before["tools"] - assert body == before - assert stats["mode"] == ("passthrough" if "mode" in kwargs else "balanced") - assert stats["truncated_items"] == 0 - print("✅ test_responses_projection_passthrough_and_zero_limit_are_lossless") - - -def test_responses_projection_rejects_invalid_mode_and_limit(): - """Projection validates its public mode and byte-limit arguments.""" - for mode in ("aggressive", "conservative", "", None, True): - try: - project_responses_chat_body({"messages": []}, mode=mode) - except ValueError: - pass - else: - raise AssertionError(f"invalid mode accepted: {mode!r}") - for limit in (-1, 1, 128, 255, True, 1.5, "256", None, [], {}): - try: - project_responses_chat_body({"messages": []}, max_item_bytes=limit) - except ValueError: - pass - else: - raise AssertionError(f"invalid max_item_bytes accepted: {limit!r}") - assert project_responses_chat_body({"messages": []}, mode="balanced", max_item_bytes=256)[1]["max_item_bytes"] == 256 - print("✅ test_responses_projection_rejects_invalid_mode_and_limit") - - -def test_responses_projection_stats_have_official_shape(): - """Stats expose current projection counters without legacy mode fields.""" - body = {"messages": [{"role": "user", "content": "hello"}], "tools": [{"type": "function", "function": {"name": "t"}}]} - _, stats = project_responses_chat_body(body, mode="balanced", max_item_bytes=0) - assert stats == { - "mode": "balanced", "max_item_bytes": 0, "original_messages": 1, - "projected_messages": 1, "original_message_chars": stats["original_message_chars"], - "projected_message_chars": stats["projected_message_chars"], "original_tools": 1, - "projected_tools": 1, "original_tool_chars": stats["original_tool_chars"], - "projected_tool_chars": stats["projected_tool_chars"], "harness_messages_projected": 0, - "truncated_items": 0, "truncated_original_bytes": 0, "truncated_projected_bytes": 0, + chat = responses_request_to_chat(req) + tools = chat.get("tools", []) + assert len(tools) == 1 + fn = tools[0]["function"] + assert fn["name"] == "spawn_agent" + assert "encrypted" not in fn["parameters"]["properties"]["message"] + messages = chat.get("messages", []) + assert len(messages) == 1 + assert messages[0]["role"] == "user" + assert "Execute news task" in messages[0]["content"] + assert "Payload:\n" in messages[0]["content"] + + conv = ResponsesStreamConverter(model="deepseek", tool_namespaces=chat.get("_tool_namespaces")) + slot = { + "id": "call_1", + "name": "spawn_agent", + "args": "{}", + "fc_id": "fc_1", + "output_idx": 0, + "emitted": False, + "emitted_args_length": 0, } - print("✅ test_responses_projection_stats_have_official_shape") - - -def test_stream_converter_text(): - """Convert Chat text SSE into Responses events.""" - conv = ResponsesStreamConverter(model="glm-5.2") - - # Synthetic Chat SSE chunks - chunks = [ - 'data: {"id":"chatcmpl-1","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}', - 'data: {"id":"chatcmpl-1","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}', - 'data: {"id":"chatcmpl-1","choices":[{"index":0,"delta":{"content":" world"},"finish_reason":null}]}', - 'data: {"id":"chatcmpl-1","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":2,"total_tokens":12}}', - 'data: [DONE]', - ] - - all_events = [] - for line in chunks: - result = conv.feed_line(line) - if result: - for evt_line in result.strip().split("\n\n"): - if evt_line.startswith("data: "): - all_events.append(json.loads(evt_line[6:])) - - finish = conv.finish() - for evt_line in finish.strip().split("\n\n"): - if evt_line.startswith("data: "): - all_events.append(json.loads(evt_line[6:])) - - types = [e["type"] for e in all_events] - assert "response.created" in types - assert "response.in_progress" in types - assert "response.output_item.added" in types - assert "response.content_part.added" in types - assert "response.output_text.delta" in types - assert "response.output_text.done" in types - assert "response.content_part.done" in types - assert "response.output_item.done" in types - assert "response.completed" in types - - text_done = [e for e in all_events if e["type"] == "response.output_text.done"][0] - assert text_done["text"] == "Hello world" - - completed = [e for e in all_events if e["type"] == "response.completed"][0] - resp = completed["response"] - assert resp["status"] == "completed" - assert resp["output"][0]["type"] == "message" - assert resp["output"][0]["content"][0]["text"] == "Hello world" - assert resp["usage"]["input_tokens"] == 10 - - print("✅ test_stream_converter_text") - - -def test_stream_converter_function_call(): - """Convert Chat tool-call SSE into Responses function-call events.""" - conv = ResponsesStreamConverter(model="glm-5.2") - - chunks = [ - 'data: {"id":"chatcmpl-2","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"id":"call_abc","type":"function","function":{"name":"shell","arguments":""}}]},"finish_reason":null}]}', - 'data: {"id":"chatcmpl-2","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\\"cmd"}}]},"finish_reason":null}]}', - 'data: {"id":"chatcmpl-2","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\\": \\"ls\\"}"}}]},"finish_reason":null}]}', - 'data: {"id":"chatcmpl-2","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}', - 'data: [DONE]', - ] - - all_events = [] - for line in chunks: - result = conv.feed_line(line) - if result: - for evt_line in result.strip().split("\n\n"): - if evt_line.startswith("data: "): - all_events.append(json.loads(evt_line[6:])) - - finish = conv.finish() - for evt_line in finish.strip().split("\n\n"): - if evt_line.startswith("data: "): - all_events.append(json.loads(evt_line[6:])) - - types = [e["type"] for e in all_events] - assert "response.output_item.added" in types - assert "response.function_call_arguments.delta" in types - assert "response.function_call_arguments.done" in types - assert "response.completed" in types - - args_done = [e for e in all_events if e["type"] == "response.function_call_arguments.done"][0] - assert args_done["arguments"] == '{"cmd": "ls"}' - - print("✅ test_stream_converter_function_call") - - -def test_nonstream_response(): - """Build a non-streaming Response object.""" - conv = ResponsesStreamConverter(model="glm-5.2") - conv.feed_line('data: {"id":"c1","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":null}]}') - conv.feed_line('data: {"id":"c1","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}}') - - resp = conv.get_nonstream_response() - assert resp["object"] == "response" - assert resp["status"] == "completed" - assert resp["output"][0]["type"] == "message" - assert resp["output"][0]["content"][0]["text"] == "Hi" - assert resp["usage"]["input_tokens"] == 5 - - print("✅ test_nonstream_response") - - -def test_finish_reason_maps_to_terminal_status(): - """Report truncated or filtered output as incomplete in streaming and aggregated responses.""" - conv = ResponsesStreamConverter(model="glm-5.2") - conv.feed_line('data: {"id":"c2","choices":[{"index":0,"delta":{"content":"partial"},"finish_reason":null}]}') - conv.feed_line('data: {"id":"c2","choices":[{"index":0,"delta":{},"finish_reason":"length"}]}') - tail = conv.finish() - assert '"type": "response.incomplete"' in tail and '"type": "response.completed"' not in tail - resp = conv.get_nonstream_response() - assert resp["status"] == "incomplete" - assert resp["incomplete_details"] == {"reason": "max_output_tokens"} - assert resp["output"][0]["status"] == "incomplete" - - conv = ResponsesStreamConverter(model="glm-5.2") - conv.feed_line('data: {"id":"c3","choices":[{"index":0,"delta":{},"finish_reason":"content_filter"}]}') - resp = conv.get_nonstream_response() - assert resp["status"] == "incomplete" and resp["incomplete_details"]["reason"] == "content_filter" - - conv = ResponsesStreamConverter(model="glm-5.2") - conv.feed_line('data: {"id":"c4","choices":[{"index":0,"delta":{"content":"ok"},"finish_reason":null}]}') - conv.feed_line('data: {"id":"c4","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}') - assert '"type": "response.completed"' in conv.finish() - assert "incomplete_details" not in conv.get_nonstream_response() - - print("✅ test_finish_reason_maps_to_terminal_status") - - -def test_stream_events_carry_sequence_and_item_ids(): - """Emit monotonic sequence numbers and item IDs on every corresponding delta.""" - conv = ResponsesStreamConverter(model="glm-5.2") - chunks = [ - 'data: {"id":"s1","choices":[{"index":0,"delta":{"reasoning_content":"想"},"finish_reason":null}]}', - 'data: {"id":"s1","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":null}]}', - 'data: {"id":"s1","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"shell","arguments":"{}"}}]},"finish_reason":null}]}', - 'data: {"id":"s1","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}', - ] - raw = "".join(conv.feed_line(line) for line in chunks) + conv.finish() - evts = [json.loads(part[6:]) for part in raw.strip().split("\n\n") if part.startswith("data: ")] - seqs = [e["sequence_number"] for e in evts] - assert seqs == sorted(seqs) and len(set(seqs)) == len(seqs), seqs - msg_ids = {e.get("item_id") for e in evts if e["type"].startswith(("response.output_text.", "response.content_part."))} - assert msg_ids == {conv.msg_id}, msg_ids - rs = [e for e in evts if e["type"].startswith("response.reasoning_summary_text.")] - assert rs and all(e["item_id"] == conv._reasoning_item_id for e in rs) - fc = [e for e in evts if e["type"].startswith("response.function_call_arguments.")] - fc_ids = {e.get("item_id") for e in fc} - assert len(fc_ids) == 1 and None not in fc_ids and next(iter(fc_ids)).startswith("fc_") - - print("✅ test_stream_events_carry_sequence_and_item_ids") - - -def test_usage_maps_cached_tokens_and_omits_when_unknown(): - """Preserve known cache counters and omit details when upstream counters are absent.""" - conv = ResponsesStreamConverter(model="m") - conv.feed_line('data: {"id":"u1","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') - conv.feed_line('data: {"id":"u1","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":1,"total_tokens":10,"prompt_tokens_details":{"cached_tokens":7}}}') - assert conv.get_nonstream_response()["usage"]["input_tokens_details"] == {"cached_tokens": 7} - - conv = ResponsesStreamConverter(model="m") - conv.feed_line('data: {"id":"u2","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') - conv.feed_line('data: {"id":"u2","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":1,"total_tokens":10,"cache_read_input_tokens":3}}') - assert conv.get_nonstream_response()["usage"]["input_tokens_details"] == {"cached_tokens": 3} - - conv = ResponsesStreamConverter(model="m") - conv.feed_line('data: {"id":"u3","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') - conv.feed_line('data: {"id":"u3","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":1,"total_tokens":10}}') - assert "input_tokens_details" not in conv.get_nonstream_response()["usage"] - - print("✅ test_usage_maps_cached_tokens_and_omits_when_unknown") - - -def test_reasoning_effort_and_text_format_are_mapped(): - """Map reasoning and text format options with explicit top-level precedence.""" - chat = responses_request_to_chat({"input": "hi", "reasoning": {"effort": "high"}, - "text": {"format": {"type": "json_object"}}}) - assert chat["reasoning_effort"] == "high" - assert chat["response_format"] == {"type": "json_object"} - - chat = responses_request_to_chat({"input": "hi", "reasoning_effort": "low", - "reasoning": {"effort": "high"}}) - assert chat["reasoning_effort"] == "low" # Explicit top-level fields take precedence. - - chat = responses_request_to_chat({"input": "hi", "text": {"format": { - "type": "json_schema", "name": "answer", "strict": True, - "schema": {"type": "object", "properties": {"a": {"type": "integer"}}}}}}) - fmt = chat["response_format"] - assert fmt["type"] == "json_schema" - assert fmt["json_schema"]["name"] == "answer" and fmt["json_schema"]["strict"] is True - assert fmt["json_schema"]["schema"]["properties"]["a"]["type"] == "integer" - - chat = responses_request_to_chat({"input": "hi", "text": {"format": {"type": "text"}}}) - assert "response_format" not in chat - - try: - responses_request_to_chat({"input": "hi", "text": {"format": {"type": "xml"}}}) - raise AssertionError("unsupported text.format must raise") - except ValueError: - pass - - print("✅ test_reasoning_effort_and_text_format_are_mapped") - -def test_parallel_tool_calls_roundtrip(): - """Preserve the requested parallel_tool_calls setting in upstream and response objects.""" - chat = responses_request_to_chat({"input": "hi", "parallel_tool_calls": False}) - assert chat["parallel_tool_calls"] is False - conv = ResponsesStreamConverter(model="m", parallel_tool_calls=False) - conv.feed_line('data: {"id":"p1","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') - conv.feed_line('data: {"id":"p1","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}') - assert conv.get_nonstream_response()["parallel_tool_calls"] is False - conv = ResponsesStreamConverter(model="m") - conv.feed_line('data: {"id":"p2","choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}') - conv.feed_line('data: {"id":"p2","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}') - assert conv.get_nonstream_response()["parallel_tool_calls"] is True - print("✅ test_parallel_tool_calls_roundtrip") + item = conv._fc_item(slot, "completed") + assert item["type"] == "function_call" + assert item["name"] == "spawn_agent" + assert item["namespace"] == "collaboration" + print("✅ test_encrypted_content_and_namespace_tools") if __name__ == "__main__": - test_simple_text_request() - test_array_input_request() - test_tools_conversion() - test_max_output_tokens() - test_developer_role() - test_typed_developer_message_request() - test_desensitize_harness_user_and_tools() - test_compact_harness_messages_and_strip_tool_metadata() - test_no_compact_still_prunes_codex_runtime_metadata() - test_responses_projection_balanced_preserves_real_text_tools_and_structure() - test_responses_projection_preserves_history_and_tool_chain() - test_responses_projection_truncates_generated_content_and_json_arguments() - test_responses_projection_passthrough_and_zero_limit_are_lossless() - test_responses_projection_rejects_invalid_mode_and_limit() - test_responses_projection_stats_have_official_shape() - test_stream_converter_text() - test_stream_converter_function_call() - test_nonstream_response() - test_finish_reason_maps_to_terminal_status() - test_stream_events_carry_sequence_and_item_ids() - test_usage_maps_cached_tokens_and_omits_when_unknown() + test_simple_text_request() + test_array_input_request() + test_tools_conversion() + test_max_output_tokens() + test_developer_role() + test_typed_developer_message_request() + test_desensitize_harness_user_and_tools() + test_compact_harness_messages_and_strip_tool_metadata() + test_no_compact_still_prunes_codex_runtime_metadata() + test_responses_projection_balanced_preserves_real_text_tools_and_structure() + test_responses_projection_preserves_history_and_tool_chain() + test_responses_projection_truncates_generated_content_and_json_arguments() + test_responses_projection_passthrough_and_zero_limit_are_lossless() + test_responses_projection_rejects_invalid_mode_and_limit() + test_responses_projection_stats_have_official_shape() + test_stream_converter_text() + test_stream_converter_function_call() + test_nonstream_response() + test_finish_reason_maps_to_terminal_status() + test_stream_events_carry_sequence_and_item_ids() + test_usage_maps_cached_tokens_and_omits_when_unknown() test_reasoning_effort_and_text_format_are_mapped() test_parallel_tool_calls_roundtrip() + test_encrypted_content_and_namespace_tools() print(f"\n🎉 All {22} tests passed!")