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:
Krrish Dholakia 2026-07-15 18:40:32 +00:00
parent 4580ad003a
commit ebb0f7e4cf
3 changed files with 148 additions and 14 deletions

View file

@ -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:

View file

@ -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]],

View file

@ -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.