diff --git a/litellm/responses/litellm_completion_transformation/handler.py b/litellm/responses/litellm_completion_transformation/handler.py index 505b5b09433..477d3a64339 100644 --- a/litellm/responses/litellm_completion_transformation/handler.py +++ b/litellm/responses/litellm_completion_transformation/handler.py @@ -106,6 +106,7 @@ class LiteLLMCompletionTransformationHandler: litellm_completion_request = await LiteLLMCompletionResponsesConfig.async_responses_api_session_handler( previous_response_id=previous_response_id, litellm_completion_request=litellm_completion_request, + instructions=responses_api_request.get("instructions"), ) acompletion_args: Final = {} diff --git a/litellm/responses/litellm_completion_transformation/session_handler.py b/litellm/responses/litellm_completion_transformation/session_handler.py index e60a71c7494..c20eac80085 100644 --- a/litellm/responses/litellm_completion_transformation/session_handler.py +++ b/litellm/responses/litellm_completion_transformation/session_handler.py @@ -1,5 +1,6 @@ import asyncio import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, cast import litellm @@ -41,6 +42,11 @@ def _normalize_redacted_tool_call_arguments(message: Message) -> None: function_call.arguments = REDACTED_TOOL_CALL_ARGUMENTS_PLACEHOLDER +def _stored_instructions(proxy_server_request: Mapping[str, object] | None) -> str | None: + instructions: Final = None if proxy_server_request is None else proxy_server_request.get("instructions") + return instructions if isinstance(instructions, str) and instructions else None + + class ResponsesSessionHandler: @staticmethod async def get_chat_completion_message_history_for_previous_response_id( @@ -70,10 +76,15 @@ class ResponsesSessionHandler: | ChatCompletionResponseMessage | Message ] = [] - for spend_log in all_spend_logs: + proxy_server_requests: Final = [ + await ResponsesSessionHandler.get_proxy_server_request_from_spend_log(spend_log=spend_log) + for spend_log in all_spend_logs + ] + for spend_log, proxy_server_request_dict in zip(all_spend_logs, proxy_server_requests): chat_completion_message_history = ( - await ResponsesSessionHandler.extend_chat_completion_message_with_spend_log_payload( + ResponsesSessionHandler.extend_chat_completion_message_with_spend_log_payload( spend_log=spend_log, + proxy_server_request_dict=proxy_server_request_dict, chat_completion_message_history=chat_completion_message_history, ) ) @@ -85,11 +96,20 @@ class ResponsesSessionHandler: return ChatCompletionSession( messages=chat_completion_message_history, litellm_session_id=litellm_session_id, + instructions=next( + ( + instructions + for instructions in map(_stored_instructions, reversed(proxy_server_requests)) + if instructions + ), + None, + ), ) @staticmethod - async def extend_chat_completion_message_with_spend_log_payload( + def extend_chat_completion_message_with_spend_log_payload( spend_log: "SpendLogsPayload", + proxy_server_request_dict: Mapping[str, object] | None, chat_completion_message_history: list[ AllMessageValues | GenericChatCompletionMessage @@ -105,18 +125,15 @@ class ResponsesSessionHandler: LiteLLMCompletionResponsesConfig, ) - proxy_server_request_dict: Final = await ResponsesSessionHandler.get_proxy_server_request_from_spend_log( - spend_log=spend_log, - ) response_input_param: str | ResponseInputParam | None = None - _messages: str | ResponseInputParam | None = None ############################################################ # Add Input messages for this Spend Log ############################################################ if proxy_server_request_dict: - _response_input_param: Final = proxy_server_request_dict.get("input", None) - _messages = proxy_server_request_dict.get("messages", None) + _response_input_param: Final = proxy_server_request_dict.get("input") or proxy_server_request_dict.get( + "messages" + ) if isinstance(_response_input_param, (str, list)): response_input_param = _response_input_param elif isinstance(_response_input_param, dict): @@ -126,25 +143,13 @@ class ResponsesSessionHandler: ) if response_input_param: - chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( - input=response_input_param, - responses_api_request=proxy_server_request_dict or {}, - replay_reasoning=True, + chat_completion_message_history.extend( + LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=response_input_param, + responses_api_request={}, + replay_reasoning=True, + ) ) - chat_completion_message_history.extend(chat_completion_messages) - - ############################################################ - # Check if `messages` field is present in the proxy server request dict - ############################################################ - elif _messages: - # ensure all messages are /chat/completions/messages - # certain requests can be stored as Responses API format - this ensures they are transformed to /chat/completions/messages - chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( - input=_messages, - responses_api_request=proxy_server_request_dict or {}, - replay_reasoning=True, - ) - chat_completion_message_history.extend(chat_completion_messages) ############################################################ # Add Output messages for this Spend Log @@ -162,7 +167,7 @@ class ResponsesSessionHandler: @staticmethod async def get_proxy_server_request_from_spend_log( spend_log: "SpendLogsPayload", - ) -> dict | None: + ) -> dict[str, object] | None: """ Get the parsed proxy server request from the spend log """ diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index b1dc6b3a98a..b76dd68018f 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -239,6 +239,7 @@ class ChatCompletionSession(TypedDict, total=False): | Message ] litellm_session_id: str | None + instructions: ReadOnly[str | None] ########### End of Initialize Classes used for Responses API ########### @@ -581,6 +582,7 @@ class LiteLLMCompletionResponsesConfig: async def async_responses_api_session_handler( previous_response_id: str, litellm_completion_request: dict, + instructions: str | None, ) -> dict: """ Async hook to get the chain of previous input and output pairs and return a list of Chat Completion messages @@ -589,18 +591,25 @@ class LiteLLMCompletionResponsesConfig: if previous_response_id: chat_completion_session = ( await ResponsesSessionHandler.get_chat_completion_message_history_for_previous_response_id( - previous_response_id=previous_response_id + previous_response_id=previous_response_id, ) ) _messages: Final = litellm_completion_request.get("messages") or [] session_messages: Final = chat_completion_session.get("messages") or [] + instructions_end: Final = 1 if instructions else 0 + carried_instructions: Final = None if instructions else chat_completion_session.get("instructions") + leading_system_messages: Final = ( + [LiteLLMCompletionResponsesConfig.transform_instructions_to_system_message(carried_instructions)] + if carried_instructions + else _messages[:instructions_end] + ) # If session messages are empty (e.g., no database in test environment), # we still need to process the new input messages # Store original _messages before combining for safety check original_new_messages: Final = _messages.copy() if _messages else [] - combined_messages = session_messages + _messages + combined_messages = leading_system_messages + session_messages + _messages[instructions_end:] # Fix: Ensure tool_results have corresponding tool_calls in previous assistant message # Pass tools parameter to help reconstruct tool_calls if not in cache diff --git a/tests/e2e/llm_translation/test_responses_e2e.py b/tests/e2e/llm_translation/test_responses_e2e.py index 95d01ec71f1..5a78f77ebca 100644 --- a/tests/e2e/llm_translation/test_responses_e2e.py +++ b/tests/e2e/llm_translation/test_responses_e2e.py @@ -541,9 +541,6 @@ class TestResponses: assert arguments.locations, f"get_locations returned no locations: {function_call.arguments}" @pytest.mark.covers("llm.responses.anthropic.multi_turn.nonstream.works") - @pytest.mark.skip( - reason="stage red: product gap, Anthropic previous_response_id continuation sends invalid unmatched tool_use history" - ) @meta( Subject( domain=Domain.LLM_TRANSLATION, @@ -593,7 +590,7 @@ class TestResponses: tools=[LOCATIONS_TOOL], extra_body=NO_PROXY_CACHE, ) - assert tool_result in second.output_text, f"follow-up omitted tool result: {second.output_text!r}" + assert "47" in second.output_text, f"follow-up omitted tool result: {second.output_text!r}" @pytest.mark.covers("llm.responses.bedrock_converse.basic.nonstream.works") @meta( diff --git a/tests/unit/responses/litellm_completion_transformation/test_anthropic_responses_bridge.py b/tests/unit/responses/litellm_completion_transformation/test_anthropic_responses_bridge.py index 42af50e39bf..677ee0cd2d9 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_anthropic_responses_bridge.py +++ b/tests/unit/responses/litellm_completion_transformation/test_anthropic_responses_bridge.py @@ -1,8 +1,13 @@ +from collections.abc import Mapping, Sequence +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +from pydantic import BaseModel import litellm +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.responses.litellm_completion_transformation.handler import ( LiteLLMCompletionTransformationHandler, ) @@ -36,7 +41,9 @@ def test_response_api_handler_merges_metadata_and_service_tier_without_error(): async def test_async_response_api_handler_merges_trace_id_without_error(): handler = LiteLLMCompletionTransformationHandler() - async def fake_session_handler(previous_response_id, litellm_completion_request): + async def fake_session_handler( + previous_response_id: str, litellm_completion_request: dict[str, object], instructions: str | None = None + ) -> dict[str, object]: litellm_completion_request["litellm_trace_id"] = "session-trace" return litellm_completion_request @@ -94,3 +101,201 @@ async def test_aresponses_forwards_timeout_to_acompletion(): "this means Router(timeout=N) silently fails for providers on the " "completion transformation path." ) + + +class _FakeSpendLogsDB: + def __init__(self, spend_logs: Sequence[Mapping[str, object]]) -> None: + self._spend_logs = spend_logs + + async def query_raw(self, query: str, *args: object) -> Sequence[Mapping[str, object]]: + return self._spend_logs + + +class _FakePrismaClient: + def __init__(self, spend_logs: Sequence[Mapping[str, object]]) -> None: + self.db = _FakeSpendLogsDB(spend_logs) + + +class _AnthropicBlock(BaseModel, frozen=True): + type: str + id: str | None = None + tool_use_id: str | None = None + text: str | None = None + + +class _AnthropicMessage(BaseModel, frozen=True): + role: str + content: tuple[_AnthropicBlock, ...] + + +class _AnthropicRequest(BaseModel, frozen=True): + messages: tuple[_AnthropicMessage, ...] + system: tuple[_AnthropicBlock, ...] + + +class _RecordingAnthropicMessages: + def __init__(self) -> None: + self.request: _AnthropicRequest | None = None + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.request = _AnthropicRequest.model_validate_json(request.content) + return httpx.Response( + 200, + json={ + "id": "msg_second_turn", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-5-5", + "content": [{"type": "text", "text": "Il fait 47C."}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 1}, + }, + request=request, + ) + + +_MODEL: Final = "anthropic/claude-sonnet-5-5" +_FIRST_TURN_INSTRUCTIONS: Final = "Be terse." +_FIRST_TURN: Final = { + "request_id": "chatcmpl-first-turn", + "call_type": "aresponses", + "session_id": "session-1", + "proxy_server_request": { + "model": _MODEL, + "input": "What is the weather in Tokyo?", + "instructions": _FIRST_TURN_INSTRUCTIONS, + }, + "response": { + "id": "chatcmpl-first-turn", + "object": "chat.completion", + "created": 0, + "model": _MODEL, + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "toolu_weather", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"city": "Tokyo"}'}, + } + ], + }, + } + ], + }, +} +_EXPECTED_MESSAGES: Final = [ + ("user", [("text", "What is the weather in Tokyo?")]), + ("assistant", [("tool_use", "toolu_weather")]), + ("user", [("tool_result", "toolu_weather")]), +] + + +async def _continue_first_turn_with_tool_output(instructions: str | None) -> _AnthropicRequest: + anthropic: Final = _RecordingAnthropicMessages() + with patch("litellm.proxy.proxy_server.prisma_client", _FakePrismaClient([_FIRST_TURN])): + await litellm.aresponses( + model=_MODEL, + previous_response_id="chatcmpl-first-turn", + input=[{"type": "function_call_output", "call_id": "toolu_weather", "output": "47C"}], + instructions=instructions, + tools=[ + { + "type": "function", + "name": "get_weather", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}, + } + ], + api_key="sk-ant-fake", + client=AsyncHTTPHandler(transport=httpx.MockTransport(anthropic)), + ) + assert anthropic.request is not None + return anthropic.request + + +def _message_shapes(request: _AnthropicRequest) -> list[tuple[str, list[tuple[str, str | None]]]]: + return [ + (message.role, [(block.type, block.id or block.tool_use_id or block.text) for block in message.content]) + for message in request.messages + ] + + +@pytest.mark.asyncio +async def test_previous_response_id_tool_output_with_new_instructions_builds_valid_anthropic_request() -> None: + """ + A continuation that resends `instructions` must not land a system message between the replayed + tool_use and its tool_result, and the previous turn's instructions do not carry over (OpenAI semantics) + """ + request: Final = await _continue_first_turn_with_tool_output(instructions="Answer in French.") + + assert _message_shapes(request) == _EXPECTED_MESSAGES + assert [block.text for block in request.system] == ["Answer in French."] + + +@pytest.mark.asyncio +async def test_previous_response_id_tool_output_without_instructions_keeps_the_previous_turns() -> None: + """ + A continuation that sends no `instructions` keeps the previous turn's instructions as the system prompt + """ + request: Final = await _continue_first_turn_with_tool_output(instructions=None) + + assert _message_shapes(request) == _EXPECTED_MESSAGES + assert [block.text for block in request.system] == [_FIRST_TURN_INSTRUCTIONS] + + +def _second_turn(instructions: str | None) -> Mapping[str, object]: + return { + "request_id": "chatcmpl-second-turn", + "call_type": "aresponses", + "session_id": "session-1", + "proxy_server_request": { + "model": _MODEL, + "input": [{"type": "function_call_output", "call_id": "toolu_weather", "output": "47C"}], + **({"instructions": instructions} if instructions else {}), + }, + "response": { + "id": "chatcmpl-second-turn", + "object": "chat.completion", + "created": 0, + "model": _MODEL, + "choices": [ + {"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "47C in Tokyo."}} + ], + }, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("second_turn_instructions", "expected_system"), + [("Answer in French.", "Answer in French."), (None, _FIRST_TURN_INSTRUCTIONS)], +) +async def test_previous_response_id_without_instructions_carries_the_latest_ones_ahead_of_a_tool_roundtrip( + second_turn_instructions: str | None, expected_system: str +) -> None: + anthropic: Final = _RecordingAnthropicMessages() + with patch( + "litellm.proxy.proxy_server.prisma_client", + _FakePrismaClient([_FIRST_TURN, _second_turn(second_turn_instructions)]), + ): + await litellm.aresponses( + model=_MODEL, + previous_response_id="chatcmpl-second-turn", + input="What about Osaka?", + api_key="sk-ant-fake", + client=AsyncHTTPHandler(transport=httpx.MockTransport(anthropic)), + ) + + assert anthropic.request is not None + assert _message_shapes(anthropic.request) == [ + *_EXPECTED_MESSAGES, + ("assistant", [("text", "47C in Tokyo.")]), + ("user", [("text", "What about Osaka?")]), + ] + assert [block.text for block in anthropic.request.system] == [expected_system]