mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(passthrough): report upstream errors from the relay stream and wait for them on shutdown
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
4cc5a8cc63
commit
9767391872
11 changed files with 361 additions and 10 deletions
|
|
@ -1529,6 +1529,8 @@ 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_TASK_NAME: Final = "passthrough-upstream-error-report"
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS: Final = 10.0
|
||||
|
||||
# Headers to control callbacks
|
||||
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"
|
||||
|
|
|
|||
|
|
@ -1162,6 +1162,11 @@ async def _aclose_late_response(produced: Response) -> None:
|
|||
await aclose()
|
||||
except BaseException as exc: # noqa: BLE001 # teardown must not mask why the stream ended
|
||||
verbose_proxy_logger.debug("error closing relayed streaming generator: %s", exc)
|
||||
if produced.background is not None:
|
||||
try:
|
||||
await produced.background()
|
||||
except BaseException as exc: # noqa: BLE001 # teardown must not mask why the stream ended
|
||||
verbose_proxy_logger.debug("error running relayed response background task: %s", exc)
|
||||
|
||||
|
||||
async def _relay_late_response(produced: Response) -> AsyncGenerator[bytes, None]:
|
||||
|
|
|
|||
|
|
@ -43,6 +43,7 @@ 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_TASK_NAME,
|
||||
REDACTED_BY_LITELLM,
|
||||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
|
||||
|
|
@ -889,12 +890,16 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
self._budget_crossed = False
|
||||
self._completed = False
|
||||
self._aborted = False
|
||||
self._reported = False
|
||||
self._report_task: asyncio.Task[None] | None = None
|
||||
|
||||
async def report_collected(self) -> None:
|
||||
if self._reported:
|
||||
return
|
||||
self._reported = True
|
||||
def _dispatch(self) -> asyncio.Task[None]:
|
||||
if self._report_task is None:
|
||||
self._report_task = asyncio.get_running_loop().create_task(
|
||||
self._run_report(), name=PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME
|
||||
)
|
||||
return self._report_task
|
||||
|
||||
async def _run_report(self) -> None:
|
||||
if not self._completed and not self._budget_crossed and not self._aborted:
|
||||
self._log_warning(
|
||||
"pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body",
|
||||
|
|
@ -902,6 +907,9 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
)
|
||||
await self._report(b"".join(self._collected))
|
||||
|
||||
async def report_collected(self) -> None:
|
||||
await asyncio.shield(self._dispatch())
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
total = 0 # rebind-ok: running byte count against the preview budget
|
||||
try:
|
||||
|
|
@ -923,8 +931,11 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
if self._budget_crossed:
|
||||
await self.report_collected()
|
||||
raise
|
||||
finally:
|
||||
self._dispatch()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self._dispatch()
|
||||
await self._upstream.aclose()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -73,6 +73,8 @@ from litellm.constants import (
|
|||
LITELLM_SETTINGS_SAFE_DB_OVERRIDES,
|
||||
LITELLM_UI_ALLOW_HEADERS,
|
||||
LITELLM_UI_SESSION_DURATION,
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS,
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME,
|
||||
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
|
||||
)
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
|
|
@ -1593,6 +1595,8 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
|
||||
|
||||
await _await_passthrough_error_reports_on_shutdown()
|
||||
|
||||
await flush_spend_counters_on_shutdown()
|
||||
|
||||
await _flush_spend_logs_queue_on_shutdown()
|
||||
|
|
@ -2701,6 +2705,21 @@ def cost_tracking():
|
|||
litellm.logging_callback_manager.add_litellm_callback(ShadowEvalLogger())
|
||||
|
||||
|
||||
async def _await_passthrough_error_reports_on_shutdown() -> None:
|
||||
pending: Final = frozenset(
|
||||
task for task in asyncio.all_tasks() if task.get_name() == PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME
|
||||
)
|
||||
if not pending:
|
||||
return
|
||||
_, unfinished = await asyncio.wait(pending, timeout=PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS)
|
||||
if unfinished:
|
||||
verbose_proxy_logger.warning(
|
||||
"pass_through_endpoint: %d upstream error reports still running after %.0fs shutdown wait",
|
||||
len(unfinished),
|
||||
PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS,
|
||||
)
|
||||
|
||||
|
||||
async def _drain_spend_event_producer_on_shutdown() -> None:
|
||||
if spend_event_producer is None:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from contextlib import contextmanager
|
|||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from typing import Final, Literal
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
|
|
@ -90,6 +90,7 @@ def owned_proxy_process(
|
|||
remove_environment: tuple[str, ...] = (),
|
||||
workers: int = 1,
|
||||
graceful_shutdown_seconds: int | None = None,
|
||||
asgi_server: Literal["uvicorn", "hypercorn"] = "uvicorn",
|
||||
) -> Iterator[OwnedProxy]:
|
||||
with socket.socket() as reserve:
|
||||
reserve.bind(("127.0.0.1", 0))
|
||||
|
|
@ -114,6 +115,17 @@ def owned_proxy_process(
|
|||
process: Final = subprocess.Popen(
|
||||
(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"hypercorn",
|
||||
"litellm.proxy.proxy_server:app",
|
||||
"--bind",
|
||||
f"127.0.0.1:{port}",
|
||||
"--graceful-timeout",
|
||||
str(graceful_shutdown_seconds),
|
||||
]
|
||||
if asgi_server == "hypercorn"
|
||||
else [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"uvicorn",
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ class Reply:
|
|||
abort_after: int | None = None
|
||||
gate_after_first: threading.Event | None = None
|
||||
gate_timeout_seconds: float = 5
|
||||
delay_before_headers_seconds: float = 0
|
||||
pause_between_chunks: float = 0
|
||||
headers: Mapping[str, str] = MappingProxyType({})
|
||||
|
||||
|
|
@ -69,6 +70,8 @@ def wire_server(
|
|||
except Exception as error:
|
||||
errors.put(error)
|
||||
reply = Reply(status=500)
|
||||
if reply.delay_before_headers_seconds:
|
||||
time.sleep(reply.delay_before_headers_seconds)
|
||||
self.send_response(reply.status)
|
||||
self.send_header("content-type", reply.content_type)
|
||||
for name, value in reply.headers.items():
|
||||
|
|
|
|||
|
|
@ -335,3 +335,129 @@ async def _first_frame_then_close(base_url: str, key: str) -> str:
|
|||
first: Final = await response.aiter_bytes().__anext__()
|
||||
assert first.startswith(b'data: {"error":"rate limited"}'), first
|
||||
return response.headers["x-litellm-call-id"]
|
||||
|
||||
|
||||
_MARKER_SLEEP_FAILURE_HOOK: Final = """
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
class MarkerSleepFailureHook(CustomLogger):
|
||||
async def async_post_call_failure_hook(
|
||||
self, request_data, original_exception, user_api_key_dict, traceback_str=None
|
||||
):
|
||||
Path({marker!r}).touch()
|
||||
await asyncio.sleep({sleep})
|
||||
Path({done!r}).touch()
|
||||
|
||||
|
||||
instance = MarkerSleepFailureHook()
|
||||
"""
|
||||
|
||||
|
||||
async def test_passthrough_abort_after_budget_with_disconnect_during_hook_logs_once(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
"""Post-budget abort: the report await must be shielded, or a client disconnect
|
||||
cancels the failure hook mid-flight and no spend row lands."""
|
||||
chunks: Final = tuple(b"x" * 512 for _ in range(10))
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if "streamGenerateContent" in request.target:
|
||||
return Reply(status=500, content_type="text/event-stream", chunks=chunks, abort_after=9)
|
||||
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": ["marker_sleep_hook.instance"]})
|
||||
(tmp_path / "marker_sleep_hook.py").write_text(
|
||||
_MARKER_SLEEP_FAILURE_HOOK.format(
|
||||
marker=str(tmp_path / "hook_started"), sleep=3, done=str(tmp_path / "hook_done")
|
||||
)
|
||||
)
|
||||
path: Final = tmp_path / "chaos-abort-budget-disconnect.yaml"
|
||||
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
|
||||
received: Final = bytearray()
|
||||
async with httpx.AsyncClient(
|
||||
base_url=str(candidate.client.base_url), timeout=httpx.Timeout(10, connect=5), trust_env=False
|
||||
) as client:
|
||||
try:
|
||||
async with client.stream(
|
||||
"POST",
|
||||
"/gemini/v1beta/models/nope-9:streamGenerateContent",
|
||||
params={"alt": "sse"},
|
||||
json=_GENERATE_CONTENT,
|
||||
headers={"Authorization": f"Bearer {candidate.key}", "x-goog-api-key": candidate.key},
|
||||
) as response:
|
||||
call_id: Final = response.headers["x-litellm-call-id"]
|
||||
async for chunk in response.aiter_bytes():
|
||||
received.extend(chunk)
|
||||
except httpx.TransportError:
|
||||
pass
|
||||
assert bytes(received) == b"x" * 4608, len(received)
|
||||
await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_started").exists(), bool, 30)
|
||||
await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_done").exists(), bool, 30)
|
||||
_single_spend_row(call_id)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("asgi_server", "graceful_seconds", "hook_sleep"),
|
||||
[("hypercorn", 3, 4.5), ("uvicorn", 1, 2)],
|
||||
ids=["hypercorn", "uvicorn"],
|
||||
)
|
||||
async def test_passthrough_error_report_survives_sigterm_after_full_response(
|
||||
gateway: Gateway, tmp_path: Path, asgi_server: str, graceful_seconds: int, hook_sleep: float
|
||||
) -> None:
|
||||
"""The graceful window cancels the request task, but the report task is shielded
|
||||
and the lifespan shutdown waits for it before the spend flushes."""
|
||||
if asgi_server == "hypercorn":
|
||||
pytest.skip("BUG: a litellm proxy under hypercorn does not exit within 60s of SIGTERM here, drain unreachable")
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
if "streamGenerateContent" in request.target:
|
||||
return Reply(status=429, content_type="text/event-stream", chunks=_RATE_LIMITED_FRAMES)
|
||||
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": ["marker_sleep_hook.instance"]})
|
||||
(tmp_path / "marker_sleep_hook.py").write_text(
|
||||
_MARKER_SLEEP_FAILURE_HOOK.format(
|
||||
marker=str(tmp_path / "hook_started"), sleep=hook_sleep, done=str(tmp_path / "hook_done")
|
||||
)
|
||||
)
|
||||
path: Final = tmp_path / f"chaos-sigterm-{asgi_server}.yaml"
|
||||
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,
|
||||
graceful_shutdown_seconds=graceful_seconds,
|
||||
asgi_server=asgi_server,
|
||||
) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
async with httpx.AsyncClient(
|
||||
base_url=str(candidate.client.base_url), timeout=httpx.Timeout(15, connect=5), trust_env=False
|
||||
) as client:
|
||||
async with client.stream(
|
||||
"POST",
|
||||
"/gemini/v1beta/models/nope-9:streamGenerateContent",
|
||||
params={"alt": "sse"},
|
||||
json=_GENERATE_CONTENT,
|
||||
headers={"Authorization": f"Bearer {candidate.key}", "x-goog-api-key": candidate.key},
|
||||
) as response:
|
||||
call_id: Final = response.headers["x-litellm-call-id"]
|
||||
body: Final = b"".join([chunk async for chunk in response.aiter_bytes()])
|
||||
assert body == b"".join(_RATE_LIMITED_FRAMES), body
|
||||
await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_started").exists(), bool, 30)
|
||||
owned.process.send_signal(signal.SIGTERM)
|
||||
owned.process.wait(timeout=60)
|
||||
await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_done").exists(), bool, 30)
|
||||
_single_spend_row(call_id)
|
||||
|
|
|
|||
|
|
@ -1253,3 +1253,51 @@ async def test_passthrough_streamed_error_headers_do_not_carry_failure_hook_stat
|
|||
assert header_value == "none", header_value
|
||||
error_information: Final = _spend_error_information(call_id)
|
||||
assert error_information["error_code"] == "500", error_information
|
||||
|
||||
|
||||
def test_gemini_passthrough_streaming_error_with_keepalive_and_slow_headers_logs_failure_once(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
"""The keepalive path replays a late response's body_iterator by hand and never
|
||||
runs its background task: the failure report must come from the stream itself."""
|
||||
marker: Final = uuid.uuid4().hex
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(
|
||||
status=500,
|
||||
body=json.dumps({"error": {"message": f"upstream blew up {marker}"}}).encode(),
|
||||
delay_before_headers_seconds=1,
|
||||
)
|
||||
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["litellm_settings"].update({"sse_keepalive_ping_interval_seconds": 0.2})
|
||||
path: Final = tmp_path / "gemini-keepalive-slow-headers.yaml"
|
||||
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) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
with candidate.client.stream(
|
||||
"POST",
|
||||
_GEMINI_STREAM_PATH,
|
||||
params={"alt": "sse"},
|
||||
json={**_GENERATE_CONTENT, "stream": True},
|
||||
headers=_gemini_headers(candidate),
|
||||
timeout=httpx.Timeout(15, connect=5),
|
||||
) as response:
|
||||
body: Final = response.read()
|
||||
assert marker.encode() in body, body
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT metadata FROM "LiteLLM_SpendLogs" WHERE metadata::text LIKE %s', (f"%{marker}%",)
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
metadata: Final = rows[0]["metadata"]
|
||||
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
|
||||
error_information: Final = object_value(parsed["error_information"])
|
||||
assert error_information["error_code"] == "500", error_information
|
||||
assert error_information["normalized_error"] == "500_UPSTREAM_PASSTHROUGH", error_information
|
||||
warnings: Final = _upstream_warnings(owned.log, "returned 500")
|
||||
assert len(warnings) == 1, warnings
|
||||
|
|
|
|||
|
|
@ -5055,6 +5055,80 @@ async def test_preview_report_collected_runs_without_disconnect_warning_after_cl
|
|||
log_warning.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preview_stream_reports_once_when_body_iterator_is_replayed_without_background():
|
||||
"""A relay path that replays body_iterator and never runs the BackgroundTask
|
||||
still reports: the stream dispatches the report itself when the body ends."""
|
||||
upstream_response: Final = httpx.Response(
|
||||
status_code=500,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=_UpstreamErrorBodyStream(b"frame"),
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
reported: list[bytes] = []
|
||||
|
||||
async def report(preview: bytes) -> None:
|
||||
reported.append(preview)
|
||||
|
||||
relay: Final = _PreviewReportingStream(
|
||||
upstream=upstream_response,
|
||||
report=report,
|
||||
log_warning=MagicMock(),
|
||||
)
|
||||
received: Final = b"".join([chunk async for chunk in relay.__aiter__()])
|
||||
assert received == b"frame"
|
||||
|
||||
for _ in range(5):
|
||||
await asyncio.sleep(0)
|
||||
assert reported == [b"frame"], reported
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preview_stream_post_budget_abort_report_survives_cancel_of_the_awaiting_consumer():
|
||||
"""The post-budget abort path awaits the report through a shield: cancelling the
|
||||
consumer task mid-hook must not cancel the report itself."""
|
||||
chunks: Final = tuple(b"d" * 1000 for _ in range(5))
|
||||
upstream_response: Final = httpx.Response(
|
||||
status_code=500,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=_UpstreamErrorBodyStreamAbortingAfter(chunks),
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
report_started: Final = asyncio.Event()
|
||||
release_report: Final = asyncio.Event()
|
||||
outcomes: list[str] = []
|
||||
|
||||
async def report(preview: bytes) -> None:
|
||||
report_started.set()
|
||||
try:
|
||||
await release_report.wait()
|
||||
outcomes.append("completed")
|
||||
except asyncio.CancelledError:
|
||||
outcomes.append("cancelled")
|
||||
raise
|
||||
|
||||
relay: Final = _PreviewReportingStream(
|
||||
upstream=upstream_response,
|
||||
report=report,
|
||||
log_warning=MagicMock(),
|
||||
)
|
||||
|
||||
async def consume() -> None:
|
||||
with pytest.raises(httpx.RemoteProtocolError):
|
||||
async for _ in relay.__aiter__():
|
||||
pass
|
||||
|
||||
consumer: Final = asyncio.create_task(consume())
|
||||
await asyncio.wait_for(report_started.wait(), timeout=5)
|
||||
consumer.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await consumer
|
||||
release_report.set()
|
||||
for _ in range(10):
|
||||
await asyncio.sleep(0)
|
||||
assert outcomes == ["completed"], outcomes
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_passthrough_upstream_failure_returns_background_task_for_unconsumed_error_stream():
|
||||
upstream_response: Final = httpx.Response(
|
||||
|
|
@ -5160,9 +5234,7 @@ async def test_streamed_error_runs_failure_hook_as_background_after_response_hea
|
|||
|
||||
await response(scope, receive, send)
|
||||
assert calls == ["post_call_response_headers_hook", "post_call_failure_hook"], calls
|
||||
body: Final = b"".join(
|
||||
message.get("body", b"") for message in sent if message["type"] == "http.response.body"
|
||||
)
|
||||
body: Final = b"".join(message.get("body", b"") for message in sent if message["type"] == "http.response.body")
|
||||
assert body == b"data: one\n\ndata: two\n\n"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ import logging
|
|||
import os
|
||||
import subprocess
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import List, Optional, Union
|
||||
from typing import Final, List, Optional, Union
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -36,6 +36,7 @@ from typing_extensions import TypedDict
|
|||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.proxy.proxy_server import (
|
||||
ProxyStartupEvent,
|
||||
_await_passthrough_error_reports_on_shutdown,
|
||||
_initialize_shared_aiohttp_session,
|
||||
_resolve_pydantic_type,
|
||||
_resolve_typed_dict_type,
|
||||
|
|
@ -1271,3 +1272,36 @@ async def test_prometheus_fallback_stats_job_runs_when_the_lock_is_free_or_absen
|
|||
await jobs["prometheus_fallback_stats_job"]()
|
||||
|
||||
assert send_fallback_stats.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_await_passthrough_error_reports_on_shutdown_waits_for_named_tasks():
|
||||
from litellm.constants import PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME
|
||||
|
||||
finished: Final = asyncio.Event()
|
||||
|
||||
async def named_report() -> None:
|
||||
await asyncio.sleep(0.05)
|
||||
finished.set()
|
||||
|
||||
async def unrelated() -> None:
|
||||
await asyncio.Event().wait()
|
||||
|
||||
named_task: Final = asyncio.get_running_loop().create_task(
|
||||
named_report(), name=PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME
|
||||
)
|
||||
stray: Final = asyncio.create_task(unrelated())
|
||||
try:
|
||||
await asyncio.wait_for(_await_passthrough_error_reports_on_shutdown(), timeout=10)
|
||||
assert named_task.done()
|
||||
assert finished.is_set()
|
||||
assert not stray.done()
|
||||
finally:
|
||||
stray.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await stray
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_await_passthrough_error_reports_on_shutdown_returns_without_pending_tasks():
|
||||
await asyncio.wait_for(_await_passthrough_error_reports_on_shutdown(), timeout=5)
|
||||
|
|
|
|||
|
|
@ -10226,3 +10226,22 @@ class TestStreamingContainerOwnershipRecordedBeforeDone:
|
|||
assert tuple(chunk for chunk, _ in observed) == self.CHUNKS
|
||||
assert tuple(count for _, count in observed) == (0, 0, 0, 0)
|
||||
recorder.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclose_late_response_runs_background_task():
|
||||
from starlette.background import BackgroundTask
|
||||
|
||||
from litellm.proxy.common_request_processing import _aclose_late_response
|
||||
|
||||
ran: list[bool] = []
|
||||
|
||||
async def body():
|
||||
yield b"x"
|
||||
|
||||
async def mark() -> None:
|
||||
ran.append(True)
|
||||
|
||||
produced: Final = StreamingResponse(body(), background=BackgroundTask(mark))
|
||||
await _aclose_late_response(produced)
|
||||
assert ran == [True]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue