mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
ba3536520f
commit
f38bbad108
2 changed files with 18 additions and 24 deletions
|
|
@ -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__(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue