mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(router): coordinate async and sync failure handlers at remaining router call sites (#39887)
* fix(router): coordinate async and sync failure handlers at remaining router call sites Five router failure paths still scheduled logging_obj.async_failure_handler as a task while starting logging_obj.failure_handler on a raw thread, so both handlers mutated the same logging object concurrently. Route them through dispatch_failure_handlers like the streaming paths already do. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(router): wait on the real logging executor and justify the callbacks global patch Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(logging): submit sync failure handler even when the dispatch task is cancelled Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(logging): justify the executor submit patch Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
95a1d301b6
commit
0ad361a728
4 changed files with 216 additions and 40 deletions
|
|
@ -1918,7 +1918,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
two paths cannot mutate it at the same time. ``prefer_async_handlers`` only
|
||||
bypasses the sync-SDK-only shortcut (e.g. ``async for`` on a stream from
|
||||
``completion()``); legacy string callbacks still run via
|
||||
``executor.submit(failure_handler)`` when configured.
|
||||
``executor.submit(failure_handler)`` when configured, and still get submitted
|
||||
when the awaiting task is cancelled (e.g. the event loop shuts down right after
|
||||
the request failed).
|
||||
"""
|
||||
litellm_params: Final = self.model_call_details.get("litellm_params", {}) or {}
|
||||
sync_sdk: Final = self._is_sync_litellm_request(litellm_params)
|
||||
|
|
@ -1927,12 +1929,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.failure_handler(exception, traceback_exception)
|
||||
return
|
||||
|
||||
await self.async_failure_handler(exception, traceback_exception)
|
||||
|
||||
if not self._should_run_sync_failure_callbacks_for_async_calls():
|
||||
return
|
||||
|
||||
executor.submit(self.failure_handler, exception, traceback_exception)
|
||||
try:
|
||||
await self.async_failure_handler(exception, traceback_exception)
|
||||
finally:
|
||||
if self._should_run_sync_failure_callbacks_for_async_calls():
|
||||
executor.submit(self.failure_handler, exception, traceback_exception)
|
||||
|
||||
def should_run_logging(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -8380,17 +8380,12 @@ class Router:
|
|||
## LOG FAILURE EVENT
|
||||
if logging_obj is not None:
|
||||
asyncio.create_task(
|
||||
logging_obj.async_failure_handler(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback.format_exc(),
|
||||
end_time=time.time(),
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback.format_exc()),
|
||||
).start() # log response
|
||||
_set_cooldown_deployments(
|
||||
litellm_router_instance=self,
|
||||
exception_status=e.status_code,
|
||||
|
|
@ -8403,17 +8398,12 @@ class Router:
|
|||
## LOG FAILURE EVENT
|
||||
if logging_obj is not None:
|
||||
asyncio.create_task(
|
||||
logging_obj.async_failure_handler(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback.format_exc(),
|
||||
end_time=time.time(),
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback.format_exc()),
|
||||
).start() # log response
|
||||
raise e
|
||||
|
||||
async def async_callback_filter_deployments(
|
||||
|
|
@ -8451,17 +8441,12 @@ class Router:
|
|||
## LOG FAILURE EVENT
|
||||
if logging_obj is not None:
|
||||
asyncio.create_task(
|
||||
logging_obj.async_failure_handler(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback.format_exc(),
|
||||
end_time=time.time(),
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback.format_exc()),
|
||||
).start() # log response
|
||||
raise e
|
||||
return returned_healthy_deployments
|
||||
|
||||
|
|
@ -12643,13 +12628,13 @@ class Router:
|
|||
logging_obj: Final = request_kwargs.get("litellm_logging_obj", None)
|
||||
|
||||
if logging_obj is not None:
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback_exception),
|
||||
).start() # log response
|
||||
# Handle any exceptions that might occur during streaming
|
||||
asyncio.create_task(logging_obj.async_failure_handler(e, traceback_exception))
|
||||
asyncio.create_task(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback_exception,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
raise e
|
||||
|
||||
async def async_get_available_deployment_for_pass_through(
|
||||
|
|
@ -12777,11 +12762,13 @@ class Router:
|
|||
if request_kwargs is not None:
|
||||
logging_obj: Final = request_kwargs.get("litellm_logging_obj", None)
|
||||
if logging_obj is not None:
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback_exception),
|
||||
).start()
|
||||
asyncio.create_task(logging_obj.async_failure_handler(e, traceback_exception))
|
||||
asyncio.create_task(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback_exception,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
raise e
|
||||
|
||||
async def _run_routing_plugins(
|
||||
|
|
|
|||
|
|
@ -1546,6 +1546,62 @@ async def test_dispatch_failure_handlers_async_completes_before_sync_submit(
|
|||
assert events == ["async_start", "async_end", "sync_submit"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_failure_handlers_submits_sync_handler_when_task_is_cancelled(
|
||||
logging_obj,
|
||||
):
|
||||
"""Cancelling the dispatch task mid-await still submits the sync failure_handler.
|
||||
|
||||
Router failure paths fire the dispatcher with ``asyncio.create_task`` and raise
|
||||
right away. When the event loop is torn down before the task finishes (a short
|
||||
``asyncio.run`` in the SDK), the cancelled task must still hand the sync callbacks
|
||||
to the executor, as the old raw-thread path did, and only once the async handler
|
||||
has stopped.
|
||||
"""
|
||||
exception = ValueError("boom")
|
||||
traceback_exception = "traceback"
|
||||
events: list[str] = []
|
||||
async_started = asyncio.Event()
|
||||
|
||||
async def _async_failure(exc, tb, **kwargs):
|
||||
events.append("async_start")
|
||||
async_started.set()
|
||||
await asyncio.sleep(10)
|
||||
events.append("async_end")
|
||||
|
||||
def _submit(*args, **kwargs):
|
||||
events.append("sync_submit")
|
||||
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
|
||||
with (
|
||||
patch.object(logging_obj, "async_failure_handler", side_effect=_async_failure),
|
||||
patch.object(logging_obj, "failure_handler", new_callable=MagicMock),
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_should_run_sync_failure_callbacks_for_async_calls",
|
||||
return_value=True,
|
||||
),
|
||||
patch( # test-quality-ok: the executor submit is the observable
|
||||
"litellm.litellm_core_utils.litellm_logging.executor.submit",
|
||||
side_effect=_submit,
|
||||
),
|
||||
):
|
||||
task = asyncio.create_task(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception,
|
||||
traceback_exception,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
await async_started.wait()
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
assert events == ["async_start", "sync_submit"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_failure_handlers_submits_sync_handler_for_failure_only_callbacks(
|
||||
logging_obj,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import json
|
|||
import logging
|
||||
import os
|
||||
import threading
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -20,6 +21,7 @@ import litellm
|
|||
from litellm import Router
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES,
|
||||
|
|
@ -12940,6 +12942,136 @@ async def test_router_retry_policy_controls_upstream_attempt_count(
|
|||
assert upstream.call_count == expected_upstream_calls
|
||||
|
||||
|
||||
def _make_failure_logging_obj():
|
||||
return LiteLLMLogging(
|
||||
model="gpt-5.6",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="acompletion",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="lit-6960",
|
||||
function_id="f",
|
||||
)
|
||||
|
||||
|
||||
async def _assert_router_failure_logging_is_coordinated(logging_obj, trigger, expected_exception):
|
||||
"""The sync failure_handler must not start until async_failure_handler has finished on the shared logging_obj."""
|
||||
events: list[str] = []
|
||||
sync_done = threading.Event()
|
||||
|
||||
async def _async_failure(*args, **kwargs):
|
||||
events.append("async_start")
|
||||
await asyncio.sleep(0.05)
|
||||
events.append("async_end")
|
||||
|
||||
def _sync_failure(*args, **kwargs):
|
||||
events.append("sync_start")
|
||||
sync_done.set()
|
||||
|
||||
with (
|
||||
patch.object(logging_obj, "async_failure_handler", side_effect=_async_failure),
|
||||
patch.object(logging_obj, "failure_handler", side_effect=_sync_failure),
|
||||
patch.object(logging_obj, "_should_run_sync_failure_callbacks_for_async_calls", return_value=True),
|
||||
):
|
||||
with pytest.raises(expected_exception):
|
||||
await trigger()
|
||||
pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
|
||||
await asyncio.gather(*pending)
|
||||
assert await asyncio.to_thread(sync_done.wait, 5), "failure_handler never ran"
|
||||
|
||||
assert events == ["async_start", "async_end", "sync_start"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"hook_error",
|
||||
[
|
||||
litellm.RateLimitError(message="rpm exceeded", llm_provider="openai", model="gpt-5.6"),
|
||||
RuntimeError("pre call check blew up"),
|
||||
],
|
||||
)
|
||||
async def test_async_routing_strategy_pre_call_checks_failure_logging_is_coordinated(hook_error):
|
||||
class _RaisingPreCallCheck(CustomLogger):
|
||||
async def async_pre_call_check(self, deployment, parent_otel_span):
|
||||
raise hook_error
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake"}}]
|
||||
)
|
||||
deployment = router.model_list[0]
|
||||
logging_obj = _make_failure_logging_obj()
|
||||
|
||||
with patch.object(litellm, "callbacks", [_RaisingPreCallCheck()]): # test-quality-ok: router reads this global
|
||||
await _assert_router_failure_logging_is_coordinated(
|
||||
logging_obj,
|
||||
lambda: router.async_routing_strategy_pre_call_checks(
|
||||
deployment=deployment, parent_otel_span=None, logging_obj=logging_obj
|
||||
),
|
||||
type(hook_error),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_callback_filter_deployments_failure_logging_is_coordinated():
|
||||
class _RaisingFilter(CustomLogger):
|
||||
async def async_filter_deployments(self, *args, **kwargs):
|
||||
raise RuntimeError("filter blew up")
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake"}}]
|
||||
)
|
||||
logging_obj = _make_failure_logging_obj()
|
||||
|
||||
with patch.object(litellm, "callbacks", [_RaisingFilter()]): # test-quality-ok: router reads this global
|
||||
await _assert_router_failure_logging_is_coordinated(
|
||||
logging_obj,
|
||||
lambda: router.async_callback_filter_deployments(
|
||||
model="gpt-5.6",
|
||||
healthy_deployments=router.model_list,
|
||||
messages=None,
|
||||
parent_otel_span=None,
|
||||
request_kwargs={},
|
||||
logging_obj=logging_obj,
|
||||
),
|
||||
RuntimeError,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_available_deployment_failure_logging_is_coordinated():
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake"}}]
|
||||
)
|
||||
logging_obj = _make_failure_logging_obj()
|
||||
|
||||
await _assert_router_failure_logging_is_coordinated(
|
||||
logging_obj,
|
||||
lambda: router.async_get_available_deployment(
|
||||
model="model-that-is-not-configured",
|
||||
request_kwargs={"litellm_logging_obj": logging_obj},
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
),
|
||||
litellm.BadRequestError,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_available_deployment_for_pass_through_failure_logging_is_coordinated():
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake"}}]
|
||||
)
|
||||
logging_obj = _make_failure_logging_obj()
|
||||
|
||||
await _assert_router_failure_logging_is_coordinated(
|
||||
logging_obj,
|
||||
lambda: router.async_get_available_deployment_for_pass_through(
|
||||
model="gpt-5.6",
|
||||
request_kwargs={"litellm_logging_obj": logging_obj},
|
||||
),
|
||||
litellm.BadRequestError,
|
||||
)
|
||||
|
||||
|
||||
class _InFlightTracker:
|
||||
def __init__(self) -> None:
|
||||
self.current = 0
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue