mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(bedrock_mantle): send OpenAI explicit prompt cache breakpoints for GPT-5.6 and newer (#38729)
* fix(cache_control): let a hosted deployment opt into the OpenAI cache dialect _targets_openai_prompt_cache_breakpoint gated on custom_llm_provider == "openai" unconditionally, so an OpenAI-shaped model served by another provider could never qualify, even with supports_prompt_cache_breakpoint set explicitly on its own cost-map entry. bedrock_mantle honours prompt_cache_breakpoint end to end, and had no way to say so. model_cost is keyed per exact deployment string, so a flag on the deployment's own entry states the dialect more precisely than a provider name can. Consult it before the provider check, and require the entry's litellm_provider to match the serving provider so an openai entry cannot license another provider. Entries for the openai provider keep their api_base check, so an OpenAI-compatible third-party host is still not assumed to speak the dialect. Fixes #38666 * test: drop a laziness assertion the existing suite already makes test_provider_lookup_skipped_for_models_below_gpt_5_6 already pins that _resolve_provider is not called for an unflagged model, and it covers the new flag lookup unchanged. The duplicate patched a litellm internal for no added coverage, which the test-quality gate counts against TQ008. * fix(bedrock_mantle): flag GPT-5.6 and newer rows for OpenAI explicit prompt cache breakpoints * fix(responses): predict the chat-completions bridge with the model name the dispatch resolves * fix(cache_control): read a region-prefixed Mantle GPT deployment through its region-free price-map row * fix(cache_control): let a bare hosted deployment name read its own provider's price-map row * fix(responses): hand the hook the provider the router resolved for a Foundry GPT deployment * fix(cache_control): read the deployment breakpoint flag strictly and default a null prompt_cache_options A cost-map or deployment `supports_prompt_cache_breakpoint` that is not the boolean `true` (the string "true", the integer 1, "false", 0, a long string) no longer opts a deployment into the OpenAI prompt cache dialect; only `true` does, the same reading Bedrock Converse applies to its own flag. A client that sends `"prompt_cache_options": null` on `/v1/chat/completions` or `/v1/messages` now gets the implicit default the hook already applied on `/v1/responses` when it placed a breakpoint; before, the null suppressed the default and the request left with a breakpoint and no options. * test(integration): audit cells for the Mantle GPT prompt cache breakpoint dialect Deterministic cells for the configured-breakpoint path on Bedrock Mantle GPT and Azure AI Foundry GPT-6 deployments: every endpoint, streaming and not, sync and async SDKs and raw httpx, the hostile option shapes, flag precedence, the response cache twin, a two-worker burst, a worker kill and a graceful restart. * fix(cache_control): ignore a malformed cache_control_injection_points value instead of failing the request A deployment or client `cache_control_injection_points` that is not a list of points (a string, an integer, a bare point dict, a list of strings) raised inside the prompt hook (`'str' object has no attribute 'get'`, `'int' object is not iterable`) and turned every request to that deployment into a 500 on `/v1/chat/completions`, `/v1/messages` and `/v1/responses`. Every entry point now reads such a value as no configured points, the way a `null` already read, and the request leaves without a breakpoint; entries of a list that are not points are dropped and the point entries kept. * test(integration): cover int, dict and string-list injection point shapes in the Mantle audit cells The D9 cell now also sends an integer, a bare point dict and a list of strings as the deployment's `cache_control_injection_points`, which the merge base answered with a 500 on every request. * fix(cache_control): stamp the OpenAI dialect for a bare deployment name served by a flagged provider A deployment written as `model: openai.gpt-5.6-sol` with `custom_llm_provider: bedrock_mantle` has no cost-map row of its own and no openai row of the same name, so the dialect stamp's cheap gate (`supports_openai_prompt_cache_breakpoint`) returned early and the configured point was never stamped. On `/v1/responses` and on the chat seeding path the hook sees no provider, so it fell back to the Anthropic `cache_control` marker, which Bedrock Mantle strips. The gate now also passes a model whose serving provider is already known and whose provider-keyed row carries the flag, which is the same row the dialect resolution reads and costs no provider lookup. A bare name without a provider is still left alone, so models below GPT-5.6 keep their points untouched and resolve nothing. * test(integration): cover a bare Mantle deployment name with its provider in the audit cells A deployment configured as `model: openai.gpt-5.6-sol` plus `custom_llm_provider: bedrock_mantle` sends the breakpoint and the implicit options on all three endpoints; the merge base leaves it on the Anthropic dialect. * fix(cache_control): keep a configured point that sits beside junk entries on every chat seed path `configured_injection_points` kept the point dicts of a mixed `cache_control_injection_points` list as a tuple, but only read a `list` back. The chat seed and the `/v1/responses` dialect stamp write the normalized value back onto the request, and on a deployment the OpenAI dialect does not stamp (an Anthropic model, Claude on Bedrock Mantle) that value is the tuple itself, so the hook then read it as no configured points and the valid point was silently dropped. The stamped OpenAI dialect and the `/v1/messages` path kept it only because they build a new list. The normalizer now reads back the tuple it wrote, so the point reaches the wire on every path; an all-dict list still passes through as the same object. * test(integration): cover a configured point beside junk entries on a Claude Mantle deployment A deployment configured with `cache_control_injection_points: ["system", {system point}, 3]` marks the system block on the Anthropic dialect on all three endpoints; the merge base answers 500 on the junk entry. --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
7a659973e3
commit
7c7b0ea85b
9 changed files with 2450 additions and 19 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
959
tests/integration/providers/_mantle_gpt_prompt_cache_support.py
Normal file
959
tests/integration/providers/_mantle_gpt_prompt_cache_support.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)})
|
||||
|
|
@ -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/<region>/<model>`` 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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue