From c026dd153c8f68635b0b7bc28ea17f8222c0eb01 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 15 Jul 2026 22:17:59 +0000 Subject: [PATCH] fix(responses): apply deployment defaults to websocket frames --- litellm/llms/custom_httpx/llm_http_handler.py | 2 + litellm/responses/main.py | 21 +++++++ litellm/responses/streaming_iterator.py | 21 ++++++- .../test_responses_websocket_all_providers.py | 55 +++++++++++++++++++ 4 files changed, 96 insertions(+), 3 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index b47fc50e196..3d7cd46eef6 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6010,6 +6010,7 @@ class BaseLLMHTTPHandler: litellm_metadata: Optional[Dict[str, Any]] = None, custom_llm_provider: Optional[str] = None, first_message: Optional[str] = None, + request_defaults: dict[str, object] | None = None, **kwargs: Any, ): """ @@ -6141,6 +6142,7 @@ class BaseLLMHTTPHandler: guardrail_callbacks=_ws_guardrail_callbacks, output_guardrail_callbacks=_ws_output_guardrail_callbacks, authorized_model=model, + request_defaults=request_defaults, ) await streaming.bidirectional_forward() diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 12f9be970c7..030291449b6 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -1,5 +1,6 @@ import asyncio import contextvars +from collections.abc import Mapping from functools import partial from typing import ( TYPE_CHECKING, @@ -1969,6 +1970,25 @@ def _build_litellm_metadata_for_ws(kwargs: dict) -> dict: return metadata +def _build_responses_websocket_request_defaults(kwargs: Mapping[str, object]) -> dict[str, object]: + valid_keys = ResponsesAPIOptionalRequestParams.__annotations__.keys() + optional_params = {key: value for key, value in kwargs.items() if key in valid_keys and value is not None} + reasoning_effort = kwargs.get("reasoning_effort") + if "reasoning" not in optional_params and isinstance(reasoning_effort, str): + reasoning = LiteLLMResponsesTransformationHandler()._map_reasoning_effort(reasoning_effort) + if reasoning is not None: + optional_params["reasoning"] = reasoning + elif "reasoning" not in optional_params and isinstance(reasoning_effort, dict): + optional_params["reasoning"] = reasoning_effort + extra_body = kwargs.get("extra_body") + provider_defaults = ( + {key: value for key, value in extra_body.items() if isinstance(key, str)} + if isinstance(extra_body, dict) + else {} + ) + return {**optional_params, **provider_defaults} + + @client async def _aresponses_websocket( model: str, @@ -2058,5 +2078,6 @@ async def _aresponses_websocket( user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata_for_ws(kwargs), custom_llm_provider=_custom_llm_provider, + request_defaults=_build_responses_websocket_request_defaults(kwargs), **remaining_kwargs, ) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index eb78e6f9c8d..8e36270c603 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1270,6 +1270,7 @@ class ResponsesWebSocketStreaming: guardrail_callbacks: Optional[List[Any]] = None, output_guardrail_callbacks: Optional[List[Any]] = None, authorized_model: Optional[str] = None, + request_defaults: dict[str, object] | None = None, ): self.websocket = websocket self.backend_ws = backend_ws @@ -1284,6 +1285,7 @@ class ResponsesWebSocketStreaming: # Model name authorized at connection time; enforced on every # response.create frame to prevent deployment-substitution attacks. self.authorized_model: Optional[str] = authorized_model + self.request_defaults: dict[str, object] = request_defaults or {} def _should_store_event(self, event_obj: dict) -> bool: return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES @@ -1425,6 +1427,19 @@ class ResponsesWebSocketStreaming: modified = True return modified + def _apply_request_defaults(self, msg_obj: dict[str, object]) -> bool: + nested = msg_obj.get("response") + request = ( + {key: value for key, value in nested.items() if isinstance(key, str)} + if isinstance(nested, dict) + else msg_obj + ) + if request is not msg_obj: + msg_obj["response"] = request + missing_defaults = {key: value for key, value in self.request_defaults.items() if key not in request} + request.update(missing_defaults) + return bool(missing_defaults) + async def _mask_response_create(self, message: str) -> str: """ Enforce the authorized model and apply Presidio PII masking to a @@ -1447,16 +1462,16 @@ class ResponsesWebSocketStreaming: if msg_obj.get("type") != "response.create": return message - # Always enforce the authorized model, even when PII masking is off. + defaults_modified = self._apply_request_defaults(msg_obj) model_modified = self._enforce_authorized_model(msg_obj) if not self.guardrail_callbacks: - return json.dumps(msg_obj) if model_modified else message + return json.dumps(msg_obj) if defaults_modified or model_modified else message if "metadata" not in self.request_data: self.request_data["metadata"] = {} - modified = model_modified + modified = defaults_modified or model_modified for cb in self.guardrail_callbacks: presidio_config = cb.get_presidio_settings_from_request_data(self.request_data) # response.create carries client text in two shapes: diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/test_litellm/responses/test_responses_websocket_all_providers.py index 4509abc7749..517837c4a40 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -1057,6 +1057,61 @@ class TestNativeWebSocketGuardrails: json.loads(nested_message)["response"]["model"] == "authorized-deployment" ) + @pytest.mark.asyncio + async def test_native_websocket_merges_deployment_defaults(self): + import asyncio + from unittest.mock import AsyncMock, patch + + from litellm.responses.main import _aresponses_websocket + + class FakeBackendWebSocket: + def __init__(self): + self.send = AsyncMock() + self.close = AsyncMock() + + async def recv(self, decode=False): + await asyncio.Future() + + backend_websocket = FakeBackendWebSocket() + + class FakeConnect: + def __init__(self, url, **kwargs): + pass + + async def __aenter__(self): + return backend_websocket + + async def __aexit__(self, *args): + pass + + first_message = json.dumps( + { + "type": "response.create", + "model": "gpt-4o-mini", + "input": "hi", + "service_tier": "default", + } + ) + websocket = MagicMock() + websocket.receive_text = AsyncMock(side_effect=RuntimeError("disconnect")) + + with patch("websockets.connect", FakeConnect): + await _aresponses_websocket.__wrapped__( + model="openai/gpt-4o-mini", + websocket=websocket, + api_key="sk-test", + litellm_logging_obj=MagicMock(), + first_message=first_message, + reasoning_effort="high", + service_tier="priority", + extra_body={"provider_default": "configured"}, + ) + + sent_message = json.loads(backend_websocket.send.await_args.args[0]) + assert sent_message["reasoning"] == {"effort": "high"} + assert sent_message["service_tier"] == "default" + assert sent_message["provider_default"] == "configured" + @pytest.mark.asyncio async def test_completed_event_with_null_response_passes_through(self): from unittest.mock import MagicMock