diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 687df9b0348..11cafd12c0c 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1250,6 +1250,13 @@ async def open_sse_before_first_byte( if interval is None: return await produce_response + # The slot record must exist before this task is forked, or the release + # looks in the original task and never sees it. + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + get_or_create_request_stash, + ) + + get_or_create_request_stash() produce_task: Final = asyncio.ensure_future(produce_response) await asyncio.wait((produce_task,), timeout=interval) if produce_task.done(): diff --git a/tests/test_litellm/proxy/test_sse_keepalive_request_stash.py b/tests/test_litellm/proxy/test_sse_keepalive_request_stash.py new file mode 100644 index 00000000000..d7c582d1ebe --- /dev/null +++ b/tests/test_litellm/proxy/test_sse_keepalive_request_stash.py @@ -0,0 +1,34 @@ +"""The SSE keepalive task must share the parallel-request slot with its caller.""" + +import asyncio + +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _request_stash, + get_or_create_request_stash, + get_request_stash, +) + + +def test_sse_keepalive_fork_shares_the_parallel_request_slot(): + async def process(): + get_or_create_request_stash().owner_litellm_call_id = "call-1" + + async def without_seed(): + _request_stash.set(None) + task = asyncio.ensure_future(process()) + await asyncio.wait((task,)) + return get_request_stash() + + async def with_seed(): + _request_stash.set(None) + get_or_create_request_stash() + task = asyncio.ensure_future(process()) + await asyncio.wait((task,)) + return get_request_stash() + + leaked = asyncio.run(without_seed()) + shared = asyncio.run(with_seed()) + + assert leaked is None + assert shared is not None + assert shared.owner_litellm_call_id == "call-1"