diff --git a/app/adapters/anthropic_adapter.py b/app/adapters/anthropic_adapter.py index d33d65a..e65a1d4 100644 --- a/app/adapters/anthropic_adapter.py +++ b/app/adapters/anthropic_adapter.py @@ -7,6 +7,7 @@ import time from typing import Any +from app.reasoning import extract_reasoning_text, map_reasoning_controls from app.upstream_io import (StreamOutputBudget, merge_tool_call_delta, new_tool_state, seal_tool_identities, tool_identity_complete) @@ -93,6 +94,7 @@ def anthropic_request_to_chat(body: dict) -> dict: raise ValueError("stop_sequences must be an array of strings") chat["stop"] = stop_sequences + map_reasoning_controls(body, chat, protocol="messages") return chat @@ -124,6 +126,9 @@ def _convert_anthropic_message(msg: dict) -> list[dict]: # Structured content blocks blocks = content + if role != "assistant" and any(isinstance(block, dict) and block.get("type") in + ("thinking", "redacted_thinking") for block in blocks): + raise ValueError("thinking requires an assistant message") # User messages may contain tool results. if role == "user": @@ -157,11 +162,14 @@ def _convert_anthropic_message(msg: dict) -> list[dict]: if role == "assistant": content_out = _convert_content_blocks(blocks) tool_calls: list[dict] = [] + thoughts: list[str] = [] for block in blocks: if not isinstance(block, dict): continue - bt = block.get("type", "") - if bt == "tool_use": + thought = extract_reasoning_text(block) + if thought is not None: + thoughts.append(thought) + elif block.get("type") == "tool_use": tc = { "id": block.get("id", _rand_id("call_")), "type": "function", @@ -176,6 +184,8 @@ def _convert_anthropic_message(msg: dict) -> list[dict]: msg_out["content"] = content_out if content_out or has_text else None if tool_calls: msg_out["tool_calls"] = tool_calls + if thoughts: + msg_out["reasoning_content"] = "".join(thoughts) return [msg_out] content_out = _convert_content_blocks(blocks) diff --git a/app/adapters/chat_input.py b/app/adapters/chat_input.py index c40bda4..f452b1b 100644 --- a/app/adapters/chat_input.py +++ b/app/adapters/chat_input.py @@ -3,6 +3,8 @@ from fastapi import HTTPException +from app.reasoning import ReasoningInputError, extract_reasoning_text + _ANTHROPIC_BLOCKS = ("tool_use", "tool_result", "thinking", "redacted_thinking", "image") @@ -116,17 +118,15 @@ def _convert_message(message, index, pending): raise _invalid(location + ".tool_use_id", "tool_result must match one preceding, unanswered tool call") result_ids.add(identifier) results.append(result) - elif kind == "thinking": + elif kind in ("thinking", "redacted_thinking"): if role != "assistant": raise _invalid(location, "thinking requires an assistant message") if message.get("reasoning_content") not in (None, ""): raise _invalid(location, "thinking conflicts with existing reasoning_content") - if not isinstance(block.get("thinking"), str): - raise _invalid(location + ".thinking", "Thinking content must be a string") - # Anthropic signatures have no Chat equivalent and must not become visible text. - thoughts.append(block["thinking"]) - elif kind == "redacted_thinking": - raise _invalid(location, "redacted_thinking cannot be converted to Chat; use the Messages protocol") + try: + thoughts.append(extract_reasoning_text(block)) + except ReasoningInputError as error: + raise _invalid(location + ("." + error.field if error.field else ""), str(error)) from None else: parts.append(_content_part(block, location)) if results: diff --git a/app/adapters/responses_adapter.py b/app/adapters/responses_adapter.py index 6dbd900..517fa50 100644 --- a/app/adapters/responses_adapter.py +++ b/app/adapters/responses_adapter.py @@ -2,11 +2,15 @@ from __future__ import annotations +from copy import deepcopy +import hashlib import json import os +import re import time from typing import Any +from app.reasoning import extract_reasoning_text, map_reasoning_controls from app.upstream_io import (StreamOutputBudget, merge_tool_call_delta, new_tool_state, seal_tool_identities, tool_identity_complete) @@ -47,6 +51,19 @@ def _text_format_to_response_format(fmt) -> dict | None: def responses_request_to_chat(body: dict) -> dict: """Convert Responses input, instructions and tools to a Chat request.""" + # Collect declared tools from both top-level tools and input additional_tools items. + raw_tools = [] + if isinstance(body.get("tools"), list): + raw_tools.extend(body["tools"]) + inp = body.get("input", []) + if isinstance(inp, list): + for item in inp: + if isinstance(item, dict) and item.get("type") == "additional_tools" and isinstance(item.get("tools"), list): + raw_tools.extend(item["tools"]) + + registry, chat_tools = _build_tool_registry_and_chat_tools( + raw_tools, input_items=inp if isinstance(inp, list) else None) + messages: list[dict] = [] # instructions → system message @@ -55,11 +72,10 @@ def responses_request_to_chat(body: dict) -> dict: 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)) + messages.extend(_convert_input_items(inp, tool_registry=registry)) # Build the Chat request body. chat: dict[str, Any] = {"messages": messages, "stream": True} @@ -68,26 +84,22 @@ def responses_request_to_chat(body: dict) -> dict: if "model" in body: chat["model"] = body["model"] - # Normalize function tool definitions. - tools = body.get("tools") - if tools: - chat["tools"] = _convert_tools_for_chat(tools) + if chat_tools: + chat["tools"] = chat_tools + if registry.upstream_to_identity: + chat["_tool_registry"] = registry.to_dict() if "tool_choice" in body: - chat["tool_choice"] = body["tool_choice"] + chat["tool_choice"] = _convert_tool_choice_for_chat(body["tool_choice"], tool_registry=registry) # 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"): + "response_format", "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 + map_reasoning_controls(body, chat, protocol="responses") text = body.get("text") if isinstance(text, dict) and "response_format" not in chat: mapped = _text_format_to_response_format(text.get("format")) @@ -102,23 +114,42 @@ def responses_request_to_chat(body: dict) -> dict: return chat -def _convert_input_items(items: list) -> list[dict]: +def _convert_input_items(items: list, tool_registry: ToolRegistry | None = None) -> 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] = [] + pending_reasoning: list[str] = [] def _flush_assistant(): nonlocal pending_assistant_content, pending_tool_calls - if pending_assistant_content is not None or pending_tool_calls: + if pending_assistant_content is not None or pending_tool_calls or pending_reasoning: msg: dict[str, Any] = {"role": "assistant", "content": pending_assistant_content or ""} if pending_tool_calls: msg["tool_calls"] = pending_tool_calls[:] + if pending_reasoning: + msg["reasoning_content"] = "".join(pending_reasoning) messages.append(msg) pending_assistant_content = None pending_tool_calls.clear() + pending_reasoning.clear() + + def _set_assistant_content(content): + nonlocal pending_assistant_content + if pending_assistant_content is not None and not pending_tool_calls: + _flush_assistant() + # Realtime output can place text after a tool call within the same turn. + if pending_tool_calls and pending_assistant_content: + if isinstance(pending_assistant_content, str) and isinstance(content, str): + content = pending_assistant_content + content + else: + previous = (pending_assistant_content if isinstance(pending_assistant_content, list) + else [{"type": "text", "text": pending_assistant_content}]) + following = content if isinstance(content, list) else [{"type": "text", "text": content}] + content = previous + following + pending_assistant_content = content for item in items: if not isinstance(item, dict): @@ -127,6 +158,24 @@ def _flush_assistant(): item_type = item.get("type") role = item.get("role", "") + # Ignore additional_tools in message sequence. + if item_type == "additional_tools": + continue + + # Agent message from multi-agent collaboration. + if item_type == "agent_message": + _flush_assistant() + content = _extract_content(item.get("content", "")) + messages.append({"role": "user", "content": content}) + continue + + if item_type == "reasoning": + if role not in ("", "assistant"): + raise ValueError("reasoning requires an assistant item") + text = extract_reasoning_text(item) + pending_reasoning.append(text) + continue + # Untyped role messages if item_type is None and role in ("user", "system", "developer"): _flush_assistant() @@ -145,17 +194,15 @@ def _flush_assistant(): # 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 + _set_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 + _set_assistant_content(content) continue # Merge calls into the preceding assistant message. @@ -165,11 +212,14 @@ def _flush_assistant(): raise ValueError("function_call.arguments must be a JSON string") if pending_assistant_content is None: pending_assistant_content = "" + raw_name = item.get("name", "") + ns = item.get("namespace") + mapped_name = tool_registry.get_upstream_name(ns, raw_name) if tool_registry else raw_name pending_tool_calls.append({ "id": item.get("call_id", item.get("id", _rand_id("call_"))), "type": "function", "function": { - "name": item.get("name", ""), + "name": mapped_name, "arguments": arguments, }, }) @@ -208,6 +258,10 @@ def _extract_content(content) -> str | list[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") @@ -242,28 +296,210 @@ def _extract_output_text(content_parts: list) -> str | list[dict]: return "".join(texts) -def _convert_tools_for_chat(tools: list) -> list: - """Convert Responses tool definitions to Chat function objects.""" - result = [] +_SCHEMA_MAPS = frozenset(("properties", "patternProperties", "$defs", "definitions", + "dependentSchemas", "dependencies")) +_SCHEMA_ARRAYS = frozenset(("allOf", "anyOf", "oneOf", "prefixItems", "items")) +_SCHEMA_CHILDREN = frozenset(("items", "additionalItems", "contains", "unevaluatedItems", + "additionalProperties", "unevaluatedProperties", "propertyNames", + "not", "if", "then", "else", "contentSchema")) + + +def _sanitize_schema(schema: Any) -> Any: + """Strip encryption markers in schema positions, preserving names and instance data.""" + if not isinstance(schema, dict): + return deepcopy(schema) + result = {} + for key, value in schema.items(): + if key == "encrypted" and isinstance(value, bool): + continue + if key in _SCHEMA_MAPS and isinstance(value, dict): + result[key] = {name: _sanitize_schema(child) for name, child in value.items()} + elif key in _SCHEMA_ARRAYS and isinstance(value, list): + result[key] = [_sanitize_schema(child) for child in value] + elif key in _SCHEMA_CHILDREN: + result[key] = _sanitize_schema(value) + else: + result[key] = deepcopy(value) + return result + + +_CHAT_FUNCTION_NAME_LIMIT = 64 + + +class ToolRegistry: + """Bidirectional mapping between Responses (namespace, name) and upstream Chat function names.""" + + def __init__(self, mappings: dict[str, dict[str, Any]] | None = None): + self.upstream_to_identity: dict[str, tuple[str | None, str]] = {} + self.identity_to_upstream: dict[tuple[str | None, str], str] = {} + + if mappings: + for up_name, ident in mappings.items(): + ns = ident.get("namespace") if isinstance(ident, dict) else None + nm = ident.get("name", up_name) if isinstance(ident, dict) else up_name + self._record(up_name, ns, nm) + + @staticmethod + def _identity(namespace, name) -> tuple[str | None, str]: + if not isinstance(name, str) or not name.strip(): + raise ValueError("function name must be a non-empty string") + if namespace is not None and (not isinstance(namespace, str) or not namespace.strip()): + raise ValueError("namespace must be a non-empty string") + return namespace, name + + def _record(self, upstream_name: str, namespace: str | None, name: str) -> None: + ident = self._identity(namespace, name) + if upstream_name in self.upstream_to_identity and self.upstream_to_identity[upstream_name] != ident: + raise ValueError("tool identities must have distinct upstream names") + self.upstream_to_identity[upstream_name] = ident + self.identity_to_upstream[ident] = upstream_name + + def register(self, namespace: str | None, name: str) -> str: + """Register a tool identity (namespace, name) and return its unique upstream name.""" + ident = self._identity(namespace, name) + if ident in self.identity_to_upstream: + return self.identity_to_upstream[ident] + + if namespace is None: + upstream_name = name + else: + base = f"{namespace}__{name}" + if len(base) > _CHAT_FUNCTION_NAME_LIMIT or re.fullmatch(r"[A-Za-z0-9_-]+", base) is None: + # Different identity pairs can share the same concatenated name. + identity = json.dumps(ident, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + digest = hashlib.sha256(identity).hexdigest()[:16] + prefix = re.sub(r"[^A-Za-z0-9_-]", "_", base[:_CHAT_FUNCTION_NAME_LIMIT - len(digest) - 1]) + base = f"{prefix}_{digest}" + candidate = base + counter = 1 + while candidate in self.upstream_to_identity: + suffix = f"_{counter}" + candidate = f"{base[:_CHAT_FUNCTION_NAME_LIMIT - len(suffix)]}{suffix}" + counter += 1 + upstream_name = candidate + + self._record(upstream_name, namespace, name) + return upstream_name + + def get_identity(self, upstream_name: str) -> tuple[str | None, str]: + """Map upstream function name back to (namespace, name). Never guess by prefix.""" + if upstream_name in self.upstream_to_identity: + return self.upstream_to_identity[upstream_name] + return None, upstream_name + + def get_upstream_name(self, namespace: str | None, name: str) -> str: + """Resolve an exact identity already reserved from declarations or history.""" + ident = self._identity(namespace, name) + if ident not in self.identity_to_upstream: + raise ValueError("unknown tool identity in this request") + return self.identity_to_upstream[ident] + + def to_dict(self) -> dict: + return { + up_name: {"namespace": ident[0], "name": ident[1]} + for up_name, ident in self.upstream_to_identity.items() + } + + @classmethod + def from_dict(cls, data: dict | None) -> ToolRegistry: + if isinstance(data, cls): + return data + if not isinstance(data, dict): + return cls() + return cls(data) + + +def _collect_raw_tools(tools: list, current_ns: str = "") -> list[tuple[str | None, str, dict]]: + """Recursively collect (namespace, name, tool_def) from Responses tools.""" + collected = [] for t in tools: if not isinstance(t, dict): continue - if t.get("type") != "function": - continue - # Already in Chat format. - if "function" in t: - result.append(t) + tool_type = t.get("type") + if tool_type == "namespace" and isinstance(t.get("tools"), list): + ns_name = t.get("name", "") + if not isinstance(ns_name, str) or not ns_name.strip(): + raise ValueError("namespace.name must be a non-empty string") + full_ns = f"{current_ns}.{ns_name}" if current_ns else ns_name + collected.extend(_collect_raw_tools(t["tools"], full_ns)) + elif tool_type == "function" or "function" in t: + fn = t.get("function") if isinstance(t.get("function"), dict) else t + name = fn.get("name") or t.get("name") + if isinstance(name, str) and name: + collected.append((current_ns or None, name, t)) + return collected + + +def _format_chat_tool(upstream_name: str, tool_def: dict) -> dict: + """Format tool definition for Chat Completions, ensuring sanitized schema.""" + if "function" in tool_def: + tool_obj = json.loads(json.dumps(tool_def)) + fn_dict = tool_obj.get("function") + if isinstance(fn_dict, dict): + fn_dict["name"] = upstream_name + if "parameters" in fn_dict: + fn_dict["parameters"] = _sanitize_schema(fn_dict["parameters"]) + return tool_obj + + fn: dict[str, Any] = {"name": upstream_name} + if "description" in tool_def: + fn["description"] = tool_def["description"] + if "parameters" in tool_def: + fn["parameters"] = _sanitize_schema(tool_def["parameters"]) + if "strict" in tool_def: + fn["strict"] = tool_def["strict"] + return {"type": "function", "function": fn} + + +def _build_tool_registry_and_chat_tools(raw_tools: list, *, input_items: list | None = None) -> tuple[ToolRegistry, list[dict]]: + """Reserve current and historical identities, advertising only currently declared tools.""" + collected = _collect_raw_tools(raw_tools) + registry = ToolRegistry() + result = [] + + identities = [(ns, name) for ns, name, _ in collected] + identities.extend((item.get("namespace"), item.get("name")) for item in (input_items or []) + if isinstance(item, dict) and item.get("type") == "function_call") + # Historical global names must also remain available before allocating namespace aliases. + for ns, name in identities: + if ns is None: + registry.register(ns, name) + for ns, name in identities: + if ns is not None: + registry.register(ns, name) + + plain = [item for item in collected if item[0] is None] + namespaced = [item for item in collected if item[0] is not None] + + seen_identities: set[tuple[str | None, str]] = set() + for ns, name, tool_def in plain + namespaced: + ident = (ns, name) + if ident in seen_identities: continue - # Nest the flat Responses function fields. - fn: dict[str, Any] = {"name": t.get("name", "")} - if "description" in t: - fn["description"] = t["description"] - if "parameters" in t: - fn["parameters"] = t["parameters"] - if "strict" in t: - fn["strict"] = t["strict"] - result.append({"type": "function", "function": fn}) - return result + seen_identities.add(ident) + upstream_name = registry.get_upstream_name(ns, name) + result.append(_format_chat_tool(upstream_name, tool_def)) + + return registry, result + + +def _convert_tools_for_chat(tools: list) -> list: + """Backwards-compatible helper for converting Responses tools.""" + _, chat_tools = _build_tool_registry_and_chat_tools(tools) + return chat_tools + + +def _convert_tool_choice_for_chat(choice: Any, tool_registry: ToolRegistry | None = None) -> Any: + """Convert Responses tool_choice to Chat Completions format.""" + if not isinstance(choice, dict) or choice.get("type") != "function": + return choice + fn_obj = choice.get("function") if isinstance(choice.get("function"), dict) else choice + name = fn_obj.get("name") if isinstance(fn_obj, dict) else None + if not isinstance(name, str) or not name.strip(): + return choice + namespace = fn_obj.get("namespace", choice.get("namespace")) if isinstance(fn_obj, dict) else choice.get("namespace") + mapped_name = tool_registry.get_upstream_name(namespace, name) if tool_registry else name + return {"type": "function", "function": {"name": mapped_name}} # --------------------------------------------------------------------------- @@ -275,7 +511,9 @@ class ResponsesStreamConverter: 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_registry: ToolRegistry | dict | None = None, + tool_namespaces: dict[str, str] | None = None): self.resp_id = _rand_id("resp_") self.msg_id = _rand_id("msg_") self.model = model @@ -284,6 +522,13 @@ 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] = {} + if isinstance(tool_registry, ToolRegistry): + self._tool_registry = tool_registry + elif isinstance(tool_registry, dict): + self._tool_registry = ToolRegistry.from_dict(tool_registry) + else: + self._tool_registry = None + 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()) @@ -698,14 +943,26 @@ def _msg_item(self, status: str = "in_progress", empty: bool = False) -> dict: } 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 self._tool_registry: + ns, name = self._tool_registry.get_identity(raw_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, } + if ns: + item["namespace"] = ns + return item def _response_obj(self, status: str, incomplete_reason: str | None = None) -> dict: output = [] diff --git a/app/model_capabilities.py b/app/model_capabilities.py index 4637dd2..f5c4a13 100644 --- a/app/model_capabilities.py +++ b/app/model_capabilities.py @@ -9,6 +9,7 @@ from app.safe_logging import sanitize_log_text from app.model_catalog_view import SharedModel +from app.reasoning import resolve_reasoning_effort, thinking_mode PROFILES = frozenset(("cn-cli", "cn-work", "intl-cli", "intl-work")) _TEXT = frozenset(("id", "name", "vendor", "description", "descriptionZh", "descriptionEn", "credits", "summary")) _BOOL = frozenset(("supportsImages", "disabledMultimodal", "supportsToolCall", "supportsReasoning", @@ -209,8 +210,7 @@ def from_request(cls, body, payload=None, protocol="chat"): tools = bool(body.get("tools")) or any(isinstance(message, dict) and (message.get("role") == "tool" or message.get("tool_calls")) for message in messages) effort = body.get("reasoning_effort") - thinking = (payload or {}).get("thinking") if protocol == "messages" else None - thinking = thinking.get("type") if isinstance(thinking, dict) else None + thinking = thinking_mode(payload or {}) if protocol == "messages" else None output = body.get("max_tokens") param = "max_output_tokens" if protocol == "responses" else "max_tokens" if output is not None and _positive(output) is None: @@ -227,16 +227,17 @@ def violations(self, model): failures.append(("unsupported_image_input", self.image_param, "当前路由的模型声明不支持图片输入")) if self.tools and caps["tools"] is False: failures.append(("unsupported_tools", "tools", "当前路由的模型声明不支持工具调用")) - enabled = self.effort not in (None, "none") or self.thinking in ("enabled", "adaptive") - disabled = self.effort == "none" or self.thinking == "disabled" + effort = resolve_reasoning_effort(self.effort, self.thinking, model) + enabled = effort not in (None, "none") + disabled = effort == "none" if enabled and caps["reasoning"] is False: failures.append(("unsupported_reasoning", "reasoning_effort", "当前路由的模型声明不支持思考")) if disabled and caps["thinking_disable"] is False: failures.append(("reasoning_required", "reasoning_effort", "当前路由的模型声明不能关闭思考")) reasoning = model.get("reasoning") if isinstance(model.get("reasoning"), dict) else {} efforts = reasoning.get("supportedEfforts") - if (self.effort not in (None, "none") and isinstance(efforts, list) and (efforts or isinstance(model, SharedModel)) - and all(isinstance(value, str) for value in efforts) and self.effort not in efforts): + if (effort not in (None, "none") and isinstance(efforts, list) and (efforts or isinstance(model, SharedModel)) + and all(isinstance(value, str) for value in efforts) and effort not in efforts): failures.append(("unsupported_reasoning_effort", "reasoning_effort", "思考强度不在当前模型声明的选项中")) maximum = _positive(model.get("maxOutputTokens")) if maximum is not None and self.max_output is not None and self.max_output > maximum: diff --git a/app/reasoning.py b/app/reasoning.py new file mode 100644 index 0000000..542f6bf --- /dev/null +++ b/app/reasoning.py @@ -0,0 +1,121 @@ +"""Normalize readable reasoning and request controls for Chat upstreams.""" + + +class ReasoningInputError(ValueError): + """Identify an unsupported reasoning field without exposing its contents.""" + + def __init__(self, message, field=""): + self.field = field + super().__init__(f"{field}: {message}" if field else message) + + +def _text_parts(parts, kinds, field): + if parts is None: + return [] + if not isinstance(parts, list): + raise ReasoningInputError("must be an array", field) + texts = [] + for index, part in enumerate(parts): + location = f"{field}[{index}]" + if not isinstance(part, dict) or part.get("type") not in kinds: + raise ReasoningInputError("unsupported reasoning text block", location) + if not isinstance(part.get("text"), str): + raise ReasoningInputError("must be a string", location + ".text") + texts.append(part["text"]) + return texts + + +def extract_reasoning_text(block): + """Return readable thinking without interpreting signatures or encrypted state.""" + kind = block.get("type") + if kind == "redacted_thinking": + raise ReasoningInputError("encrypted reasoning cannot be converted to Chat", "data") + if kind == "thinking": + text = block.get("thinking") + if not isinstance(text, str): + raise ReasoningInputError("must be a string", "thinking") + if not text and block.get("signature") not in (None, ""): + raise ReasoningInputError("encrypted-only thinking cannot be converted to Chat", "signature") + return text + if kind == "reasoning": + if block.get("encrypted_content") not in (None, ""): + raise ReasoningInputError("encrypted reasoning cannot be converted to Chat", "encrypted_content") + content = _text_parts(block.get("content"), ("reasoning_text", "text"), "content") + summary = _text_parts(block.get("summary"), ("summary_text",), "summary") + return "".join(content if content else summary) + return None + + +def thinking_mode(body): + """Read the already-validated Messages thinking mode for account selection.""" + thinking = body.get("thinking") + return thinking.get("type") if isinstance(thinking, dict) else None + + +def _object(body, field): + value = body.get(field) + if value is None: + return {} + if not isinstance(value, dict): + raise ReasoningInputError("must be an object", field) + return value + + +def _effort(value, field, allowed=None): + if value is not None and (not isinstance(value, str) or not value.strip() + or (allowed is not None and value not in allowed)): + raise ReasoningInputError("unsupported reasoning effort", field) + return value + + +def map_reasoning_controls(body, chat, *, protocol): + """Map explicit controls, leaving implicit activation to the selected account.""" + explicit = "reasoning_effort" in body + effort = _effort(body.get("reasoning_effort"), "reasoning_effort") + if protocol == "responses": + reasoning = _object(body, "reasoning") + if not explicit and "effort" in reasoning: + effort = _effort(reasoning["effort"], "reasoning.effort") + explicit = True + elif protocol == "messages": + thinking = _object(body, "thinking") + mode = thinking_mode(body) + if body.get("thinking") is not None and mode not in ("enabled", "adaptive", "disabled"): + raise ReasoningInputError("must be enabled, adaptive or disabled", "thinking.type") + if mode == "enabled": + budget = thinking.get("budget_tokens") + if type(budget) is not int or budget < 1024: + raise ReasoningInputError("must be an integer >= 1024", "thinking.budget_tokens") + elif "budget_tokens" in thinking: + raise ReasoningInputError("requires thinking.type=enabled", "thinking.budget_tokens") + if thinking.get("display") not in (None, "summarized"): + raise ReasoningInputError("only summarized display is supported by Chat upstreams", "thinking.display") + output = _object(body, "output_config") + if not explicit and "effort" in output: + effort = _effort(output["effort"], "output_config.effort", ("low", "medium", "high", "xhigh", "max")) + explicit = True + if mode == "disabled": + effort, explicit = "none", True + else: + raise ValueError("unsupported reasoning protocol") + if explicit: + chat["reasoning_effort"] = effort + + +def resolve_reasoning_effort(effort, mode, model): + """Use the same account-owned default for capability checks and upstream requests.""" + if mode == "disabled": + return "none" + if effort is not None or mode not in ("enabled", "adaptive"): + return effort + model = model or {} + reasoning = model.get("reasoning") if isinstance(model.get("reasoning"), dict) else {} + supported = reasoning.get("supportedEfforts") + supported = supported if isinstance(supported, list) and supported and all(isinstance(v, str) for v in supported) else None + candidates = [reasoning.get("defaultEffort"), reasoning.get("effort"), + "high", "medium", "low", "xhigh", "max", "minimal", *(supported or [])] + for candidate in candidates: + if (isinstance(candidate, str) and candidate.strip() and candidate != "none" + and (supported is None or candidate in supported)): + return candidate + return "high" diff --git a/converter.py b/converter.py index 543c27a..cbd8fc3 100644 --- a/converter.py +++ b/converter.py @@ -66,6 +66,7 @@ def desensitize_body(body, roles=("system",), desensitize_harness_user=False, from app.inference_resources import (AccountCapacity, InferenceResourcesMiddleware, inference_lifespan, request_resources, release_credential) from app.request_context import SessionIdentifierError, current_context +from app.reasoning import resolve_reasoning_effort, thinking_mode from app import model_capabilities from app.message_normalization import merge_intl_user_images from app.adapters.chat_input import normalize_chat_messages @@ -1931,21 +1932,28 @@ def _cred_for(payload: dict, model: str | None = None, *, region=None, tried=(), def _route_chat(payload, body, rid, *, tried=()): """Validate account capabilities and derive each routed body from canonical input.""" context = current_context() + protocol = context.protocol if context is not None else "chat" + mode = thinking_mode(payload) if protocol == "messages" else None 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) + requirements = model_capabilities.Requirements.from_request(body, payload, protocol) 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: + metadata = None + if mode in ("enabled", "adaptive") or (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"))) + metadata = model_capabilities.entry_model(sys.modules[__name__], entry, body.get("model")) + if requirements is not None and CONFIG.get("cred_pool") is None: + failures = requirements.violations(metadata) if failures: raise model_capabilities.capability_error(failures) canonical = body + effort = resolve_reasoning_effort(body.get("reasoning_effort"), mode, metadata) + if effort != body.get("reasoning_effort"): + body = {**body, "reasoning_effort": effort} if routed_model != body.get("model"): body = {**body, "model": routed_model} body, merged_runs, merged_messages = merge_intl_user_images(body, profile) @@ -2716,12 +2724,18 @@ def _prepare_chat_body(body: dict, *, region=None, session_payload=None) -> dict return body +def _upstream_chat_body(body: dict) -> dict: + """Exclude process-local metadata from both wire JSON and its byte budget.""" + return {key: value for key, value in body.items() + if key is not _REQUEST_POLICY_KEY and not (isinstance(key, str) and key.startswith("_"))} + + 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): + for part in json.JSONEncoder(ensure_ascii=False, separators=(",", ":"), allow_nan=False).iterencode(_upstream_chat_body(body)): size += len(part.encode("utf-8")) if size > limit: _log(f"[limit] 请求体超限,拒绝请求 | limit_bytes={limit}") @@ -3089,7 +3103,7 @@ def retry(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, + async with open_backend_stream(url, headers, _upstream_chat_body(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 @@ -3644,7 +3658,11 @@ def attempt(routed, cred, headers, url): 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_registry = body.get("_tool_registry") + converter = (AnthropicStreamConverter(model=model_name) if anthropic else + ResponsesStreamConverter(model=model_name, + parallel_tool_calls=body.get("parallel_tool_calls", True), + tool_registry=tool_registry)) async def fetch(routed, cred, headers, url): return await _fetch_checked_chat(url, headers, routed, model_name, rid, cred, @@ -3671,6 +3689,7 @@ async def _stream_adapted(url, headers, body, model_name, t0, rid, cred=None, *, """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_registry = body.get("_tool_registry") state = {} tracker = None declared_names = _declared_tool_names(body) @@ -3684,11 +3703,13 @@ 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_registry=tool_registry)) 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_registry=tool_registry)) sent = False upstream = _chat_sse_lines( url, headers, body, model_name, t0, rid, cred, diff --git a/docs/advanced.md b/docs/advanced.md index 7557048..40f3df7 100644 --- a/docs/advanced.md +++ b/docs/advanced.md @@ -42,7 +42,7 @@ Compose explicitly passes some environment variables and CLI flags, so deleting | `--request-context-mode` | `legacy` | `scoped` enables explicit sessions and per-attempt tracing; changes apply to new requests | | `--failover-max` | `0` | Extra credentials tried when a request fails before the first response byte reaches the client; `0` keeps the upstream behaviour of surfacing the failure directly | | `--retry-write-timeout` | `false` | Opt a request-body write timeout into replay (fresh connection and `--failover-max`), accepting that bytes already sent may have been processed | -| `--max-request-bytes` | `33554432` | Positive byte limit for the processed upstream JSON | +| `--max-request-bytes` | `33554432` | Positive byte limit for processed upstream JSON, excluding gateway-only metadata | | `--log-body-limit` | `65536` | Legacy text-preview option; text output is retired and SQLite diagnostics use their own budget | Environment variables include `CODEBUDDY_AUTH_DIR`, `CODEBUDDY_IMPORT_DIR`, `CODEBUDDY2API_KEY`, `CODEBUDDY2API_ADMIN_CSRF`, `CODEBUDDY2API_ADMIN_ORIGINS`, `CODEBUDDY2API_KEEP_TOOL_METADATA`, `CODEBUDDY2API_STREAM_MODE`, `CODEBUDDY2API_LOG`, `CODEBUDDY2API_RESPONSES_PROJECTION_MODE`, `CODEBUDDY2API_RESPONSES_PROJECTION_MAX_BYTES`, `CODEBUDDY2API_MAX_IMAGES`, `CODEBUDDY2API_IMAGE_POLICY`, `CODEBUDDY2API_MAX_REQUEST_BYTES`, `CODEBUDDY2API_LOG_BODY_LIMIT`, `CODEBUDDY2API_FAILOVER_MAX` and `CODEBUDDY2API_RETRY_WRITE_TIMEOUT`. See [deployment](deployment.md) for startup examples. @@ -243,6 +243,16 @@ Credential domain / token issuer determine the product identity. Chat and refres Both international profiles merge image-bearing consecutive `user` runs only after routing, preserving content order and image data. Domestic bodies, text-only runs and system/assistant/tool boundaries remain unchanged. Conflicting message attributes or unrepresentable content return `400 / image_user_run_not_mergeable`; final byte limits still apply. This compatibility step remains enabled when capability preflight is disabled; it neither adds retries nor makes a text model natively visual. +## Reasoning compatibility + +Messages `enabled` / `adaptive` activate Chat reasoning; `output_config.effort` and Responses `reasoning.effort` map to `reasoning_effort`. An explicit top-level `reasoning_effort` takes precedence, except Messages `disabled` always selects `none`; model capability checks still apply. Omitted controls leave upstream defaults unchanged. + +Without an explicit effort, Messages activation uses the selected account's `reasoning.defaultEffort` or legacy `reasoning.effort`, restricted to its declared options. Otherwise it prefers `high`, then an available option; unknown declarations fall back to `high`. Failover resolves the replacement account's default again. + +Manual `enabled` requires an integer `budget_tokens >= 1024`, but the budget is not an exact upstream token limit; `max_tokens` is forwarded unchanged. This mapping does not reproduce native adaptive scheduling. Only `display: summarized` is supported. + +Readable history is kept in `reasoning_content`, never ordinary answer text. Responses uses readable `content` before `summary`; summaries cannot reconstruct native hidden reasoning. Signatures are not forwarded. `redacted_thinking`, encrypted-only thinking and non-empty Responses `encrypted_content` return 400 before routing; upstreams decide which readable history they use. + ## Request boundaries - All three generation protocols normalize `developer` to `system`, move an existing system message first or insert a default. This normalization does not mutate the caller's payload. Optional [Responses projection](#responses-projection) and desensitization process content separately; the whole pipeline is not a verbatim pass-through by default. diff --git a/docs/advanced.zh-CN.md b/docs/advanced.zh-CN.md index 58414d3..b5aaa40 100644 --- a/docs/advanced.zh-CN.md +++ b/docs/advanced.zh-CN.md @@ -42,7 +42,7 @@ Compose 会显式传入部分环境变量及 CLI 参数,删除 `.env` 中的 | `--request-context-mode` | `legacy` | `scoped` 启用显式会话与逐尝试追踪;变更只影响新请求 | | `--failover-max` | `0` | 请求在「一个字节都还没发给下游」之前失败时,最多再换几个凭证就地重放;`0` 表示如实把失败回给下游 | | `--retry-write-timeout` | `false` | 让「写请求体超时」也参与重放(换新连接与 `--failover-max` 换凭证),代价是已发出的那半截正文可能已被上游处理 | -| `--max-request-bytes` | `33554432` | 处理后的上游 JSON 字节上限,须为正整数 | +| `--max-request-bytes` | `33554432` | 处理后的上游 JSON 字节上限,不含网关内部元数据,须为正整数 | | `--log-body-limit` | `65536` | 旧文本预览兼容项;文本输出已停用,SQLite 诊断使用独立预算 | 环境变量包括 `CODEBUDDY_AUTH_DIR`、`CODEBUDDY_IMPORT_DIR`、`CODEBUDDY2API_KEY`、`CODEBUDDY2API_ADMIN_CSRF`、`CODEBUDDY2API_ADMIN_ORIGINS`、`CODEBUDDY2API_KEEP_TOOL_METADATA`、`CODEBUDDY2API_STREAM_MODE`、`CODEBUDDY2API_LOG`、`CODEBUDDY2API_RESPONSES_PROJECTION_MODE`、`CODEBUDDY2API_RESPONSES_PROJECTION_MAX_BYTES`、`CODEBUDDY2API_MAX_IMAGES`、`CODEBUDDY2API_IMAGE_POLICY`、`CODEBUDDY2API_MAX_REQUEST_BYTES`、`CODEBUDDY2API_LOG_BODY_LIMIT`、`CODEBUDDY2API_FAILOVER_MAX`、`CODEBUDDY2API_RETRY_WRITE_TIMEOUT`。启动示例见[部署指南](deployment.zh-CN.md)。 @@ -243,6 +243,16 @@ WebUI 可以直接上传文件;以下限制针对 `POST /admin/credentials` 两种国际产品在选路后归并含图的连续 `user` 段,保留内容顺序和图片数据;国内请求、纯文本段及 system/assistant/tool 边界不变。消息级属性冲突或内容无法无损表达时返回 `400 / image_user_run_not_mergeable`,最终字节限制仍生效。图片兼容不随能力开关关闭,不增加重试,也不让文本模型获得原生视觉。 +## 思考兼容 + +Messages 的 `enabled` / `adaptive` 启用 Chat 推理,`output_config.effort` 和 Responses 的 `reasoning.effort` 映射为 `reasoning_effort`。显式顶层 `reasoning_effort` 优先,但 Messages 的 `disabled` 始终使用 `none`;模型能力检查仍生效。未提供控制参数时保留上游默认行为。 + +未指定强度时,Messages 从实际选中账号的 `reasoning.defaultEffort` 或旧 `reasoning.effort` 选择声明支持的默认值;否则优先 `high`,再选可用选项,声明未知时回落为 `high`。换号后重新解析新账号的默认值。 + +手工 `enabled` 要求整数 `budget_tokens >= 1024`,但该预算不等于上游精确 token 限额,`max_tokens` 原样转发;兼容映射不模拟原生自适应调度。仅支持 `display: summarized`。 + +可读历史统一放入 `reasoning_content`,不混入普通正文。Responses 优先使用可读 `content`,其次使用 `summary`;摘要不能还原原生隐藏推理。签名不转发,`redacted_thinking`、仅含加密签名的思考及非空 Responses `encrypted_content` 在选路前返回 400;上游自行决定使用哪些可读历史。 + ## 请求边界 - 三个生成协议统一将 `developer` 归一为 `system`,已有 system 移到首位,缺失时补默认值;归一化不修改调用方 payload。可选的 [Responses 投影](#responses-投影)和脱敏会另行处理内容,因此默认链路并非逐字透传。 diff --git a/docs/clients.md b/docs/clients.md index 29db1f9..93de1c1 100644 --- a/docs/clients.md +++ b/docs/clients.md @@ -75,7 +75,11 @@ Streaming policy is a server setting, not a client request field. The default `c - `developer` messages become `system`; the first system message is placed first before matching tool results, without mutating the original payload. - Chat accepts mixed Anthropic `tool_use` / `tool_result` history, preserving call IDs, arguments, result images and error markers; ordinary `thinking` becomes `reasoning_content`, not visible text. Native Chat fields stay unchanged. - Conflicting fields, unmatched tool results, unsupported mixed blocks and `redacted_thinking` return HTTP 400 before routing. Split user messages accept only `role` and `content`, with all `tool_result` blocks before ordinary text/images; Anthropic thinking signatures are not forwarded. -- Named function choices are sent upstream as `required` with only that function available; invalid names are rejected locally. +- Messages `thinking` blocks and Responses readable `reasoning` items are retained as assistant `reasoning_content`, including full-history tool continuations; encrypted-only history returns HTTP 400. +- Messages `thinking` / `output_config.effort` and Responses `reasoning.effort` map to upstream reasoning controls. See [reasoning compatibility](advanced.md#reasoning-compatibility) for defaults and budget limits. +- Named function choices are sent upstream as `required` with only that function available. Responses matches the exact `(namespace, name)`; unknown or historical-only choices return HTTP 400. +- Responses keeps current and historical tool identities distinct. Schema cleanup removes only boolean `encrypted` markers in schema positions, preserving property names, definitions and literal data. +- Namespaced tools use Chat-safe aliases of at most 64 characters, including nested and long identities; Responses output and history retain the original namespace and name. Plain tool names stay unchanged. - Errors follow the client protocol's own shape (OpenAI `error` object vs Anthropic `{"type":"error"}`), and status codes are preserved. Realtime mode can deliver useful deltas before a later invalid terminal, disconnect or size error; valid truncation/filter distinctions remain native. Clients must not treat an opened SSE connection as proof of successful completion. - `POST /v1/messages/count_tokens` returns a character-based heuristic estimate for budgeting, not an exact count. diff --git a/docs/clients.zh-CN.md b/docs/clients.zh-CN.md index 71d589c..8d7e906 100644 --- a/docs/clients.zh-CN.md +++ b/docs/clients.zh-CN.md @@ -75,7 +75,11 @@ Cherry Studio、ZCode、LobeChat、NextChat、Open WebUI 或自研 SDK 客户端 - 先将 `developer` 转为 `system` 并置顶首条系统消息,再关联工具结果;不改动调用方原始载荷。 - Chat 兼容混入的 Anthropic `tool_use` / `tool_result` 历史,保留调用 ID、参数、结果图片与错误标记;普通 `thinking` 转为 `reasoning_content`,不混入正文,原生 Chat 字段保持不变。 - 字段冲突、工具结果无法关联、不支持的混合内容块及 `redacted_thinking` 在选路前返回 HTTP 400。需拆分的用户消息只能包含 `role`、`content`,且 `tool_result` 必须在普通文本/图片之前;Anthropic 思考签名不转发。 -- 指定名称的函数选择会以 `required` 且仅含该函数的形式发往上游;无效名称在本地拒绝。 +- Messages 的 `thinking` 块和 Responses 的可读 `reasoning` 项统一保留为 assistant `reasoning_content`,支持完整历史及工具续接;仅含加密状态的历史返回 HTTP 400。 +- Messages 的 `thinking` / `output_config.effort` 与 Responses 的 `reasoning.effort` 会映射到上游思考控制;默认值和预算限制见[思考兼容](advanced.zh-CN.md#思考兼容)。 +- 指定名称的函数选择会以 `required` 且仅含该函数的形式发往上游。Responses 按 `(namespace, name)` 精确匹配;未知或仅在历史中出现的选择返回 HTTP 400。 +- Responses 分别保留当前与历史工具身份。Schema 清理仅移除 schema 位置的布尔 `encrypted` 标记,保留属性名、定义及字面数据。 +- 命名空间工具(含嵌套和长名称)使用不超过 64 字符的合法 Chat 别名;Responses 输出与历史保留原命名空间和名称,普通工具名称保持不变。 - 错误按客户端协议各自的形态返回(OpenAI 的 `error` 对象与 Anthropic 的 `{"type":"error"}`),状态码保留。实时模式可能先送出有效增量,随后才遇到非法终端状态、断连或大小错误;合法截断/过滤仍保留协议原生区别。客户端不能仅凭 SSE 已开启就认定最终成功。 - `POST /v1/messages/count_tokens` 返回按字符估算的启发式结果,用于预算参考,不是精确计数。 diff --git a/tests/test_reasoning_requests.py b/tests/test_reasoning_requests.py new file mode 100644 index 0000000..94ff06d --- /dev/null +++ b/tests/test_reasoning_requests.py @@ -0,0 +1,376 @@ +"""Verify shared reasoning controls, protocol round trips and account-owned defaults.""" +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from copy import deepcopy +from itertools import permutations, product +import json +import unittest +from unittest.mock import patch + +import httpx +import converter as gateway +from app.control_store import ControlStore +import test_api_flow as fixtures +import test_runtime_endpoints as runtime + +MODES = ((False, "compatible"), (True, "compatible"), (True, "realtime")) +PROTOCOLS = ("chat/completions", "messages", "responses") + + +def model(effort="medium", **extra): + return {"id": "shared-model", "credits": "x0.00", "supportsReasoning": True, + "supportsToolCall": True, "canDisableThinking": True, + "reasoning": {"defaultEffort": effort, "supportedEfforts": ["low", "medium", "high", "xhigh", "max"]}, **extra} + + +def payload(protocol, **extra): + body = {"model": "shared-model", "max_tokens": 2048} + body["input" if protocol == "responses" else "messages"] = [{"role": "user", "content": "question"}] + body.update(extra) + return body + + +def reasoning(text): + return {"type": "reasoning", "summary": [{"type": "summary_text", "text": text}]} + + +def events(response): + return [json.loads(line[6:]) for line in response.text.splitlines() if line.startswith("data: ") and line[6:] != "[DONE]"] + + +def output(response, protocol, stream): + if not stream: + return response.json() + parsed = events(response) + if protocol == "responses": + return next(event["response"] for event in parsed if event.get("type") == "response.completed") + blocks = [] + for event in parsed: + if event["type"] == "content_block_start": + blocks.append(deepcopy(event["content_block"])) + elif event["type"] == "content_block_delta": + block, delta = blocks[event["index"]], event["delta"] + if delta["type"] == "thinking_delta": + block["thinking"] += delta["thinking"] + elif delta["type"] == "text_delta": + block["text"] += delta["text"] + elif delta["type"] == "input_json_delta": + block["_arguments"] = block.get("_arguments", "") + delta["partial_json"] + for block in blocks: + if "_arguments" in block: + block["input"] = json.loads(block.pop("_arguments")) + return {"content": blocks} + + +class ReasoningEndpointTests(unittest.TestCase): + def setUp(self): + self.fx = runtime.EndpointTests("test_stateful_responses_fields_are_rejected") + self.addCleanup(self.fx.doCleanups) + self.fx.setUp() + self.metadata = self.enterContext(patch.object(gateway.model_capabilities, "entry_model", return_value=model())) + self.fx.respond = self.respond + + def respond(self, request): + body = json.loads(request.content) + delta = {"content": "answer"} + if body.get("reasoning_effort") not in (None, "none"): + delta["reasoning_content"] = "new reasoning" + return httpx.Response(200, content=runtime.sse(delta)) + + def post(self, protocol, body, *, stream=False, mode="compatible", projection="balanced", desensitize=False): + original = deepcopy(body) + before = len(self.fx.requests) + with patch.dict(gateway.CONFIG, stream_mode=mode, responses_projection_mode=projection, desensitize=desensitize): + response = self.fx.client.post("/v1/" + protocol, json={**body, "stream": stream}) + self.assertEqual(body, original) + upstream = [json.loads(request.content) for request in self.fx.requests[before:]] + return response, upstream + + def test_controls_and_output_in_all_modes(self): + cases = [ + ("chat/completions", {"reasoning_effort": "high"}, "high"), + ("responses", {"reasoning": {"effort": "high"}}, "high"), + ("responses", {"reasoning": {"effort": "high"}, "reasoning_effort": "low"}, "low"), + ("responses", {"reasoning": {"effort": "none"}}, "none"), + ("messages", {"thinking": {"type": "enabled", "budget_tokens": 1024}}, "medium"), + ("messages", {"thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}}, "low"), + ("messages", {"thinking": {"type": "enabled", "budget_tokens": 1024}, "output_config": {"effort": "high"}}, "high"), + ("messages", {"thinking": {"type": "disabled"}, "output_config": {"effort": "high"}}, "none"), + ("messages", {"reasoning_effort": "xhigh", "output_config": {"effort": "low"}}, "xhigh"), + ("messages", {"output_config": {"effort": "high"}}, "high"), + ] + [(protocol, {}, None) for protocol in PROTOCOLS] + for (protocol, fields, expected), (stream, mode) in product(cases, MODES): + with self.subTest(protocol=protocol, fields=fields, stream=stream, mode=mode): + response, sent = self.post(protocol, payload(protocol, **fields), stream=stream, mode=mode) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(len(sent), 1) + self.assertEqual(sent[0].get("reasoning_effort"), expected) + self.assertEqual(sent[0]["max_tokens"], 2048) + self.assertTrue({"thinking", "output_config", "reasoning", "budget_tokens"}.isdisjoint(sent[0])) + self.assertEqual("new reasoning" in response.text, expected not in (None, "none")) + + def test_readable_history_is_separate_from_visible_text(self): + thoughts = [{"type": "thinking", "thinking": "first", "signature": "signature-canary"}, + {"type": "thinking", "thinking": "second"}] + for protocol, (stream, mode), only, projection, desensitize in product( + PROTOCOLS, MODES, (False, True), ("balanced", "passthrough"), (False, True)): + with self.subTest(protocol=protocol, stream=stream, mode=mode, only=only, projection=projection, desensitize=desensitize): + body = payload(protocol) + history = [{"role": "user", "content": "question"}] + if protocol == "responses": + history += [reasoning("first"), reasoning("second")] + if not only: + history.append({"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "answer"}]}) + else: + history.append({"role": "assistant", "content": deepcopy(thoughts) + ([] if only else [{"type": "text", "text": "answer"}])}) + history.append({"role": "user", "content": "continue"}) + body["input" if protocol == "responses" else "messages"] = history + response, sent = self.post(protocol, body, stream=stream, mode=mode, projection=projection, desensitize=desensitize) + self.assertEqual(response.status_code, 200, response.text) + assistants = [m for m in sent[0]["messages"] if m["role"] == "assistant"] + self.assertEqual(len(assistants), 1) + self.assertEqual(assistants[0]["reasoning_content"], "firstsecond") + self.assertNotIn("first", json.dumps(assistants[0]["content"])) + self.assertNotIn("signature-canary", json.dumps(sent)) + + def test_gateway_output_roundtrips_through_tool_results(self): + call = {"index": 0, "id": "call_shared", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + first = runtime.sse({"reasoning_content": "firstsecond", "tool_calls": [call]}, finish="tool_calls") + for protocol, (stream, mode), projection in product(("messages", "responses"), MODES, ("balanced", "passthrough")): + with self.subTest(protocol=protocol, stream=stream, mode=mode, projection=projection): + self.fx.respond = lambda request: httpx.Response(200, content=first) + body = payload(protocol) + body["tools"] = ([{"name": "lookup", "input_schema": {"type": "object", "properties": {}}}] + if protocol == "messages" else [{"type": "function", "name": "lookup", "parameters": {"type": "object", "properties": {}}}]) + response, _ = self.post(protocol, body, stream=stream, mode=mode, projection=projection) + self.assertEqual(response.status_code, 200, response.text) + result = output(response, protocol, stream) + if protocol == "messages": + body["messages"] += [{"role": "assistant", "content": result["content"]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "call_shared", "content": "43"}]}] + else: + body["input"] += result["output"] + [{"type": "function_call_output", "call_id": "call_shared", "output": "43"}] + self.fx.respond = self.respond + response, sent = self.post(protocol, body, stream=stream, mode=mode, projection=projection) + self.assertEqual(response.status_code, 200, response.text) + messages = sent[0]["messages"] + assistant = next(m for m in messages if m["role"] == "assistant") + self.assertEqual(assistant["reasoning_content"], "firstsecond") + self.assertEqual(assistant["tool_calls"][0]["id"], "call_shared") + self.assertEqual(messages[messages.index(assistant) + 1], {"role": "tool", "tool_call_id": "call_shared", "content": "43"}) + + def test_realtime_tool_first_output_preserves_one_assistant_turn(self): + deltas = [{"tool_calls": [{"index": 0, "id": "call_first", "type": "function", + "function": {"name": "lookup", "arguments": "{}"}}]}, + {"reasoning_content": "late reasoning"}, {"content": "after tool"}] + chunks = [{"choices": [{"index": 0, "delta": delta, "finish_reason": None}]} for delta in deltas] + chunks.append({"choices": [{"index": 0, "delta": {}, "finish_reason": "tool_calls"}]}) + raw = ("".join("data: " + json.dumps(chunk) + "\n\n" for chunk in chunks) + "data: [DONE]\n\n").encode() + self.fx.respond = lambda request: httpx.Response(200, content=raw) + body = payload("responses", tools=[{"type": "function", "name": "lookup", "parameters": {"type": "object", "properties": {}}}]) + response, _ = self.post("responses", body, stream=True, mode="realtime") + self.assertEqual(response.status_code, 200, response.text) + result = output(response, "responses", True) + self.assertEqual([item["type"] for item in result["output"]], ["function_call", "reasoning", "message"]) + body["input"] += result["output"] + [{"type": "function_call_output", "call_id": "call_first", "output": "43"}] + self.fx.respond = self.respond + response, sent = self.post("responses", body, stream=True, mode="realtime") + self.assertEqual(response.status_code, 200, response.text) + messages = sent[0]["messages"] + assistants = [m for m in messages if m["role"] == "assistant"] + self.assertEqual(len(assistants), 1) + self.assertEqual(assistants[0]["reasoning_content"], "late reasoning") + self.assertEqual(assistants[0]["content"], "after tool") + self.assertEqual(assistants[0]["tool_calls"][0]["id"], "call_first") + self.assertEqual(messages[messages.index(assistants[0]) + 1]["tool_call_id"], "call_first") + + def test_responses_keeps_upstream_declared_effort_extensions(self): + self.metadata.return_value = model(reasoning={"supportedEfforts": ["ultra"]}) + response, sent = self.post("responses", payload("responses", reasoning={"effort": "ultra"})) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(sent[0]["reasoning_effort"], "ultra") + response, sent = self.post("responses", payload("responses", reasoning={"effort": "low"})) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(response.json()["error"]["code"], "unsupported_reasoning_effort") + self.assertEqual(sent, []) + + def test_responses_assistant_item_orders_keep_tools_and_reasoning_together(self): + items = [reasoning("same turn"), + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "answer"}]}, + {"type": "function_call", "name": "lookup", "call_id": "call_1", "arguments": "{}"}] + for ordered in permutations(items): + with self.subTest(order=[item["type"] for item in ordered]): + body = payload("responses", input=[{"role": "user", "content": "question"}, *ordered, + {"type": "function_call_output", "call_id": "call_1", "output": "43"}]) + response, sent = self.post("responses", body) + self.assertEqual(response.status_code, 200, response.text) + assistants = [m for m in sent[0]["messages"] if m["role"] == "assistant"] + self.assertEqual(len(assistants), 1) + self.assertEqual(assistants[0]["reasoning_content"], "same turn") + self.assertEqual(assistants[0]["content"], "answer") + self.assertEqual(assistants[0]["tool_calls"][0]["id"], "call_1") + + def test_responses_keeps_reasoning_with_its_assistant_turn(self): + body = payload("responses", input=[ + {"role": "user", "content": "question"}, reasoning("before answer"), + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "answer"}]}, + {"role": "user", "content": "look it up next"}, + reasoning("before tool"), {"type": "function_call", "name": "lookup", "call_id": "call_1", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "call_1", "output": "43"}, + reasoning("after tool"), {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "done"}]}, + {"role": "user", "content": "continue"}, + ]) + response, sent = self.post("responses", body) + self.assertEqual(response.status_code, 200, response.text) + assistants = [m for m in sent[0]["messages"] if m["role"] == "assistant"] + self.assertEqual([m["reasoning_content"] for m in assistants], ["before answer", "before tool", "after tool"]) + self.assertEqual([m["content"] for m in assistants], ["answer", "", "done"]) + + def test_responses_prefers_readable_content_over_its_summary(self): + item = {**reasoning("summary"), "content": [{"type": "reasoning_text", "text": "full "}, {"type": "text", "text": "content"}]} + body = payload("responses", input=[{"role": "user", "content": "question"}, item, + {"role": "assistant", "content": "answer"}, {"role": "user", "content": "continue"}]) + response, sent = self.post("responses", body) + self.assertEqual(response.status_code, 200, response.text) + assistant = next(m for m in sent[0]["messages"] if m["role"] == "assistant") + self.assertEqual(assistant["reasoning_content"], "full content") + + def test_opaque_or_malformed_reasoning_is_rejected_without_upstream(self): + blocks = [{"type": "redacted_thinking", "data": "opaque-canary"}, {"type": "thinking", "thinking": 7}, + {"type": "thinking", "thinking": "", "signature": "opaque-canary"}] + for protocol in ("chat/completions", "messages"): + for block in blocks: + body = payload(protocol, messages=[{"role": "assistant", "content": [block]}]) + response, sent = self.post(protocol, body) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(sent, []) + self.assertNotIn("opaque-canary", response.text) + for item in ({**reasoning("readable summary"), "encrypted_content": "opaque-canary"}, + {"type": "reasoning", "summary": "opaque-canary"}, + {"type": "reasoning", "summary": [{"type": "summary_text", "text": 7}]}): + response, sent = self.post("responses", payload("responses", input=[item])) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(sent, []) + self.assertNotIn("opaque-canary", response.text) + self.assertEqual(self.fx.credentials.call_count, 0) + + def test_invalid_controls_fail_before_routing(self): + for controls in ({"thinking": []}, {"thinking": {"type": "unknown"}}, + {"thinking": {"type": "enabled", "budget_tokens": True}}, + {"thinking": {"type": "enabled", "budget_tokens": 100}}, + {"thinking": {"type": "adaptive", "budget_tokens": 1024}}, + {"thinking": {"type": "adaptive", "display": "omitted"}}, + {"output_config": {"effort": "unknown"}}, {"output_config": []}): + response, sent = self.post("messages", payload("messages", **controls)) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(sent, []) + for control in ([], {"effort": True}, {"effort": " "}): + response, sent = self.post("responses", payload("responses", reasoning=control)) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(sent, []) + self.assertEqual(self.fx.credentials.call_count, 0) + + def test_activation_uses_legacy_defaults_and_skips_disabled_defaults(self): + for metadata, expected in (({"reasoning": {"effort": "low"}}, "low"), + ({"reasoning": {"defaultEffort": "low", "supportedEfforts": []}}, "low"), + ({"reasoning": {"defaultEffort": "none", "supportedEfforts": ["low", "high"]}}, "high"), + ({"reasoning": {"supportedEfforts": ["low"]}}, "low"), ({}, "high")): + self.metadata.return_value = metadata + response, sent = self.post("messages", payload("messages", thinking={"type": "adaptive"})) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(sent[0]["reasoning_effort"], expected) + + def test_disabled_preserves_capability_checks_and_does_not_request_reasoning(self): + self.metadata.return_value = model(supportsReasoning=False) + body = payload("messages", thinking={"type": "disabled"}, output_config={"effort": "high"}) + response, sent = self.post("messages", body) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(sent[0]["reasoning_effort"], "none") + self.metadata.return_value = model(canDisableThinking=False) + response, sent = self.post("messages", body) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(response.json()["error"]["code"], "reasoning_required") + self.assertEqual(sent, []) + + def test_reasoning_counts_toward_request_size_limit(self): + for protocol in ("messages", "responses"): + body = payload(protocol) + if protocol == "messages": + body["messages"] = [{"role": "assistant", "content": [{"type": "thinking", "thinking": "x" * 2048}]}] + else: + body["input"] = [reasoning("x" * 2048)] + with patch.dict(gateway.CONFIG, max_request_bytes=1024): + response, sent = self.post(protocol, body) + self.assertEqual(response.status_code, 413, response.text) + self.assertEqual(sent, []) + + +class ReasoningRoutingTests(fixtures.GatewayFixture, unittest.TestCase): + def test_each_account_uses_its_own_default_even_within_one_profile(self): + self.fx.add_account("second", "intl-cli") + self.fx.configure(profiles=("intl-cli", "second")) + self.fx.account_catalogs({"intl-cli": [model("low")], "second": [model("high")]}) + seen = set() + for guard in (True, False): + for _ in range(4): + body = self.fx.payload("messages") + body["thinking"] = {"type": "adaptive"} + with patch.dict(gateway.CONFIG, model_capability_guard=guard): + request, sent = self.fx.post_ok("messages", body, {"intl-cli", "second"}) + uid = request.headers["x-user-id"] + seen.add(uid) + self.assertEqual(sent["reasoning_effort"], {"intl-cli": "low", "second": "high"}[uid]) + self.assertEqual(seen, {"intl-cli", "second"}) + body["output_config"] = {"effort": "medium"} + _, sent = self.fx.post_ok("messages", body, seen) + self.assertEqual(sent["reasoning_effort"], "medium") + + def test_failover_recomputes_defaults_without_mutating_canonical_input(self): + for stream, mode in MODES: + with self.subTest(stream=stream, mode=mode): + self.fx.configure(profiles=("cn-cli", "intl-work")) + self.fx.account_catalogs({"cn-cli": [model("low")], "intl-work": [model("high")]}) + seen = [] + def respond(request): + seen.append(request) + if len(seen) == 1: + return httpx.Response(429, json={"error": {"message": "synthetic quota"}}) + return httpx.Response(200, content=fixtures.fixtures.success_sse()) + body = self.fx.payload("messages", stream=stream) + body["thinking"] = {"type": "adaptive"} + original = deepcopy(body) + with self.responder(respond), patch.dict(gateway.CONFIG, failover_max=1, stream_mode=mode): + response = self.fx.client.post("/v1/messages", json=body) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(body, original) + self.assertEqual(len(seen), 2) + self.assertNotEqual(seen[0].headers["x-user-id"], seen[1].headers["x-user-id"]) + for request in seen: + expected = {"cn-cli": "low", "intl-work": "high"}[request.headers["x-user-id"]] + self.assertEqual(json.loads(request.content)["reasoning_effort"], expected) + self.assertEqual(self.fx.pool._capacity._counts, {}) + + def test_activation_cannot_escape_free_tier_or_strict_binding(self): + self.fx.configure(profiles=("cn-cli", "intl-work")) + self.fx.account_catalogs({"intl-work": [model(supportsReasoning=False)], "cn-cli": [model(credits="x1.00")]}) + body = self.fx.payload("messages") + body["thinking"] = {"type": "adaptive"} + response = self.fx.client.post("/v1/messages", json=body) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(self.fx.requests, []) + self.fx.account_catalogs({"intl-work": [model(supportsReasoning=False)], "cn-cli": [model()]}) + store = ControlStore(self.fx.root / "reasoning-control.sqlite3") + self.addCleanup(store.close) + store.update_model("shared-model", {"profile": "intl-work"}, store.snapshot()["revision"], {"shared-model"}) + with patch.dict(gateway.CONFIG, control_store=store): + response = self.fx.client.post("/v1/messages", json=body) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(self.fx.requests, []) + self.assertEqual(self.fx.pool._capacity._counts, {}) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_responses_adapter.py b/tests/test_responses_adapter.py index eab1ff2..57f6e41 100644 --- a/tests/test_responses_adapter.py +++ b/tests/test_responses_adapter.py @@ -623,6 +623,193 @@ def test_parallel_tool_calls_roundtrip(): print("✅ test_parallel_tool_calls_roundtrip") +def test_tool_registry_and_identity_mapping(): + """Verify ToolRegistry handles plain tools, namespaced tools, collision avoidance, and reversible lookup.""" + from app.adapters.responses_adapter import ToolRegistry + + registry = ToolRegistry() + # 1. Plain tool: must keep original name + name1 = registry.register(None, "collaboration__spawn_agent") + assert name1 == "collaboration__spawn_agent" + assert registry.get_identity("collaboration__spawn_agent") == (None, "collaboration__spawn_agent") + + # 2. Namespaced tool colliding with existing plain tool: must disambiguate safely + name2 = registry.register("collaboration", "spawn_agent") + assert name2 == "collaboration__spawn_agent_1" + assert registry.get_identity("collaboration__spawn_agent_1") == ("collaboration", "spawn_agent") + # Plain tool lookup still returns original identity + assert registry.get_identity("collaboration__spawn_agent") == (None, "collaboration__spawn_agent") + + # 3. Tool with __ in name inside namespace + name3 = registry.register("custom_ns", "exec__command") + assert name3 == "custom_ns__exec__command" + assert registry.get_identity("custom_ns__exec__command") == ("custom_ns", "exec__command") + + # 4. Tools with same name in different namespaces + name_ns1 = registry.register("ns1", "search") + name_ns2 = registry.register("ns2", "search") + assert name_ns1 == "ns1__search" + assert name_ns2 == "ns2__search" + assert registry.get_identity("ns1__search") == ("ns1", "search") + assert registry.get_identity("ns2__search") == ("ns2", "search") + + # 5. Serialization and round-trip + dumped = registry.to_dict() + restored = ToolRegistry.from_dict(dumped) + assert restored.get_identity("custom_ns__exec__command") == ("custom_ns", "exec__command") + assert restored.get_identity("collaboration__spawn_agent") == (None, "collaboration__spawn_agent") + assert restored.get_identity("collaboration__spawn_agent_1") == ("collaboration", "spawn_agent") + print("✅ test_tool_registry_and_identity_mapping") + + +def test_responses_multiagent_request_conversion_and_schema_sanitization(): + """Verify responses_request_to_chat converts tools, strips encrypted markers, and maps tool_choice and history.""" + req = { + "model": "gpt-4o", + "tools": [ + { + "type": "function", + "name": "collaboration__spawn_agent", + "description": "Standard unnamespaced tool", + "parameters": {"type": "object", "properties": {"prompt": {"type": "string"}}}, + }, + { + "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"}, + "meta": { + "type": "object", + "properties": { + "secret": {"type": "string", "encrypted": True}, + "public": {"type": "string"}, + }, + }, + }, + }, + } + ], + }, + { + "type": "namespace", + "name": "ns1", + "tools": [{"type": "function", "name": "search", "parameters": {"type": "object"}}], + }, + { + "type": "namespace", + "name": "ns2", + "tools": [{"type": "function", "name": "search", "parameters": {"type": "object"}}], + }, + { + "type": "namespace", + "name": "custom_ns", + "tools": [{"type": "function", "name": "exec__command", "parameters": {"type": "object"}}], + }, + ], + "tool_choice": {"type": "function", "name": "spawn_agent", "namespace": "collaboration"}, + "input": [ + { + "type": "additional_tools", + "tools": [ + { + "type": "namespace", + "name": "extra", + "tools": [{"type": "function", "name": "ping", "parameters": {"type": "object"}}], + } + ], + }, + { + "type": "agent_message", + "author": "/root", + "recipient": "/root/worker", + "content": [ + {"type": "input_text", "text": "Task header\n"}, + {"type": "encrypted_content", "encrypted_content": "Sensitive payload text"}, + ], + }, + { + "type": "function_call", + "call_id": "call_hist_1", + "name": "spawn_agent", + "namespace": "collaboration", + "arguments": '{"task_name": "t1"}', + }, + { + "type": "function_call_output", + "call_id": "call_hist_1", + "output": "ok", + }, + ], + } + + chat = responses_request_to_chat(req) + chat_tools = chat.get("tools", []) + tool_names = [t["function"]["name"] for t in chat_tools] + + # 1. Verify plain tool preserved original name + assert "collaboration__spawn_agent" in tool_names + # 2. Verify namespaced tool disambiguated + assert "collaboration__spawn_agent_1" in tool_names + # 3. Verify same name under different namespaces preserved + assert "ns1__search" in tool_names + assert "ns2__search" in tool_names + # 4. Verify tool with __ in name inside namespace + assert "custom_ns__exec__command" in tool_names + # 5. Verify additional_tools registered + assert "extra__ping" in tool_names + + # 6. Verify parameter schema cleaning (no 'encrypted' key anywhere) + collab_tool = next(t for t in chat_tools if t["function"]["name"] == "collaboration__spawn_agent_1") + props = collab_tool["function"]["parameters"]["properties"] + assert "encrypted" not in props["message"] + assert "encrypted" not in props["meta"]["properties"]["secret"] + + # 7. Verify tool_choice mapped + assert chat.get("tool_choice") == {"type": "function", "function": {"name": "collaboration__spawn_agent_1"}} + + # 8. Verify messages conversion + messages = chat.get("messages", []) + # User message from agent_message with encrypted_content + user_msg = next(m for m in messages if m["role"] == "user") + assert "Task header\n" in user_msg["content"] + assert "Sensitive payload text" in user_msg["content"] + + # Assistant message from function_call + asst_msg = next(m for m in messages if m["role"] == "assistant") + assert asst_msg["tool_calls"][0]["function"]["name"] == "collaboration__spawn_agent_1" + + # 9. Verify stream converter output reconstruction + conv = ResponsesStreamConverter(model="test", tool_registry=chat.get("_tool_registry")) + + # Namespaced tool call + slot_collab = {"id": "c1", "name": "collaboration__spawn_agent_1", "args": "{}", "fc_id": "fc_1", "output_idx": 0, "emitted": False, "emitted_args_length": 0} + item_collab = conv._fc_item(slot_collab, "completed") + assert item_collab["name"] == "spawn_agent" + assert item_collab["namespace"] == "collaboration" + + # Plain tool call + slot_plain = {"id": "c2", "name": "collaboration__spawn_agent", "args": "{}", "fc_id": "fc_2", "output_idx": 1, "emitted": False, "emitted_args_length": 0} + item_plain = conv._fc_item(slot_plain, "completed") + assert item_plain["name"] == "collaboration__spawn_agent" + assert "namespace" not in item_plain + + # Custom namespace with __ in name + slot_custom = {"id": "c3", "name": "custom_ns__exec__command", "args": "{}", "fc_id": "fc_3", "output_idx": 2, "emitted": False, "emitted_args_length": 0} + item_custom = conv._fc_item(slot_custom, "completed") + assert item_custom["name"] == "exec__command" + assert item_custom["namespace"] == "custom_ns" + + print("✅ test_responses_multiagent_request_conversion_and_schema_sanitization") + + if __name__ == "__main__": test_simple_text_request() test_array_input_request() @@ -647,4 +834,6 @@ def test_parallel_tool_calls_roundtrip(): test_usage_maps_cached_tokens_and_omits_when_unknown() test_reasoning_effort_and_text_format_are_mapped() test_parallel_tool_calls_roundtrip() + test_tool_registry_and_identity_mapping() + test_responses_multiagent_request_conversion_and_schema_sanitization() print(f"\n🎉 All {22} tests passed!") diff --git a/tests/test_responses_multiagent.py b/tests/test_responses_multiagent.py new file mode 100644 index 0000000..973370a --- /dev/null +++ b/tests/test_responses_multiagent.py @@ -0,0 +1,765 @@ +#!/usr/bin/env python3 +"""Interface-level regression tests for Responses multi-agent tools and namespace identity mapping.""" + +from copy import deepcopy +import json +import re +import sys +from pathlib import Path +import unittest +from unittest.mock import patch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import httpx +from fastapi.testclient import TestClient + +import converter +from app import upstream_io + + +def _chat_completion_response(tool_calls=None, content="hello"): + """Build a standard non-streaming Chat completions JSON response.""" + msg = {"role": "assistant", "content": content} + if tool_calls is not None: + msg["tool_calls"] = tool_calls + return { + "id": "chatcmpl_test", + "object": "chat.completion", + "created": 1234567890, + "model": "auto", + "choices": [ + { + "index": 0, + "message": msg, + "finish_reason": "tool_calls" if tool_calls else "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, + } + + +def _chat_sse_stream(tool_calls=None, content=None): + """Build Chat SSE streaming lines for buffered or realtime streams.""" + lines = [] + # Initial chunk + lines.append("data: " + json.dumps({ + "id": "chatcmpl_chunk", + "object": "chat.completion.chunk", + "created": 1234567890, + "model": "auto", + "choices": [{"index": 0, "delta": {"role": "assistant"}, "finish_reason": None}], + }) + "\n\n") + + if content: + lines.append("data: " + json.dumps({ + "id": "chatcmpl_chunk", + "object": "chat.completion.chunk", + "created": 1234567890, + "model": "auto", + "choices": [{"index": 0, "delta": {"content": content}, "finish_reason": None}], + }) + "\n\n") + + if tool_calls: + for idx, tc in enumerate(tool_calls): + fn = tc.get("function", {}) + lines.append("data: " + json.dumps({ + "id": "chatcmpl_chunk", + "object": "chat.completion.chunk", + "created": 1234567890, + "model": "auto", + "choices": [{ + "index": 0, + "delta": { + "tool_calls": [{ + "index": idx, + "id": tc.get("id", f"call_{idx}"), + "type": "function", + "function": {"name": fn.get("name", ""), "arguments": fn.get("arguments", "{}")}, + }] + }, + "finish_reason": None, + }], + }) + "\n\n") + + lines.append("data: " + json.dumps({ + "id": "chatcmpl_chunk", + "object": "chat.completion.chunk", + "created": 1234567890, + "model": "auto", + "choices": [{"index": 0, "delta": {}, "finish_reason": "tool_calls"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, + }) + "\n\n") + else: + lines.append("data: " + json.dumps({ + "id": "chatcmpl_chunk", + "object": "chat.completion.chunk", + "created": 1234567890, + "model": "auto", + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, + }) + "\n\n") + + lines.append("data: [DONE]\n\n") + return "".join(lines).encode("utf-8") + + +def _parse_responses_sse_events(raw_text: str) -> list[dict]: + """Parse SSE event stream from /v1/responses.""" + events = [] + for block in raw_text.strip().split("\n\n"): + if not block: + continue + data_str = next((line[6:] for line in block.splitlines() if line.startswith("data: ")), None) + if data_str and data_str != "[DONE]": + events.append(json.loads(data_str)) + return events + + +class ResponsesMultiAgentInterfaceTests(unittest.TestCase): + def setUp(self): + self.enterContext(patch.dict(converter.CONFIG, { + "api_key": "", "cred": None, "cred_pool": None, "model_guard": False, + "max_images": 16, "image_policy": "truncate", "max_request_bytes": 32 * 1024 * 1024, + "log_body_limit": 65536, "log_path": None, "desensitize": False, "no_compact": False, + "stream_mode": "compatible", + })) + self.enterContext(patch.object(converter, "_cred_for", return_value=(None, {}))) + self.enterContext(patch.object(converter, "_log")) + self.enterContext(patch.object(converter, "_note_cred_status")) + + self.upstream_requests = [] + self.mock_response_generator = None + + def handle_request(request: httpx.Request): + self.upstream_requests.append(request) + if self.mock_response_generator: + return self.mock_response_generator(request) + return httpx.Response(200, json=_chat_completion_response()) + + real_client = httpx.AsyncClient + transport = httpx.MockTransport(handle_request) + self.enterContext(patch.object(upstream_io.httpx, "AsyncClient", + side_effect=lambda **kw: real_client(transport=transport, **kw))) + self.client = self.enterContext(TestClient(converter.app)) + + def _sample_multiagent_payload(self, stream: bool = False): + return { + "model": "auto", + "stream": stream, + "tools": [ + { + "type": "function", + "name": "collaboration__spawn_agent", + "description": "Standard ordinary tool without namespace", + "parameters": {"type": "object", "properties": {"task": {"type": "string"}}}, + }, + { + "type": "namespace", + "name": "ns1", + "tools": [ + { + "type": "function", + "name": "search", + "description": "Search in ns1", + "parameters": {"type": "object", "properties": {"q": {"type": "string"}}}, + } + ], + }, + { + "type": "namespace", + "name": "ns2", + "tools": [ + { + "type": "function", + "name": "search", + "description": "Search in ns2", + "parameters": {"type": "object", "properties": {"q": {"type": "string"}}}, + } + ], + }, + { + "type": "namespace", + "name": "custom_ns", + "tools": [ + { + "type": "function", + "name": "exec__command", + "description": "Tool whose name contains double underscores", + "parameters": { + "type": "object", + "properties": {"cmd": {"type": "string", "encrypted": True}}, + }, + } + ], + }, + ], + "input": [{"role": "user", "content": "Execute multi-agent subtasks"}], + } + + def test_nonstream_tools_isolation_identity_and_metadata_sanitation(self): + """Cover non-streaming mode: verify identity restoration, ordinary tool preservation, and metadata purity.""" + # 1. Test upstream call to ns1.search + tc_ns1 = [{"id": "call_1", "type": "function", "function": {"name": "ns1__search", "arguments": '{"q":"weather"}'}}] + self.mock_response_generator = lambda req: httpx.Response(200, content=_chat_sse_stream(tool_calls=tc_ns1)) + + res = self.client.post("/v1/responses", json=self._sample_multiagent_payload(stream=False)) + self.assertEqual(res.status_code, 200) + data = res.json() + output = data.get("output", []) + self.assertEqual(len(output), 1) + self.assertEqual(output[0]["name"], "search") + self.assertEqual(output[0]["namespace"], "ns1") + + # Verify internal metadata was not sent to upstream + upstream_req = self.upstream_requests[-1] + upstream_json = json.loads(upstream_req.content.decode("utf-8")) + self.assertNotIn("_tool_registry", upstream_json) + self.assertNotIn("_tool_namespaces", upstream_json) + self.assertFalse(any(k.startswith("_") for k in upstream_json.keys())) + + # Verify upstream tools contains both ns1__search and ns2__search + declared = [t["function"]["name"] for t in upstream_json.get("tools", [])] + self.assertIn("ns1__search", declared) + self.assertIn("ns2__search", declared) + self.assertIn("collaboration__spawn_agent", declared) + self.assertIn("custom_ns__exec__command", declared) + + # 2. Test upstream call to ns2.search + tc_ns2 = [{"id": "call_2", "type": "function", "function": {"name": "ns2__search", "arguments": '{"q":"news"}'}}] + self.mock_response_generator = lambda req: httpx.Response(200, content=_chat_sse_stream(tool_calls=tc_ns2)) + res2 = self.client.post("/v1/responses", json=self._sample_multiagent_payload(stream=False)) + self.assertEqual(res2.status_code, 200) + out2 = res2.json().get("output", []) + self.assertEqual(out2[0]["name"], "search") + self.assertEqual(out2[0]["namespace"], "ns2") + + # 3. Test upstream call to plain tool collaboration__spawn_agent (MUST NOT be altered) + tc_collab = [{"id": "call_3", "type": "function", "function": {"name": "collaboration__spawn_agent", "arguments": '{"task":"inspect"}'}}] + self.mock_response_generator = lambda req: httpx.Response(200, content=_chat_sse_stream(tool_calls=tc_collab)) + res3 = self.client.post("/v1/responses", json=self._sample_multiagent_payload(stream=False)) + self.assertEqual(res3.status_code, 200) + out3 = res3.json().get("output", []) + self.assertEqual(out3[0]["name"], "collaboration__spawn_agent") + self.assertNotIn("namespace", out3[0]) + + # 4. Test upstream call to tool with __ in name + tc_custom = [{"id": "call_4", "type": "function", "function": {"name": "custom_ns__exec__command", "arguments": '{"cmd":"ls"}'}}] + self.mock_response_generator = lambda req: httpx.Response(200, content=_chat_sse_stream(tool_calls=tc_custom)) + res4 = self.client.post("/v1/responses", json=self._sample_multiagent_payload(stream=False)) + self.assertEqual(res4.status_code, 200) + out4 = res4.json().get("output", []) + self.assertEqual(out4[0]["name"], "exec__command") + self.assertEqual(out4[0]["namespace"], "custom_ns") + + def test_buffered_stream_tools_isolation_identity_and_metadata_sanitation(self): + """Cover buffered streaming mode (compatible): verify events and tool identities.""" + tc = [ + {"id": "call_a", "type": "function", "function": {"name": "ns1__search", "arguments": '{"q":"a"}'}}, + {"id": "call_b", "type": "function", "function": {"name": "collaboration__spawn_agent", "arguments": '{"task":"b"}'}}, + {"id": "call_c", "type": "function", "function": {"name": "custom_ns__exec__command", "arguments": '{"cmd":"c"}'}}, + ] + self.mock_response_generator = lambda req: httpx.Response( + 200, + content=_chat_sse_stream(tool_calls=tc), + headers={"Content-Type": "text/event-stream"}, + ) + + with patch.dict(converter.CONFIG, {"stream_mode": "compatible"}): + res = self.client.post("/v1/responses", json=self._sample_multiagent_payload(stream=True)) + self.assertEqual(res.status_code, 200) + events = _parse_responses_sse_events(res.text) + + # Check output_item.done events + done_items = [e["item"] for e in events if e.get("type") == "response.output_item.done" and e.get("item", {}).get("type") == "function_call"] + self.assertEqual(len(done_items), 3) + + # 1. ns1.search + self.assertEqual(done_items[0]["name"], "search") + self.assertEqual(done_items[0]["namespace"], "ns1") + + # 2. collaboration__spawn_agent + self.assertEqual(done_items[1]["name"], "collaboration__spawn_agent") + self.assertNotIn("namespace", done_items[1]) + + # 3. custom_ns.exec__command + self.assertEqual(done_items[2]["name"], "exec__command") + self.assertEqual(done_items[2]["namespace"], "custom_ns") + + # Check response.completed terminal + completed = next(e for e in events if e.get("type") == "response.completed") + outputs = completed["response"]["output"] + self.assertEqual(len(outputs), 3) + self.assertEqual(outputs[0]["name"], "search") + self.assertEqual(outputs[0]["namespace"], "ns1") + self.assertEqual(outputs[1]["name"], "collaboration__spawn_agent") + self.assertNotIn("namespace", outputs[1]) + self.assertEqual(outputs[2]["name"], "exec__command") + self.assertEqual(outputs[2]["namespace"], "custom_ns") + + # Verify no internal metadata sent upstream + upstream_req = self.upstream_requests[-1] + upstream_json = json.loads(upstream_req.content.decode("utf-8")) + self.assertNotIn("_tool_registry", upstream_json) + self.assertNotIn("_tool_namespaces", upstream_json) + self.assertFalse(any(k.startswith("_") for k in upstream_json.keys())) + + def test_realtime_stream_tools_isolation_identity_and_metadata_sanitation(self): + """Cover realtime streaming mode: verify incremental events and tool identities.""" + tc = [ + {"id": "call_1", "type": "function", "function": {"name": "ns2__search", "arguments": '{"q":"live"}'}}, + {"id": "call_2", "type": "function", "function": {"name": "collaboration__spawn_agent", "arguments": '{"task":"live"}'}}, + {"id": "call_3", "type": "function", "function": {"name": "custom_ns__exec__command", "arguments": '{"cmd":"live"}'}}, + ] + self.mock_response_generator = lambda req: httpx.Response( + 200, + content=_chat_sse_stream(tool_calls=tc), + headers={"Content-Type": "text/event-stream"}, + ) + + with patch.dict(converter.CONFIG, {"stream_mode": "realtime"}): + res = self.client.post("/v1/responses", json=self._sample_multiagent_payload(stream=True)) + self.assertEqual(res.status_code, 200) + events = _parse_responses_sse_events(res.text) + + # Check output_item.added events + added_items = [e["item"] for e in events if e.get("type") == "response.output_item.added" and e.get("item", {}).get("type") == "function_call"] + self.assertEqual(len(added_items), 3) + + # 1. ns2.search + self.assertEqual(added_items[0]["name"], "search") + self.assertEqual(added_items[0]["namespace"], "ns2") + + # 2. collaboration__spawn_agent + self.assertEqual(added_items[1]["name"], "collaboration__spawn_agent") + self.assertNotIn("namespace", added_items[1]) + + # 3. custom_ns.exec__command + self.assertEqual(added_items[2]["name"], "exec__command") + self.assertEqual(added_items[2]["namespace"], "custom_ns") + + # Check response.completed terminal + completed = next(e for e in events if e.get("type") == "response.completed") + outputs = completed["response"]["output"] + self.assertEqual(len(outputs), 3) + self.assertEqual(outputs[0]["name"], "search") + self.assertEqual(outputs[0]["namespace"], "ns2") + self.assertEqual(outputs[1]["name"], "collaboration__spawn_agent") + self.assertNotIn("namespace", outputs[1]) + self.assertEqual(outputs[2]["name"], "exec__command") + self.assertEqual(outputs[2]["namespace"], "custom_ns") + + # Verify no internal metadata sent upstream + upstream_req = self.upstream_requests[-1] + upstream_json = json.loads(upstream_req.content.decode("utf-8")) + self.assertNotIn("_tool_registry", upstream_json) + self.assertNotIn("_tool_namespaces", upstream_json) + self.assertFalse(any(k.startswith("_") for k in upstream_json.keys())) + + def test_historical_tool_calls_and_tool_choice_integration(self): + """Verify historical function_call and tool_choice with namespace are correctly mapped to upstream.""" + payload = { + "model": "auto", + "stream": False, + "tools": [ + { + "type": "namespace", + "name": "collaboration", + "tools": [{"type": "function", "name": "spawn_agent", "parameters": {"type": "object"}}], + } + ], + "tool_choice": {"type": "function", "name": "spawn_agent", "namespace": "collaboration"}, + "input": [ + { + "type": "agent_message", + "content": [ + {"type": "input_text", "text": "Task: "}, + {"type": "encrypted_content", "encrypted_content": "Run analysis worker"}, + ], + }, + { + "type": "function_call", + "call_id": "call_hist_1", + "name": "spawn_agent", + "namespace": "collaboration", + "arguments": '{"task_name": "worker_1"}', + }, + { + "type": "function_call_output", + "call_id": "call_hist_1", + "output": "worker spawned", + }, + ], + } + + tc = [{"id": "call_resp_1", "type": "function", "function": {"name": "collaboration__spawn_agent", "arguments": '{"task_name":"worker_2"}'}}] + self.mock_response_generator = lambda req: httpx.Response(200, content=_chat_sse_stream(tool_calls=tc)) + res = self.client.post("/v1/responses", json=payload) + self.assertEqual(res.status_code, 200) + + upstream_req = self.upstream_requests[-1] + upstream_json = json.loads(upstream_req.content.decode("utf-8")) + + # 1. tool_choice mapped + self.assertEqual(upstream_json.get("tool_choice"), "required") + # And tools filtered to that single tool by _normalize_tool_choice + self.assertEqual(upstream_json["tools"][0]["function"]["name"], "collaboration__spawn_agent") + + # 2. Historical tool call in assistant message mapped + messages = upstream_json["messages"] + asst = next(m for m in messages if m["role"] == "assistant" and "tool_calls" in m) + self.assertEqual(asst["tool_calls"][0]["function"]["name"], "collaboration__spawn_agent") + + # 3. Encrypted content extracted into user message + user = next(m for m in messages if m["role"] == "user") + self.assertIn("Task: ", user["content"]) + self.assertIn("Run analysis worker", user["content"]) + + # 4. Response returned by gateway has namespace and name restored + out = res.json().get("output", []) + self.assertEqual(len(out), 1) + self.assertEqual(out[0]["name"], "spawn_agent") + self.assertEqual(out[0]["namespace"], "collaboration") + + + def _boundary_request(self, payload, stream, mode, projection="balanced"): + self.upstream_requests.clear() + + def respond(request): + body = json.loads(request.content) + names = [tool["function"]["name"] for tool in body.get("tools", [])] + names.extend(call["function"]["name"] for message in body["messages"] for call in message.get("tool_calls", [])) + if any(re.fullmatch(r"[A-Za-z0-9_-]{1,64}", name) is None for name in names): + return httpx.Response(400, json={"error": {"message": "Invalid Chat function name", "type": "invalid_request_error"}}) + calls = None + if body.get("tool_choice") == "required" and body.get("tools"): + calls = [{"id": "next_call", "function": { + "name": body["tools"][0]["function"]["name"], "arguments": "{}"}}] + return httpx.Response(200, content=_chat_sse_stream(tool_calls=calls, content=None if calls else "OK")) + + self.mock_response_generator = respond + body = {"input": [{"role": "user", "content": "Use the tool"}], + **deepcopy(payload), "model": "auto", "stream": stream} + with patch.dict(converter.CONFIG, {"stream_mode": mode, "responses_projection_mode": projection, + "keep_tool_metadata": True}): + response = self.client.post("/v1/responses", json=body) + sent = json.loads(self.upstream_requests[-1].content) if self.upstream_requests else None + return response, sent + + def test_schema_cleanup_preserves_names_and_instance_data(self): + literal = {"encrypted": True, "payload": "fixture", + "metadata": {"properties": {"encrypted": {"encrypted": False}}}} + schema = { + "type": "object", "encrypted": True, + "properties": {"encrypted": {"type": "boolean", "encrypted": True}, + "payload": {"type": "string", "encrypted": False}, + "metadata": {"type": "object", "additionalProperties": True}}, + "required": ["encrypted", "payload"], "additionalProperties": False, + "$defs": {"encrypted": {"type": "boolean", "encrypted": True}}, + "allOf": [{"encrypted": True, "properties": {"encrypted": {"$ref": "#/$defs/encrypted"}}}], + "default": literal, "const": literal, "enum": [literal], "examples": [literal], + } + expected = deepcopy(schema) + del expected["encrypted"] + del expected["properties"]["encrypted"]["encrypted"] + del expected["properties"]["payload"]["encrypted"] + del expected["$defs"]["encrypted"]["encrypted"] + del expected["allOf"][0]["encrypted"] + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + for projection in ("balanced", "passthrough"): + with self.subTest(stream=stream, mode=mode, projection=projection): + response, sent = self._boundary_request({"tools": [{ + "type": "function", "name": "save", "parameters": schema}]}, stream, mode, projection) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(sent["tools"][0]["function"]["parameters"], expected) + + def test_schema_cleanup_handles_nested_schema_positions(self): + nested = {"type": "object", "encrypted": True, + "properties": {"encrypted": {"type": "boolean", "encrypted": True}}} + cleaned = {"type": "object", "properties": {"encrypted": {"type": "boolean"}}} + schema = { + "type": "object", "definitions": {"encrypted": nested}, + "patternProperties": {"encrypted": nested}, "dependentSchemas": {"encrypted": nested}, + "dependencies": {"encrypted": ["payload"], "payload": nested}, + "additionalProperties": nested, "propertyNames": nested, "unevaluatedProperties": nested, + "items": nested, "contains": nested, "additionalItems": nested, "unevaluatedItems": nested, + "contentSchema": nested, "not": nested, "if": nested, "then": nested, "else": nested, + "anyOf": [nested], "oneOf": [nested], "prefixItems": [nested], + } + response, sent = self._boundary_request({"tools": [{"type": "function", "function": { + "name": "save", "parameters": schema}}]}, False, "compatible") + self.assertEqual(response.status_code, 200, response.text) + actual = sent["tools"][0]["function"]["parameters"] + self.assertEqual(actual["definitions"]["encrypted"], cleaned) + self.assertEqual(actual["patternProperties"]["encrypted"], cleaned) + self.assertEqual(actual["dependentSchemas"]["encrypted"], cleaned) + self.assertEqual(actual["dependencies"], {"encrypted": ["payload"], "payload": cleaned}) + for key in ("additionalProperties", "propertyNames", "unevaluatedProperties", "items", "contains", + "additionalItems", "unevaluatedItems", "contentSchema", "not", "if", "then", "else"): + self.assertEqual(actual[key], cleaned, key) + for key in ("anyOf", "oneOf", "prefixItems"): + self.assertEqual(actual[key], [cleaned], key) + response, sent = self._boundary_request({"tools": [{"type": "function", "name": "tuple", + "parameters": {"type": "array", "items": [nested, False]}}]}, False, "compatible") + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(sent["tools"][0]["function"]["parameters"]["items"], [cleaned, False]) + + def test_explicit_choices_require_exact_current_identity(self): + function = {"type": "function", "name": "read", "parameters": {"type": "object"}} + namespaced = {"type": "namespace", "name": "vault", "tools": [function]} + cases = [ + ([{**function, "name": "vault__read"}], {"type": "function", "namespace": "vault", "name": "read"}), + ([function], {"type": "custom", "name": "read"}), + ([function], {"name": "read"}), + ([namespaced], {"type": "function", "name": "read"}), + ([namespaced], {"type": "function", "namespace": "missing", "name": "read"}), + ] + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + for tools, choice in cases: + with self.subTest(stream=stream, mode=mode, choice=choice): + response, _ = self._boundary_request({"tools": tools, "tool_choice": choice}, stream, mode) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(self.upstream_requests, []) + + def test_history_collisions_preserve_reasoning_and_agent_messages(self): + cases = [((None, "vault__read"), ("vault", "read")), + (("vault", "read"), (None, "vault__read")), + (("vault", "inner__read"), ("vault__inner", "read"))] + for current, prior in cases: + namespace, name = current + tool = {"type": "function", "name": name, "parameters": {"type": "object"}} + if namespace is not None: + tool = {"type": "namespace", "name": namespace, "tools": [tool]} + history = [ + {"role": "user", "content": "Earlier request"}, + {"type": "reasoning", "summary": [{"type": "summary_text", "text": "prior thought"}]}, + {"type": "function_call", "call_id": "prior_call", "namespace": prior[0], + "name": prior[1], "arguments": "{}"}, + {"type": "function_call_output", "call_id": "prior_call", "output": "prior result"}, + {"type": "additional_tools", "tools": [tool]}, + {"type": "agent_message", "content": [{"type": "encrypted_content", + "encrypted_content": "next task"}]}, + ] + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + for projection in ("balanced", "passthrough"): + with self.subTest(current=current, prior=prior, stream=stream, mode=mode, projection=projection): + payload = {"input": history, "tool_choice": {"type": "function", "name": name, "namespace": namespace}} + response, sent = self._boundary_request(payload, stream, mode, projection) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(len(sent["tools"]), 1) + active_name = sent["tools"][0]["function"]["name"] + previous = next(message for message in sent["messages"] if message.get("tool_calls")) + previous_name = previous["tool_calls"][0]["function"]["name"] + self.assertNotEqual(previous_name, active_name) + if prior[0] is None: + self.assertEqual(previous_name, prior[1]) + self.assertEqual(previous["reasoning_content"], "prior thought") + self.assertEqual(sent["messages"][-1]["content"], "next task") + self.assertFalse(any(key.startswith("_") for key in sent)) + data = (next(event["response"] for event in _parse_responses_sse_events(response.text) + if event["type"] == "response.completed") if stream else response.json()) + call = next(item for item in data["output"] if item["type"] == "function_call") + self.assertEqual((call.get("namespace"), call["name"]), current) + + def test_retired_tools_are_not_exposed_or_selectable(self): + history = [{"role": "user", "content": "Earlier request"}, + {"type": "function_call", "call_id": "past", "name": "read", "namespace": "retired", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "past", "output": "prior result"}, + {"role": "user", "content": "Continue"}] + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + with self.subTest(stream=stream, mode=mode): + response, sent = self._boundary_request({"input": history}, stream, mode) + self.assertEqual(response.status_code, 200, response.text) + self.assertNotIn("tools", sent) + self.assertNotIn("_tool_registry", sent) + response, _ = self._boundary_request({"input": history, "tool_choice": { + "type": "function", "name": "read", "namespace": "retired"}}, stream, mode) + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(self.upstream_requests, []) + + def test_upstream_retired_identity_cannot_become_an_executable_call(self): + payload = {"model": "auto", + "tools": [{"type": "namespace", "name": "probe", "tools": [{"type": "function", + "name": "record", "parameters": {"type": "object"}}]}], + "tool_choice": {"type": "function", "name": "record", "namespace": "probe"}, + "input": [{"role": "user", "content": "Earlier request"}, + {"type": "function_call", "call_id": "past", "name": "probe__record", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "past", "output": "done"}, + {"role": "user", "content": "Use the new tool"}]} + calls = [{"id": "invalid_call", "function": {"name": "probe__record", "arguments": "{}"}}] + self.mock_response_generator = lambda request: httpx.Response(200, content=_chat_sse_stream(tool_calls=calls)) + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + with self.subTest(stream=stream, mode=mode), patch.dict(converter.CONFIG, { + "stream_mode": mode, "tool_call_max_retry": 0}): + response = self.client.post("/v1/responses", json={**payload, "stream": stream}) + if not stream or response.status_code != 200: + self.assertEqual(response.status_code, 502, response.text) + continue + events = _parse_responses_sse_events(response.text) + self.assertTrue(any(event.get("type") in ("error", "response.failed") or event.get("error") for event in events)) + self.assertFalse(any(event.get("type") in ("response.completed", "response.function_call_arguments.done") for event in events)) + self.assertFalse(any(event.get("item", {}).get("type") == "function_call" for event in events)) + + + def test_malformed_namespaces_return_400_without_upstream(self): + for namespace in ([], {}, True, 7, ""): + with self.subTest(namespace=namespace): + response, _ = self._boundary_request({"tools": [{"type": "function", "name": "read"}], + "tool_choice": {"type": "function", "name": "read", "namespace": namespace}}, False, "compatible") + self.assertEqual(response.status_code, 400, response.text) + self.assertEqual(self.upstream_requests, []) + + + + def test_request_budget_counts_wire_bytes_without_tool_registry(self): + payloads = [self._sample_multiagent_payload(), { + "input": [{"role": "user", "content": "Previous task"}, + {"type": "function_call", "call_id": "old", "namespace": "retired", + "name": "lookup", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "old", "output": "done"}, + {"role": "user", "content": "汉字" * 800}]}] + payloads[0]["input"] = [{"role": "user", "content": "汉字" * 800}] + payloads[0]["tool_choice"] = {"type": "function", "namespace": "ns2", "name": "search"} + for index, payload in enumerate(payloads): + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + with self.subTest(payload=index, stream=stream, mode=mode): + response, _ = self._boundary_request(payload, stream, mode) + self.assertEqual(response.status_code, 200, response.text) + wire_size = len(self.upstream_requests[-1].content) + with patch.dict(converter.CONFIG, {"max_request_bytes": wire_size}): + response, sent = self._boundary_request(payload, stream, mode) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(len(self.upstream_requests[-1].content), wire_size) + self.assertNotIn("_tool_registry", sent) + data = (next(event["response"] for event in _parse_responses_sse_events(response.text) + if event["type"] == "response.completed") if stream else response.json()) + if index == 0: + call = next(item for item in data["output"] if item["type"] == "function_call") + self.assertEqual((call.get("namespace"), call["name"]), ("ns2", "search")) + with patch.dict(converter.CONFIG, {"max_request_bytes": wire_size - 1}): + response, _ = self._boundary_request(payload, stream, mode) + self.assertEqual(response.status_code, 413, response.text) + self.assertEqual(self.upstream_requests, []) + + + + def test_request_budget_rechecks_routed_model_without_registry(self): + payload = self._sample_multiagent_payload() + payload["input"] = [{"role": "user", "content": "汉字" * 800}] + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + with self.subTest(stream=stream, mode=mode), patch.object(converter, "_upstream_model", + return_value="routed-" + "model" * 50): + response, _ = self._boundary_request(payload, stream, mode) + self.assertEqual(response.status_code, 200, response.text) + size = len(self.upstream_requests[-1].content) + with patch.dict(converter.CONFIG, {"max_request_bytes": size}): + response, sent = self._boundary_request(payload, stream, mode) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(len(self.upstream_requests[-1].content), size) + self.assertNotIn("_tool_registry", sent) + with patch.dict(converter.CONFIG, {"max_request_bytes": size - 1}): + response, _ = self._boundary_request(payload, stream, mode) + self.assertEqual(response.status_code, 413, response.text) + self.assertEqual(self.upstream_requests, []) + + + + def test_nested_and_long_namespace_aliases_round_trip(self): + cases = [(["parent", "child"], "tool"), (["space group", "中文"], "run-command"), + (["n" * 40], "f" * 40), (["n" * 31], "f" * 31)] + for parts, name in cases: + namespace = ".".join(parts) + tools = [{"type": "function", "name": name, "parameters": {"type": "object"}}] + for part in reversed(parts): + tools = [{"type": "namespace", "name": part, "tools": tools}] + payload = {"tool_choice": {"type": "function", "namespace": namespace, "name": name}, + "input": [{"role": "user", "content": "Previous task"}, + {"type": "function_call", "namespace": namespace, "name": name, + "call_id": "past", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "past", "output": "done"}, + {"type": "additional_tools", "tools": tools}, + {"role": "user", "content": "Call the tool again"}]} + original = deepcopy(payload) + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + with self.subTest(namespace=namespace, name=name, stream=stream, mode=mode): + request = deepcopy(payload) + for turn in range(2): + response, sent = self._boundary_request(request, stream, mode) + self.assertEqual(response.status_code, 200, response.text) + alias = sent["tools"][0]["function"]["name"] + self.assertRegex(alias, r"\A[A-Za-z0-9_-]{1,64}\Z") + historical = [call for message in sent["messages"] for call in message.get("tool_calls", [])] + self.assertEqual(len(historical), turn + 1) + self.assertTrue(all(call["function"]["name"] == alias for call in historical)) + self.assertEqual(sent["tool_choice"], "required") + self.assertNotIn("_tool_registry", sent) + data = (next(event["response"] for event in _parse_responses_sse_events(response.text) + if event["type"] == "response.completed") if stream else response.json()) + calls = [item for item in data["output"] if item["type"] == "function_call"] + self.assertEqual(len(calls), 1) + self.assertEqual((calls[0].get("namespace"), calls[0]["name"]), (namespace, name)) + request["input"].extend([*data["output"], {"type": "function_call_output", + "call_id": calls[0]["call_id"], "output": "done"}, {"role": "user", "content": "Again"}]) + self.assertEqual(payload, original) + + def test_alias_collisions_remain_bounded_and_preserve_global_names(self): + from app.adapters.responses_adapter import ToolRegistry + for namespace, name in (("n" * 31, "f" * 31), ("parent.child" * 8, "function" * 8)): + base = ToolRegistry().register(namespace, name) + reserved = [base, *(base[:64 - len(f"_{index}")] + f"_{index}" for index in range(1, 13))] + for order in (reserved, list(reversed(reserved))): + with self.subTest(namespace=namespace, order=order): + registry = ToolRegistry() + for global_name in order: + self.assertEqual(registry.register(None, global_name), global_name) + alias = registry.register(namespace, name) + self.assertRegex(alias, r"\A[A-Za-z0-9_-]{1,64}\Z") + self.assertNotIn(alias, reserved) + self.assertEqual(registry.get_identity(alias), (namespace, name)) + self.assertEqual(registry.register(namespace, name), alias) + restored = ToolRegistry.from_dict(registry.to_dict()) + self.assertEqual(restored.get_upstream_name(namespace, name), alias) + for global_name in reserved: + self.assertEqual(restored.get_identity(global_name), (None, global_name)) + + def test_encoded_aliases_distinguish_normalized_and_truncated_identities(self): + from app.adapters.responses_adapter import ToolRegistry + identities = [("parent.child", "tool"), ("parent_child", "tool"), ("parent/child", "tool"), + ("n" * 70 + "first", "tool"), ("n" * 70 + "second", "tool"), + ("n" * 70, "__tool"), ("n" * 70 + "__", "tool")] + registry = ToolRegistry() + aliases = [registry.register(*identity) for identity in identities] + self.assertEqual(len(set(aliases)), len(identities)) + for identity, alias in zip(identities, aliases): + self.assertRegex(alias, r"\A[A-Za-z0-9_-]{1,64}\Z") + self.assertEqual(ToolRegistry().register(*identity), alias) + self.assertEqual(registry.get_identity(alias), identity) + + + def test_bounded_aliases_do_not_reuse_retired_global_names(self): + from app.adapters.responses_adapter import ToolRegistry + for namespace, name in (("n" * 31, "f" * 31), ("parent.child" * 8, "tool")): + global_name = ToolRegistry().register(namespace, name) + payload = {"tools": [{"type": "namespace", "name": namespace, "tools": [{"type": "function", "name": name}]}], + "tool_choice": {"type": "function", "namespace": namespace, "name": name}, + "input": [{"role": "user", "content": "Previous task"}, + {"type": "function_call", "name": global_name, "call_id": "past", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "past", "output": "done"}, + {"role": "user", "content": "Use the namespaced tool"}]} + for stream, mode in ((False, "compatible"), (True, "compatible"), (True, "realtime")): + with self.subTest(namespace=namespace, stream=stream, mode=mode): + response, sent = self._boundary_request(payload, stream, mode) + self.assertEqual(response.status_code, 200, response.text) + alias = sent["tools"][0]["function"]["name"] + self.assertRegex(alias, r"\A[A-Za-z0-9_-]{1,64}\Z") + self.assertNotEqual(alias, global_name) + history = next(message["tool_calls"] for message in sent["messages"] if message.get("tool_calls")) + self.assertEqual(history[0]["function"]["name"], global_name) + data = (next(event["response"] for event in _parse_responses_sse_events(response.text) + if event["type"] == "response.completed") if stream else response.json()) + call = next(item for item in data["output"] if item["type"] == "function_call") + self.assertEqual((call.get("namespace"), call["name"]), (namespace, name)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime_endpoints.py b/tests/test_runtime_endpoints.py index 91a237b..000ad56 100644 --- a/tests/test_runtime_endpoints.py +++ b/tests/test_runtime_endpoints.py @@ -384,6 +384,25 @@ def test_request_byte_budget_uses_actual_utf8_encoding(self): converter._guard_request_size(body) self.assertEqual(caught.exception.status_code, 413) + def test_request_byte_budget_excludes_only_top_level_local_metadata(self): + public = {"messages": [{"role": "user", "content": "汉字"}], + "tools": [{"type": "function", "function": {"name": "echo", "parameters": { + "type": "object", "properties": {"_tool_registry": {"type": "string", "default": "汉字" * 200}}}}}]} + registry = {"internal": "x" * 10000} + local_policy = object() + body = {**public, "_tool_registry": registry, "_runtime_note": object(), + converter._REQUEST_POLICY_KEY: local_policy} + size = len(json.dumps(public, ensure_ascii=False, separators=(",", ":"), allow_nan=False).encode("utf-8")) + converter.CONFIG["max_request_bytes"] = size + self.assertEqual(converter._guard_request_size(body), size) + self.assertIs(body["_tool_registry"], registry) + self.assertIs(body[converter._REQUEST_POLICY_KEY], local_policy) + converter.CONFIG["max_request_bytes"] = size - 1 + with self.assertRaises(converter.HTTPException) as caught: + converter._guard_request_size(body) + self.assertEqual(caught.exception.status_code, 413) + + def test_bad_payloads_and_auth_fail_without_upstream(self): for route in ROUTES: for value in (None, [], "not an object"):