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
6 changes: 3 additions & 3 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ dependencies = [
"pydantic>2,<3",
"openai>=2.18.0",
"openresponses-types>=2.4.0,<3",
"anthropic>=0.119.0",
"anthropic>=0.119.0,<1",
"rich",
"httpx",
"typing_extensions>=4.5.0",
Expand Down Expand Up @@ -49,11 +49,11 @@ vertexai = [
]

vertexaianthropic = [
"anthropic[vertex]>=0.119.0",
"anthropic[vertex]>=0.119.0,<1",
]

azureanthropic = [
"anthropic>=0.119.0",
"anthropic>=0.119.0,<1",
]

huggingface = [
Expand Down
154 changes: 117 additions & 37 deletions tests/unit/providers/test_anthropic_messages.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
from typing import Any, Self, cast
from unittest.mock import AsyncMock, MagicMock, Mock, patch

import httpx2
import httpx
import pytest
from anthropic.types import Message, TextBlock, ThinkingBlock, ToolUseBlock, Usage
from anthropic.types.beta import BetaMCPToolUseBlock, BetaMessage, BetaThinkingBlock, BetaUsage
Expand All @@ -16,6 +16,7 @@
from any_llm.exceptions import InvalidRequestError, UnsupportedParameterError
from any_llm.providers.anthropic.anthropic import AnthropicProvider
from any_llm.providers.anthropic.base import BaseAnthropicProvider, _messages_betas, _pop_anthropic_beta_header
from any_llm.types.completion import CompletionParams
from any_llm.types.messages import (
CompactionDelta,
ContentBlockDeltaEvent,
Expand Down Expand Up @@ -66,6 +67,19 @@ def _make_message(**overrides: Any) -> Message:
return Message(**defaults)


def _sdk_message_response() -> dict[str, Any]:
return {
"id": "msg_test123",
"type": "message",
"role": "assistant",
"model": "claude-3-5-sonnet",
"stop_reason": "end_turn",
"stop_sequence": None,
"content": [{"type": "text", "text": "Hello!"}],
"usage": {"input_tokens": 10, "output_tokens": 5},
}


def test_convert_native_message_to_response_text() -> None:
"""Test converting an Anthropic Message with text content."""
msg = _make_message(content=[TextBlock(type="text", text="Hello!")])
Expand Down Expand Up @@ -260,15 +274,74 @@ async def test_amessages_non_streaming() -> None:
assert call_kwargs["container"] == "container_123"


@pytest.mark.asyncio
async def test_anthropic_sdk_accepts_completion_sampling_parameters() -> None:
requests: list[httpx.Request] = []

async def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_sdk_message_response())

async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as http_client:
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
await provider._acompletion(
CompletionParams(
model_id="claude-3-5-sonnet",
messages=[{"role": "user", "content": "Hello"}],
max_tokens=1024,
temperature=0.7,
top_p=0.9,
)
)

assert len(requests) == 1
request_body = json.loads(requests[0].content)
assert request_body["temperature"] == 0.7
assert request_body["top_p"] == 0.9


@pytest.mark.asyncio
async def test_anthropic_sdk_accepts_native_messages_parameters() -> None:
requests: list[httpx.Request] = []

async def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_sdk_message_response())

async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as http_client:
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
result = await provider._amessages(
MessagesParams(
model="claude-3-5-sonnet",
messages=[{"role": "user", "content": "Hello"}],
max_tokens=1024,
temperature=0.7,
top_p=0.9,
top_k=40,
container="container_123",
service_tier="standard_only",
)
)

assert isinstance(result, MessageResponse)
assert len(requests) == 1
request_body = json.loads(requests[0].content)
assert request_body["temperature"] == 0.7
assert request_body["top_p"] == 0.9
assert request_body["top_k"] == 40
assert request_body["container"] == "container_123"
assert request_body["service_tier"] == "standard_only"


@pytest.mark.asyncio
async def test_amessages_rejects_prompt_cache_key_before_client_call() -> None:
requests: list[httpx2.Request] = []
requests: list[httpx.Request] = []

async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx2.Response(500)
return httpx.Response(500)

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
try:
with pytest.raises(UnsupportedParameterError, match="prompt_cache_key"):
Expand All @@ -286,13 +359,13 @@ async def handler(request: httpx2.Request) -> httpx2.Response:

@pytest.mark.asyncio
async def test_acompletion_rejects_prompt_cache_key_before_client_call() -> None:
requests: list[httpx2.Request] = []
requests: list[httpx.Request] = []

async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx2.Response(500)
return httpx.Response(500)

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
try:
with pytest.raises(UnsupportedParameterError, match="prompt_cache_key"):
Expand All @@ -309,11 +382,11 @@ async def handler(request: httpx2.Request) -> httpx2.Response:

@pytest.mark.asyncio
async def test_amessages_context_compaction_uses_beta_resource_and_preserves_response() -> None:
requests: list[httpx2.Request] = []
requests: list[httpx.Request] = []

async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx2.Response(
return httpx.Response(
200,
headers={"request-id": "req_test"},
json={
Expand All @@ -340,7 +413,7 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
},
)

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
context_management = {"edits": [{"type": "compact_20260112"}]}
params = MessagesParams(
Expand Down Expand Up @@ -535,11 +608,11 @@ def test_pop_anthropic_beta_header_preserves_unparseable_values(value: object) -

@pytest.mark.asyncio
async def test_amessages_merges_beta_extra_header_with_inferred_betas() -> None:
requests: list[httpx2.Request] = []
requests: list[httpx.Request] = []

async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx2.Response(
return httpx.Response(
200,
json={
"id": "msg_test",
Expand All @@ -558,7 +631,7 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
"x-custom-header": "custom-value",
}
original_extra_headers = extra_headers.copy()
http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
params = MessagesParams(
model="claude-opus-5",
Expand All @@ -581,11 +654,11 @@ async def handler(request: httpx2.Request) -> httpx2.Response:

@pytest.mark.asyncio
async def test_amessages_routes_beta_extra_header_through_beta_resource() -> None:
requests: list[httpx2.Request] = []
requests: list[httpx.Request] = []

async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx2.Response(
return httpx.Response(
200,
json={
"id": "msg_test",
Expand All @@ -599,7 +672,7 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
},
)

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
params = MessagesParams(
model="claude-opus-5",
Expand All @@ -620,8 +693,8 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
async def test_amessages_beta_extra_header_suppresses_unknown_edit_warning(
caplog: pytest.LogCaptureFixture,
) -> None:
async def handler(request: httpx2.Request) -> httpx2.Response:
return httpx2.Response(
async def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(
200,
json={
"id": "msg_test",
Expand All @@ -635,7 +708,7 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
},
)

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
params = MessagesParams(
model="claude-opus-5",
Expand Down Expand Up @@ -691,11 +764,11 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
async def test_amessages_selects_betas_for_context_management(
context_management: dict[str, Any] | None, betas: list[str] | None, expected_betas: str
) -> None:
requests: list[httpx2.Request] = []
requests: list[httpx.Request] = []

async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx2.Response(
return httpx.Response(
200,
headers={"request-id": "req_test"},
json={
Expand All @@ -710,7 +783,7 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
},
)

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
params = MessagesParams(
model="claude-opus-5",
Expand All @@ -730,7 +803,7 @@ async def handler(request: httpx2.Request) -> httpx2.Response:

@pytest.mark.asyncio
async def test_amessages_streams_beta_compaction_events() -> None:
async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
assert request.url.query == b"beta=true"
events = [
(
Expand Down Expand Up @@ -782,9 +855,9 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
("message_stop", {"type": "message_stop"}),
]
body = "".join(f"event: {name}\ndata: {json.dumps(data)}\n\n" for name, data in events)
return httpx2.Response(200, text=body, headers={"content-type": "text/event-stream"})
return httpx.Response(200, text=body, headers={"content-type": "text/event-stream"})

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
params = MessagesParams(
model="claude-opus-5",
Expand Down Expand Up @@ -839,7 +912,7 @@ async def handler(request: httpx2.Request) -> httpx2.Response:

@pytest.mark.asyncio
async def test_amessages_streams_beta_only_content_block() -> None:
async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
assert request.url.query == b"beta=true"
events = [
(
Expand Down Expand Up @@ -884,9 +957,9 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
("message_stop", {"type": "message_stop"}),
]
body = "".join(f"event: {name}\ndata: {json.dumps(data)}\n\n" for name, data in events)
return httpx2.Response(200, text=body, headers={"content-type": "text/event-stream"})
return httpx.Response(200, text=body, headers={"content-type": "text/event-stream"})

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
params = MessagesParams(
model="claude-opus-5",
Expand Down Expand Up @@ -1297,8 +1370,12 @@ async def __anext__(self) -> Any:
async def test_amessages_stream_preserves_accumulated_stop_event_payloads() -> None:
"""The SDK's stream helper attaches the accumulated message and block to the stop events."""

async def handler(request: httpx2.Request) -> httpx2.Response:
async def handler(request: httpx.Request) -> httpx.Response:
assert request.url.query == b""
request_body = json.loads(request.content)
assert request_body["temperature"] == 0.7
assert request_body["top_p"] == 0.9
assert request_body["top_k"] == 40
events = [
(
"message_start",
Expand Down Expand Up @@ -1336,14 +1413,17 @@ async def handler(request: httpx2.Request) -> httpx2.Response:
("message_stop", {"type": "message_stop"}),
]
body = "".join(f"event: {name}\ndata: {json.dumps(payload)}\n\n" for name, payload in events)
return httpx2.Response(200, headers={"content-type": "text/event-stream"}, content=body.encode())
return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=body.encode())

http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
provider = AnthropicProvider(api_key="test-key", http_client=http_client)
params = MessagesParams(
model="claude-opus-5",
messages=[{"role": "user", "content": "Hello"}],
max_tokens=1024,
temperature=0.7,
top_p=0.9,
top_k=40,
stream=True,
)

Expand Down
Loading