This commit is contained in:
ryan-crabbe-berri 2026-10-04 23:16:49 +08:00 • committed by GitHub
commit d197d13743
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1404 additions and 0 deletions

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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:<model_id>`` 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})

View file

@ -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:<provider>`` 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

View file

@ -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