mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(responses): preserve reasoning through prompt hooks
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
4580ad003a
commit
ebb0f7e4cf
3 changed files with 148 additions and 14 deletions
|
|
@ -494,7 +494,14 @@ async def aresponses(
|
|||
prompt_label=kwargs.get("prompt_label", None),
|
||||
prompt_version=kwargs.get("prompt_version", None),
|
||||
)
|
||||
input = cast(Union[str, ResponseInputParam], merged_input)
|
||||
input = cast(
|
||||
Union[str, ResponseInputParam],
|
||||
ResponsesAPIRequestUtils.merge_prompt_management_input(
|
||||
original_input=input,
|
||||
client_input=client_input,
|
||||
merged_input=merged_input,
|
||||
),
|
||||
)
|
||||
if model != original_model:
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
|
||||
kwargs.pop("prompt_id", None)
|
||||
|
|
@ -609,7 +616,14 @@ def _apply_prompt_management_to_responses_call(
|
|||
prompt_label=kwargs.get("prompt_label", None),
|
||||
prompt_version=kwargs.get("prompt_version", None),
|
||||
)
|
||||
input = cast(Union[str, ResponseInputParam], merged_input)
|
||||
input = cast(
|
||||
Union[str, ResponseInputParam],
|
||||
ResponsesAPIRequestUtils.merge_prompt_management_input(
|
||||
original_input=input,
|
||||
client_input=client_input,
|
||||
merged_input=merged_input,
|
||||
),
|
||||
)
|
||||
local_vars["input"] = input
|
||||
local_vars["model"] = model
|
||||
if model != original_model:
|
||||
|
|
|
|||
|
|
@ -19,7 +19,9 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ResponseAPIUsage,
|
||||
ResponseInputParam,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
ResponseText,
|
||||
|
|
@ -36,6 +38,54 @@ from litellm.types.utils import (
|
|||
class ResponsesAPIRequestUtils:
|
||||
"""Helper utils for constructing ResponseAPI requests"""
|
||||
|
||||
@staticmethod
|
||||
def merge_prompt_management_input(
|
||||
original_input: str | ResponseInputParam,
|
||||
client_input: list[AllMessageValues],
|
||||
merged_input: list[AllMessageValues],
|
||||
) -> list[object]:
|
||||
if isinstance(original_input, str):
|
||||
return [*merged_input]
|
||||
|
||||
original_items = tuple(original_input)
|
||||
client_item_ids = frozenset(id(item) for item in client_input)
|
||||
message_positions = tuple(index for index, item in enumerate(original_items) if id(item) in client_item_ids)
|
||||
|
||||
if len(message_positions) == len(original_items):
|
||||
return [*merged_input]
|
||||
if not message_positions:
|
||||
return [*merged_input, *original_items]
|
||||
|
||||
corresponding_messages = len(client_input) == len(merged_input) and all(
|
||||
original.get("role") == merged.get("role")
|
||||
and (not isinstance(original.get("id"), str) or original.get("id") == merged.get("id"))
|
||||
for original, merged in zip(client_input, merged_input)
|
||||
)
|
||||
if corresponding_messages:
|
||||
merged_by_position = dict(zip(message_positions, merged_input))
|
||||
return [
|
||||
merged_by_position[index] if index in merged_by_position else item
|
||||
for index, item in enumerate(original_items)
|
||||
]
|
||||
|
||||
all_messages_preserved = all(any(original is merged for merged in merged_input) for original in client_input)
|
||||
if all_messages_preserved:
|
||||
prefixes = {
|
||||
id(original_items[position]): original_items[
|
||||
message_positions[index - 1] + 1 if index else 0 : position
|
||||
]
|
||||
for index, position in enumerate(message_positions)
|
||||
}
|
||||
trailing_items = original_items[message_positions[-1] + 1 :]
|
||||
return [item for merged in merged_input for item in (*prefixes.get(id(merged), ()), merged)] + list(
|
||||
trailing_items
|
||||
)
|
||||
|
||||
verbose_logger.warning(
|
||||
"Prompt management hook replaced Responses API messages; non-message input items were dropped"
|
||||
)
|
||||
return [*merged_input]
|
||||
|
||||
@staticmethod
|
||||
def _check_valid_arg(
|
||||
supported_params: Optional[List[str]],
|
||||
|
|
|
|||
|
|
@ -14,13 +14,19 @@ Covers:
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import List
|
||||
from typing import List, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.anthropic_cache_control_hook import (
|
||||
AnthropicCacheControlHook,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ResponseInputParam,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
|
|
@ -54,18 +60,15 @@ def _patch_responses_dispatch():
|
|||
return_value=("gpt-4o", "openai", None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.responses.mcp.litellm_proxy_mcp_handler."
|
||||
"LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway",
|
||||
"litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.responses.main.ProviderConfigManager"
|
||||
".get_provider_responses_api_config",
|
||||
"litellm.responses.main.ProviderConfigManager.get_provider_responses_api_config",
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"litellm.responses.main.litellm_completion_transformation_handler"
|
||||
".response_api_handler",
|
||||
"litellm.responses.main.litellm_completion_transformation_handler.response_api_handler",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
]
|
||||
|
|
@ -77,7 +80,6 @@ def _patch_responses_dispatch():
|
|||
|
||||
|
||||
class TestResponsesAPIPromptManagement:
|
||||
|
||||
def test_str_input_coerced_and_merged(self):
|
||||
"""[A] str input is wrapped into a message list before being passed to the hook."""
|
||||
template_messages: List[AllMessageValues] = [
|
||||
|
|
@ -108,9 +110,7 @@ class TestResponsesAPIPromptManagement:
|
|||
logging_obj.get_chat_completion_prompt.assert_called_once()
|
||||
call_kwargs = logging_obj.get_chat_completion_prompt.call_args.kwargs
|
||||
# str was coerced to a single user message before being passed to the hook
|
||||
assert call_kwargs["messages"] == [
|
||||
{"role": "user", "content": "Tell me about AI."}
|
||||
]
|
||||
assert call_kwargs["messages"] == [{"role": "user", "content": "Tell me about AI."}]
|
||||
assert call_kwargs["prompt_id"] == "summariser-prompt"
|
||||
|
||||
def test_list_input_merged_with_template(self):
|
||||
|
|
@ -256,6 +256,76 @@ class TestResponsesAPIPromptManagement:
|
|||
assert all(isinstance(m, dict) and "role" in m for m in passed_messages)
|
||||
assert len(passed_messages) == 1
|
||||
|
||||
def test_cache_control_hook_preserves_reasoning_items(self):
|
||||
system_message = cast(
|
||||
AllMessageValues,
|
||||
{"role": "system", "content": "Analyze the request"},
|
||||
)
|
||||
assistant_message = cast(
|
||||
AllMessageValues,
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_1",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [
|
||||
{
|
||||
"type": "output_text",
|
||||
"text": "The code has a bug",
|
||||
"annotations": [],
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
user_message = cast(
|
||||
AllMessageValues,
|
||||
{"role": "user", "content": "Check for security issues"},
|
||||
)
|
||||
reasoning_item = {
|
||||
"type": "reasoning",
|
||||
"id": "rs_1",
|
||||
"summary": [],
|
||||
"encrypted_content": "encrypted",
|
||||
}
|
||||
original_input = cast(
|
||||
ResponseInputParam,
|
||||
[system_message, reasoning_item, assistant_message, user_message],
|
||||
)
|
||||
_, merged_messages, _ = AnthropicCacheControlHook().get_chat_completion_prompt(
|
||||
model="azure/gpt-5-codex",
|
||||
messages=[system_message, assistant_message, user_message],
|
||||
non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]},
|
||||
prompt_id=None,
|
||||
prompt_variables=None,
|
||||
dynamic_callback_params={},
|
||||
)
|
||||
logging_obj = _make_logging_obj(
|
||||
merged_model="azure/gpt-5-codex",
|
||||
merged_messages=merged_messages,
|
||||
)
|
||||
|
||||
patches = _patch_responses_dispatch()
|
||||
with patches[0], patches[1], patches[2], patches[3] as mock_handler:
|
||||
import litellm
|
||||
|
||||
litellm.responses(
|
||||
input=original_input,
|
||||
model="azure/gpt-5-codex",
|
||||
litellm_logging_obj=logging_obj,
|
||||
cache_control_injection_points=[{"location": "message", "role": "system"}],
|
||||
)
|
||||
|
||||
sent_input = mock_handler.call_args.kwargs["input"]
|
||||
assert [item.get("type") for item in sent_input] == [
|
||||
None,
|
||||
"reasoning",
|
||||
"message",
|
||||
None,
|
||||
]
|
||||
assert sent_input[0]["cache_control"] == {"type": "ephemeral"}
|
||||
assert sent_input[1] == reasoning_item
|
||||
assert sent_input[2]["id"] == "msg_1"
|
||||
|
||||
def test_model_override_re_resolves_provider(self):
|
||||
"""[G] When the prompt template overrides the model to a different provider,
|
||||
custom_llm_provider is re-resolved so downstream routing uses the correct provider.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue