mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #38407 from BerriAI/litellm_fix_dotprompt_model_swap
fix(prompts): apply prompt templates before routing on /v1/responses and honor ignore_prompt_manager_model
This commit is contained in:
commit
5175fda0af
11 changed files with 411 additions and 67 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 ""
|
||||
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,
|
||||
|
|
|
|||
|
|
@ -28,7 +28,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,
|
||||
|
|
@ -519,10 +518,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)
|
||||
with _prompt_management_sees_a_provisional_message_list(
|
||||
kwargs,
|
||||
bridged=_will_bridge_to_chat_completions(
|
||||
|
|
@ -551,7 +547,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
|
||||
|
||||
|
|
@ -621,6 +623,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,
|
||||
|
|
@ -640,10 +671,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
|
||||
|
|
@ -676,7 +704,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
|
||||
|
|
@ -994,6 +1028,33 @@ def responses(
|
|||
# Update local_vars to include the converted text parameter
|
||||
local_vars["text"] = text
|
||||
|
||||
#########################################################
|
||||
# PROMPT MANAGEMENT
|
||||
# If aresponses() already ran the async hook, it pops prompt_id and
|
||||
# passes the result via _async_prompt_merged_params — apply those
|
||||
# directly and skip the sync hook to avoid double-merging.
|
||||
#########################################################
|
||||
_stripped_model, _from_chat_completions_prefix = _normalize_openai_chat_completions_responses_model(model)
|
||||
model = _stripped_model
|
||||
local_vars["model"] = model
|
||||
use_chat_completions_api = use_chat_completions_api or _from_chat_completions_prefix
|
||||
|
||||
if custom_llm_provider is None:
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model, api_base=local_vars.get("base_url", None)
|
||||
)
|
||||
local_vars["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
input, model, custom_llm_provider = _apply_prompt_management_to_responses_call(
|
||||
input=input,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
kwargs=kwargs,
|
||||
local_vars=local_vars,
|
||||
use_chat_completions_api=use_chat_completions_api,
|
||||
)
|
||||
|
||||
# get llm provider logic
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
|
|
@ -1003,11 +1064,6 @@ def responses(
|
|||
if litellm_params.mock_response and isinstance(litellm_params.mock_response, str):
|
||||
return mock_responses_api_response(mock_response=litellm_params.mock_response)
|
||||
|
||||
_stripped_model, _from_chat_completions_prefix = _normalize_openai_chat_completions_responses_model(model)
|
||||
model = _stripped_model
|
||||
local_vars["model"] = model
|
||||
use_chat_completions_api = use_chat_completions_api or _from_chat_completions_prefix
|
||||
|
||||
model, custom_llm_provider = _resolve_model_provider_for_responses(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1015,22 +1071,6 @@ def responses(
|
|||
local_vars=local_vars,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# PROMPT MANAGEMENT
|
||||
# If aresponses() already ran the async hook, it pops prompt_id and
|
||||
# passes the result via _async_prompt_merged_params — apply those
|
||||
# directly and skip the sync hook to avoid double-merging.
|
||||
#########################################################
|
||||
input, model, custom_llm_provider = _apply_prompt_management_to_responses_call(
|
||||
input=input,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
kwargs=kwargs,
|
||||
local_vars=local_vars,
|
||||
use_chat_completions_api=use_chat_completions_api,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Update input and tools with provider-specific file IDs if managed files are used
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -12,7 +12,9 @@ from unittest.mock import MagicMock, Mock, patch
|
|||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.dotprompt.dotprompt_manager import DotpromptManager
|
||||
from litellm.integrations.dotprompt.prompt_manager import PromptManager, PromptTemplate
|
||||
from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec
|
||||
|
||||
|
||||
def test_prompt_manager_initialization():
|
||||
|
|
@ -657,3 +659,90 @@ def test_prompt_initializer_registers_flat_db_prompt_under_base_id():
|
|||
template = dotprompt_manager.prompt_manager.get_prompt("agent-prompt")
|
||||
assert template is not None
|
||||
assert template.content == "AHOY {{name}}"
|
||||
|
||||
|
||||
def _swap_prompt_manager_and_spec(ignore_prompt_manager_model: bool) -> tuple[DotpromptManager, 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"
|
||||
|
|
|
|||
|
|
@ -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, object] = {"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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -52,12 +52,19 @@ def _make_logging_obj(
|
|||
return logging_obj
|
||||
|
||||
|
||||
def _provider_by_model(model: str, **_: object) -> tuple[str, str, None, None]:
|
||||
provider, _, bare_model = model.partition("/")
|
||||
if not bare_model:
|
||||
return (model, "anthropic" if "claude" in model else "openai", None, None)
|
||||
return (bare_model, provider, None, None)
|
||||
|
||||
|
||||
def _patch_responses_dispatch():
|
||||
"""Patch everything after the prompt management block so tests stay unit-level."""
|
||||
return [
|
||||
patch(
|
||||
"litellm.responses.main.litellm.get_llm_provider",
|
||||
return_value=("gpt-4o", "openai", None, None),
|
||||
side_effect=_provider_by_model,
|
||||
),
|
||||
patch(
|
||||
"litellm.responses.mcp.litellm_proxy_mcp_handler."
|
||||
|
|
@ -278,7 +285,7 @@ class TestResponsesAPIPromptManagement:
|
|||
|
||||
# The model passed to the downstream handler should be the overridden one
|
||||
handler_call_kwargs = mock_handler.call_args.kwargs
|
||||
assert handler_call_kwargs.get("model") == "openai/gpt-4o-mini"
|
||||
assert handler_call_kwargs.get("model") == "gpt-4o-mini"
|
||||
|
||||
def test_non_message_input_items_filtered(self):
|
||||
"""[F] Non-message items in ResponseInputParam (e.g. function_call_output) are
|
||||
|
|
@ -388,10 +395,7 @@ class TestResponsesAPIPromptManagement:
|
|||
with (
|
||||
patch(
|
||||
"litellm.responses.main.litellm.get_llm_provider",
|
||||
side_effect=[
|
||||
("gpt-4o", "openai", None, None),
|
||||
("claude-3-5-sonnet", "anthropic", None, None),
|
||||
],
|
||||
side_effect=_provider_by_model,
|
||||
),
|
||||
patches[1],
|
||||
patches[2],
|
||||
|
|
@ -539,3 +543,102 @@ 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_resolves_credentials_for_swapped_provider(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setenv("XAI_API_KEY", "sk-xai-test")
|
||||
logging_obj = _make_logging_obj("gpt-4o-mini", [{"role": "user", "content": "hi"}])
|
||||
with patch( # test-quality-ok: handler boundary stub proves creds resolve for the swapped provider without network
|
||||
"litellm.responses.main.base_llm_http_handler.response_api_handler", return_value=MagicMock()
|
||||
) as mock_handler:
|
||||
litellm.responses(input="hi", model="xai/grok-4", prompt_id="p1", litellm_logging_obj=logging_obj)
|
||||
|
||||
handler_kwargs = mock_handler.call_args.kwargs
|
||||
assert handler_kwargs["model"] == "gpt-4o-mini"
|
||||
assert handler_kwargs["custom_llm_provider"] == "openai"
|
||||
assert handler_kwargs["litellm_params"].api_base is None
|
||||
assert handler_kwargs["litellm_params"].api_key != "sk-xai-test"
|
||||
|
||||
|
||||
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={},
|
||||
use_chat_completions_api=False,
|
||||
)
|
||||
|
||||
|
||||
@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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue