diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 22d76242f2f..6047c625665 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -12,7 +12,7 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and import copy import os import re -from collections.abc import Iterable, Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, cast from urllib.parse import urlparse @@ -90,6 +90,14 @@ def _validated_object_list(value: object) -> list[object] | None: return None +def configured_injection_points(value: object) -> Sequence[CacheControlInjectionPoint]: + if not isinstance(value, (list, tuple)): + return () + if all(isinstance(entry, dict) for entry in value): + return cast(Sequence[CacheControlInjectionPoint], value) + return tuple(cast(CacheControlInjectionPoint, entry) for entry in value if isinstance(entry, dict)) + + def supports_openai_prompt_cache_breakpoint(model: str) -> bool: model_map_flag: Final = _model_map_prompt_cache_breakpoint_flag(model) if model_map_flag is not None: @@ -106,7 +114,60 @@ def _model_map_prompt_cache_breakpoint_flag(model: str) -> bool | None: entries: Final = (litellm.model_cost.get(key) for key in (model, model.rsplit("/", 1)[-1])) flags: Final = (entry.get("supports_prompt_cache_breakpoint") for entry in entries if isinstance(entry, dict)) - return next((bool(flag) for flag in flags if flag is not None), None) + return next((flag is True for flag in flags if flag is not None), None) + + +def _hosted_entry_flag(entry: Mapping[str, object], resolve_provider: Callable[[], str | None]) -> bool | None: + flag: Final = entry.get("supports_prompt_cache_breakpoint") + entry_provider: Final = entry.get("litellm_provider") + if flag is None or entry_provider is None or entry_provider == "openai": + return None + return (flag is True) if entry_provider == resolve_provider() else None + + +def _hosted_openai_dialect_flag( + model: str, custom_llm_provider: str | None, resolve_provider: Callable[[str], str | None] +) -> bool | None: + """ + Explicit opt-in for an OpenAI-shaped deployment served by another provider. + + ``model_cost`` is keyed per deployment string, so a flag on the deployment's own + entry states the dialect directly, which a provider name cannot express. A routing + form the map does not key verbatim (``bedrock_mantle/us-east-1/openai.gpt-5.6-sol``) + is read through the candidate keys ``get_model_info`` resolves it with, an entry + keyed with the region outranking the region-free one. A bare name the map does not + key has no such candidates, so it costs no provider lookup. Entries for the openai + provider are left to the caller's api_base check, so an OpenAI-compatible + third-party host is still not assumed to speak the dialect. A bare deployment name + another provider serves (``gpt-6-astra`` on azure_ai) may collide with the openai + row of the same name, so an exact entry only speaks for a deployment when it is + keyed for that deployment's provider. + """ + import litellm + + exact_entry: Final = litellm.model_cost.get(model) + exact_provider: Final = exact_entry.get("litellm_provider") if isinstance(exact_entry, dict) else None + if custom_llm_provider is None and (exact_provider == "openai" or ("/" not in model and exact_provider is None)): + return None + provider: Final = custom_llm_provider or resolve_provider(model) + if provider is None or provider == "openai": + return None + if isinstance(exact_entry, dict) and exact_provider == provider: + return _hosted_entry_flag(exact_entry, lambda: provider) + from litellm.utils import get_potential_model_names + + names: Final = get_potential_model_names(model, provider) + candidates: Final = ( + names["combined_model_name"], + names["region_free_combined_model_name"], + names["split_model"], + names["combined_stripped_model_name"], + names["stripped_model_name"], + names["provider_prefixed_model_name"], + ) + entries: Final = (litellm.model_cost.get(candidate) for candidate in candidates) + flags: Final = (_hosted_entry_flag(entry, lambda: provider) for entry in entries if isinstance(entry, dict)) + return next((flag for flag in flags if flag is not None), None) def targets_openai_api(api_base: object) -> bool: @@ -227,8 +288,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): """ # Extract cache control injection points carry_unmatched: Final = bool(non_default_params.pop(CARRY_UNMATCHED_MESSAGE_POINTS, False)) - injection_points: Final[list[CacheControlInjectionPoint]] = non_default_params.pop( - "cache_control_injection_points", [] + injection_points: Final = configured_injection_points( + non_default_params.pop("cache_control_injection_points", None) ) if not injection_points: return model, messages, non_default_params @@ -282,8 +343,9 @@ class AnthropicCacheControlHook(CustomPromptManagement): if ( openai_dialect and AnthropicCacheControlHook.count_request_cache_breakpoints(processed_messages) > breakpoints_before + and non_default_params.get("prompt_cache_options") is None ): - non_default_params.setdefault("prompt_cache_options", PromptCacheOptions(mode="implicit")) + non_default_params["prompt_cache_options"] = PromptCacheOptions(mode="implicit") # Points this pass did not place: non-message ones for the provider transform, and # the deferred role-targeted ones. Deferring is what reaches the Responses API's @@ -310,7 +372,14 @@ class AnthropicCacheControlHook(CustomPromptManagement): api_base: object = None, prompt_cache_options: object = None, ) -> bool: - if model is None or not supports_openai_prompt_cache_breakpoint(model): + if model is None: + return False + hosted_flag: Final = _hosted_openai_dialect_flag( + model, custom_llm_provider, AnthropicCacheControlHook._resolve_provider + ) + if hosted_flag is not None: + return hosted_flag + if not supports_openai_prompt_cache_breakpoint(model): return False if (custom_llm_provider or AnthropicCacheControlHook._resolve_provider(model)) != "openai": return False @@ -701,7 +770,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): api_base: object, prompt_cache_options: object, ) -> Sequence[Mapping[str, object]]: - if not supports_openai_prompt_cache_breakpoint(model): + if not AnthropicCacheControlHook._may_target_openai_prompt_cache_breakpoint(model, custom_llm_provider): return points return AnthropicCacheControlHook._stamped( points, @@ -711,6 +780,20 @@ class AnthropicCacheControlHook(CustomPromptManagement): ), ) + @staticmethod + def _may_target_openai_prompt_cache_breakpoint(model: str, custom_llm_provider: str | None) -> bool: + """Cheap gate before the dialect is resolved: the model's own row or version, or, when the + serving provider is already known, a row keyed for that provider (``openai.gpt-5.6-sol`` + served by ``bedrock_mantle``), which costs no provider lookup.""" + if supports_openai_prompt_cache_breakpoint(model): + return True + if custom_llm_provider is None: + return False + return ( + _hosted_openai_dialect_flag(model, custom_llm_provider, AnthropicCacheControlHook._resolve_provider) + is not None + ) + @staticmethod def _stamped(points: Sequence[Mapping[str, object]], key: str, value: object) -> Sequence[Mapping[str, object]]: return [{**point, key: value} for point in points] @@ -859,7 +942,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): """ import litellm - configured: Final = non_default_params.get("cache_control_injection_points") + configured: Final = configured_injection_points(non_default_params.get("cache_control_injection_points")) if configured: tools_keeping_marks: Final = tuple( tool for tool in tools or () if not _chat_transform_drops_tool_cache_control(tool) @@ -990,9 +1073,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): bool | None, kwargs.pop("enable_prompt_caching", None) ) cache_control: Final = kwargs.get("cache_control") - configured: Final = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list - list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None) - ) + configured: Final = configured_injection_points(kwargs.pop("cache_control_injection_points", None)) injection_points: Final[Sequence[CacheControlInjectionPoint]] = configured or ( AnthropicCacheControlHook.get_default_injection_points( messages=typed_messages, @@ -1028,8 +1109,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system) - breakpoints_before ) AnthropicCacheControlHook.record_gateway_injection(kwargs, breakpoints_added) - if openai_dialect and breakpoints_added > 0: - kwargs.setdefault("prompt_cache_options", PromptCacheOptions(mode="implicit")) + if openai_dialect and breakpoints_added > 0 and kwargs.get("prompt_cache_options") is None: + kwargs["prompt_cache_options"] = PromptCacheOptions(mode="implicit") if remaining: kwargs["cache_control_injection_points"] = remaining return messages, system diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index f34606ce127..32fd15de9e7 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -58344,6 +58344,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58387,6 +58388,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58419,6 +58421,7 @@ "supports_function_calling": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58495,6 +58498,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58882,6 +58886,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58920,6 +58925,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58958,6 +58964,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -80349,6 +80356,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, diff --git a/litellm/responses/main.py b/litellm/responses/main.py index c4632f09723..7298f4b1a36 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -18,7 +18,11 @@ from litellm.completion_extras.litellm_responses_transformation.transformation i LiteLLMResponsesTransformationHandler, ) from litellm.constants import DEFAULT_CHAT_COMPLETION_PARAM_VALUES, request_timeout -from litellm.integrations.anthropic_cache_control_hook import CARRY_UNMATCHED_MESSAGE_POINTS +from litellm.integrations.anthropic_cache_control_hook import ( + CARRY_UNMATCHED_MESSAGE_POINTS, + AnthropicCacheControlHook, + configured_injection_points, +) from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -526,6 +530,13 @@ def _api_base_kwarg(kwargs: Mapping[str, object]) -> str | None: return api_base if isinstance(api_base, str) else None +def _dispatched_model_name(model: str, custom_llm_provider: str, api_base: str | None) -> str: + provider_model, _, _, _ = litellm.get_llm_provider( + model=model, custom_llm_provider=custom_llm_provider, api_base=api_base + ) + return _strip_responses_routing_prefix(provider_model) + + def _will_bridge_to_chat_completions( model: str, custom_llm_provider: str | None, @@ -536,20 +547,50 @@ def _will_bridge_to_chat_completions( """``_bridges_to_chat_completions`` for callers running before the provider config is resolved. Resolving the config is a pure lookup, so this asks the same question the dispatch - asks rather than restating its condition. Both callers resolve the provider before - this runs, so the only way to be wrong is a prompt manager that moves the model - across the bridge boundary, which would leave the deferred points to a pass that - never comes. + asks rather than restating its condition, with the model name the dispatch hands the + lookup: a provider whose config is keyed by model name (Bedrock Mantle reads the + price map) answers nothing for ``bedrock_mantle/openai.gpt-5.6-sol`` and would read + as bridged. Both callers resolve the provider before this runs, so the only way to be + wrong is a prompt manager that moves the model across the bridge boundary, which + would leave the deferred points to a pass that never comes. """ normalized_model: Final = _normalize_openai_chat_completions_responses_model(model) if custom_llm_provider is None: return True return _bridges_to_chat_completions( - _resolve_responses_api_provider_config(normalized_model[0], custom_llm_provider, model_info, api_base), + _resolve_responses_api_provider_config( + _dispatched_model_name(normalized_model[0], custom_llm_provider, api_base), + custom_llm_provider, + model_info, + api_base, + ), use_chat_completions_api or normalized_model[1], ) +def _stamp_injection_points_with_dialect( + kwargs: dict[str, object], # mutable-ok: the points are rewritten in the caller's own kwargs for the hook to read + model: str, + custom_llm_provider: str | None, +) -> None: + """Carry the provider this layer resolved onto the points. + + The hook reads ``custom_llm_provider`` from the request kwargs, which never hold the one + resolved here, and resolving the model name alone reads a Foundry deployment of an OpenAI + model (``azure_ai/gpt-6-astra``) as Azure OpenAI, which left it on the Anthropic dialect. + """ + points: Final = configured_injection_points(kwargs.get("cache_control_injection_points")) + if not points: + return + kwargs["cache_control_injection_points"] = AnthropicCacheControlHook._stamped_with_dialect( + points, + model, + custom_llm_provider, + kwargs.get("api_base") or kwargs.get("base_url"), + kwargs.get("prompt_cache_options"), + ) + + @contextmanager def _prompt_management_sees_a_provisional_message_list( kwargs: dict[str, object], # mutable-ok: the signal is read and popped out of the caller's own kwargs @@ -659,6 +700,7 @@ async def aresponses( _api_base_kwarg(kwargs), ), ): + _stamp_injection_points_with_dialect(kwargs, model, custom_llm_provider) ( model, merged_input, @@ -829,6 +871,7 @@ def _apply_prompt_management_to_responses_call( _api_base_kwarg(kwargs), ), ): + _stamp_injection_points_with_dialect(kwargs, model, custom_llm_provider) ( model, merged_input, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index f34606ce127..32fd15de9e7 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -58344,6 +58344,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58387,6 +58388,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58419,6 +58421,7 @@ "supports_function_calling": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58495,6 +58498,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58882,6 +58886,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58920,6 +58925,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58958,6 +58964,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -80349,6 +80356,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, diff --git a/tests/integration/providers/_mantle_gpt_prompt_cache_support.py b/tests/integration/providers/_mantle_gpt_prompt_cache_support.py new file mode 100644 index 00000000000..8f497fea74f --- /dev/null +++ b/tests/integration/providers/_mantle_gpt_prompt_cache_support.py @@ -0,0 +1,959 @@ +import base64 +import binascii +import json +import math +import os +import uuid +from collections.abc import Callable, Iterable, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final, Literal, Protocol +from urllib.parse import urlsplit + +import anthropic +import httpx +import openai +from integration._support.client import Gateway, Scenario, eventually +from integration._support.database import read_rows +from integration._support.responses_vendor import answer, error, newest_marker, sse +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with + +Endpoint = Literal["chat", "responses", "messages"] + + +JSON: Final = TypeAdapter(dict[str, JsonValue]) + + +ITEMS: Final = TypeAdapter(list[dict[str, JsonValue]]) + + +SIGNING_KEY: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt") + + +TOKEN: Final = "synthetic-mantle-bearer" + + +GPT: Final = "bedrock_mantle/openai.gpt-5.6-sol" + + +GPT_REGION: Final = "bedrock_mantle/us-east-1/openai.gpt-5.6-sol" + +GPT_BARE: Final = "openai.gpt-5.6-sol" + + +GPT_UNFLAGGED_ROW: Final = "bedrock_mantle/openai.gpt-5.4" + + +GPT_FLAGGED_ROW: Final = "bedrock_mantle/openai.gpt-6-luna" + + +GPT_ODD_STRING_FLAG: Final = "bedrock_mantle/openai.gpt-5.5" + + +GPT_ODD_INT_FLAG: Final = "bedrock_mantle/openai.gpt-daybreak-blue-5.6-sol" +ODD_FLAGS: Final[tuple[tuple[str, str, JsonValue], ...]] = ( + ("string-true", GPT_ODD_STRING_FLAG, "true"), + ("int-one", GPT_ODD_INT_FLAG, 1), + ("int-zero", GPT_ODD_INT_FLAG, 0), + ("5kb-string", GPT_ODD_STRING_FLAG, "x" * 5120), +) + + +CLAUDE: Final = "bedrock_mantle/anthropic.claude-haiku-4-5" + + +AZURE: Final = "azure/gpt-5.6" + + +THIRD_PARTY: Final = "openai/gpt-5.6" + + +SYSTEM: Final = "Reply with the signature the user gives you." + + +SYSTEM_POINT: Final[list[JsonValue]] = [{"location": "message", "role": "system"}] + + +EXPLICIT: Final[dict[str, JsonValue]] = {"mode": "explicit"} + + +IMPLICIT: Final[dict[str, JsonValue]] = {"mode": "implicit"} + + +EPHEMERAL: Final[dict[str, JsonValue]] = {"type": "ephemeral"} + + +NO_CACHE: Final[dict[str, JsonValue]] = {"no-cache": True} + + +MAX_TOKENS: Final = 64 + + +INPUT_TOKENS: Final = 2730 + + +CACHED_TOKENS: Final = 1024 + + +WRITTEN_TOKENS: Final = 1700 + + +UNCACHED_TOKENS: Final = INPUT_TOKENS - CACHED_TOKENS - WRITTEN_TOKENS + + +OUTPUT_TOKENS: Final = 7 + + +USAGE: Final[dict[str, JsonValue]] = { + "input_tokens": INPUT_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "total_tokens": INPUT_TOKENS + OUTPUT_TOKENS, + "input_tokens_details": {"cached_tokens": CACHED_TOKENS, "cache_write_tokens": WRITTEN_TOKENS}, + "output_tokens_details": {"reasoning_tokens": 0}, +} + + +ANTHROPIC_USAGE: Final[dict[str, JsonValue]] = { + "input_tokens": UNCACHED_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "cache_read_input_tokens": CACHED_TOKENS, + "cache_creation_input_tokens": WRITTEN_TOKENS, +} + + +CHAT_USAGE: Final[dict[str, JsonValue]] = { + "prompt_tokens": INPUT_TOKENS, + "completion_tokens": OUTPUT_TOKENS, + "total_tokens": INPUT_TOKENS + OUTPUT_TOKENS, + "prompt_tokens_details": {"cached_tokens": CACHED_TOKENS}, +} + + +COST_MAP: Final = JSON.validate_json(Path("model_prices_and_context_window.json").read_bytes()) + + +RESPONSES_PATH: Final = "/openai/v1/responses" + + +ANTHROPIC_PATH: Final = "/anthropic/v1/messages" + + +BURST: Final = 30 + + +ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "responses", "messages") + + +def cost_rate(model: str, field: str) -> float: + value: Final = JSON.validate_python(COST_MAP[model])[field] + assert isinstance(value, int | float), (model, field, value) + return float(value) + + +def expected_spend(model: str) -> float: + return ( + UNCACHED_TOKENS * cost_rate(model, "input_cost_per_token") + + CACHED_TOKENS * cost_rate(model, "cache_read_input_token_cost") + + WRITTEN_TOKENS * cost_rate(model, "cache_creation_input_token_cost") + + OUTPUT_TOKENS * cost_rate(model, "output_cost_per_token") + ) + + +def fresh_marker() -> str: + return uuid.uuid4().hex + + +def prompt_text(marker: str) -> str: + return f"Return the signature marker-{marker}." + + +def valid_options(options: JsonValue) -> bool: + if options is None: + return True + if not isinstance(options, dict) or not set(options) <= {"mode", "ttl"}: + return False + return options.get("mode", "implicit") in ("implicit", "explicit") + + +def message_item(identity: str, text: str, status: str) -> dict[str, JsonValue]: + return { + "id": f"msg_{identity}", + "type": "message", + "role": "assistant", + "status": status, + "content": [{"type": "output_text", "text": text, "annotations": []}] if status == "completed" else [], + } + + +def responses_object(identity: str, model: str, text: str) -> dict[str, JsonValue]: + return { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": model, + "output": [message_item(identity, text, "completed")], + "usage": USAGE, + } + + +def responses_stream(identity: str, model: str, text: str) -> tuple[bytes, ...]: + response: Final = responses_object(identity, model, text) + return ( + sse({"type": "response.created", "sequence_number": 0, "response": {**response, "status": "in_progress"}}), + sse( + { + "type": "response.output_item.added", + "sequence_number": 1, + "output_index": 0, + "item": message_item(identity, text, "in_progress"), + } + ), + sse( + { + "type": "response.output_text.delta", + "sequence_number": 2, + "item_id": f"msg_{identity}", + "output_index": 0, + "content_index": 0, + "delta": text, + } + ), + sse( + { + "type": "response.output_item.done", + "sequence_number": 3, + "output_index": 0, + "item": message_item(identity, text, "completed"), + } + ), + sse({"type": "response.completed", "sequence_number": 4, "response": response}), + ) + + +def anthropic_message(identity: str, model: str, text: str) -> dict[str, JsonValue]: + return { + "id": identity, + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": text}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": ANTHROPIC_USAGE, + } + + +def anthropic_stream(identity: str, model: str, text: str) -> tuple[bytes, ...]: + started: Final = {**anthropic_message(identity, model, text), "content": [], "stop_reason": None} + return ( + sse({"type": "message_start", "message": {**started, "usage": {**ANTHROPIC_USAGE, "output_tokens": 0}}}), + sse({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + sse({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}), + sse({"type": "content_block_stop", "index": 0}), + sse( + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": OUTPUT_TOKENS}, + } + ), + sse({"type": "message_stop"}), + ) + + +def chat_completion(identity: str, model: str, text: str) -> dict[str, JsonValue]: + return { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": model, + "choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}], + "usage": CHAT_USAGE, + } + + +def issued_id(prefix: str, marker: str | None) -> str: + return f"{prefix}{marker or 'unmarked'}-{uuid.uuid4().hex[:8]}" + + +def mantle_peer(*, pause: float = 0) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.headers.get("authorization") == f"Bearer {TOKEN}", request.headers + body: Final = JSON.validate_json(request.body) + model: Final = str(body["model"]) + marker: Final = newest_marker(request.body.decode()) + text: Final = answer(marker) + streaming: Final = body.get("stream") is True + path: Final = urlsplit(request.target).path + if path == ANTHROPIC_PATH: + identity: Final = issued_id("msg_", marker) + if streaming: + return Reply( + content_type="text/event-stream", + chunks=anthropic_stream(identity, model, text), + pause_between_chunks=pause, + ) + return Reply(body=json.dumps(anthropic_message(identity, model, text)).encode()) + assert path == RESPONSES_PATH, request.target + if not valid_options(body.get("prompt_cache_options")): + return error(400, "Invalid prompt_cache_options", "invalid_prompt_cache_options") + response_id: Final = issued_id("resp_", marker) + if streaming: + return Reply( + content_type="text/event-stream", + chunks=responses_stream(response_id, model, text), + pause_between_chunks=pause, + ) + return Reply(body=json.dumps(responses_object(response_id, model, text)).encode()) + + return respond + + +def failing_peer(status: int, message: str, code: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert urlsplit(request.target).path == RESPONSES_PATH, request.target + return error(status, message, code) + + return respond + + +def openai_shaped_peer() -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + body: Final = JSON.validate_json(request.body) + model: Final = str(body["model"]) + marker: Final = newest_marker(request.body.decode()) + text: Final = answer(marker) + path: Final = urlsplit(request.target).path + if path.endswith("/chat/completions"): + return Reply(body=json.dumps(chat_completion(issued_id("chatcmpl-", marker), model, text)).encode()) + assert path.endswith("/responses"), request.target + return Reply(body=json.dumps(responses_object(issued_id("resp_", marker), model, text)).encode()) + + return respond + + +def system_item(*, marked: bool, endpoint: Endpoint) -> dict[str, JsonValue]: + part: Final[dict[str, JsonValue]] = {"type": "input_text", "text": SYSTEM} + content: Final[list[JsonValue]] = [{**part, "prompt_cache_breakpoint": EXPLICIT} if marked else part] + if endpoint == "responses": + return {"role": "system", "content": content} + return {"type": "message", "role": "system", "content": content} + + +def user_item(text: str, endpoint: Endpoint) -> dict[str, JsonValue]: + if endpoint == "responses": + return {"role": "user", "content": text} + return {"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]} + + +def expected_wire( + model: str, + prompt: str, + *, + endpoint: Endpoint, + marked: bool, + options: JsonValue | None = IMPLICIT, + **extra: JsonValue, +) -> dict[str, JsonValue]: + body: Final[dict[str, JsonValue]] = { + "model": model.rsplit("/", 1)[-1], + "input": [system_item(marked=marked, endpoint=endpoint), user_item(prompt, endpoint)], + "max_output_tokens": MAX_TOKENS, + **extra, + } + return body if options is None else {**body, "prompt_cache_options": options} + + +def body_of(request: Request) -> dict[str, JsonValue]: + return JSON.validate_json(request.body) + + +def without_stream(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return {key: value for key, value in body.items() if key != "stream"} + + +def only_received(wire: Wire) -> Request: + (request,) = wire.drain() + return request + + +def assert_wire(request: Request, expected: Mapping[str, JsonValue], *, streaming: bool) -> None: + body: Final = body_of(request) + assert urlsplit(request.target).path == RESPONSES_PATH, request.target + assert without_stream(body) == expected, json.dumps(body, sort_keys=True) + assert (body.get("stream") is True) is streaming, body.get("stream") + + +def breakpoint_count(body: Mapping[str, JsonValue]) -> int: + parts: Final = ( + part + for item in ITEMS.validate_python(body["input"]) + for part in (item["content"] if isinstance(item["content"], list) else ()) + ) # comprehension-ok: nested input items + return sum(1 for part in parts if isinstance(part, dict) and "prompt_cache_breakpoint" in part) + + +def endpoint_path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "responses": + return "/v1/responses" + case "messages": + return "/v1/messages" + + +def request_body( + endpoint: Endpoint, model: str, prompt: str, *, stream: bool = False, **extra: JsonValue +) -> dict[str, JsonValue]: + match endpoint: + case "chat": + return { + "model": model, + "messages": [{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + "max_tokens": MAX_TOKENS, + "stream": stream, + **({"stream_options": {"include_usage": True}} if stream else {}), + **extra, + } + case "responses": + return { + "model": model, + "input": [{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + "max_output_tokens": MAX_TOKENS, + "stream": stream, + **extra, + } + case "messages": + return { + "model": model, + "system": SYSTEM, + "messages": [{"role": "user", "content": prompt}], + "max_tokens": MAX_TOKENS, + "stream": stream, + **extra, + } + + +@dataclass(frozen=True, slots=True) +class Outcome: + status: int + call_id: str + response_id: str + text: str + usage: dict[str, JsonValue] + headers: Mapping[str, str] + raw: str + + +def sse_payloads(lines: Iterable[str]) -> tuple[dict[str, JsonValue], ...]: + return tuple(JSON.validate_json(line[6:]) for line in lines if line.startswith("data: ") and line != "data: [DONE]") + + +def unwrapped(identity: str) -> str | None: + try: + decoded: Final = base64.b64decode(identity.removeprefix("resp_"), validate=True).decode() + except (binascii.Error, UnicodeDecodeError): + return None + return decoded.rsplit("response_id:", 1)[1] if decoded.startswith("litellm:") else None + + +def upstream_id_of(identity: str) -> str: + managed: Final = decrypt_if_encrypted_with(identity.removeprefix("resp_"), SIGNING_KEY) + wrapped: Final = identity if managed is None else managed.rsplit("response_id:", 1)[1].split(";", 1)[0] + inner: Final = unwrapped(wrapped) + return wrapped if inner is None else inner + + +def row_key(request_id: str) -> str: + return upstream_id_of(request_id.split("_cache_hit", 1)[0]) + + +def caller_sees_upstream_id(endpoint: Endpoint, *, stream: bool) -> bool: + return (endpoint, stream) != ("messages", True) + + +def usage_of(payload: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + usage: Final = payload.get("usage") + return JSON.validate_python(usage) if isinstance(usage, dict) else {} + + +def chat_text(payload: Mapping[str, JsonValue]) -> str: + (choice,) = ITEMS.validate_python(payload["choices"]) + return str(JSON.validate_python(choice["message"])["content"]) + + +def chat_stream_fields(payloads: Sequence[Mapping[str, JsonValue]]) -> tuple[str, str, dict[str, JsonValue]]: + (identity,) = {str(chunk["id"]) for chunk in payloads if "id" in chunk} + choices: Final = tuple(ITEMS.validate_python(chunk["choices"]) for chunk in payloads if chunk.get("choices")) + deltas: Final = tuple(JSON.validate_python(choice[0]["delta"]) for choice in choices if choice) + text: Final = "".join(str(delta["content"]) for delta in deltas if isinstance(delta.get("content"), str)) + usages: Final = tuple(usage_of(chunk) for chunk in payloads if isinstance(chunk.get("usage"), dict)) + return identity, text, usages[-1] if usages else {} + + +def responses_text(payload: Mapping[str, JsonValue]) -> str: + items: Final = ITEMS.validate_python(payload["output"]) + parts: Final = tuple(ITEMS.validate_python(item["content"]) for item in items if item.get("type") == "message") + return "".join(str(part["text"]) for part in parts[0] if part.get("type") == "output_text") + + +def completed_response(events: Iterable[Mapping[str, JsonValue]]) -> dict[str, JsonValue]: + (completed,) = tuple(event for event in events if event.get("type") == "response.completed") + return JSON.validate_python(completed["response"]) + + +def messages_text(payload: Mapping[str, JsonValue]) -> str: + return "".join(str(block["text"]) for block in ITEMS.validate_python(payload["content"]) if "text" in block) + + +def messages_stream_fields(payloads: Sequence[Mapping[str, JsonValue]]) -> tuple[str, str, dict[str, JsonValue]]: + (started,) = tuple(payload for payload in payloads if payload.get("type") == "message_start") + message: Final = JSON.validate_python(started["message"]) + deltas: Final = tuple(JSON.validate_python(payload["delta"]) for payload in payloads if "delta" in payload) + text: Final = "".join(str(delta["text"]) for delta in deltas if delta.get("type") == "text_delta") + final_usages: Final = tuple(usage_of(payload) for payload in payloads if payload.get("type") == "message_delta") + return str(message["id"]), text, {**usage_of(message), **(final_usages[-1] if final_usages else {})} + + +def parse_outcome( + endpoint: Endpoint, *, stream: bool, status: int, headers: Mapping[str, str], lines: tuple[str, ...] +) -> Outcome: + call_id: Final = headers.get("x-litellm-call-id", "") + raw: Final = "\n".join(lines) + if status != 200: + return Outcome(status, call_id, "", "", {}, headers, raw) + payloads: Final = sse_payloads(lines) if stream else (JSON.validate_json(raw),) + match endpoint, stream: + case "chat", True: + identity, text, usage = chat_stream_fields(payloads) + return Outcome(status, call_id, upstream_id_of(identity), text, usage, headers, raw) + case "chat", False: + return Outcome( + status, + call_id, + upstream_id_of(str(payloads[0]["id"])), + chat_text(payloads[0]), + usage_of(payloads[0]), + headers, + raw, + ) + case "responses", True: + completed: Final = completed_response(payloads) + return Outcome( + status, + call_id, + upstream_id_of(str(completed["id"])), + responses_text(completed), + usage_of(completed), + headers, + raw, + ) + case "responses", False: + return Outcome( + status, + call_id, + upstream_id_of(str(payloads[0]["id"])), + responses_text(payloads[0]), + usage_of(payloads[0]), + headers, + raw, + ) + case "messages", True: + identity, text, usage = messages_stream_fields(payloads) + return Outcome(status, call_id, upstream_id_of(identity), text, usage, headers, raw) + case "messages", False: + return Outcome( + status, + call_id, + upstream_id_of(str(payloads[0]["id"])), + messages_text(payloads[0]), + usage_of(payloads[0]), + headers, + raw, + ) + raise AssertionError((endpoint, stream)) + + +def send(gateway: Gateway, endpoint: Endpoint, body: Mapping[str, JsonValue], *, key: str | None = None) -> Outcome: + stream: Final = body.get("stream") is True + headers: Final = {"Authorization": f"Bearer {gateway.key if key is None else key}"} + with gateway.client.stream("POST", endpoint_path(endpoint), json=body, headers=headers, timeout=60) as response: + lines: Final = tuple(line for line in response.iter_lines() if line) + return parse_outcome(endpoint, stream=stream, status=response.status_code, headers=response.headers, lines=lines) + + +def send_raw(gateway: Gateway, endpoint: Endpoint, content: str, *, key: str | None = None) -> Outcome: + headers: Final = { + "Authorization": f"Bearer {gateway.key if key is None else key}", + "Content-Type": "application/json", + } + response: Final = gateway.client.post( + endpoint_path(endpoint), content=content.encode(), headers=headers, timeout=60 + ) + lines: Final = tuple(line for line in response.text.splitlines() if line) + return parse_outcome(endpoint, stream=False, status=response.status_code, headers=response.headers, lines=lines) + + +def matching_rows(name: str, markers: frozenset[str]) -> tuple[dict[str, JsonValue], ...]: + rows: Final = read_rows( + 'SELECT request_id, status, spend, prompt_tokens, completion_tokens, cache_hit FROM "LiteLLM_SpendLogs" ' + "WHERE model_group = %s", + (name,), + ) + return tuple(row for row in rows if any(marker in row_key(str(row["request_id"])) for marker in markers)) + + +def spend_rows( + name: str, markers: frozenset[str], *, expected: int, seconds: float = 90 +) -> tuple[dict[str, JsonValue], ...]: + return eventually(lambda: matching_rows(name, markers), lambda found: len(found) >= expected, seconds=seconds) + + +def success_row(name: str, *needles: str) -> dict[str, JsonValue]: + (row,) = tuple(row for row in spend_rows(name, frozenset(needles), expected=1) if row["cache_hit"] != "True") + assert row["status"] == "success", row + return row + + +def failure_row(call_id: str) -> dict[str, JsonValue]: + assert call_id, "No call id to look the failure row up by" + rows: Final = eventually( + lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (call_id,)), + lambda found: len(found) >= 1, + seconds=90, + ) + (row,) = tuple(rows) + assert row["status"] == "failure", row + return row + + +def assert_priced_row(row: Mapping[str, JsonValue], model: str) -> None: + assert row["prompt_tokens"] == INPUT_TOKENS, row + assert row["completion_tokens"] == OUTPUT_TOKENS, row + assert math.isclose(float(str(row["spend"])), expected_spend(model), rel_tol=1e-9), (row, expected_spend(model)) + + +def assert_usage(endpoint: Endpoint, usage: Mapping[str, JsonValue]) -> None: + match endpoint: + case "chat": + assert usage["prompt_tokens"] == INPUT_TOKENS, usage + assert usage["completion_tokens"] == OUTPUT_TOKENS, usage + details: Final = JSON.validate_python(usage["prompt_tokens_details"]) + assert details["cached_tokens"] == CACHED_TOKENS, usage + assert details["cache_write_tokens"] == WRITTEN_TOKENS, usage + assert details["cache_creation_tokens"] == WRITTEN_TOKENS, usage + case "responses": + assert usage["input_tokens"] == INPUT_TOKENS, usage + assert usage["output_tokens"] == OUTPUT_TOKENS, usage + input_details: Final = JSON.validate_python(usage["input_tokens_details"]) + assert input_details["cached_tokens"] == CACHED_TOKENS, usage + assert input_details["cache_write_tokens"] == WRITTEN_TOKENS, usage + case "messages": + assert usage["input_tokens"] == UNCACHED_TOKENS, usage + assert usage["output_tokens"] == OUTPUT_TOKENS, usage + assert usage["cache_read_input_tokens"] == CACHED_TOKENS, usage + assert usage["cache_creation_input_tokens"] == WRITTEN_TOKENS, usage + + +def assert_answered(outcome: Outcome, marker: str) -> None: + assert outcome.status == 200, (outcome.status, outcome.raw) + assert outcome.text == answer(marker), outcome.raw + + +def deployment( + scenario: Scenario, + wire: Wire, + model: str, + *, + model_info: Mapping[str, JsonValue] | None = None, + points: JsonValue = SYSTEM_POINT, + **extra: JsonValue, +) -> str: + return scenario.model( + model_info=model_info, + model=model, + api_base=wire.url, + api_key=TOKEN, + aws_region_name="us-east-1", + cache_control_injection_points=points, + **extra, + ) + + +def settled(gateway: Gateway, name: str, wire: Wire, *, accepted: frozenset[int] = frozenset({200})) -> None: + eventually( + lambda: tuple( + gateway.request( + "POST", "/v1/chat/completions", request_body("chat", name, prompt_text(fresh_marker()), cache=NO_CACHE) + ).status_code + for _ in range(12) + ), + lambda codes: all(code in accepted for code in codes), + seconds=90, + ) + wire.drain() + + +def mantle_deployment( + gateway: Gateway, + scenario: Scenario, + wire: Wire, + model: str = GPT, + *, + model_info: Mapping[str, JsonValue] | None = None, + **extra: JsonValue, +) -> str: + name: Final = deployment(scenario, wire, model, model_info=model_info, **extra) + settled(gateway, name, wire) + return name + + +def observe( + gateway: Gateway, wire: Wire, endpoint: Endpoint, body: Mapping[str, JsonValue], *, key: str | None = None +) -> tuple[Outcome, Request]: + wire.drain() + outcome: Final = send(gateway, endpoint, body, key=key) + return outcome, only_received(wire) + + +def marked_cell(gateway: Gateway, endpoint: Endpoint, *, stream: bool) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome, received = observe( + gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker), stream=stream) + ) + assert_answered(outcome, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=stream) + assert_usage(endpoint, outcome.usage) + row: Final = success_row(name, marker) + assert_priced_row(row, GPT) + if caller_sees_upstream_id(endpoint, stream=stream): + assert row_key(str(row["request_id"])) == outcome.response_id, (row, outcome.response_id) + if not stream: + assert math.isclose(float(outcome.headers["x-litellm-response-cost"]), expected_spend(GPT), rel_tol=1e-9), ( + outcome.headers + ) + + +def openai_client(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI( + base_url=f"{gateway.client.base_url}/v1", + api_key=gateway.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=60), + ) + + +def async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI( + base_url=f"{gateway.client.base_url}/v1", + api_key=gateway.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=60), + ) + + +def anthropic_client(gateway: Gateway) -> anthropic.Anthropic: + return anthropic.Anthropic( + base_url=str(gateway.client.base_url), + api_key=gateway.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=60), + ) + + +def async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic( + base_url=str(gateway.client.base_url), + api_key=gateway.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=60), + ) + + +class Dumpable(Protocol): + def model_dump(self, *, exclude_none: bool = ...) -> Mapping[str, object]: ... + + +def dump_model(value: Dumpable) -> dict[str, JsonValue]: + return JSON.validate_python(value.model_dump(exclude_none=True)) + + +def sdk_outcome(endpoint: Endpoint, payloads: Sequence[Mapping[str, JsonValue]], *, stream: bool) -> Outcome: + match endpoint, stream: + case "chat", True: + identity, text, usage = chat_stream_fields(payloads) + return Outcome(200, "", upstream_id_of(identity), text, usage, {}, "") + case "chat", False: + return Outcome( + 200, "", upstream_id_of(str(payloads[0]["id"])), chat_text(payloads[0]), usage_of(payloads[0]), {}, "" + ) + case "responses", _: + completed: Final = completed_response(payloads) if stream else dict(payloads[0]) + return Outcome( + 200, + "", + upstream_id_of(str(completed["id"])), + responses_text(completed), + usage_of(completed), + {}, + "", + ) + case "messages", True: + identity, text, usage = messages_stream_fields(payloads) + return Outcome(200, "", upstream_id_of(identity), text, usage, {}, "") + case "messages", False: + return Outcome( + 200, + "", + upstream_id_of(str(payloads[0]["id"])), + messages_text(payloads[0]), + usage_of(payloads[0]), + {}, + "", + ) + raise AssertionError((endpoint, stream)) + + +def sdk_call(gateway: Gateway, endpoint: Endpoint, name: str, prompt: str, *, stream: bool) -> Outcome: + match endpoint, stream: + case "chat", False: + completion: Final = openai_client(gateway).chat.completions.create( + model=name, + messages=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_tokens=MAX_TOKENS, + ) + return sdk_outcome(endpoint, (dump_model(completion),), stream=False) + case "chat", True: + chunks: Final = openai_client(gateway).chat.completions.create( + model=name, + messages=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_tokens=MAX_TOKENS, + stream=True, + stream_options={"include_usage": True}, + ) + return sdk_outcome(endpoint, tuple(dump_model(chunk) for chunk in chunks), stream=True) + case "responses", False: + response: Final = openai_client(gateway).responses.create( + model=name, + input=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_output_tokens=MAX_TOKENS, + ) + return sdk_outcome(endpoint, (dump_model(response),), stream=False) + case "responses", True: + events: Final = openai_client(gateway).responses.create( + model=name, + input=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_output_tokens=MAX_TOKENS, + stream=True, + ) + return sdk_outcome(endpoint, tuple(dump_model(event) for event in events), stream=True) + case "messages", False: + message: Final = anthropic_client(gateway).messages.create( + model=name, system=SYSTEM, messages=[{"role": "user", "content": prompt}], max_tokens=MAX_TOKENS + ) + return sdk_outcome(endpoint, (dump_model(message),), stream=False) + case "messages", True: + with anthropic_client(gateway).messages.stream( + model=name, system=SYSTEM, messages=[{"role": "user", "content": prompt}], max_tokens=MAX_TOKENS + ) as events_stream: + raw_events: Final = tuple(dump_model(event) for event in events_stream) + return sdk_outcome(endpoint, raw_events, stream=True) + raise AssertionError((endpoint, stream)) + + +async def async_sdk_call(gateway: Gateway, endpoint: Endpoint, name: str, prompt: str) -> Outcome: + match endpoint: + case "chat": + completion: Final = await async_openai_client(gateway).chat.completions.create( + model=name, + messages=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_tokens=MAX_TOKENS, + ) + return sdk_outcome(endpoint, (dump_model(completion),), stream=False) + case "responses": + response: Final = await async_openai_client(gateway).responses.create( + model=name, + input=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_output_tokens=MAX_TOKENS, + ) + return sdk_outcome(endpoint, (dump_model(response),), stream=False) + case "messages": + message: Final = await async_anthropic_client(gateway).messages.create( + model=name, system=SYSTEM, messages=[{"role": "user", "content": prompt}], max_tokens=MAX_TOKENS + ) + return sdk_outcome(endpoint, (dump_model(message),), stream=False) + + +def assert_marked_sdk_cell( + wire: Wire, endpoint: Endpoint, name: str, marker: str, outcome: Outcome, *, stream: bool +) -> None: + assert_answered(outcome, marker) + assert_wire( + only_received(wire), expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=stream + ) + assert_usage(endpoint, outcome.usage) + assert_priced_row(success_row(name, marker), GPT) + + +HOSTILE_OPTIONS: Final[dict[str, JsonValue]] = { + "int": 5, + "list": [{"mode": "explicit"}], + "empty-string": "", + "5kb-string": "x" * 5120, +} + + +MALFORMED_POINTS: Final[dict[str, JsonValue]] = { + "null": None, + "string": "system", + "int": 5, + "dict": {"location": "message", "role": "system"}, + "string-list": ["system"], + "no-location": [{"role": "system"}], +} +MIXED_POINTS: Final[list[JsonValue]] = ["system", *SYSTEM_POINT, 3] + + +Step = tuple[Endpoint, bool, str] + + +def plan_burst(size: int) -> tuple[Step, ...]: + return tuple((ENDPOINTS[index % 3], index % 3 == 0, fresh_marker()) for index in range(size)) + + +def burst(gateway: Gateway, name: str, size: int) -> tuple[tuple[str, Outcome], ...]: + def call(step: Step) -> tuple[str, Outcome]: + endpoint, stream, marker = step + return marker, send(gateway, endpoint, request_body(endpoint, name, prompt_text(marker), stream=stream)) + + with ThreadPoolExecutor(max_workers=10) as pool: + return tuple(pool.map(call, plan_burst(size))) + + +def assert_marked(request: Request, *, marked: bool) -> None: + body: Final = body_of(request) + assert breakpoint_count(body) == (1 if marked else 0), request.body + assert ("prompt_cache_options" in body) is marked, request.body + + +def assert_burst_landed(wire: Wire, name: str, outcomes: Sequence[tuple[str, Outcome]], *, marked: bool) -> None: + for marker, outcome in outcomes: + assert_answered(outcome, marker) + identities: Final = frozenset(outcome.response_id for _, outcome in outcomes) + assert len(identities) == len(outcomes), identities + received: Final = wire.drain() + assert len(received) == len(outcomes), (len(received), len(outcomes)) + for request in received: + assert_marked(request, marked=marked) + markers: Final = frozenset(marker for marker, _ in outcomes) + rows: Final = spend_rows(name, markers, expected=len(outcomes), seconds=120) + assert sorted(row_key(str(row["request_id"])) for row in rows) == sorted(identities), rows diff --git a/tests/integration/providers/test_bedrock_mantle_gpt_prompt_cache_breakpoint_wire.py b/tests/integration/providers/test_bedrock_mantle_gpt_prompt_cache_breakpoint_wire.py new file mode 100644 index 00000000000..2af41c2b683 --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_gpt_prompt_cache_breakpoint_wire.py @@ -0,0 +1,531 @@ +import asyncio +import json +from typing import Final +from urllib.parse import urlsplit + +import pytest +from integration._support.client import Gateway, eventually +from integration._support.wire import Request, wire_server +from integration.providers._mantle_gpt_prompt_cache_support import ( + ANTHROPIC_PATH, + AZURE, + BURST, + CLAUDE, + ENDPOINTS, + EPHEMERAL, + EXPLICIT, + GPT, + GPT_FLAGGED_ROW, + GPT_BARE, + GPT_REGION, + GPT_UNFLAGGED_ROW, + HOSTILE_OPTIONS, + IMPLICIT, + ITEMS, + MALFORMED_POINTS, + MAX_TOKENS, + MIXED_POINTS, + NO_CACHE, + ODD_FLAGS, + SYSTEM, + SYSTEM_POINT, + THIRD_PARTY, + Endpoint, + Outcome, + assert_answered, + assert_burst_landed, + assert_marked_sdk_cell, + assert_priced_row, + assert_wire, + async_sdk_call, + body_of, + breakpoint_count, + burst, + deployment, + expected_wire, + failing_peer, + failure_row, + fresh_marker, + mantle_deployment, + mantle_peer, + marked_cell, + observe, + only_received, + openai_shaped_peer, + prompt_text, + request_body, + row_key, + sdk_call, + send, + send_raw, + settled, + spend_rows, + success_row, + system_item, +) +from pydantic import JsonValue + +pytestmark = pytest.mark.timeout(600) + + +@pytest.mark.parametrize("stream", [False, True], ids=["json", "stream"]) +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_a_configured_system_point_reaches_mantle_as_a_breakpoint_over_httpx( + gateway: Gateway, endpoint: Endpoint, stream: bool +) -> None: + marked_cell(gateway, endpoint, stream=stream) + + +@pytest.mark.parametrize("stream", [False, True], ids=["json", "stream"]) +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_a_the_sync_sdks_see_the_breakpoint_and_the_mapped_cache_usage( + gateway: Gateway, endpoint: Endpoint, stream: bool +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome: Final = sdk_call(gateway, endpoint, name, prompt_text(marker), stream=stream) + assert_marked_sdk_cell(wire, endpoint, name, marker, outcome, stream=stream) + + +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_a_the_async_sdks_see_the_breakpoint_and_the_mapped_cache_usage(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome: Final = asyncio.run(async_sdk_call(gateway, endpoint, name, prompt_text(marker))) + assert_marked_sdk_cell(wire, endpoint, name, marker, outcome, stream=False) + + +@pytest.mark.parametrize("endpoint", ["responses", "chat"]) +def test_b1_b2_a_pinned_explicit_mode_reaches_the_wire_with_the_breakpoint( + gateway: Gateway, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, prompt_cache_options=EXPLICIT) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True, options=EXPLICIT), + streaming=False, + ) + assert_priced_row(success_row(name, marker), GPT) + + +@pytest.mark.parametrize("endpoint", ["responses", "chat"]) +def test_b3_b4_a_region_prefixed_deployment_reads_the_region_free_row(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, GPT_REGION) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + assert_priced_row(success_row(name, marker), GPT) + + +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_b15_to_b17_a_bare_deployment_name_with_its_provider_reads_the_provider_keyed_row( + gateway: Gateway, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, GPT_BARE, custom_llm_provider="bedrock_mantle") + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + assert_priced_row(success_row(name, marker), GPT) + + +def assert_anthropic_marked_wire(received: Request, marker: str) -> None: + body: Final = body_of(received) + assert urlsplit(received.target).path == ANTHROPIC_PATH, received.target + assert body["system"] == [{"type": "text", "text": SYSTEM, "cache_control": EPHEMERAL}], received.body + assert body["messages"] == [{"role": "user", "content": [{"type": "text", "text": prompt_text(marker)}]}], ( + received.body + ) + assert "prompt_cache_options" not in body, received.body + assert "prompt_cache_breakpoint" not in received.body.decode(), received.body + + +@pytest.mark.parametrize("endpoint", ["messages", "chat"]) +def test_b5_b6_claude_on_mantle_keeps_the_anthropic_dialect(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, CLAUDE) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_anthropic_marked_wire(received, marker) + success_row(name, marker, outcome.response_id) + + +def test_b7_a_deployment_flag_true_opts_an_unflagged_mantle_row_in(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment( + gateway, scenario, wire, GPT_UNFLAGGED_ROW, model_info={"supports_prompt_cache_breakpoint": True} + ) + outcome, received = observe(gateway, wire, "responses", request_body("responses", name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire(GPT_UNFLAGGED_ROW, prompt_text(marker), endpoint="responses", marked=True), + streaming=False, + ) + success_row(name, marker) + + +def test_b8_a_deployment_flag_false_opts_a_flagged_mantle_row_out(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment( + gateway, scenario, wire, GPT_FLAGGED_ROW, model_info={"supports_prompt_cache_breakpoint": False} + ) + outcome, received = observe(gateway, wire, "responses", request_body("responses", name, prompt_text(marker))) + assert_answered(outcome, marker) + body: Final = body_of(received) + assert breakpoint_count(body) == 0, received.body + assert "prompt_cache_options" not in body, received.body + assert "cache_control" not in received.body.decode(), received.body + success_row(name, marker) + + +@pytest.mark.parametrize("endpoint", ["chat", "responses"]) +def test_b13_azure_openai_stays_ineligible_for_the_breakpoint_dialect(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(openai_shaped_peer()) as wire, gateway.scenario() as scenario: + name: Final = scenario.model( + model=AZURE, + api_base=wire.url, + api_key="synthetic-azure-key", + api_version="2025-04-01-preview", + cache_control_injection_points=SYSTEM_POINT, + ) + settled(gateway, name, wire) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert urlsplit(received.target).path.startswith("/openai/"), received.target + assert "prompt_cache_breakpoint" not in received.body.decode(), received.body + assert "prompt_cache_options" not in body_of(received), received.body + assert SYSTEM in received.body.decode(), received.body + success_row(name, marker) + + +def test_b14_an_openai_entry_on_a_third_party_host_keeps_the_anthropic_dialect(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(openai_shaped_peer()) as wire, gateway.scenario() as scenario: + name: Final = scenario.model( + model=THIRD_PARTY, + api_base=wire.url, + api_key="synthetic-third-party-key", + cache_control_injection_points=SYSTEM_POINT, + ) + settled(gateway, name, wire) + outcome, received = observe(gateway, wire, "chat", request_body("chat", name, prompt_text(marker))) + assert_answered(outcome, marker) + body: Final = body_of(received) + assert urlsplit(received.target).path == "/chat/completions", received.target + assert body["messages"] == [ + {"role": "system", "content": SYSTEM, "cache_control": EPHEMERAL}, + {"role": "user", "content": prompt_text(marker)}, + ], received.body + assert "prompt_cache_options" not in body, received.body + success_row(name, marker) + + +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_c_a_response_cache_hit_serves_the_marked_request_again(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + body: Final = request_body(endpoint, name, prompt_text(marker)) + first, received = observe(gateway, wire, endpoint, body) + assert_answered(first, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + sends: Final[list[Outcome]] = [] + + def resend() -> Outcome: + served: Final = send(gateway, endpoint, body) + sends.append(served) + return served + + hit: Final = eventually( + resend, lambda served: served.status == 200 and served.response_id == first.response_id, seconds=30 + ) + assert_answered(hit, marker) + misses: Final = wire.drain() + assert len(misses) == len(sends) - 1, (len(misses), len(sends)) + for miss in misses: + assert_wire(miss, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + rows: Final = spend_rows(name, frozenset({marker}), expected=len(sends) + 1, seconds=70) + assert len(rows) == len(sends) + 1, rows + (hit_row,) = tuple(row for row in rows if row["cache_hit"] == "True") + assert row_key(str(hit_row["request_id"])) == first.response_id, (hit_row, first.response_id) + + +@pytest.mark.parametrize("endpoint", ["chat", "responses"]) +@pytest.mark.parametrize("shape", sorted(HOSTILE_OPTIONS)) +def test_d1_to_d4_hostile_prompt_cache_options_reach_the_wire_verbatim_and_the_providers_400_reaches_the_caller( + gateway: Gateway, endpoint: Endpoint, shape: str +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + hostile: Final = HOSTILE_OPTIONS[shape] + outcome, received = observe( + gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker), prompt_cache_options=hostile) + ) + assert outcome.status == 400, (outcome.status, outcome.raw) + assert "Invalid prompt_cache_options" in outcome.raw, outcome.raw + assert_wire( + received, + expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True, options=hostile), + streaming=False, + ) + failure_row(outcome.call_id) + + +@pytest.mark.parametrize("endpoint", ["chat", "responses"]) +def test_d5_a_duplicated_prompt_cache_options_key_resolves_to_the_last_value( + gateway: Gateway, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + encoded: Final = json.dumps(request_body(endpoint, name, prompt_text(marker), prompt_cache_options=IMPLICIT)) + duplicated: Final = encoded[:-1] + ', "prompt_cache_options": {"mode": "explicit"}}' + wire.drain() + outcome: Final = send_raw(gateway, endpoint, duplicated) + assert_answered(outcome, marker) + assert_wire( + only_received(wire), + expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True, options=EXPLICIT), + streaming=False, + ) + success_row(name, marker) + + +@pytest.mark.parametrize("endpoint", ["chat", "responses"]) +def test_d6_a_null_prompt_cache_options_is_treated_as_unset(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome, received = observe( + gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker), prompt_cache_options=None) + ) + assert_answered(outcome, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + success_row(name, marker) + + +def test_d7_an_unauthenticated_hostile_request_never_reaches_the_wire(gateway: Gateway) -> None: + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + wire.drain() + outcome: Final = send( + gateway, + "responses", + request_body("responses", name, prompt_text(fresh_marker()), prompt_cache_options=5), + key="sk-not-a-key", + ) + assert outcome.status == 401, (outcome.status, outcome.raw) + assert wire.drain() == (), "the upstream saw an unauthenticated request" + + +@pytest.mark.parametrize("shape", sorted(MALFORMED_POINTS)) +def test_d8_to_d10_malformed_injection_points_never_crash_the_request(gateway: Gateway, shape: str) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = deployment(scenario, wire, GPT, points=MALFORMED_POINTS[shape]) + settled(gateway, name, wire) + outcome, received = observe(gateway, wire, "responses", request_body("responses", name, prompt_text(marker))) + assert_answered(outcome, marker) + body: Final = body_of(received) + assert breakpoint_count(body) == 0, received.body + assert "prompt_cache_options" not in body, received.body + success_row(name, marker) + + +@pytest.mark.parametrize( + ("status", "endpoint"), + [(400, "chat"), (400, "responses"), (401, "responses")], + ids=["400-chat", "400-responses", "401-responses"], +) +def test_d11_d12_a_provider_error_on_the_marked_request_reaches_the_caller_after_one_attempt( + gateway: Gateway, status: int, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + message: Final = f"scripted provider failure {marker}" + with wire_server(failing_peer(status, message, "scripted_failure")) as wire, gateway.scenario() as scenario: + name: Final = deployment(scenario, wire, GPT) + settled(gateway, name, wire, accepted=frozenset({status})) + outcome: Final = send(gateway, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert outcome.status == status, (outcome.status, outcome.raw) + assert message in outcome.raw, outcome.raw + attempts: Final = wire.drain() + assert len(attempts) == 1, [attempt.target for attempt in attempts] + assert breakpoint_count(body_of(attempts[0])) == 1, attempts[0].body + failure_row(outcome.call_id) + + +@pytest.mark.parametrize(("label", "model", "flag"), ODD_FLAGS, ids=[label for label, _, _ in ODD_FLAGS]) +def test_d13_an_odd_typed_deployment_flag_never_opts_a_mantle_row_in( + gateway: Gateway, label: str, model: str, flag: JsonValue +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment( + gateway, scenario, wire, model, model_info={"supports_prompt_cache_breakpoint": flag} + ) + outcome, received = observe(gateway, wire, "responses", request_body("responses", name, prompt_text(marker))) + assert_answered(outcome, marker) + body: Final = body_of(received) + assert breakpoint_count(body) == 0, (label, received.body) + assert "prompt_cache_options" not in body, (label, received.body) + success_row(name, marker) + + +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_d14_a_point_beside_junk_entries_still_marks_the_anthropic_dialect( + gateway: Gateway, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, CLAUDE, points=MIXED_POINTS) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_anthropic_marked_wire(received, marker) + success_row(name, marker, outcome.response_id) + + +def test_e1_client_breakpoint_prompt_cache_key_and_explicit_mode_pass_through_with_no_second_breakpoint( + gateway: Gateway, +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + body: Final[dict[str, JsonValue]] = { + "model": name, + "input": [system_item(marked=True, endpoint="responses"), {"role": "user", "content": prompt_text(marker)}], + "max_output_tokens": MAX_TOKENS, + "prompt_cache_key": f"key-{marker}", + "prompt_cache_options": EXPLICIT, + } + outcome, received = observe(gateway, wire, "responses", body) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire( + GPT, + prompt_text(marker), + endpoint="responses", + marked=True, + options=EXPLICIT, + prompt_cache_key=f"key-{marker}", + ), + streaming=False, + ) + success_row(name, marker) + + +def test_e2_four_client_breakpoints_leave_no_slot_for_the_configured_point(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + turns: Final = tuple(f"Earlier turn {index} marker-{marker}." for index in range(3)) + body: Final[dict[str, JsonValue]] = { + "model": name, + "input": [ + {"role": "system", "content": SYSTEM}, + *( + { + "role": "user", + "content": [{"type": "input_text", "text": turn, "prompt_cache_breakpoint": EXPLICIT}], + } + for turn in turns + ), + { + "role": "user", + "content": [ + {"type": "input_text", "text": prompt_text(marker), "prompt_cache_breakpoint": EXPLICIT} + ], + }, + ], + "max_output_tokens": MAX_TOKENS, + } + outcome, received = observe(gateway, wire, "responses", body) + assert_answered(outcome, marker) + wire_body: Final = body_of(received) + items: Final = ITEMS.validate_python(wire_body["input"]) + assert breakpoint_count(wire_body) == 4, received.body + assert items[0]["role"] == "system" and "prompt_cache_breakpoint" not in json.dumps(items[0]), items[0] + assert "prompt_cache_options" not in wire_body, received.body + success_row(name, marker) + + +def test_e3_a_client_implicit_mode_on_a_pinned_explicit_deployment_wins_on_the_wire(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, prompt_cache_options=EXPLICIT) + outcome, received = observe( + gateway, + wire, + "responses", + request_body("responses", name, prompt_text(marker), prompt_cache_options=IMPLICIT), + ) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire(GPT, prompt_text(marker), endpoint="responses", marked=True, options=IMPLICIT), + streaming=False, + ) + success_row(name, marker) + + +def test_e4_an_empty_prompt_cache_options_object_is_kept_as_the_clients_value(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome, received = observe( + gateway, wire, "responses", request_body("responses", name, prompt_text(marker), prompt_cache_options={}) + ) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire(GPT, prompt_text(marker), endpoint="responses", marked=True, options={}), + streaming=False, + ) + success_row(name, marker) + + +def test_e5_the_same_request_twice_writes_two_rows_and_two_marked_upstream_requests(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + body: Final = request_body("chat", name, prompt_text(marker), cache=NO_CACHE) + first, first_received = observe(gateway, wire, "chat", body) + second, second_received = observe(gateway, wire, "chat", body) + assert_answered(first, marker) + assert_answered(second, marker) + assert first.response_id != second.response_id, (first.response_id, second.response_id) + for received in (first_received, second_received): + assert_wire( + received, expected_wire(GPT, prompt_text(marker), endpoint="chat", marked=True), streaming=False + ) + rows: Final = spend_rows(name, frozenset({marker}), expected=2) + assert {row_key(str(row["request_id"])) for row in rows} == {first.response_id, second.response_id}, rows + + +def test_f1_a_mixed_burst_marks_every_request_and_lands_every_id_once(gateway: Gateway) -> None: + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + wire.drain() + assert_burst_landed(wire, name, burst(gateway, name, BURST), marked=True) + + +def test_f4_a_slow_mantle_stream_during_a_burst_completes_with_every_id_landing_once(gateway: Gateway) -> None: + with wire_server(mantle_peer(pause=0.3)) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + wire.drain() + assert_burst_landed(wire, name, burst(gateway, name, 12), marked=True) diff --git a/tests/integration/providers/test_openai_dialect_prompt_cache_breakpoint_owned_proxy.py b/tests/integration/providers/test_openai_dialect_prompt_cache_breakpoint_owned_proxy.py new file mode 100644 index 00000000000..cc7ab06da90 --- /dev/null +++ b/tests/integration/providers/test_openai_dialect_prompt_cache_breakpoint_owned_proxy.py @@ -0,0 +1,385 @@ +import asyncio +import re +import signal +import threading +from collections.abc import Callable, Iterator +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process +from integration._support.responses_vendor import newest_marker +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.providers._cache_control_marks_support import owned_config +from integration.providers._mantle_gpt_prompt_cache_support import ( + EXPLICIT, + GPT, + IMPLICIT, + SYSTEM, + SYSTEM_POINT, + TOKEN, + Outcome, + assert_answered, + assert_burst_landed, + assert_wire, + body_of, + breakpoint_count, + expected_wire, + fresh_marker, + mantle_deployment, + mantle_peer, + observe, + openai_shaped_peer, + parse_outcome, + plan_burst, + prompt_text, + request_body, + row_key, + send, + settled, + spend_rows, + success_row, +) +from pydantic import JsonValue + +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_FOUNDRY_BASE: Final = "http://prompt-cache-breakpoint-audit.services.ai.azure.com" +_FOUNDRY_HOST: Final = urlsplit(_FOUNDRY_BASE).netloc +_FOUNDRY_MODEL: Final = "gpt-6-astra" +_FOUNDRY_KEY: Final = "synthetic-foundry-key" +_FOUNDRY_CHAT_PATH: Final = "/models/chat/completions" +_CELL_TIMEOUT: Final = int(2 * graceful_stop_seconds() + 120) +_RESTART_TIMEOUT: Final = int(4 * graceful_stop_seconds() + 240) +_HELD_BURST: Final = 20 +_SPREAD_BURST: Final = 12 +_SPREAD_ATTEMPTS: Final = 20 +_MANTLE_NAME: Final = "mantle-gpt-cache-owned" + + +def _peer(release: threading.Event, held: SimpleQueue[str]) -> Callable[[Request], Reply]: + mantle: Final = mantle_peer() + foundry: Final = openai_shaped_peer() + + def respond(request: Request) -> Reply: + if urlsplit(request.target).netloc == _FOUNDRY_HOST: + return foundry(request) + marker: Final = newest_marker(request.body.decode()) + assert marker is not None, request.body + held.put(marker) + assert release.wait(timeout=120), "The held burst was never released" + return mantle(request) + + return respond + + +@dataclass(frozen=True, slots=True) +class _Rig: + gateway: Gateway + wire: Wire + owned: OwnedProxy + release: threading.Event + held: SimpleQueue[str] + + +def _overrides(wire: Wire) -> MappingProxyType[str, str]: + return MappingProxyType({"HTTP_PROXY": wire.url, "NO_PROXY": "127.0.0.1,localhost", "AIOHTTP_TRUST_ENV": "True"}) + + +def _started_workers(log: Path) -> tuple[int, ...]: + return tuple(int(match.group(1)) for match in _STARTED_WORKER.finditer(log.read_text())) + + +def _live_workers(log: Path) -> tuple[int, ...]: + return tuple(pid for pid in _started_workers(log) if psutil.pid_exists(pid)) + + +def _open_peer_connections(pid: int, peer_url: str) -> int: + port: Final = urlsplit(peer_url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +def _drain(queue: SimpleQueue[str]) -> None: + while not queue.empty(): + queue.get_nowait() + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]: + directory: Final = tmp_path_factory.mktemp("openai-dialect-prompt-cache-breakpoint") + release: Final = threading.Event() + release.set() + held: Final[SimpleQueue[str]] = SimpleQueue() + with gateway_from_environment() as environment, wire_server(_peer(release, held)) as wire: + with owned_proxy_process(environment, directory, _overrides(wire), workers=2) as owned: + eventually(lambda: len(_started_workers(owned.log)), lambda count: count == 2, seconds=60) + wire.drain() + yield _Rig(owned.gateway, wire, owned, release, held) + + +def _foundry_deployment(rig: _Rig, scenario: Scenario, *, bare: bool) -> str: + name: Final = ( + scenario.model( + model=_FOUNDRY_MODEL, + custom_llm_provider="azure_ai", + api_base=_FOUNDRY_BASE, + api_key=_FOUNDRY_KEY, + cache_control_injection_points=SYSTEM_POINT, + ) + if bare + else scenario.model( + model=f"azure_ai/{_FOUNDRY_MODEL}", + api_base=_FOUNDRY_BASE, + api_key=_FOUNDRY_KEY, + cache_control_injection_points=SYSTEM_POINT, + ) + ) + settled(rig.gateway, name, rig.wire) + return name + + +def _assert_foundry_chat_wire(received: Request, prompt: str, *, options: JsonValue | None) -> None: + body: Final = body_of(received) + assert urlsplit(received.target).netloc == _FOUNDRY_HOST, received.target + assert urlsplit(received.target).path == _FOUNDRY_CHAT_PATH, received.target + assert "prompt_cache_breakpoint" not in received.body.decode(), received.body + assert body["model"] == _FOUNDRY_MODEL, received.body + assert body["messages"] == [{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], ( + received.body + ) + assert body.get("prompt_cache_options") == options, received.body + + +@pytest.mark.timeout(_CELL_TIMEOUT) +@pytest.mark.parametrize("bare", [False, True], ids=["prefixed", "bare-provider-field"]) +def test_b9_b12_foundry_gpt6_responses_reach_the_openai_dialect_on_the_foundry_host(rig: _Rig, bare: bool) -> None: + marker: Final = fresh_marker() + with rig.gateway.scenario() as scenario: + name: Final = _foundry_deployment(rig, scenario, bare=bare) + outcome, received = observe( + rig.gateway, rig.wire, "responses", request_body("responses", name, prompt_text(marker)) + ) + assert_answered(outcome, marker) + assert urlsplit(received.target).netloc == _FOUNDRY_HOST, received.target + assert_wire( + received, + expected_wire(_FOUNDRY_MODEL, prompt_text(marker), endpoint="responses", marked=True), + streaming=False, + ) + success_row(name, marker) + + +@pytest.mark.timeout(_CELL_TIMEOUT) +def test_b10_foundry_gpt6_messages_bridge_to_chat_without_a_breakpoint(rig: _Rig) -> None: + marker: Final = fresh_marker() + with rig.gateway.scenario() as scenario: + name: Final = _foundry_deployment(rig, scenario, bare=False) + outcome, received = observe( + rig.gateway, rig.wire, "messages", request_body("messages", name, prompt_text(marker)) + ) + assert_answered(outcome, marker) + _assert_foundry_chat_wire(received, prompt_text(marker), options=IMPLICIT) + success_row(name, marker) + + +@pytest.mark.timeout(_CELL_TIMEOUT) +def test_b11_foundry_gpt6_chat_stays_on_the_plain_chat_wire(rig: _Rig) -> None: + marker: Final = fresh_marker() + with rig.gateway.scenario() as scenario: + name: Final = _foundry_deployment(rig, scenario, bare=False) + outcome, received = observe(rig.gateway, rig.wire, "chat", request_body("chat", name, prompt_text(marker))) + assert_answered(outcome, marker) + _assert_foundry_chat_wire(received, prompt_text(marker), options=IMPLICIT) + success_row(name, marker) + + +def _spread_burst(rig: _Rig, name: str, workers: tuple[int, ...]) -> dict[int, int]: + plan: Final = plan_burst(_SPREAD_BURST) + _drain(rig.held) + rig.wire.drain() + rig.release.clear() + with ThreadPoolExecutor(max_workers=_SPREAD_BURST) as pool: + futures: Final = tuple( + pool.submit(send, rig.gateway, endpoint, request_body(endpoint, name, prompt_text(marker), stream=stream)) + for endpoint, stream, marker in plan + ) + eventually(rig.held.qsize, lambda size: size == _SPREAD_BURST, seconds=60) + held_by: Final = {pid: _open_peer_connections(pid, rig.wire.url) for pid in workers} + rig.release.set() + outcomes: Final = tuple(future.result() for future in futures) + assert sum(held_by.values()) == _SPREAD_BURST, held_by + assert_burst_landed( + rig.wire, name, tuple((marker, outcome) for (_, _, marker), outcome in zip(plan, outcomes)), marked=True + ) + return held_by + + +@pytest.mark.timeout(_CELL_TIMEOUT) +def test_e6_every_worker_of_a_two_worker_proxy_marks_the_mantle_request(rig: _Rig) -> None: + with rig.gateway.scenario() as scenario: + name: Final = mantle_deployment(rig.gateway, scenario, rig.wire) + workers: Final = _live_workers(rig.owned.log) + assert len(workers) == 2, workers + for _ in range(_SPREAD_ATTEMPTS): + if all(count > 0 for count in _spread_burst(rig, name, workers).values()): + break + else: + raise AssertionError("One worker never took a marked request") + + +def _mantle_config(wire: Wire, directory: Path, **litellm_params: JsonValue) -> Path: + return owned_config( + directory, + [ + { + "model_name": _MANTLE_NAME, + "litellm_params": { + "model": GPT, + "api_base": wire.url, + "api_key": TOKEN, + "aws_region_name": "us-east-1", + "cache_control_injection_points": SYSTEM_POINT, + **litellm_params, + }, + } + ], + ) + + +async def _one(client: httpx.AsyncClient, key: str, body: dict[str, JsonValue]) -> Outcome: + response: Final = await client.post("/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {key}"}) + lines: Final = tuple(line for line in response.text.splitlines() if line) + return parse_outcome("chat", stream=False, status=response.status_code, headers=response.headers, lines=lines) + + +async def _held_burst(url: str, key: str, markers: tuple[str, ...]) -> tuple[tuple[str, Outcome], ...]: + async with httpx.AsyncClient(base_url=url, timeout=180, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_one(client, key, request_body("chat", _MANTLE_NAME, prompt_text(marker))) for marker in markers), + return_exceptions=True, + ) + return tuple((marker, result) for marker, result in zip(markers, results) if isinstance(result, Outcome)) + + +def _assert_served_and_landed(wire: Wire, served: tuple[tuple[str, Outcome], ...], *, expected_received: int) -> None: + for marker, outcome in served: + assert_answered(outcome, marker) + received: Final = wire.drain() + assert len(received) == expected_received, (len(received), expected_received) + for request in received: + assert breakpoint_count(body_of(request)) == 1, request.body + identities: Final = frozenset(outcome.response_id for _, outcome in served) + assert len(identities) == len(served), identities + rows: Final = spend_rows(_MANTLE_NAME, frozenset(marker for marker, _ in served), expected=len(served), seconds=120) + assert sorted(row_key(str(row["request_id"])) for row in rows) == sorted(identities), rows + + +@pytest.mark.timeout(_CELL_TIMEOUT) +async def test_f2_a_worker_killed_mid_burst_leaves_the_survivor_marking_requests( + gateway: Gateway, tmp_path: Path +) -> None: + release: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + markers: Final = tuple(fresh_marker() for _ in range(_HELD_BURST)) + with wire_server(_peer(release, held)) as wire: + config: Final = _mantle_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + url: Final = str(candidate.client.base_url).rstrip("/") + workers: Final = eventually(lambda: _live_workers(owned.log), lambda pids: len(pids) == 2, seconds=60) + burst: Final = asyncio.create_task(_held_burst(url, candidate.key, markers)) + await asyncio.to_thread(eventually, held.qsize, lambda size: size == _HELD_BURST, 60) + held_by: Final = MappingProxyType({pid: _open_peer_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == _HELD_BURST, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = await burst + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + follow_up: Final = fresh_marker() + (answered,) = await _held_burst(url, candidate.key, (follow_up,)) + _assert_served_and_landed(wire, (*served, answered), expected_received=_HELD_BURST + 1) + eventually(lambda: len(_started_workers(owned.log)), lambda count: count == 3, seconds=90) + + +def _model_id(name: str) -> str: + (row,) = read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_name=%s', (name,)) + return str(row["model_id"]) + + +@pytest.mark.timeout(_RESTART_TIMEOUT) +async def test_f3_a_graceful_restart_mid_burst_drains_the_held_requests_and_keeps_the_stored_explicit_mode( + gateway: Gateway, tmp_path: Path +) -> None: + release: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + markers: Final = tuple(fresh_marker() for _ in range(_HELD_BURST)) + stored: Final = f"mantle-gpt-explicit-stored-{fresh_marker()}" + with wire_server(_peer(release, held)) as wire: + config: Final = _mantle_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + url: Final = str(candidate.client.base_url).rstrip("/") + eventually(lambda: len(_live_workers(owned.log)), lambda count: count == 2, seconds=60) + candidate.post( + "/model/new", + { + "model_name": stored, + "litellm_params": { + "model": GPT, + "api_base": wire.url, + "api_key": TOKEN, + "aws_region_name": "us-east-1", + "cache_control_injection_points": SYSTEM_POINT, + "prompt_cache_options": EXPLICIT, + }, + }, + ) + release.set() + settled(candidate, stored, wire) + before, before_received = observe( + candidate, wire, "responses", request_body("responses", stored, prompt_text(markers[0])) + ) + assert_answered(before, markers[0]) + assert_wire( + before_received, + expected_wire(GPT, prompt_text(markers[0]), endpoint="responses", marked=True, options=EXPLICIT), + streaming=False, + ) + success_row(stored, markers[0]) + _drain(held) + release.clear() + burst: Final = asyncio.create_task(_held_burst(url, candidate.key, markers)) + await asyncio.to_thread(eventually, held.qsize, lambda size: size == _HELD_BURST, 60) + owned.process.terminate() + release.set() + served: Final = await burst + assert len(served) == _HELD_BURST, len(served) + _assert_served_and_landed(wire, served, expected_received=_HELD_BURST) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as restarted: + settled(restarted.gateway, stored, wire) + after, after_received = observe( + restarted.gateway, wire, "responses", request_body("responses", stored, prompt_text(markers[1])) + ) + assert_answered(after, markers[1]) + assert_wire( + after_received, + expected_wire(GPT, prompt_text(markers[1]), endpoint="responses", marked=True, options=EXPLICIT), + streaming=False, + ) + success_row(stored, markers[1]) + restarted.gateway.post("/model/delete", {"id": _model_id(stored)}) diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index 8392cef8a7b..d96f00cc2b1 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -14,6 +14,7 @@ from pydantic import BaseModel, ConfigDict import litellm from litellm.integrations.anthropic_cache_control_hook import ( AnthropicCacheControlHook, + configured_injection_points, supports_openai_prompt_cache_breakpoint, ) from litellm.litellm_core_utils.prompt_templates.factory import ( @@ -3738,6 +3739,14 @@ class TestPromptCacheBreakpointCapability: ) assert supports_openai_prompt_cache_breakpoint("gpt-5.6") is False + @pytest.mark.parametrize("model", ["gpt-4.1", "gpt-5.6"]) + @pytest.mark.parametrize("flag", ["true", 1, "false", 0]) + def test_listed_model_with_an_odd_typed_flag_is_not_eligible(self, monkeypatch, model, flag): + monkeypatch.setitem( + litellm.model_cost, model, {**litellm.model_cost[model], "supports_prompt_cache_breakpoint": flag} + ) + assert supports_openai_prompt_cache_breakpoint(model) is False + def test_published_map_without_the_flag_still_injects_on_gpt_5_6(self, monkeypatch): unflagged = {k: v for k, v in litellm.model_cost["gpt-5.6"].items() if k != "supports_prompt_cache_breakpoint"} @@ -3769,6 +3778,277 @@ class TestPromptCacheBreakpointCapability: assert model not in litellm.model_cost assert supports_openai_prompt_cache_breakpoint(model) is expected + def test_a_null_prompt_cache_options_takes_the_implicit_default_on_both_paths(self): + points = [{"location": "message", "role": "system"}] + + _, _, chat_params = AnthropicCacheControlHook().get_chat_completion_prompt( + model="openai/gpt-5.6", + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + non_default_params={"cache_control_injection_points": copy.deepcopy(points), "prompt_cache_options": None}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert chat_params["prompt_cache_options"] == {"mode": "implicit"} + + kwargs = {"cache_control_injection_points": copy.deepcopy(points), "prompt_cache_options": None} + AnthropicCacheControlHook.maybe_inject_cache_control( + [{"role": "user", "content": "hi"}], "sys", kwargs, model="gpt-5.6", custom_llm_provider="openai" + ) + assert kwargs["prompt_cache_options"] == {"mode": "implicit"} + + +class TestHostedOpenAIDialectFlag: + """#38666: an OpenAI-shaped model served by another provider can opt in through its own + model-map entry, instead of being excluded by the openai-only provider check.""" + + MANTLE_MODEL = "bedrock_mantle/openai.gpt-5.6-sol" + + def _register(self, monkeypatch, key, provider, flag=True, **extra): + entry = {"litellm_provider": provider, "mode": "chat", **extra} + if flag is not None: + entry["supports_prompt_cache_breakpoint"] = flag + monkeypatch.setitem(litellm.model_cost, key, entry) + + def test_flagged_non_openai_deployment_is_eligible(self, monkeypatch): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is True + ) + + def test_bedrock_api_base_does_not_veto_the_explicit_flag(self, monkeypatch): + """The api_base check exists to sniff for api.openai.com, which a Bedrock host never is.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint( + self.MANTLE_MODEL, + "bedrock_mantle", + api_base="https://bedrock-runtime.us-east-1.amazonaws.com", + ) + is True + ) + + def test_flag_set_false_keeps_the_deployment_ineligible(self, monkeypatch): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle", flag=False) + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is False + ) + + @pytest.mark.parametrize("flag", ["true", 1, "false", 0]) + def test_an_odd_typed_flag_keeps_the_deployment_ineligible(self, monkeypatch, flag): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle", flag=flag) + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is False + ) + + def test_unflagged_non_openai_deployment_stays_ineligible(self, monkeypatch): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle", flag=None) + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is False + ) + + def test_entry_provider_must_match_the_request_provider(self, monkeypatch): + """A flagged entry does not license a different provider serving the same model string.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "azure") is False + ) + + def test_openai_entries_still_go_through_the_api_base_check(self, monkeypatch): + """gpt-5.6 is flagged and openai-provided, so it must not bypass the host gate.""" + assert litellm.model_cost["gpt-5.6"]["supports_prompt_cache_breakpoint"] is True + assert litellm.model_cost["gpt-5.6"]["litellm_provider"] == "openai" + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint( + "gpt-5.6", "openai", api_base="https://some-compatible-host.example.com" + ) + is False + ) + + def test_azure_hosted_gpt_5_6_remains_ineligible(self): + """Regression guard: the openai entry's flag must not leak to another provider.""" + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("gpt-5.6", "azure") is False + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("azure/gpt-5.6", None) is False + + REGIONAL_MODEL = "bedrock_mantle/us-east-1/openai.gpt-5.6-sol" + + def test_region_prefixed_deployment_reads_its_region_free_entry(self, monkeypatch): + """``bedrock_mantle//`` is a documented routing form the map keys without the region.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.REGIONAL_MODEL, "bedrock_mantle") + is True + ) + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.REGIONAL_MODEL, None) is True + + def test_region_prefixed_deployment_honors_a_flag_set_false(self, monkeypatch): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle", flag=False) + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.REGIONAL_MODEL, "bedrock_mantle") + is False + ) + + def test_region_prefixed_entry_outranks_the_region_free_one(self, monkeypatch): + """A row keyed with the region states that deployment's own dialect; GovCloud rows carry no flag.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + gov_model = "bedrock_mantle/us-gov-west-1/openai.gpt-5.6-sol" + self._register(monkeypatch, gov_model, "bedrock_mantle", flag=None) + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(gov_model, "bedrock_mantle") is False + + def test_region_free_entry_does_not_license_another_provider(self, monkeypatch): + """The candidate keys are built for the request's provider, so a flagged Mantle row stays Mantle's.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("us-east-1/openai.gpt-5.6-sol", "azure") + is False + ) + + def test_unmapped_deployment_name_stays_ineligible(self): + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("azure/my-deployment", None) is False + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("my-deployment", "azure") is False + + def test_bare_name_colliding_with_an_openai_row_reads_its_own_provider_entry(self, monkeypatch): + """The Responses layer hands the hook a bare deployment name plus its provider. The openai row keyed by + that bare name neither answers for the deployment nor stops the lookup of the provider's own entry.""" + self._register(monkeypatch, "gpt-collide", "openai") + self._register(monkeypatch, "azure_ai/gpt-collide", "azure_ai") + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("gpt-collide", "azure_ai") is True + + def test_bare_name_colliding_with_an_openai_row_stays_ineligible_without_its_own_flag(self, monkeypatch): + self._register(monkeypatch, "gpt-collide", "openai") + self._register(monkeypatch, "azure_ai/gpt-collide", "azure_ai", flag=None) + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("gpt-collide", "azure_ai") is False + + +class TestBedrockMantleGptShipsTheOpenAIDialect: + """The shipped cost map flags Bedrock Mantle's GPT-5.6 and newer OpenAI rows, so a configured injection point + on one of them reaches the wire as prompt_cache_breakpoint instead of an Anthropic cache_control the Mantle + bridge strips (verified live against bedrock-mantle.us-east-1 on 2026-10-07: cache_write_tokens then + cached_tokens on the repeat call).""" + + MANTLE_MODEL = "bedrock_mantle/openai.gpt-5.6-sol" + + @pytest.fixture(autouse=True) + def _bundled_model_map(self, monkeypatch): + bundled = os.path.join(os.path.dirname(litellm.__file__), "model_prices_and_context_window_backup.json") + with open(bundled) as handle: + monkeypatch.setattr(litellm, "model_cost", json.load(handle)) + litellm.utils.cached_get_model_info_helper.cache_clear() + yield + litellm.utils.cached_get_model_info_helper.cache_clear() + + def test_shipped_entry_makes_the_deployment_eligible(self): + assert supports_openai_prompt_cache_breakpoint(self.MANTLE_MODEL) is True + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is True + ) + + def test_every_flagged_mantle_row_is_an_openai_gpt_5_6_or_newer_model(self): + flagged = { + key for key, entry in litellm.model_cost.items() + if key.startswith("bedrock_mantle/") and entry.get("supports_prompt_cache_breakpoint") is True + } + assert self.MANTLE_MODEL in flagged + for key in flagged: + bare = key.rsplit("/", 1)[-1].removeprefix("openai.") + assert supports_openai_prompt_cache_breakpoint(bare) is True, key + + def test_seeding_stamps_the_openai_dialect_for_a_configured_point(self): + non_default_params = {"cache_control_injection_points": [{"location": "message", "role": "system"}]} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=non_default_params, + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + model=self.MANTLE_MODEL, + custom_llm_provider=None, + ) + assert non_default_params["cache_control_injection_points"][0]["_litellm_openai_dialect"] is True + + def test_configured_point_emits_the_openai_marker_and_default_options(self): + _, messages, params = AnthropicCacheControlHook().get_chat_completion_prompt( + model=self.MANTLE_MODEL, + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert messages[0]["content"] == [{"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}}] + assert params["prompt_cache_options"] == {"mode": "implicit"} + assert AnthropicCacheControlHook.count_request_cache_breakpoints(messages) == 1 + + def test_region_prefixed_deployment_emits_the_openai_marker(self): + """The region-prefixed routing form documented for Mantle lands on the same shipped row.""" + _, messages, params = AnthropicCacheControlHook().get_chat_completion_prompt( + model="bedrock_mantle/us-east-1/openai.gpt-5.6-sol", + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert messages[0]["content"] == [{"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}}] + assert params["prompt_cache_options"] == {"mode": "implicit"} + + BARE_MODEL = "openai.gpt-5.6-sol" + POINTS = [{"location": "message", "role": "system"}] + MESSAGES = [{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}] + MARKED_SYSTEM = [{"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}}] + + def test_a_bare_deployment_name_with_its_provider_is_stamped_on_the_openai_dialect(self): + """A deployment written as ``model: openai.gpt-5.6-sol`` plus ``custom_llm_provider: bedrock_mantle`` + has no row of its own and no openai row of the same name, so only the provider-keyed row can + answer; the stamp must read it the way the dialect resolution does.""" + stamped = AnthropicCacheControlHook._stamped_with_dialect( + copy.deepcopy(self.POINTS), self.BARE_MODEL, "bedrock_mantle", None, None + ) + assert stamped[0]["_litellm_openai_dialect"] is True + + def test_a_bare_deployment_name_without_its_provider_keeps_its_points_and_costs_no_lookup(self): + points = copy.deepcopy(self.POINTS) + with patch.object(AnthropicCacheControlHook, "_resolve_provider") as resolve: + assert AnthropicCacheControlHook._stamped_with_dialect(points, self.BARE_MODEL, None, None, None) is points + resolve.assert_not_called() + + def test_the_chat_seed_carries_the_resolved_provider_for_a_bare_deployment_name(self): + params: dict = {"cache_control_injection_points": copy.deepcopy(self.POINTS)} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, + messages=copy.deepcopy(self.MESSAGES), + model=self.BARE_MODEL, + custom_llm_provider="bedrock_mantle", + ) + _, messages, out = AnthropicCacheControlHook().get_chat_completion_prompt( + model=self.BARE_MODEL, + messages=copy.deepcopy(self.MESSAGES), + non_default_params=params, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert messages[0]["content"] == self.MARKED_SYSTEM + assert out["prompt_cache_options"] == {"mode": "implicit"} + + def test_the_responses_stamp_carries_the_resolved_provider_for_a_bare_deployment_name(self): + from litellm.responses.main import _stamp_injection_points_with_dialect + + kwargs: dict = {"cache_control_injection_points": copy.deepcopy(self.POINTS)} + _stamp_injection_points_with_dialect(kwargs, self.BARE_MODEL, "bedrock_mantle") + _, messages, out = AnthropicCacheControlHook().get_chat_completion_prompt( + model=self.BARE_MODEL, + messages=copy.deepcopy(self.MESSAGES), + non_default_params=kwargs, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert messages[0]["content"] == self.MARKED_SYSTEM + assert out["prompt_cache_options"] == {"mode": "implicit"} + class TestRecordGatewayInjection: """The injection marker spend accounting gates prompt-caching savings on.""" @@ -3901,3 +4181,83 @@ class TestRecordGatewayInjection: custom_llm_provider="anthropic", ) assert self.KEY not in kwargs["litellm_metadata"] + + +class TestMalformedInjectionPointsAreIgnored: + """A ``cache_control_injection_points`` value that is not a list of points (a string, an int, a bare + dict, a list of strings) raised inside the hook and turned every request to that deployment into a 500. + Every entry point now reads it as no configured points, the way ``null`` already read.""" + + SHAPES = ("system", 5, {"location": "message", "role": "system"}, ["system"], None) + MIXED = ["system", {"location": "message", "role": "system"}, 3] + MESSAGES = [{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}] + + @pytest.mark.parametrize("value", SHAPES) + def test_reads_as_no_points(self, value): + assert configured_injection_points(value) == () + + def test_keeps_the_point_entries_of_a_mixed_list(self): + assert configured_injection_points(self.MIXED) == ({"location": "message", "role": "system"},) + + def test_the_point_beside_junk_entries_survives_the_chat_seed_on_an_unstamped_deployment(self): + params: dict = {"cache_control_injection_points": copy.deepcopy(self.MIXED)} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, + messages=copy.deepcopy(self.MESSAGES), + model="claude-sonnet-4-5", + custom_llm_provider="anthropic", + ) + _, processed, _ = AnthropicCacheControlHook().get_chat_completion_prompt( + model="claude-sonnet-4-5", + messages=copy.deepcopy(self.MESSAGES), + non_default_params=params, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert processed[0] == {"role": "system", "content": "sys", "cache_control": {"type": "ephemeral"}} + + @pytest.mark.parametrize("value", SHAPES) + def test_chat_prompt_hook_leaves_the_request_untouched(self, value): + _, processed, params = AnthropicCacheControlHook().get_chat_completion_prompt( + model="openai/gpt-5.6", + messages=copy.deepcopy(self.MESSAGES), + non_default_params={"cache_control_injection_points": copy.deepcopy(value)}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert processed == self.MESSAGES + assert "cache_control_injection_points" not in params + assert "prompt_cache_options" not in params + + @pytest.mark.parametrize("value", SHAPES) + def test_chat_seeding_falls_through_to_the_defaults(self, monkeypatch, value): + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + params: dict = {"cache_control_injection_points": copy.deepcopy(value)} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, + messages=copy.deepcopy(self.MESSAGES), + model="claude-sonnet-4-5", + custom_llm_provider="anthropic", + ) + seeded = params["cache_control_injection_points"] + assert seeded and all(isinstance(point, dict) and "location" in point for point in seeded) + + @pytest.mark.parametrize("value", SHAPES) + def test_messages_path_leaves_the_request_untouched(self, value): + kwargs: dict = {"cache_control_injection_points": copy.deepcopy(value)} + messages, system = AnthropicCacheControlHook.maybe_inject_cache_control( + [{"role": "user", "content": "hi"}], "sys", kwargs, model="gpt-5.6", custom_llm_provider="openai" + ) + assert (messages, system) == ([{"role": "user", "content": "hi"}], "sys") + assert "cache_control_injection_points" not in kwargs + assert "prompt_cache_options" not in kwargs + + @pytest.mark.parametrize("value", SHAPES) + def test_responses_dialect_stamp_leaves_the_request_untouched(self, value): + from litellm.responses.main import _stamp_injection_points_with_dialect + + kwargs: dict = {"cache_control_injection_points": copy.deepcopy(value)} + _stamp_injection_points_with_dialect(kwargs, "gpt-5.6", "openai") + assert kwargs == {"cache_control_injection_points": value} diff --git a/tests/unit/responses/test_responses_api_request_body.py b/tests/unit/responses/test_responses_api_request_body.py index 3ef3cd8854a..8c8b347be82 100644 --- a/tests/unit/responses/test_responses_api_request_body.py +++ b/tests/unit/responses/test_responses_api_request_body.py @@ -827,6 +827,62 @@ async def test_injection_points_still_reach_a_native_responses_provider(): assert "cache_control_injection_points" not in body +@pytest.mark.asyncio +async def test_injection_points_reach_a_provider_prefixed_native_responses_model(): + """Bedrock Mantle resolves its Responses config from the price map by bare model name, + so predicting the bridge with the ``bedrock_mantle/``-prefixed name read as bridged and + deferred the system point to a chat-completions pass this native path never runs.""" + injected_client = AsyncHTTPHandler() + mock_post = AsyncMock( + return_value=MockResponse(_minimal_responses_api_payload("resp_mantle", "openai.gpt-5.6-sol"), 200) + ) + injected_client.post = mock_post + + await litellm.aresponses( + model="bedrock_mantle/openai.gpt-5.6-sol", + api_key="fake-bearer-token", + aws_region_name="us-east-1", + input=copy.deepcopy(_INJECTION_POINT_INPUT), + cache_control_injection_points=copy.deepcopy(_SYSTEM_INJECTION_POINT), + client=injected_client, + ) + + assert mock_post.call_args.kwargs["url"].endswith("/openai/v1/responses") + body = _sent_body(mock_post) + assert body["input"][0]["content"][0]["prompt_cache_breakpoint"] == {"mode": "explicit"} + assert body["prompt_cache_options"] == {"mode": "implicit"} + assert "cache_control_injection_points" not in body + + +@pytest.mark.asyncio +async def test_injection_points_reach_a_foundry_deployment_of_an_openai_model(monkeypatch): + """The router hands this layer the provider it resolved with the deployment's api_base, but the + hook reads the provider from the request kwargs, which never carry it, and resolving + ``azure_ai/gpt-6-astra`` by name alone reads the ambient ``AZURE_AI_API_BASE``, so next to an + Azure OpenAI one the Foundry deployment got Anthropic marks the Responses transform then stripped.""" + monkeypatch.setenv("AZURE_AI_API_BASE", "https://other-deployment.openai.azure.com") + injected_client = AsyncHTTPHandler() + mock_post = AsyncMock( + return_value=MockResponse(_minimal_responses_api_payload("resp_foundry", "gpt-6-astra"), 200) + ) + injected_client.post = mock_post + + await litellm.aresponses( + model="azure_ai/gpt-6-astra", + custom_llm_provider="azure_ai", + api_key="fake-api-key", + api_base="https://foundry.services.ai.azure.com", + input=copy.deepcopy(_INJECTION_POINT_INPUT), + cache_control_injection_points=copy.deepcopy(_SYSTEM_INJECTION_POINT), + client=injected_client, + ) + + body = _sent_body(mock_post) + assert body["input"][0]["content"][0]["prompt_cache_breakpoint"] == {"mode": "explicit"} + assert body["prompt_cache_options"] == {"mode": "implicit"} + assert "cache_control_injection_points" not in body + + async def _bridged_body(mock_post, *, points, input, instructions="You are a documentation assistant."): injected_client = AsyncHTTPHandler() injected_client.post = mock_post