diff --git a/verifiers/v1/clients/train.py b/verifiers/v1/clients/train.py index e803610231..a4a4a9dfee 100644 --- a/verifiers/v1/clients/train.py +++ b/verifiers/v1/clients/train.py @@ -17,7 +17,7 @@ from verifiers.v1.clients.client import SESSION_ID_HEADER, Client from verifiers.v1.configs.client import TrainClientConfig from verifiers.v1.dialects import FINISH_REASONS, ChatDialect, Dialect, parse_tools -from verifiers.v1.dialects.chat import message_to_wire +from verifiers.v1.dialects.chat import message_to_wire, tool_calls_to_wire from verifiers.v1.errors import ProviderError, model_error from verifiers.v1.graph import PendingTurn from verifiers.v1.types import ( @@ -55,17 +55,7 @@ def serialize_completion(response: Response, model: str) -> dict: if response.message.reasoning_content is not None: message["reasoning_content"] = response.message.reasoning_content if response.message.tool_calls: - message["tool_calls"] = [ - { - "id": c.id, - "type": c.type, - c.type: { - "name": c.name, - "input" if c.type == "custom" else "arguments": c.arguments, - }, - } - for c in response.message.tool_calls - ] + message["tool_calls"] = tool_calls_to_wire(response.message.tool_calls) usage: dict | None = None if response.usage: # Usage is validated earlier in the pipeline; building its wire dict directly saves time. diff --git a/verifiers/v1/dialects/anthropic.py b/verifiers/v1/dialects/anthropic.py index 40324110f6..ad2e1f2725 100644 --- a/verifiers/v1/dialects/anthropic.py +++ b/verifiers/v1/dialects/anthropic.py @@ -20,9 +20,12 @@ RawRequest, StreamParser, append_user_notice, + blocked_path, blocked_url, + mediate_parts, parse_sse_event, provider_allowed_domains, + user_and_tool_messages, ) from verifiers.v1.types import ( AssistantMessage, @@ -129,20 +132,11 @@ def parse_content(content) -> str | list[ContentPart]: def blocked_content_path(value, path: str, policy: NetworkPolicyConfig) -> str | None: - if isinstance(value, list): - for index, item in enumerate(value): - if blocked := blocked_content_path(item, f"{path}[{index}]", policy): - return blocked - return None - if not isinstance(value, dict): - return None + return blocked_path(value, path, policy, _blocked_block) + +def _blocked_block(value: dict, path: str, policy: NetworkPolicyConfig) -> str | None: kind = value.get("type") - caller = value.get("caller") - if caller is not None and not ( - isinstance(caller, dict) and caller.get("type") == "direct" - ): - return f"{path}.caller.type" if kind in ("image", "document"): source_path = f"{path}.source" source = value.get("source") or {} @@ -178,31 +172,7 @@ def blocked_content_path(value, path: str, policy: NetworkPolicyConfig) -> str | def mediate_content(value, path: str, policy: NetworkPolicyConfig): - if not isinstance(value, list): - blocked = blocked_content_path(value, path, policy) - return ("", [blocked]) if blocked else (value, []) - - mediated = [] - capabilities = [] - for index, block in enumerate(value): - item_path = f"{path}[{index}]" - if isinstance(block, dict) and block.get("type") in _CONTENT_WRAPPERS: - if blocked := blocked_content_path( - {**block, "content": []}, item_path, policy - ): - capabilities.append(blocked) - continue - content, removed = mediate_content( - block.get("content"), f"{item_path}.content", policy - ) - if removed: - block["content"] = content or "" - capabilities.extend(removed) - elif blocked := blocked_content_path(block, item_path, policy): - capabilities.append(blocked) - continue - mediated.append(block) - return mediated, capabilities + return mediate_parts(value, path, policy, blocked_content_path, _CONTENT_WRAPPERS) def content_to_wire(content) -> str | list[dict]: @@ -577,12 +547,8 @@ def parse_response(self, response: AnthropicMessage) -> Response: return response_from_wire(response) def rewrite_request(self, body: dict, before: Request, after: Request) -> None: - original = [ - m for m in before.messages if isinstance(m, (UserMessage, ToolMessage)) - ] - rewritten = [ - m for m in after.messages if isinstance(m, (UserMessage, ToolMessage)) - ] + original = user_and_tool_messages(before) + rewritten = user_and_tool_messages(after) targets: list[tuple[dict, dict | None]] = [] for native in body.get("messages", []): if native.get("role") == "assistant": diff --git a/verifiers/v1/dialects/base.py b/verifiers/v1/dialects/base.py index ed223916f9..bd860963ff 100644 --- a/verifiers/v1/dialects/base.py +++ b/verifiers/v1/dialects/base.py @@ -22,7 +22,14 @@ from pydantic_core import from_json from verifiers.v1.configs.runtime import NetworkPolicyConfig -from verifiers.v1.types import Request, Response, Sampling, SamplingConfig +from verifiers.v1.types import ( + Request, + Response, + Sampling, + SamplingConfig, + ToolMessage, + UserMessage, +) RespT = TypeVar("RespT", bound=BaseModel) RawRequest = dict[str, Any] @@ -49,6 +56,70 @@ def blocked_url(value: str, policy: NetworkPolicyConfig) -> bool: return not policy.permits(url.scheme, host, url.port) +def blocked_path( + value, + path: str, + policy: NetworkPolicyConfig, + blocked_item: Callable[[dict, str, NetworkPolicyConfig], str | None], +) -> str | None: + """The first policy-blocked path under `value`, or None. Lists recurse per index, + non-dicts pass, a non-`direct` `caller` is blocked in every format, and `blocked_item` + applies the format's own rules to each dict.""" + if isinstance(value, list): + for index, item in enumerate(value): + if blocked := blocked_path(item, f"{path}[{index}]", policy, blocked_item): + return blocked + return None + if not isinstance(value, dict): + return None + caller = value.get("caller") + if caller is not None and not ( + isinstance(caller, dict) and caller.get("type") == "direct" + ): + return f"{path}.caller.type" + return blocked_item(value, path, policy) + + +def mediate_parts( + value, + path: str, + policy: NetworkPolicyConfig, + blocked: Callable[[Any, str, NetworkPolicyConfig], str | None], + wrappers: tuple[str, ...] = (), +) -> tuple[Any, list[str]]: + """Drop the policy-blocked parts of a content list and report their paths; a non-list + value is kept whole or replaced by "". A part whose type is in `wrappers` is checked + without its content, then its content is mediated in place.""" + if not isinstance(value, list): + blocked_at = blocked(value, path, policy) + return ("", [blocked_at]) if blocked_at else (value, []) + + mediated = [] + capabilities = [] + for index, part in enumerate(value): + item_path = f"{path}[{index}]" + if isinstance(part, dict) and part.get("type") in wrappers: + if blocked_at := blocked({**part, "content": []}, item_path, policy): + capabilities.append(blocked_at) + continue + content, removed = mediate_parts( + part.get("content"), f"{item_path}.content", policy, blocked, wrappers + ) + if removed: + part["content"] = content or "" + capabilities.extend(removed) + elif blocked_at := blocked(part, item_path, policy): + capabilities.append(blocked_at) + continue + mediated.append(part) + return mediated, capabilities + + +def user_and_tool_messages(request: Request) -> list[UserMessage | ToolMessage]: + """The rewritable messages of a request, in order.""" + return [m for m in request.messages if isinstance(m, (UserMessage, ToolMessage))] + + def provider_allowed_domains( policy: NetworkPolicyConfig, requested: object = None ) -> list[str]: diff --git a/verifiers/v1/dialects/chat.py b/verifiers/v1/dialects/chat.py index 3add706a56..0272671326 100644 --- a/verifiers/v1/dialects/chat.py +++ b/verifiers/v1/dialects/chat.py @@ -159,6 +159,20 @@ def _content_to_wire(content): return [part.model_dump() for part in content] +def tool_calls_to_wire(tool_calls: list[ToolCall]) -> list[dict]: + return [ + { + "id": call.id, + "type": call.type, + call.type: { + "name": call.name, + "input" if call.type == "custom" else "arguments": call.arguments, + }, + } + for call in tool_calls + ] + + def message_to_wire(message: Message) -> dict: if message.role == "assistant": # Strict providers reject `content: null` without tool calls. @@ -171,19 +185,7 @@ def message_to_wire(message: Message) -> dict: elif message.reasoning_content is not None: wire["reasoning_content"] = message.reasoning_content if message.tool_calls: - wire["tool_calls"] = [ - { - "id": call.id, - "type": call.type, - call.type: { - "name": call.name, - "input" - if call.type == "custom" - else "arguments": call.arguments, - }, - } - for call in message.tool_calls - ] + wire["tool_calls"] = tool_calls_to_wire(message.tool_calls) return wire if message.role == "tool": wire = { diff --git a/verifiers/v1/dialects/responses.py b/verifiers/v1/dialects/responses.py index c558318e8c..d59f2f2308 100644 --- a/verifiers/v1/dialects/responses.py +++ b/verifiers/v1/dialects/responses.py @@ -24,9 +24,12 @@ RawRequest, StreamParser, append_user_notice, + blocked_path, blocked_url, + mediate_parts, parse_sse_event, provider_allowed_domains, + user_and_tool_messages, ) from verifiers.v1.errors import model_error from verifiers.v1.types import ( @@ -145,6 +148,18 @@ def parse_content(content) -> str | list[ContentPart]: return parts +def _content_to_input(content) -> str | list[dict]: + """Typed text/image content in Responses' native input shape.""" + if isinstance(content, str): + return content + return [ + {"type": "input_text", "text": part.text} + if isinstance(part, TextContentPart) + else {"type": "input_image", "image_url": part.image_url.url} + for part in content + ] + + def mediate_tools( tools, path: str, policy: NetworkPolicyConfig ) -> tuple[list[dict], list[str]]: @@ -199,20 +214,11 @@ def mediate_tools( def blocked_content_path(value, path: str, policy: NetworkPolicyConfig) -> str | None: - if isinstance(value, list): - for index, item in enumerate(value): - if blocked := blocked_content_path(item, f"{path}[{index}]", policy): - return blocked - return None - if not isinstance(value, dict): - return None + return blocked_path(value, path, policy, _blocked_item) + +def _blocked_item(value: dict, path: str, policy: NetworkPolicyConfig) -> str | None: kind = value.get("type") - caller = value.get("caller") - if caller is not None and not ( - isinstance(caller, dict) and caller.get("type") == "direct" - ): - return f"{path}.caller.type" if kind == "input_file": if value.get("file_id"): return f"{path}.file_id" @@ -260,18 +266,7 @@ def blocked_content_path(value, path: str, policy: NetworkPolicyConfig) -> str | def mediate_content(value, path: str, policy: NetworkPolicyConfig): - if not isinstance(value, list): - blocked = blocked_content_path(value, path, policy) - return ("", [blocked]) if blocked else (value, []) - - mediated = [] - capabilities = [] - for index, part in enumerate(value): - if blocked := blocked_content_path(part, f"{path}[{index}]", policy): - capabilities.append(blocked) - continue - mediated.append(part) - return mediated, capabilities + return mediate_parts(value, path, policy, blocked_content_path) def fold_assistant(items: list[dict] | None) -> AssistantMessage: @@ -602,29 +597,12 @@ def parse_response(self, response: OpenAIResponse) -> Response: return response_from_wire(response) def rewrite_request(self, body: dict, before: Request, after: Request) -> None: - original = [ - m for m in before.messages if isinstance(m, (UserMessage, ToolMessage)) - ] - rewritten = [ - m for m in after.messages if isinstance(m, (UserMessage, ToolMessage)) - ] + original = user_and_tool_messages(before) + rewritten = user_and_tool_messages(after) items = body.get("input") if isinstance(items, str): if original != rewritten: - message = rewritten[0] - content = ( - message.content - if isinstance(message.content, str) - else [ - {"type": "input_text", "text": part.text} - if isinstance(part, TextContentPart) - else { - "type": "input_image", - "image_url": part.image_url.url, - } - for part in message.content - ] - ) + content = _content_to_input(rewritten[0].content) body["input"] = ( content if isinstance(content, str) @@ -645,19 +623,7 @@ def rewrite_request(self, body: dict, before: Request, after: Request) -> None: for item, old, new in zip(targets, original, rewritten, strict=True): if old == new: continue - content = ( - new.content - if isinstance(new.content, str) - else [ - {"type": "input_text", "text": part.text} - if isinstance(part, TextContentPart) - else { - "type": "input_image", - "image_url": part.image_url.url, - } - for part in new.content - ] - ) + content = _content_to_input(new.content) item["output" if isinstance(new, ToolMessage) else "content"] = content def rewrite_response(self, raw: dict, text: str) -> None: