mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(router): resume sync streaming fallbacks without retrying primary (#43959)
* fix(router): resume sync streaming fallback chain Signed-off-by: amarrtech <272048731+amarrtech@users.noreply.github.com> * test(router): cover the sync mid-stream fallback walking every configured target * test(router): cover the sync stream fallback walk against a wire upstream --------- Signed-off-by: amarrtech <272048731+amarrtech@users.noreply.github.com> Co-authored-by: amarrtech <272048731+amarrtech@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
df23f11e9b
commit
88f15e7572
3 changed files with 415 additions and 10 deletions
|
|
@ -3713,11 +3713,17 @@ class Router:
|
|||
initial_kwargs["original_function"] = router_self._completion
|
||||
initial_kwargs["messages"] = messages
|
||||
router_self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs)
|
||||
fallback_response = router_self.function_with_fallbacks(
|
||||
**initial_kwargs,
|
||||
fallback_response = run_async_function(
|
||||
router_self.async_function_with_fallbacks_common_utils,
|
||||
e=e,
|
||||
disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs),
|
||||
fallbacks=fallbacks,
|
||||
context_window_fallbacks=context_window_fallbacks,
|
||||
content_policy_fallbacks=content_policy_fallbacks,
|
||||
model_group=model_group,
|
||||
args=(),
|
||||
kwargs=initial_kwargs,
|
||||
include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True,
|
||||
)
|
||||
|
||||
if hasattr(fallback_response, "__iter__"):
|
||||
|
|
|
|||
271
tests/integration/sdk/test_router_sync_stream_fallback_wire.py
Normal file
271
tests/integration/sdk/test_router_sync_stream_fallback_wire.py
Normal file
|
|
@ -0,0 +1,271 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Callable, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from litellm import Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
_MODEL: Final = "gpt-5.6"
|
||||
_API_KEY: Final = "synthetic-sync-fallback-key"
|
||||
_PROMPT: Final = "which deployment answers when the primary dies before its first chunk?"
|
||||
_ERROR_FRAME: Final = (
|
||||
b"data: " + json.dumps({"error": {"message": "overloaded", "type": "server_error", "code": 500}}).encode() + b"\n\n"
|
||||
)
|
||||
_DONE: Final = b"data: [DONE]\n\n"
|
||||
_BURST: Final = 6
|
||||
|
||||
|
||||
def _delta(text: str, finish_reason: str | None) -> bytes:
|
||||
chunk: Final = {
|
||||
"id": "chatcmpl-sync-fallback-wire",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1700000000,
|
||||
"model": _MODEL,
|
||||
"choices": [{"index": 0, "delta": {"role": "assistant", "content": text}, "finish_reason": finish_reason}],
|
||||
}
|
||||
return b"data: " + json.dumps(chunk).encode() + b"\n\n"
|
||||
|
||||
|
||||
def _serves(text: str) -> Reply:
|
||||
return Reply(content_type="text/event-stream", chunks=(_delta(text, None), _delta("", "stop"), _DONE))
|
||||
|
||||
|
||||
def _dies_after(text: str) -> Reply:
|
||||
return Reply(content_type="text/event-stream", chunks=(_delta(text, None), _ERROR_FRAME, _DONE))
|
||||
|
||||
|
||||
_DIES_BEFORE_CONTENT: Final = Reply(content_type="text/event-stream", chunks=(_ERROR_FRAME, _DONE))
|
||||
_DROPS_BEFORE_CONTENT: Final = Reply(content_type="text/event-stream", chunks=(_DONE,), abort_after=0)
|
||||
|
||||
|
||||
def _peer(replies: Mapping[str, Reply]) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST", request.method
|
||||
deployment, _, route = request.target.lstrip("/").partition("/")
|
||||
assert route == "chat/completions", request.target
|
||||
return replies[deployment]
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _deployments_hit(wire: Wire) -> tuple[str, ...]:
|
||||
return tuple(request.target.lstrip("/").partition("/")[0] for request in wire.drain())
|
||||
|
||||
|
||||
def _router(wire: Wire, deployments: tuple[str, ...], **settings: object) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": name,
|
||||
"litellm_params": {"model": f"openai/{_MODEL}", "api_base": f"{wire.url}/{name}", "api_key": _API_KEY},
|
||||
}
|
||||
for name in deployments
|
||||
],
|
||||
num_retries=0,
|
||||
disable_cooldowns=True,
|
||||
**settings,
|
||||
)
|
||||
|
||||
|
||||
class _FallbackRecorder(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.successes: tuple[str, ...] = ()
|
||||
self.failures: tuple[str, ...] = ()
|
||||
|
||||
async def log_success_fallback_event(
|
||||
self, original_model_group: str, kwargs: dict, original_exception: Exception
|
||||
) -> None:
|
||||
self.successes = (*self.successes, original_model_group)
|
||||
|
||||
async def log_failure_fallback_event(
|
||||
self, original_model_group: str, kwargs: dict, original_exception: Exception
|
||||
) -> None:
|
||||
self.failures = (*self.failures, original_model_group)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Streamed:
|
||||
text: str
|
||||
attempted_fallbacks: object
|
||||
|
||||
|
||||
def _text_of(chunk: object) -> str:
|
||||
choices: Final = getattr(chunk, "choices", None) or ()
|
||||
return "".join(str(choice.delta.content or "") for choice in choices)
|
||||
|
||||
|
||||
def _attempted_fallbacks(stream: object) -> object:
|
||||
hidden: Final = getattr(stream, "_hidden_params", None) or {}
|
||||
return (hidden.get("additional_headers") or {}).get("x-litellm-attempted-fallbacks")
|
||||
|
||||
|
||||
def _stream_sync(router: Router, **request: object) -> _Streamed:
|
||||
stream: Final = router.completion(model="primary", messages=[{"role": "user", "content": _PROMPT}], stream=True, **request)
|
||||
text: Final = "".join(_text_of(chunk) for chunk in stream)
|
||||
return _Streamed(text=text, attempted_fallbacks=_attempted_fallbacks(stream))
|
||||
|
||||
|
||||
async def _stream_async(router: Router, **request: object) -> _Streamed:
|
||||
stream: Final = await router.acompletion(
|
||||
model="primary", messages=[{"role": "user", "content": _PROMPT}], stream=True, **request
|
||||
)
|
||||
parts: Final = [_text_of(chunk) async for chunk in stream]
|
||||
return _Streamed(text="".join(parts), attempted_fallbacks=_attempted_fallbacks(stream))
|
||||
|
||||
|
||||
def _stream(client: str, router: Router, **request: object) -> _Streamed:
|
||||
if client == "async":
|
||||
return asyncio.run(_stream_async(router, **request))
|
||||
return _stream_sync(router, **request)
|
||||
|
||||
|
||||
_CLIENTS: Final = ("sync", "async")
|
||||
_PRIMARY_DIES: Final = {"primary": _DIES_BEFORE_CONTENT, "backup": _serves("answered by the backup")}
|
||||
_PRIMARY_AND_FB1_DIE: Final = {"primary": _DIES_BEFORE_CONTENT, "fb1": _DIES_BEFORE_CONTENT, "fb2": _serves("answered by fb2")}
|
||||
_PRIMARY_TO_BACKUP: Final = [{"primary": ["backup"]}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_primary_dies_before_content(client: str, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
recorder: Final = _FallbackRecorder()
|
||||
monkeypatch.setattr(litellm, "callbacks", [recorder])
|
||||
with wire_server(_peer(_PRIMARY_DIES)) as wire:
|
||||
router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP)
|
||||
streamed: Final = _stream(client, router)
|
||||
assert streamed.text == "answered by the backup", streamed
|
||||
assert streamed.attempted_fallbacks == 1, streamed
|
||||
assert _deployments_hit(wire) == ("primary", "backup")
|
||||
assert recorder.successes == ("primary",), recorder.successes
|
||||
assert recorder.failures == (), recorder.failures
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_walks_every_configured_fallback(client: str) -> None:
|
||||
with wire_server(_peer(_PRIMARY_AND_FB1_DIE)) as wire:
|
||||
router: Final = _router(wire, ("primary", "fb1", "fb2"), fallbacks=[{"primary": ["fb1", "fb2"]}])
|
||||
streamed: Final = _stream(client, router)
|
||||
assert streamed.text == "answered by fb2", streamed
|
||||
assert streamed.attempted_fallbacks == 2, streamed
|
||||
assert _deployments_hit(wire) == ("primary", "fb1", "fb2")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_every_target_dies(client: str) -> None:
|
||||
with wire_server(_peer({"primary": _DIES_BEFORE_CONTENT, "backup": _DIES_BEFORE_CONTENT})) as wire:
|
||||
router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP)
|
||||
with pytest.raises(litellm.APIConnectionError, match="overloaded"):
|
||||
_stream(client, router)
|
||||
assert _deployments_hit(wire) == ("primary", "backup")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_fallbacks_disabled(client: str) -> None:
|
||||
with wire_server(_peer(_PRIMARY_DIES)) as wire:
|
||||
router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP)
|
||||
with pytest.raises(litellm.APIConnectionError, match="overloaded"):
|
||||
_stream(client, router, disable_fallbacks=True)
|
||||
assert _deployments_hit(wire) == ("primary",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_dies_after_first_chunk(client: str) -> None:
|
||||
with wire_server(_peer({"primary": _dies_after("partial "), "backup": _serves("never asked")})) as wire:
|
||||
router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP)
|
||||
with pytest.raises(litellm.APIConnectionError, match="overloaded"):
|
||||
_stream(client, router)
|
||||
assert _deployments_hit(wire) == ("primary",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_router_retries_configured(client: str) -> None:
|
||||
with wire_server(_peer(_PRIMARY_DIES)) as wire:
|
||||
router: Final = Router(
|
||||
model_list=_router(wire, ("primary", "backup")).model_list,
|
||||
fallbacks=_PRIMARY_TO_BACKUP,
|
||||
num_retries=2,
|
||||
disable_cooldowns=True,
|
||||
)
|
||||
streamed: Final = _stream(client, router)
|
||||
assert streamed.text == "answered by the backup", streamed
|
||||
assert _deployments_hit(wire) == ("primary", "backup")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Outcome:
|
||||
text: str | None
|
||||
error: str | None
|
||||
hit: tuple[str, ...]
|
||||
|
||||
|
||||
def _outcome(client: str, wire: Wire, router: Router, **request: object) -> _Outcome:
|
||||
try:
|
||||
streamed: Final = _stream(client, router, **request)
|
||||
except litellm.APIConnectionError as error:
|
||||
return _Outcome(text=None, error=type(error).__name__, hit=_deployments_hit(wire))
|
||||
return _Outcome(text=streamed.text, error=None, hit=_deployments_hit(wire))
|
||||
|
||||
|
||||
def test_per_request_fallback_list_behaves_like_the_async_twin() -> None:
|
||||
with wire_server(_peer(_PRIMARY_DIES)) as wire:
|
||||
router: Final = _router(wire, ("primary", "backup"))
|
||||
twin: Final = _outcome("async", wire, router, fallbacks=_PRIMARY_TO_BACKUP)
|
||||
observed: Final = _outcome("sync", wire, router, fallbacks=_PRIMARY_TO_BACKUP)
|
||||
assert observed == twin, (observed, twin)
|
||||
assert observed.hit[:1] == ("primary",), observed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_max_fallbacks_caps_the_walk(client: str) -> None:
|
||||
with wire_server(_peer(_PRIMARY_AND_FB1_DIE)) as wire:
|
||||
router: Final = _router(wire, ("primary", "fb1", "fb2"), fallbacks=[{"primary": ["fb1", "fb2"]}], max_fallbacks=1)
|
||||
with pytest.raises(litellm.APIConnectionError, match="overloaded"):
|
||||
_stream(client, router)
|
||||
assert _deployments_hit(wire) == ("primary", "fb1")
|
||||
|
||||
|
||||
def test_called_inside_a_running_loop() -> None:
|
||||
with wire_server(_peer(_PRIMARY_DIES)) as wire:
|
||||
router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP)
|
||||
|
||||
async def inside_a_loop() -> _Streamed:
|
||||
return _stream_sync(router)
|
||||
|
||||
streamed: Final = asyncio.run(inside_a_loop())
|
||||
assert streamed.text == "answered by the backup", streamed
|
||||
assert _deployments_hit(wire) == ("primary", "backup")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_concurrent_burst(client: str) -> None:
|
||||
with wire_server(_peer(_PRIMARY_DIES)) as wire:
|
||||
router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP)
|
||||
if client == "async":
|
||||
|
||||
async def burst() -> tuple[_Streamed, ...]:
|
||||
return tuple(await asyncio.gather(*(_stream_async(router) for _ in range(_BURST))))
|
||||
|
||||
streamed: tuple[_Streamed, ...] = asyncio.run(burst())
|
||||
else:
|
||||
with ThreadPoolExecutor(max_workers=_BURST) as pool:
|
||||
streamed = tuple(pool.map(lambda _: _stream_sync(router), range(_BURST)))
|
||||
assert [item.text for item in streamed] == ["answered by the backup"] * _BURST, streamed
|
||||
hit: Final = _deployments_hit(wire)
|
||||
assert (hit.count("primary"), hit.count("backup"), len(hit)) == (_BURST, _BURST, 2 * _BURST), hit
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", _CLIENTS)
|
||||
def test_primary_drops_the_connection_before_content(client: str) -> None:
|
||||
with wire_server(_peer({"primary": _DROPS_BEFORE_CONTENT, "backup": _serves("answered by the backup")})) as wire:
|
||||
router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP)
|
||||
streamed: Final = _stream(client, router)
|
||||
assert streamed.text == "answered by the backup", streamed
|
||||
assert _deployments_hit(wire) == ("primary", "backup")
|
||||
|
|
@ -3650,7 +3650,11 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste
|
|||
}
|
||||
return chunk
|
||||
|
||||
with patch.object(router, "function_with_fallbacks", return_value=NestedFallbackStream()):
|
||||
with patch.object(
|
||||
router,
|
||||
"async_function_with_fallbacks_common_utils",
|
||||
new=AsyncMock(return_value=NestedFallbackStream()),
|
||||
):
|
||||
result = router._completion_streaming_iterator(
|
||||
model_response=FailedStream(),
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
|
|
@ -3667,6 +3671,126 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste
|
|||
assert result._hidden_params["model_id"] == "served-deployment"
|
||||
|
||||
|
||||
def test_completion_streaming_fallback_resumes_chain_without_retrying_primary():
|
||||
class FailingStream(CustomStreamWrapper):
|
||||
def __init__(self, model: str):
|
||||
super().__init__(
|
||||
completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock()
|
||||
)
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
raise MidStreamFallbackError(
|
||||
message=f"provider 500 from {self.model}",
|
||||
model=self.model,
|
||||
llm_provider="openai",
|
||||
generated_content="",
|
||||
is_pre_first_chunk=True,
|
||||
original_exception=litellm.InternalServerError(
|
||||
message=f"provider 500 from {self.model}", model=self.model, llm_provider="openai"
|
||||
),
|
||||
)
|
||||
|
||||
class OkStream(FailingStream):
|
||||
def __init__(self, model: str):
|
||||
super().__init__(model)
|
||||
self._chunks = iter(
|
||||
[litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": f"ok-from-{model}"}}])]
|
||||
)
|
||||
|
||||
def __next__(self):
|
||||
return next(self._chunks)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "fake-key"}},
|
||||
{"model_name": "backup", "litellm_params": {"model": "openai/backup-model", "api_key": "fake-key"}},
|
||||
],
|
||||
fallbacks=[{"primary": ["backup"]}],
|
||||
num_retries=0,
|
||||
)
|
||||
primary_calls: Final = iter(range(2))
|
||||
|
||||
def fake_completion(**kwargs):
|
||||
model_group: Final = kwargs["metadata"]["model_group"]
|
||||
if model_group == "backup":
|
||||
return OkStream(kwargs["model"])
|
||||
if next(primary_calls) > 0:
|
||||
raise RuntimeError("primary group retried")
|
||||
return FailingStream(kwargs["model"])
|
||||
|
||||
with patch("litellm.completion", side_effect=fake_completion) as provider_calls:
|
||||
response: Final = router.completion(model="primary", messages=[{"role": "user", "content": "hi"}], stream=True)
|
||||
content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in response)
|
||||
|
||||
assert content == "ok-from-openai/backup-model"
|
||||
assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == [
|
||||
"primary",
|
||||
"backup",
|
||||
]
|
||||
|
||||
|
||||
def test_completion_mid_stream_fallback_walks_every_entry_of_the_configured_list():
|
||||
class FailingStream(CustomStreamWrapper):
|
||||
def __init__(self, model: str):
|
||||
super().__init__(
|
||||
completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock()
|
||||
)
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
raise MidStreamFallbackError(
|
||||
message=f"provider 500 from {self.model}",
|
||||
model=self.model,
|
||||
llm_provider="openai",
|
||||
generated_content="",
|
||||
is_pre_first_chunk=True,
|
||||
original_exception=litellm.InternalServerError(
|
||||
message=f"provider 500 from {self.model}", model=self.model, llm_provider="openai"
|
||||
),
|
||||
)
|
||||
|
||||
class OkStream(FailingStream):
|
||||
def __init__(self, model: str):
|
||||
super().__init__(model)
|
||||
self._chunks = iter(
|
||||
[litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": f"ok-from-{model}"}}])]
|
||||
)
|
||||
|
||||
def __next__(self):
|
||||
return next(self._chunks)
|
||||
|
||||
def fake_completion(**kwargs):
|
||||
if "fb2" in kwargs["model"]:
|
||||
return OkStream(kwargs["model"])
|
||||
return FailingStream(kwargs["model"])
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "fake-key"}},
|
||||
{"model_name": "fb1", "litellm_params": {"model": "openai/fb1-model", "api_key": "fake-key"}},
|
||||
{"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "fake-key"}},
|
||||
],
|
||||
fallbacks=[{"primary": ["fb1", "fb2"]}],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
with patch("litellm.completion", side_effect=fake_completion) as provider_calls:
|
||||
response: Final = router.completion(model="primary", messages=[{"role": "user", "content": "hi"}], stream=True)
|
||||
content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in response if chunk is not None)
|
||||
|
||||
assert content == "ok-from-openai/fb2-model"
|
||||
assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == [
|
||||
"primary",
|
||||
"fb1",
|
||||
"fb2",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configured_list():
|
||||
"""LIT-7400: fallbacks=[{primary: [fb1, fb2]}] must reach fb2 when fb1 dies before its first chunk.
|
||||
|
|
@ -3828,7 +3952,11 @@ def test_completion_streaming_iterator_adopts_fallback_response_headers():
|
|||
def __iter__(self):
|
||||
return iter([])
|
||||
|
||||
with patch.object(router, "function_with_fallbacks", return_value=FallbackStream()):
|
||||
with patch.object(
|
||||
router,
|
||||
"async_function_with_fallbacks_common_utils",
|
||||
new=AsyncMock(return_value=FallbackStream()),
|
||||
):
|
||||
result = router._completion_streaming_iterator(
|
||||
model_response=FailedStream(),
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
|
|
@ -3895,8 +4023,8 @@ def test_completion_streaming_iterator_fallback_on_429():
|
|||
|
||||
with patch.object(
|
||||
router,
|
||||
"function_with_fallbacks",
|
||||
return_value=mock_fallback_response,
|
||||
"async_function_with_fallbacks_common_utils",
|
||||
new=AsyncMock(return_value=mock_fallback_response),
|
||||
) as mock_fallback:
|
||||
result = router._completion_streaming_iterator(
|
||||
model_response=mock_response,
|
||||
|
|
@ -3906,12 +4034,12 @@ def test_completion_streaming_iterator_fallback_on_429():
|
|||
|
||||
collected_chunks = list(result)
|
||||
|
||||
assert mock_fallback.called
|
||||
call_kwargs = mock_fallback.call_args
|
||||
mock_fallback.assert_awaited_once()
|
||||
call_kwargs = mock_fallback.await_args.kwargs["kwargs"]
|
||||
# Pre-first-chunk: should use original messages, no continuation prompt
|
||||
assert call_kwargs.kwargs.get("messages") == messages
|
||||
assert call_kwargs.get("messages") == messages
|
||||
# Verify original_function is _completion (sync)
|
||||
assert call_kwargs.kwargs.get("original_function") == router._completion
|
||||
assert call_kwargs.get("original_function") == router._completion
|
||||
|
||||
|
||||
def test_completion_streaming_iterator_preserves_hidden_params():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue