mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
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:
parent
f6571a653f
commit
dbc819dc77
11 changed files with 357 additions and 40 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 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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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