From f816e7194bb82cee4485db630d2477b2bd8ccc58 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 2 Oct 2026 15:54:47 -0700 Subject: [PATCH] test(router): integration cells auditing router budget limiting Covers provider, deployment and tag budgets end to end on a two-worker proxy: caps under, at and over the limit, zero and unset caps, fallbacks, every LLM endpoint through the OpenAI and Anthropic SDKs and raw HTTP, streaming charges, runtime /model/new, /provider/budgets, Prometheus, window reset, two proxies sharing Redis, restarts and a Redis outage mid burst. Rows that fail on main are strict xfails. --- .../routing/router_budgets/_rig.py | 312 ++++++++++++++++++ .../test_router_budget_endpoints.py | 258 +++++++++++++++ .../test_router_budget_multi_instance.py | 222 +++++++++++++ .../test_router_deployment_budgets.py | 253 ++++++++++++++ .../test_router_provider_budgets.py | 232 +++++++++++++ .../router_budgets/test_router_tag_budgets.py | 127 +++++++ 6 files changed, 1404 insertions(+) create mode 100644 tests/integration/routing/router_budgets/_rig.py create mode 100644 tests/integration/routing/router_budgets/test_router_budget_endpoints.py create mode 100644 tests/integration/routing/router_budgets/test_router_budget_multi_instance.py create mode 100644 tests/integration/routing/router_budgets/test_router_deployment_budgets.py create mode 100644 tests/integration/routing/router_budgets/test_router_provider_budgets.py create mode 100644 tests/integration/routing/router_budgets/test_router_tag_budgets.py diff --git a/tests/integration/routing/router_budgets/_rig.py b/tests/integration/routing/router_budgets/_rig.py new file mode 100644 index 00000000000..273606c9dba --- /dev/null +++ b/tests/integration/routing/router_budgets/_rig.py @@ -0,0 +1,312 @@ +"""Owned rig for router budget cells (provider, deployment and tag budgets in ``RouterBudgetLimiting``). + +- ``upstream``: OpenAI wire double for chat (plain and SSE) and Responses (plain and SSE); every reply + bills ``PROMPT_TOKENS`` + ``COMPLETION_TOKENS``. A body carrying ``PROVIDER_FAILURE`` gets HTTP 500. +- ``deployment(...)``: a model_list entry priced at ``PRICE`` per token, so one call costs ``CALL_COST`` + exactly in binary floating point and boundary comparisons are exact. +- ``budget_proxy(...)``: owned Redis plus an owned two-worker proxy booted from a config built here. Each + proxy owns its Redis, so provider spend keys never leak between files. +- ``send(...)``: one raw httpx call per endpoint shape. +""" + +from __future__ import annotations + +import json +import os +import uuid +from collections.abc import Iterator, Mapping, Sequence +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import yaml +from integration._support.client import JSON_OBJECT, Gateway, eventually +from integration._support.process import owned_proxy +from integration._support.redis_process import OwnedRedis, owned_redis +from integration._support.wire import Reply, Request +from pydantic import JsonValue +from redis import Redis + +PRICE: Final = 2**-10 +PROMPT_TOKENS: Final = 20 +COMPLETION_TOKENS: Final = 12 +CALL_COST: Final = (PROMPT_TOKENS + COMPLETION_TOKENS) * PRICE +PROVIDER_FAILURE: Final = "router-budget-provider-failure" +BUDGET_ERROR: Final = "No deployments available - crossed budget" +ENDPOINTS: Final = ("chat", "chat_stream", "messages", "messages_stream", "responses", "responses_stream") + + +def _sse(events: Sequence[Mapping[str, JsonValue]], *, named: bool) -> bytes: + frames: Final = ( + (f"event: {event['type']}\n" if named else "") + f"data: {json.dumps(event)}\n\n" for event in events + ) + return ("".join(frames) + ("" if named else "data: [DONE]\n\n")).encode() + + +def _chat(body: Mapping[str, JsonValue]) -> Reply: + identity: Final = f"chatcmpl-{uuid.uuid4().hex}" + usage: Final = { + "prompt_tokens": PROMPT_TOKENS, + "completion_tokens": COMPLETION_TOKENS, + "total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS, + } + base: Final = {"id": identity, "created": 1, "model": str(body.get("model"))} + if body.get("stream") is True: + chunks: Final = ( + {"choices": [{"index": 0, "delta": {"role": "assistant", "content": "budget"}, "finish_reason": None}]}, + {"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + {"choices": [], "usage": usage}, + ) + return Reply( + body=_sse([{**base, "object": "chat.completion.chunk", **chunk} for chunk in chunks], named=False), + content_type="text/event-stream", + ) + return Reply( + body=json.dumps( + { + **base, + "object": "chat.completion", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "budget"}, "finish_reason": "stop"} + ], + "usage": usage, + } + ).encode() + ) + + +def _response_object(model: str, status: str) -> dict[str, JsonValue]: + return { + "id": f"resp_{uuid.uuid4().hex}", + "object": "response", + "created_at": 1, + "status": status, + "model": model, + "output": [ + { + "type": "message", + "id": "msg_budget", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "budget", "annotations": []}], + } + ] + if status == "completed" + else [], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "usage": { + "input_tokens": PROMPT_TOKENS, + "output_tokens": COMPLETION_TOKENS, + "total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS, + } + if status == "completed" + else None, + } + + +def _responses(body: Mapping[str, JsonValue]) -> Reply: + model: Final = str(body.get("model")) + if body.get("stream") is not True: + return Reply(body=json.dumps(_response_object(model, "completed")).encode()) + events: Final = ( + {"type": "response.created", "sequence_number": 0, "response": _response_object(model, "in_progress")}, + { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": "msg_budget", + "output_index": 0, + "content_index": 0, + "delta": "budget", + }, + {"type": "response.completed", "sequence_number": 2, "response": _response_object(model, "completed")}, + ) + return Reply(body=_sse(events, named=True), content_type="text/event-stream") + + +def upstream(request: Request) -> Reply: + if PROVIDER_FAILURE.encode() in request.body: + return Reply(status=500, body=b'{"error":{"type":"server_error","message":"scripted provider failure"}}') + body: Final = json.loads(request.body or b"{}") + if request.target.split("?", 1)[0].endswith("/responses"): + return _responses(body) + return _chat(body) + + +def deployment( + model_name: str, + model: str, + upstream_url: str, + *, + model_id: str, + **litellm_params: JsonValue, +) -> dict[str, JsonValue]: + return { + "model_name": model_name, + "litellm_params": { + "model": model, + "api_base": f"{upstream_url}/v1", + "api_key": "router-budget-provider-key", + "input_cost_per_token": PRICE, + "output_cost_per_token": PRICE, + "max_retries": 0, + **litellm_params, + }, + "model_info": {"id": model_id}, + } + + +def write_config( + path: Path, + model_list: Sequence[Mapping[str, JsonValue]], + *, + provider_budget_config: Mapping[str, JsonValue] | None = None, + tag_budget_config: Mapping[str, JsonValue] | None = None, + litellm_settings: Mapping[str, JsonValue] | None = None, +) -> Path: + path.write_text( + yaml.safe_dump( + { + "model_list": list(model_list), + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "proxy_batch_write_at": 1, + }, + "litellm_settings": { + **({"tag_budget_config": dict(tag_budget_config)} if tag_budget_config is not None else {}), + **(litellm_settings or {}), + }, + "router_settings": { + "redis_host": "os.environ/REDIS_HOST", + "redis_port": "os.environ/REDIS_PORT", + "disable_cooldowns": True, + "num_retries": 0, + **( + {"provider_budget_config": dict(provider_budget_config)} + if provider_budget_config is not None + else {} + ), + }, + } + ) + ) + return path + + +@dataclass(frozen=True, slots=True) +class BudgetRig: + gateway: Gateway + redis: OwnedRedis + + def redis_float(self, key: str) -> float | None: + with Redis(host=self.redis.host, port=self.redis.port, socket_timeout=2) as client: + value: Final = client.get(key) + return None if value is None else float(value) + + def settled(self, key: str, expected: float, seconds: float = 15) -> float | None: + return eventually( + lambda: self.redis_float(key), lambda spend: spend == expected, seconds=seconds, return_last_on_timeout=True + ) + + +@contextmanager +def redis_for(tmp_path: Path) -> Iterator[OwnedRedis]: + directory: Final = tmp_path / f"redis-{uuid.uuid4().hex[:8]}" + directory.mkdir() + with owned_redis(directory) as cache: + yield cache + + +@contextmanager +def proxy_on( + gateway: Gateway, + tmp_path: Path, + config: Path, + cache: OwnedRedis, + *, + workers: int = 2, + extra: Mapping[str, str] | None = None, +) -> Iterator[Gateway]: + with owned_proxy( + gateway, + tmp_path, + { + "REDIS_HOST": cache.host, + "REDIS_PORT": str(cache.port), + "LITELLM_DISABLE_NO_REDIS_WARNING": "true", + **(extra or {}), + }, + config=config, + workers=workers, + ) as candidate: + yield candidate + + +@contextmanager +def budget_proxy(gateway: Gateway, tmp_path: Path, config: Path, *, workers: int = 2) -> Iterator[BudgetRig]: + with redis_for(tmp_path) as cache, proxy_on(gateway, tmp_path, config, cache, workers=workers) as candidate: + yield BudgetRig(candidate, cache) + + +def body_for(endpoint: str, model: str, text: str) -> dict[str, JsonValue]: + stream: Final = {"stream": True} if endpoint.endswith("_stream") else {} + if endpoint.startswith("responses"): + return {"model": model, "input": text, **stream} + if endpoint.startswith("messages"): + return {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": text}], **stream} + return { + "model": model, + "messages": [{"role": "user", "content": text}], + **({"stream": True, "stream_options": {"include_usage": True}} if stream else {}), + } + + +def path_for(endpoint: str) -> str: + if endpoint.startswith("responses"): + return "/v1/responses" + if endpoint.startswith("messages"): + return "/v1/messages" + return "/v1/chat/completions" + + +def send( + gateway: Gateway, + endpoint: str, + model: str, + text: str, + extra: Mapping[str, JsonValue] | None = None, +) -> httpx.Response: + return gateway.request("POST", path_for(endpoint), {**body_for(endpoint, model, text), **(extra or {})}) + + +def chat(gateway: Gateway, model: str, text: str, extra: Mapping[str, JsonValue] | None = None) -> httpx.Response: + return send(gateway, "chat", model, text, extra) + + +def is_budget_rejection(response: httpx.Response) -> bool: + return response.status_code != 200 and BUDGET_ERROR in response.text + + +def until_rejected( + gateway: Gateway, model: str, text: str, extra: Mapping[str, JsonValue] | None = None, seconds: float = 30 +) -> httpx.Response: + return eventually(lambda: chat(gateway, model, text, extra), is_budget_rejection, seconds=seconds) + + +def probe_text(label: str) -> str: + return f"{label} {PROVIDER_FAILURE} {uuid.uuid4().hex}" + + +def json_body(request: Request) -> dict[str, JsonValue]: + return JSON_OBJECT.validate_json(request.body) + + +def proxy_logs_mentioning(directory: Path, needle: str) -> tuple[str, ...]: + output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory))) + logs: Final = (path.read_text(errors="replace") for path in output.glob("owned-proxy-*.log")) + return tuple(log for log in logs if needle in log) diff --git a/tests/integration/routing/router_budgets/test_router_budget_endpoints.py b/tests/integration/routing/router_budgets/test_router_budget_endpoints.py new file mode 100644 index 00000000000..e05634b7e25 --- /dev/null +++ b/tests/integration/routing/router_budgets/test_router_budget_endpoints.py @@ -0,0 +1,258 @@ +"""What a caller sees from a router budget on every endpoint, streaming and not, through every client. + +Each (endpoint, client) cell owns one deployment capped at 0.01, so its first call (0.03125) crosses the cap. +The cell drives that first call to completion through the real client, checks the deployment spend in Redis +is exactly one call, then probes through the same client until the router rejects it. Probes carry +``PROVIDER_FAILURE`` so an admitted probe is a free upstream 500, never a charge. The held and abandoned +stream cells gate the upstream after the first SSE frame to observe when the charge lands. +""" + +from __future__ import annotations + +import asyncio +import threading +from collections.abc import Iterator +from dataclasses import dataclass +from typing import Final + +import anthropic +import httpx +import openai +import pytest +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.wire import Reply, Request, wire_server +from integration.routing.router_budgets import _rig as rig + +CLIENTS: Final = ("sdk_sync", "sdk_async", "httpx") +CELLS: Final = tuple((endpoint, client) for endpoint in rig.ENDPOINTS for client in CLIENTS) +HOLD_MARKER: Final = "router-budget-hold-stream" +ABANDON_MARKER: Final = "router-budget-abandon-stream" +HELD_GATE: Final = threading.Event() +ABANDON_GATE: Final = threading.Event() +TERMINAL_FRAME: Final = {"chat": "[DONE]", "messages": "message_stop", "responses": "response.completed"} + + +def _gated(reply: Reply, gate: threading.Event) -> Reply: + frames: Final = tuple(frame + b"\n\n" for frame in reply.body.split(b"\n\n") if frame) + return Reply(content_type=reply.content_type, chunks=frames, gate_after_first=gate) + + +def _respond(request: Request) -> Reply: + reply: Final = rig.upstream(request) + if HOLD_MARKER.encode() in request.body: + return _gated(reply, HELD_GATE) + if ABANDON_MARKER.encode() in request.body: + return _gated(reply, ABANDON_GATE) + return reply + + +def _model(endpoint: str, client: str) -> str: + return f"ep-{endpoint}-{client}".replace("_", "-") + + +@dataclass(frozen=True, slots=True) +class EndpointRig: + gateway: Gateway + budget: rig.BudgetRig + + +@pytest.fixture(scope="module") +def endpoints(tmp_path_factory: pytest.TempPathFactory) -> Iterator[EndpointRig]: + tmp_path: Final = tmp_path_factory.mktemp("budget-endpoints") + with gateway_from_environment() as gateway, wire_server(_respond) as upstream: + capped: Final = tuple( + rig.deployment( + name, + f"openai/{name}", + upstream.url, + model_id=name, + max_budget=0.01, + budget_duration="1d", + ) + for name in (*(_model(endpoint, client) for endpoint, client in CELLS), "held-stream", "abandoned-stream") + ) + config: Final = rig.write_config(tmp_path / "endpoints.yaml", capped) + with rig.budget_proxy(gateway, tmp_path, config) as budget: + yield EndpointRig(budget.gateway, budget) + + +@dataclass(frozen=True, slots=True) +class Outcome: + status: int + message: str + error_class: str + completed: bool + + +def _sdk_error(error: openai.APIStatusError | anthropic.APIStatusError) -> Outcome: + return Outcome(error.status_code, str(error), type(error).__name__, False) + + +def _openai_sync(base_url: str, key: str, endpoint: str, model: str, text: str) -> Outcome: + client: Final = openai.OpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0) + try: + if endpoint == "chat": + return Outcome( + 200, + client.chat.completions.create(model=model, messages=[{"role": "user", "content": text}]).id, + "", + True, + ) + if endpoint == "chat_stream": + chunks: Final = tuple( + client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": text}], + stream=True, + stream_options={"include_usage": True}, + ) + ) + return Outcome(200, chunks[-1].id, "", any(chunk.usage is not None for chunk in chunks)) + if endpoint == "responses": + return Outcome(200, client.responses.create(model=model, input=text).id, "", True) + events: Final = tuple(client.responses.create(model=model, input=text, stream=True)) + return Outcome(200, events[-1].type, "", events[-1].type == "response.completed") + except openai.APIStatusError as error: + return _sdk_error(error) + + +async def _openai_async(base_url: str, key: str, endpoint: str, model: str, text: str) -> Outcome: + client: Final = openai.AsyncOpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0) + try: + if endpoint == "chat": + completion: Final = await client.chat.completions.create( + model=model, messages=[{"role": "user", "content": text}] + ) + return Outcome(200, completion.id, "", True) + if endpoint == "chat_stream": + stream: Final = await client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": text}], + stream=True, + stream_options={"include_usage": True}, + ) + chunks: Final = tuple([chunk async for chunk in stream]) + return Outcome(200, chunks[-1].id, "", any(chunk.usage is not None for chunk in chunks)) + if endpoint == "responses": + response: Final = await client.responses.create(model=model, input=text) + return Outcome(200, response.id, "", True) + events_stream: Final = await client.responses.create(model=model, input=text, stream=True) + events: Final = tuple([event async for event in events_stream]) + return Outcome(200, events[-1].type, "", events[-1].type == "response.completed") + except openai.APIStatusError as error: + return _sdk_error(error) + finally: + await client.close() + + +def _anthropic_sync(base_url: str, key: str, endpoint: str, model: str, text: str) -> Outcome: + client: Final = anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0) + try: + if endpoint == "messages": + message: Final = client.messages.create( + model=model, max_tokens=16, messages=[{"role": "user", "content": text}] + ) + return Outcome(200, message.id, "", True) + events: Final = tuple( + client.messages.create( + model=model, max_tokens=16, messages=[{"role": "user", "content": text}], stream=True + ) + ) + return Outcome(200, events[-1].type, "", events[-1].type == "message_stop") + except anthropic.APIStatusError as error: + return _sdk_error(error) + + +async def _anthropic_async(base_url: str, key: str, endpoint: str, model: str, text: str) -> Outcome: + client: Final = anthropic.AsyncAnthropic(base_url=base_url, api_key=key, max_retries=0) + try: + if endpoint == "messages": + message: Final = await client.messages.create( + model=model, max_tokens=16, messages=[{"role": "user", "content": text}] + ) + return Outcome(200, message.id, "", True) + stream: Final = await client.messages.create( + model=model, max_tokens=16, messages=[{"role": "user", "content": text}], stream=True + ) + events: Final = tuple([event async for event in stream]) + return Outcome(200, events[-1].type, "", events[-1].type == "message_stop") + except anthropic.APIStatusError as error: + return _sdk_error(error) + finally: + await client.close() + + +def _raw(gateway: Gateway, endpoint: str, model: str, text: str) -> Outcome: + response: Final = rig.send(gateway, endpoint, model, text) + terminal: Final = TERMINAL_FRAME[endpoint.removesuffix("_stream")] + completed: Final = response.status_code == 200 and (not endpoint.endswith("_stream") or terminal in response.text) + return Outcome(response.status_code, response.text, "", completed) + + +def _invoke(gateway: Gateway, endpoint: str, client: str, model: str, text: str) -> Outcome: + base_url: Final = str(gateway.client.base_url).rstrip("/") + if client == "httpx": + return _raw(gateway, endpoint, model, text) + if endpoint.startswith("messages"): + if client == "sdk_sync": + return _anthropic_sync(base_url, gateway.key, endpoint, model, text) + return asyncio.run(_anthropic_async(base_url, gateway.key, endpoint, model, text)) + if client == "sdk_sync": + return _openai_sync(base_url, gateway.key, endpoint, model, text) + return asyncio.run(_openai_async(base_url, gateway.key, endpoint, model, text)) + + +@pytest.mark.parametrize(("endpoint", "client"), CELLS) +def test_a_deployment_over_budget_rejects_the_caller_with_429_after_one_charged_call( + endpoints: EndpointRig, endpoint: str, client: str +) -> None: + model: Final = _model(endpoint, client) + first: Final = _invoke(endpoints.gateway, endpoint, client, model, f"{model} first") + assert (first.status, first.completed) == (200, True), first + + assert endpoints.budget.settled(f"deployment_spend:{model}:1d", rig.CALL_COST) == rig.CALL_COST + + rejected: Final = eventually( + lambda: _invoke(endpoints.gateway, endpoint, client, model, rig.probe_text(model)), + lambda outcome: outcome.status != 500, + seconds=30, + ) + assert rejected.status == 429, rejected + assert rig.BUDGET_ERROR in rejected.message, rejected + assert f"model_id: {model}" in rejected.message, rejected + assert rejected.error_class == ("" if client == "httpx" else "RateLimitError"), rejected + assert endpoints.budget.redis_float(f"deployment_spend:{model}:1d") == rig.CALL_COST + + +def test_a_stream_is_charged_only_once_it_completes(endpoints: EndpointRig) -> None: + body: Final = rig.body_for("chat_stream", "held-stream", f"held {HOLD_MARKER}") + with endpoints.gateway.client.stream( + "POST", "/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {endpoints.gateway.key}"} + ) as stream: + lines: Final = stream.iter_lines() + assert next(lines).startswith("data: ") + mid_stream: Final = rig.chat(endpoints.gateway, "held-stream", rig.probe_text("held mid stream")) + assert mid_stream.status_code == 500, mid_stream.text + HELD_GATE.set() + rest: Final = tuple(lines) + assert "data: [DONE]" in rest + + assert endpoints.budget.settled("deployment_spend:held-stream:1d", rig.CALL_COST) == rig.CALL_COST + rig.until_rejected(endpoints.gateway, "held-stream", rig.probe_text("held after")) + + +def test_a_stream_the_client_abandons_midway_is_still_charged(endpoints: EndpointRig) -> None: + body: Final = rig.body_for("chat_stream", "abandoned-stream", f"abandoned {ABANDON_MARKER}") + with httpx.Client(base_url=str(endpoints.gateway.client.base_url), trust_env=False, timeout=15) as client: + with client.stream( + "POST", "/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {endpoints.gateway.key}"} + ) as stream: + assert next(stream.iter_lines()).startswith("data: ") + ABANDON_GATE.set() + + rig.until_rejected(endpoints.gateway, "abandoned-stream", rig.probe_text("abandoned after")) + charged: Final = eventually( + lambda: endpoints.budget.redis_float("deployment_spend:abandoned-stream:1d"), lambda spend: spend is not None + ) + assert charged is not None and 0 < charged <= rig.CALL_COST, charged + assert (charged / rig.PRICE).is_integer(), charged diff --git a/tests/integration/routing/router_budgets/test_router_budget_multi_instance.py b/tests/integration/routing/router_budgets/test_router_budget_multi_instance.py new file mode 100644 index 00000000000..b519ba267e2 --- /dev/null +++ b/tests/integration/routing/router_budgets/test_router_budget_multi_instance.py @@ -0,0 +1,222 @@ +"""Router budgets across processes: two two-worker proxies sharing one owned Redis, restarts, window resets, +a Redis outage mid burst and a burst that is all in flight before any spend lands. + +The burst cells hold every request at the upstream behind a barrier until the whole burst has arrived, so +every request passes the budget filter before the first success is logged; the result does not depend on +scheduling. Probes carry ``PROVIDER_FAILURE`` and are never charged. +""" + +from __future__ import annotations + +import threading +from collections.abc import Callable, Iterator +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.routing.router_budgets import _rig as rig +from pydantic import JsonValue + +BURST: Final = 12 +NEAR_CAP: Final = 3 * rig.CALL_COST +BARRIERS: Final = MappingProxyType( + {f"router-budget-burst-{label}": threading.Barrier(BURST) for label in ("accounting", "hard-cap")} +) +MIXED: Final = ("chat", "chat_stream", "messages", "responses") + + +def _respond(request: Request) -> Reply: + barrier: Final = next((barrier for marker, barrier in BARRIERS.items() if marker.encode() in request.body), None) + if barrier is not None: + barrier.wait(timeout=20) + return rig.upstream(request) + + +def _deployments(url: str) -> tuple[dict[str, JsonValue], ...]: + capped: Final = ( + ("shared-tiny", 0.01, "1d"), + ("restart-tiny", 0.01, "1d"), + ("window-tiny", 0.01, "3s"), + ("near-cap-accounting", NEAR_CAP, "1d"), + ("near-cap-hard-cap", NEAR_CAP, "1d"), + ("outage-tiny", 0.01, "1d"), + ("outage-roomy", 100, "1d"), + ) + return ( + *( + rig.deployment(name, f"hosted_vllm/{name}", url, model_id=name, max_budget=cap, budget_duration=window) + for name, cap, window in capped + ), + rig.deployment("roomy", "openai/roomy", url, model_id="roomy"), + ) + + +@dataclass(frozen=True, slots=True) +class FleetRig: + first: Gateway + second: Gateway + budget: rig.BudgetRig + upstream: Wire + config: Path + tmp_path: Path + environment: Gateway + + +def _config(tmp_path: Path, url: str) -> Path: + return rig.write_config( + tmp_path / "fleet.yaml", + _deployments(url), + provider_budget_config={"openai": {"budget_limit": 100, "time_period": "1d"}}, + ) + + +@pytest.fixture(scope="module") +def fleet(tmp_path_factory: pytest.TempPathFactory) -> Iterator[FleetRig]: + tmp_path: Final = tmp_path_factory.mktemp("budget-fleet") + with gateway_from_environment() as gateway, wire_server(_respond) as upstream, rig.redis_for(tmp_path) as cache: + config: Final = _config(tmp_path, upstream.url) + with ( + rig.proxy_on(gateway, tmp_path, config, cache) as first, + rig.proxy_on(gateway, tmp_path, config, cache) as second, + ): + yield FleetRig(first, second, rig.BudgetRig(first, cache), upstream, config, tmp_path, gateway) + + +def _fresh_send(gateway: Gateway, endpoint: str, model: str, text: str) -> httpx.Response: + with httpx.Client(base_url=str(gateway.client.base_url), trust_env=False, timeout=60) as client: + return Gateway(client, gateway.key, gateway.upstream_url).request( + "POST", rig.path_for(endpoint), rig.body_for(endpoint, model, text) + ) + + +def _burst(targets: tuple[Gateway, ...], model: str, text: Callable[[int], str]) -> tuple[httpx.Response, ...]: + with ThreadPoolExecutor(max_workers=BURST) as pool: + futures: Final = tuple( + pool.submit(_fresh_send, targets[index % len(targets)], MIXED[index % len(MIXED)], model, text(index)) + for index in range(BURST) + ) + return tuple(future.result() for future in futures) + + +def test_spend_on_one_proxy_is_enforced_by_its_peer(fleet: FleetRig) -> None: + first: Final = rig.chat(fleet.first, "shared-tiny", "shared first") + assert first.status_code == 200, first.text + + blocked: Final = rig.until_rejected(fleet.second, "shared-tiny", rig.probe_text("shared on peer")) + + assert "model_id: shared-tiny" in blocked.text + assert fleet.budget.settled("deployment_spend:shared-tiny:1d", rig.CALL_COST) == rig.CALL_COST + + +def test_a_concurrent_burst_across_both_proxies_is_counted_exactly_once(fleet: FleetRig) -> None: + responses: Final = _burst((fleet.first, fleet.second), "roomy", lambda index: f"roomy burst {index}") + + assert all(response.status_code == 200 for response in responses), [r.text for r in responses] + assert fleet.budget.settled("provider_spend:openai:1d", BURST * rig.CALL_COST) == BURST * rig.CALL_COST + for proxy in (fleet.first, fleet.second): + reported = eventually( + lambda proxy=proxy: object_value(object_value(proxy.get("/provider/budgets")["providers"])["openai"]), + lambda entry: entry["spend"] == BURST * rig.CALL_COST, + ) + assert reported["budget_limit"] == 100.0 + + +def test_a_burst_in_flight_past_the_cap_is_charged_in_full_and_then_blocked(fleet: FleetRig) -> None: + marker: Final = "router-budget-burst-accounting" + responses: Final = _burst((fleet.first, fleet.second), "near-cap-accounting", lambda index: f"{marker} {index}") + admitted: Final = sum(response.status_code == 200 for response in responses) + + assert admitted >= 3, [response.text for response in responses] + assert ( + fleet.budget.settled("deployment_spend:near-cap-accounting:1d", admitted * rig.CALL_COST) + == admitted * rig.CALL_COST + ) + for proxy in (fleet.first, fleet.second): + rig.until_rejected(proxy, "near-cap-accounting", rig.probe_text("near cap after burst")) + + +@pytest.mark.xfail( + strict=True, + reason="budgets are checked before the call and charged after it, so a concurrent burst overshoots the cap", +) +def test_a_burst_in_flight_never_admits_more_than_the_cap_allows(fleet: FleetRig) -> None: + marker: Final = "router-budget-burst-hard-cap" + responses: Final = _burst((fleet.first, fleet.second), "near-cap-hard-cap", lambda index: f"{marker} {index}") + + assert sum(response.status_code == 200 for response in responses) <= 3 + + +def test_an_exhausted_budget_survives_a_proxy_restart(fleet: FleetRig) -> None: + first: Final = rig.chat(fleet.first, "restart-tiny", "restart first") + assert first.status_code == 200, first.text + assert fleet.budget.settled("deployment_spend:restart-tiny:1d", rig.CALL_COST) == rig.CALL_COST + + with rig.proxy_on(fleet.environment, fleet.tmp_path, fleet.config, fleet.budget.redis) as restarted: + blocked: Final = rig.until_rejected(restarted, "restart-tiny", rig.probe_text("restart")) + + assert "model_id: restart-tiny" in blocked.text + + +def test_a_spent_budget_admits_traffic_again_after_its_window_resets(fleet: FleetRig) -> None: + first: Final = rig.chat(fleet.first, "window-tiny", "window first") + assert first.status_code == 200, first.text + rig.until_rejected(fleet.first, "window-tiny", rig.probe_text("window blocked")) + + def five_probes() -> tuple[int, ...]: + return tuple(rig.chat(fleet.first, "window-tiny", rig.probe_text("window reset")).status_code for _ in range(5)) + + assert eventually(five_probes, lambda statuses: statuses == (500,) * 5, seconds=30) == (500,) * 5 + + +def test_a_redis_outage_mid_burst_keeps_serving_and_reconciles_spend_exactly_once( + tmp_path: Path, fleet: FleetRig +) -> None: + with rig.redis_for(tmp_path) as cache: + config: Final = _config(tmp_path, fleet.upstream.url) + with rig.proxy_on( + fleet.environment, tmp_path, config, cache, extra={"REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1"} + ) as proxy: + budget: Final = rig.BudgetRig(proxy, cache) + warm: Final = rig.chat(proxy, "outage-roomy", "outage warm") + assert warm.status_code == 200, warm.text + assert budget.settled("deployment_spend:outage-roomy:1d", rig.CALL_COST) == rig.CALL_COST + cache.stop() + + responses: Final = _burst((proxy,), "outage-roomy", lambda index: f"outage burst {index}") + exhausted_during: Final = rig.chat(proxy, "outage-tiny", "outage exhaust") + + cache.start() + + assert all(response.status_code == 200 for response in responses), [r.text for r in responses] + assert exhausted_during.status_code == 200, exhausted_during.text + assert budget.settled("deployment_spend:outage-tiny:1d", rig.CALL_COST, seconds=30) == rig.CALL_COST + assert budget.settled("deployment_spend:outage-roomy:1d", BURST * rig.CALL_COST, seconds=30) == ( + BURST * rig.CALL_COST + ) + rig.until_rejected(proxy, "outage-tiny", rig.probe_text("outage after recovery")) + + +@pytest.mark.xfail( + strict=True, + reason="router budget spend lives only in Redis and process memory, so an emptied Redis resets every budget", +) +def test_an_exhausted_budget_survives_a_redis_restart_for_a_new_process(tmp_path: Path, fleet: FleetRig) -> None: + with rig.redis_for(tmp_path) as cache: + config: Final = _config(tmp_path, fleet.upstream.url) + with rig.proxy_on(fleet.environment, tmp_path, config, cache) as before: + exhausted: Final = rig.chat(before, "restart-tiny", "redis restart exhaust") + assert exhausted.status_code == 200, exhausted.text + spent: Final = rig.BudgetRig(before, cache).settled("deployment_spend:restart-tiny:1d", rig.CALL_COST) + assert spent == rig.CALL_COST + cache.stop() + cache.start() + with rig.proxy_on(fleet.environment, tmp_path, config, cache) as after: + probe: Final = rig.chat(after, "restart-tiny", rig.probe_text("redis restart")) + + assert rig.is_budget_rejection(probe), probe.text diff --git a/tests/integration/routing/router_budgets/test_router_deployment_budgets.py b/tests/integration/routing/router_budgets/test_router_deployment_budgets.py new file mode 100644 index 00000000000..e131ca8662f --- /dev/null +++ b/tests/integration/routing/router_budgets/test_router_deployment_budgets.py @@ -0,0 +1,253 @@ +"""Deployment budgets (``litellm_params.max_budget`` + ``budget_duration``) on a two-worker proxy. + +Each cell owns its deployments, so the ``deployment_spend:`` keys never collide. Probes carry +``PROVIDER_FAILURE`` and are never charged. Response caching is on so the cache-hit cell runs against the +same proxy; every other request carries unique text and never hits the cache. +""" + +from __future__ import annotations + +import uuid +from collections.abc import Iterator +from dataclasses import dataclass +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value, string_value +from integration._support.database import read_rows +from integration._support.wire import Wire, wire_server +from pydantic import JsonValue +from integration.routing.router_budgets import _rig as rig + +ISSUE_43214: Final = "https://github.com/BerriAI/litellm/issues/43214" + + +@dataclass(frozen=True, slots=True) +class DeploymentRig: + gateway: Gateway + upstream: Wire + budget: rig.BudgetRig + + +def _capped( + name: str, model_id: str, url: str, cap: float, duration: str | None = "1d", **extra: JsonValue +) -> dict[str, JsonValue]: + return rig.deployment( + name, + f"hosted_vllm/{model_id}", + url, + model_id=model_id, + max_budget=cap, + **({"budget_duration": duration} if duration is not None else {}), + **extra, + ) + + +@pytest.fixture(scope="module") +def deployments(tmp_path_factory: pytest.TempPathFactory) -> Iterator[DeploymentRig]: + tmp_path: Final = tmp_path_factory.mktemp("deployment-budgets") + with gateway_from_environment() as gateway, wire_server(rig.upstream) as upstream: + url: Final = upstream.url + path: Final = rig.write_config( + tmp_path / "deployments.yaml", + ( + _capped("dep-at-cap", "dep-at-cap", url, 2 * rig.CALL_COST), + _capped("dep-tiny", "dep-tiny", url, 0.01), + _capped("dep-zero", "dep-zero", url, 0), + _capped("dep-no-duration", "dep-no-duration", url, 0.01, None), + _capped("dep-pair", "dep-pair-capped", url, 0.01, order=1), + rig.deployment("dep-pair", "hosted_vllm/dep-pair-sibling", url, model_id="dep-pair-sibling", order=2), + _capped("dep-solo", "dep-solo", url, 0.01), + rig.deployment("dep-backup", "hosted_vllm/dep-backup", url, model_id="dep-backup"), + _capped("dep-message", "dep-message", url, 0.01), + _capped("dep-accounting", "dep-accounting", url, 2 * rig.CALL_COST), + _capped("dep-cached", "dep-cached", url, 2 * rig.CALL_COST), + ), + litellm_settings={ + "cache": True, + "cache_params": {"type": "redis", "host": "os.environ/REDIS_HOST", "port": "os.environ/REDIS_PORT"}, + }, + ) + config: Final = yaml.safe_load(path.read_text()) + config["router_settings"]["fallbacks"] = [{"dep-solo": ["dep-backup"]}] + path.write_text(yaml.safe_dump(config)) + with rig.budget_proxy(gateway, tmp_path, path) as budget: + yield DeploymentRig(budget.gateway, upstream, budget) + + +def test_spend_reaching_exactly_the_deployment_cap_blocks_the_next_request(deployments: DeploymentRig) -> None: + first: Final = rig.chat(deployments.gateway, "dep-at-cap", "deployment at cap first") + second: Final = rig.chat(deployments.gateway, "dep-at-cap", "deployment at cap second") + assert (first.status_code, second.status_code) == (200, 200), (first.text, second.text) + + blocked: Final = rig.until_rejected(deployments.gateway, "dep-at-cap", rig.probe_text("deployment at cap")) + + assert blocked.status_code == 429, blocked.text + assert "model_id: dep-at-cap" in blocked.text + assert deployments.budget.settled("deployment_spend:dep-at-cap:1d", 2 * rig.CALL_COST) == 2 * rig.CALL_COST + + +def test_spend_over_a_tiny_deployment_cap_rejects_with_429(deployments: DeploymentRig) -> None: + first: Final = rig.chat(deployments.gateway, "dep-tiny", "deployment tiny first") + assert first.status_code == 200, first.text + + blocked: Final = rig.until_rejected(deployments.gateway, "dep-tiny", rig.probe_text("deployment tiny")) + + assert blocked.status_code == 429, blocked.text + assert "Exceeded budget for deployment model_name: dep-tiny" in blocked.text + + +@pytest.mark.xfail(strict=True, reason=f"max_budget 0 is treated as unlimited: {ISSUE_43214}") +def test_a_zero_deployment_cap_rejects_the_first_request(deployments: DeploymentRig) -> None: + response: Final = rig.chat(deployments.gateway, "dep-zero", rig.probe_text("deployment zero")) + + assert rig.is_budget_rejection(response), response.text + + +@pytest.mark.xfail(strict=True, reason="max_budget without budget_duration loads without error and is never enforced") +def test_a_deployment_cap_without_budget_duration_is_still_enforced(deployments: DeploymentRig) -> None: + first: Final = rig.chat(deployments.gateway, "dep-no-duration", "deployment without duration") + assert first.status_code == 200, first.text + + rig.until_rejected(deployments.gateway, "dep-no-duration", rig.probe_text("deployment no duration"), seconds=10) + + +@pytest.mark.xfail(strict=True, reason="the deployment rejection prints budget_duration where the cap belongs") +def test_the_deployment_rejection_names_the_cap_it_crossed(deployments: DeploymentRig) -> None: + first: Final = rig.chat(deployments.gateway, "dep-message", "deployment message first") + assert first.status_code == 200, first.text + + blocked: Final = rig.until_rejected(deployments.gateway, "dep-message", rig.probe_text("deployment message")) + + assert f"{rig.CALL_COST} >= 0.01" in blocked.text, blocked.text + + +def test_traffic_converges_on_the_uncapped_sibling_deployment(deployments: DeploymentRig) -> None: + first: Final = rig.chat(deployments.gateway, "dep-pair", "dep pair first") + assert first.status_code == 200, first.text + assert first.headers["x-litellm-model-id"] == "dep-pair-capped" + + def batch() -> tuple[tuple[int, str | None], ...]: + return tuple( + (response.status_code, response.headers.get("x-litellm-model-id")) + for response in ( + rig.chat(deployments.gateway, "dep-pair", f"dep pair {uuid.uuid4().hex}") for _ in range(6) + ) + ) + + eventually(batch, lambda outcomes: all(outcome == (200, "dep-pair-sibling") for outcome in outcomes), seconds=30) + + eventually( + lambda: deployments.budget.redis_float("deployment_spend:dep-pair-capped:1d"), + lambda spend: spend is not None and spend >= rig.CALL_COST, + ) + assert deployments.budget.redis_float("deployment_spend:dep-pair-sibling:1d") is None + + +def test_an_over_budget_group_falls_back_to_the_configured_fallback_group(deployments: DeploymentRig) -> None: + first: Final = rig.chat(deployments.gateway, "dep-solo", "deployment solo first") + assert first.status_code == 200, first.text + assert first.headers["x-litellm-model-id"] == "dep-solo" + + fallen_back: Final = eventually( + lambda: rig.chat(deployments.gateway, "dep-solo", f"deployment solo {uuid.uuid4().hex}"), + lambda response: response.headers.get("x-litellm-model-id") == "dep-backup", + seconds=30, + ) + + assert fallen_back.status_code == 200, fallen_back.text + assert fallen_back.headers["x-litellm-attempted-fallbacks"] == "1" + + +def test_failed_calls_are_free_and_a_success_charges_exactly_its_spend_log_cost(deployments: DeploymentRig) -> None: + deployments.upstream.drain() + failures: Final = tuple( + rig.chat(deployments.gateway, "dep-accounting", rig.probe_text("accounting")) for _ in range(5) + ) + assert all(response.status_code == 500 for response in failures), [response.text for response in failures] + success: Final = rig.chat(deployments.gateway, "dep-accounting", "accounting success") + assert success.status_code == 200, success.text + + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, model_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (string_value(success.json()["id"]),) + ), + lambda found: len(found) == 1, + seconds=70, + ) + + assert rows[0]["spend"] == rig.CALL_COST + assert rows[0]["model_id"] == "dep-accounting" + assert deployments.budget.settled("deployment_spend:dep-accounting:1d", rig.CALL_COST) == rig.CALL_COST + probe: Final = rig.chat(deployments.gateway, "dep-accounting", rig.probe_text("accounting after")) + assert probe.status_code == 500, probe.text + assert len(deployments.upstream.drain()) == 7 + + +def test_a_response_cache_hit_does_not_charge_the_deployment_budget(deployments: DeploymentRig) -> None: + text: Final = f"cached {uuid.uuid4().hex}" + first: Final = rig.chat(deployments.gateway, "dep-cached", text) + assert first.status_code == 200, first.text + eventually( + lambda: deployments.budget.redis_float("deployment_spend:dep-cached:1d"), lambda spend: spend == rig.CALL_COST + ) + deployments.upstream.drain() + + hit: Final = eventually( + lambda: rig.chat(deployments.gateway, "dep-cached", text), + lambda response: response.headers.get("x-litellm-cache-key") is not None, + ) + + assert hit.json()["id"] == first.json()["id"] + assert deployments.upstream.drain() == () + second: Final = rig.chat(deployments.gateway, "dep-cached", "cached sibling request") + assert second.status_code == 200, second.text + assert deployments.budget.settled("deployment_spend:dep-cached:1d", 2 * rig.CALL_COST) == 2 * rig.CALL_COST + rig.until_rejected(deployments.gateway, "dep-cached", rig.probe_text("cached probe")) + + +def test_a_budgeted_deployment_added_and_raised_at_runtime_is_enforced(deployments: DeploymentRig) -> None: + model_id: Final = f"runtime-{uuid.uuid4().hex}" + model_name: Final = f"runtime-{uuid.uuid4().hex}" + created: Final = deployments.gateway.post( + "/model/new", + { + "model_name": model_name, + "litellm_params": { + "model": f"hosted_vllm/{model_id}", + "api_base": f"{deployments.upstream.url}/v1", + "api_key": "router-budget-provider-key", + "input_cost_per_token": rig.PRICE, + "output_cost_per_token": rig.PRICE, + "max_budget": 0.01, + "budget_duration": "1d", + }, + "model_info": {"id": model_id}, + }, + ) + assert object_value(created["model_info"])["id"] == model_id + try: + first: Final = eventually( + lambda: rig.chat(deployments.gateway, model_name, "runtime first"), + lambda response: response.status_code == 200, + seconds=30, + ) + assert first.headers["x-litellm-model-id"] == model_id + + rig.until_rejected(deployments.gateway, model_name, rig.probe_text("runtime")) + + raised: Final = deployments.gateway.request( + "POST", "/model/update", {"model_info": {"id": model_id}, "litellm_params": {"max_budget": 100}} + ) + assert raised.status_code == 200, raised.text + + def five_probes() -> tuple[int, ...]: + return tuple( + rig.chat(deployments.gateway, model_name, rig.probe_text("runtime raised")).status_code + for _ in range(5) + ) + + eventually(five_probes, lambda statuses: statuses == (500,) * 5, seconds=30) + finally: + deployments.gateway.post("/model/delete", {"id": model_id}) diff --git a/tests/integration/routing/router_budgets/test_router_provider_budgets.py b/tests/integration/routing/router_budgets/test_router_provider_budgets.py new file mode 100644 index 00000000000..7aecd7cf1c2 --- /dev/null +++ b/tests/integration/routing/router_budgets/test_router_provider_budgets.py @@ -0,0 +1,232 @@ +"""Provider budgets (``router_settings.provider_budget_config``) on a two-worker proxy. + +Every provider in the config belongs to exactly one cell, so the global ``provider_spend:`` keys +never collide and the cells run in any order. A probe carries ``PROVIDER_FAILURE``: when the filter admits +it the upstream answers 500 and nothing is charged, so probing never moves spend across the cap. The +Prometheus cell runs one worker because the registry is per process and the rig sets no multiprocess dir, +and the boot-failure cell runs one worker so the failed boot exits instead of respawning. +""" + +from __future__ import annotations + +from collections.abc import Iterator +from dataclasses import dataclass +import uuid +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value +from integration._support.wire import Wire, wire_server +from integration.routing.router_budgets import _rig as rig + +ISSUE_43214: Final = "https://github.com/BerriAI/litellm/issues/43214" + + +@dataclass(frozen=True, slots=True) +class ProviderRig: + gateway: Gateway + upstream: Wire + budget: rig.BudgetRig + + +@pytest.fixture(scope="module") +def providers(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ProviderRig]: + tmp_path: Final = tmp_path_factory.mktemp("provider-budgets") + with gateway_from_environment() as gateway, wire_server(rig.upstream) as upstream: + url: Final = upstream.url + config: Final = rig.write_config( + tmp_path / "providers.yaml", + ( + rig.deployment("at-cap", "openai/budget-at-cap", url, model_id="provider-at-cap"), + rig.deployment("tiny-cap", "hosted_vllm/budget-tiny", url, model_id="provider-tiny"), + rig.deployment("zero-cap", "deepseek/budget-zero", url, model_id="provider-zero"), + rig.deployment("no-limit", "groq/budget-no-limit", url, model_id="provider-no-limit"), + rig.deployment("no-period", "together_ai/budget-no-period", url, model_id="provider-no-period"), + rig.deployment("fw-solo", "fireworks_ai/budget-fw-solo", url, model_id="provider-fw-solo"), + rig.deployment("fw-pair", "fireworks_ai/budget-fw-pair", url, model_id="provider-fw-pair"), + rig.deployment("fw-pair", "lm_studio/budget-sibling", url, model_id="provider-sibling"), + rig.deployment("reported", "deepinfra/budget-reported", url, model_id="provider-reported"), + ), + provider_budget_config={ + "openai": {"budget_limit": 2 * rig.CALL_COST, "time_period": "1d"}, + "hosted_vllm": {"budget_limit": 0.01, "time_period": "1d"}, + "deepseek": {"budget_limit": 0, "time_period": "1d"}, + "groq": {"time_period": "1d"}, + "together_ai": {"budget_limit": 0.01}, + "fireworks_ai": {"budget_limit": 0.01, "time_period": "1d"}, + "deepinfra": {"budget_limit": 1, "time_period": "1d"}, + }, + ) + with rig.budget_proxy(gateway, tmp_path, config) as budget: + yield ProviderRig(budget.gateway, upstream, budget) + + +def _served_models(upstream: Wire) -> tuple[str, ...]: + return tuple( + str(rig.json_body(request)["model"]) + for request in upstream.drain() + if rig.PROVIDER_FAILURE.encode() not in request.body + ) + + +def test_spend_reaching_exactly_the_provider_cap_blocks_the_next_request(providers: ProviderRig) -> None: + providers.upstream.drain() + first: Final = rig.chat(providers.gateway, "at-cap", "provider at cap first") + second: Final = rig.chat(providers.gateway, "at-cap", "provider at cap second") + assert (first.status_code, second.status_code) == (200, 200), (first.text, second.text) + + blocked: Final = rig.until_rejected(providers.gateway, "at-cap", rig.probe_text("provider at cap")) + + assert blocked.status_code == 429, blocked.text + assert object_value(blocked.json()["error"])["message"] == ( + f"{rig.BUDGET_ERROR}: Exceeded budget for provider openai: {2 * rig.CALL_COST} >= {2 * rig.CALL_COST}\n" + ) + assert _served_models(providers.upstream) == ("budget-at-cap", "budget-at-cap") + assert providers.budget.settled("provider_spend:openai:1d", 2 * rig.CALL_COST) == 2 * rig.CALL_COST + + +def test_spend_over_a_tiny_provider_cap_rejects_with_429(providers: ProviderRig) -> None: + first: Final = rig.chat(providers.gateway, "tiny-cap", "provider tiny first") + assert first.status_code == 200, first.text + + blocked: Final = rig.until_rejected(providers.gateway, "tiny-cap", rig.probe_text("provider tiny")) + + assert blocked.status_code == 429, blocked.text + assert "Exceeded budget for provider hosted_vllm" in blocked.text + assert ">= 0.01" in blocked.text + + +@pytest.mark.xfail(strict=True, reason=f"budget_limit 0 is treated as unlimited: {ISSUE_43214}") +def test_a_zero_provider_cap_rejects_the_first_request(providers: ProviderRig) -> None: + response: Final = rig.chat(providers.gateway, "zero-cap", rig.probe_text("provider zero")) + + assert rig.is_budget_rejection(response), response.text + + +@pytest.mark.xfail( + strict=True, + reason=f"a provider entry without budget_limit drops every deployment of that provider: {ISSUE_43214}", +) +def test_a_provider_entry_without_budget_limit_leaves_the_provider_uncapped(providers: ProviderRig) -> None: + response: Final = rig.chat(providers.gateway, "no-limit", "provider without budget_limit") + + assert response.status_code == 200, response.text + + +@pytest.mark.xfail( + strict=True, + reason="a provider budget_limit without time_period is accepted at boot but never counted or enforced", +) +def test_a_provider_cap_without_time_period_is_still_enforced(providers: ProviderRig) -> None: + first: Final = rig.chat(providers.gateway, "no-period", "provider without time_period") + assert first.status_code == 200, first.text + + rig.until_rejected(providers.gateway, "no-period", rig.probe_text("provider no period"), seconds=10) + + +def test_traffic_moves_to_an_uncapped_sibling_once_the_provider_is_over_budget(providers: ProviderRig) -> None: + exhaust: Final = rig.chat(providers.gateway, "fw-solo", "fireworks spend") + assert exhaust.status_code == 200, exhaust.text + rig.until_rejected(providers.gateway, "fw-solo", rig.probe_text("fireworks solo")) + + def batch() -> tuple[tuple[int, str | None], ...]: + return tuple( + (response.status_code, response.headers.get("x-litellm-model-id")) + for response in (rig.chat(providers.gateway, "fw-pair", f"pair {index}") for index in range(6)) + ) + + served: Final = eventually( + batch, lambda outcomes: all(outcome == (200, "provider-sibling") for outcome in outcomes), seconds=30 + ) + + assert len(served) == 6 + + +def test_provider_budgets_endpoint_reports_the_spend_redis_holds(providers: ProviderRig) -> None: + before: Final = datetime.now(timezone.utc) + response: Final = rig.chat(providers.gateway, "reported", "reported spend") + assert response.status_code == 200, response.text + + report: Final = eventually( + lambda: object_value(object_value(providers.gateway.get("/provider/budgets")["providers"])["deepinfra"]), + lambda entry: entry["spend"] == rig.CALL_COST, + ) + + assert report["budget_limit"] == 1.0 + assert report["time_period"] == "1d" + assert providers.budget.redis_float("provider_spend:deepinfra:1d") == rig.CALL_COST + reset_at: Final = datetime.fromisoformat(str(report["budget_reset_at"])) + assert before < reset_at <= before + timedelta(days=1, minutes=1) + + +def test_a_null_provider_entry_fails_the_proxy_boot_with_a_named_error(tmp_path: Path) -> None: + provider: Final = f"null-provider-{uuid.uuid4().hex}" + with gateway_from_environment() as gateway, wire_server(rig.upstream) as upstream: + config: Final = rig.write_config( + tmp_path / "null-provider.yaml", + (rig.deployment("null-provider", "openai/null-provider", upstream.url, model_id=provider),), + provider_budget_config={provider: None}, + ) + with pytest.raises(AssertionError, match="Owned proxy exited before readiness"): + with rig.budget_proxy(gateway, tmp_path, config, workers=1): + pass + + assert len(rig.proxy_logs_mentioning(tmp_path, f"No budget config found for provider {provider}")) == 1 + + +@pytest.mark.xfail(strict=True, reason="GET /provider/budgets without provider_budget_config answers 500") +def test_provider_budgets_endpoint_without_provider_config_is_a_client_error(tmp_path: Path) -> None: + with gateway_from_environment() as gateway, wire_server(rig.upstream) as upstream: + config: Final = rig.write_config( + tmp_path / "deployment-only.yaml", + ( + rig.deployment( + "deployment-only", + "openai/deployment-only", + upstream.url, + model_id="deployment-only", + max_budget=1, + budget_duration="1d", + ), + ), + ) + with rig.budget_proxy(gateway, tmp_path, config) as budget: + response: Final[httpx.Response] = budget.gateway.request("GET", "/provider/budgets") + + assert 400 <= response.status_code < 500, response.text + + +def _remaining_budget(gateway: Gateway) -> float | None: + scrape: Final = gateway.request("GET", "/metrics/").text + prefix: Final = 'litellm_provider_remaining_budget_metric{api_provider="openai"} ' + return next((float(line.removeprefix(prefix)) for line in scrape.splitlines() if line.startswith(prefix)), None) + + +@pytest.mark.xfail( + strict=True, + reason="litellm_provider_remaining_budget_metric is set only while routing, so it lags one request behind spend", +) +def test_the_prometheus_remaining_budget_follows_spend_after_a_call(tmp_path: Path) -> None: + with gateway_from_environment() as gateway, wire_server(rig.upstream) as upstream: + config: Final = rig.write_config( + tmp_path / "prometheus.yaml", + (rig.deployment("metered", "openai/metered", upstream.url, model_id="metered"),), + provider_budget_config={"openai": {"budget_limit": 1, "time_period": "1d"}}, + litellm_settings={"callbacks": ["prometheus"]}, + ) + with rig.budget_proxy(gateway, tmp_path, config, workers=1) as budget: + response: Final = rig.chat(budget.gateway, "metered", "metered call") + assert response.status_code == 200, response.text + assert budget.settled("provider_spend:openai:1d", rig.CALL_COST) == rig.CALL_COST + + remaining: Final = eventually( + lambda: _remaining_budget(budget.gateway), + lambda value: value == 1 - rig.CALL_COST, + seconds=10, + return_last_on_timeout=True, + ) + + assert remaining == 1 - rig.CALL_COST diff --git a/tests/integration/routing/router_budgets/test_router_tag_budgets.py b/tests/integration/routing/router_budgets/test_router_tag_budgets.py new file mode 100644 index 00000000000..81cde020b7a --- /dev/null +++ b/tests/integration/routing/router_budgets/test_router_tag_budgets.py @@ -0,0 +1,127 @@ +"""Tag budgets (``litellm_settings.tag_budget_config``, an enterprise feature) on a two-worker proxy. + +Tags are per request, so one uncapped deployment serves every cell and each cell owns its tags. Probes carry +``PROVIDER_FAILURE`` and are never charged. +""" + +from __future__ import annotations + +from collections.abc import Iterator +from dataclasses import dataclass +from typing import Final + +import pytest +from integration._support.client import Gateway, gateway_from_environment +from integration._support.wire import wire_server +from integration.routing.router_budgets import _rig as rig +from pydantic import JsonValue + +ISSUE_43214: Final = "https://github.com/BerriAI/litellm/issues/43214" +MODEL: Final = "tagged" + + +@dataclass(frozen=True, slots=True) +class TagRig: + gateway: Gateway + budget: rig.BudgetRig + + +@pytest.fixture(scope="module") +def tags(tmp_path_factory: pytest.TempPathFactory) -> Iterator[TagRig]: + tmp_path: Final = tmp_path_factory.mktemp("tag-budgets") + with gateway_from_environment() as gateway, wire_server(rig.upstream) as upstream: + config: Final = rig.write_config( + tmp_path / "tags.yaml", + (rig.deployment(MODEL, "hosted_vllm/tagged", upstream.url, model_id="tagged"),), + tag_budget_config={ + "tag-at-cap": {"max_budget": 2 * rig.CALL_COST, "budget_duration": "1d"}, + "tag-tiny": {"max_budget": 0.01, "budget_duration": "1d"}, + "tag-zero": {"max_budget": 0, "budget_duration": "1d"}, + "tag-uncapped": {"budget_duration": "1d"}, + "tag-no-duration": {"max_budget": 0.01}, + "tag-roomy": {"max_budget": 10, "budget_duration": "1d"}, + "tag-multi": {"max_budget": 0.01, "budget_duration": "1d"}, + "tag-header": {"max_budget": 0.01, "budget_duration": "1d"}, + }, + ) + with rig.budget_proxy(gateway, tmp_path, config) as budget: + yield TagRig(budget.gateway, budget) + + +def _tagged(*names: str) -> dict[str, JsonValue]: + return {"metadata": {"tags": list(names)}} + + +def test_spend_over_a_tiny_tag_cap_rejects_requests_carrying_that_tag(tags: TagRig) -> None: + first: Final = rig.chat(tags.gateway, MODEL, "tag tiny first", _tagged("tag-tiny")) + assert first.status_code == 200, first.text + + blocked: Final = rig.until_rejected(tags.gateway, MODEL, rig.probe_text("tag tiny"), _tagged("tag-tiny")) + + assert blocked.status_code == 429, blocked.text + assert f"Exceeded budget for tag='tag-tiny', tag_spend={rig.CALL_COST}, tag_budget_limit=0.01" in blocked.text + + +def test_spend_reaching_exactly_the_tag_cap_blocks_the_next_request(tags: TagRig) -> None: + first: Final = rig.chat(tags.gateway, MODEL, "tag at cap first", _tagged("tag-at-cap")) + second: Final = rig.chat(tags.gateway, MODEL, "tag at cap second", _tagged("tag-at-cap")) + assert (first.status_code, second.status_code) == (200, 200), (first.text, second.text) + + blocked: Final = rig.until_rejected(tags.gateway, MODEL, rig.probe_text("tag at cap"), _tagged("tag-at-cap")) + + assert f"tag_spend={2 * rig.CALL_COST}, tag_budget_limit={2 * rig.CALL_COST}" in blocked.text + assert tags.budget.settled("tag_spend:tag-at-cap:1d", 2 * rig.CALL_COST) == 2 * rig.CALL_COST + + +@pytest.mark.xfail(strict=True, reason=f"max_budget 0 is treated as unlimited: {ISSUE_43214}") +def test_a_zero_tag_cap_rejects_the_first_request(tags: TagRig) -> None: + response: Final = rig.chat(tags.gateway, MODEL, rig.probe_text("tag zero"), _tagged("tag-zero")) + + assert rig.is_budget_rejection(response), response.text + + +def test_a_tag_without_max_budget_stays_uncapped(tags: TagRig) -> None: + responses: Final = tuple( + rig.chat(tags.gateway, MODEL, f"tag uncapped {index}", _tagged("tag-uncapped")) for index in range(3) + ) + + assert all(response.status_code == 200 for response in responses), [response.text for response in responses] + assert tags.budget.settled("tag_spend:tag-uncapped:1d", 3 * rig.CALL_COST) == 3 * rig.CALL_COST + + +@pytest.mark.xfail(strict=True, reason="a tag max_budget without budget_duration is never counted or enforced") +def test_a_tag_cap_without_budget_duration_is_still_enforced(tags: TagRig) -> None: + first: Final = rig.chat(tags.gateway, MODEL, "tag without duration", _tagged("tag-no-duration")) + assert first.status_code == 200, first.text + + rig.until_rejected(tags.gateway, MODEL, rig.probe_text("tag no duration"), _tagged("tag-no-duration"), seconds=10) + + +def test_two_tags_are_each_charged_once_and_only_the_exhausted_one_blocks(tags: TagRig) -> None: + first: Final = rig.chat(tags.gateway, MODEL, "two tags", _tagged("tag-roomy", "tag-multi")) + assert first.status_code == 200, first.text + + blocked: Final = rig.until_rejected( + tags.gateway, MODEL, rig.probe_text("two tags"), _tagged("tag-roomy", "tag-multi") + ) + + assert "tag='tag-multi'" in blocked.text and "tag='tag-roomy'" not in blocked.text, blocked.text + assert tags.budget.settled("tag_spend:tag-multi:1d", rig.CALL_COST) == rig.CALL_COST + assert tags.budget.settled("tag_spend:tag-roomy:1d", rig.CALL_COST) == rig.CALL_COST + roomy_only: Final = rig.chat(tags.gateway, MODEL, rig.probe_text("roomy only"), _tagged("tag-roomy")) + untagged: Final = rig.chat(tags.gateway, MODEL, rig.probe_text("untagged")) + assert (roomy_only.status_code, untagged.status_code) == (500, 500), (roomy_only.text, untagged.text) + + +def test_a_tag_spent_through_the_header_blocks_the_same_tag_in_body_metadata(tags: TagRig) -> None: + first: Final = tags.gateway.request( + "POST", + "/v1/chat/completions", + rig.body_for("chat", MODEL, "tag header"), + headers={"x-litellm-tags": "tag-header"}, + ) + assert first.status_code == 200, first.text + + blocked: Final = rig.until_rejected(tags.gateway, MODEL, rig.probe_text("tag header"), _tagged("tag-header")) + + assert "tag='tag-header'" in blocked.text