diff --git a/litellm/constants.py b/litellm/constants.py index 5751e6e46af..7d86080086a 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -311,6 +311,10 @@ REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float( # RFC 6455 caps the close frame payload at 125 bytes, 2 of which carry the status code WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123 +# ASGI websocket scope keys the Bedrock realtime bridge uses to carry state across router fallback attempts +BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY: Final = "litellm.bedrock_realtime.pending_session_update" +BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY: Final = "litellm.bedrock_realtime.session_committed" + # SSL/TLS cipher configuration for faster handshakes # Strategy: Strongly prefer fast modern ciphers, but allow fallback to commonly supported ones # This balances performance with broad compatibility diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index ca5b87d1700..126a7ae5c3e 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -16,6 +16,10 @@ from pydantic import JsonValue, TypeAdapter import litellm from litellm._logging import _redact_string, verbose_proxy_logger +from litellm.constants import ( + BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY, + BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY, +) from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER @@ -31,8 +35,6 @@ _CLIENT_MODALITIES_ADAPTER: Final[TypeAdapter["list[str] | None"]] = TypeAdapter _CLIENT_MESSAGE_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) _EMPTY_JSON_OBJECT: Final[Mapping[str, JsonValue]] = MappingProxyType({}) -_PENDING_UPDATE_KEY: Final = "litellm.bedrock_realtime.pending_session_update" -_COMMITTED_KEY: Final = "litellm.bedrock_realtime.session_committed" _BEDROCK_STREAM_ERROR_STATUS: Final[Mapping[str, int]] = MappingProxyType( { @@ -145,16 +147,14 @@ async def _client_messages(client_ws: RealtimeClientWebSocket, initial_message: return -def _take_pending_session_update( - scope: MutableMapping[str, object], # mutable-ok: the ASGI scope is the per-connection state store -) -> str | None: +def _pending_session_update(scope: Mapping[str, object]) -> str | None: """A fallback attempt on the same websocket replays the session.update the failed attempt never acked.""" - if scope.get(_COMMITTED_KEY) is True: + if scope.get(BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY) is True: raise BedrockError( status_code=409, message="Bedrock realtime session already committed to a provider stream; it cannot be replayed", ) - pending: Final = scope.pop(_PENDING_UPDATE_KEY, None) # rebind-ok: the ASGI scope outlives this attempt + pending: Final = scope.get(BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY) return pending if isinstance(pending, str) else None @@ -175,8 +175,8 @@ async def _ack_session_update( ) -> bool: """Ack the client's session.update once Bedrock accepted the stream; False means the client is gone.""" await bedrock_stream.await_output() - client_ws.scope.pop(_PENDING_UPDATE_KEY, None) - client_ws.scope[_COMMITTED_KEY] = True # rebind-ok: the ASGI scope outlives this attempt + client_ws.scope.pop(BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY, None) + client_ws.scope[BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY] = True # rebind-ok: scope outlives the attempt if logging_obj is None: return True requested_modalities: Final = _CLIENT_MODALITIES_ADAPTER.validate_python( @@ -239,7 +239,7 @@ class BedrockRealtime(BaseAWSLLM): except ImportError: raise ImportError("Missing aws_sdk_bedrock_runtime. Install with: pip install aws-sdk-bedrock-runtime") - pending_session_update: Final = _take_pending_session_update(websocket.scope) + pending_session_update: Final = _pending_session_update(websocket.scope) # Get AWS region if aws_region_name is None: @@ -445,7 +445,9 @@ class BedrockRealtime(BaseAWSLLM): parsed_client_message = _parse_client_message(message) is_session_update = _json_str(parsed_client_message.get("type")) == "session.update" if is_session_update: - client_ws.scope[_PENDING_UPDATE_KEY] = message # rebind-ok: scope outlives the attempt + client_ws.scope[BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY] = ( + message # rebind-ok: scope outlives the attempt + ) transformed_messages = transformation_config.transform_realtime_request( message=message, diff --git a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py index 0439a5090f5..085ffeb5129 100644 --- a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py +++ b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py @@ -217,7 +217,10 @@ def stub_aws_sdk_client(monkeypatch): async def invoke_model_with_bidirectional_stream(self, operation_input): captured["operation_input"] = operation_input if captured.get("streams"): - return captured["streams"].pop(0) + stream = captured["streams"].pop(0) + if isinstance(stream, Exception): + raise stream + return stream return ScriptedBedrockStream(captured.get("scripted_payloads", [])) package = types.ModuleType("aws_sdk_bedrock_runtime") @@ -531,7 +534,8 @@ class TestBedrockRealtimeProviderFailurePropagation: handler = BedrockRealtime() websocket = ConnectedClientWS([self.SESSION_UPDATE]) healthy_stream = ScriptedBedrockStream([]) - stub_aws_sdk_client["streams"] = [UnavailableBedrockStream(), healthy_stream] + eager_failure = ServiceUnavailableException("fault injected before the stream was returned") + stub_aws_sdk_client["streams"] = [UnavailableBedrockStream(), eager_failure, healthy_stream] with pytest.raises(BedrockError) as failure: await handler.async_realtime( @@ -542,6 +546,11 @@ class TestBedrockRealtimeProviderFailurePropagation: assert [json.loads(m)["type"] for m in websocket.sent_to_client] == ["session.created"] assert not websocket.closed, "the proxy route owns the client-facing error event and 1011 close" + with pytest.raises(ServiceUnavailableException): + await handler.async_realtime( + model="amazon.nova-sonic-v1:0", websocket=websocket, logging_obj=FakeLogging(), **self.AWS_PARAMS + ) + await handler.async_realtime( model="amazon.nova-sonic-v1:0", websocket=websocket, logging_obj=FakeLogging(), **self.AWS_PARAMS )