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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 2 additions & 12 deletions verifiers/v1/clients/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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.
Expand Down
52 changes: 9 additions & 43 deletions verifiers/v1/dialects/anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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 {}
Expand Down Expand Up @@ -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]:
Expand Down Expand Up @@ -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":
Expand Down
73 changes: 72 additions & 1 deletion verifiers/v1/dialects/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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]:
Expand Down
28 changes: 15 additions & 13 deletions verifiers/v1/dialects/chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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 = {
Expand Down
80 changes: 23 additions & 57 deletions verifiers/v1/dialects/responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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]]:
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand All @@ -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:
Expand Down