fix(realtime): surface an upstream handshake refusal as an error event and policy close (#42388)

* fix(realtime): surface an upstream handshake refusal as an error event and policy close

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(realtime): tidy the handshake refusal e2e

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(realtime): keep upstream exception text out of the Azure client error

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(realtime): map handshake refusal close codes with a lookup

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-22 11:31:15 -07:00 • committed by GitHub
parent 8332f0848a
commit 2bab39e374
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 299 additions and 17 deletions

View file

@ -9,13 +9,20 @@ frame itself fail, which is how a loud failure turns back into a silent one.
"""
import json
from typing import Final
from types import MappingProxyType
from typing import Final, Protocol
from litellm.types.realtime import RealtimeErrorDetail, RealtimeErrorEvent
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
class _ClientWebSocket(Protocol):
async def send_text(self, data: str) -> None: ...
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
def realtime_error_event(message: str, error_type: str) -> str:
detail: Final[RealtimeErrorDetail] = {"type": error_type, "message": message}
event: Final[RealtimeErrorEvent] = {"type": "error", "error": detail}
@ -37,3 +44,28 @@ def client_close_code(upstream_code: int) -> int:
if upstream_code in EXTERNAL_CLOSE_CODES or 3000 <= upstream_code < 5000:
return upstream_code
return int(CloseCode.INTERNAL_ERROR)
def upstream_handshake_close_code(status_code: int) -> int:
from websockets.frames import CloseCode
refusal_codes: Final = MappingProxyType(
{
401: int(CloseCode.POLICY_VIOLATION),
403: int(CloseCode.POLICY_VIOLATION),
429: int(CloseCode.TRY_AGAIN_LATER),
}
)
return refusal_codes.get(status_code, int(CloseCode.INTERNAL_ERROR))
async def close_after_upstream_handshake_refusal(websocket: _ClientWebSocket, status_code: int) -> None:
message: Final = f"Upstream realtime handshake rejected with HTTP {status_code}"
try:
await websocket.send_text(realtime_error_event(message, error_type="server_error"))
except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below
pass
await websocket.close(
code=upstream_handshake_close_code(status_code),
reason=websocket_close_reason(message, fallback="Upstream handshake rejected"),
)

View file

@ -8,11 +8,15 @@ from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final, Protocol, cast
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm._logging import verbose_proxy_logger
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.types.realtime import RealtimeQueryParams
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ....litellm_core_utils.realtime_errors import (
close_after_upstream_handshake_refusal,
realtime_error_event,
)
from ....litellm_core_utils.realtime_streaming import (
RealTimeStreaming,
ScopedWebSocket,
@ -49,7 +53,9 @@ def azure_realtime_protocol_for_client(
class _ProxyClientWebSocket(Protocol):
"""Client-facing websocket handle: this path only closes it after a failed handshake."""
"""Client-facing websocket handle: this path only writes to it after a failed handshake."""
async def send_text(self, data: str) -> None: ...
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
@ -181,7 +187,16 @@ class AzureOpenAIRealtime(AzureChatCompletion):
)
await realtime_streaming.bidirectional_forward()
except websockets.exceptions.InvalidStatusCode as e:
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
except websockets.exceptions.InvalidStatus as e:
verbose_proxy_logger.exception("Error in AzureOpenAIRealtime.async_realtime")
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
except Exception:
verbose_proxy_logger.exception("Error in AzureOpenAIRealtime.async_realtime")
try:
await websocket.send_text(realtime_error_event("Internal server error", error_type="server_error"))
except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below
pass
try:
await websocket.close(code=1011, reason="Internal server error")
except Exception: # noqa: BLE001 # the lower layer may have closed the socket already; closing twice is not an error
pass

View file

@ -45,7 +45,11 @@ from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
)
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
from litellm.litellm_core_utils.realtime_errors import realtime_error_event, websocket_close_reason
from litellm.litellm_core_utils.realtime_errors import (
close_after_upstream_handshake_refusal,
realtime_error_event,
websocket_close_reason,
)
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.anthropic_messages.transformation import (
@ -6384,9 +6388,9 @@ class BaseLLMHTTPHandler:
await realtime_streaming.bidirectional_forward()
except websockets.exceptions.InvalidStatusCode as e:
except websockets.exceptions.InvalidStatus as e:
verbose_logger.exception("Error connecting to backend: %s", e)
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
except Exception as e:
verbose_logger.exception("Error connecting to backend: %s", e)
redacted_error: Final = _redact_string(str(e))
@ -6799,9 +6803,9 @@ class BaseLLMHTTPHandler:
)
return await streaming.bidirectional_forward()
except websockets.exceptions.InvalidStatusCode as e:
except websockets.exceptions.InvalidStatus as e:
verbose_logger.exception("Error connecting to responses WS backend: %s", e)
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
except Exception as e:
verbose_logger.exception("Error in responses WS: %s", e)
try:

View file

@ -12,6 +12,7 @@ from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.types.realtime import RealtimeQueryParams
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ....litellm_core_utils.realtime_errors import close_after_upstream_handshake_refusal
from ....litellm_core_utils.realtime_streaming import (
RealtimeEventNormalizer,
RealTimeStreaming,
@ -175,8 +176,8 @@ class OpenAIRealtime(OpenAIChatCompletion):
)
await realtime_streaming.bidirectional_forward()
except websockets.exceptions.InvalidStatusCode as e:
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
except websockets.exceptions.InvalidStatus as e:
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
except Exception as e:
try:
await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e}"))

View file

@ -310,6 +310,7 @@ from litellm.litellm_core_utils.core_helpers import (
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.realtime_errors import (
close_after_upstream_handshake_refusal,
realtime_error_event,
websocket_close_reason,
)
@ -12531,9 +12532,9 @@ async def realtime_websocket_endpoint(
user_model=user_model,
)
await llm_call
except websockets.exceptions.InvalidStatusCode as e:
except websockets.exceptions.InvalidStatus as e:
verbose_proxy_logger.exception("Invalid status code")
await websocket.close(code=e.status_code, reason="Invalid status code")
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
except Exception as e:
verbose_proxy_logger.exception("Internal server error")
redacted_error: Final = _redact_string(str(e))

View file

@ -286,7 +286,7 @@ def function_call_item(events: tuple[ReceivedEvent, ...]) -> OutputItem | None:
# ---- session + client --------------------------------------------------
def _as_text(message: str | bytes) -> str:
def as_text(message: str | bytes) -> str:
return message.decode("utf-8") if isinstance(message, bytes) else message
@ -304,7 +304,7 @@ class RealtimeSession:
collected: list[ReceivedEvent] = []
while time.monotonic() < deadline:
try:
text = _as_text(
text = as_text(
self.connection.recv(timeout=deadline - time.monotonic())
)
except TimeoutError:

View file

@ -13,8 +13,9 @@ failure. See REALTIME_COVERAGE_MATRIX.md.
"""
import pytest
from lifecycle import ResourceManager
from models import LiteLLMParamsBody
from pydantic import BaseModel
from realtime_client import (
PROVIDERS,
ConversationItemCreate,
@ -27,14 +28,17 @@ from realtime_client import (
RealtimeProvider,
ResponseCreate,
ResponseDone,
ServerEnvelope,
SessionConfig,
SessionUpdate,
as_text,
function_call_item,
parse_last,
realtime_model,
transcript,
user_message,
)
from websockets.exceptions import ConnectionClosedError
pytestmark = pytest.mark.e2e
@ -147,3 +151,39 @@ def test_tool_call_round_trip(
second = session.collect_until("response.done", timeout=60)
assert "72" in transcript(second), "follow-up did not use the tool result"
_REFUSED_UPSTREAMS = (
RealtimeProvider(
"azure-bad-key",
"azure-realtime-refused",
LiteLLMParamsBody(
model="azure/gpt-realtime",
api_key="invalid-e2e-key",
api_version="2025-08-28",
realtime_protocol="GA",
),
),
)
@pytest.mark.parametrize("provider", _REFUSED_UPSTREAMS, ids=[p.id for p in _REFUSED_UPSTREAMS])
def test_upstream_handshake_refusal_is_an_error_event_and_policy_close(
client: RealtimeClient,
resources: ResourceManager,
scoped_key: str,
provider: RealtimeProvider,
) -> None:
model_name, model_id = client.provision(provider)
resources.defer(lambda: client.proxy.delete_model(model_id))
with client.connect(key=scoped_key, model=model_name) as session:
first = ServerEnvelope.model_validate_json(
as_text(session.connection.recv(timeout=15))
)
assert first.type == "error", first
with pytest.raises(ConnectionClosedError) as closed:
session.connection.recv(timeout=15)
assert closed.value.rcvd is not None
assert closed.value.rcvd.code == 1008, closed.value

View file

@ -1,13 +1,17 @@
import json
from typing import cast
import pytest
from litellm.litellm_core_utils.realtime_errors import (
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
client_close_code,
close_after_upstream_handshake_refusal,
realtime_error_event,
upstream_handshake_close_code,
websocket_close_reason,
)
from litellm.types.realtime import RealtimeErrorEvent
def test_realtime_error_event_shape():
@ -52,3 +56,52 @@ def test_websocket_close_reason_truncates_multibyte_message_by_bytes():
)
def test_client_close_code_only_forwards_codes_a_server_may_send(upstream_code, expected):
assert client_close_code(upstream_code) == expected
@pytest.mark.parametrize(
("status_code", "expected"),
[(401, 1008), (403, 1008), (429, 1013), (500, 1011)],
)
def test_upstream_handshake_close_code_maps_http_status_to_close_code(status_code: int, expected: int):
assert upstream_handshake_close_code(status_code) == expected
class _RecordingWebSocket:
def __init__(self) -> None:
self.sent: list[str] = []
self.closed: tuple[int, str | None] | None = None
async def send_text(self, data: str) -> None:
self.sent.append(data)
async def close(self, code: int = 1000, reason: str | None = None) -> None:
self.closed = (code, reason)
@pytest.mark.asyncio
async def test_close_after_upstream_handshake_refusal_sends_error_event_then_policy_close():
websocket = _RecordingWebSocket()
await close_after_upstream_handshake_refusal(websocket, 401)
assert len(websocket.sent) == 1
event = cast(RealtimeErrorEvent, json.loads(websocket.sent[0]))
assert event["type"] == "error"
assert event["error"]["type"] == "server_error"
assert "401" in event["error"]["message"]
assert websocket.closed is not None
assert websocket.closed[0] == 1008
assert websocket.closed[1]
@pytest.mark.asyncio
async def test_close_after_upstream_handshake_refusal_still_closes_when_send_fails():
class _DeadWebSocket(_RecordingWebSocket):
async def send_text(self, data: str) -> None:
raise RuntimeError("socket gone")
websocket = _DeadWebSocket()
await close_after_upstream_handshake_refusal(websocket, 500)
assert websocket.closed == (1011, "Upstream realtime handshake rejected with HTTP 500")

View file

@ -0,0 +1 @@

View file

@ -0,0 +1,86 @@
import json
from typing import cast
from unittest.mock import MagicMock, patch
import pytest
class _RecordingClientWebSocket:
scope: dict[str, list[tuple[bytes, bytes]]] = {"headers": []}
def __init__(self) -> None:
self.sent: list[str] = []
self.closed: list[tuple[int, str | None]] = []
async def send_text(self, data: str) -> None:
self.sent.append(data)
async def close(self, code: int = 1000, reason: str | None = None) -> None:
self.closed.append((code, reason))
@pytest.mark.asyncio
async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_policy_close():
import websockets
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
from litellm.types.realtime import RealtimeErrorEvent
handler = AzureOpenAIRealtime()
model = "gpt-realtime"
dummy_websocket = _RecordingClientWebSocket()
dummy_logging_obj = MagicMock()
refused = websockets.exceptions.InvalidStatus(
websockets.http11.Response(401, "Unauthorized", websockets.datastructures.Headers())
)
with patch("websockets.connect", side_effect=refused):
await handler.async_realtime( # pyright: ignore[reportUnknownMemberType] # handler's websocket param is a Protocol here but the mock connect type is incomplete
model=model,
websocket=dummy_websocket,
logging_obj=dummy_logging_obj,
api_base="https://example.openai.azure.com",
api_key="bad-key",
api_version="2025-08-28",
query_params={"model": model},
)
assert len(dummy_websocket.sent) == 1
event = cast(RealtimeErrorEvent, json.loads(dummy_websocket.sent[0]))
assert event["type"] == "error"
assert event["error"]["type"] == "server_error"
assert "401" in event["error"]["message"]
assert dummy_websocket.closed and dummy_websocket.closed[0][0] == 1008
@pytest.mark.asyncio
async def test_async_realtime_unexpected_error_sends_error_event_then_internal_close():
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
from litellm.types.realtime import RealtimeErrorEvent
handler = AzureOpenAIRealtime()
model = "gpt-realtime"
dummy_websocket = _RecordingClientWebSocket()
dummy_logging_obj = MagicMock()
with patch("websockets.connect", side_effect=OSError("connection reset")):
await handler.async_realtime( # pyright: ignore[reportUnknownMemberType] # same as above
model=model,
websocket=dummy_websocket,
logging_obj=dummy_logging_obj,
api_base="https://example.openai.azure.com",
api_key="bad-key",
api_version="2025-08-28",
query_params={"model": model},
)
assert len(dummy_websocket.sent) == 1
event = cast(RealtimeErrorEvent, json.loads(dummy_websocket.sent[0]))
assert event["type"] == "error"
assert event["error"]["type"] == "server_error"
assert event["error"]["message"] == "Internal server error"
assert "connection reset" not in dummy_websocket.sent[0]
assert dummy_websocket.closed and dummy_websocket.closed[0] == (1011, "Internal server error")

View file

@ -416,3 +416,52 @@ async def test_async_realtime_ws_url_has_no_ssl():
# Verify ssl is None for ws:// URLs (the fix for issue #19222)
assert called_kwargs["ssl"] is None
@pytest.mark.asyncio
async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_policy_close():
from typing import cast
import websockets
from litellm.llms.openai.realtime.handler import OpenAIRealtime
from litellm.types.realtime import RealtimeErrorEvent
handler = OpenAIRealtime()
model = "gpt-realtime"
sent: list[str] = []
closed: list[tuple[int, str | None]] = []
class RecordingClientWebSocket:
scope: dict[str, list[tuple[bytes, bytes]]] = {"headers": []}
async def send_text(self, data: str) -> None:
sent.append(data)
async def close(self, code: int = 1000, reason: str | None = None) -> None:
closed.append((code, reason))
dummy_websocket = RecordingClientWebSocket()
dummy_logging_obj = MagicMock()
refused = websockets.exceptions.InvalidStatus(
websockets.http11.Response(401, "Unauthorized", websockets.datastructures.Headers())
)
with patch("websockets.connect", side_effect=refused):
await handler.async_realtime( # pyright: ignore[reportUnknownMemberType] # handler's websocket param is Any
model=model,
websocket=dummy_websocket,
logging_obj=dummy_logging_obj,
api_base="https://api.openai.com/",
api_key="bad-key",
query_params={"model": model},
)
assert len(sent) == 1
event = cast(RealtimeErrorEvent, json.loads(sent[0]))
assert event["type"] == "error"
assert event["error"]["type"] == "server_error"
assert "401" in event["error"]["message"]
assert closed and closed[0][0] == 1008