mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(router): forward include_fallback_errors through multi-hop fallbacks
run_async_fallback received include_fallback_errors as an explicit named parameter, so it was bound out of **kwargs and never reached the nested async_function_with_fallbacks call. Multi-hop fallback chains (a fallback group that itself fails over) therefore stopped collecting fallback errors beyond the first hop when a caller opted in. Re-inject the flag into kwargs before the nested call so inner hops keep accumulating errors, which add_fallback_headers_to_response already merges across levels.
This commit is contained in:
parent
d093dc2e3b
commit
dc6611b4f7
2 changed files with 47 additions and 0 deletions
|
|
@ -139,6 +139,8 @@ async def run_async_fallback(
|
|||
fallback_depth = fallback_depth + 1
|
||||
kwargs["fallback_depth"] = fallback_depth
|
||||
kwargs["max_fallbacks"] = max_fallbacks
|
||||
if include_fallback_errors:
|
||||
kwargs["include_fallback_errors"] = include_fallback_errors
|
||||
response = await litellm_router.async_function_with_fallbacks(
|
||||
*args, **kwargs
|
||||
)
|
||||
|
|
|
|||
|
|
@ -80,6 +80,51 @@ async def test_run_async_fallback_raises_when_all_fallbacks_fail():
|
|||
)
|
||||
|
||||
|
||||
class RecordingRouter:
|
||||
def __init__(self):
|
||||
self.received_kwargs = None
|
||||
|
||||
def log_retry(self, kwargs, e):
|
||||
return kwargs
|
||||
|
||||
async def async_function_with_fallbacks(self, *args, **kwargs):
|
||||
self.received_kwargs = kwargs
|
||||
return StreamingWrapper()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_forwards_include_fallback_errors_to_nested_call():
|
||||
"""A nested fallback (multi-hop) must keep collecting errors, so the opt-in
|
||||
flag has to reach the nested async_function_with_fallbacks call."""
|
||||
router = RecordingRouter()
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["fallback-model"],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("upstream limited request"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
include_fallback_errors=True,
|
||||
)
|
||||
|
||||
assert router.received_kwargs.get("include_fallback_errors") is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_does_not_forward_flag_without_opt_in():
|
||||
router = RecordingRouter()
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["fallback-model"],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("upstream limited request"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
assert "include_fallback_errors" not in router.received_kwargs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_skips_original_model_group():
|
||||
response = await run_async_fallback(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue