mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
fix(responses): merge deployment litellm_params into native websocket response.create frames
This commit is contained in:
parent
6fa34a299b
commit
4a951847bb
6 changed files with 275 additions and 9 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -19394,7 +19394,7 @@
|
|||
}
|
||||
}
|
||||
},
|
||||
"description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n "
|
||||
"description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n"
|
||||
},
|
||||
"500": {
|
||||
"content": {
|
||||
|
|
|
|||
|
|
@ -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,27 @@ def _build_litellm_metadata_for_ws(kwargs: dict) -> dict:
|
|||
return metadata
|
||||
|
||||
|
||||
_EXTRA_BODY_ADAPTER: Final = TypeAdapter(dict[str, object] | None)
|
||||
|
||||
|
||||
def _build_responses_websocket_request_defaults(kwargs: Mapping[str, object]) -> ResponsesWebSocketRequestDefaults:
|
||||
reasoning_effort: Final = kwargs.get("reasoning_effort")
|
||||
mapped_reasoning: Final = (
|
||||
LiteLLMResponsesTransformationHandler()._map_reasoning_effort(reasoning_effort)
|
||||
if kwargs.get("reasoning") is None and isinstance(reasoning_effort, str)
|
||||
else None
|
||||
)
|
||||
candidate_params: Final[dict[str, object]] = {
|
||||
**kwargs,
|
||||
**({"reasoning": mapped_reasoning} if mapped_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(_EXTRA_BODY_ADAPTER.validate_python(kwargs.get("extra_body")) or {}),
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
async def _aresponses_websocket(
|
||||
model: str,
|
||||
|
|
@ -2352,5 +2374,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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -1204,6 +1204,216 @@ 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) == {}
|
||||
|
||||
@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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue