fix(router): initialize shared Fusion replay stream state

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Moe Khalil 2026-09-22 03:29:18 +00:00
parent ba3536520f
commit f38bbad108
2 changed files with 18 additions and 24 deletions

View file

@ -829,26 +829,22 @@ class FusionReplayStream(CustomStreamWrapper):
chunks: Sequence[ModelResponseStream],
fusion_metadata: Mapping[str, object],
) -> None:
# Deliberately do not call CustomStreamWrapper.__init__. The source
# wrapper already normalized and logged these chunks while Fusion
# buffered them to determine whether its private tool was invoked.
source_model: Final = getattr(source, "model", "")
self.model = source_model if isinstance(source_model, str) else ""
self.custom_llm_provider = source.custom_llm_provider
self.logging_obj = source.logging_obj
self._hidden_params = dict( # mutable-ok: local provider payload
getattr(
source,
"_hidden_params",
{}, # mutable-ok: stream metadata defaults to a native mapping
) # mutable-ok: local provider payload
) # mutable-ok: local provider payload
self._hidden_params["fusion"] = dict( # mutable-ok: local provider payload
fusion_metadata
) # mutable-ok: local provider payload
self._source = source
super().__init__(
completion_stream=source,
model=source.model,
logging_obj=source.logging_obj,
custom_llm_provider=source.custom_llm_provider,
stream_options=source.stream_options,
)
self._hidden_params = { # mutable-ok: stream consumers attach response metadata
**source._hidden_params,
"fusion": dict(fusion_metadata),
}
self._iterator = iter(chunks)
def __next__(self) -> ModelResponseStream:
return next(self._iterator)
def __aiter__(self) -> FusionReplayStream:
return self
@ -858,10 +854,6 @@ class FusionReplayStream(CustomStreamWrapper):
except StopIteration as exc:
raise StopAsyncIteration from exc
async def aclose(self) -> None:
if hasattr(self._source, "aclose"):
await self._source.aclose()
class FusionRouter:
def __init__(

View file

@ -1118,7 +1118,8 @@ async def test_proxy_fusion_fails_closed_without_authorization_context() -> None
@pytest.mark.asyncio
async def test_router_replays_direct_outer_response_as_an_async_stream() -> None:
@pytest.mark.parametrize("consume_sync", [True, False])
async def test_router_replays_direct_outer_response_without_reprocessing(consume_sync: bool) -> None:
router = Router(model_list=_router_model_list())
response = await router.acompletion(
model="fusion/test",
@ -1127,11 +1128,12 @@ async def test_router_replays_direct_outer_response_as_an_async_stream() -> None
)
assert isinstance(response, CustomStreamWrapper)
chunks = [chunk async for chunk in response]
chunks = list(response) if consume_sync else [chunk async for chunk in response]
rebuilt = litellm.stream_chunk_builder(chunks=chunks)
assert isinstance(rebuilt, ModelResponse)
assert rebuilt.choices[0].message.content == "Final"
assert response._hidden_params["fusion"]["invoked"] is False
await response.aclose()
@pytest.mark.asyncio