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
36 changes: 34 additions & 2 deletions src/any_llm/any_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
ContentFilterFinishReasonError,
LengthFinishReasonError,
MissingApiKeyError,
UnsupportedParameterError,
UnsupportedProviderError,
)
from any_llm.tools import prepare_tools
Expand Down Expand Up @@ -132,6 +133,9 @@ class AnyLLM(ABC):
SUPPORTS_MESSAGES: bool = True
"""Anthropic Messages API (all providers support it via conversion)"""

PROMPT_CACHE_KEY_SUPPORT: Literal["unsupported", "supported", "passthrough"] = "unsupported"
"""Whether prompt_cache_key is supported, forwarded to a router, or rejected."""

API_BASE: str | None = None
"""This is used to set the API base for the provider.
It is not required but may prove useful for providers that have overridable api bases.
Expand Down Expand Up @@ -572,6 +576,7 @@ def completion(
*,
response_format: dict[str, Any] | type | None = None,
stream: bool | None = None,
prompt_cache_key: str | None = None,
allow_running_loop: bool | None = None,
**kwargs: Any,
) -> ChatCompletion | Iterator[ChatCompletionChunk] | ParsedChatCompletion[Any]:
Expand All @@ -588,13 +593,21 @@ def completion(
messages=messages,
response_format=response_format,
stream=stream,
prompt_cache_key=prompt_cache_key,
**kwargs,
),
allow_running_loop=allow_running_loop,
)

response = run_async_in_sync(
self.acompletion(model=model, messages=messages, response_format=response_format, stream=stream, **kwargs),
self.acompletion(
model=model,
messages=messages,
response_format=response_format,
stream=stream,
prompt_cache_key=prompt_cache_key,
**kwargs,
),
allow_running_loop=allow_running_loop,
)
if isinstance(response, ChatCompletion):
Expand Down Expand Up @@ -673,6 +686,7 @@ async def acompletion(
stream_options: dict[str, Any] | None = None,
max_completion_tokens: int | None = None,
reasoning_effort: ReasoningEffort | None = "auto",
prompt_cache_key: str | None = None,
**kwargs: Any,
) -> ChatCompletion | AsyncIterator[ChatCompletionChunk] | ParsedChatCompletion[Any]:
"""Create a chat completion asynchronously.
Expand Down Expand Up @@ -701,6 +715,7 @@ async def acompletion(
stream_options: Additional options controlling streaming behavior
max_completion_tokens: Maximum number of tokens for the completion
reasoning_effort: Reasoning effort level for models that support it. "auto" will map to each provider's default.
prompt_cache_key: A key to use when reading from or writing to a provider's prompt cache.
**kwargs: Additional provider-specific arguments that will be passed to the provider's API call.

Returns:
Expand Down Expand Up @@ -742,8 +757,10 @@ async def acompletion(
stream_options=stream_options,
max_completion_tokens=max_completion_tokens,
reasoning_effort=reasoning_effort,
prompt_cache_key=prompt_cache_key,
)

self._validate_prompt_cache_key(prompt_cache_key)
result = await self._acompletion(params, **kwargs)

if is_structured_output_type(response_format):
Expand All @@ -767,6 +784,11 @@ async def acompletion(

return result

def _validate_prompt_cache_key(self, prompt_cache_key: str | None) -> None:
if prompt_cache_key is not None and self.PROMPT_CACHE_KEY_SUPPORT == "unsupported":
parameter_name = "prompt_cache_key"
raise UnsupportedParameterError(parameter_name, self.PROVIDER_NAME)

async def _acompletion(
self, params: CompletionParams, **kwargs: Any
) -> ChatCompletion | AsyncIterator[ChatCompletionChunk]:
Expand All @@ -780,6 +802,7 @@ def messages(
self,
*,
allow_running_loop: bool | None = None,
prompt_cache_key: str | None = None,
context_management: dict[str, Any] | None = None,
betas: list[str] | None = None,
**kwargs: Any,
Expand All @@ -791,7 +814,12 @@ def messages(
if allow_running_loop is None:
allow_running_loop = INSIDE_NOTEBOOK
response = run_async_in_sync(
self.amessages(context_management=context_management, betas=betas, **kwargs),
self.amessages(
prompt_cache_key=prompt_cache_key,
context_management=context_management,
betas=betas,
**kwargs,
),
allow_running_loop=allow_running_loop,
)
if isinstance(response, (MessageResponse, ParsedMessage, ParsedBetaMessage)):
Expand All @@ -816,6 +844,7 @@ async def amessages(
metadata: dict[str, Any] | None = None,
thinking: dict[str, Any] | None = None,
cache_control: dict[str, Any] | None = None,
prompt_cache_key: str | None = None,
context_management: dict[str, Any] | None = None,
betas: list[str] | None = None,
output_format: type | dict[str, Any] | None = None,
Expand All @@ -841,6 +870,7 @@ async def amessages(
metadata: Request metadata.
thinking: Thinking/reasoning configuration.
cache_control: Cache control configuration for prompt caching.
prompt_cache_key: A key to use when reading from or writing to a provider's prompt cache.
context_management: Anthropic context management configuration. The
`compact_20260112` strategy requires a supported model. Its `input_tokens`
trigger value must be at least 50,000 when provided; see
Expand Down Expand Up @@ -882,10 +912,12 @@ async def amessages(
metadata=metadata,
thinking=thinking,
cache_control=cache_control,
prompt_cache_key=prompt_cache_key,
context_management=context_management,
betas=betas,
output_format=output_format,
)
self._validate_prompt_cache_key(prompt_cache_key)
result = await self._amessages(params, **kwargs)

# The Anthropic provider already returns a ParsedMessage via native messages.parse (typed
Expand Down
12 changes: 12 additions & 0 deletions src/any_llm/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ def completion(
stream_options: dict[str, Any] | None = None,
max_completion_tokens: int | None = None,
reasoning_effort: ReasoningEffort | None = "auto",
prompt_cache_key: str | None = None,
client_args: dict[str, Any] | None = None,
**kwargs: Any,
) -> ChatCompletion | Iterator[ChatCompletionChunk]:
Expand Down Expand Up @@ -85,6 +86,7 @@ def completion(
stream_options: Additional options controlling streaming behavior
max_completion_tokens: Maximum number of tokens for the completion
reasoning_effort: Reasoning effort level for models that support it. "auto" will map to each provider's default.
prompt_cache_key: A key to use when reading from or writing to a provider's prompt cache.
client_args: Additional provider-specific arguments that will be passed to the provider's client instantiation.
**kwargs: Additional provider-specific arguments that will be passed to the provider's API call.

Expand Down Expand Up @@ -127,6 +129,7 @@ def completion(
stream_options=stream_options,
max_completion_tokens=max_completion_tokens,
reasoning_effort=reasoning_effort,
prompt_cache_key=prompt_cache_key,
**kwargs,
)

Expand Down Expand Up @@ -159,6 +162,7 @@ async def acompletion(
stream_options: dict[str, Any] | None = None,
max_completion_tokens: int | None = None,
reasoning_effort: ReasoningEffort | None = "auto",
prompt_cache_key: str | None = None,
client_args: dict[str, Any] | None = None,
**kwargs: Any,
) -> ChatCompletion | AsyncIterator[ChatCompletionChunk]:
Expand Down Expand Up @@ -194,6 +198,7 @@ async def acompletion(
stream_options: Additional options controlling streaming behavior
max_completion_tokens: Maximum number of tokens for the completion
reasoning_effort: Reasoning effort level for models that support it. "auto" will map to each provider's default.
prompt_cache_key: A key to use when reading from or writing to a provider's prompt cache.
client_args: Additional provider-specific arguments that will be passed to the provider's client instantiation.
**kwargs: Additional provider-specific arguments that will be passed to the provider's API call.

Expand Down Expand Up @@ -236,6 +241,7 @@ async def acompletion(
stream_options=stream_options,
max_completion_tokens=max_completion_tokens,
reasoning_effort=reasoning_effort,
prompt_cache_key=prompt_cache_key,
**kwargs,
)

Expand Down Expand Up @@ -555,6 +561,7 @@ def messages(
metadata: dict[str, Any] | None = None,
thinking: dict[str, Any] | None = None,
cache_control: dict[str, Any] | None = None,
prompt_cache_key: str | None = None,
context_management: dict[str, Any] | None = None,
betas: list[str] | None = None,
output_format: type | dict[str, Any] | None = None,
Expand Down Expand Up @@ -582,6 +589,7 @@ def messages(
metadata: Request metadata.
thinking: Thinking/reasoning configuration.
cache_control: Cache control configuration for prompt caching.
prompt_cache_key: A key to use when reading from or writing to a provider's prompt cache.
context_management: Anthropic context management configuration. The `compact_20260112`
strategy requires a supported model. Its `input_tokens` trigger value must be at
least 50,000 when provided; see [Anthropic's compaction documentation](https://platform.claude.com/docs/en/build-with-claude/compaction).
Expand Down Expand Up @@ -623,6 +631,7 @@ def messages(
metadata=metadata,
thinking=thinking,
cache_control=cache_control,
prompt_cache_key=prompt_cache_key,
context_management=context_management,
betas=betas,
output_format=output_format,
Expand All @@ -647,6 +656,7 @@ async def amessages(
metadata: dict[str, Any] | None = None,
thinking: dict[str, Any] | None = None,
cache_control: dict[str, Any] | None = None,
prompt_cache_key: str | None = None,
context_management: dict[str, Any] | None = None,
betas: list[str] | None = None,
output_format: type | dict[str, Any] | None = None,
Expand Down Expand Up @@ -674,6 +684,7 @@ async def amessages(
metadata: Request metadata.
thinking: Thinking/reasoning configuration.
cache_control: Cache control configuration for prompt caching.
prompt_cache_key: A key to use when reading from or writing to a provider's prompt cache.
context_management: Anthropic context management configuration. The `compact_20260112`
strategy requires a supported model. Its `input_tokens` trigger value must be at
least 50,000 when provided; see [Anthropic's compaction documentation](https://platform.claude.com/docs/en/build-with-claude/compaction).
Expand Down Expand Up @@ -715,6 +726,7 @@ async def amessages(
metadata=metadata,
thinking=thinking,
cache_control=cache_control,
prompt_cache_key=prompt_cache_key,
context_management=context_management,
betas=betas,
output_format=output_format,
Expand Down
1 change: 1 addition & 0 deletions src/any_llm/providers/openai/custom.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@ class OpenAICompatibleProvider(BaseOpenAIProvider):
PROVIDER_NAME = "openai_compatible"
PROVIDER_DOCUMENTATION_URL = "https://platform.openai.com/docs/api-reference"
ENV_API_KEY_NAME = "OPENAI_COMPATIBLE_API_KEY"
PROMPT_CACHE_KEY_SUPPORT = "passthrough"

def __init__(self, api_base: str, api_key: str | None = None, **kwargs: Any) -> None:
if not api_base:
Expand Down
1 change: 1 addition & 0 deletions src/any_llm/providers/openai/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ class OpenaiProvider(BaseOpenAIProvider):
ENV_API_BASE_NAME = "OPENAI_BASE_URL"
PROVIDER_NAME = "openai"
PROVIDER_DOCUMENTATION_URL = "https://platform.openai.com/docs/api-reference"
PROMPT_CACHE_KEY_SUPPORT = "supported"
SUPPORTS_RESPONSES = True
SUPPORTS_LIST_MODELS = True
SUPPORTS_BATCH = True
Expand Down
1 change: 1 addition & 0 deletions src/any_llm/providers/otari/otari.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,7 @@ class OtariProvider(BaseOpenAIProvider):
PROVIDER_NAME = "otari"
PROVIDER_DOCUMENTATION_URL = "https://mozilla-ai.github.io/otari/"
MISSING_PACKAGES_ERROR = _MISSING_PACKAGES_ERROR
PROMPT_CACHE_KEY_SUPPORT = "passthrough"

SUPPORTS_COMPLETION_STREAMING = True
SUPPORTS_COMPLETION = True
Expand Down
3 changes: 3 additions & 0 deletions src/any_llm/types/completion.py
Original file line number Diff line number Diff line change
Expand Up @@ -232,3 +232,6 @@ def check_messages_not_empty(cls, v: list[dict[str, Any]]) -> list[dict[str, Any

reasoning_effort: ReasoningEffort | None = "auto"
"""Reasoning effort level for models that support it. "auto" will map to each provider's default."""

prompt_cache_key: str | None = None
"""A key to use when reading from or writing to a provider's prompt cache."""
3 changes: 3 additions & 0 deletions src/any_llm/types/messages.py
Original file line number Diff line number Diff line change
Expand Up @@ -212,6 +212,9 @@ class MessagesParams(BaseModel):
cache_control: dict[str, Any] | None = None
"""Cache control configuration for prompt caching"""

prompt_cache_key: str | None = None
"""A key to use when reading from or writing to a provider's prompt cache."""

context_management: dict[str, Any] | None = None
"""Anthropic context management configuration"""

Expand Down
2 changes: 2 additions & 0 deletions src/any_llm/utils/messages_compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,8 @@ def messages_params_to_completion_params(params: MessagesParams) -> dict[str, An
"max_tokens": params.max_tokens,
}

if params.prompt_cache_key is not None:
result["prompt_cache_key"] = params.prompt_cache_key
if params.temperature is not None:
result["temperature"] = params.temperature
if params.top_p is not None:
Expand Down
48 changes: 48 additions & 0 deletions tests/unit/providers/test_anthropic_messages.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from anthropic.types.beta import BetaMCPToolUseBlock, BetaMessage, BetaThinkingBlock, BetaUsage
from pydantic import BaseModel

from any_llm.exceptions import 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.messages import (
Expand Down Expand Up @@ -257,6 +258,53 @@ async def test_amessages_non_streaming() -> None:
assert call_kwargs["max_tokens"] == 1024


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

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

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"):
await provider.amessages(
model="claude-opus-5",
messages=[{"role": "user", "content": "Hello"}],
max_tokens=1024,
prompt_cache_key="tenant-1",
)
finally:
await http_client.aclose()

assert requests == []


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

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

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"):
await provider.acompletion(
model="claude-opus-5",
messages=[{"role": "user", "content": "Hello"}],
prompt_cache_key="tenant-1",
)
finally:
await http_client.aclose()

assert requests == []


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