diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index 43138b2c526..2c1ce6068b2 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -384,7 +384,7 @@ class BedrockRealtime(BaseAWSLLM): ) bedrock_task: Final = asyncio.create_task(collect_logged_events()) - await asyncio.wait((client_task, bedrock_task), return_when=asyncio.FIRST_EXCEPTION) + await asyncio.wait((client_task, bedrock_task), return_when=asyncio.FIRST_COMPLETED) client_disconnected: Final = ( client_task.done() and not client_task.cancelled() and client_task.exception() is None ) 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 6c7659d84cd..ac3a43b742f 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 @@ -711,6 +711,25 @@ class TestBedrockRealtimeProviderFailurePropagation: assert stream.input_stream.closed + @pytest.mark.asyncio + async def test_client_disconnect_ends_the_session_while_bedrock_output_stays_open(self, stub_aws_sdk_client): + receiver = DrainedThenOpenBedrockReceiver([]) + stream = ScriptedBedrockStream([], receiver_type=lambda _payloads: receiver) + stub_aws_sdk_client["streams"] = [stream] + + await asyncio.wait_for( + BedrockRealtime().async_realtime( + model="amazon.nova-sonic-v1:0", + websocket=RealtimeClientWS(), + logging_obj=FakeLogging(), + **self.AWS_PARAMS, + ), + timeout=1, + ) + + assert receiver.drained.is_set(), "the handler must have been waiting on the open provider stream" + assert stream.input_stream.closed + @pytest.mark.asyncio async def test_session_updated_is_not_sent_before_bedrock_is_ready(self, stub_aws_models): handler = BedrockRealtime()