mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(responses): apply deployment defaults to websocket frames
This commit is contained in:
parent
9121ae3024
commit
c026dd153c
4 changed files with 96 additions and 3 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue