mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(responses): keep the addressed response id off bridged provider requests
Backport of #41689 to rc/1.102.0.
Cherry-picked from merge commit 07b5051c0d (main), originally by app/devin-ai-integration.
Both conflicts were in test files: the rc line lacks the Codex additional_tools tests that sit next to the new test in test_handler.py on main, and its test_utils.py imports all_litellm_params on a separate line, so only this PR's own additions (the imports, the recording handler, and the two new tests) are taken.
This commit is contained in:
parent
707a81aeae
commit
3a7a204b9b
4 changed files with 75 additions and 6 deletions
|
|
@ -22,7 +22,7 @@ from litellm.types.llms.openai import (
|
|||
BaseLiteLLMOpenAIResponseObject,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.utils import CallTypesLiteral, LLMResponseTypes, SpecialEnums
|
||||
from litellm.types.utils import ADDRESSED_RESPONSE_ID_FIELD, CallTypesLiteral, LLMResponseTypes, SpecialEnums
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -32,7 +32,6 @@ if TYPE_CHECKING:
|
|||
_RESPONSES_API_PROVIDER_PREFIX: Final = "/openai"
|
||||
_RESPONSES_API_CREATE_ROUTES: Final = frozenset({"/v1/responses", "/responses"})
|
||||
|
||||
_ADDRESSED_RESPONSE_ID_KEY: Final = "_litellm_addressed_response_id"
|
||||
_UNMANAGED_RESPONSE_ID_DETAIL: Final = (
|
||||
"Forbidden. This response id was not issued by this proxy, so the proxy cannot tell who owns it. "
|
||||
"To let keys address responses this proxy did not issue, set "
|
||||
|
|
@ -132,7 +131,7 @@ class ResponsesIDSecurity(CustomLogger):
|
|||
if call_type not in responses_api_call_types:
|
||||
return None
|
||||
addressed_id_field: Final = "previous_response_id" if call_type == "aresponses" else "response_id"
|
||||
retained_id: Final = data.get(_ADDRESSED_RESPONSE_ID_KEY)
|
||||
retained_id: Final = data.get(ADDRESSED_RESPONSE_ID_FIELD)
|
||||
addressed_id: Final = (
|
||||
retained_id if isinstance(retained_id, str) and retained_id else data.get(addressed_id_field)
|
||||
)
|
||||
|
|
@ -140,7 +139,7 @@ class ResponsesIDSecurity(CustomLogger):
|
|||
return data
|
||||
authorized_id: Final = self._authorize_response_id(addressed_id, user_api_key_dict)
|
||||
data[addressed_id_field] = authorized_id
|
||||
data[_ADDRESSED_RESPONSE_ID_KEY] = addressed_id
|
||||
data[ADDRESSED_RESPONSE_ID_FIELD] = addressed_id
|
||||
return data
|
||||
|
||||
def _authorize_response_id(
|
||||
|
|
|
|||
|
|
@ -3706,6 +3706,8 @@ agentic_loop_internal_litellm_params: Final = [
|
|||
# the provider.
|
||||
TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars"
|
||||
|
||||
ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id"
|
||||
|
||||
# Bedrock managed-batch deployment config, read from litellm_params by the batch and
|
||||
# files transformations. Listed for the same reason as the fields above: these sit on
|
||||
# a deployment that also serves chat, so leaking them into extra_body makes Bedrock
|
||||
|
|
@ -3720,7 +3722,7 @@ bedrock_batch_litellm_params: Final = (
|
|||
|
||||
all_litellm_params = (
|
||||
agentic_loop_internal_litellm_params
|
||||
+ [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params]
|
||||
+ [TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD, *bedrock_batch_litellm_params]
|
||||
+ [
|
||||
"metadata",
|
||||
"litellm_metadata",
|
||||
|
|
|
|||
|
|
@ -12,14 +12,21 @@ capture the forwarded kwargs; if the flag-setting line is removed the captured
|
|||
kwargs lack the flag and these tests fail.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.responses.litellm_completion_transformation.handler import (
|
||||
LiteLLMCompletionTransformationHandler,
|
||||
)
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.utils import ADDRESSED_RESPONSE_ID_FIELD
|
||||
|
||||
|
||||
class _StopForwarding(Exception):
|
||||
|
|
@ -68,3 +75,49 @@ async def test_async_fallback_tags_skip_responses_api_bridge():
|
|||
await coro
|
||||
|
||||
assert captured.get("_skip_responses_api_bridge") is True
|
||||
|
||||
|
||||
class _RecordingAnthropicHandler:
|
||||
def __init__(self, reply: Mapping[str, object]) -> None:
|
||||
self.reply: Final = reply
|
||||
self.request_body: Mapping[str, object] | None = None
|
||||
|
||||
def __call__(self, request: httpx.Request) -> httpx.Response:
|
||||
self.request_body = json.loads(request.content)
|
||||
return httpx.Response(200, json=dict(self.reply), request=request)
|
||||
|
||||
|
||||
_ANTHROPIC_MESSAGE_PAYLOAD: Final = {
|
||||
"id": "msg_turn_two",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-6",
|
||||
"content": [{"type": "text", "text": "14"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 12, "output_tokens": 1},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridged_follow_up_turn_keeps_the_addressed_response_id_off_the_provider_body():
|
||||
provider: Final = _RecordingAnthropicHandler(_ANTHROPIC_MESSAGE_PAYLOAD)
|
||||
client: Final = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=httpx.MockTransport(provider))
|
||||
|
||||
response = await litellm.aresponses(
|
||||
model="azure_ai/claude-sonnet-4-6",
|
||||
api_base="https://fake-foundry-resource.services.ai.azure.com",
|
||||
api_key="fake-api-key",
|
||||
input="Double it",
|
||||
previous_response_id="resp_turn_one",
|
||||
client=client,
|
||||
**{ADDRESSED_RESPONSE_ID_FIELD: "resp_turn_one"},
|
||||
)
|
||||
|
||||
assert provider.request_body is not None, "the bridged turn never reached the provider"
|
||||
assert ADDRESSED_RESPONSE_ID_FIELD not in provider.request_body, (
|
||||
f"the addressed response id reached the provider body: {sorted(provider.request_body)}"
|
||||
)
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
assert [item.type for item in response.output] == ["message"]
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.types.utils import (
|
|||
PromptTokensDetailsWrapper,
|
||||
StreamingChoices,
|
||||
Usage,
|
||||
ADDRESSED_RESPONSE_ID_FIELD,
|
||||
)
|
||||
from litellm.types.utils import all_litellm_params, bedrock_batch_litellm_params
|
||||
from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams
|
||||
|
|
@ -5603,6 +5604,20 @@ def test_get_litellm_params_keys_never_reach_the_provider():
|
|||
)
|
||||
|
||||
|
||||
def test_addressed_response_id_never_reaches_the_provider():
|
||||
kwargs = {
|
||||
"a_real_provider_specific_param": 1,
|
||||
ADDRESSED_RESPONSE_ID_FIELD: "resp_addressed-by-the-client",
|
||||
}
|
||||
|
||||
non_default = get_non_default_completion_params(kwargs)
|
||||
|
||||
assert non_default == {"a_real_provider_specific_param": 1}, (
|
||||
"the addressed response id leaked into the provider params: "
|
||||
f"{sorted(set(non_default) - {'a_real_provider_specific_param'})}"
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_batch_params_never_reach_the_provider():
|
||||
"""A Bedrock managed-batch deployment carries aws_batch_role_arn / s3_* /
|
||||
bedrock_tags in its litellm_params, and the same deployment also serves chat.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue