fix(responses): apply deployment defaults to websocket frames

This commit is contained in:
Devin AI 2026-07-15 22:17:59 +00:00
parent 9121ae3024
commit c026dd153c
4 changed files with 96 additions and 3 deletions

View file

@ -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()

View file

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

View file

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

View file

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