From c13746c951811e9235e8bba7e4b078ce6a7c2ad5 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 05:31:09 +0000 Subject: [PATCH] fix(router): price a model group from the deployments that serve it (#44732) * fix(router): price a model group from the deployments that serve it A model group's info read composed the group's own model_group_alias entry into its deployments, a hop the router never takes: an alias is resolved exactly once at request time, so a group reached as an alias target is served by its own deployments. For the chain X -> T -> U the price read for X included U's deployments too, and the free-model budget waiver refused a free request to X on an over-budget key, while GET /model_group/info reported U's providers and price for X. The group info read now prices a group from the deployments routing serves it with: the ones named after it, the routing group of that name, or the wildcard route matching it when neither exists. The budget waiver, GET /model_group/info, the rate limiters, and the response headers all read the same set as routing. get_model_list keeps its behavior for every other caller. * test(router): give the paid fixtures explicit per-token prices * test(router): call the routed-group read by name so the router coverage gate sees it * test(integration): audit the alias chain budget waiver on every route, shape, and outage Thirty-four cells under the management group prove an over-budget key is served through an alias chain entry at the price of the deployment that serves it, on chat, responses, and messages, sync and streamed, through the OpenAI and Anthropic SDKs and raw httpx on both replicas, with the chain middle, the reverse chain, a cost-map priced middle, a ghost middle, wildcard and routing-group targets, malformed and hostile inputs, a cached reply, a repointed alias, a provider failure, and two chaos bursts (a killed worker, a scripted outage) --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/router.py | 50 +- .../router_code_coverage.py | 1 + .../test_alias_chain_budget_waiver.py | 1018 +++++++++++++++++ tests/unit/proxy/auth/test_auth_checks.py | 42 + tests/unit/test_router/test_router.py | 111 ++ 5 files changed, 1204 insertions(+), 18 deletions(-) create mode 100644 tests/integration/authorization/test_alias_chain_budget_waiver.py diff --git a/litellm/router.py b/litellm/router.py index 21709bde1f8..c57302c5645 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -11276,9 +11276,7 @@ class Router: configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None reasoning_efforts_initialized = False reasoning_efforts_unknown = False - model_list: Final = self.get_model_list(model_name=model_group) - if model_list is None: - return None + model_list: Final = self.get_model_list_of_routed_group(model_group) for model in model_list: is_match = False if ( @@ -12409,27 +12407,43 @@ class Router: returned_models.extend(self.get_model_list_from_model_alias(model_name=model_name)) returned_models.extend(self.get_model_list_from_routing_groups(model_name=model_name)) - if len(returned_models) == 0: # check if wildcard route - potential_wildcard_models: Final = self.pattern_router.get_deployments_by_pattern(model=model_name or "") - - ## check for team-specific wildcard models - if team_id is not None and team_id in self.team_pattern_routers: - potential_team_only_wildcard_models: Final = self.team_pattern_routers[ - team_id - ].get_deployments_by_pattern(model=model_name or "") - potential_wildcard_models.extend(potential_team_only_wildcard_models) - - if model_name is not None and potential_wildcard_models is not None: - for m in potential_wildcard_models: - deployment_typed_dict = DeploymentTypedDict(**m) - deployment_typed_dict["model_name"] = model_name - returned_models.append(deployment_typed_dict) + if len(returned_models) == 0 and model_name is not None: + returned_models.extend(self._get_wildcard_deployments(model_name=model_name, team_id=team_id)) if model_name is None: returned_models += self.model_list return returned_models + def _get_wildcard_deployments(self, model_name: str, team_id: str | None = None) -> list[DeploymentTypedDict]: + """ + The deployments of the wildcard routes matching model_name (the proxy-wide + ones, plus team_id's own when given), each emitted under model_name. + """ + team_router: Final = self.team_pattern_routers.get(team_id) if team_id is not None else None + matches: Final = [ + *self.pattern_router.get_deployments_by_pattern(model=model_name), + *(team_router.get_deployments_by_pattern(model=model_name) if team_router is not None else ()), + ] + return [{**DeploymentTypedDict(**m), "model_name": model_name} for m in matches] + + def get_model_list_of_routed_group(self, model_group: str) -> list[DeploymentTypedDict]: + """ + The deployments a request the router has already resolved to model_group is + served from: the ones named model_group, the routing group of that name, or + the wildcard route matching it when neither exists. + + Unlike get_model_list, model_group's own model_group_alias entry is not + followed. The router resolves an alias exactly once, so a group reached as + an alias target is served by its own deployments, never by a second hop: + in the chain X -> T -> U a request to X is served from T's deployments. + """ + named: Final = [ + *self._get_all_deployments(model_name=model_group), + *self.get_model_list_from_routing_groups(model_name=model_group), + ] + return named or self._get_wildcard_deployments(model_name=model_group) + def resolved_litellm_models(self, model_name: str, team_id: str | None = None) -> tuple[str, ...]: """The provider model strings `model_name` can actually be served by on this proxy. diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index e989c40d095..baa46aa5331 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -105,6 +105,7 @@ ignored_function_names = [ "_aanthropic_messages_retry_same_group", # Tested through the dropped-before-content retry tests in test_router.py "_aanthropic_messages_yield_recovered", # Tested through every mid-stream retry and fallback test in test_router.py "_anthropic_messages_policy_retries", # Tested through the retry budget precedence test in test_router.py + "_get_wildcard_deployments", # Tested through the get_model_list_of_routed_group wildcard test in test_router.py ] diff --git a/tests/integration/authorization/test_alias_chain_budget_waiver.py b/tests/integration/authorization/test_alias_chain_budget_waiver.py new file mode 100644 index 00000000000..175d4765a24 --- /dev/null +++ b/tests/integration/authorization/test_alias_chain_budget_waiver.py @@ -0,0 +1,1018 @@ +import json +import math +import os +import re +import signal +import threading +import uuid +from collections.abc import Callable, Generator, Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass +from functools import partial +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support.anthropic_thinking import JSON_OBJECT +from integration._support.client import ( + GATEWAY_LIMITS, + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, +) +from integration._support.database import read_rows +from integration._support.openai_wire import chat_reply +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from integration.authorization._hidden_alias_budget import ( + BUDGET, + BUDGET_EXCEEDED, + CHAT_REPLY, + GATEWAY_BURST, + NO_EXTRA, + PEER_BURST, + anthropic_client, + assert_free_row, + async_anthropic_client, + async_openai_client, + base_url, + error_type, + fresh_chat, + fresh_message, + fresh_response, + hidden, + install_aliases, + landed, + landed_all_once, + landed_once, + openai_client, + remove_aliases, + settle_candidate, + spend_marker, + upstream_hits, + upstream_requests, +) +from pydantic import JsonValue, TypeAdapter + +pytestmark: Final = pytest.mark.timeout(240) + +PAID_INPUT_COST_PER_TOKEN: Final = 0.001 +PAID_OUTPUT_COST_PER_TOKEN: Final = 0.002 +UPSTREAM_PROMPT_TOKENS: Final = 20 +UPSTREAM_COMPLETION_TOKENS: Final = 20 +PAID_REPLY_COST: Final = ( + UPSTREAM_PROMPT_TOKENS * PAID_INPUT_COST_PER_TOKEN + UPSTREAM_COMPLETION_TOKENS * PAID_OUTPUT_COST_PER_TOKEN +) +HEADROOM: Final = 5.0 +BURST_FAILURES: Final = 12 +LIVELINESS_PROBES: Final = 4 +BURST_KINDS: Final = ("chat", "chat-stream", "responses", "responses-stream", "messages", "messages-stream") +ROUTES: Final = (("chat", fresh_chat), ("responses", fresh_response), ("messages", fresh_message)) +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +WORKER_PIDS: Final = TypeAdapter(tuple[int, ...]) + + +@dataclass(frozen=True, slots=True) +class ChainRig: + gateway: Gateway + peer: Gateway + paid: str + mid_free: str + tail_free: str + mid_paid: str + unpriced: str + ghost_mid: str + failing_free: str + failing_provider_model: str + entry: str + hidden_entry: str + shown_entry: str + null_entry: str + reverse_entry: str + unpriced_entry: str + ghost_entry: str + failing_entry: str + + @property + def candidates(self) -> tuple[Gateway, Gateway]: + return (self.gateway, self.peer) + + +@dataclass(frozen=True, slots=True) +class GroupPrice: + input_cost_per_token: float + output_cost_per_token: float + providers: tuple[str, ...] + + +def _free_deployment(scenario: Scenario, label: str) -> str: + return scenario.model( + model=f"lemonade/{label}-{uuid.uuid4().hex[:8]}", input_cost_per_token=0, output_cost_per_token=0 + ) + + +def _paid_deployment(scenario: Scenario) -> str: + return scenario.model( + input_cost_per_token=PAID_INPUT_COST_PER_TOKEN, output_cost_per_token=PAID_OUTPUT_COST_PER_TOKEN + ) + + +def _wildcard_route(rig: ChainRig, scenario: Scenario, prefix: str, input_cost: float, output_cost: float) -> str: + created: Final = rig.gateway.post( + "/model/new", + { + "model_name": f"{prefix}/*", + "litellm_params": { + "model": "openai/*", + "api_key": "integration-provider-key", + "api_base": f"{rig.gateway.upstream_url}/v1", + "input_cost_per_token": input_cost, + "output_cost_per_token": output_cost, + }, + "model_info": {}, + }, + ) + identity: Final = str(object_value(created["model_info"])["id"]) + scenario.cleanups.callback(scenario.delete_model, identity) + return identity + + +def _deployment_id(name: str) -> str: + rows: Final = read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_name = %s', (name,)) + assert len(rows) == 1, rows + return str(rows[0]["model_id"]) + + +def _serving_deployments(candidate: Gateway, model: str, count: int) -> frozenset[str | None]: + def serving(marker: str) -> str | None: + response: Final = fresh_chat(candidate, model, candidate.key, marker) + return response.headers.get("x-litellm-model-id") if response.status_code == 200 else None + + with ThreadPoolExecutor(max_workers=count) as pool: + return frozenset(pool.map(serving, (f"route-{uuid.uuid4().hex}" for _ in range(count)))) + + +def _await_served_by(candidate: Gateway, name: str, deployment: str) -> None: + eventually( + lambda: _serving_deployments(candidate, name, 16), + lambda seen: seen == frozenset({deployment}), + seconds=120, + ) + + +def _settle_both(rig: ChainRig, model: str, key: str, status: int, extra: Mapping[str, JsonValue] = NO_EXTRA) -> None: + settle_candidate(rig.gateway, model, key, status, extra=extra) + settle_candidate(rig.peer, model, key, status, burst=PEER_BURST, extra=extra) + + +def _exhausted_key(rig: ChainRig, scenario: Scenario) -> str: + key: Final = scenario.key(max_budget=BUDGET) + first: Final = fresh_chat(rig.gateway, rig.paid, key, "exhaust-" + uuid.uuid4().hex) + assert first.status_code == 200, first.text + _settle_both(rig, rig.paid, key, BUDGET_EXCEEDED) + return key + + +def _refusal(response: httpx.Response) -> tuple[int, str]: + return response.status_code, error_type(response) + + +def _assert_refused( + rig: ChainRig, + candidate: Gateway, + key: str, + model: str, + prefix: str, + send: Callable[[Gateway, str, str, str], httpx.Response] = fresh_chat, +) -> None: + marker: Final = f"{prefix}-" + uuid.uuid4().hex + refused: Final = send(candidate, model, key, marker) + assert _refusal(refused) == (BUDGET_EXCEEDED, "budget_exceeded"), refused.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0 + + +def _assert_refused_on_every_route(rig: ChainRig, candidate: Gateway, key: str, model: str) -> None: + for route, send in ROUTES: + _assert_refused(rig, candidate, key, model, f"refused-{route}", send) + + +def _assert_provider_failure(response: httpx.Response) -> None: + assert response.status_code == 500, response.text + assert "Controlled provider failure" in response.text, response.text + assert "budget" not in response.text.lower(), response.text + + +def _assert_no_served_row(key: str, marker: str) -> None: + assert all(row["status"] != "success" for row in landed(key, marker)), landed(key, marker) + + +def _group_rows(candidate: Gateway, name: str) -> tuple[Mapping[str, JsonValue], ...]: + response: Final = candidate.request("GET", "/model_group/info", params={"model_group": name}) + assert response.status_code == 200, response.text + rows: Final = JSON_OBJECT.validate_json(response.content)["data"] + assert isinstance(rows, list), response.text + return tuple(object_value(row) for row in rows) + + +def _cost(value: JsonValue) -> float | None: + return float(str(value)) if value is not None else None + + +def _price_of(candidate: Gateway, name: str) -> GroupPrice: + rows: Final = _group_rows(candidate, name) + assert len(rows) == 1, rows + providers: Final = rows[0]["providers"] + assert isinstance(providers, list), rows + input_cost: Final = _cost(rows[0]["input_cost_per_token"]) + output_cost: Final = _cost(rows[0]["output_cost_per_token"]) + assert input_cost is not None and output_cost is not None, rows + return GroupPrice(input_cost, output_cost, tuple(str(provider) for provider in providers)) + + +def _assert_price(price: GroupPrice, input_cost: float, output_cost: float, providers: tuple[str, ...]) -> None: + assert (price.input_cost_per_token, price.output_cost_per_token) == (input_cost, output_cost), price + assert price.providers == providers, price + + +def _listed_input_prices(candidate: Gateway) -> Mapping[str, float | None]: + rows: Final = candidate.get("/model_group/info")["data"] + assert isinstance(rows, list), rows + return {str(object_value(row)["model_group"]): _cost(object_value(row)["input_cost_per_token"]) for row in rows} + + +def _listed_models(candidate: Gateway) -> frozenset[str]: + models: Final = candidate.get("/v1/models")["data"] + assert isinstance(models, list), models + return frozenset(str(object_value(entry)["id"]) for entry in models) + + +def _upstream_model_of(observed: tuple[str, ...], marker: str) -> str: + entries: Final = tuple(entry for entry in observed if marker in entry) + assert len(entries) == 1, entries + return str(object_value(JSON_OBJECT.validate_json(entries[0])["body"])["model"]) + + +def _script_failing(rig: ChainRig, statuses: tuple[int, ...]) -> None: + scripted: Final = httpx.post( + f"{rig.gateway.upstream_url}/__scripts/{rig.failing_provider_model}", + json={"statuses": list(statuses)}, + timeout=15, + trust_env=False, + ) + assert scripted.status_code == 200, scripted.text + + +def _clear_failing(rig: ChainRig) -> None: + cleared: Final = httpx.delete( + f"{rig.gateway.upstream_url}/__scripts/{rig.failing_provider_model}", timeout=15, trust_env=False + ) + assert cleared.status_code in (200, 404), cleared.text + + +def _await_rig(rig: ChainRig) -> None: + admin: Final = rig.gateway.key + for name in ( + rig.entry, + rig.hidden_entry, + rig.shown_entry, + rig.null_entry, + rig.reverse_entry, + rig.unpriced_entry, + rig.failing_entry, + ): + _settle_both(rig, name, admin, 200) + paid_id: Final = _deployment_id(rig.paid) + tail_id: Final = _deployment_id(rig.tail_free) + for candidate in rig.candidates: + _await_served_by(candidate, rig.mid_free, paid_id) + _await_served_by(candidate, rig.failing_free, paid_id) + _await_served_by(candidate, rig.mid_paid, tail_id) + _await_served_by(candidate, rig.unpriced, tail_id) + + +@contextmanager +def _chain_rig() -> Generator[ChainRig]: + suffix: Final = uuid.uuid4().hex[:12] + with ( + gateway_from_environment() as gateway, + httpx.Client( + base_url=os.environ["INTEGRATION_PEER_URL"], timeout=15, trust_env=False, limits=GATEWAY_LIMITS + ) as peer_client, + gateway.scenario() as scenario, + ): + failing_provider_model: Final = f"chain-failing-{suffix}" + rig: Final = ChainRig( + gateway=gateway, + peer=Gateway(peer_client, gateway.key, gateway.upstream_url), + paid=_paid_deployment(scenario), + mid_free=_free_deployment(scenario, "chain-mid"), + tail_free=_free_deployment(scenario, "chain-tail"), + mid_paid=_paid_deployment(scenario), + unpriced=scenario.model(), + ghost_mid=f"chain-ghost-mid-{suffix}", + failing_free=scenario.model( + model=f"lemonade/{failing_provider_model}", input_cost_per_token=0, output_cost_per_token=0 + ), + failing_provider_model=failing_provider_model, + entry=f"chain-entry-{suffix}", + hidden_entry=f"chain-hidden-entry-{suffix}", + shown_entry=f"chain-shown-entry-{suffix}", + null_entry=f"chain-null-entry-{suffix}", + reverse_entry=f"chain-reverse-entry-{suffix}", + unpriced_entry=f"chain-unpriced-entry-{suffix}", + ghost_entry=f"chain-ghost-entry-{suffix}", + failing_entry=f"chain-failing-entry-{suffix}", + ) + aliases: Final[Mapping[str, JsonValue]] = { + rig.entry: rig.mid_free, + rig.hidden_entry: hidden(rig.mid_free), + rig.shown_entry: {"model": rig.mid_free, "hidden": False}, + rig.null_entry: {"model": rig.mid_free, "hidden": None}, + rig.mid_free: rig.paid, + rig.reverse_entry: rig.mid_paid, + rig.mid_paid: rig.tail_free, + rig.unpriced_entry: rig.unpriced, + rig.unpriced: rig.tail_free, + rig.ghost_entry: rig.ghost_mid, + rig.ghost_mid: rig.paid, + rig.failing_entry: rig.failing_free, + rig.failing_free: rig.paid, + } + install_aliases(gateway, aliases) + scenario.cleanups.callback(remove_aliases, gateway, frozenset(aliases)) + _await_rig(rig) + yield rig + + +@pytest.fixture(scope="module") +def rig() -> Iterator[ChainRig]: + with _chain_rig() as built: + yield built + + +def test_exhausted_key_reaches_chain_entry_through_openai_chat(rig: ChainRig) -> None: + marker: Final = "chat-sync-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + with openai_client(rig.gateway, key) as client: + completion: Final = client.chat.completions.create( + model=rig.entry, + messages=[{"role": "user", "content": marker}], + extra_headers=spend_marker(marker), + ) + assert completion.choices[0].message.content == CHAT_REPLY, completion + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == completion.id, row + assert_free_row(row, rig.entry) + + +async def test_exhausted_key_reaches_chain_entry_through_streamed_openai_chat(rig: ChainRig) -> None: + marker: Final = "chat-stream-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + async with async_openai_client(rig.gateway, key) as client: + stream: Final = await client.chat.completions.create( + model=rig.entry, + messages=[{"role": "user", "content": marker}], + stream=True, + stream_options={"include_usage": True}, + extra_headers=spend_marker(marker), + ) + chunks: Final = tuple([chunk async for chunk in stream]) + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == CHAT_REPLY + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == chunks[0].id, row + assert_free_row(row, rig.entry) + + +def test_exhausted_key_reaches_chain_entry_through_openai_responses(rig: ChainRig) -> None: + marker: Final = "responses-sync-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + with openai_client(rig.gateway, key) as client: + response: Final = client.responses.create(model=rig.entry, input=marker, extra_headers=spend_marker(marker)) + assert response.output_text == CHAT_REPLY, response + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + assert_free_row(landed_once(key, marker), rig.entry) + + +async def test_exhausted_key_reaches_chain_entry_through_streamed_openai_responses(rig: ChainRig) -> None: + marker: Final = "responses-stream-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + async with async_openai_client(rig.gateway, key) as client: + stream: Final = await client.responses.create( + model=rig.entry, input=marker, stream=True, extra_headers=spend_marker(marker) + ) + events: Final = tuple([event async for event in stream]) + assert events[-1].type == "response.completed", events + assert "".join(event.delta for event in events if event.type == "response.output_text.delta") == CHAT_REPLY + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + assert_free_row(landed_once(key, marker), rig.entry) + + +def test_exhausted_key_reaches_chain_entry_through_anthropic_messages(rig: ChainRig) -> None: + marker: Final = "messages-sync-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + with anthropic_client(rig.gateway, key) as client: + message: Final = client.messages.create( + model=rig.entry, + max_tokens=16, + messages=[{"role": "user", "content": marker}], + extra_headers=spend_marker(marker), + ) + assert [block.text for block in message.content if block.type == "text"] == [CHAT_REPLY], message + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + assert_free_row(landed_once(key, marker), rig.entry) + + +async def test_exhausted_key_reaches_chain_entry_through_streamed_anthropic_messages(rig: ChainRig) -> None: + marker: Final = "messages-stream-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + async with ( + async_anthropic_client(rig.gateway, key) as client, + client.messages.stream( + model=rig.entry, + max_tokens=16, + messages=[{"role": "user", "content": marker}], + extra_headers=spend_marker(marker), + ) as stream, + ): + text: Final = "".join([piece async for piece in stream.text_stream]) + final: Final = await stream.get_final_message() + assert text == CHAT_REPLY, final + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + assert_free_row(landed_once(key, marker), rig.entry) + + +def _assert_raw_chat_served(rig: ChainRig, candidate: Gateway, entry: str, key: str, prefix: str) -> None: + marker: Final = f"{prefix}-" + uuid.uuid4().hex + response: Final = fresh_chat(candidate, entry, key, marker) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == CHAT_REPLY, response.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == response.json()["id"], row + assert_free_row(row, entry) + + +def test_exhausted_key_reaches_chain_entry_over_raw_http_on_both_replicas(rig: ChainRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + for candidate in rig.candidates: + _assert_raw_chat_served(rig, candidate, rig.entry, key, "raw") + + +@pytest.mark.parametrize("shape", ["string", "hidden_false", "hidden_null", "hidden_true"]) +def test_chain_entry_shape_keeps_the_waiver(rig: ChainRig, shape: str) -> None: + entries: Final = { + "string": rig.entry, + "hidden_false": rig.shown_entry, + "hidden_null": rig.null_entry, + "hidden_true": rig.hidden_entry, + } + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + for candidate in rig.candidates: + _assert_raw_chat_served(rig, candidate, entries[shape], key, f"shape-{shape}") + + +def _assert_chain_prices(rig: ChainRig, candidate: Gateway) -> None: + _assert_price(_price_of(candidate, rig.entry), 0.0, 0.0, ("lemonade",)) + _assert_price( + _price_of(candidate, rig.mid_free), PAID_INPUT_COST_PER_TOKEN, PAID_OUTPUT_COST_PER_TOKEN, ("openai",) + ) + _assert_price(_price_of(candidate, rig.paid), PAID_INPUT_COST_PER_TOKEN, PAID_OUTPUT_COST_PER_TOKEN, ("openai",)) + listed: Final = _listed_input_prices(candidate) + assert listed[rig.entry] == 0.0, listed + assert rig.hidden_entry not in listed, listed + assert _group_rows(candidate, rig.hidden_entry) == () + models: Final = _listed_models(candidate) + assert rig.entry in models and rig.hidden_entry not in models, models + + +def test_chain_entry_is_priced_by_the_hop_the_router_takes(rig: ChainRig) -> None: + for candidate in rig.candidates: + _assert_chain_prices(rig, candidate) + + +def test_chain_middle_by_name_is_judged_by_its_own_alias(rig: ChainRig) -> None: + billed_marker: Final = "middle-billed-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + exhausted: Final = _exhausted_key(rig, scenario) + for candidate in rig.candidates: + _assert_refused(rig, candidate, exhausted, rig.mid_free, "middle-refused") + _assert_refused(rig, candidate, exhausted, rig.paid, "middle-refused") + budgeted: Final = scenario.key(max_budget=HEADROOM) + served: Final = fresh_chat(rig.gateway, rig.mid_free, budgeted, billed_marker) + assert served.status_code == 200, served.text + assert served.headers["x-litellm-model-id"] == _deployment_id(rig.paid), dict(served.headers) + row: Final = landed_once(budgeted, billed_marker) + assert math.isclose(float(str(row["spend"])), PAID_REPLY_COST), row + assert row["model_group"] == rig.mid_free, row + + +def test_reverse_chain_is_priced_by_the_middle_it_is_served_from(rig: ChainRig) -> None: + free_marker: Final = "reverse-free-" + uuid.uuid4().hex + billed_marker: Final = "reverse-billed-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + exhausted: Final = _exhausted_key(rig, scenario) + for candidate in rig.candidates: + _assert_refused_on_every_route(rig, candidate, exhausted, rig.reverse_entry) + served: Final = fresh_chat(rig.gateway, rig.mid_paid, exhausted, free_marker) + assert served.status_code == 200, served.text + assert served.headers["x-litellm-model-id"] == _deployment_id(rig.tail_free), dict(served.headers) + assert_free_row(landed_once(exhausted, free_marker), rig.mid_paid) + budgeted: Final = scenario.key(max_budget=HEADROOM) + billed: Final = fresh_chat(rig.gateway, rig.reverse_entry, budgeted, billed_marker) + assert billed.status_code == 200, billed.text + assert billed.headers["x-litellm-model-id"] == _deployment_id(rig.mid_paid), dict(billed.headers) + row: Final = landed_once(budgeted, billed_marker) + assert math.isclose(float(str(row["spend"])), PAID_REPLY_COST), row + assert row["model_group"] == rig.reverse_entry, row + + +def test_chain_through_a_cost_map_priced_middle_stays_budgeted(rig: ChainRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + for candidate in rig.candidates: + _assert_refused_on_every_route(rig, candidate, key, rig.unpriced_entry) + + +def _assert_ghost_chain(rig: ChainRig, candidate: Gateway, key: str) -> None: + assert _group_rows(candidate, rig.ghost_entry) == () + assert rig.ghost_entry not in _listed_input_prices(candidate) + refused: Final = fresh_chat(candidate, rig.ghost_entry, key, "ghost-exhausted-" + uuid.uuid4().hex) + assert refused.status_code == BUDGET_EXCEEDED, refused.text + unroutable: Final = fresh_chat(candidate, rig.ghost_entry, rig.gateway.key, "ghost-admin-" + uuid.uuid4().hex) + assert unroutable.status_code == 400, unroutable.text + assert "no healthy deployments" in unroutable.text, unroutable.text + for path in ("/health/liveliness", "/model/info", "/v1/models"): + assert candidate.request("GET", path).status_code == 200, path + + +def test_ghost_chain_reports_no_group_and_stays_unroutable(rig: ChainRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + for candidate in rig.candidates: + _assert_ghost_chain(rig, candidate, key) + control: Final = fresh_chat(rig.gateway, rig.entry, key, "ghost-control-" + uuid.uuid4().hex) + assert control.status_code == 200, control.text + + +def test_provider_failure_behind_a_chain_entry_reaches_the_caller(rig: ChainRig) -> None: + failed_marker: Final = "provider-failure-" + uuid.uuid4().hex + recovered_marker: Final = "provider-recovered-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + scenario.cleanups.callback(_clear_failing, rig) + _script_failing(rig, (500,)) + _assert_provider_failure(fresh_chat(rig.gateway, rig.failing_entry, key, failed_marker)) + _clear_failing(rig) + recovered: Final = fresh_chat(rig.gateway, rig.failing_entry, key, recovered_marker) + assert recovered.status_code == 200, recovered.text + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert upstream_hits(observed, failed_marker) == 1 + assert upstream_hits(observed, recovered_marker) == 1 + assert_free_row(landed_once(key, recovered_marker), rig.failing_entry) + _assert_no_served_row(key, failed_marker) + + +def test_chain_through_a_free_wildcard_route_is_priced_by_that_route(rig: ChainRig) -> None: + prefix: Final = "wc" + uuid.uuid4().hex[:8] + middle: Final = f"{prefix}/gpt-4o-mini" + other: Final = f"{prefix}/gpt-4o" + entry: Final = "wildcard-entry-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + route: Final = _wildcard_route(rig, scenario, prefix, 0.0, 0.0) + install_aliases(rig.gateway, {entry: middle, middle: rig.paid}) + scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({entry, middle})) + paid_id: Final = _deployment_id(rig.paid) + for candidate in rig.candidates: + _await_served_by(candidate, entry, route) + _await_served_by(candidate, middle, paid_id) + key: Final = _exhausted_key(rig, scenario) + for candidate in rig.candidates: + _assert_served_by_route(candidate, other, key, route, "wildcard-other") + _assert_price(_price_of(candidate, other), 0.0, 0.0, ("openai",)) + _assert_price( + _price_of(candidate, middle), PAID_INPUT_COST_PER_TOKEN, PAID_OUTPUT_COST_PER_TOKEN, ("openai",) + ) + _assert_served_by_route(candidate, entry, key, route, "wildcard-served") + _assert_refused(rig, candidate, key, middle, "wildcard-middle") + + +def _assert_served_by_route(candidate: Gateway, model: str, key: str, route: str, prefix: str) -> None: + marker: Final = f"{prefix}-" + uuid.uuid4().hex + served: Final = fresh_chat(candidate, model, key, marker) + assert served.status_code == 200, served.text + assert served.headers["x-litellm-model-id"] == route, dict(served.headers) + assert_free_row(landed_once(key, marker), model) + + +def test_chain_through_a_priced_wildcard_route_stays_budgeted(rig: ChainRig) -> None: + prefix: Final = "pwc" + uuid.uuid4().hex[:8] + middle: Final = f"{prefix}/gpt-4o-mini" + entry: Final = "priced-wildcard-entry-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + route: Final = _wildcard_route(rig, scenario, prefix, PAID_INPUT_COST_PER_TOKEN, PAID_OUTPUT_COST_PER_TOKEN) + install_aliases(rig.gateway, {entry: middle, middle: rig.tail_free}) + scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({entry, middle})) + tail_id: Final = _deployment_id(rig.tail_free) + for candidate in rig.candidates: + _await_served_by(candidate, entry, route) + _await_served_by(candidate, middle, tail_id) + key: Final = _exhausted_key(rig, scenario) + for candidate in rig.candidates: + _assert_refused_on_every_route(rig, candidate, key, entry) + + +def _assert_hostile_group_queries_answered(rig: ChainRig, candidate: Gateway) -> None: + for value in ("7", "", "a,b", "x" * 5120): + assert isinstance(_group_rows(candidate, value), tuple), value + repeated: Final = httpx.get( + f"{base_url(candidate)}/model_group/info", + params=[("model_group", "7"), ("model_group", rig.entry)], + headers={"Authorization": f"Bearer {candidate.key}"}, + timeout=15, + trust_env=False, + ) + assert repeated.status_code == 200, repeated.text + assert isinstance(JSON_OBJECT.validate_json(repeated.content)["data"], list), repeated.text + assert _group_rows(candidate, rig.entry) == _group_rows(candidate, rig.entry) + unauthenticated: Final = httpx.get( + f"{base_url(candidate)}/model_group/info", params={"model_group": rig.entry}, timeout=15, trust_env=False + ) + assert unauthenticated.status_code == 401, unauthenticated.text + assert candidate.request("GET", "/health/liveliness").status_code == 200 + + +def test_hostile_model_group_queries_leave_the_proxy_healthy(rig: ChainRig) -> None: + for candidate in rig.candidates: + _assert_hostile_group_queries_answered(rig, candidate) + + +@pytest.mark.parametrize( + ("shape", "status"), + [ + ("int", BUDGET_EXCEEDED), + ("list", 400), + ("free_list", 400), + ("empty", BUDGET_EXCEEDED), + ("oversized", BUDGET_EXCEEDED), + ], +) +def test_malformed_model_value_never_takes_the_chain_waiver(rig: ChainRig, shape: str, status: int) -> None: + marker: Final = f"malformed-{shape}-" + uuid.uuid4().hex + models: Final[Mapping[str, JsonValue]] = { + "int": 5, + "list": [rig.entry], + "free_list": [rig.tail_free], + "empty": "", + "oversized": rig.entry + "x" * 5120, + } + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + response: Final = rig.gateway.request( + "POST", + "/v1/chat/completions", + {"model": models[shape], "messages": [{"role": "user", "content": marker}]}, + key=key, + ) + assert response.status_code == status, response.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0 + control: Final = fresh_chat(rig.gateway, rig.tail_free, key, "malformed-control-" + uuid.uuid4().hex) + assert control.status_code == 200, control.text + + +def test_unauthenticated_request_to_chain_entry_is_rejected(rig: ChainRig) -> None: + marker: Final = "unauthenticated-" + uuid.uuid4().hex + response: Final = httpx.post( + f"{base_url(rig.gateway)}/v1/chat/completions", + json={"model": rig.entry, "messages": [{"role": "user", "content": marker}]}, + timeout=60, + trust_env=False, + ) + assert response.status_code == 401, response.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0 + + +def _duplicate_model_post(rig: ChainRig, key: str, first: str, last: str, marker: str) -> httpx.Response: + messages: Final = json.dumps([{"role": "user", "content": marker}]) + return httpx.post( + f"{base_url(rig.gateway)}/v1/chat/completions", + content=f'{{"model": {json.dumps(first)}, "model": {json.dumps(last)}, "messages": {messages}}}'.encode(), + headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json", **spend_marker(marker)}, + timeout=60, + trust_env=False, + ) + + +def test_duplicate_model_field_is_judged_by_its_last_value(rig: ChainRig) -> None: + free_marker: Final = "duplicate-free-" + uuid.uuid4().hex + paid_marker: Final = "duplicate-paid-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + served: Final = _duplicate_model_post(rig, key, rig.paid, rig.entry, free_marker) + assert served.status_code == 200, served.text + refused: Final = _duplicate_model_post(rig, key, rig.entry, rig.paid, paid_marker) + assert refused.status_code == BUDGET_EXCEEDED, refused.text + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert upstream_hits(observed, free_marker) == 1 + assert upstream_hits(observed, paid_marker) == 0 + assert_free_row(landed_once(key, free_marker), rig.entry) + + +def test_key_restricted_to_the_chain_entry_keeps_the_waiver(rig: ChainRig) -> None: + marker: Final = "restricted-served-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + outsider: Final = scenario.key(models=[rig.tail_free, rig.paid], max_budget=BUDGET) + denied: Final = fresh_chat(rig.gateway, rig.entry, outsider, "restricted-denied-" + uuid.uuid4().hex) + assert _refusal(denied) == (403, "key_model_access_denied"), denied.text + insider: Final = scenario.key(models=[rig.entry, rig.paid], max_budget=BUDGET) + first: Final = fresh_chat(rig.gateway, rig.paid, insider, "restricted-exhaust-" + uuid.uuid4().hex) + assert first.status_code == 200, first.text + _settle_both(rig, rig.paid, insider, BUDGET_EXCEEDED) + served: Final = fresh_chat(rig.gateway, rig.entry, insider, marker) + assert served.status_code == 200, served.text + assert_free_row(landed_once(insider, marker), rig.entry) + + +def test_repointing_the_chain_middle_follows_the_served_deployment(rig: ChainRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + scenario.cleanups.callback(install_aliases, rig.gateway, {rig.entry: rig.mid_free, rig.mid_free: rig.paid}) + install_aliases(rig.gateway, {rig.mid_free: rig.tail_free}) + _settle_both(rig, rig.mid_free, key, 200) + _settle_both(rig, rig.entry, key, 200) + install_aliases(rig.gateway, {rig.mid_free: rig.paid}) + _settle_both(rig, rig.mid_free, key, BUDGET_EXCEEDED) + _settle_both(rig, rig.entry, key, 200) + install_aliases(rig.gateway, {rig.entry: rig.paid}) + _settle_both(rig, rig.entry, key, BUDGET_EXCEEDED) + install_aliases(rig.gateway, {rig.entry: rig.mid_free}) + _settle_both(rig, rig.entry, key, 200) + + +def test_exhausted_key_is_served_a_cached_reply_through_chain_entry(rig: ChainRig) -> None: + marker: Final = "cache-twin-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + first: Final = fresh_chat(rig.gateway, rig.entry, key, marker) + assert first.status_code == 200, first.text + second: Final = fresh_chat(rig.gateway, rig.entry, key, marker) + assert second.status_code == 200, second.text + first_id: Final = str(JSON_OBJECT.validate_json(first.content)["id"]) + assert str(JSON_OBJECT.validate_json(second.content)["id"]) == first_id, second.text + assert "x-litellm-cache-key" in second.headers, dict(second.headers) + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + rows: Final = eventually(lambda: landed(key, marker), lambda found: len(found) >= 2, seconds=70) + assert len(rows) == 2, rows + assert {str(row["request_id"]).split("_cache_hit")[0] for row in rows} == {first_id}, rows + assert sum(1 for row in rows if "_cache_hit" in str(row["request_id"])) == 1, rows + for row in rows: + assert_free_row(row, rig.entry) + + +@dataclass(frozen=True, slots=True) +class BurstCall: + candidate: Gateway + kind: str + marker: str + + +def _burst_calls(candidate: Gateway, count: int, prefix: str) -> tuple[BurstCall, ...]: + return tuple( + BurstCall(candidate, BURST_KINDS[index % len(BURST_KINDS)], f"{prefix}-" + uuid.uuid4().hex) + for index in range(count) + ) + + +def _send_kind(call: BurstCall, model: str, key: str) -> httpx.Response: + stream: Final[Mapping[str, JsonValue]] = {"stream": True} + match call.kind: + case "chat": + return fresh_chat(call.candidate, model, key, call.marker) + case "chat-stream": + return fresh_chat(call.candidate, model, key, call.marker, stream) + case "responses": + return fresh_response(call.candidate, model, key, call.marker) + case "responses-stream": + return fresh_response(call.candidate, model, key, call.marker, stream) + case "messages": + return fresh_message(call.candidate, model, key, call.marker) + case "messages-stream": + return fresh_message(call.candidate, model, key, call.marker, stream) + case _: + raise AssertionError(call.kind) + + +def test_chain_entry_burst_lands_every_served_id_once(rig: ChainRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + calls: Final = _burst_calls(rig.gateway, GATEWAY_BURST, "burst") + _burst_calls(rig.peer, PEER_BURST, "burst") + with ThreadPoolExecutor(max_workers=len(calls)) as pool: + responses: Final = dict( + zip(calls, pool.map(partial(_send_kind, model=rig.entry, key=key), calls), strict=True) + ) + failed: Final = { + call.marker: response.status_code for call, response in responses.items() if response.status_code != 200 + } + assert failed == {}, failed + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert all(upstream_hits(observed, call.marker) == 1 for call in calls), observed + for row in landed_all_once(key, frozenset(call.marker for call in calls)): + assert_free_row(row, rig.entry) + + +def _base_config() -> Mapping[str, JsonValue]: + return JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + + +def _routing_group_config(rig: ChainRig, directory: Path, group: str, member: str) -> Path: + base: Final = _base_config() + model_list: Final = base["model_list"] + assert isinstance(model_list, list), model_list + deployment: Final[JsonValue] = { + "model_name": member, + "litellm_params": { + "model": f"lemonade/{member}", + "api_key": "integration-provider-key", + "api_base": f"{rig.gateway.upstream_url}/v1", + "input_cost_per_token": 0, + "output_cost_per_token": 0, + }, + } + routing_groups: Final[JsonValue] = [{"group_name": group, "models": [member], "routing_strategy": "simple-shuffle"}] + path: Final = directory / "routing-group-chain.yaml" + path.write_text( + yaml.safe_dump( + { + **base, + "model_list": [*model_list, deployment], + "router_settings": {**object_value(base["router_settings"]), "routing_groups": routing_groups}, + } + ) + ) + return path + + +@pytest.mark.timeout(480) +def test_chain_through_a_routing_group_keeps_the_waiver(rig: ChainRig, tmp_path: Path) -> None: + suffix: Final = uuid.uuid4().hex[:8] + group: Final = f"rg-chain-{suffix}" + member: Final = f"rg-member-{suffix}" + entry: Final = f"rg-entry-{suffix}" + marker: Final = "routing-group-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + install_aliases(rig.gateway, {entry: group}) + scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({entry})) + key: Final = _exhausted_key(rig, scenario) + with owned_proxy( + rig.gateway, tmp_path, {}, config=_routing_group_config(rig, tmp_path, group, member) + ) as candidate: + settle_candidate(candidate, entry, key, 200, seconds=120) + _assert_price(_price_of(candidate, entry), 0.0, 0.0, ("lemonade",)) + served: Final = fresh_chat(candidate, entry, key, marker) + assert served.status_code == 200, served.text + assert _upstream_model_of(upstream_requests(rig.gateway.upstream_url), marker) == member + assert_free_row(landed_once(key, marker), entry) + + +def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]: + text: Final = log.read_text() + return WORKER_PIDS.validate_python(STARTED_WORKER.findall(text)), text.count("Application startup complete.") + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +def _marker_of(request: Request) -> str: + messages: Final = JSON_OBJECT.validate_json(request.body)["messages"] + assert isinstance(messages, list) and messages, request.body + return str(object_value(messages[0])["content"]) + + +def _response_or_none(send: Callable[[], httpx.Response]) -> httpx.Response | None: + try: + return send() + except httpx.TransportError: + return None + + +@pytest.mark.timeout(600) +def test_killing_one_worker_mid_burst_keeps_the_chain_entry_served(rig: ChainRig, tmp_path: Path) -> None: + suffix: Final = uuid.uuid4().hex[:8] + entry: Final = f"chaos-entry-{suffix}" + release: Final = threading.Event() + held_markers: Final[SimpleQueue[str]] = SimpleQueue() + + def respond(request: Request) -> Reply: + if request.method != "POST" or not request.target.endswith("/chat/completions"): + return Reply(status=404) + marker: Final = _marker_of(request) + if marker.startswith("hold-"): + held_markers.put(marker) + assert release.wait(timeout=120), "The burst was never released" + model: Final = str(JSON_OBJECT.validate_json(request.body)["model"]) + return chat_reply(f"chatcmpl-{marker}", model, CHAT_REPLY, stream=False) + + with wire_server(respond) as wire, rig.gateway.scenario() as scenario: + chaos_free: Final = scenario.model( + model=f"lemonade/chaos-free-{suffix}", + api_base=f"{wire.url}/v1", + input_cost_per_token=0, + output_cost_per_token=0, + ) + install_aliases(rig.gateway, {entry: chaos_free, chaos_free: rig.paid}) + scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({entry, chaos_free})) + key: Final = _exhausted_key(rig, scenario) + with owned_proxy_process(rig.gateway, tmp_path, {}, workers=2) as owned: + candidate: Final = owned.gateway + workers, _ = eventually( + lambda: _worker_startups(owned.log), + lambda found: len(found[0]) == 2 and found[1] == 2, + seconds=120, + ) + settle_candidate(candidate, entry, key, 200, seconds=120) + held_calls: Final = tuple("hold-" + uuid.uuid4().hex for _ in range(GATEWAY_BURST)) + with ThreadPoolExecutor(max_workers=GATEWAY_BURST) as pool: + pending: Final = { + marker: pool.submit(_response_or_none, partial(fresh_chat, candidate, entry, key, marker)) + for marker in held_calls + } + eventually(held_markers.qsize, lambda size: size == GATEWAY_BURST, seconds=60) + held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == GATEWAY_BURST, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + responses: Final = {marker: future.result() for marker, future in pending.items()} + served: Final = frozenset( + marker for marker, response in responses.items() if response is not None and response.status_code == 200 + ) + lost: Final = frozenset(responses) - served + assert len(served) == held_by[survivor_pid], (held_by, len(served), len(lost)) + follow_up_marker: Final = "after-kill-" + uuid.uuid4().hex + follow_up: Final = fresh_chat(candidate, entry, key, follow_up_marker) + assert follow_up.status_code == 200, follow_up.text + eventually( + lambda: _worker_startups(owned.log), + lambda found: len(found[0]) == 3 and found[1] == 3, + seconds=180, + ) + for row in landed_all_once(key, served | {follow_up_marker}): + assert_free_row(row, entry) + for marker in lost: + _assert_no_served_row(key, marker) + + +def test_chain_entry_burst_with_provider_outage_lands_every_served_id_once(rig: ChainRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = _exhausted_key(rig, scenario) + scenario.cleanups.callback(_clear_failing, rig) + _script_failing(rig, (500,) * BURST_FAILURES + (200,) * (GATEWAY_BURST - BURST_FAILURES)) + calls: Final = _burst_calls(rig.gateway, GATEWAY_BURST, "outage") + with ThreadPoolExecutor(max_workers=len(calls) + LIVELINESS_PROBES) as pool: + probes: Final = tuple( + pool.submit(httpx.get, f"{base_url(rig.gateway)}/health/liveliness", timeout=60, trust_env=False) + for _ in range(LIVELINESS_PROBES) + ) + responses: Final = dict( + zip(calls, pool.map(partial(_send_kind, model=rig.failing_entry, key=key), calls), strict=True) + ) + failed: Final = {call: response for call, response in responses.items() if response.status_code != 200} + assert len(failed) == BURST_FAILURES, {call.kind: response.status_code for call, response in failed.items()} + for response in failed.values(): + _assert_provider_failure(response) + for probe in probes: + assert probe.result().status_code == 200, probe.result().text + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert all(upstream_hits(observed, call.marker) == 1 for call in calls), observed + served_markers: Final = frozenset(call.marker for call in calls if call not in failed) + assert len(served_markers) == GATEWAY_BURST - BURST_FAILURES + for row in landed_all_once(key, served_markers): + assert_free_row(row, rig.failing_entry) + for call in failed: + _assert_no_served_row(key, call.marker) diff --git a/tests/unit/proxy/auth/test_auth_checks.py b/tests/unit/proxy/auth/test_auth_checks.py index 448211978d1..b698a3d02fa 100644 --- a/tests/unit/proxy/auth/test_auth_checks.py +++ b/tests/unit/proxy/auth/test_auth_checks.py @@ -32,6 +32,7 @@ from litellm.proxy._types import ( from litellm.proxy.utils import PrismaClient from litellm.proxy.auth.auth_checks import ( can_team_access_model, + _is_model_cost_zero, _virtual_key_soft_budget_check, _team_soft_budget_check, ) @@ -1595,3 +1596,44 @@ async def test_get_user_object_cache_miss_emits_exactly_one_postgres_get_user_ob assert result is not None and result.user_id == user_id assert await _db_service_call_types(db_success_hook) == ("get_user_object",) assert db_success_hook.await_args_list[0].kwargs["parent_otel_span"] == "auth-span" + + +@pytest.mark.parametrize("entry_first", [True, False]) +def test_is_model_cost_zero_judges_an_alias_chain_by_the_deployment_its_entry_routes_to( + monkeypatch: pytest.MonkeyPatch, entry_first: bool +) -> None: + """chain-entry resolves one hop to local-free and is served by local-free's own free + deployment, so an over-budget key is waived for it; local-free by name resolves to paid-gpt + and stays enforced. local-free's alias is a hop the router never takes for chain-entry, and + the per-name verdict cache must not let either name's verdict leak into the other's.""" + monkeypatch.setattr(litellm, "model_cost", dict(litellm.model_cost)) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "local-free", + "litellm_params": { + "model": "ollama/qwen3:0.6b", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0, + "output_cost_per_token": 0, + }, + }, + { + "model_name": "paid-gpt", + "litellm_params": { + "model": "gpt-4o", + "api_key": "fake", + "input_cost_per_token": 3e-06, + "output_cost_per_token": 1.5e-05, + }, + }, + ], + model_group_alias={"chain-entry": "local-free", "local-free": "paid-gpt"}, + ) + expected: Final = {"chain-entry": True, "local-free": False, "paid-gpt": False} + order: Final = ("chain-entry", "local-free", "paid-gpt") if entry_first else ("local-free", "paid-gpt", "chain-entry") + + verdicts: Final = {name: _is_model_cost_zero(model=name, llm_router=router) for name in order} + + assert verdicts == expected + assert {name: _is_model_cost_zero(model=name, llm_router=router) for name in order} == expected diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 053dd730dab..3d9ab07350e 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -64,6 +64,7 @@ from litellm.types.router import ( Deployment, DeploymentTypedDict, LiteLLM_Params, + ModelGroupInfo, ModelInfo, PreRoutingHookResponse, RetryPolicy, @@ -2333,6 +2334,116 @@ def test_update_settings_model_group_alias_drops_cached_group_info(): assert after.input_cost_per_token is not None and after.input_cost_per_token > 0 +_PAID_INPUT_COST_PER_TOKEN: Final = 3e-06 +_PAID_OUTPUT_COST_PER_TOKEN: Final = 1.5e-05 + + +def _free_ollama_deployment(model_name: str) -> dict: + return { + "model_name": model_name, + "litellm_params": { + "model": "ollama/qwen3:0.6b", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0, + "output_cost_per_token": 0, + }, + } + + +def _paid_openai_deployment(model_name: str, model: str) -> dict: + return { + "model_name": model_name, + "litellm_params": { + "model": model, + "api_key": "fake", + "input_cost_per_token": _PAID_INPUT_COST_PER_TOKEN, + "output_cost_per_token": _PAID_OUTPUT_COST_PER_TOKEN, + }, + } + + +def _assert_priced(info: ModelGroupInfo | None, provider: str) -> None: + assert info is not None + assert info.providers == [provider] + assert info.input_cost_per_token == _PAID_INPUT_COST_PER_TOKEN + assert info.output_cost_per_token == _PAID_OUTPUT_COST_PER_TOKEN + + +def _assert_free(info: ModelGroupInfo | None, provider: str) -> None: + assert info is not None + assert info.providers == [provider] + assert info.input_cost_per_token == 0 + assert info.output_cost_per_token == 0 + + +def test_get_model_group_info_prices_an_alias_chain_from_the_group_it_routes_to(): + """chain-entry resolves one hop to local-free and is served by local-free's own + deployment, so its price is that deployment's; local-free's own alias to gpt-priced + is a hop the router takes only for a request to local-free by name.""" + router = Router( + model_list=[ + _free_ollama_deployment("local-free"), + _paid_openai_deployment("gpt-priced", "gpt-4o"), + ], + model_group_alias={"chain-entry": "local-free", "local-free": "gpt-priced"}, + ) + + _assert_free(router.get_model_group_info(model_group="chain-entry"), "ollama") + _assert_priced(router.get_model_group_info(model_group="local-free"), "openai") + + +def test_get_model_group_info_prices_an_alias_chain_from_the_wildcard_route_serving_it(): + """When the routed group has no deployment of its own, the wildcard route matching it + serves the request, so the price is the wildcard's and never the routed group's own alias + target's.""" + router = Router( + model_list=[ + _paid_openai_deployment("openai/*", "openai/*"), + _free_ollama_deployment("local-free"), + ], + model_group_alias={"wildcard-entry": "openai/gpt-4o", "openai/gpt-4o": "local-free"}, + ) + + _assert_priced(router.get_model_group_info(model_group="wildcard-entry"), "openai") + _assert_free(router.get_model_group_info(model_group="openai/gpt-4o"), "ollama") + + +def _served_models(deployments: list[DeploymentTypedDict] | None) -> list[str]: + return [deployment["litellm_params"]["model"] for deployment in deployments or ()] + + +def test_get_model_list_of_routed_group_reads_the_groups_own_deployments_only(): + """The router resolves an alias once, so a group reached as an alias target is served by + its own deployments. get_model_list composes the group's own alias target too, the hop a + request to that group by name takes.""" + router = Router( + model_list=[ + _free_ollama_deployment("local-free"), + _paid_openai_deployment("gpt-priced", "gpt-4o"), + ], + model_group_alias={"local-free": "gpt-priced"}, + ) + + assert _served_models(router.get_model_list_of_routed_group("local-free")) == ["ollama/qwen3:0.6b"] + assert _served_models(router.get_model_list(model_name="local-free")) == ["ollama/qwen3:0.6b", "gpt-4o"] + + +def test_get_model_list_of_routed_group_falls_back_to_the_wildcard_route_serving_it(): + router = Router( + model_list=[ + _paid_openai_deployment("openai/*", "openai/*"), + _free_ollama_deployment("local-free"), + ], + model_group_alias={"openai/gpt-4o": "local-free"}, + ) + + routed = router.get_model_list_of_routed_group("openai/gpt-4o") + + assert [deployment["model_name"] for deployment in routed] == ["openai/gpt-4o"] + assert _served_models(routed) == ["openai/gpt-4o"] + assert _served_models(router.get_model_list(model_name="openai/gpt-4o")) == ["ollama/qwen3:0.6b"] + + def test_switch_routing_strategy_installs_lar1_then_restores_the_default_selector(): router = _alias_cost_router()