mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(passthrough): run the upstream error report as a response background task
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
a614d85c92
commit
3f708fa930
9 changed files with 578 additions and 642 deletions
|
|
@ -1529,11 +1529,6 @@ 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_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
|
||||
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ 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
|
||||
|
|
@ -27,6 +26,7 @@ from fastapi import (
|
|||
status,
|
||||
)
|
||||
from fastapi.responses import StreamingResponse
|
||||
from starlette.background import BackgroundTask
|
||||
from starlette.datastructures import UploadFile as StarletteUploadFile
|
||||
from starlette.websockets import WebSocketState
|
||||
from websockets.asyncio.client import connect
|
||||
|
|
@ -43,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_DRAIN_SECONDS,
|
||||
REDACTED_BY_LITELLM,
|
||||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
|
||||
|
|
@ -882,99 +881,53 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
upstream: httpx.Response,
|
||||
report: _ReportPreview,
|
||||
log_warning: Callable[..., None],
|
||||
spawn: Callable[[Awaitable[None]], asyncio.Future[None]],
|
||||
) -> None:
|
||||
self._upstream: Final = upstream
|
||||
self._report: Final = report
|
||||
self._log_warning: Final = log_warning
|
||||
self._spawn: Final = spawn
|
||||
self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying
|
||||
self._dispatched = False
|
||||
self._pending: asyncio.Future[None] | None = None
|
||||
self._budget_crossed = False
|
||||
self._completed = False
|
||||
self._aborted = False
|
||||
self._reported = False
|
||||
|
||||
def _dispatch_report(self) -> None:
|
||||
if self._dispatched:
|
||||
async def report_collected(self) -> None:
|
||||
if self._reported:
|
||||
return
|
||||
self._dispatched = True
|
||||
self._pending = self._spawn(self._report(b"".join(self._collected)))
|
||||
|
||||
async def _drain_pending_report(self) -> None:
|
||||
pending: Final = self._pending
|
||||
if pending is not None:
|
||||
await asyncio.shield(pending)
|
||||
|
||||
def _dispatch_disconnect_report(self) -> None:
|
||||
if not self._completed and not self._dispatched:
|
||||
self._reported = True
|
||||
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",
|
||||
sum(len(part) for part in self._collected),
|
||||
)
|
||||
self._dispatch_report()
|
||||
await self._report(b"".join(self._collected))
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
total = 0 # rebind-ok: running byte count against the preview budget
|
||||
try:
|
||||
async for chunk in self._upstream.aiter_bytes():
|
||||
if not self._dispatched:
|
||||
if not self._budget_crossed:
|
||||
self._collected.append(chunk)
|
||||
total += len(chunk)
|
||||
if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
|
||||
self._dispatch_report()
|
||||
self._budget_crossed = True
|
||||
yield chunk
|
||||
self._completed = True
|
||||
self._dispatch_report()
|
||||
await self._drain_pending_report()
|
||||
except httpx.HTTPError as err:
|
||||
dispatched_before: Final = self._dispatched
|
||||
self._aborted = True
|
||||
self._log_warning(
|
||||
"pass_through_endpoint: upstream error body read failed after %d bytes: %s",
|
||||
sum(len(part) for part in self._collected),
|
||||
type(err).__name__,
|
||||
)
|
||||
self._dispatch_report()
|
||||
await self._drain_pending_report()
|
||||
if dispatched_before:
|
||||
if self._budget_crossed:
|
||||
await self.report_collected()
|
||||
raise
|
||||
finally:
|
||||
self._dispatch_disconnect_report()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self._dispatch_disconnect_report()
|
||||
await self._upstream.aclose()
|
||||
|
||||
|
||||
_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(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 | None = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
|
||||
log_warning: Callable[..., None] = verbose_proxy_logger.warning,
|
||||
) -> None:
|
||||
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:
|
||||
return httpx.Headers(
|
||||
[(name, value) for name, value in headers.raw if name.lower() not in (b"content-encoding", b"content-length")]
|
||||
|
|
@ -1033,14 +986,20 @@ def _passthrough_upstream_failure_reporter(
|
|||
return report
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _UpstreamRelay:
|
||||
response: httpx.Response
|
||||
background: BackgroundTask | None
|
||||
|
||||
|
||||
async def _log_passthrough_upstream_failure(
|
||||
response: httpx.Response,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_payload: dict,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> httpx.Response:
|
||||
) -> _UpstreamRelay:
|
||||
if response.status_code < 400:
|
||||
return response
|
||||
return _UpstreamRelay(response=response, background=None)
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
log_warning: Final = verbose_proxy_logger.warning
|
||||
|
|
@ -1049,18 +1008,21 @@ async def _log_passthrough_upstream_failure(
|
|||
)
|
||||
if response.is_stream_consumed:
|
||||
await report(response.content)
|
||||
return response
|
||||
return httpx.Response(
|
||||
status_code=response.status_code,
|
||||
headers=_headers_without_body_framing(response.headers),
|
||||
stream=_PreviewReportingStream(
|
||||
upstream=response,
|
||||
report=report,
|
||||
log_warning=log_warning,
|
||||
spawn=_spawn_report_task,
|
||||
return _UpstreamRelay(response=response, background=None)
|
||||
stream: Final = _PreviewReportingStream(
|
||||
upstream=response,
|
||||
report=report,
|
||||
log_warning=log_warning,
|
||||
)
|
||||
return _UpstreamRelay(
|
||||
response=httpx.Response(
|
||||
status_code=response.status_code,
|
||||
headers=_headers_without_body_framing(response.headers),
|
||||
stream=stream,
|
||||
request=response.request,
|
||||
extensions=response.extensions,
|
||||
),
|
||||
request=response.request,
|
||||
extensions=response.extensions,
|
||||
background=BackgroundTask(stream.report_collected),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1491,7 +1453,7 @@ async def pass_through_request(
|
|||
headers=response.headers,
|
||||
)
|
||||
|
||||
relay_response: Final = await _log_passthrough_upstream_failure(
|
||||
relay: Final = await _log_passthrough_upstream_failure(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_payload=_build_passthrough_failure_request_payload(
|
||||
|
|
@ -1506,13 +1468,13 @@ async def pass_through_request(
|
|||
|
||||
# Call response headers hook for streaming pass-through
|
||||
_response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
|
||||
headers=relay_response.headers,
|
||||
headers=relay.response.headers,
|
||||
litellm_call_id=litellm_call_id,
|
||||
)
|
||||
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
|
||||
data=_parsed_body or {},
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=relay_response,
|
||||
response=relay.response,
|
||||
request_headers=dict(request.headers),
|
||||
)
|
||||
if callback_headers:
|
||||
|
|
@ -1523,7 +1485,7 @@ async def pass_through_request(
|
|||
stream=_own_streamed_managed_ids(
|
||||
stream=_relay_reporting_failures(
|
||||
stream=PassThroughStreamingHandler.chunk_processor(
|
||||
response=relay_response,
|
||||
response=relay.response,
|
||||
request_body=_parsed_body,
|
||||
litellm_logging_obj=logging_obj,
|
||||
endpoint_type=endpoint_type,
|
||||
|
|
@ -1531,7 +1493,7 @@ async def pass_through_request(
|
|||
passthrough_success_handler_obj=pass_through_endpoint_logging,
|
||||
url_route=str(url),
|
||||
),
|
||||
upstream_status=relay_response.status_code,
|
||||
upstream_status=relay.response.status_code,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_payload=_build_passthrough_failure_request_payload(
|
||||
parsed_body=_parsed_body,
|
||||
|
|
@ -1545,10 +1507,11 @@ async def pass_through_request(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
),
|
||||
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
|
||||
upstream_headers=relay_response.headers,
|
||||
upstream_headers=relay.response.headers,
|
||||
),
|
||||
headers=_response_headers,
|
||||
status_code=relay_response.status_code,
|
||||
status_code=relay.response.status_code,
|
||||
background=relay.background,
|
||||
)
|
||||
|
||||
if state_raw_body is not None:
|
||||
|
|
@ -1583,7 +1546,7 @@ async def pass_through_request(
|
|||
logging_obj.stream = True
|
||||
logging_obj.model_call_details["stream"] = True
|
||||
|
||||
detected_relay_response: Final = await _log_passthrough_upstream_failure(
|
||||
detected_relay: Final = await _log_passthrough_upstream_failure(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_payload=_build_passthrough_failure_request_payload(
|
||||
|
|
@ -1598,13 +1561,13 @@ async def pass_through_request(
|
|||
|
||||
# Call response headers hook for detected streaming pass-through
|
||||
_response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
|
||||
headers=detected_relay_response.headers,
|
||||
headers=detected_relay.response.headers,
|
||||
litellm_call_id=litellm_call_id,
|
||||
)
|
||||
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
|
||||
data=_parsed_body or {},
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=detected_relay_response,
|
||||
response=detected_relay.response,
|
||||
request_headers=dict(request.headers),
|
||||
)
|
||||
if callback_headers:
|
||||
|
|
@ -1615,7 +1578,7 @@ async def pass_through_request(
|
|||
stream=_own_streamed_managed_ids(
|
||||
stream=_relay_reporting_failures(
|
||||
stream=PassThroughStreamingHandler.chunk_processor(
|
||||
response=detected_relay_response,
|
||||
response=detected_relay.response,
|
||||
request_body=_parsed_body,
|
||||
litellm_logging_obj=logging_obj,
|
||||
endpoint_type=endpoint_type,
|
||||
|
|
@ -1623,7 +1586,7 @@ async def pass_through_request(
|
|||
passthrough_success_handler_obj=pass_through_endpoint_logging,
|
||||
url_route=str(url),
|
||||
),
|
||||
upstream_status=detected_relay_response.status_code,
|
||||
upstream_status=detected_relay.response.status_code,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_payload=_build_passthrough_failure_request_payload(
|
||||
parsed_body=_parsed_body,
|
||||
|
|
@ -1637,10 +1600,11 @@ async def pass_through_request(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
),
|
||||
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
|
||||
upstream_headers=detected_relay_response.headers,
|
||||
upstream_headers=detected_relay.response.headers,
|
||||
),
|
||||
headers=_response_headers,
|
||||
status_code=detected_relay_response.status_code,
|
||||
status_code=detected_relay.response.status_code,
|
||||
background=detected_relay.background,
|
||||
)
|
||||
|
||||
if not _should_buffer_passthrough_response(response):
|
||||
|
|
|
|||
|
|
@ -733,7 +733,6 @@ from litellm.proxy.pass_through_endpoints.openai_passthrough_endpoints import (
|
|||
router as openai_passthrough_router,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
drain_passthrough_upstream_error_reports,
|
||||
initialize_pass_through_endpoints,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
|
|
@ -1097,27 +1096,6 @@ async def _flush_spend_logs_queue_on_shutdown() -> None:
|
|||
verbose_proxy_logger.exception("Error flushing spend logs queue on shutdown: %s", e)
|
||||
|
||||
|
||||
async def _drain_reports_then_flush_spend(
|
||||
drain_reports: Callable[[], Awaitable[None]],
|
||||
drain_spend_events: Callable[[], Awaitable[None]],
|
||||
stop_scheduler_jobs: Callable[[], Awaitable[None]] | None,
|
||||
flush_spend_counters: Callable[[], Awaitable[None]],
|
||||
flush_spend_logs: Callable[[], Awaitable[None]],
|
||||
) -> None:
|
||||
try:
|
||||
await drain_reports()
|
||||
except Exception as e: # noqa: BLE001 # shutdown must continue when a report drain fails
|
||||
verbose_proxy_logger.error("Error draining passthrough upstream error reports: %s", e)
|
||||
await drain_spend_events()
|
||||
if stop_scheduler_jobs is not None:
|
||||
try:
|
||||
await stop_scheduler_jobs()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
|
||||
await flush_spend_counters()
|
||||
await flush_spend_logs()
|
||||
|
||||
|
||||
async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = None) -> None:
|
||||
global prisma_client, master_key, user_custom_auth, user_custom_key_generate, user_custom_key_update
|
||||
verbose_proxy_logger.info("Shutting down LiteLLM Proxy Server")
|
||||
|
|
@ -1606,18 +1584,18 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e)
|
||||
|
||||
await _drain_spend_event_producer_on_shutdown()
|
||||
|
||||
# Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect
|
||||
await _drain_reports_then_flush_spend(
|
||||
drain_reports=drain_passthrough_upstream_error_reports,
|
||||
drain_spend_events=_drain_spend_event_producer_on_shutdown,
|
||||
stop_scheduler_jobs=(
|
||||
partial(stop_in_flight_scheduler_jobs, scheduler, scheduler_executor)
|
||||
if scheduler is not None and scheduler_executor is not None
|
||||
else None
|
||||
),
|
||||
flush_spend_counters=flush_spend_counters_on_shutdown,
|
||||
flush_spend_logs=_flush_spend_logs_queue_on_shutdown,
|
||||
)
|
||||
if scheduler is not None and scheduler_executor is not None:
|
||||
try:
|
||||
await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
|
||||
|
||||
await flush_spend_counters_on_shutdown()
|
||||
|
||||
await _flush_spend_logs_queue_on_shutdown()
|
||||
|
||||
await proxy_config.stop_config_sync_subscriber()
|
||||
|
||||
|
|
@ -6364,6 +6342,11 @@ class ProxyConfig:
|
|||
if general_settings is None:
|
||||
general_settings = {}
|
||||
|
||||
if general_settings.get("mcp_advertised_versions") is not None:
|
||||
from litellm.types.mcp import MCPAdvertisedVersions
|
||||
|
||||
TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"])
|
||||
|
||||
if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None:
|
||||
warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1"))
|
||||
if declared_proxy_ranges(general_settings) is None:
|
||||
|
|
|
|||
|
|
@ -89,6 +89,7 @@ def owned_proxy_process(
|
|||
config: Path | None = None,
|
||||
remove_environment: tuple[str, ...] = (),
|
||||
workers: int = 1,
|
||||
graceful_shutdown_seconds: int | None = None,
|
||||
) -> Iterator[OwnedProxy]:
|
||||
with socket.socket() as reserve:
|
||||
reserve.bind(("127.0.0.1", 0))
|
||||
|
|
@ -104,27 +105,43 @@ def owned_proxy_process(
|
|||
"LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"),
|
||||
"STORE_MODEL_IN_DB": "True",
|
||||
**overrides,
|
||||
**({"CONFIG_FILE_PATH": str(config)} if graceful_shutdown_seconds is not None and config else {}),
|
||||
}
|
||||
output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory)))
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
log_path: Final = output / f"owned-proxy-{uuid.uuid4().hex}.log"
|
||||
with log_path.open("w") as log:
|
||||
process: Final = subprocess.Popen(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"integration._support.proxy",
|
||||
"--config",
|
||||
str(config or "tests/integration/proxy_config.yaml"),
|
||||
"--host",
|
||||
"127.0.0.1",
|
||||
"--port",
|
||||
str(port),
|
||||
"--num_workers",
|
||||
str(workers),
|
||||
"--use_prisma_db_push",
|
||||
"--enforce_prisma_migration_check",
|
||||
],
|
||||
(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"uvicorn",
|
||||
"litellm.proxy.proxy_server:app",
|
||||
"--host",
|
||||
"127.0.0.1",
|
||||
"--port",
|
||||
str(port),
|
||||
"--timeout-graceful-shutdown",
|
||||
str(graceful_shutdown_seconds),
|
||||
]
|
||||
if graceful_shutdown_seconds is not None
|
||||
else [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"integration._support.proxy",
|
||||
"--config",
|
||||
str(config or "tests/integration/proxy_config.yaml"),
|
||||
"--host",
|
||||
"127.0.0.1",
|
||||
"--port",
|
||||
str(port),
|
||||
"--num_workers",
|
||||
str(workers),
|
||||
"--use_prisma_db_push",
|
||||
"--enforce_prisma_migration_check",
|
||||
]
|
||||
),
|
||||
cwd=root,
|
||||
env=environment,
|
||||
stdout=log,
|
||||
|
|
|
|||
|
|
@ -243,6 +243,73 @@ async def test_passthrough_sigterm_drains_reports_parked_on_a_slow_failure_hook(
|
|||
_single_spend_row(call_id)
|
||||
|
||||
|
||||
_PARKING_FAILURE_HOOK: Final = """
|
||||
import asyncio
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
class ParkingFailureHook(CustomLogger):
|
||||
async def async_post_call_failure_hook(
|
||||
self, request_data, original_exception, user_api_key_dict, traceback_str=None
|
||||
):
|
||||
await asyncio.Event().wait()
|
||||
|
||||
|
||||
instance = ParkingFailureHook()
|
||||
"""
|
||||
|
||||
|
||||
async def test_passthrough_sigterm_with_graceful_timeout_exits_and_flushes_spend(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
"""A streamed upstream error whose failure report can never finish must not hold
|
||||
the proxy open: uvicorn cancels the stuck request task after the graceful
|
||||
window and the shutdown spend flush still lands rows buffered behind the
|
||||
3600 s batch interval."""
|
||||
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(
|
||||
{
|
||||
"candidates": [{"content": {"parts": [{"text": "ok"}], "role": "model"}}],
|
||||
"usageMetadata": {"promptTokenCount": 3, "candidatesTokenCount": 2, "totalTokenCount": 5},
|
||||
}
|
||||
).encode(),
|
||||
)
|
||||
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["litellm_settings"].update({"callbacks": ["park_hook.instance"]})
|
||||
config["general_settings"]["proxy_batch_write_at"] = 3600
|
||||
(tmp_path / "park_hook.py").write_text(_PARKING_FAILURE_HOOK)
|
||||
path: Final = tmp_path / "chaos-graceful-sigterm.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=3) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
await _first_frame_then_close(str(candidate.client.base_url), candidate.key)
|
||||
gate.set()
|
||||
follow_up: Final = candidate.request(
|
||||
"POST",
|
||||
"/gemini/v1beta/models/claude-nope-9:generateContent",
|
||||
_GENERATE_CONTENT,
|
||||
headers={"x-goog-api-key": candidate.key},
|
||||
)
|
||||
assert follow_up.status_code == 200, follow_up.text
|
||||
call_id: Final = follow_up.headers["x-litellm-call-id"]
|
||||
await asyncio.sleep(1.5)
|
||||
owned.process.send_signal(signal.SIGTERM)
|
||||
owned.process.wait(timeout=20)
|
||||
_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(
|
||||
|
|
|
|||
|
|
@ -1133,3 +1133,117 @@ def test_llm_endpoint_upstream_quota_429_normalized_error_unchanged(
|
|||
assert error_information["normalized_error"] == _LLM_429_NORMALIZED, error_information
|
||||
finally:
|
||||
delete_scenario(handle)
|
||||
|
||||
|
||||
_TWO_SECOND_FAILURE_HOOK: Final = """
|
||||
import asyncio
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
class TwoSecondFailureHook(CustomLogger):
|
||||
async def async_post_call_failure_hook(
|
||||
self, request_data, original_exception, user_api_key_dict, traceback_str=None
|
||||
):
|
||||
await asyncio.sleep(2)
|
||||
|
||||
|
||||
instance = TwoSecondFailureHook()
|
||||
"""
|
||||
|
||||
|
||||
async def test_passthrough_keepalive_pings_never_follow_the_upstream_error_body(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
"""Keepalive pings fill idle time while the relay is still producing; once the
|
||||
upstream error body is fully relayed the response ends. If the failure
|
||||
report sat on the response path, pings emitted during the slow hook would
|
||||
land after the body's last frame."""
|
||||
frames: Final = (b'data: {"error":"one"}\n\n', b'data: {"error":"two"}\n\n')
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(status=500, content_type="text/event-stream", chunks=frames)
|
||||
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["litellm_settings"].update(
|
||||
{"callbacks": ["two_sec_hook.instance"], "sse_keepalive_ping_interval_seconds": 0.2}
|
||||
)
|
||||
(tmp_path / "two_sec_hook.py").write_text(_TWO_SECOND_FAILURE_HOOK)
|
||||
path: Final = tmp_path / "gemini-keepalive-error.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,
|
||||
headers=_gemini_headers(candidate),
|
||||
) as response:
|
||||
assert response.status_code == 500, response.text
|
||||
call_id: Final = response.headers["x-litellm-call-id"]
|
||||
streamed: Final = response.read()
|
||||
assert streamed == b"".join(frames), streamed
|
||||
_spend_error_information(call_id)
|
||||
|
||||
|
||||
_HEADER_STATE_HOOK: Final = """
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
class HeaderStateHook(CustomLogger):
|
||||
def __init__(self):
|
||||
self.last_status = None
|
||||
|
||||
async def async_post_call_failure_hook(
|
||||
self, request_data, original_exception, user_api_key_dict, traceback_str=None
|
||||
):
|
||||
self.last_status = getattr(original_exception, "status_code", None)
|
||||
|
||||
async def async_post_call_response_headers_hook(
|
||||
self, data, user_api_key_dict, response, request_headers=None, litellm_call_info=None
|
||||
):
|
||||
return {"x-failure-for-this-request": str(self.last_status or "none")}
|
||||
|
||||
|
||||
instance = HeaderStateHook()
|
||||
"""
|
||||
|
||||
|
||||
async def test_passthrough_streamed_error_headers_do_not_carry_failure_hook_state(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
"""The streamed error report runs after response headers are sent, so a
|
||||
callback-derived header sees pre-request state; the failure hook still
|
||||
records the upstream status in the spend row."""
|
||||
frames: Final = (b'data: {"error":"quota"}\n\n',)
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(status=500, content_type="text/event-stream", chunks=frames)
|
||||
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["litellm_settings"].update({"callbacks": ["header_state_hook.instance"]})
|
||||
(tmp_path / "header_state_hook.py").write_text(_HEADER_STATE_HOOK)
|
||||
path: Final = tmp_path / "gemini-header-state.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,
|
||||
headers=_gemini_headers(candidate),
|
||||
) as response:
|
||||
assert response.status_code == 500, response.text
|
||||
call_id: Final = response.headers["x-litellm-call-id"]
|
||||
header_value: Final = response.headers["x-failure-for-this-request"]
|
||||
streamed: Final = response.read()
|
||||
assert streamed == b"".join(frames), streamed
|
||||
assert header_value == "none", header_value
|
||||
error_information: Final = _spend_error_information(call_id)
|
||||
assert error_information["error_code"] == "500", error_information
|
||||
|
|
|
|||
|
|
@ -1,11 +1,9 @@
|
|||
import asyncio
|
||||
import gc
|
||||
import gzip
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import zlib
|
||||
from collections.abc import Callable, Mapping
|
||||
from contextlib import ExitStack, contextmanager
|
||||
|
|
@ -20,6 +18,7 @@ import pytest
|
|||
from fastapi import HTTPException, Request, Response, UploadFile
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from starlette.background import BackgroundTask
|
||||
from starlette.datastructures import FormData, Headers, QueryParams
|
||||
from starlette.datastructures import UploadFile as StarletteUploadFile
|
||||
|
||||
|
|
@ -29,19 +28,17 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
_REPORT_TASKS,
|
||||
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS,
|
||||
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
|
||||
HttpPassThroughEndpointHelpers,
|
||||
InitPassThroughEndpointHelpers,
|
||||
_PreviewReportingStream,
|
||||
_registered_pass_through_routes,
|
||||
_spawn_report_task,
|
||||
_truncate_upstream_error_body,
|
||||
_with_trace_context,
|
||||
chat_completion_pass_through_endpoint,
|
||||
create_pass_through_route,
|
||||
drain_passthrough_upstream_error_reports,
|
||||
_log_passthrough_upstream_failure,
|
||||
initialize_pass_through_endpoints,
|
||||
pass_through_request,
|
||||
resolve_llm_passthrough_timeout,
|
||||
|
|
@ -4266,7 +4263,7 @@ async def test_pass_through_request_streaming_upstream_error_body_reaches_client
|
|||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
|
||||
await _poll(lambda: mock_proxy_logging.post_call_failure_hook.called)
|
||||
await response.background()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
original_exception: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"]
|
||||
assert "was not found or your project" in original_exception.detail
|
||||
|
|
@ -4462,42 +4459,30 @@ async def test_pass_through_request_streaming_upstream_error_reads_only_preview_
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
served_at_dispatch: list[int] = []
|
||||
|
||||
def _recording_spawn(coro):
|
||||
served_at_dispatch.append(body_stream.served)
|
||||
enqueued.append(asyncio.ensure_future(coro))
|
||||
return enqueued[-1]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=_recording_spawn,
|
||||
):
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 500
|
||||
|
|
@ -4506,12 +4491,12 @@ async def test_pass_through_request_streaming_upstream_error_reads_only_preview_
|
|||
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
|
||||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
assert body_stream.served == 10
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
|
||||
assert served_at_dispatch == [5], (
|
||||
"the report is dispatched after the fifth 1024-byte chunk, the first point the preview budget is exceeded"
|
||||
)
|
||||
assert len(enqueued) == 1, enqueued
|
||||
await enqueued[0]
|
||||
assert response.background is not None
|
||||
await response.background()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
|
||||
assert (
|
||||
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
|
||||
|
|
@ -4532,42 +4517,30 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
served_at_dispatch: list[int] = []
|
||||
|
||||
def _recording_spawn(coro):
|
||||
served_at_dispatch.append(body_stream.served)
|
||||
enqueued.append(asyncio.ensure_future(coro))
|
||||
return enqueued[-1]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=_recording_spawn,
|
||||
):
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 500
|
||||
|
|
@ -4576,12 +4549,11 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
|
|||
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
|
||||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
|
||||
assert served_at_dispatch == [1], (
|
||||
"the report is dispatched after the first raw chunk crosses the preview budget; the second is not pulled first"
|
||||
)
|
||||
assert len(enqueued) == 1, enqueued
|
||||
await enqueued[0]
|
||||
assert response.background is not None
|
||||
await response.background()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
|
||||
assert (
|
||||
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
|
||||
|
|
@ -4589,13 +4561,6 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
|
|||
)
|
||||
|
||||
|
||||
async def _poll(condition: Callable[[], bool], seconds: float = 5) -> None:
|
||||
deadline: Final = asyncio.get_running_loop().time() + seconds
|
||||
while not condition():
|
||||
assert asyncio.get_running_loop().time() < deadline, "condition not met in time"
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
|
||||
class _GatedUpstreamErrorBodyStream(httpx.AsyncByteStream):
|
||||
def __init__(self, first: bytes, second: bytes) -> None:
|
||||
self._first: Final = first
|
||||
|
|
@ -4625,8 +4590,6 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
|
|
@ -4634,31 +4597,27 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
|
|||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
|
||||
):
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 429
|
||||
iterator: Final = response.body_iterator.__aiter__()
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 429
|
||||
iterator: Final = response.body_iterator.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert not body_stream.gate.is_set()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
|
|
@ -4670,8 +4629,9 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
|
|||
assert relayed == first_chunk + second_chunk
|
||||
await upstream_response.aclose()
|
||||
|
||||
assert len(enqueued) == 1, enqueued
|
||||
await enqueued[0]
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
assert response.background is not None
|
||||
await response.background()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
|
||||
assert (
|
||||
|
|
@ -4680,12 +4640,11 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_streaming_upstream_error_client_disconnect_enqueues_failure_report():
|
||||
async def test_pass_through_request_streaming_upstream_error_client_disconnect_reports_relayed_bytes():
|
||||
"""
|
||||
Regression: when the client disconnects mid-relay the response task is
|
||||
cancelled, so the preview report cannot be awaited inline; it must be
|
||||
handed to the logging worker, which then fires the failure hook once
|
||||
with the chunks already relayed.
|
||||
Regression: when the client disconnects mid-relay the response ends early,
|
||||
so the report runs as the response background task; it must fire the
|
||||
failure hook once with only the chunks already relayed.
|
||||
"""
|
||||
first_chunk: Final = b'data: {"error":"rate limited"}\n\n'
|
||||
second_chunk: Final = b"data: [DONE]\n\n"
|
||||
|
|
@ -4697,8 +4656,6 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
|
|
@ -4706,42 +4663,33 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
|
|||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
|
||||
):
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
iterator = response.body_iterator.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert first_chunk in (first if isinstance(first, bytes) else first.encode())
|
||||
await iterator.aclose()
|
||||
del iterator
|
||||
gc.collect()
|
||||
for _ in range(40):
|
||||
if enqueued:
|
||||
break
|
||||
await asyncio.sleep(0.05)
|
||||
assert isinstance(response, StreamingResponse)
|
||||
iterator = response.body_iterator.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert first_chunk in (first if isinstance(first, bytes) else first.encode())
|
||||
await iterator.aclose()
|
||||
|
||||
assert len(enqueued) == 1, enqueued
|
||||
await enqueued[0]
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
assert response.background is not None
|
||||
await response.background()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
|
||||
assert detail == 'Upstream passthrough request failed with status 429: data: {"error":"rate limited"}', detail
|
||||
|
|
@ -4750,9 +4698,9 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
|
|||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_streaming_upstream_error_yields_over_budget_chunk_before_report_finishes():
|
||||
"""
|
||||
Regression: the failure report is never awaited inside the relay, so a
|
||||
Regression: the failure report never sits inside the relay, so a
|
||||
single chunk that crosses the preview budget reaches the client even
|
||||
while the report coroutine is still running.
|
||||
while the report would still be running.
|
||||
"""
|
||||
first_chunk: Final = b"x" * 6144
|
||||
release: Final = asyncio.Event()
|
||||
|
|
@ -4764,8 +4712,6 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
|
||||
async def _held_hook(**kwargs):
|
||||
await release.wait()
|
||||
|
||||
|
|
@ -4776,43 +4722,40 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
|
|||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
|
||||
):
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_held_hook)
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_held_hook)
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 500
|
||||
iterator: Final = response.body_iterator.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert not release.is_set()
|
||||
assert not enqueued[0].done()
|
||||
release.set()
|
||||
relayed: Final = b"".join(
|
||||
[first]
|
||||
+ [chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") async for chunk in iterator]
|
||||
)
|
||||
assert relayed == first_chunk
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 500
|
||||
iterator: Final = response.body_iterator.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert not release.is_set()
|
||||
relayed: Final = b"".join(
|
||||
[first]
|
||||
+ [chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") async for chunk in iterator]
|
||||
)
|
||||
assert relayed == first_chunk
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
background: Final = asyncio.ensure_future(response.background())
|
||||
assert not background.done()
|
||||
release.set()
|
||||
await background
|
||||
|
||||
assert len(enqueued) == 1, enqueued
|
||||
await enqueued[0]
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
|
||||
assert (
|
||||
|
|
@ -4821,7 +4764,7 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_streaming_upstream_error_drains_the_report_before_stream_end():
|
||||
async def test_pass_through_request_streaming_upstream_error_reports_via_background_task_after_stream_end():
|
||||
upstream_response: Final = httpx.Response(
|
||||
status_code=500,
|
||||
headers={"content-type": "text/plain"},
|
||||
|
|
@ -4858,6 +4801,9 @@ async def test_pass_through_request_streaming_upstream_error_drains_the_report_b
|
|||
assert response.status_code == 500
|
||||
drained: Final = [chunk async for chunk in response.body_iterator]
|
||||
assert drained == [b"x" * 512]
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
assert response.background is not None
|
||||
await response.background()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
|
||||
|
||||
|
|
@ -4924,11 +4870,7 @@ async def test_pass_through_request_streaming_upstream_error_body_read_failure_k
|
|||
assert streamed_bytes == b'{"error": "half'
|
||||
await upstream_response.aclose()
|
||||
|
||||
await _poll(
|
||||
lambda: any(
|
||||
args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" for args in recorded_warnings
|
||||
)
|
||||
)
|
||||
await response.background()
|
||||
rendered: Final = [str(args[0]) for args in recorded_warnings]
|
||||
formats: Final = [args[0] for args in recorded_warnings]
|
||||
assert any(
|
||||
|
|
@ -4974,40 +4916,30 @@ async def test_pass_through_request_streaming_upstream_abort_after_preview_budge
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
|
||||
def _recording_spawn(coro):
|
||||
enqueued.append(asyncio.ensure_future(coro))
|
||||
return enqueued[-1]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=_recording_spawn,
|
||||
):
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
received: list[bytes] = []
|
||||
|
|
@ -5020,18 +4952,21 @@ async def test_pass_through_request_streaming_upstream_abort_after_preview_budge
|
|||
await consume_response()
|
||||
assert b"".join(received) == b"d" * 5000
|
||||
|
||||
assert len(enqueued) == 1, enqueued
|
||||
await enqueued[0]
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
expected_body: Final = f"{'d' * 4096}... (truncated at 4096 chars)"
|
||||
assert (
|
||||
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
|
||||
== f"Upstream passthrough request failed with status 500: {expected_body}"
|
||||
)
|
||||
assert response.background is not None
|
||||
await response.background()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
|
||||
async def test_preview_report_collected_reports_once_and_warns_on_disconnect():
|
||||
"""Closing the relay early leaves the preview partial: report_collected logs the
|
||||
disconnect warning once and reports the collected bytes, and a second call is a no-op."""
|
||||
upstream_hold: Final = asyncio.Event()
|
||||
chunks: Final = (b"first",)
|
||||
upstream_response: Final = httpx.Response(
|
||||
|
|
@ -5040,189 +4975,171 @@ async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
|
|||
stream=_UpstreamErrorBodyStreamHeld(chunks, upstream_hold),
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
release: Final = asyncio.Event()
|
||||
hook_done: Final = asyncio.Event()
|
||||
reported: list[bytes] = []
|
||||
log_warning: Final = MagicMock()
|
||||
|
||||
async def report(preview: bytes) -> None:
|
||||
await release.wait()
|
||||
hook_done.set()
|
||||
reported.append(preview)
|
||||
|
||||
relay: Final = _PreviewReportingStream(
|
||||
upstream=upstream_response,
|
||||
report=report,
|
||||
log_warning=log_warning,
|
||||
spawn=_spawn_report_task,
|
||||
)
|
||||
|
||||
async def consume() -> None:
|
||||
async for _ in relay.__aiter__():
|
||||
pass
|
||||
iterator: Final = relay.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert first == b"first"
|
||||
await iterator.aclose()
|
||||
|
||||
consumer: Final = asyncio.ensure_future(consume())
|
||||
for _ in range(20):
|
||||
await asyncio.sleep(0)
|
||||
consumer.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await consumer
|
||||
|
||||
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()
|
||||
await relay.report_collected()
|
||||
log_warning.assert_called_once_with(
|
||||
"pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body", 5
|
||||
)
|
||||
assert reported == [b"first"], reported
|
||||
|
||||
release.set()
|
||||
await drain_passthrough_upstream_error_reports(timeout=5)
|
||||
assert hook_done.is_set()
|
||||
assert pending[0].done()
|
||||
assert not _REPORT_TASKS.get(asyncio.get_running_loop())
|
||||
await relay.report_collected()
|
||||
assert reported == [b"first"], reported
|
||||
log_warning.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_drain_returns_with_report_still_pending_on_timeout():
|
||||
upstream_hold: Final = asyncio.Event()
|
||||
chunks: Final = (b"first",)
|
||||
async def test_preview_report_collected_runs_without_disconnect_warning_after_clean_end():
|
||||
upstream_response: Final = httpx.Response(
|
||||
status_code=500,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=_UpstreamErrorBodyStreamHeld(chunks, upstream_hold),
|
||||
stream=_UpstreamErrorBodyStream(b"frame"),
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
release: Final = asyncio.Event()
|
||||
reported: list[bytes] = []
|
||||
log_warning: Final = MagicMock()
|
||||
|
||||
async def report(preview: bytes) -> None:
|
||||
await release.wait()
|
||||
reported.append(preview)
|
||||
|
||||
relay: Final = _PreviewReportingStream(
|
||||
upstream=upstream_response,
|
||||
report=report,
|
||||
log_warning=MagicMock(),
|
||||
spawn=_spawn_report_task,
|
||||
log_warning=log_warning,
|
||||
)
|
||||
|
||||
async def consume() -> None:
|
||||
async for _ in relay.__aiter__():
|
||||
pass
|
||||
|
||||
consumer: Final = asyncio.ensure_future(consume())
|
||||
for _ in range(20):
|
||||
await asyncio.sleep(0)
|
||||
consumer.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await consumer
|
||||
|
||||
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
|
||||
assert len(pending) == 1, pending
|
||||
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_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()
|
||||
|
||||
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
|
||||
assert received == b"frame"
|
||||
assert reported == []
|
||||
|
||||
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()
|
||||
|
||||
spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70))
|
||||
assert len(_REPORT_TASKS.get(loop, ())) == 70
|
||||
|
||||
release.set()
|
||||
await drain_passthrough_upstream_error_reports()
|
||||
assert all(task.done() for task in spawned)
|
||||
assert not _REPORT_TASKS.get(loop)
|
||||
|
||||
asyncio.run(run_once())
|
||||
asyncio.run(run_once())
|
||||
assert all(len(pending) == 0 for pending in _REPORT_TASKS.values())
|
||||
await relay.report_collected()
|
||||
assert reported == [b"frame"], reported
|
||||
log_warning.assert_not_called()
|
||||
|
||||
|
||||
@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
|
||||
async def test_log_passthrough_upstream_failure_returns_background_task_for_unconsumed_error_stream():
|
||||
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"),
|
||||
)
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj"):
|
||||
relay: Final = await _log_passthrough_upstream_failure(
|
||||
response=upstream_response,
|
||||
user_api_key_dict=MagicMock(),
|
||||
request_payload={},
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert isinstance(relay.background, BackgroundTask)
|
||||
assert relay.background.func == relay.response.stream.report_collected
|
||||
|
||||
ok_response: Final = httpx.Response(
|
||||
status_code=200,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=_UpstreamErrorBodyStream(b"ok"),
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
ok_relay: Final = await _log_passthrough_upstream_failure(
|
||||
response=ok_response,
|
||||
user_api_key_dict=MagicMock(),
|
||||
request_payload={},
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
assert ok_relay.response is ok_response
|
||||
assert ok_relay.background is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streamed_error_runs_failure_hook_as_background_after_response_headers_hook():
|
||||
"""The streamed error report runs via the response's background task, so the
|
||||
failure hook must fire after the response headers hook and after the last
|
||||
body byte, still exactly once."""
|
||||
upstream_response: Final = httpx.Response(
|
||||
status_code=500,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=_ChunkedUpstreamErrorBodyStream((b"data: one\n\n", b"data: two\n\n")),
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
calls: list[str] = []
|
||||
|
||||
async def _headers_hook(**kwargs):
|
||||
calls.append("post_call_response_headers_hook")
|
||||
return None
|
||||
|
||||
async def _failure_hook(**kwargs):
|
||||
calls.append("post_call_failure_hook")
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
|
||||
) as mock_success_handler:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_failure_hook)
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(side_effect=_headers_hook)
|
||||
mock_success_handler.return_value = None
|
||||
|
||||
async_client: Final = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
response: Final = await pass_through_request(
|
||||
request=_upstream_error_request(),
|
||||
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
scope: Final = {
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0", "spec_version": "2.3"},
|
||||
"http_version": "1.1",
|
||||
"method": "POST",
|
||||
"scheme": "http",
|
||||
"path": "/",
|
||||
"query_string": b"",
|
||||
"root_path": "",
|
||||
"headers": [],
|
||||
"client": ("127.0.0.1", 1234),
|
||||
"server": ("127.0.0.1", 80),
|
||||
}
|
||||
sent: list = []
|
||||
|
||||
async def receive():
|
||||
await asyncio.sleep(3600)
|
||||
|
||||
async def send(message):
|
||||
sent.append(message)
|
||||
|
||||
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"
|
||||
)
|
||||
assert body == b"data: one\n\ndata: two\n\n"
|
||||
|
||||
|
||||
class _UpstreamErrorGzipStreamDropping(httpx.AsyncByteStream):
|
||||
|
|
@ -5295,11 +5212,7 @@ async def test_pass_through_request_streaming_upstream_error_gzip_read_failure_r
|
|||
assert streamed_bytes == plaintext
|
||||
await upstream_response.aclose()
|
||||
|
||||
await _poll(
|
||||
lambda: any(
|
||||
args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" for args in recorded_warnings
|
||||
)
|
||||
)
|
||||
await response.background()
|
||||
rendered: Final = [str(args[0]) for args in recorded_warnings]
|
||||
assert any(
|
||||
args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" and plaintext.decode() in str(args[4])
|
||||
|
|
@ -5351,12 +5264,7 @@ async def test_pass_through_request_streaming_upstream_error_gzip_body_decoded_f
|
|||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
|
||||
await _poll(
|
||||
lambda: any(
|
||||
call.args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s"
|
||||
for call in mock_warning.call_args_list
|
||||
)
|
||||
)
|
||||
await response.background()
|
||||
upstream_warnings: Final = [
|
||||
call
|
||||
for call in mock_warning.call_args_list
|
||||
|
|
|
|||
|
|
@ -343,96 +343,6 @@ async def test_flush_spend_counters_on_shutdown_logs_and_swallows_commit_errors(
|
|||
assert "Error flushing spend counters on shutdown: db gone" in caplog.text
|
||||
|
||||
|
||||
def _recorded_step(calls: List[str], name: str, report_done: "asyncio.Event") -> Callable[[], Awaitable[None]]:
|
||||
async def _step() -> None:
|
||||
calls.append(f"{name}:report_done={report_done.is_set()}")
|
||||
|
||||
return _step
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_flush_runs_after_in_flight_passthrough_error_reports():
|
||||
"""A passthrough error report callback writes its spend row through the
|
||||
logging worker, so every spend flush must run only after in-flight reports
|
||||
have finished: a flush that ran while the report was still pending could
|
||||
skip its row.
|
||||
"""
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
_spawn_report_task,
|
||||
drain_passthrough_upstream_error_reports,
|
||||
)
|
||||
|
||||
report_done = asyncio.Event()
|
||||
calls: List[str] = []
|
||||
|
||||
async def report() -> None:
|
||||
await asyncio.sleep(0.05)
|
||||
report_done.set()
|
||||
|
||||
_spawn_report_task(report())
|
||||
await ps._drain_reports_then_flush_spend(
|
||||
drain_reports=drain_passthrough_upstream_error_reports,
|
||||
drain_spend_events=_recorded_step(calls, "drain_spend_events", report_done),
|
||||
stop_scheduler_jobs=_recorded_step(calls, "stop_scheduler_jobs", report_done),
|
||||
flush_spend_counters=_recorded_step(calls, "flush_spend_counters", report_done),
|
||||
flush_spend_logs=_recorded_step(calls, "flush_spend_logs", report_done),
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
"drain_spend_events:report_done=True",
|
||||
"stop_scheduler_jobs:report_done=True",
|
||||
"flush_spend_counters:report_done=True",
|
||||
"flush_spend_logs:report_done=True",
|
||||
], calls
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_flush_continues_when_report_drain_fails():
|
||||
report_done = asyncio.Event()
|
||||
calls: List[str] = []
|
||||
|
||||
async def failing_drain() -> None:
|
||||
raise RuntimeError("drain gone")
|
||||
|
||||
await ps._drain_reports_then_flush_spend(
|
||||
drain_reports=failing_drain,
|
||||
drain_spend_events=_recorded_step(calls, "drain_spend_events", report_done),
|
||||
stop_scheduler_jobs=_recorded_step(calls, "stop_scheduler_jobs", report_done),
|
||||
flush_spend_counters=_recorded_step(calls, "flush_spend_counters", report_done),
|
||||
flush_spend_logs=_recorded_step(calls, "flush_spend_logs", report_done),
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
"drain_spend_events:report_done=False",
|
||||
"stop_scheduler_jobs:report_done=False",
|
||||
"flush_spend_counters:report_done=False",
|
||||
"flush_spend_logs:report_done=False",
|
||||
], calls
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_flush_skips_scheduler_step_when_no_scheduler():
|
||||
report_done = asyncio.Event()
|
||||
calls: List[str] = []
|
||||
|
||||
async def noop_drain() -> None:
|
||||
report_done.set()
|
||||
|
||||
await ps._drain_reports_then_flush_spend(
|
||||
drain_reports=noop_drain,
|
||||
drain_spend_events=_recorded_step(calls, "drain_spend_events", report_done),
|
||||
stop_scheduler_jobs=None,
|
||||
flush_spend_counters=_recorded_step(calls, "flush_spend_counters", report_done),
|
||||
flush_spend_logs=_recorded_step(calls, "flush_spend_logs", report_done),
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
"drain_spend_events:report_done=True",
|
||||
"flush_spend_counters:report_done=True",
|
||||
"flush_spend_logs:report_done=True",
|
||||
], calls
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _initialize_shared_aiohttp_session
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -68,25 +68,3 @@ def _build_constant_env_var_map() -> dict[str, str]:
|
|||
env_var_map[constant_name] = env_var_name
|
||||
|
||||
return env_var_map
|
||||
|
||||
|
||||
@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_DRAIN_SECONDS == expected
|
||||
finally:
|
||||
monkeypatch.undo()
|
||||
importlib.reload(constants)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue