mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(responses): keep a flagged hosted deployment's prompt cache breakpoint on the chat bridge (#45500)
* fix(responses): keep a flagged hosted deployment's prompt cache breakpoint on the chat bridge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(responses): read the serving provider's row for prompt cache breakpoints on the chat bridge and base_model * fix(anthropic_cache_control_hook): rename the module-level provider resolver so the recursion check passes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(ui): follow the credential dialog's new name field and auth method id in the federation e2e spec Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: mateo <mateo@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
dd31339e05
commit
bbe5595a99
5 changed files with 187 additions and 22 deletions
|
|
@ -175,9 +175,13 @@ def _without_audio_input_parts(input_items: list[object]) -> list[object]:
|
|||
return [_without_audio_input_parts_in_item(item) for item in input_items]
|
||||
|
||||
|
||||
def _supports_audio_input(model: str, litellm_params: Mapping[str, object]) -> bool:
|
||||
def _served_provider(litellm_params: Mapping[str, object]) -> str | None:
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
provider: Final = custom_llm_provider if isinstance(custom_llm_provider, str) else None
|
||||
return custom_llm_provider if isinstance(custom_llm_provider, str) else None
|
||||
|
||||
|
||||
def _supports_audio_input(model: str, litellm_params: Mapping[str, object]) -> bool:
|
||||
provider: Final = _served_provider(litellm_params)
|
||||
base_model: Final = litellm_params.get("base_model")
|
||||
return litellm.supports_audio_input(model=model, custom_llm_provider=provider) or (
|
||||
isinstance(base_model, str)
|
||||
|
|
@ -737,14 +741,17 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
model: str,
|
||||
messages: list["AllMessageValues"],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
litellm_params: dict[str, object],
|
||||
headers: dict,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
client: object | None = None,
|
||||
) -> dict:
|
||||
provider: Final = _served_provider(litellm_params)
|
||||
base_model: Final = litellm_params.get("base_model")
|
||||
supports_prompt_cache_breakpoint: Final = supports_openai_prompt_cache_breakpoint(model) or (
|
||||
isinstance(base_model, str) and bool(base_model) and supports_openai_prompt_cache_breakpoint(base_model)
|
||||
supports_prompt_cache_breakpoint: Final = supports_openai_prompt_cache_breakpoint(model, provider) or (
|
||||
isinstance(base_model, str)
|
||||
and bool(base_model)
|
||||
and supports_openai_prompt_cache_breakpoint(base_model, provider)
|
||||
)
|
||||
converted_input_items, converted_instructions = self.convert_chat_completion_messages_to_responses_api(
|
||||
messages,
|
||||
|
|
|
|||
|
|
@ -98,7 +98,25 @@ def configured_injection_points(value: object) -> Sequence[CacheControlInjection
|
|||
return tuple(cast(CacheControlInjectionPoint, entry) for entry in value if isinstance(entry, dict))
|
||||
|
||||
|
||||
def supports_openai_prompt_cache_breakpoint(model: str) -> bool:
|
||||
def _served_provider_for_model(model: str) -> str | None:
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
try:
|
||||
_, provider, _, _ = get_llm_provider(model=model)
|
||||
except BadRequestError:
|
||||
return None
|
||||
return provider
|
||||
|
||||
|
||||
def supports_openai_prompt_cache_breakpoint(model: str, custom_llm_provider: str | None = None) -> bool:
|
||||
hosted_flag: Final = (
|
||||
None
|
||||
if custom_llm_provider is None
|
||||
else _hosted_openai_dialect_flag(model, custom_llm_provider, _served_provider_for_model)
|
||||
)
|
||||
if hosted_flag is not None:
|
||||
return hosted_flag
|
||||
model_map_flag: Final = _model_map_prompt_cache_breakpoint_flag(model)
|
||||
if model_map_flag is not None:
|
||||
return model_map_flag
|
||||
|
|
@ -387,14 +405,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
@staticmethod
|
||||
def _resolve_provider(model: str) -> str | None:
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
try:
|
||||
_, provider, _, _ = get_llm_provider(model=model)
|
||||
except BadRequestError:
|
||||
return None
|
||||
return provider
|
||||
return _served_provider_for_model(model)
|
||||
|
||||
@staticmethod
|
||||
def count_request_cache_breakpoints(messages: Iterable[object], system: object = None) -> int:
|
||||
|
|
|
|||
|
|
@ -109,7 +109,11 @@ async function pickAuthMethod(
|
|||
dialog: Locator,
|
||||
label: string,
|
||||
): Promise<void> {
|
||||
await pickOption(page, dialog.locator("#anthropic_auth_method"), label);
|
||||
await pickOption(
|
||||
page,
|
||||
dialog.getByRole("combobox", { name: "Authentication:", exact: true }),
|
||||
label,
|
||||
);
|
||||
}
|
||||
|
||||
async function pickIdentitySource(
|
||||
|
|
@ -142,9 +146,7 @@ async function openAddCredentialDialog(
|
|||
.click();
|
||||
const dialog = page.getByRole("dialog", { name: "Add New Credential" });
|
||||
await expect(dialog).toBeVisible();
|
||||
await dialog
|
||||
.getByPlaceholder("Enter a friendly name for these credentials")
|
||||
.fill(name);
|
||||
await dialog.getByLabel("Credential Name:", { exact: true }).fill(name);
|
||||
await pickProvider(page, dialog, ANTHROPIC_LABEL);
|
||||
await pickAuthMethod(page, dialog, FEDERATION_BADGE);
|
||||
return dialog;
|
||||
|
|
@ -161,6 +163,11 @@ async function openEditDialog(
|
|||
await page.getByTestId("credential-action-edit").click();
|
||||
const dialog = page.getByRole("dialog", { name: "Edit Credential" });
|
||||
await expect(dialog).toBeVisible();
|
||||
const credentialNameField = dialog.getByLabel("Credential Name:", {
|
||||
exact: true,
|
||||
});
|
||||
await expect(credentialNameField).toHaveValue(name);
|
||||
await expect(credentialNameField).toBeDisabled();
|
||||
return dialog;
|
||||
}
|
||||
|
||||
|
|
@ -637,11 +644,14 @@ test("a proxy admin saves a federation credential from the Add Model tab and the
|
|||
const provider = dialog.getByPlaceholder("Select a provider");
|
||||
await expect(provider).toHaveValue(ANTHROPIC_LABEL);
|
||||
await expect(provider).toBeDisabled();
|
||||
await expect(dialog.locator("#anthropic_auth_method")).toContainText(
|
||||
FEDERATION_BADGE,
|
||||
);
|
||||
await expect(
|
||||
dialog.getByRole("combobox", {
|
||||
name: "Authentication:",
|
||||
exact: true,
|
||||
}),
|
||||
).toContainText(FEDERATION_BADGE);
|
||||
await dialog
|
||||
.getByPlaceholder("Enter a friendly name for these credentials")
|
||||
.getByLabel("Credential Name:", { exact: true })
|
||||
.fill(credentialName);
|
||||
await fillFields(dialog, {
|
||||
...ids,
|
||||
|
|
|
|||
|
|
@ -4247,6 +4247,50 @@ def test_prompt_cache_breakpoint_survives_chat_to_responses_conversion(
|
|||
assert request["prompt_cache_options"] == cache_breakpoint
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "custom_llm_provider", "keep_marker"),
|
||||
[
|
||||
("openai.gpt-5.6-sol", "bedrock_mantle", True),
|
||||
("openai.gpt-5.4", "bedrock_mantle", False),
|
||||
("openai.gpt-5.6-sol", "azure_ai", False),
|
||||
],
|
||||
ids=("flagged-hosted-deployment", "unflagged-hosted-deployment", "provider-without-flag"),
|
||||
)
|
||||
def test_transform_request_uses_provider_keyed_prompt_cache_breakpoint_flag(
|
||||
model: str, custom_llm_provider: str, keep_marker: bool
|
||||
) -> None:
|
||||
handler: Final = LiteLLMResponsesTransformationHandler()
|
||||
cache_breakpoint: Final = {"mode": "explicit"}
|
||||
messages: Final = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [{"type": "text", "text": "Stable prefix", "prompt_cache_breakpoint": cache_breakpoint}],
|
||||
},
|
||||
{"role": "user", "content": "Hi"},
|
||||
]
|
||||
|
||||
request: Final = handler.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={"custom_llm_provider": custom_llm_provider},
|
||||
headers={},
|
||||
litellm_logging_obj=Mock(),
|
||||
)
|
||||
|
||||
assert request["input"][0] == {
|
||||
"type": "message",
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "input_text",
|
||||
"text": "Stable prefix",
|
||||
**({"prompt_cache_breakpoint": cache_breakpoint} if keep_marker else {}),
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def test_prompt_cache_breakpoint_read_tolerates_non_string_content_block_keys() -> None:
|
||||
handler: Final = LiteLLMResponsesTransformationHandler()
|
||||
# Non-string keys are not JSON-representable but are accepted by chat completion
|
||||
|
|
@ -4534,6 +4578,81 @@ def test_prompt_cache_breakpoint_supports_model_alias_with_base_model(
|
|||
]
|
||||
|
||||
|
||||
_MANTLE_GPT_ROW: Final = {
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"mode": "responses",
|
||||
"supports_prompt_cache_breakpoint": True,
|
||||
}
|
||||
_MARKED_SYSTEM_PART: Final = {"type": "text", "text": "Stable prefix", "prompt_cache_breakpoint": {"mode": "explicit"}}
|
||||
|
||||
|
||||
def _bridged_input_for_served_provider(
|
||||
model: str, custom_llm_provider: str, **litellm_params: object
|
||||
) -> list[dict[str, object]]:
|
||||
request: Final = LiteLLMResponsesTransformationHandler().transform_request(
|
||||
model=model,
|
||||
messages=cast(
|
||||
List[AllMessageValues],
|
||||
[{"role": "system", "content": [_MARKED_SYSTEM_PART]}, {"role": "user", "content": "hi"}],
|
||||
), # cast-ok: the test builds chat messages as plain mappings
|
||||
optional_params={},
|
||||
litellm_params={"custom_llm_provider": custom_llm_provider, **litellm_params},
|
||||
headers={},
|
||||
litellm_logging_obj=Mock(),
|
||||
)
|
||||
return cast(list[dict[str, object]], request["input"]) # cast-ok: the bridge emits message item mappings
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["openai.gpt-5.6-sol", "us-east-1/openai.gpt-5.6-sol"])
|
||||
def test_transform_request_keeps_the_breakpoint_for_a_hosted_openai_model_by_its_served_provider(
|
||||
monkeypatch: pytest.MonkeyPatch, model: str
|
||||
) -> None:
|
||||
"""The bridge is handed the provider-stripped deployment name, which has no cost-map row and no
|
||||
``gpt-`` version of its own: the row keyed for the serving provider states the dialect."""
|
||||
monkeypatch.setitem(litellm.model_cost, "bedrock_mantle/openai.gpt-5.6-sol", _MANTLE_GPT_ROW)
|
||||
|
||||
assert _bridged_input_for_served_provider(model, "bedrock_mantle") == [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "system",
|
||||
"content": [{**_MARKED_SYSTEM_PART, "type": "input_text"}],
|
||||
},
|
||||
{"type": "message", "role": "user", "content": [{"type": "input_text", "text": "hi"}]},
|
||||
]
|
||||
|
||||
|
||||
def test_transform_request_keeps_the_breakpoint_when_the_served_providers_row_is_the_base_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setitem(litellm.model_cost, "bedrock_mantle/openai.gpt-5.6-sol", _MANTLE_GPT_ROW)
|
||||
|
||||
assert _bridged_input_for_served_provider("my-alias", "bedrock_mantle", base_model="openai.gpt-5.6-sol")[0][
|
||||
"content"
|
||||
] == [{**_MARKED_SYSTEM_PART, "type": "input_text"}]
|
||||
|
||||
|
||||
def test_transform_request_strips_the_breakpoint_when_the_served_providers_row_opts_out(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost, "bedrock_mantle/gpt-6-luna", {**_MANTLE_GPT_ROW, "supports_prompt_cache_breakpoint": False}
|
||||
)
|
||||
|
||||
assert _bridged_input_for_served_provider("gpt-6-luna", "bedrock_mantle")[0]["content"] == [
|
||||
{"type": "input_text", "text": "Stable prefix"}
|
||||
]
|
||||
|
||||
|
||||
def test_transform_request_strips_the_breakpoint_for_a_served_provider_without_a_row_of_its_own(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setitem(litellm.model_cost, "bedrock_mantle/openai.gpt-5.6-sol", _MANTLE_GPT_ROW)
|
||||
|
||||
assert _bridged_input_for_served_provider("openai.gpt-5.6-sol", "azure")[0]["content"] == [
|
||||
{"type": "input_text", "text": "Stable prefix"}
|
||||
]
|
||||
|
||||
|
||||
def test_mid_conversation_system_string_stays_in_input_after_a_user_turn():
|
||||
handler: Final = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
|
|
|
|||
|
|
@ -3923,6 +3923,24 @@ class TestHostedOpenAIDialectFlag:
|
|||
self._register(monkeypatch, "azure_ai/gpt-collide", "azure_ai", flag=None)
|
||||
assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("gpt-collide", "azure_ai") is False
|
||||
|
||||
def test_the_support_check_reads_the_served_providers_row_for_a_provider_stripped_name(self, monkeypatch):
|
||||
"""The chat-to-Responses bridge holds the provider-stripped deployment name and the provider that
|
||||
serves it, which has no row and no ``gpt-`` version of its own."""
|
||||
self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle")
|
||||
assert supports_openai_prompt_cache_breakpoint("openai.gpt-5.6-sol") is False
|
||||
assert supports_openai_prompt_cache_breakpoint("openai.gpt-5.6-sol", "bedrock_mantle") is True
|
||||
assert supports_openai_prompt_cache_breakpoint("us-east-1/openai.gpt-5.6-sol", "bedrock_mantle") is True
|
||||
|
||||
def test_the_support_check_honors_the_served_providers_opt_out_over_the_version_rule(self, monkeypatch):
|
||||
self._register(monkeypatch, "bedrock_mantle/gpt-6-luna", "bedrock_mantle", flag=False)
|
||||
assert supports_openai_prompt_cache_breakpoint("gpt-6-luna") is True
|
||||
assert supports_openai_prompt_cache_breakpoint("gpt-6-luna", "bedrock_mantle") is False
|
||||
|
||||
def test_the_support_check_ignores_a_served_provider_without_a_row_of_its_own(self, monkeypatch):
|
||||
self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle")
|
||||
assert supports_openai_prompt_cache_breakpoint("openai.gpt-5.6-sol", "azure") is False
|
||||
assert supports_openai_prompt_cache_breakpoint("gpt-5.6", "azure") is True
|
||||
|
||||
|
||||
class TestBedrockMantleGptShipsTheOpenAIDialect:
|
||||
"""The shipped cost map flags Bedrock Mantle's GPT-5.6 and newer OpenAI rows, so a configured injection point
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue