From bf9f59d76b0a917858b2f618676a7a7240029fac Mon Sep 17 00:00:00 2001 From: RachelHuangZW Date: Wed, 7 Oct 2026 20:34:38 -0400 Subject: [PATCH] fix(scheduler): remove a request's queue entry once it stops waiting (#43061) * fix(scheduler): remove a request's queue entry once it stops waiting Requests admitted while a healthy deployment existed never left the priority queue, and neither did requests cancelled while waiting. The stale entries blocked later requests during cooldown and, with Redis, made add_request raise a TypeError on queues read back as JSON lists Both scheduling paths now share one polling helper that removes the entry in a finally block, whether the request was admitted, timed out or cancelled Related to #43059 * test(router): allowlist _wait_for_scheduler_turn in the router coverage check The coverage script only counts direct calls in test files. The helper is exercised through prioritized acompletion and atext_completion in test_router.py, like the other allowlisted entries * fix(scheduler): admit healthy requests before reading the queue, and clean up cancelled enqueues poll() raised on an empty queue before checking for healthy deployments. With the cleanup now rewriting the queue after every admission, a concurrent write from another replica can erase a waiting request's entry, and that request then failed while a deployment was healthy. poll() now admits as soon as a deployment is healthy and only reads the queue during cooldown add_request also moved inside the try block, so a request cancelled while its queue write is in flight still has its entry removed * refactor(scheduler): move the wait loop into Scheduler.wait_for_turn The router passes a healthy-deployments callable into the scheduler, so tests inject a Scheduler directly instead of replacing the router's scheduler attribute. The cancelled-mid-enqueue test moves to test_scheduler.py with the other scheduler tests, and poll() takes the deployments as a Sequence since it only checks whether any are healthy * test(scheduler): move scheduler tests into tests/unit * fix(scheduler): finish the queue removal when cancellation is delivered again during cleanup A second cancel, or an anyio cancel scope that re-cancels on every await, interrupted remove_request mid-write and left the entry in Redis. * fix(scheduler): read queue entries back from redis as tuples * fix(scheduler): admit a request whose queue entry vanished during cooldown * fix(scheduler): re-enqueue a request whose queue entry vanished during cooldown * test(integration): cover priority scheduler queue cleanup across instances Adds tests/integration/routing/test_priority_scheduler_queue_cleanup.py: 27 cells against the real proxy (two workers, Postgres, Redis, a scripted upstream) for every prioritized surface (/v1/chat/completions, /v1/completions, /queue/chat/completions, streaming and not, OpenAI SDK sync and async, raw httpx), the non-integer priority pass-through, the in-memory queue on one proxy, two proxies sharing a Redis queue (served requests leave no entry, a dead replica's entry is skipped or expires, a waiter that times out or disconnects removes only itself), a Redis outage mid burst and a SIGKILLed worker. Each cell asserts the caller's response, the upstream's requests by marker and the Redis queue contents. On the merge base the cross-instance cells fail with the list-of-lists TypeError and the in-memory cell with 408s behind the leaked entry; on this branch every cell passes --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/router.py | 143 +-- litellm/scheduler.py | 84 +- .../router_code_coverage.py | 1 + .../test_priority_scheduler_queue_cleanup.py | 860 ++++++++++++++++++ tests/unit/test_router/test_router.py | 59 ++ tests/unit/test_scheduler.py | 193 +++- 6 files changed, 1203 insertions(+), 137 deletions(-) create mode 100644 tests/integration/routing/test_priority_scheduler_queue_cleanup.py diff --git a/litellm/router.py b/litellm/router.py index 7b89d5eb5ac..d6730db8748 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4534,61 +4534,21 @@ class Router: stream=False, **kwargs, ): - parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) - ### FLOW ITEM ### - _request_id: Final = str(uuid.uuid4()) - item: Final = FlowItem( - priority=priority, # 👈 SET PRIORITY FOR REQUEST - request_id=_request_id, # 👈 SET REQUEST ID - model_name=model, # 👈 SAME as 'Router' + await self._wait_for_scheduler_turn( + model=model, priority=priority, parent_otel_span=get_parent_otel_span_from_kwargs(kwargs) ) - ### [fin] ### - - ## ADDS REQUEST TO QUEUE ## - await self.scheduler.add_request(request=item) - - ## POLL QUEUE - end_time: Final = time.monotonic() + self.timeout - curr_time = time.monotonic() - poll_interval: Final = self.scheduler.polling_interval # poll every 3ms - make_request = False - - while curr_time < end_time: - _healthy_deployments, _ = await self._async_get_healthy_deployments( - model=model, parent_otel_span=parent_otel_span - ) - make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue - id=item.request_id, - model_name=item.model_name, - health_deployments=_healthy_deployments, - ) - if make_request: ## IF TRUE -> MAKE REQUEST - break - else: ## ELSE -> loop till default_timeout - await asyncio.sleep(poll_interval) - curr_time = time.monotonic() - - if make_request: - try: - _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) - response_hidden_params: Final = get_hidden_params(_response) - if response_hidden_params is not None: - additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping - dict[str, object], response_hidden_params.setdefault("additional_headers", {}) - ) - additional_headers.update({"x-litellm-request-prioritization-used": True}) - return _response - except Exception as e: - setattr(e, "priority", priority) - raise e - else: - # Clean up the request from the scheduler queue also before raising the timeout exception - await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) - raise litellm.Timeout( - message="Request timed out while polling queue", - model=model, - llm_provider="openai", - ) + try: + _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) + response_hidden_params: Final = get_hidden_params(_response) + if response_hidden_params is not None: + additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping + dict[str, object], response_hidden_params.setdefault("additional_headers", {}) + ) + additional_headers.update({"x-litellm-request-prioritization-used": True}) + return _response + except Exception as e: + setattr(e, "priority", priority) + raise e async def _schedule_factory( self, @@ -4598,61 +4558,32 @@ class Router: args: tuple[object, ...], kwargs: dict[str, object], ): - parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) - ### FLOW ITEM ### - _request_id: Final = str(uuid.uuid4()) - item: Final = FlowItem( - priority=priority, # 👈 SET PRIORITY FOR REQUEST - request_id=_request_id, # 👈 SET REQUEST ID - model_name=model, # 👈 SAME as 'Router' + await self._wait_for_scheduler_turn( + model=model, priority=priority, parent_otel_span=get_parent_otel_span_from_kwargs(kwargs) ) - ### [fin] ### + try: + _response: Final = await original_function(*args, **kwargs) + response_hidden_params: Final = get_hidden_params(_response) + if response_hidden_params is not None: + additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping + dict[str, object], response_hidden_params.setdefault("additional_headers", {}) + ) + additional_headers.update({"x-litellm-request-prioritization-used": True}) + return _response + except Exception as e: + setattr(e, "priority", priority) + raise e - ## ADDS REQUEST TO QUEUE ## - await self.scheduler.add_request(request=item) + async def _wait_for_scheduler_turn(self, model: str, priority: int, parent_otel_span: Span | None) -> None: + async def healthy_deployments() -> Sequence[object]: + deployments, _ = await self._async_get_healthy_deployments(model=model, parent_otel_span=parent_otel_span) + return deployments - ## POLL QUEUE - end_time: Final = time.monotonic() + self.timeout - curr_time = time.monotonic() - poll_interval: Final = self.scheduler.polling_interval # poll every 3ms - make_request = False - - while curr_time < end_time: - _healthy_deployments, _ = await self._async_get_healthy_deployments( - model=model, parent_otel_span=parent_otel_span - ) - make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue - id=item.request_id, - model_name=item.model_name, - health_deployments=_healthy_deployments, - ) - if make_request: ## IF TRUE -> MAKE REQUEST - break - else: ## ELSE -> loop till default_timeout - await asyncio.sleep(poll_interval) - curr_time = time.monotonic() - - if make_request: - try: - _response: Final = await original_function(*args, **kwargs) - response_hidden_params: Final = get_hidden_params(_response) - if response_hidden_params is not None: - additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping - dict[str, object], response_hidden_params.setdefault("additional_headers", {}) - ) - additional_headers.update({"x-litellm-request-prioritization-used": True}) - return _response - except Exception as e: - setattr(e, "priority", priority) - raise e - else: - # Clean up the request from the scheduler queue also before raising the timeout exception - await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) - raise litellm.Timeout( - message="Request timed out while polling queue", - model=model, - llm_provider="openai", - ) + await self.scheduler.wait_for_turn( + request=FlowItem(priority=priority, request_id=str(uuid.uuid4()), model_name=model), + timeout=self.timeout, + get_healthy_deployments=healthy_deployments, + ) def _is_prompt_management_model(self, model: str) -> bool: model_list: Final = self.get_model_list(model_name=model) diff --git a/litellm/scheduler.py b/litellm/scheduler.py index c88cc4ce0c0..73196d24e4d 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -1,14 +1,22 @@ +import asyncio import enum import heapq -from typing import Final +import time +from collections.abc import Awaitable, Callable, Sequence +from typing import Final, TypeAlias + +from pydantic import TypeAdapter from litellm import print_verbose from litellm._internal_context import with_service_target from litellm.caching.caching import DualCache, RedisCache from litellm.constants import DEFAULT_IN_MEMORY_TTL, DEFAULT_POLLING_INTERVAL +from litellm.exceptions import Timeout from litellm.types.llms.base import LiteLLMBaseModel SCHEDULER_QUEUE_TARGET: Final = "scheduler_queue" +QueueEntry: TypeAlias = tuple[int, str] +_QUEUE_ENTRIES: Final = TypeAdapter(list[QueueEntry]) class SchedulerCacheKeys(enum.Enum): @@ -33,7 +41,7 @@ class Scheduler: """ polling_interval: float or null - frequency of polling queue. Default is 3ms. """ - self.queue: list = [] + self.queue: list[QueueEntry] = [] default_in_memory_ttl: float | None = None if redis_cache is not None: # if redis-cache available frequently poll that instead of using in-memory. @@ -51,7 +59,7 @@ class Scheduler: # save the queue await self.save_queue(queue=queue, model_name=request.model_name) - async def poll(self, id: str, model_name: str, health_deployments: list) -> bool: + async def poll(self, request: FlowItem, health_deployments: Sequence[object]) -> bool: """ Return if request can be processed. @@ -62,30 +70,48 @@ class Scheduler: - False: * If no healthy deployments available * AND request not at the top of queue + + A request the queue no longer holds (its cache key expired or a concurrent writer erased the entry) + is put back at its priority so it keeps its place in the order instead of failing or jumping ahead """ - queue: Final = await self.get_queue(model_name=model_name) - if not queue: - raise Exception(f"Incorrectly setup. Queue is invalid. Queue={queue}") - - # ------------ - # Setup values - # ------------ - print_verbose(f"len(health_deployments): {len(health_deployments)}") - if len(health_deployments) == 0: - print_verbose(f"queue: {queue}, seeking id={id}") - # Check if the id is at the top of the heap - if queue[0][1] == id: - # Remove the item from the queue - heapq.heappop(queue) - await self.save_queue(queue=queue, model_name=model_name) - print_verbose(f"Popped id: {id}") - return True - else: - return False + if len(health_deployments) > 0: + return True + queue: Final = await self.get_queue(model_name=request.model_name) + entry: Final = (request.priority, request.request_id) + print_verbose(f"queue: {queue}, seeking {entry}") + if entry not in queue: + print_verbose(f"queue no longer holds {entry}, re-enqueueing it") + heapq.heappush(queue, entry) + if queue[0] != entry: + await self.save_queue(queue=queue, model_name=request.model_name) + return False + elif queue[0] != entry: + return False + + heapq.heappop(queue) + await self.save_queue(queue=queue, model_name=request.model_name) + print_verbose(f"Popped id: {request.request_id}") return True + async def wait_for_turn( + self, + request: FlowItem, + timeout: float, + get_healthy_deployments: Callable[[], Awaitable[Sequence[object]]], + ) -> None: + try: + await self.add_request(request=request) + end_time: Final = time.monotonic() + timeout + while time.monotonic() < end_time: + if await self.poll(request=request, health_deployments=await get_healthy_deployments()): + return + await asyncio.sleep(self.polling_interval) + finally: + await asyncio.shield(self.remove_request(request_id=request.request_id, model_name=request.model_name)) + raise Timeout(message="Request timed out while polling queue", model=request.model_name, llm_provider="openai") + async def remove_request(self, request_id: str, model_name: str) -> None: """ Remove a specific request from the priority queue for a model. @@ -118,21 +144,23 @@ class Scheduler: return self.queue @with_service_target(SCHEDULER_QUEUE_TARGET) - async def get_queue(self, model_name: str) -> list: + async def get_queue(self, model_name: str) -> list[QueueEntry]: """ - Return a queue for that specific model group + Return a queue for that specific model group. + + Redis hands the queue back as JSON lists, so every entry is validated into the + (priority, request_id) tuple the heap operations compare against. """ if self.cache is not None: _cache_key: Final = f"{SchedulerCacheKeys.queue.value}:{model_name}" response: Final = await self.cache.async_get_cache(key=_cache_key) - if response is None or not isinstance(response, list): + if not isinstance(response, list): return [] - elif isinstance(response, list): - return response + return _QUEUE_ENTRIES.validate_python(response) return self.queue @with_service_target(SCHEDULER_QUEUE_TARGET) - async def save_queue(self, queue: list, model_name: str) -> None: + async def save_queue(self, queue: list[QueueEntry], model_name: str) -> None: """ Save the updated queue of the model group """ diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 07521fef965..cbc09edc357 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -89,6 +89,7 @@ ignored_function_names = [ "_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py "_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py "_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name) + "_wait_for_scheduler_turn", # Tested through prioritized acompletion and atext_completion in test_router.py "arm_routing_read_prefetch", # Tested in tests/unit/caching/test_request_redis_batch_pre_call.py (file lacks "router" in name) "_configured_model_info", # Tested through get_configured_service_tiers in test_router.py "_routable_deployments", # Tested through get_configured_service_tiers and get_routable_upstream_model in test_router.py diff --git a/tests/integration/routing/test_priority_scheduler_queue_cleanup.py b/tests/integration/routing/test_priority_scheduler_queue_cleanup.py new file mode 100644 index 00000000000..a4606554d34 --- /dev/null +++ b/tests/integration/routing/test_priority_scheduler_queue_cleanup.py @@ -0,0 +1,860 @@ +from __future__ import annotations + +import asyncio +import http.client +import json +import re +import signal +import threading +import time +import uuid +from collections.abc import Callable, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from typing import Final, Literal, TypeVar +from urllib.parse import urlsplit + +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.openai_wire import answering_model_discovery, chat_reply, openai_error, responses_reply +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.redis_process import OwnedRedis, owned_redis +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter +from redis import Redis + +T = TypeVar("T") +Endpoint = Literal["chat", "chat-stream", "completions", "completions-stream", "queue", "queue-stream"] +Instance = Literal["first", "second"] + +MARKER: Final = re.compile(r"sched-[0-9a-f]{32}") +UNKNOWN_MODEL: Final = "Invalid model name passed in model=" +NO_DEPLOYMENTS: Final = "No deployments available" +QUEUE_TIMEOUT: Final = "Request timed out while polling queue" +ROUTER_TIMEOUT_SECONDS: Final = 8 +PROMPT_SECONDS: Final = 4 +COOLDOWN_SECONDS: Final = 60 +USAGE: Final[dict[str, JsonValue]] = {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8} +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +OWNED_CELL_TIMEOUT: Final = 2 * graceful_stop_seconds() + 120 +PAIR_CELL_TIMEOUT: Final = 3 * graceful_stop_seconds() + 120 +WORKER_HEALTHCHECK_ARGUMENTS: Final = ("--timeout_worker_healthcheck", str(int(graceful_stop_seconds()))) +PINNED_CONNECTION_TIMEOUT_SECONDS: Final = 30 +PINNED_LIMITS: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=30) +ENDPOINTS: Final[tuple[Endpoint, ...]] = ( + "chat", + "chat-stream", + "completions", + "completions-stream", + "queue", + "queue-stream", +) +PAIR_GROUPS: Final = tuple(f"sched-pair-{row}" for row in ("p1", "p2", "p3", "p4", "p5", "p6", "p7", "p8")) +INMEM_GROUP: Final = "sched-inmem" +OUTAGE_GROUP: Final = "sched-outage" +KILL_GROUP: Final = "sched-kill" +GHOST_ENTRY: Final[list[JsonValue]] = [1, "ghost"] +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +JSON_LIST: Final = TypeAdapter(list[JsonValue]) + + +def new_marker() -> str: + return f"sched-{uuid.uuid4().hex}" + + +def marker_in(text: str) -> str: + found: Final = MARKER.search(text) + assert found is not None, text[:300] + return found.group(0) + + +def markers_in(text: str) -> frozenset[str]: + return frozenset(MARKER.findall(text)) + + +def upstream_name(group: str) -> str: + return f"{group}-upstream" + + +def queue_key(group: str) -> str: + return f"scheduler:queue:{group}" + + +def data_frame(frame: Mapping[str, JsonValue]) -> bytes: + return b"data: " + json.dumps(frame).encode() + b"\n\n" + + +def text_completion_reply(marker: str, model: str, *, stream: bool) -> Reply: + choice: Final[dict[str, JsonValue]] = { + "text": f"served {marker}", + "index": 0, + "logprobs": None, + "finish_reason": "stop", + } + body: Final[dict[str, JsonValue]] = { + "id": f"cmpl-{marker}", + "object": "text_completion", + "created": 1, + "model": model, + "choices": [choice], + "usage": USAGE, + } + if not stream: + return Reply(body=json.dumps(body).encode()) + first: Final = data_frame({**body, "choices": [{**choice, "finish_reason": None}]}) + last: Final = data_frame({**body, "choices": [{**choice, "text": ""}]}) + return Reply(content_type="text/event-stream", chunks=(first, last + b"data: [DONE]\n\n")) + + +@dataclass(frozen=True, slots=True) +class Upstream: + refusing: Mapping[str, threading.Event] + held: SimpleQueue[str] + release: threading.Event + hold: frozenset[str] + + def respond(self, request: Request) -> Reply: + body: Final = JSON_OBJECT.validate_json(request.body) + model: Final = str(body["model"]) + refusal: Final = self.refusing.get(model) + if refusal is not None and refusal.is_set(): + return openai_error(401) + marker: Final = marker_in(request.body.decode()) + if model in self.hold: + self.held.put(marker) + assert self.release.wait(timeout=60), "The held burst was never released" + stream: Final = body.get("stream") is True + if request.target.endswith("/responses"): + return responses_reply(f"resp_{marker}", model, f"served {marker}", stream=stream) + if request.target.endswith("/completions") and not request.target.endswith("/chat/completions"): + return text_completion_reply(marker, model, stream=stream) + return chat_reply(f"chatcmpl-{marker}", model, f"served {marker}", stream=stream) + + +def upstream_for(groups: Sequence[str], *, hold: Sequence[str] = ()) -> Upstream: + return Upstream( + {upstream_name(group): threading.Event() for group in groups}, + SimpleQueue(), + threading.Event(), + frozenset(upstream_name(group) for group in hold), + ) + + +def received_markers(wire: Wire) -> tuple[str, ...]: + return tuple(marker_in(request.body.decode()) for request in wire.drain() if request.method == "POST") + + +def assert_once(received: Sequence[str], markers: Sequence[str]) -> None: + counts: Final = {marker: received.count(marker) for marker in markers} + assert all(count == 1 for count in counts.values()), counts + + +def assert_never(received: Sequence[str], markers: Sequence[str]) -> None: + assert not set(markers) & set(received), (markers, received) + + +def deployment_entry(group: str, wire: Wire) -> dict[str, JsonValue]: + return { + "model_name": group, + "litellm_params": { + "model": f"openai/{upstream_name(group)}", + "api_base": wire.url + "/v1", + "api_key": "synthetic-openai-key", + }, + "model_info": {"id": f"{group}-deployment"}, + } + + +def owned_config( + directory: Path, + wire: Wire, + groups: Sequence[str], + settings: Mapping[str, JsonValue], + *, + cancel_on_disconnect: bool = False, +) -> Path: + base: Final = JSON_OBJECT.validate_python( + yaml.safe_load((Path(__file__).parents[1] / "proxy_config.yaml").read_text()) + ) + general_settings: Final = base.get("general_settings", {}) + assert isinstance(general_settings, dict) + general: Final[dict[str, JsonValue]] = { + **general_settings, + **({"cancel_on_disconnect": True} if cancel_on_disconnect else {}), + } + config: Final[dict[str, JsonValue]] = { + **base, + "general_settings": general, + "router_settings": dict(settings), + "model_list": [deployment_entry(group, wire) for group in groups], + } + path: Final = directory / f"scheduler-{uuid.uuid4().hex[:8]}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def cooldown_settings(cache: OwnedRedis | None) -> dict[str, JsonValue]: + redis: Final = {} if cache is None else {"redis_host": cache.host, "redis_port": cache.port} + return {"num_retries": 0, "timeout": ROUTER_TIMEOUT_SECONDS, "cooldown_time": COOLDOWN_SECONDS, **redis} + + +def redis_settings(cache: OwnedRedis) -> dict[str, JsonValue]: + return {"num_retries": 0, "redis_host": cache.host, "redis_port": cache.port} + + +def chat_body(group: str, marker: str, **extra: JsonValue) -> dict[str, JsonValue]: + return {"model": group, "messages": [{"role": "user", "content": marker}], **extra} + + +def completion_body(group: str, marker: str, **extra: JsonValue) -> dict[str, JsonValue]: + return {"model": group, "prompt": marker, **extra} + + +@dataclass(frozen=True, slots=True) +class Answer: + status: int + text: str + headers: Mapping[str, str] + seconds: float + + @property + def identity(self) -> str: + return str(JSON_OBJECT.validate_json(self.text)["id"]) + + +def timed(send: Callable[[], httpx.Response]) -> Answer: + started: Final = time.monotonic() + response: Final = send() + return Answer(response.status_code, response.text, dict(response.headers), time.monotonic() - started) + + +def settled(send: Callable[[], T], text: Callable[[T], str]) -> T: + return eventually(send, lambda observed: UNKNOWN_MODEL not in text(observed), seconds=30) + + +def sdk_settled(call: Callable[[], T]) -> T: + def attempt() -> T | openai.APIStatusError: + try: + return call() + except openai.APIStatusError as error: + if UNKNOWN_MODEL in error.message: + return error + raise + + outcome: Final = eventually(attempt, lambda observed: not isinstance(observed, openai.APIStatusError), seconds=30) + assert not isinstance(outcome, openai.APIStatusError) + return outcome + + +def post(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> Answer: + return settled(lambda: timed(lambda: gateway.request("POST", path, body)), lambda answer: answer.text) + + +def stream_text(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> Answer: + def send() -> Answer: + started: Final = time.monotonic() + with gateway.client.stream( + "POST", path, json=body, headers={"Authorization": f"Bearer {gateway.key}"} + ) as response: + text: Final = "".join(response.iter_text()) + return Answer(response.status_code, text, dict(response.headers), time.monotonic() - started) + + return settled(send, lambda answer: answer.text) + + +def stream_identities(text: str) -> frozenset[str]: + frames: Final = tuple( + JSON_OBJECT.validate_json(line[len("data: ") :]) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + assert frames, text + return frozenset(str(frame["id"]) for frame in frames) + + +def call(gateway: Gateway, endpoint: Endpoint, group: str, marker: str, priority: int = 1) -> Answer: + match endpoint: + case "chat": + return post(gateway, "/v1/chat/completions", chat_body(group, marker, priority=priority)) + case "chat-stream": + return stream_text( + gateway, "/v1/chat/completions", chat_body(group, marker, priority=priority, stream=True) + ) + case "completions": + return post(gateway, "/v1/completions", completion_body(group, marker, priority=priority)) + case "completions-stream": + return stream_text( + gateway, "/v1/completions", completion_body(group, marker, priority=priority, stream=True) + ) + case "queue": + return post(gateway, "/queue/chat/completions", chat_body(group, marker, priority=priority)) + case "queue-stream": + return stream_text( + gateway, "/queue/chat/completions", chat_body(group, marker, priority=priority, stream=True) + ) + + +def served_identity(status: int, text: str, marker: str) -> str: + assert status == 200, (status, text) + identities: Final = ( + stream_identities(text) if text.startswith("data:") else frozenset({str(JSON_OBJECT.validate_json(text)["id"])}) + ) + assert identities in ({f"chatcmpl-{marker}"}, {f"cmpl-{marker}"}), (identities, marker) + assert markers_in(text) == {marker}, text + return next(iter(identities)) + + +def assert_served(answer: Answer, marker: str) -> str: + return served_identity(answer.status, answer.text, marker) + + +def spend_row_lands(identity: str) -> None: + rows: Final = eventually( + lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (identity,)), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["spend"] is not None, rows + + +def sdk(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=30) + + +def queue_entries(cache: OwnedRedis, group: str) -> list[JsonValue] | None: + with Redis(host=cache.host, port=cache.port) as client: + raw: Final = client.get(queue_key(group)) + if raw is None: + return None + assert isinstance(raw, bytes), raw + return JSON_LIST.validate_json(raw) + + +def write_queue(cache: OwnedRedis, group: str, entries: Sequence[JsonValue]) -> None: + with Redis(host=cache.host, port=cache.port) as client: + client.set(queue_key(group), json.dumps(list(entries))) + + +def persist_queue(cache: OwnedRedis, group: str) -> None: + with Redis(host=cache.host, port=cache.port) as client: + client.persist(queue_key(group)) + + +def entry_priority(entry: JsonValue) -> int: + assert isinstance(entry, list), entry + priority: Final = entry[0] + assert isinstance(priority, int), entry + return priority + + +def waiting_entries(cache: OwnedRedis, group: str, priority: int) -> list[JsonValue]: + def read() -> list[JsonValue]: + return queue_entries(cache, group) or [] + + return eventually( + read, lambda found: any(entry_priority(entry) == priority for entry in found), seconds=PROMPT_SECONDS + ) + + +def cooled(gateway: Gateway, upstream: Upstream, group: str) -> None: + upstream.refusing[upstream_name(group)].set() + trip: Final = post(gateway, "/v1/chat/completions", chat_body(group, new_marker())) + assert trip.status == 401, (trip.status, trip.text) + eventually( + lambda: post(gateway, "/v1/chat/completions", chat_body(group, new_marker())), + lambda answer: answer.status == 429 and NO_DEPLOYMENTS in answer.text, + seconds=10, + ) + + +def assert_refused_at_once(answer: Answer) -> None: + assert answer.status == 429, (answer.status, answer.text) + assert NO_DEPLOYMENTS in answer.text, answer.text + assert answer.seconds < PROMPT_SECONDS, answer.seconds + + +def assert_timed_out_polling(answer: Answer) -> None: + assert answer.status == 408, (answer.status, answer.text) + assert QUEUE_TIMEOUT in answer.text, answer.text + assert answer.seconds >= ROUTER_TIMEOUT_SECONDS, answer.seconds + + +@pytest.fixture(scope="module") +def rig_upstream() -> Iterator[tuple[Upstream, Wire]]: + upstream: Final = upstream_for(()) + with wire_server(answering_model_discovery(upstream.respond)) as wire: + yield upstream, wire + + +@pytest.fixture +def rig_model(gateway: Gateway, rig_upstream: tuple[Upstream, Wire]) -> Iterator[tuple[str, Wire]]: + _, wire = rig_upstream + with gateway.scenario() as scenario: + yield scenario.model(api_base=wire.url + "/v1"), wire + + +def test_sdk_chat_with_priority_is_served_once_and_billed(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + client: Final = sdk(gateway) + raw: Final = sdk_settled( + lambda: client.chat.completions.with_raw_response.create( + model=model, messages=[{"role": "user", "content": marker}], extra_body={"priority": 1} + ) + ) + assert served_identity(raw.status_code, raw.text, marker) == f"chatcmpl-{marker}" + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert "priority" not in JSON_OBJECT.validate_json(upstream_request.body), upstream_request.body + spend_row_lands(f"chatcmpl-{marker}") + + +def test_sdk_chat_stream_with_priority_is_served_once_and_billed(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + client: Final = sdk(gateway) + chunks: Final = sdk_settled( + lambda: tuple( + client.chat.completions.create( + model=model, messages=[{"role": "user", "content": marker}], stream=True, extra_body={"priority": 1} + ) + ) + ) + assert {chunk.id for chunk in chunks} == {f"chatcmpl-{marker}"}, chunks + content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + assert content == f"served {marker}", chunks + assert_once(received_markers(wire), (marker,)) + spend_row_lands(f"chatcmpl-{marker}") + + +def test_async_sdk_chat_with_priority_is_served_once(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + + async def send() -> tuple[int, str]: + async with openai.AsyncOpenAI( + base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=30 + ) as client: + raw: Final = await client.chat.completions.with_raw_response.create( + model=model, messages=[{"role": "user", "content": marker}], extra_body={"priority": 1} + ) + return raw.status_code, raw.text + + status, text = sdk_settled(lambda: asyncio.run(send())) + assert served_identity(status, text, marker) == f"chatcmpl-{marker}" + assert_once(received_markers(wire), (marker,)) + + +def test_raw_chat_with_duplicate_priority_keys_takes_the_last_one( + gateway: Gateway, rig_model: tuple[str, Wire] +) -> None: + model, wire = rig_model + marker: Final = new_marker() + body: Final = ( + '{"model": "%s", "messages": [{"role": "user", "content": "%s"}], "priority": "not-an-int", "priority": 1}' + % (model, marker) + ) + + def send() -> httpx.Response: + return gateway.client.post( + "/v1/chat/completions", + content=body.encode(), + headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"}, + ) + + answer: Final = settled(lambda: timed(send), lambda observed: observed.text) + assert_served(answer, marker) + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert "priority" not in JSON_OBJECT.validate_json(upstream_request.body), upstream_request.body + + +@pytest.mark.parametrize("endpoint", ("completions", "completions-stream", "queue", "queue-stream")) +def test_other_prioritized_endpoints_are_served_once_and_billed( + gateway: Gateway, rig_model: tuple[str, Wire], endpoint: Endpoint +) -> None: + model, wire = rig_model + marker: Final = new_marker() + answer: Final = call(gateway, endpoint, model, marker) + identity: Final = assert_served(answer, marker) + if endpoint == "queue": + assert answer.headers.get("x-litellm-priority") == "1", answer.headers + assert_once(received_markers(wire), (marker,)) + spend_row_lands(identity) + + +@pytest.mark.parametrize("priority", ("1", [1], "", "p" * 5120, 0), ids=("string", "list", "empty", "5kb", "zero")) +def test_non_integer_priority_bypasses_the_scheduler_and_is_forwarded( + gateway: Gateway, rig_model: tuple[str, Wire], priority: JsonValue +) -> None: + model, wire = rig_model + marker: Final = new_marker() + answer: Final = post(gateway, "/v1/chat/completions", chat_body(model, marker, priority=priority)) + assert_served(answer, marker) + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert JSON_OBJECT.validate_json(upstream_request.body).get("priority") == priority + + +def test_unauthenticated_prioritized_request_never_reaches_the_upstream( + gateway: Gateway, rig_model: tuple[str, Wire] +) -> None: + model, wire = rig_model + marker: Final = new_marker() + response: Final = gateway.client.post("/v1/chat/completions", json=chat_body(model, marker, priority=1)) + assert response.status_code == 401, response.text + assert_never(received_markers(wire), (marker,)) + + +@pytest.mark.parametrize("path", ("/v1/messages", "/v1/responses"), ids=("messages", "responses")) +def test_priority_on_unscheduled_routes_is_served_once( + gateway: Gateway, rig_model: tuple[str, Wire], path: str +) -> None: + model, wire = rig_model + marker: Final = new_marker() + body: Final[dict[str, JsonValue]] = ( + {"model": model, "max_tokens": 32, "messages": [{"role": "user", "content": marker}], "priority": 1} + if path == "/v1/messages" + else {"model": model, "input": marker, "priority": 1} + ) + answer: Final = post(gateway, path, body) + assert answer.status == 200, (answer.status, answer.text) + assert markers_in(answer.text) == {marker}, answer.text + assert_once(received_markers(wire), (marker,)) + + +def pinned(gateway: Gateway, stack: ExitStack) -> Gateway: + client: Final = stack.enter_context( + httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False, limits=PINNED_LIMITS) + ) + return Gateway(client, gateway.key, gateway.upstream_url) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_in_memory_queue_forgets_served_requests_before_a_cooldown(gateway: Gateway, tmp_path: Path) -> None: + upstream: Final = upstream_for((INMEM_GROUP,)) + with ExitStack() as stack: + wire: Final = stack.enter_context(wire_server(answering_model_discovery(upstream.respond))) + config: Final = owned_config(tmp_path, wire, (INMEM_GROUP,), cooldown_settings(None)) + owned: Final = stack.enter_context( + owned_proxy_process( + gateway, tmp_path, {}, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + worker: Final = pinned(owned.gateway, stack) + served: Final = new_marker() + assert_served(post(worker, "/v1/chat/completions", chat_body(INMEM_GROUP, served, priority=1)), served) + cooled(worker, upstream, INMEM_GROUP) + waiting: Final = (new_marker(), new_marker()) + for marker in waiting: + assert_refused_at_once(post(worker, "/v1/chat/completions", chat_body(INMEM_GROUP, marker, priority=2))) + received: Final = received_markers(wire) + assert_once(received, (served,)) + assert_never(received, waiting) + + +@dataclass(frozen=True, slots=True) +class Pair: + first: Gateway + second: Gateway + cache: OwnedRedis + wire: Wire + upstream: Upstream + + def at(self, instance: Instance) -> Gateway: + return self.first if instance == "first" else self.second + + +@pytest.fixture(scope="module") +def pair(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Pair]: + directory: Final = tmp_path_factory.mktemp("scheduler-pair") + upstream: Final = upstream_for(PAIR_GROUPS) + with ExitStack() as stack: + gateway: Final = stack.enter_context(gateway_from_environment()) + cache: Final = stack.enter_context(owned_redis(directory)) + wire: Final = stack.enter_context(wire_server(answering_model_discovery(upstream.respond))) + config: Final = owned_config(directory, wire, PAIR_GROUPS, cooldown_settings(cache), cancel_on_disconnect=True) + overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} + first: Final = stack.enter_context( + owned_proxy_process( + gateway, directory, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + second: Final = stack.enter_context( + owned_proxy_process( + gateway, directory, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + yield Pair(first.gateway, second.gateway, cache, wire, upstream) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_served_request_leaves_no_queue_entry_for_the_other_instance(pair: Pair) -> None: + group: Final = PAIR_GROUPS[0] + markers: Final = (new_marker(), new_marker(), new_marker()) + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, markers[0], priority=1)), markers[0]) + assert queue_entries(pair.cache, group) == [] + assert_served(post(pair.second, "/v1/chat/completions", chat_body(group, markers[1], priority=1)), markers[1]) + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, markers[2], priority=1)), markers[2]) + assert_once(received_markers(pair.wire), markers) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_concurrent_prioritized_burst_across_instances_and_endpoints_is_served(pair: Pair) -> None: + group: Final = PAIR_GROUPS[1] + primer: Final = new_marker() + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, primer, priority=1)), primer) + instances: Final[tuple[Instance, ...]] = ("first", "second") + plan: Final = tuple((instance, endpoint, new_marker()) for instance in instances for endpoint in ENDPOINTS) + + def one(item: tuple[Instance, Endpoint, str]) -> Answer: + instance, endpoint, marker = item + return call(pair.at(instance), endpoint, group, marker) + + with ThreadPoolExecutor(max_workers=len(plan)) as pool: + answers: Final = tuple(pool.map(one, plan)) + for (_, _, marker), answer in zip(plan, answers, strict=True): + assert_served(answer, marker) + closing: Final = new_marker() + assert_served(post(pair.second, "/v1/chat/completions", chat_body(group, closing, priority=1)), closing) + assert_once(received_markers(pair.wire), (primer, *(marker for _, _, marker in plan), closing)) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_fresh_queue_during_a_cooldown_is_refused_at_once(pair: Pair) -> None: + group: Final = PAIR_GROUPS[2] + cooled(pair.second, pair.upstream, group) + client: Final = sdk(pair.second) + started: Final = time.monotonic() + with pytest.raises(openai.APIStatusError) as refused: + client.completions.create(model=group, prompt=new_marker(), extra_body={"priority": 1}) + assert refused.value.status_code == 429, refused.value.message + assert NO_DEPLOYMENTS in refused.value.message + assert time.monotonic() - started < PROMPT_SECONDS + assert queue_entries(pair.cache, group) == [] + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_served_request_on_one_instance_does_not_block_a_waiter_on_the_other(pair: Pair) -> None: + group: Final = PAIR_GROUPS[3] + served: Final = new_marker() + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, served, priority=1)), served) + cooled(pair.second, pair.upstream, group) + assert_refused_at_once(post(pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2))) + + +def assert_refused_before_the_router_timeout(answer: Answer) -> None: + assert answer.status == 429, (answer.status, answer.text) + assert NO_DEPLOYMENTS in answer.text, answer.text + assert answer.seconds < ROUTER_TIMEOUT_SECONDS, answer.seconds + + +def assert_drained(cache: OwnedRedis, group: str) -> None: + assert queue_entries(cache, group) in ([], None) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_behind_an_expired_dead_replica_entry_is_re_enqueued_and_proceeds(pair: Pair) -> None: + group: Final = PAIR_GROUPS[4] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + answer: Final = post(pair.second, "/queue/chat/completions", chat_body(group, new_marker(), priority=2)) + assert_refused_before_the_router_timeout(answer) + assert_drained(pair.cache, group) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_behind_a_live_entry_times_out_and_removes_only_itself(pair: Pair) -> None: + group: Final = PAIR_GROUPS[5] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + with ThreadPoolExecutor(max_workers=1) as pool: + waiter: Final = pool.submit( + post, pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2) + ) + waiting_entries(pair.cache, group, 2) + persist_queue(pair.cache, group) + assert_timed_out_polling(waiter.result(timeout=30)) + assert queue_entries(pair.cache, group) == [GHOST_ENTRY] + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_proceeds_once_the_entry_ahead_is_cleared(pair: Pair) -> None: + group: Final = PAIR_GROUPS[6] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + with ThreadPoolExecutor(max_workers=1) as pool: + waiter: Final = pool.submit( + post, pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2) + ) + entries: Final = waiting_entries(pair.cache, group, 2) + write_queue(pair.cache, group, [entry for entry in entries if entry_priority(entry) == 2]) + assert_refused_before_the_router_timeout(waiter.result(timeout=30)) + assert_drained(pair.cache, group) + + +def pinned_connection(gateway: Gateway) -> http.client.HTTPConnection: + address: Final = urlsplit(str(gateway.client.base_url)) + assert address.hostname is not None and address.port is not None, address + return http.client.HTTPConnection(address.hostname, address.port, timeout=PINNED_CONNECTION_TIMEOUT_SECONDS) + + +def send_pinned(connection: http.client.HTTPConnection, key: str, body: Mapping[str, JsonValue]) -> None: + connection.request( + "POST", + "/v1/chat/completions", + body=json.dumps(body).encode(), + headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json"}, + ) + + +def post_pinned(connection: http.client.HTTPConnection, key: str, body: Mapping[str, JsonValue]) -> tuple[int, str]: + send_pinned(connection, key, body) + response: Final = connection.getresponse() + text: Final = response.read().decode() + assert connection.sock is not None, "the proxy closed the pinned connection" + return response.status, text + + +def cooled_over(connection: http.client.HTTPConnection, key: str, upstream: Upstream, group: str) -> None: + upstream.refusing[upstream_name(group)].set() + trip: Final = post_pinned(connection, key, chat_body(group, new_marker())) + assert trip[0] == 401, trip + eventually( + lambda: post_pinned(connection, key, chat_body(group, new_marker())), + lambda answer: answer[0] == 429 and NO_DEPLOYMENTS in answer[1], + seconds=10, + ) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_disconnected_waiter_is_removed_from_the_queue(pair: Pair) -> None: + group: Final = PAIR_GROUPS[7] + waiter: Final = pinned_connection(pair.second) + cooled_over(waiter, pair.second.key, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + send_pinned(waiter, pair.second.key, chat_body(group, new_marker(), priority=2)) + waiting_entries(pair.cache, group, 2) + persist_queue(pair.cache, group) + waiter.close() + eventually( + lambda: queue_entries(pair.cache, group), lambda entries: entries == [GHOST_ENTRY], seconds=PROMPT_SECONDS + ) + + +Outcome = Answer | httpx.TransportError + + +def burst(gateway: Gateway, group: str, count: int) -> tuple[tuple[Endpoint, str, Outcome], ...]: + plan: Final[tuple[tuple[Endpoint, str], ...]] = tuple( + (ENDPOINTS[index % len(ENDPOINTS)], new_marker()) for index in range(count) + ) + + def one(item: tuple[Endpoint, str]) -> Outcome: + endpoint, marker = item + with httpx.Client(base_url=gateway.client.base_url, timeout=60, trust_env=False) as client: + try: + return call(Gateway(client, gateway.key, gateway.upstream_url), endpoint, group, marker) + except httpx.TransportError as error: + return error + + with ThreadPoolExecutor(max_workers=count) as pool: + outcomes: Final = tuple(pool.map(one, plan)) + return tuple((endpoint, marker, outcome) for (endpoint, marker), outcome in zip(plan, outcomes, strict=True)) + + +def assert_all_served(outcomes: Sequence[tuple[Endpoint, str, Outcome]]) -> tuple[str, ...]: + for _, marker, outcome in outcomes: + assert isinstance(outcome, Answer), repr(outcome) + assert_served(outcome, marker) + return tuple(marker for _, marker, _ in outcomes) + + +def served_eventually(gateway: Gateway, group: str, marker: str) -> None: + def send() -> Outcome: + try: + return post(gateway, "/v1/chat/completions", chat_body(group, marker, priority=1)) + except httpx.TransportError as error: + return error + + answer: Final = eventually(send, lambda outcome: isinstance(outcome, Answer), seconds=30) + assert isinstance(answer, Answer) + assert_served(answer, marker) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_prioritized_requests_survive_a_redis_outage(gateway: Gateway, tmp_path: Path) -> None: + upstream: Final = upstream_for((OUTAGE_GROUP,)) + with owned_redis(tmp_path) as cache, wire_server(answering_model_discovery(upstream.respond)) as wire: + config: Final = owned_config(tmp_path, wire, (OUTAGE_GROUP,), redis_settings(cache)) + overrides: Final = { + "REDIS_HOST": cache.host, + "REDIS_PORT": str(cache.port), + "REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1", + } + with owned_proxy_process( + gateway, tmp_path, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) as owned: + before: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + cache.stop() + during: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + cache.start() + after: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + assert_once(received_markers(wire), (*before, *during, *after)) + eventually(lambda: queue_entries(cache, OUTAGE_GROUP), lambda entries: entries in (None, []), seconds=10) + + +def established_upstream_connections(pid: int, wire: Wire) -> int: + port: Final = urlsplit(wire.url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_sibling_worker_keeps_serving_prioritized_requests_after_a_worker_is_killed( + gateway: Gateway, tmp_path: Path +) -> None: + upstream: Final = upstream_for((KILL_GROUP,), hold=(KILL_GROUP,)) + with owned_redis(tmp_path) as cache, wire_server(answering_model_discovery(upstream.respond)) as wire: + config: Final = owned_config(tmp_path, wire, (KILL_GROUP,), redis_settings(cache)) + overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} + with owned_proxy_process( + gateway, tmp_path, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) as owned: + workers: Final = eventually( + lambda: tuple(int(found.group(1)) for found in STARTED_WORKER.finditer(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=graceful_stop_seconds(), + ) + with ThreadPoolExecutor(max_workers=1) as pool: + pending: Final = pool.submit(burst, owned.gateway, KILL_GROUP, 20) + eventually(upstream.held.qsize, lambda size: size == 20, seconds=60) + held_by: Final = {pid: established_upstream_connections(pid, wire) for pid in workers} + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + upstream.release.set() + outcomes: Final = pending.result(timeout=90) + answered: Final = tuple(outcome for outcome in outcomes if isinstance(outcome[2], Answer)) + failed: Final = tuple(outcome for outcome in outcomes if not isinstance(outcome[2], Answer)) + assert len(answered) == held_by[survivor_pid], (held_by, len(answered)) + assert len(failed) == held_by[victim_pid], (held_by, len(failed)) + assert_all_served(answered) + follow_up: Final = new_marker() + served_eventually(owned.gateway, KILL_GROUP, follow_up) + eventually( + lambda: len(STARTED_WORKER.findall(owned.log.read_text())), + lambda count: count == 3, + seconds=graceful_stop_seconds(), + ) + assert_once(received_markers(wire), (*(marker for _, marker, _ in outcomes), follow_up)) diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 8f3133dcaa4..ec9e4861fdd 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -65,6 +65,7 @@ from litellm.router_utils.cooldown_handlers import ( ) from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute +from litellm.scheduler import FlowItem from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import ( CustomRoutingStrategyBase, @@ -19949,6 +19950,64 @@ def test_access_windows_filter_reserved_deployments_method(): ] == ["reserved-deployment", "open-deployment"] +def _scheduled_router(timeout: float) -> Router: + return Router( + model_list=[ + { + "model_name": "sched-model", + "litellm_params": {"model": "openai/sched-model", "api_key": "sk-fake", "mock_response": "hi"}, + "model_info": {"id": "sched-deployment"}, + } + ], + timeout=timeout, + ) + + +async def _send_scheduled_chat(router: Router, priority: int) -> object: + return await router.acompletion( + model="sched-model", messages=[{"role": "user", "content": "hi"}], priority=priority + ) + + +async def _send_scheduled_text(router: Router, priority: int) -> object: + return await router.atext_completion(model="sched-model", prompt="hi", priority=priority) + + +@pytest.mark.parametrize( + "send", [_send_scheduled_chat, _send_scheduled_text], ids=["schedule_acompletion", "schedule_factory"] +) +@pytest.mark.asyncio +async def test_admitted_prioritized_request_does_not_block_later_request_during_cooldown( + send: Callable[[Router, int], Awaitable[object]], +): + from litellm.types.router import RouterRateLimitError + + router: Final = _scheduled_router(timeout=1) + await send(router, 1) + _cool_down(router, "sched-deployment") + + with pytest.raises(RouterRateLimitError, match="cooldown"): + await send(router, 2) + + +@pytest.mark.parametrize("stop_waiting", ["cancel", "timeout"]) +@pytest.mark.asyncio +async def test_prioritized_request_leaves_queue_when_it_stops_waiting(stop_waiting: Literal["cancel", "timeout"]): + router: Final = _scheduled_router(timeout=0.5) + _cool_down(router, "sched-deployment") + await router.scheduler.add_request(FlowItem(priority=0, request_id="head-of-queue", model_name="sched-model")) + waiting: Final = asyncio.create_task(_send_scheduled_chat(router, 5)) + await asyncio.sleep(0.05) + assert len(await router.scheduler.get_queue("sched-model")) == 2 + + if stop_waiting == "cancel": + waiting.cancel() + with pytest.raises(asyncio.CancelledError if stop_waiting == "cancel" else litellm.Timeout): + await waiting + + assert await router.scheduler.get_queue("sched-model") == [(0, "head-of-queue")] + + @pytest.mark.asyncio async def test_bare_model_group_served_by_wildcard_deployment_uses_provider_prefixed_fallback_key() -> None: """Claude Code sends the bare "claude-sonnet-4-6" to /v1/messages; routing serves it through the diff --git a/tests/unit/test_scheduler.py b/tests/unit/test_scheduler.py index 553fb44cb27..e0968e9b2cf 100644 --- a/tests/unit/test_scheduler.py +++ b/tests/unit/test_scheduler.py @@ -3,7 +3,10 @@ import asyncio import importlib +import json import os +from collections.abc import Sequence +from typing import Final import pytest @@ -26,8 +29,8 @@ async def test_scheduler_diff_model_names(): await scheduler.add_request(item1) await scheduler.add_request(item2) - assert await scheduler.poll(id="10", model_name="gpt-3.5-turbo", health_deployments=[{"key": "value"}]) == True - assert await scheduler.poll(id="11", model_name="gpt-4", health_deployments=[{"key": "value"}]) == True + assert await scheduler.poll(request=item1, health_deployments=[{"key": "value"}]) == True + assert await scheduler.poll(request=item2, health_deployments=[{"key": "value"}]) == True @pytest.mark.asyncio @@ -50,7 +53,7 @@ async def test_scheduler_poll_persists_queue_to_cache(): await scheduler.add_request(item1) await scheduler.add_request(item2) - await scheduler.poll(id="10", model_name="gpt-3.5-turbo", health_deployments=[]) + await scheduler.poll(request=item1, health_deployments=[]) queue_key = f"{SchedulerCacheKeys.queue.value}:{item1.model_name}" updated_queue = redis_cache.store[queue_key] @@ -145,6 +148,190 @@ async def test_scheduler_queue_cleanup_on_timeout(): assert queue_after[0][1] == "req-0", "Expected req-0 (priority 0) to be at front" +@pytest.mark.asyncio +async def test_poll_admits_request_missing_from_queue_while_a_deployment_is_healthy(): + scheduler: Final = Scheduler() + + assert await scheduler.poll( + request=FlowItem(priority=1, request_id="erased-by-concurrent-write", model_name="sched-model"), + health_deployments=[{"model_info": {"id": "a"}}], + ) + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_admits_only_the_head_of_the_queue(): + scheduler: Final = Scheduler() + later: Final = FlowItem(priority=2, request_id="later", model_name="sched-model") + head: Final = FlowItem(priority=1, request_id="head", model_name="sched-model") + await scheduler.add_request(later) + await scheduler.add_request(head) + + assert not await scheduler.poll(request=later, health_deployments=[]) + assert await scheduler.poll(request=head, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [(2, "later")] + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_re_enqueues_a_request_a_concurrent_writer_erased_behind_the_head(): + scheduler: Final = Scheduler() + still_queued: Final = FlowItem(priority=0, request_id="still-queued", model_name="sched-model") + erased: Final = FlowItem(priority=1, request_id="erased-by-concurrent-write", model_name="sched-model") + await scheduler.add_request(still_queued) + + assert not await scheduler.poll(request=erased, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [(0, "still-queued"), (1, "erased-by-concurrent-write")] + assert await scheduler.poll(request=still_queued, health_deployments=[]) + assert await scheduler.poll(request=erased, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [] + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_admits_an_erased_request_that_outranks_the_queue(): + scheduler: Final = Scheduler() + await scheduler.add_request(FlowItem(priority=2, request_id="still-queued", model_name="sched-model")) + + assert await scheduler.poll( + request=FlowItem(priority=0, request_id="erased-urgent", model_name="sched-model"), health_deployments=[] + ) + assert await scheduler.get_queue("sched-model") == [(2, "still-queued")] + + +class _ExpiringCache: + def __init__(self) -> None: + self.store: dict[str, object] = {} + + async def async_get_cache(self, key: str, **kwargs: object) -> object: + return self.store.get(key) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.store[key] = value + + +@pytest.mark.asyncio +async def test_wait_for_turn_during_cooldown_survives_the_queue_key_expiring(): + cache: Final = _ExpiringCache() + scheduler: Final = Scheduler(redis_cache=cache) + scheduler.cache.in_memory_cache.cache_dict.clear() + + async def no_healthy_deployments_after_the_key_expired() -> Sequence[object]: + cache.store.clear() + scheduler.cache.in_memory_cache.cache_dict.clear() + return () + + await scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="sole-waiter", model_name="sched-model"), + timeout=5, + get_healthy_deployments=no_healthy_deployments_after_the_key_expired, + ) + + assert await scheduler.get_queue("sched-model") == [] + + + +class _JsonRoundTripRedisCache: + def __init__(self) -> None: + self.store: dict[str, str] = {} + + async def async_get_cache(self, key: str, **kwargs: object) -> object: + raw: Final = self.store.get(key) + return None if raw is None else json.loads(raw) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.store[key] = json.dumps(value) + + +@pytest.mark.asyncio +async def test_second_replica_enqueues_behind_a_queue_decoded_from_redis(): + redis_cache: Final = _JsonRoundTripRedisCache() + replica_a: Final = Scheduler(redis_cache=redis_cache) + replica_b: Final = Scheduler(redis_cache=redis_cache) + waiting_on_a: Final = FlowItem(priority=1, request_id="waiting-on-a", model_name="sched-model") + urgent_on_b: Final = FlowItem(priority=0, request_id="urgent-on-b", model_name="sched-model") + await replica_a.add_request(waiting_on_a) + + await replica_b.add_request(urgent_on_b) + + assert await replica_b.get_queue("sched-model") == [(0, "urgent-on-b"), (1, "waiting-on-a")] + assert await replica_b.poll(request=urgent_on_b, health_deployments=[]) + assert await replica_b.poll(request=waiting_on_a, health_deployments=[]) + await replica_b.remove_request(request_id="waiting-on-a", model_name="sched-model") + assert await replica_b.get_queue("sched-model") == [] + +class _PausesAfterEnqueueScheduler(Scheduler): + def __init__(self) -> None: + super().__init__() + self.enqueued: Final = asyncio.Event() + + async def add_request(self, request: FlowItem) -> None: + await super().add_request(request) + self.enqueued.set() + await asyncio.Event().wait() + + +async def _no_healthy_deployments() -> Sequence[object]: + return () + + +@pytest.mark.asyncio +async def test_wait_for_turn_removes_entry_when_cancelled_mid_enqueue(): + scheduler: Final = _PausesAfterEnqueueScheduler() + waiting: Final = asyncio.create_task( + scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="cancelled", model_name="sched-model"), + timeout=5, + get_healthy_deployments=_no_healthy_deployments, + ) + ) + await scheduler.enqueued.wait() + assert await scheduler.get_queue("sched-model") == [(1, "cancelled")] + + waiting.cancel() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == [] + + +class _HeldRemovalScheduler(Scheduler): + def __init__(self) -> None: + super().__init__() + self.removing: Final = asyncio.Event() + self.finish_removal: Final = asyncio.Event() + + async def remove_request(self, request_id: str, model_name: str) -> None: + self.removing.set() + await self.finish_removal.wait() + await super().remove_request(request_id=request_id, model_name=model_name) + + +@pytest.mark.asyncio +async def test_wait_for_turn_finishes_removal_when_cancelled_again_during_cleanup(): + scheduler: Final = _HeldRemovalScheduler() + await scheduler.add_request(FlowItem(priority=0, request_id="head", model_name="sched-model")) + polling: Final = asyncio.Event() + + async def no_healthy_deployments() -> Sequence[object]: + polling.set() + return () + + waiting: Final = asyncio.create_task( + scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="cancelled", model_name="sched-model"), + timeout=5, + get_healthy_deployments=no_healthy_deployments, + ) + ) + await polling.wait() + waiting.cancel() + await scheduler.removing.wait() + waiting.cancel() + scheduler.finish_removal.set() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == [(0, "head")] + + @pytest.fixture(autouse=True) def _vcr_outcome_gate(request, vcr): install_live_call_probe(request, vcr)