Skip to content
Merged
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
96 changes: 62 additions & 34 deletions src/ai/models/anthropic/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,16 +13,10 @@

from ... import types
from ...types import events
from ...types import messages as messages_
from .. import core
from . import metadata as anthropic_metadata
from . import params
from . import tools as anthropic_tools
from .params import (
AnthropicContainer,
AnthropicCustomSkill,
AnthropicDisabledThinking,
AnthropicParams,
AnthropicProviderSkill,
)

PROVIDER_NAME = "anthropic"

Expand All @@ -38,6 +32,7 @@
}
)


# ---------------------------------------------------------------------------
# Message / tool conversion — internal types → Anthropic wire format
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -201,8 +196,22 @@ async def _messages_to_anthropic(
for part in msg.parts:
match part:
case types.messages.ReasoningPart(
text=text, signature=signature
text=text,
provider_metadata=provider_metadata,
):
part_metadata = (
provider_metadata
if isinstance(
provider_metadata,
anthropic_metadata.AnthropicProviderMetadata,
)
else None
)
signature = (
part_metadata.signature
if part_metadata is not None
else None
)
if send_reasoning and signature:
content.append(
{
Expand Down Expand Up @@ -239,12 +248,21 @@ async def _messages_to_anthropic(
)
case types.messages.BuiltinToolReturnPart():
# Result block type comes from the original wire
# event ("web_search_tool_result", etc.); stored
# in provider_details when emitted.
details = part.provider_details or {}
wire_type = details.get(
"result_type",
f"{part.tool_name}_tool_result",
# event ("web_search_tool_result", etc.); stored in
# provider metadata when emitted.
part_metadata = (
part.provider_metadata
if isinstance(
part.provider_metadata,
anthropic_metadata.AnthropicProviderMetadata,
)
else None
)
wire_type = (
part_metadata.result_type
if part_metadata is not None
and part_metadata.result_type is not None
else f"{part.tool_name}_tool_result"
)
content.append(
{
Expand Down Expand Up @@ -343,19 +361,23 @@ def _make_client(
)


def _coerce_anthropic_params(value: Any) -> AnthropicParams:
def _coerce_anthropic_params(value: Any) -> params.AnthropicParams:
if value is None:
return AnthropicParams()
if isinstance(value, AnthropicParams):
return params.AnthropicParams()
if isinstance(value, params.AnthropicParams):
return value
if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray):
items = list(value)
if len(items) == 1 and isinstance(items[0], AnthropicParams):
if len(items) == 1 and isinstance(items[0], params.AnthropicParams):
return items[0]
raise TypeError(f"anthropic stream params must be {AnthropicParams.__name__}")
raise TypeError(
f"anthropic stream params must be {params.AnthropicParams.__name__}"
)


def _container_to_wire(container: AnthropicContainer) -> str | dict[str, Any] | None:
def _container_to_wire(
container: params.AnthropicContainer,
) -> str | dict[str, Any] | None:
if not container.skills:
return container.id

Expand All @@ -367,9 +389,9 @@ def _container_to_wire(container: AnthropicContainer) -> str | dict[str, Any] |


def _skill_to_wire(
skill: AnthropicProviderSkill | AnthropicCustomSkill,
skill: params.AnthropicProviderSkill | params.AnthropicCustomSkill,
) -> dict[str, Any]:
if isinstance(skill, AnthropicProviderSkill):
if isinstance(skill, params.AnthropicProviderSkill):
result: dict[str, Any] = {
"type": "anthropic",
"skill_id": skill.skill_id,
Expand Down Expand Up @@ -466,7 +488,7 @@ async def stream(
api_kwargs["tools"] = wire_tools

if anthropic_params.thinking is not None:
if not isinstance(anthropic_params.thinking, AnthropicDisabledThinking):
if not isinstance(anthropic_params.thinking, params.AnthropicDisabledThinking):
api_kwargs["thinking"] = anthropic_params.thinking.model_dump(
exclude_none=True,
)
Expand Down Expand Up @@ -569,7 +591,7 @@ async def stream(
yield events.BuiltinToolStart(
tool_call_id=block.id,
tool_name=block.name,
provider_name=PROVIDER_NAME,
provider_metadata=anthropic_metadata.AnthropicProviderMetadata(),
)
# Result blocks (web_search_tool_result etc.) arrive
# complete; we emit on stop so we have full content.
Expand Down Expand Up @@ -614,26 +636,33 @@ async def stream(
if block_type == "text":
yield events.TextEnd(block_id=str(idx))
elif block_type == "thinking":
signature = signature_buffer.get(idx)
yield events.ReasoningEnd(
block_id=str(idx),
signature=signature_buffer.get(idx),
provider_metadata=(
anthropic_metadata.AnthropicProviderMetadata(
signature=signature
)
if signature is not None
else None
),
)
elif block_type == "tool_use":
tool_id = tool_ids.get(idx)
if tool_id:
yield events.ToolEnd(
tool_call_id=tool_id,
tool_call=messages_.DUMMY_TOOL_CALL,
tool_call=types.messages.DUMMY_TOOL_CALL,
)
elif block_type == "server_tool_use":
tool_id = tool_ids.get(idx)
if tool_id:
yield events.BuiltinToolEnd(
tool_call_id=tool_id,
tool_call=messages_.BuiltinToolCallPart(
tool_call=types.messages.BuiltinToolCallPart(
tool_call_id=tool_id,
tool_name=tool_names.get(idx, ""),
provider_name=PROVIDER_NAME,
provider_metadata=anthropic_metadata.AnthropicProviderMetadata(),
),
)
elif block_type in _TOOL_RESULT_BLOCK_TYPES:
Expand Down Expand Up @@ -662,14 +691,13 @@ async def stream(
break
yield events.BuiltinToolResult(
tool_call_id=tool_use_id,
result=messages_.BuiltinToolReturnPart(
result=types.messages.BuiltinToolReturnPart(
tool_call_id=tool_use_id,
tool_name=tool_name,
result=content_payload,
provider_name=PROVIDER_NAME,
provider_details={
"result_type": block_type or "",
},
provider_metadata=anthropic_metadata.AnthropicProviderMetadata(
result_type=block_type or ""
),
),
)

Expand Down
167 changes: 167 additions & 0 deletions src/ai/models/anthropic/metadata.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,167 @@
from __future__ import annotations

from typing import Annotated, Any, Literal

import pydantic

from ... import types
from . import params

_METADATA_CONFIG = pydantic.ConfigDict(frozen=True, populate_by_name=True)


class AnthropicCitations(pydantic.BaseModel):
model_config = _METADATA_CONFIG

enabled: bool


class AnthropicCaller(pydantic.BaseModel):
model_config = _METADATA_CONFIG

type: str
tool_id: str | None = pydantic.Field(
default=None,
validation_alias="toolId",
serialization_alias="toolId",
)


class AnthropicUsageIteration(pydantic.BaseModel):
model_config = _METADATA_CONFIG

type: Literal["compaction", "message"]
input_tokens: int = pydantic.Field(
validation_alias="inputTokens",
serialization_alias="inputTokens",
)
output_tokens: int = pydantic.Field(
validation_alias="outputTokens",
serialization_alias="outputTokens",
)


class AnthropicContainerSkill(pydantic.BaseModel):
model_config = _METADATA_CONFIG

type: Literal["anthropic", "custom"]
skill_id: str = pydantic.Field(
validation_alias="skillId",
serialization_alias="skillId",
)
version: str


class AnthropicContainerMetadata(pydantic.BaseModel):
model_config = _METADATA_CONFIG

expires_at: str = pydantic.Field(
validation_alias="expiresAt",
serialization_alias="expiresAt",
)
id: str
skills: list[AnthropicContainerSkill] | None = None


class AnthropicClearToolUsesEdit(pydantic.BaseModel):
model_config = _METADATA_CONFIG

type: Literal["clear_tool_uses_20250919"]
cleared_tool_uses: int = pydantic.Field(
validation_alias="clearedToolUses",
serialization_alias="clearedToolUses",
)
cleared_input_tokens: int = pydantic.Field(
validation_alias="clearedInputTokens",
serialization_alias="clearedInputTokens",
)


class AnthropicClearThinkingEdit(pydantic.BaseModel):
model_config = _METADATA_CONFIG

type: Literal["clear_thinking_20251015"]
cleared_thinking_turns: int = pydantic.Field(
validation_alias="clearedThinkingTurns",
serialization_alias="clearedThinkingTurns",
)
cleared_input_tokens: int = pydantic.Field(
validation_alias="clearedInputTokens",
serialization_alias="clearedInputTokens",
)


class AnthropicCompactEdit(pydantic.BaseModel):
model_config = _METADATA_CONFIG

type: Literal["compact_20260112"]


type AnthropicContextManagementEdit = Annotated[
AnthropicClearToolUsesEdit | AnthropicClearThinkingEdit | AnthropicCompactEdit,
pydantic.Field(discriminator="type"),
]


class AnthropicContextManagementMetadata(pydantic.BaseModel):
model_config = _METADATA_CONFIG

applied_edits: list[AnthropicContextManagementEdit] = pydantic.Field(
validation_alias="appliedEdits",
serialization_alias="appliedEdits",
)


class AnthropicProviderMetadata(types.metadata.ProviderMetadata):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is it worth putting all this in a class, instead of handling a json blob or something? How does it get provided to us by the anthropic SDK?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

That's an excellent question, actually

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The Anthropic SDK does provide types, and then we map from those to ours. Whether or not it's worth it, I'm not sure: on one hand, having a reasoning.signarture is nice, on the other, you as a user will probably never interact with it.

"""Anthropic-specific metadata for messages, parts, and stream events."""

provider: Literal["anthropic"] = "anthropic"

# Input/replay options.
cache_control: params.AnthropicCacheControl | None = pydantic.Field(
default=None,
validation_alias="cacheControl",
serialization_alias="cacheControl",
)
citations: AnthropicCitations | None = None
title: str | None = None
context: str | None = None

# Special part markers.
type: Literal["compaction", "mcp-tool-use"] | None = None
server_name: str | None = pydantic.Field(
default=None,
validation_alias="serverName",
serialization_alias="serverName",
)

# Reasoning replay metadata.
signature: str | None = None
redacted_data: str | None = pydantic.Field(
default=None,
validation_alias="redactedData",
serialization_alias="redactedData",
)

# Tool replay metadata.
result_type: str | None = pydantic.Field(
default=None,
validation_alias="resultType",
serialization_alias="resultType",
)
caller: AnthropicCaller | dict[str, Any] | None = None

# Response/message metadata.
usage: dict[str, Any] | None = None
stop_sequence: str | None = pydantic.Field(
default=None,
validation_alias="stopSequence",
serialization_alias="stopSequence",
)
iterations: list[AnthropicUsageIteration] | None = None
container: AnthropicContainerMetadata | None = None
context_management: AnthropicContextManagementMetadata | None = pydantic.Field(
default=None,
validation_alias="contextManagement",
serialization_alias="contextManagement",
)
Loading
Loading