fix(prompts): apply prompt templates before routing on /v1/responses and honor ignore_prompt_manager_model

On /v1/responses the prompt template ran inside litellm.aresponses, after the
router had already resolved a deployment and injected its api_key/api_base, so a
prompt whose metadata.model pointed at another provider sent the old
deployment's credentials cross-provider (401). The proxy now runs the prompt
template for aresponses in the pre-call hook, before routing, so the router
picks the deployment that matches the swapped model. As a backstop, the SDK
refuses a cross-provider swap when explicit credentials are already present
instead of forwarding them.

ignore_prompt_manager_model and ignore_prompt_manager_optional_params saved on
a prompt were only read by the generic manager, so dotprompt prompts ignored
them on every endpoint. PromptManagementBase now merges the prompt spec's flags
with the per-request ones for every manager, and the generic manager no longer
drops caller flags when no spec is present.
This commit is contained in:
mateo-berri 2026-08-26 14:12:28 -07:00
parent f6571a653f
commit dbc819dc77
11 changed files with 357 additions and 40 deletions

View file

@ -209,6 +209,8 @@ class DotpromptManager(CustomPromptManagement):
prompt_spec=prompt_spec,
prompt_label=prompt_label,
prompt_version=prompt_version,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
async def async_get_chat_completion_prompt(

View file

@ -416,17 +416,8 @@ class GenericPromptManager(CustomPromptManagement):
tools=tools,
prompt_label=prompt_label,
prompt_version=prompt_version,
ignore_prompt_manager_model=(
ignore_prompt_manager_model or prompt_spec.litellm_params.ignore_prompt_manager_model
if prompt_spec
else False
),
ignore_prompt_manager_optional_params=(
ignore_prompt_manager_optional_params
or prompt_spec.litellm_params.ignore_prompt_manager_optional_params
if prompt_spec
else False
),
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
def get_chat_completion_prompt(
@ -457,17 +448,8 @@ class GenericPromptManager(CustomPromptManagement):
prompt_spec=prompt_spec,
prompt_label=prompt_label,
prompt_version=prompt_version,
ignore_prompt_manager_model=(
ignore_prompt_manager_model or prompt_spec.litellm_params.ignore_prompt_manager_model
if prompt_spec
else False
),
ignore_prompt_manager_optional_params=(
ignore_prompt_manager_optional_params
or prompt_spec.litellm_params.ignore_prompt_manager_optional_params
if prompt_spec
else False
),
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
def clear_cache(self) -> None:

View file

@ -19,6 +19,19 @@ class PromptManagementClient(TypedDict):
completed_messages: list[AllMessageValues] | None
def resolve_prompt_manager_ignore_flags(
prompt_spec: PromptSpec | None,
ignore_prompt_manager_model: bool | None,
ignore_prompt_manager_optional_params: bool | None,
) -> tuple[bool, bool]:
spec_params: Final = prompt_spec.litellm_params if prompt_spec is not None else None
return (
bool(ignore_prompt_manager_model) or bool(spec_params is not None and spec_params.ignore_prompt_manager_model),
bool(ignore_prompt_manager_optional_params)
or bool(spec_params is not None and spec_params.ignore_prompt_manager_optional_params),
)
class PromptManagementBase(ABC):
@property
@abstractmethod
@ -182,13 +195,18 @@ class PromptManagementBase(ABC):
prompt_version=prompt_version,
)
resolved_ignore_model, resolved_ignore_optional_params = resolve_prompt_manager_ignore_flags(
prompt_spec=prompt_spec,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
return self.post_compile_prompt_processing(
prompt_template=prompt_template,
messages=messages,
non_default_params=non_default_params,
model=model,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
ignore_prompt_manager_model=resolved_ignore_model,
ignore_prompt_manager_optional_params=resolved_ignore_optional_params,
)
async def async_get_chat_completion_prompt(
@ -224,11 +242,16 @@ class PromptManagementBase(ABC):
prompt_version=prompt_version,
)
resolved_ignore_model, resolved_ignore_optional_params = resolve_prompt_manager_ignore_flags(
prompt_spec=prompt_spec,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
return self.post_compile_prompt_processing(
prompt_template=prompt_template,
messages=messages,
non_default_params=non_default_params,
model=model,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
ignore_prompt_manager_model=resolved_ignore_model,
ignore_prompt_manager_optional_params=resolved_ignore_optional_params,
)

View file

@ -1402,6 +1402,7 @@ class ProxyLogging:
get_latest_version_prompt_id,
)
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.utils import get_non_default_completion_params
if prompt_version is None:
@ -1420,13 +1421,20 @@ class ProxyLogging:
data.pop("prompt_id", None)
if custom_logger and prompt_spec is not None:
is_responses_call: Final = call_type == "aresponses"
original_responses_input: Final = data.get("input", "") if is_responses_call else None
client_messages: Final = (
ResponsesAPIRequestUtils.responses_input_to_chat_messages(original_responses_input)
if is_responses_call
else data.get("messages", [])
)
(
model,
messages,
optional_params,
) = await litellm_logging_obj.async_get_chat_completion_prompt(
model=data.get("model", ""),
messages=data.get("messages", []),
messages=client_messages,
non_default_params=get_non_default_completion_params(kwargs=data) or {},
prompt_id=litellm_prompt_id,
prompt_spec=prompt_spec,
@ -1438,7 +1446,14 @@ class ProxyLogging:
data.update(optional_params)
data["model"] = model
data["messages"] = messages
if is_responses_call:
data["input"] = ResponsesAPIRequestUtils.merge_prompt_management_input(
original_input=original_responses_input,
client_input=client_messages,
merged_input=messages,
)
else:
data["messages"] = messages
# prevent re-processing the prompt template
data.pop("prompt_id", None)
data.pop("prompt_variables", None)
@ -1653,7 +1668,7 @@ class ProxyLogging:
not guardrails_only
and litellm_logging_obj is not None
and prompt_id is not None
and (call_type == "completion" or call_type == "acompletion")
and (call_type == "completion" or call_type == "acompletion" or call_type == "aresponses")
):
await self._process_prompt_template(
data=data,

View file

@ -26,7 +26,6 @@ from litellm.responses.litellm_completion_transformation.handler import (
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
AllMessageValues,
PromptObject,
Reasoning,
ResponseIncludable,
@ -463,10 +462,7 @@ async def aresponses(
if isinstance(
litellm_logging_obj, LiteLLMLoggingObj
) and litellm_logging_obj.should_run_prompt_management_hooks(prompt_id=prompt_id, non_default_params=kwargs):
if isinstance(input, str):
client_input: list[AllMessageValues] = [{"role": "user", "content": input}]
else:
client_input = [item for item in input if isinstance(item, dict) and "role" in item]
client_input: Final = ResponsesAPIRequestUtils.responses_input_to_chat_messages(input)
(
model,
merged_input,
@ -489,7 +485,13 @@ async def aresponses(
),
)
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
custom_llm_provider = _resolve_prompt_swapped_provider(
original_model=original_model,
swapped_model=model,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
prompt_id=prompt_id,
)
kwargs.pop("prompt_id", None)
kwargs["_async_prompt_merged_params"] = merged_optional_params
@ -559,6 +561,35 @@ async def aresponses(
)
def _resolve_prompt_swapped_provider(
original_model: str,
swapped_model: str,
custom_llm_provider: str | None,
kwargs: Mapping[str, object],
prompt_id: str | None,
) -> str:
swapped_provider: Final = litellm.get_llm_provider(model=swapped_model)[1]
if kwargs.get("api_key") is None and kwargs.get("api_base") is None:
return swapped_provider
try:
original_provider: Final = custom_llm_provider or litellm.get_llm_provider(model=original_model)[1]
except litellm.BadRequestError:
return swapped_provider
if swapped_provider == original_provider:
return swapped_provider
raise litellm.BadRequestError(
message=(
f"prompt_id '{prompt_id}' swaps model '{original_model}' -> '{swapped_model}', which changes the "
f"provider from '{original_provider}' to '{swapped_provider}' after credentials for "
f"'{original_provider}' were already resolved. Refusing to send them to '{swapped_provider}'. "
"Point the request at a model whose provider matches the prompt's metadata.model, or set "
"ignore_prompt_manager_model on the prompt to keep the requested model."
),
model=swapped_model,
llm_provider=swapped_provider,
)
def _apply_prompt_management_to_responses_call(
input: str | ResponseInputParam,
model: str,
@ -577,10 +608,7 @@ def _apply_prompt_management_to_responses_call(
prompt_variables: Final = cast(dict | None, kwargs.get("prompt_variables", None))
original_model: Final = model
if isinstance(input, str):
client_input: list[AllMessageValues] = [{"role": "user", "content": input}]
else:
client_input = [item for item in input if isinstance(item, dict) and "role" in item]
client_input: Final = ResponsesAPIRequestUtils.responses_input_to_chat_messages(input)
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=kwargs
@ -609,7 +637,13 @@ def _apply_prompt_management_to_responses_call(
local_vars["input"] = input
local_vars["model"] = model
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
custom_llm_provider = _resolve_prompt_swapped_provider(
original_model=original_model,
swapped_model=model,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
prompt_id=prompt_id,
)
local_vars["custom_llm_provider"] = custom_llm_provider
for key, value in merged_optional_params.items():
local_vars[key] = value

View file

@ -72,6 +72,16 @@ class ResponsesAPIRequestUtils:
shaped_content: Final = [_as_input_text_part(part) for part in content] # mutable-ok: Responses-shaped copy
return {**message, "content": shaped_content} # mutable-ok: copy, the hook's message stays untouched
@staticmethod
def responses_input_to_chat_messages(
input: str | ResponseInputParam | None,
) -> list[AllMessageValues]:
if input is None:
return []
if isinstance(input, str):
return [{"role": "user", "content": input}]
return [item for item in input if isinstance(item, dict) and "role" in item]
@staticmethod
def merge_prompt_management_input(
original_input: str | ResponseInputParam,

View file

@ -577,3 +577,93 @@ async def test_dotprompt_with_prompt_version():
)
assert "Version 2:" in v2_rendered
assert "Test v2" in v2_rendered
def _swap_prompt_manager_and_spec(ignore_prompt_manager_model: bool):
from litellm.integrations.dotprompt.dotprompt_manager import DotpromptManager
from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec
manager = DotpromptManager(
prompt_data={"content": "You are a pirate assistant.", "metadata": {"model": "gpt-4o-mini"}},
prompt_id="swap-prompt",
)
spec = PromptSpec(
prompt_id="swap-prompt",
litellm_params=PromptLiteLLMParams(
prompt_id="swap-prompt",
prompt_integration="dotprompt",
ignore_prompt_manager_model=ignore_prompt_manager_model,
),
)
return manager, spec
@pytest.mark.asyncio
async def test_async_prompt_spec_ignore_prompt_manager_model_keeps_requested_model():
from litellm.types.utils import StandardCallbackDynamicParams
manager, spec = _swap_prompt_manager_and_spec(ignore_prompt_manager_model=True)
model, messages, _ = await manager.async_get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
non_default_params={},
prompt_id="swap-prompt",
prompt_variables=None,
dynamic_callback_params=StandardCallbackDynamicParams(),
litellm_logging_obj=MagicMock(),
prompt_spec=spec,
)
assert model == "anthropic/claude-haiku-4-5"
assert len(messages) == 2
assert "pirate" in str(messages[0]["content"])
@pytest.mark.asyncio
async def test_async_prompt_spec_without_ignore_flag_swaps_model():
from litellm.types.utils import StandardCallbackDynamicParams
manager, spec = _swap_prompt_manager_and_spec(ignore_prompt_manager_model=False)
model, _, _ = await manager.async_get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
non_default_params={},
prompt_id="swap-prompt",
prompt_variables=None,
dynamic_callback_params=StandardCallbackDynamicParams(),
litellm_logging_obj=MagicMock(),
prompt_spec=spec,
)
assert model == "gpt-4o-mini"
def test_sync_prompt_spec_ignore_prompt_manager_model_keeps_requested_model():
from litellm.types.utils import StandardCallbackDynamicParams
manager, spec = _swap_prompt_manager_and_spec(ignore_prompt_manager_model=True)
model, _, _ = manager.get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
non_default_params={},
prompt_id="swap-prompt",
prompt_variables=None,
dynamic_callback_params=StandardCallbackDynamicParams(),
prompt_spec=spec,
)
assert model == "anthropic/claude-haiku-4-5"
def test_sync_caller_ignore_flag_survives_missing_prompt_spec():
from litellm.types.utils import StandardCallbackDynamicParams
manager, _ = _swap_prompt_manager_and_spec(ignore_prompt_manager_model=False)
model, _, _ = manager.get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
non_default_params={},
prompt_id="swap-prompt",
prompt_variables=None,
dynamic_callback_params=StandardCallbackDynamicParams(),
prompt_spec=None,
ignore_prompt_manager_model=True,
)
assert model == "anthropic/claude-haiku-4-5"

View file

@ -818,3 +818,50 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
prompt_version=None,
call_type="completion",
)
@pytest.mark.asyncio
async def test_process_prompt_template_aresponses_swaps_model_and_merges_input(proxy_logging, monkeypatch):
from litellm.proxy.prompts import prompt_registry
custom_logger = MagicMock()
prompt_spec = MagicMock()
prompt_spec.litellm_params = MagicMock(prompt_id="resolved-id")
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
"get_prompt_callback_by_id",
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
)
logging_obj = MagicMock()
logging_obj.async_get_chat_completion_prompt = AsyncMock(
return_value=(
"gpt-4o-mini",
[
{"role": "user", "content": "You are a pirate."},
{"role": "user", "content": "Who are you?"},
],
{},
)
)
data: Dict[str, Any] = {"input": "Who are you?", "model": "anthropic-haiku-4-5", "prompt_id": "x"}
await proxy_logging._process_prompt_template(
data=data,
litellm_logging_obj=logging_obj,
prompt_id="x",
prompt_version=None,
call_type="aresponses",
)
assert data["model"] == "gpt-4o-mini"
assert data["input"] == [
{"role": "user", "content": "You are a pirate."},
{"role": "user", "content": "Who are you?"},
]
assert "messages" not in data
assert "prompt_id" not in data
hook_kwargs = logging_obj.async_get_chat_completion_prompt.await_args.kwargs
assert hook_kwargs["messages"] == [{"role": "user", "content": "Who are you?"}]
assert hook_kwargs["prompt_spec"] is prompt_spec

View file

@ -298,6 +298,22 @@ async def test_default_path_still_applies_prompt_templates(proxy_logging, make_u
process.assert_awaited_once()
@pytest.mark.asyncio
async def test_aresponses_call_type_applies_prompt_templates_before_routing(proxy_logging, make_user_api_key_auth, monkeypatch):
"""The responses surface must process registry prompts pre-routing so credentials follow the swapped model."""
monkeypatch.setattr(litellm, "callbacks", [])
proxy_logging.slack_alerting_instance = MagicMock(alerting=None)
process = AsyncMock()
monkeypatch.setattr(proxy_logging, "_process_prompt_template", process)
await proxy_logging.pre_call_hook(
user_api_key_dict=make_user_api_key_auth(),
data={"input": "hi", "model": "m", "prompt_id": "p1", "litellm_logging_obj": MagicMock()},
call_type="aresponses",
)
process.assert_awaited_once()
# ---------------------------------------------------------------------------
# enforces_request_content: which CustomLoggers a guardrails-only walk reaches
# ---------------------------------------------------------------------------

View file

@ -539,3 +539,84 @@ class TestAsyncResponsesAPIPromptManagement:
assert sent_input[0]["cache_control"] == {"type": "ephemeral"}
assert sent_input[1] == reasoning_item
assert sent_input[2]["id"] == "msg_1"
# ---------------------------------------------------------------------------
# Cross-provider model swap guard (prompt swaps model after credential resolution)
# ---------------------------------------------------------------------------
def test_resolve_prompt_swapped_provider_raises_cross_provider_with_credentials():
import litellm
from litellm.responses.main import _resolve_prompt_swapped_provider
with pytest.raises(litellm.BadRequestError, match="Refusing to send"):
_resolve_prompt_swapped_provider(
original_model="anthropic/claude-haiku-4-5",
swapped_model="gpt-4o-mini",
custom_llm_provider="anthropic",
kwargs={"api_key": "sk-ant-test"},
prompt_id="p1",
)
def test_resolve_prompt_swapped_provider_allows_swap_without_credentials():
from litellm.responses.main import _resolve_prompt_swapped_provider
assert (
_resolve_prompt_swapped_provider(
original_model="anthropic/claude-haiku-4-5",
swapped_model="gpt-4o-mini",
custom_llm_provider="anthropic",
kwargs={},
prompt_id="p1",
)
== "openai"
)
def test_resolve_prompt_swapped_provider_allows_same_provider_swap_with_credentials():
from litellm.responses.main import _resolve_prompt_swapped_provider
assert (
_resolve_prompt_swapped_provider(
original_model="openai/gpt-4o",
swapped_model="gpt-4o-mini",
custom_llm_provider="openai",
kwargs={"api_key": "sk-test", "api_base": "https://api.openai.com/v1"},
prompt_id="p1",
)
== "openai"
)
def test_sync_prompt_swap_cross_provider_with_credentials_raises():
import litellm
from litellm.responses.main import _apply_prompt_management_to_responses_call
logging_obj = _make_logging_obj("gpt-4o-mini", [{"role": "user", "content": "hi"}])
with pytest.raises(litellm.BadRequestError, match="Refusing to send"):
_apply_prompt_management_to_responses_call(
input="hi",
model="anthropic/claude-haiku-4-5",
custom_llm_provider="anthropic",
litellm_logging_obj=logging_obj,
kwargs={"prompt_id": "p1", "api_key": "sk-ant-test"},
local_vars={},
)
@pytest.mark.asyncio
async def test_aresponses_prompt_swap_cross_provider_with_credentials_raises():
import litellm
logging_obj = _make_logging_obj("gpt-4o-mini", [{"role": "user", "content": "hi"}])
logging_obj.async_failure_handler = AsyncMock()
with pytest.raises(litellm.BadRequestError, match="Refusing to send"):
await litellm.aresponses(
input="hi",
model="anthropic/claude-haiku-4-5",
litellm_logging_obj=logging_obj,
prompt_id="p1",
api_key="sk-ant-test",
)

View file

@ -724,3 +724,20 @@ class TestMergePromptManagementInputReshape:
)
assert result == merged
class TestResponsesInputToChatMessages:
def test_none_input_returns_empty_list(self):
assert ResponsesAPIRequestUtils.responses_input_to_chat_messages(None) == []
def test_str_input_becomes_user_message(self):
assert ResponsesAPIRequestUtils.responses_input_to_chat_messages("hi") == [
{"role": "user", "content": "hi"}
]
def test_list_input_keeps_only_role_items(self):
reasoning_item = {"type": "reasoning", "id": "rs_1", "summary": []}
user_message = {"role": "user", "content": "hi"}
assert ResponsesAPIRequestUtils.responses_input_to_chat_messages(
[reasoning_item, user_message, "stray"]
) == [user_message]