This commit is contained in:
Priyansh Nandwana 2026-09-23 14:52:00 +00:00 • committed by GitHub
commit 2a6c5e92d0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 99 additions and 2 deletions

View file

@ -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
@ -108,6 +108,27 @@ def _model_map_prompt_cache_breakpoint_flag(model: str) -> bool | None:
return next((bool(flag) for flag in flags if flag is not None), None)
def _hosted_openai_dialect_flag(model: str, resolve_provider: Callable[[], str | None]) -> bool | None:
"""
Explicit opt-in for an OpenAI-shaped deployment served by another provider.
``model_cost`` is keyed per exact deployment string, so a flag on the deployment's
own entry states the dialect directly, which a provider name cannot express. 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.
"""
import litellm
entry: Final = litellm.model_cost.get(model)
if not isinstance(entry, dict):
return 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 bool(flag) if entry_provider == resolve_provider() else None
def targets_openai_api(api_base: object) -> bool:
import litellm
@ -286,7 +307,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, lambda: custom_llm_provider or AnthropicCacheControlHook._resolve_provider(model)
)
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

View file

@ -3475,6 +3475,75 @@ class TestPromptCacheBreakpointCapability:
assert supports_openai_prompt_cache_breakpoint(model) is expected
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
)
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
class TestRecordGatewayInjection:
"""The injection marker spend accounting gates prompt-caching savings on."""