mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(responses): keep tool_result next to tool_use on Anthropic previous_response_id continuations (#45322)
* fix(responses): keep tool_result next to tool_use on Anthropic previous_response_id continuations Replayed history no longer re-adds each stored turn's instructions, and the current request's instructions go before the replayed history instead of between the last tool_use and its tool_result * test(responses): fake the Anthropic transport instead of acompletion in the continuation regression test * fix(responses): keep the previous turn's instructions when a continuation sends none * fix(responses): carry only the latest stored instructions when a continuation sends none Replaying each spend log's own instructions put a system message between a replayed tool_use and its tool_result, which Anthropic rejects with a 400
This commit is contained in:
parent
92654b69a9
commit
ce6889bbf7
5 changed files with 252 additions and 35 deletions
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue