mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(passthrough): await own report without a shared limiter, register reports per loop and drain without a timeout by default
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
bf803cde01
commit
391d8f0f93
5 changed files with 191 additions and 47 deletions
|
|
@ -1529,11 +1529,10 @@ CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTER
|
|||
MCP_TOOL_NAME_PREFIX: Final = "mcp_tool"
|
||||
MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))
|
||||
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY: Final = max(
|
||||
1, int(os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY", "64"))
|
||||
)
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS: Final = int(
|
||||
os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", "10")
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS: Final[float | None] = (
|
||||
max(0.0, float(_raw_drain))
|
||||
if (_raw_drain := os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS"))
|
||||
else None
|
||||
)
|
||||
|
||||
# Headers to control callbacks
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import copy
|
|||
import json
|
||||
import posixpath
|
||||
import traceback
|
||||
import weakref
|
||||
from base64 import b64encode
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -42,7 +43,6 @@ from litellm._uuid import uuid
|
|||
from litellm.constants import (
|
||||
MAXIMUM_TRACEBACK_LINES_TO_LOG,
|
||||
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS,
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY,
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
|
||||
REDACTED_BY_LITELLM,
|
||||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
|
|
@ -944,28 +944,35 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
await self._upstream.aclose()
|
||||
|
||||
|
||||
_REPORT_CONCURRENCY: Final = asyncio.Semaphore(PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY)
|
||||
_REPORT_TASKS: Final[set[asyncio.Future[None]]] = set() # mutable-ok: in-flight report registry drained at shutdown
|
||||
|
||||
|
||||
async def _bounded_report(report: Awaitable[None], limiter: asyncio.Semaphore) -> None:
|
||||
async with limiter:
|
||||
await report
|
||||
_REPORT_TASKS: Final[ # mutable-ok: in-flight report registry per loop, drained at shutdown
|
||||
weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, set[asyncio.Future[None]]]
|
||||
] = weakref.WeakKeyDictionary()
|
||||
|
||||
|
||||
def _spawn_report_task(report: Awaitable[None]) -> asyncio.Future[None]:
|
||||
task: Final = asyncio.ensure_future(_bounded_report(report, _REPORT_CONCURRENCY))
|
||||
_REPORT_TASKS.add(task)
|
||||
task.add_done_callback(_REPORT_TASKS.discard)
|
||||
task: Final = asyncio.ensure_future(report)
|
||||
registry: Final = _REPORT_TASKS.setdefault(
|
||||
asyncio.get_running_loop(),
|
||||
set(), # mutable-ok: per-loop task set, tasks discard themselves on completion
|
||||
)
|
||||
registry.add(task)
|
||||
task.add_done_callback(registry.discard)
|
||||
return task
|
||||
|
||||
|
||||
async def drain_passthrough_upstream_error_reports(
|
||||
timeout: float = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
|
||||
timeout: float | None = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
|
||||
log_warning: Callable[..., None] = verbose_proxy_logger.warning,
|
||||
) -> None:
|
||||
pending: Final = tuple(_REPORT_TASKS)
|
||||
if pending:
|
||||
await asyncio.wait(pending, timeout=timeout)
|
||||
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
|
||||
if not pending:
|
||||
return
|
||||
_, still_pending = await asyncio.wait(pending, timeout=timeout)
|
||||
if still_pending:
|
||||
log_warning(
|
||||
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending",
|
||||
len(still_pending),
|
||||
)
|
||||
|
||||
|
||||
def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers:
|
||||
|
|
|
|||
|
|
@ -193,6 +193,56 @@ async def test_passthrough_disconnect_burst_logs_every_failure_once(gateway: Gat
|
|||
assert follow_up.status_code == 200, follow_up.text
|
||||
|
||||
|
||||
_SLOW_FAILURE_HOOK: Final = """
|
||||
import asyncio
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
class SlowFailureHook(CustomLogger):
|
||||
async def async_post_call_failure_hook(
|
||||
self, request_data, original_exception, user_api_key_dict, traceback_str=None
|
||||
):
|
||||
await asyncio.sleep(15)
|
||||
|
||||
|
||||
instance = SlowFailureHook()
|
||||
"""
|
||||
|
||||
|
||||
async def test_passthrough_sigterm_drains_reports_parked_on_a_slow_failure_hook(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
gate: Final = threading.Event()
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if "streamGenerateContent" in request.target:
|
||||
return Reply(
|
||||
status=429, content_type="text/event-stream", chunks=_RATE_LIMITED_FRAMES, gate_after_first=gate
|
||||
)
|
||||
return Reply(status=200, body=json.dumps({"ok": True}).encode())
|
||||
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["litellm_settings"].update({"callbacks": ["slow_hook.instance"]})
|
||||
(tmp_path / "slow_hook.py").write_text(_SLOW_FAILURE_HOOK)
|
||||
path: Final = tmp_path / "chaos-sigterm-drain.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with wire_server(respond) as wire:
|
||||
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=1) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
call_ids: Final = await asyncio.gather(
|
||||
*(_first_frame_then_close(str(candidate.client.base_url), candidate.key) for _ in range(20))
|
||||
)
|
||||
assert len(set(call_ids)) == 20, call_ids
|
||||
gate.set()
|
||||
owned.process.send_signal(signal.SIGTERM)
|
||||
owned.process.wait(timeout=90)
|
||||
for call_id in call_ids:
|
||||
_single_spend_row(call_id)
|
||||
|
||||
|
||||
async def _first_frame_then_close(base_url: str, key: str) -> str:
|
||||
async with httpx.AsyncClient(base_url=base_url, timeout=httpx.Timeout(5, connect=5), trust_env=False) as client:
|
||||
async with client.stream(
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import json
|
|||
import logging
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import zlib
|
||||
from collections.abc import Callable, Mapping
|
||||
from contextlib import ExitStack, contextmanager
|
||||
|
|
@ -33,7 +34,6 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
|||
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
|
||||
HttpPassThroughEndpointHelpers,
|
||||
InitPassThroughEndpointHelpers,
|
||||
_bounded_report,
|
||||
_PreviewReportingStream,
|
||||
_registered_pass_through_routes,
|
||||
_spawn_report_task,
|
||||
|
|
@ -5066,7 +5066,7 @@ async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
|
|||
with pytest.raises(asyncio.CancelledError):
|
||||
await consumer
|
||||
|
||||
pending: Final = tuple(_REPORT_TASKS)
|
||||
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
|
||||
assert len(pending) == 1, pending
|
||||
assert not pending[0].done()
|
||||
assert not hook_done.is_set()
|
||||
|
|
@ -5078,7 +5078,7 @@ async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
|
|||
await drain_passthrough_upstream_error_reports(timeout=5)
|
||||
assert hook_done.is_set()
|
||||
assert pending[0].done()
|
||||
assert not _REPORT_TASKS
|
||||
assert not _REPORT_TASKS.get(asyncio.get_running_loop())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -5114,38 +5114,115 @@ async def test_shutdown_drain_returns_with_report_still_pending_on_timeout():
|
|||
with pytest.raises(asyncio.CancelledError):
|
||||
await consumer
|
||||
|
||||
pending: Final = tuple(_REPORT_TASKS)
|
||||
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
|
||||
assert len(pending) == 1, pending
|
||||
await drain_passthrough_upstream_error_reports(timeout=0.05)
|
||||
warnings: Final = MagicMock()
|
||||
await drain_passthrough_upstream_error_reports(timeout=0.05, log_warning=warnings)
|
||||
warnings.assert_called_once_with(
|
||||
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", 1
|
||||
)
|
||||
assert not pending[0].done()
|
||||
pending[0].cancel()
|
||||
await asyncio.gather(pending[0], return_exceptions=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_report_concurrency_is_bounded_by_the_semaphore():
|
||||
limiter: Final = asyncio.Semaphore(2)
|
||||
release: Final = asyncio.Event()
|
||||
entered: list[int] = []
|
||||
finished: list[int] = []
|
||||
async def test_parked_reports_do_not_block_the_error_stream():
|
||||
"""A pile of still-running reports must not sit on the response path: a stream whose
|
||||
own report finishes immediately delivers its body without waiting on the others."""
|
||||
hold: Final = asyncio.Event()
|
||||
|
||||
def make_report(index: int):
|
||||
async def report() -> None:
|
||||
entered.append(index)
|
||||
async def parked() -> None:
|
||||
await hold.wait()
|
||||
|
||||
spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70))
|
||||
assert len(_REPORT_TASKS.get(asyncio.get_running_loop(), ())) == 70
|
||||
|
||||
upstream_response: Final = httpx.Response(
|
||||
status_code=500,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=_UpstreamErrorBodyStream(b"d" * 5000),
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
hook_done: Final = asyncio.Event()
|
||||
|
||||
async def report(preview: bytes) -> None:
|
||||
hook_done.set()
|
||||
|
||||
relay: Final = _PreviewReportingStream(
|
||||
upstream=upstream_response,
|
||||
report=report,
|
||||
log_warning=MagicMock(),
|
||||
spawn=_spawn_report_task,
|
||||
)
|
||||
|
||||
started: Final = time.monotonic()
|
||||
received: Final = b"".join([chunk async for chunk in relay.__aiter__()])
|
||||
elapsed: Final = time.monotonic() - started
|
||||
assert received == b"d" * 5000
|
||||
assert hook_done.is_set()
|
||||
assert elapsed < 0.5, elapsed
|
||||
|
||||
for task in spawned:
|
||||
task.cancel()
|
||||
await asyncio.gather(*spawned, return_exceptions=True)
|
||||
|
||||
|
||||
def test_report_registry_is_scoped_to_each_event_loop():
|
||||
"""Two consecutive asyncio.run calls: each loop's reports register and drain on their
|
||||
own loop, so a shared module-level primitive bound to the first loop never breaks the second."""
|
||||
|
||||
async def run_once() -> None:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
async def parked() -> None:
|
||||
await release.wait()
|
||||
finished.append(index)
|
||||
|
||||
return report
|
||||
spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70))
|
||||
assert len(_REPORT_TASKS.get(loop, ())) == 70
|
||||
|
||||
tasks: Final = tuple(asyncio.ensure_future(_bounded_report(make_report(index)(), limiter)) for index in range(4))
|
||||
for _ in range(20):
|
||||
await asyncio.sleep(0)
|
||||
assert len(entered) == 2, entered
|
||||
release.set()
|
||||
await drain_passthrough_upstream_error_reports()
|
||||
assert all(task.done() for task in spawned)
|
||||
assert not _REPORT_TASKS.get(loop)
|
||||
|
||||
release.set()
|
||||
await asyncio.gather(*tasks)
|
||||
assert len(entered) == 4, entered
|
||||
assert len(finished) == 4, finished
|
||||
asyncio.run(run_once())
|
||||
asyncio.run(run_once())
|
||||
assert all(len(pending) == 0 for pending in _REPORT_TASKS.values())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_drain_waits_without_a_timeout():
|
||||
finished: Final = asyncio.Event()
|
||||
|
||||
async def report() -> None:
|
||||
await asyncio.sleep(0.3)
|
||||
finished.set()
|
||||
|
||||
task: Final = _spawn_report_task(report())
|
||||
warnings: Final = MagicMock()
|
||||
await drain_passthrough_upstream_error_reports(timeout=None, log_warning=warnings)
|
||||
assert task.done()
|
||||
assert finished.is_set()
|
||||
warnings.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_drain_timeout_warns_with_the_pending_count():
|
||||
hold: Final = asyncio.Event()
|
||||
|
||||
async def parked() -> None:
|
||||
await hold.wait()
|
||||
|
||||
task: Final = _spawn_report_task(parked())
|
||||
warnings: Final = MagicMock()
|
||||
await drain_passthrough_upstream_error_reports(timeout=0.05, log_warning=warnings)
|
||||
warnings.assert_called_once_with(
|
||||
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", 1
|
||||
)
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
|
||||
class _UpstreamErrorGzipStreamDropping(httpx.AsyncByteStream):
|
||||
|
|
|
|||
|
|
@ -70,12 +70,23 @@ def _build_constant_env_var_map() -> dict[str, str]:
|
|||
return env_var_map
|
||||
|
||||
|
||||
def test_passthrough_error_report_concurrency_env_zero_clamps_to_one(monkeypatch):
|
||||
"""A zero/negative concurrency would deadlock every report; the constant clamps to 1."""
|
||||
monkeypatch.setenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY", "0")
|
||||
@pytest.mark.parametrize(
|
||||
"env_value, expected",
|
||||
[
|
||||
(None, None),
|
||||
("2.5", 2.5),
|
||||
("-1", 0.0),
|
||||
],
|
||||
)
|
||||
def test_passthrough_error_report_drain_seconds_env_parsing(monkeypatch, env_value, expected):
|
||||
"""Unset waits for every report; a value clamps to >= 0 seconds."""
|
||||
if env_value is None:
|
||||
monkeypatch.delenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", raising=False)
|
||||
else:
|
||||
monkeypatch.setenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", env_value)
|
||||
try:
|
||||
reloaded = importlib.reload(constants)
|
||||
assert reloaded.PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY == 1
|
||||
assert reloaded.PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS == expected
|
||||
finally:
|
||||
monkeypatch.undo()
|
||||
importlib.reload(constants)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue