Merge pull request #41881 from BerriAI/litellm_responses_ws_deployment_defaults

fix(responses): merge deployment litellm_params into native websocket response.create frames
This commit is contained in:
Mateo Wang 2026-09-18 16:04:39 -07:00 committed by GitHub
commit ec9435cbf4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 314 additions and 8 deletions

View file

@ -144,6 +144,7 @@ from litellm.types.llms.openai import (
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.rerank import RerankResponse
from litellm.types.responses.main import DeleteResponseResult
from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
CallTypes,
@ -6588,6 +6589,7 @@ class BaseLLMHTTPHandler:
litellm_metadata: dict[str, object] | None = None,
custom_llm_provider: str | None = None,
first_message: str | None = None,
request_defaults: ResponsesWebSocketRequestDefaults | None = None,
**kwargs: Any,
):
"""
@ -6742,6 +6744,7 @@ class BaseLLMHTTPHandler:
output_guardrail_callbacks=_ws_output_guardrail_callbacks,
quota_callbacks=_ws_quota_callbacks,
authorized_model=model,
request_defaults=request_defaults,
)
await streaming.bidirectional_forward()

View file

@ -8,7 +8,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias, cast
import httpx
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter
from typing_extensions import assert_never
import litellm
@ -53,6 +53,7 @@ from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.llms.openai.data_residency import infer_openai_data_residency
from litellm.secret_managers.main import get_secret_str
from litellm.types.responses.main import *
from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import all_litellm_params
from litellm.utils import (
@ -2261,6 +2262,31 @@ def _build_litellm_metadata_for_ws(kwargs: dict) -> dict:
return metadata
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object] | None)
def _deployment_reasoning_default(kwargs: Mapping[str, object]) -> Reasoning | dict[str, object] | None:
if kwargs.get("reasoning") is not None:
return None
reasoning_effort: Final = kwargs.get("reasoning_effort")
if isinstance(reasoning_effort, str):
return LiteLLMResponsesTransformationHandler()._map_reasoning_effort(reasoning_effort)
return _JSON_OBJECT_ADAPTER.validate_python(reasoning_effort) if isinstance(reasoning_effort, Mapping) else None
def _build_responses_websocket_request_defaults(kwargs: Mapping[str, object]) -> ResponsesWebSocketRequestDefaults:
default_reasoning: Final = _deployment_reasoning_default(kwargs)
candidate_params: Final[dict[str, object]] = {
**kwargs,
**({"reasoning": default_reasoning} if default_reasoning is not None else {}),
}
fill_missing: Final = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(candidate_params)
return ResponsesWebSocketRequestDefaults(
fill_missing=MappingProxyType(dict(fill_missing)),
overrides=MappingProxyType(_JSON_OBJECT_ADAPTER.validate_python(kwargs.get("extra_body")) or {}),
)
@client
async def _aresponses_websocket(
model: str,
@ -2352,5 +2378,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

@ -54,6 +54,7 @@ if TYPE_CHECKING:
PresidioGuardrailCallback,
ResponsesBackendWebSocket,
ResponsesClientWebSocket,
ResponsesWebSocketRequestDefaults,
)
from litellm.types.router import LiteLLM_Params
@ -1717,6 +1718,7 @@ class ResponsesWebSocketStreaming:
output_guardrail_callbacks: list[PresidioGuardrailCallback] | None = None,
quota_callbacks: Sequence[ProjectQuotaCallback] | None = None,
authorized_model: str | None = None,
request_defaults: ResponsesWebSocketRequestDefaults | None = None,
):
self.websocket = websocket
self.backend_ws = backend_ws
@ -1732,6 +1734,7 @@ class ResponsesWebSocketStreaming:
# Model name authorized at connection time; enforced on every
# response.create frame to prevent deployment-substitution attacks.
self.authorized_model: str | None = authorized_model
self.request_defaults: ResponsesWebSocketRequestDefaults | None = request_defaults
def _should_store_event(self, event_obj: _MutableJsonObject) -> bool:
return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES
@ -1874,12 +1877,23 @@ class ResponsesWebSocketStreaming:
modified = True
return modified
def _with_request_defaults(self, msg_obj: dict[str, object]) -> dict[str, object]:
if self.request_defaults is None:
return msg_obj
nested: Final = msg_obj.get("response")
if _is_json_object(nested):
return {**msg_obj, "response": self.request_defaults.merged_into(nested)}
return {**self.request_defaults.merged_into(msg_obj), "type": msg_obj["type"]}
async def _mask_response_create(self, message: str) -> str:
"""
Enforce the authorized model and apply Presidio PII masking to a
``response.create`` message before it is forwarded to the upstream
provider.
Merge deployment defaults, enforce the authorized model, and apply
Presidio PII masking to a ``response.create`` message before it is
forwarded to the upstream provider.
- Fills the deployment's ``litellm_params`` request defaults into the
frame the way the HTTP ``/v1/responses`` path does: client-set keys
win, ``extra_body`` entries override.
- Overwrites any ``model`` field with the connection-authorized model
to prevent deployment-substitution attacks (always applied).
- Walks the ``input`` and ``instructions`` fields, calls ``check_pii``
@ -1889,23 +1903,26 @@ class ResponsesWebSocketStreaming:
Non-``response.create`` messages are returned unchanged.
"""
try:
msg_obj: Final = _load_json_object(message)
parsed: Final = _load_json_object(message)
except (json.JSONDecodeError, TypeError):
return message
if msg_obj.get("type") != "response.create":
if parsed.get("type") != "response.create":
return message
msg_obj: Final = self._with_request_defaults(parsed)
defaults_applied: Final = msg_obj != parsed
# Always enforce the authorized model, even when PII masking is off.
model_modified: Final = 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 model_modified or defaults_applied else message
if "metadata" not in self.request_data:
self.request_data["metadata"] = {}
modified = model_modified
modified = model_modified or defaults_applied
guardrail_cbs: Final[tuple[PresidioGuardrailCallback, ...]] = tuple(self.guardrail_callbacks)
for cb in guardrail_cbs:
presidio_config = cb.get_presidio_settings_from_request_data(self.request_data)

View file

@ -1,5 +1,7 @@
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Protocol
from litellm.types.guardrails import PresidioPerRequestConfig
@ -39,3 +41,14 @@ class PresidioGuardrailCallback(Protocol):
presidio_config: PresidioPerRequestConfig | None,
request_data: dict[str, object],
) -> str: ...
@dataclass(frozen=True, slots=True)
class ResponsesWebSocketRequestDefaults:
"""Deployment-level request parameters merged into every ``response.create`` frame relayed over a native websocket."""
fill_missing: Mapping[str, object]
overrides: Mapping[str, object]
def merged_into(self, request: Mapping[str, object]) -> dict[str, object]:
return {**self.fill_missing, **request, **self.overrides}

View file

@ -1257,6 +1257,252 @@ class TestWebSocketProjectQuotaEnforcement:
quota_callback.enforce_project_io_token_quota_for_frame.assert_awaited_once()
def _deployment_defaults():
from types import MappingProxyType
from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults
return ResponsesWebSocketRequestDefaults(
fill_missing=MappingProxyType({"reasoning": {"effort": "high"}, "service_tier": "priority"}),
overrides=MappingProxyType({"provider_default": "configured"}),
)
class TestNativeWebSocketDeploymentDefaults:
"""The native relay merges deployment litellm_params into every response.create like HTTP does."""
def test_builder_maps_router_kwargs_like_the_http_path(self):
from litellm.responses.main import _build_responses_websocket_request_defaults
defaults = _build_responses_websocket_request_defaults(
{
"model": "gpt-5-pro",
"reasoning_effort": "high",
"service_tier": "priority",
"extra_body": {"provider_default": "configured"},
"temperature": None,
"timeout": 600,
"max_retries": 2,
"caching": False,
"custom_llm_provider": "openai",
"litellm_metadata": {"user_api_key": "hashed"},
"user_api_key_dict": MagicMock(),
"litellm_logging_obj": MagicMock(),
"websocket": MagicMock(),
}
)
assert dict(defaults.fill_missing) == {"reasoning": {"effort": "high"}, "service_tier": "priority"}
assert dict(defaults.overrides) == {"provider_default": "configured"}
def test_builder_keeps_explicit_reasoning_over_reasoning_effort(self):
from litellm.responses.main import _build_responses_websocket_request_defaults
defaults = _build_responses_websocket_request_defaults(
{"model": "gpt-5-pro", "reasoning": {"effort": "low"}, "reasoning_effort": "high"}
)
assert dict(defaults.fill_missing) == {"reasoning": {"effort": "low"}}
assert dict(defaults.overrides) == {}
def test_builder_copies_dict_valued_reasoning_effort_like_the_http_path(self):
from litellm.responses.main import _build_responses_websocket_request_defaults
defaults = _build_responses_websocket_request_defaults(
{"model": "gpt-5-pro", "reasoning_effort": {"effort": "xhigh", "summary": "auto"}}
)
assert dict(defaults.fill_missing) == {"reasoning": {"effort": "xhigh", "summary": "auto"}}
@pytest.mark.asyncio
async def test_extra_body_type_key_never_replaces_the_frame_type(self):
from types import MappingProxyType
from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults
handler = _make_streaming(
authorized_model="gpt-5-pro",
request_defaults=ResponsesWebSocketRequestDefaults(
fill_missing=MappingProxyType({}),
overrides=MappingProxyType({"type": "session.update", "provider_default": "configured"}),
),
)
forwarded = json.loads(
await handler._mask_response_create(
json.dumps({"type": "response.create", "model": "gpt-5-pro", "input": "hi"})
)
)
assert forwarded == {
"type": "response.create",
"model": "gpt-5-pro",
"input": "hi",
"provider_default": "configured",
}
@pytest.mark.asyncio
async def test_flat_frame_gets_defaults_client_keys_win_extra_body_overrides(self):
handler = _make_streaming(authorized_model="gpt-5-pro", request_defaults=_deployment_defaults())
forwarded = json.loads(
await handler._mask_response_create(
json.dumps(
{
"type": "response.create",
"model": "gpt-5-pro",
"input": "Say hello",
"service_tier": "default",
"provider_default": "client",
}
)
)
)
assert forwarded == {
"type": "response.create",
"model": "gpt-5-pro",
"input": "Say hello",
"service_tier": "default",
"provider_default": "configured",
"reasoning": {"effort": "high"},
}
@pytest.mark.asyncio
async def test_nested_response_frame_gets_defaults_inside_response(self):
handler = _make_streaming(authorized_model="gpt-5-pro", request_defaults=_deployment_defaults())
forwarded = json.loads(
await handler._mask_response_create(
json.dumps({"type": "response.create", "response": {"model": "gpt-5-pro", "input": "hi"}})
)
)
assert forwarded == {
"type": "response.create",
"response": {
"model": "gpt-5-pro",
"input": "hi",
"reasoning": {"effort": "high"},
"service_tier": "priority",
"provider_default": "configured",
},
}
@pytest.mark.asyncio
async def test_frames_that_need_nothing_pass_through_untouched(self):
handler = _make_streaming(authorized_model="gpt-5-pro", request_defaults=_deployment_defaults())
cancel_frame = json.dumps({"type": "response.cancel"})
complete_frame = json.dumps(
{
"type": "response.create",
"model": "gpt-5-pro",
"input": "hi",
"reasoning": {"effort": "high"},
"service_tier": "priority",
"provider_default": "configured",
}
)
assert await handler._mask_response_create(cancel_frame) is cancel_frame
assert await handler._mask_response_create(complete_frame) is complete_frame
@pytest.mark.asyncio
async def test_handler_applies_defaults_to_the_first_frame_sent_upstream(self):
import asyncio
from unittest.mock import AsyncMock, patch
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
class FakeBackend:
def __init__(self):
self.sent = []
async def send(self, message):
self.sent.append(message)
async def recv(self, decode=False):
raise RuntimeError("backend closed")
async def close(self):
pass
backend = FakeBackend()
class FakeConnect:
def __init__(self, url, **kwargs):
pass
async def __aenter__(self):
return backend
async def __aexit__(self, *args):
pass
mock_config = MagicMock(spec=OpenAIResponsesAPIConfig)
mock_config.supports_native_websocket.return_value = True
mock_config.model_in_websocket_url.return_value = True
mock_config.get_websocket_url.return_value = "wss://api.openai.com/v1/responses"
mock_config.validate_environment.return_value = {}
mock_logging = MagicMock()
mock_logging.pre_call = MagicMock()
mock_logging.dispatch_success_handlers = AsyncMock()
client_ws = MagicMock()
client_ws.receive_text = AsyncMock(side_effect=RuntimeError("client closed"))
client_ws.send_text = AsyncMock()
client_ws.close = AsyncMock()
with patch("websockets.connect", FakeConnect):
await BaseLLMHTTPHandler().async_responses_websocket(
model="gpt-5-pro",
websocket=client_ws,
logging_obj=mock_logging,
responses_api_provider_config=mock_config,
api_key="sk-test",
first_message=json.dumps({"type": "response.create", "model": "gpt-5-pro", "input": "Say hello"}),
request_defaults=_deployment_defaults(),
)
await asyncio.sleep(0)
assert [json.loads(frame) for frame in backend.sent] == [
{
"type": "response.create",
"model": "gpt-5-pro",
"input": "Say hello",
"reasoning": {"effort": "high"},
"service_tier": "priority",
"provider_default": "configured",
}
]
@pytest.mark.asyncio
async def test_aresponses_websocket_builds_defaults_from_deployment_kwargs(self, monkeypatch):
import importlib
from unittest.mock import AsyncMock
responses_main = importlib.import_module("litellm.responses.main")
stub = MagicMock()
stub.async_responses_websocket = AsyncMock()
monkeypatch.setattr(responses_main, "base_llm_http_handler", stub)
await responses_main._aresponses_websocket.__wrapped__(
model="openai/gpt-5-pro",
websocket=MagicMock(),
api_key="sk-test",
litellm_logging_obj=MagicMock(),
reasoning_effort="high",
service_tier="priority",
extra_body={"provider_default": "configured"},
)
request_defaults = stub.async_responses_websocket.call_args.kwargs["request_defaults"]
assert dict(request_defaults.fill_missing) == {"reasoning": {"effort": "high"}, "service_tier": "priority"}
assert dict(request_defaults.overrides) == {"provider_default": "configured"}
class TestNativeWebSocketGuardrails:
@pytest.mark.asyncio
async def test_response_create_injects_authorized_model(self):