mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(auth): resolve hidden model_group_alias entries in the zero-cost budget check (#43741)
* fix(auth): resolve hidden model_group_alias entries in the zero-cost budget check Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(auth): key the zero-cost cache by the resolved model group Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(auth): keep the zero-cost verdict per requested name and include hidden aliases * test(integration): audit the zero-cost bypass through hidden model_group_alias names * test(router): cover the extracted routing strategy switch * test(integration): record a pre-flip burst before the alias flip --------- Co-authored-by: jesus <jesus@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
cb17588276
commit
1ef0fe9790
7 changed files with 1430 additions and 32 deletions
|
|
@ -484,7 +484,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None
|
|||
continue
|
||||
try:
|
||||
# Use router's get_model_group_info method directly for better reliability
|
||||
model_group_info = llm_router.get_model_group_info(model_group=model_name)
|
||||
model_group_info = llm_router.get_model_group_info(model_group=model_name, include_hidden=True)
|
||||
|
||||
if model_group_info is None:
|
||||
# Model not found or no pricing info available
|
||||
|
|
|
|||
|
|
@ -11246,13 +11246,13 @@ class Router:
|
|||
|
||||
return model_group_info
|
||||
|
||||
def get_model_group_info(self, model_group: str) -> ModelGroupInfo | None:
|
||||
def get_model_group_info(self, model_group: str, *, include_hidden: bool = False) -> ModelGroupInfo | None:
|
||||
"""
|
||||
For a given model group name, return the combined model info
|
||||
|
||||
Returns:
|
||||
- ModelGroupInfo if able to construct a model group
|
||||
- None if error constructing model group info or hidden model group
|
||||
- None if error constructing model group info or hidden model group (unless include_hidden)
|
||||
"""
|
||||
## Check if model group alias
|
||||
if model_group in self.model_group_alias:
|
||||
|
|
@ -11260,7 +11260,7 @@ class Router:
|
|||
if isinstance(item, str):
|
||||
_router_model_group = item
|
||||
elif isinstance(item, dict):
|
||||
if item["hidden"] is True:
|
||||
if item["hidden"] is True and not include_hidden:
|
||||
return None
|
||||
else:
|
||||
_router_model_group = item["model"]
|
||||
|
|
@ -12302,6 +12302,16 @@ class Router:
|
|||
]
|
||||
return _settings_to_return
|
||||
|
||||
def _switch_routing_strategy(self, routing_strategy: str | None, kwargs: Mapping[str, object]) -> None:
|
||||
if routing_strategy == "lar1":
|
||||
from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy
|
||||
|
||||
apply_lar1_routing_strategy(self, kwargs.get("routing_strategy_args"))
|
||||
return
|
||||
self.routing_strategy_init(
|
||||
routing_strategy=routing_strategy, routing_strategy_args=kwargs.get("routing_strategy_args", {})
|
||||
)
|
||||
|
||||
def update_settings(self, **kwargs):
|
||||
"""
|
||||
Update the router settings.
|
||||
|
|
@ -12315,6 +12325,7 @@ class Router:
|
|||
]
|
||||
|
||||
_existing_router_settings: Final = self.get_settings()
|
||||
model_group_alias_before: Final = self.model_group_alias
|
||||
rebuild_routing_groups = False
|
||||
routing_args_updated = False
|
||||
for var in kwargs:
|
||||
|
|
@ -12338,20 +12349,7 @@ class Router:
|
|||
if var == "routing_strategy":
|
||||
value = self._normalize_strategy(value)
|
||||
if _existing_router_settings["routing_strategy"] != value:
|
||||
if value == "lar1":
|
||||
from litellm.router_strategy.lar1_routing import (
|
||||
apply_lar1_routing_strategy,
|
||||
)
|
||||
|
||||
apply_lar1_routing_strategy(
|
||||
self,
|
||||
kwargs.get("routing_strategy_args"),
|
||||
)
|
||||
else:
|
||||
self.routing_strategy_init(
|
||||
routing_strategy=value,
|
||||
routing_strategy_args=kwargs.get("routing_strategy_args", {}),
|
||||
)
|
||||
self._switch_routing_strategy(value, kwargs)
|
||||
rebuild_routing_groups = True
|
||||
elif var == "routing_strategy_args":
|
||||
routing_args_updated = value != self.routing_strategy_args
|
||||
|
|
@ -12362,6 +12360,9 @@ class Router:
|
|||
if routing_args_updated:
|
||||
self._apply_updated_routing_strategy_args()
|
||||
|
||||
if self.model_group_alias != model_group_alias_before:
|
||||
self._invalidate_model_group_info_cache()
|
||||
|
||||
if rebuild_routing_groups:
|
||||
routing_groups_input: Final = kwargs.get("routing_groups", self._routing_groups_input)
|
||||
self._init_routing_groups(routing_groups_input)
|
||||
|
|
|
|||
423
tests/integration/authorization/_hidden_alias_budget.py
Normal file
423
tests/integration/authorization/_hidden_alias_budget.py
Normal file
|
|
@ -0,0 +1,423 @@
|
|||
import json
|
||||
import os
|
||||
import uuid
|
||||
from collections import Counter
|
||||
from collections.abc import Callable, Generator, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from hashlib import sha256
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from anthropic import Anthropic, AsyncAnthropic
|
||||
from integration._support.anthropic_thinking import JSON_LIST, 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.upstream import delete_scenario, register_scenario
|
||||
from integration.cost_calculation.cost_tracking_case import JsonResponse, SseResponse
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
from pydantic import JsonValue
|
||||
|
||||
BUDGET: Final = 0.05
|
||||
CHAT_REPLY: Final = "Hello! This is a mock response from the fake OpenAI endpoint."
|
||||
RESPONSES_REPLY: Final = "free reply"
|
||||
BUDGET_EXCEEDED: Final = 422
|
||||
GATEWAY_BURST: Final = 24
|
||||
PEER_BURST: Final = 4
|
||||
SPEND_MARKER_HEADER: Final = "x-litellm-spend-logs-metadata"
|
||||
PROXY_BUDGET_USER: Final = "litellm-proxy-budget"
|
||||
|
||||
_RESPONSE: Final[JsonValue] = {
|
||||
"id": "resp_$UNIQUE_ID",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": [
|
||||
{
|
||||
"id": "msg_$UNIQUE_ID",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": RESPONSES_REPLY, "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 20, "output_tokens": 20, "total_tokens": 40},
|
||||
}
|
||||
_RESPONSE_EVENTS: Final[tuple[Mapping[str, JsonValue], ...]] = (
|
||||
{
|
||||
"type": "response.created",
|
||||
"sequence_number": 0,
|
||||
"response": {**_RESPONSE, "status": "in_progress", "output": []},
|
||||
},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"sequence_number": 1,
|
||||
"item_id": "msg_$UNIQUE_ID",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": RESPONSES_REPLY,
|
||||
},
|
||||
{"type": "response.completed", "sequence_number": 2, "response": _RESPONSE},
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AliasRig:
|
||||
gateway: Gateway
|
||||
peer: Gateway
|
||||
free: str
|
||||
paid: str
|
||||
failing_free: str
|
||||
failing_provider_model: str
|
||||
hidden_free: str
|
||||
visible_free: str
|
||||
hidden_paid: str
|
||||
hidden_responses: str
|
||||
hidden_responses_stream: str
|
||||
hidden_failing: str
|
||||
shown_free: str
|
||||
null_hidden_free: str
|
||||
hidden_unpriced: str
|
||||
hidden_missing: str
|
||||
|
||||
|
||||
def base_url(candidate: Gateway) -> str:
|
||||
return str(candidate.client.base_url).rstrip("/")
|
||||
|
||||
|
||||
def spend_marker(marker: str) -> Mapping[str, str]:
|
||||
return {SPEND_MARKER_HEADER: json.dumps({"marker": marker})}
|
||||
|
||||
|
||||
def fresh_post(candidate: Gateway, path: str, body: Mapping[str, JsonValue], key: str, marker: str) -> httpx.Response:
|
||||
return httpx.post(
|
||||
f"{base_url(candidate)}{path}",
|
||||
json=dict(body),
|
||||
headers={"Authorization": f"Bearer {key}", **spend_marker(marker)},
|
||||
timeout=60,
|
||||
trust_env=False,
|
||||
)
|
||||
|
||||
|
||||
NO_EXTRA: Final[Mapping[str, JsonValue]] = MappingProxyType({})
|
||||
|
||||
|
||||
def chat_body(model: JsonValue, marker: str, extra: Mapping[str, JsonValue] = NO_EXTRA) -> Mapping[str, JsonValue]:
|
||||
return {"model": model, "messages": [{"role": "user", "content": marker}], **extra}
|
||||
|
||||
|
||||
def fresh_chat(
|
||||
candidate: Gateway, model: str, key: str, marker: str, extra: Mapping[str, JsonValue] = NO_EXTRA
|
||||
) -> httpx.Response:
|
||||
return fresh_post(candidate, "/v1/chat/completions", chat_body(model, marker, extra), key, marker)
|
||||
|
||||
|
||||
def fresh_response(
|
||||
candidate: Gateway, model: str, key: str, marker: str, extra: Mapping[str, JsonValue] = NO_EXTRA
|
||||
) -> httpx.Response:
|
||||
return fresh_post(candidate, "/v1/responses", {"model": model, "input": marker, **extra}, key, marker)
|
||||
|
||||
|
||||
def fresh_message(
|
||||
candidate: Gateway, model: str, key: str, marker: str, extra: Mapping[str, JsonValue] = NO_EXTRA
|
||||
) -> httpx.Response:
|
||||
return fresh_post(
|
||||
candidate,
|
||||
"/v1/messages",
|
||||
{"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": marker}], **extra},
|
||||
key,
|
||||
marker,
|
||||
)
|
||||
|
||||
|
||||
def statuses(send: Callable[[str], httpx.Response], count: int) -> frozenset[int]:
|
||||
markers: Final = tuple(uuid.uuid4().hex for _ in range(count))
|
||||
with ThreadPoolExecutor(max_workers=count) as pool:
|
||||
return frozenset(response.status_code for response in pool.map(send, markers))
|
||||
|
||||
|
||||
def error_type(response: httpx.Response) -> str:
|
||||
return str(object_value(JSON_OBJECT.validate_json(response.content)["error"])["type"])
|
||||
|
||||
|
||||
def settle(
|
||||
send: Callable[[str], httpx.Response], status: int, *, burst: int = GATEWAY_BURST, seconds: float = 60
|
||||
) -> None:
|
||||
eventually(
|
||||
lambda: statuses(send, burst) | statuses(send, burst),
|
||||
lambda seen: seen == frozenset({status}),
|
||||
seconds=seconds,
|
||||
)
|
||||
|
||||
|
||||
def chat_statuses(
|
||||
candidate: Gateway, model: str, key: str, count: int, extra: Mapping[str, JsonValue] = NO_EXTRA
|
||||
) -> frozenset[int]:
|
||||
return statuses(lambda marker: fresh_chat(candidate, model, key, marker, extra), count)
|
||||
|
||||
|
||||
def settle_candidate(
|
||||
candidate: Gateway,
|
||||
model: str,
|
||||
key: str,
|
||||
status: int,
|
||||
*,
|
||||
burst: int = GATEWAY_BURST,
|
||||
seconds: float = 60,
|
||||
extra: Mapping[str, JsonValue] = NO_EXTRA,
|
||||
) -> None:
|
||||
settle(lambda marker: fresh_chat(candidate, model, key, marker, extra), status, burst=burst, seconds=seconds)
|
||||
|
||||
|
||||
def settle_chat(
|
||||
rig: AliasRig,
|
||||
model: str,
|
||||
key: str,
|
||||
status: int,
|
||||
*,
|
||||
seconds: float = 60,
|
||||
extra: Mapping[str, JsonValue] = NO_EXTRA,
|
||||
) -> None:
|
||||
settle_candidate(rig.gateway, model, key, status, seconds=seconds, extra=extra)
|
||||
settle_candidate(rig.peer, model, key, status, burst=PEER_BURST, seconds=seconds, extra=extra)
|
||||
|
||||
|
||||
def exhausted_key(rig: AliasRig, 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_chat(rig, rig.paid, key, BUDGET_EXCEEDED)
|
||||
return key
|
||||
|
||||
|
||||
def upstream_requests(upstream_url: str) -> tuple[str, ...]:
|
||||
drained: Final = JSON_OBJECT.validate_json(
|
||||
httpx.get(f"{upstream_url}/__observations", timeout=15, trust_env=False).content
|
||||
)
|
||||
return tuple(json.dumps(entry) for entry in JSON_LIST.validate_python(drained["requests"]))
|
||||
|
||||
|
||||
def upstream_hits(observed: tuple[str, ...], marker: str) -> int:
|
||||
return sum(1 for entry in observed if marker in entry)
|
||||
|
||||
|
||||
def script_provider(rig: AliasRig, failures: int) -> None:
|
||||
scripted: Final = httpx.post(
|
||||
f"{rig.gateway.upstream_url}/__scripts/{rig.failing_provider_model}",
|
||||
json={"statuses": [500] * failures},
|
||||
timeout=15,
|
||||
trust_env=False,
|
||||
)
|
||||
assert scripted.status_code == 200, scripted.text
|
||||
|
||||
|
||||
def clear_provider_script(rig: AliasRig) -> 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 spend_rows(key: str) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
return tuple(
|
||||
read_rows(
|
||||
"SELECT request_id, spend, model_group, status, call_type, "
|
||||
"metadata->'spend_logs_metadata'->>'marker' AS marker "
|
||||
'FROM "LiteLLM_SpendLogs" WHERE api_key = %s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def landed(key: str, marker: str) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
return tuple(row for row in spend_rows(key) if row["marker"] == marker)
|
||||
|
||||
|
||||
def landed_once(key: str, marker: str) -> Mapping[str, JsonValue]:
|
||||
rows: Final = eventually(lambda: landed(key, marker), lambda found: len(found) >= 1, seconds=70)
|
||||
assert len(rows) == 1, rows
|
||||
return rows[0]
|
||||
|
||||
|
||||
def marker_counts(key: str, markers: frozenset[str]) -> Mapping[str, int]:
|
||||
return dict(Counter(str(row["marker"]) for row in spend_rows(key) if row["marker"] in markers))
|
||||
|
||||
|
||||
def landed_all_once(key: str, markers: frozenset[str]) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
counts: Final = eventually(
|
||||
lambda: marker_counts(key, markers), lambda found: frozenset(found) == markers, seconds=90
|
||||
)
|
||||
assert counts == dict.fromkeys(markers, 1), counts
|
||||
return tuple(row for row in spend_rows(key) if row["marker"] in markers)
|
||||
|
||||
|
||||
def assert_free_row(row: Mapping[str, JsonValue], model_group: str) -> None:
|
||||
assert float(str(row["spend"])) == 0.0, row
|
||||
assert row["model_group"] == model_group, row
|
||||
assert row["status"] == "success", row
|
||||
|
||||
|
||||
def alias_map(gateway: Gateway) -> Mapping[str, JsonValue]:
|
||||
current: Final = object_value(gateway.get("/router/settings")["current_values"]).get("model_group_alias")
|
||||
return object_value(current) if current is not None else {}
|
||||
|
||||
|
||||
def write_aliases(gateway: Gateway, aliases: Mapping[str, JsonValue]) -> None:
|
||||
gateway.post("/config/update", {"router_settings": {"model_group_alias": dict(aliases)}})
|
||||
|
||||
|
||||
def install_aliases(gateway: Gateway, aliases: Mapping[str, JsonValue]) -> None:
|
||||
write_aliases(gateway, {**alias_map(gateway), **aliases})
|
||||
|
||||
|
||||
def remove_aliases(gateway: Gateway, names: frozenset[str]) -> None:
|
||||
write_aliases(gateway, {name: target for name, target in alias_map(gateway).items() if name not in names})
|
||||
|
||||
|
||||
def hidden(group: str) -> JsonValue:
|
||||
return {"model": group, "hidden": True}
|
||||
|
||||
|
||||
def openai_client(candidate: Gateway, key: str) -> OpenAI:
|
||||
return OpenAI(
|
||||
api_key=key,
|
||||
base_url=base_url(candidate) + "/v1",
|
||||
max_retries=0,
|
||||
http_client=httpx.Client(timeout=60, trust_env=False),
|
||||
)
|
||||
|
||||
|
||||
def async_openai_client(candidate: Gateway, key: str) -> AsyncOpenAI:
|
||||
return AsyncOpenAI(
|
||||
api_key=key,
|
||||
base_url=base_url(candidate) + "/v1",
|
||||
max_retries=0,
|
||||
http_client=httpx.AsyncClient(timeout=60, trust_env=False),
|
||||
)
|
||||
|
||||
|
||||
def anthropic_client(candidate: Gateway, key: str) -> Anthropic:
|
||||
return Anthropic(
|
||||
api_key=key,
|
||||
base_url=base_url(candidate),
|
||||
max_retries=0,
|
||||
http_client=httpx.Client(timeout=60, trust_env=False),
|
||||
)
|
||||
|
||||
|
||||
def async_anthropic_client(candidate: Gateway, key: str) -> AsyncAnthropic:
|
||||
return AsyncAnthropic(
|
||||
api_key=key,
|
||||
base_url=base_url(candidate),
|
||||
max_retries=0,
|
||||
http_client=httpx.AsyncClient(timeout=60, trust_env=False),
|
||||
)
|
||||
|
||||
|
||||
def _zero_cost_responses_group(
|
||||
scenario: Scenario, gateway: Gateway, name: str, response: JsonResponse | SseResponse
|
||||
) -> str:
|
||||
handle: Final = register_scenario(name, response, control_url=gateway.upstream_url)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
return scenario.model(api_base=f"{handle.api_base()}/v1", input_cost_per_token=0, output_cost_per_token=0)
|
||||
|
||||
|
||||
def settle_responses(
|
||||
rig: AliasRig, model: str, key: str, status: int, extra: Mapping[str, JsonValue] = NO_EXTRA
|
||||
) -> None:
|
||||
settle(lambda marker: fresh_response(rig.gateway, model, key, marker, extra), status)
|
||||
settle(lambda marker: fresh_response(rig.peer, model, key, marker, extra), status, burst=PEER_BURST)
|
||||
|
||||
|
||||
def _await_rig(rig: AliasRig) -> None:
|
||||
admin: Final = rig.gateway.key
|
||||
for model in (
|
||||
rig.hidden_free,
|
||||
rig.visible_free,
|
||||
rig.hidden_paid,
|
||||
rig.hidden_failing,
|
||||
rig.shown_free,
|
||||
rig.null_hidden_free,
|
||||
rig.hidden_unpriced,
|
||||
):
|
||||
settle_chat(rig, model, admin, 200)
|
||||
settle_responses(rig, rig.hidden_responses, admin, 200)
|
||||
settle_responses(rig, rig.hidden_responses_stream, admin, 200, extra={"stream": True})
|
||||
|
||||
|
||||
@contextmanager
|
||||
def alias_rig() -> Generator[AliasRig]:
|
||||
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,
|
||||
):
|
||||
free: Final = scenario.model(input_cost_per_token=0, output_cost_per_token=0)
|
||||
paid: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
unpriced: Final = scenario.model()
|
||||
failing_provider_model: Final = f"hidden-alias-failing-{suffix}"
|
||||
failing_free: Final = scenario.model(
|
||||
model=f"openai/{failing_provider_model}", input_cost_per_token=0, output_cost_per_token=0
|
||||
)
|
||||
responses_free: Final = _zero_cost_responses_group(
|
||||
scenario,
|
||||
gateway,
|
||||
f"hidden-alias-json-{suffix}",
|
||||
JsonResponse(content_type="application/json", body=_RESPONSE),
|
||||
)
|
||||
responses_stream_free: Final = _zero_cost_responses_group(
|
||||
scenario,
|
||||
gateway,
|
||||
f"hidden-alias-sse-{suffix}",
|
||||
SseResponse(
|
||||
content_type="text/event-stream",
|
||||
frames=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}" for event in _RESPONSE_EVENTS),
|
||||
),
|
||||
)
|
||||
rig: Final = AliasRig(
|
||||
gateway=gateway,
|
||||
peer=Gateway(peer_client, gateway.key, gateway.upstream_url),
|
||||
free=free,
|
||||
paid=paid,
|
||||
failing_free=failing_free,
|
||||
failing_provider_model=failing_provider_model,
|
||||
hidden_free=f"hidden-free-{suffix}",
|
||||
visible_free=f"visible-free-{suffix}",
|
||||
hidden_paid=f"hidden-paid-{suffix}",
|
||||
hidden_responses=f"hidden-responses-{suffix}",
|
||||
hidden_responses_stream=f"hidden-responses-stream-{suffix}",
|
||||
hidden_failing=f"hidden-failing-{suffix}",
|
||||
shown_free=f"shown-free-{suffix}",
|
||||
null_hidden_free=f"null-hidden-free-{suffix}",
|
||||
hidden_unpriced=f"hidden-unpriced-{suffix}",
|
||||
hidden_missing=f"hidden-missing-{suffix}",
|
||||
)
|
||||
aliases: Final[Mapping[str, JsonValue]] = {
|
||||
rig.hidden_free: hidden(free),
|
||||
rig.visible_free: free,
|
||||
rig.hidden_paid: hidden(paid),
|
||||
rig.hidden_responses: hidden(responses_free),
|
||||
rig.hidden_responses_stream: hidden(responses_stream_free),
|
||||
rig.hidden_failing: hidden(failing_free),
|
||||
rig.shown_free: {"model": free, "hidden": False},
|
||||
rig.null_hidden_free: {"model": free, "hidden": None},
|
||||
rig.hidden_unpriced: hidden(unpriced),
|
||||
rig.hidden_missing: hidden(f"missing-group-{suffix}"),
|
||||
}
|
||||
install_aliases(gateway, aliases)
|
||||
scenario.cleanups.callback(remove_aliases, gateway, frozenset(aliases))
|
||||
_await_rig(rig)
|
||||
yield rig
|
||||
|
|
@ -0,0 +1,577 @@
|
|||
import json
|
||||
import math
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.anthropic_thinking import JSON_OBJECT
|
||||
from integration._support.client import Gateway, object_value
|
||||
from integration._support.process import owned_proxy
|
||||
from integration.authorization._hidden_alias_budget import (
|
||||
BUDGET,
|
||||
BUDGET_EXCEEDED,
|
||||
CHAT_REPLY,
|
||||
PROXY_BUDGET_USER,
|
||||
RESPONSES_REPLY,
|
||||
AliasRig,
|
||||
alias_rig,
|
||||
anthropic_client,
|
||||
assert_free_row,
|
||||
async_anthropic_client,
|
||||
async_openai_client,
|
||||
base_url,
|
||||
chat_statuses,
|
||||
clear_provider_script,
|
||||
error_type,
|
||||
exhausted_key,
|
||||
fresh_chat,
|
||||
hidden,
|
||||
install_aliases,
|
||||
landed,
|
||||
landed_once,
|
||||
openai_client,
|
||||
remove_aliases,
|
||||
script_provider,
|
||||
settle_candidate,
|
||||
settle_chat,
|
||||
spend_marker,
|
||||
upstream_hits,
|
||||
upstream_requests,
|
||||
)
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(240)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig() -> Iterator[AliasRig]:
|
||||
with alias_rig() as built:
|
||||
yield built
|
||||
|
||||
|
||||
def test_exhausted_key_reaches_hidden_free_alias_through_openai_chat(rig: AliasRig) -> 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.hidden_free,
|
||||
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.hidden_free)
|
||||
|
||||
|
||||
async def test_exhausted_key_reaches_hidden_free_alias_through_streamed_openai_chat(rig: AliasRig) -> 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.hidden_free,
|
||||
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.hidden_free)
|
||||
|
||||
|
||||
def test_exhausted_key_reaches_hidden_free_alias_through_anthropic_messages(rig: AliasRig) -> 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.hidden_responses,
|
||||
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"] == [RESPONSES_REPLY], message
|
||||
assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1
|
||||
row: Final = landed_once(key, marker)
|
||||
assert row["request_id"] == message.id, row
|
||||
assert_free_row(row, rig.hidden_responses)
|
||||
|
||||
|
||||
async def test_exhausted_key_reaches_hidden_free_alias_through_streamed_anthropic_messages(rig: AliasRig) -> 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.hidden_responses_stream,
|
||||
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 == RESPONSES_REPLY, final
|
||||
assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1
|
||||
row: Final = landed_once(key, marker)
|
||||
assert row["request_id"] == final.id, row
|
||||
assert_free_row(row, rig.hidden_responses_stream)
|
||||
|
||||
|
||||
def test_exhausted_key_reaches_hidden_free_alias_through_openai_responses(rig: AliasRig) -> 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.hidden_responses, input=marker, extra_headers=spend_marker(marker)
|
||||
)
|
||||
assert response.output_text == RESPONSES_REPLY, response
|
||||
assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1
|
||||
row: Final = landed_once(key, marker)
|
||||
assert row["request_id"] == response.id, row
|
||||
assert_free_row(row, rig.hidden_responses)
|
||||
|
||||
|
||||
async def test_exhausted_key_reaches_hidden_free_alias_through_streamed_openai_responses(rig: AliasRig) -> 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.hidden_responses_stream, 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") == RESPONSES_REPLY
|
||||
assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1
|
||||
assert_free_row(landed_once(key, marker), rig.hidden_responses_stream)
|
||||
|
||||
|
||||
def _assert_raw_chat_served(rig: AliasRig, candidate: Gateway, key: str) -> None:
|
||||
marker: Final = "raw-" + uuid.uuid4().hex
|
||||
response: Final = fresh_chat(candidate, rig.hidden_free, 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, rig.hidden_free)
|
||||
|
||||
|
||||
def test_exhausted_key_reaches_hidden_free_alias_over_raw_http_on_both_replicas(rig: AliasRig) -> None:
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
_assert_raw_chat_served(rig, rig.gateway, key)
|
||||
_assert_raw_chat_served(rig, rig.peer, key)
|
||||
|
||||
|
||||
def _duplicate_model_post(rig: AliasRig, 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: AliasRig) -> 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.hidden_paid, rig.hidden_free, free_marker)
|
||||
assert served.status_code == 200, served.text
|
||||
refused: Final = _duplicate_model_post(rig, key, rig.hidden_free, rig.hidden_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.hidden_free)
|
||||
|
||||
|
||||
def test_provider_failure_behind_hidden_free_alias_reaches_the_caller(rig: AliasRig) -> 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_provider_script, rig)
|
||||
script_provider(rig, 1)
|
||||
failed: Final = fresh_chat(rig.gateway, rig.hidden_failing, key, failed_marker)
|
||||
assert failed.status_code == 500, failed.text
|
||||
assert "Controlled provider failure" in failed.text, failed.text
|
||||
assert "budget" not in failed.text.lower(), failed.text
|
||||
clear_provider_script(rig)
|
||||
recovered: Final = fresh_chat(rig.gateway, rig.hidden_failing, 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.hidden_failing)
|
||||
|
||||
|
||||
def test_exhausted_user_budget_still_reaches_hidden_free_alias(rig: AliasRig) -> None:
|
||||
with rig.gateway.scenario() as scenario:
|
||||
user: Final = scenario.user(max_budget=BUDGET)
|
||||
key: Final = scenario.key(user_id=user, max_budget=5.0)
|
||||
first: Final = fresh_chat(rig.gateway, rig.paid, key, "user-exhaust-" + uuid.uuid4().hex)
|
||||
assert first.status_code == 200, first.text
|
||||
settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED, seconds=90)
|
||||
assert chat_statuses(rig.gateway, rig.hidden_free, key, 8) == {200}
|
||||
assert chat_statuses(rig.peer, rig.hidden_free, key, 8) == {200}
|
||||
settle_chat(rig, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=90)
|
||||
refused: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "user-refused-" + uuid.uuid4().hex)
|
||||
assert f"User={user}" in refused.text, refused.text
|
||||
|
||||
|
||||
def test_exhausted_team_budget_still_reaches_hidden_free_alias(rig: AliasRig) -> None:
|
||||
with rig.gateway.scenario() as scenario:
|
||||
team: Final = scenario.team(max_budget=BUDGET)
|
||||
key: Final = scenario.key(team_id=team, max_budget=5.0)
|
||||
first: Final = fresh_chat(rig.gateway, rig.paid, key, "team-exhaust-" + uuid.uuid4().hex)
|
||||
assert first.status_code == 200, first.text
|
||||
settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED, seconds=90)
|
||||
assert chat_statuses(rig.gateway, rig.hidden_free, key, 8) == {200}
|
||||
assert chat_statuses(rig.peer, rig.hidden_free, key, 8) == {200}
|
||||
settle_chat(rig, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=90)
|
||||
refused: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "team-refused-" + uuid.uuid4().hex)
|
||||
assert f"Team={team}" in refused.text, refused.text
|
||||
|
||||
|
||||
def test_exhausted_tag_budget_still_reaches_hidden_free_alias(rig: AliasRig) -> None:
|
||||
tag: Final = "hidden-alias-tag-" + uuid.uuid4().hex
|
||||
tagged: Final[Mapping[str, JsonValue]] = {"metadata": {"tags": [tag]}}
|
||||
with rig.gateway.scenario() as scenario:
|
||||
rig.gateway.post("/tag/new", {"name": tag, "max_budget": BUDGET})
|
||||
scenario.cleanups.callback(rig.gateway.post, "/tag/delete", {"name": tag})
|
||||
key: Final = scenario.key(max_budget=5.0)
|
||||
first: Final = fresh_chat(rig.gateway, rig.paid, key, "tag-exhaust-" + uuid.uuid4().hex, tagged)
|
||||
assert first.status_code == 200, first.text
|
||||
settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED, seconds=90, extra=tagged)
|
||||
assert chat_statuses(rig.gateway, rig.hidden_free, key, 8, tagged) == {200}
|
||||
assert chat_statuses(rig.peer, rig.hidden_free, key, 8, tagged) == {200}
|
||||
settle_chat(rig, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=90, extra=tagged)
|
||||
refused: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "tag-refused-" + uuid.uuid4().hex, tagged)
|
||||
assert f"Tag={tag}" in refused.text, refused.text
|
||||
untagged: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "tag-untagged-" + uuid.uuid4().hex)
|
||||
assert untagged.status_code == 200, untagged.text
|
||||
|
||||
|
||||
def test_hidden_alias_repointed_between_paid_and_free_groups_follows_the_target(rig: AliasRig) -> None:
|
||||
alias: Final = "hidden-repoint-" + uuid.uuid4().hex
|
||||
with rig.gateway.scenario() as scenario:
|
||||
install_aliases(rig.gateway, {alias: hidden(rig.paid)})
|
||||
scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({alias}))
|
||||
settle_chat(rig, alias, rig.gateway.key, 200)
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
settle_chat(rig, alias, key, BUDGET_EXCEEDED)
|
||||
install_aliases(rig.gateway, {alias: hidden(rig.free)})
|
||||
settle_chat(rig, alias, key, 200)
|
||||
install_aliases(rig.gateway, {alias: hidden(rig.paid)})
|
||||
settle_chat(rig, alias, key, BUDGET_EXCEEDED)
|
||||
|
||||
|
||||
def test_visible_alias_repointed_to_a_paid_group_loses_the_bypass(rig: AliasRig) -> None:
|
||||
alias: Final = "visible-repoint-" + uuid.uuid4().hex
|
||||
with rig.gateway.scenario() as scenario:
|
||||
install_aliases(rig.gateway, {alias: rig.free})
|
||||
scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({alias}))
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
settle_chat(rig, alias, key, 200)
|
||||
install_aliases(rig.gateway, {alias: rig.paid})
|
||||
settle_chat(rig, alias, key, BUDGET_EXCEEDED)
|
||||
install_aliases(rig.gateway, {alias: rig.free})
|
||||
settle_chat(rig, alias, key, 200)
|
||||
|
||||
|
||||
def test_failed_free_primary_falls_back_to_hidden_free_alias_for_exhausted_key(rig: AliasRig) -> None:
|
||||
free_marker: Final = "fallback-free-" + uuid.uuid4().hex
|
||||
paid_marker: Final = "fallback-paid-" + uuid.uuid4().hex
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
scenario.cleanups.callback(clear_provider_script, rig)
|
||||
script_provider(rig, 2)
|
||||
served: Final = fresh_chat(rig.gateway, rig.failing_free, key, free_marker, {"fallbacks": [rig.hidden_free]})
|
||||
assert served.status_code == 200, served.text
|
||||
assert served.headers["x-litellm-model-group"] == rig.hidden_free, dict(served.headers)
|
||||
refused: Final = fresh_chat(rig.gateway, rig.failing_free, key, paid_marker, {"fallbacks": [rig.hidden_paid]})
|
||||
assert refused.status_code == 500, refused.text
|
||||
assert "Controlled provider failure" in refused.text, refused.text
|
||||
observed: Final = upstream_requests(rig.gateway.upstream_url)
|
||||
assert upstream_hits(observed, free_marker) == 2
|
||||
assert upstream_hits(observed, paid_marker) == 1
|
||||
assert_free_row(landed_once(key, free_marker), rig.hidden_free)
|
||||
|
||||
|
||||
def test_exhausted_key_is_served_a_cached_reply_through_hidden_free_alias(rig: AliasRig) -> 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.hidden_free, key, marker)
|
||||
assert first.status_code == 200, first.text
|
||||
second: Final = fresh_chat(rig.gateway, rig.hidden_free, key, marker)
|
||||
assert second.status_code == 200, second.text
|
||||
assert second.json()["id"] == first.json()["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 = landed(key, marker)
|
||||
assert all(float(str(row["spend"])) == 0.0 for row in rows), rows
|
||||
|
||||
|
||||
def _base_config() -> Mapping[str, JsonValue]:
|
||||
return JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
|
||||
|
||||
|
||||
def _own_config(directory: Path, name: str, section: str, value: JsonValue) -> Path:
|
||||
path: Final = directory / name
|
||||
path.write_text(yaml.safe_dump({**_base_config(), section: value}))
|
||||
return path
|
||||
|
||||
|
||||
def _delete_proxy_budget_row(rig: AliasRig) -> None:
|
||||
deleted: Final = rig.gateway.request("POST", "/user/delete", {"user_ids": [PROXY_BUDGET_USER]})
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
|
||||
|
||||
@pytest.mark.timeout(480)
|
||||
def test_exhausted_proxy_budget_still_reaches_hidden_free_alias(rig: AliasRig, tmp_path: Path) -> None:
|
||||
settings: Final = object_value(_base_config()["litellm_settings"])
|
||||
config: Final = _own_config(
|
||||
tmp_path,
|
||||
"proxy-budget.yaml",
|
||||
"litellm_settings",
|
||||
{**settings, "max_budget": BUDGET, "budget_duration": "30d"},
|
||||
)
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = scenario.key()
|
||||
scenario.cleanups.callback(_delete_proxy_budget_row, rig)
|
||||
with owned_proxy(rig.gateway, tmp_path, {}, config=config, workers=2) as candidate:
|
||||
settle_candidate(candidate, rig.hidden_free, key, 200, seconds=120)
|
||||
first: Final = fresh_chat(candidate, rig.paid, key, "proxy-exhaust-" + uuid.uuid4().hex)
|
||||
assert first.status_code == 200, first.text
|
||||
settle_candidate(candidate, rig.paid, key, BUDGET_EXCEEDED, seconds=120)
|
||||
refused: Final = fresh_chat(candidate, rig.paid, key, "proxy-refused-" + uuid.uuid4().hex)
|
||||
assert error_type(refused) == "budget_exceeded", refused.text
|
||||
assert "Key=" not in refused.text, refused.text
|
||||
assert chat_statuses(candidate, rig.hidden_free, key, 8) == {200}
|
||||
settle_candidate(candidate, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=120)
|
||||
unbudgeted: Final = fresh_chat(rig.gateway, rig.paid, key, "proxy-unbudgeted-" + uuid.uuid4().hex)
|
||||
assert unbudgeted.status_code == 200, unbudgeted.text
|
||||
|
||||
|
||||
_TAG_ADDER: Final = """from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
||||
|
||||
class TagAdder(CustomGuardrail):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
metadata = data.setdefault("metadata", {})
|
||||
metadata["tags"] = [*(metadata.get("tags") or []), "__TAG__"]
|
||||
return data
|
||||
"""
|
||||
|
||||
|
||||
@pytest.mark.timeout(480)
|
||||
def test_guardrail_added_tag_over_budget_still_reaches_hidden_free_alias(rig: AliasRig, tmp_path: Path) -> None:
|
||||
tag: Final = "hidden-alias-guardrail-tag-" + uuid.uuid4().hex
|
||||
module: Final = "tag_adder_" + uuid.uuid4().hex
|
||||
served_marker: Final = "guardrail-served-" + uuid.uuid4().hex
|
||||
(tmp_path / f"{module}.py").write_text(_TAG_ADDER.replace("__TAG__", tag))
|
||||
config: Final = _own_config(
|
||||
tmp_path,
|
||||
"guardrail-tag.yaml",
|
||||
"guardrails",
|
||||
[
|
||||
{
|
||||
"guardrail_name": "tag-adder-" + uuid.uuid4().hex,
|
||||
"litellm_params": {"guardrail": f"{module}.TagAdder", "mode": "pre_call", "default_on": True},
|
||||
}
|
||||
],
|
||||
)
|
||||
tagged: Final[Mapping[str, JsonValue]] = {"metadata": {"tags": [tag]}}
|
||||
with rig.gateway.scenario() as scenario:
|
||||
rig.gateway.post("/tag/new", {"name": tag, "max_budget": BUDGET})
|
||||
scenario.cleanups.callback(rig.gateway.post, "/tag/delete", {"name": tag})
|
||||
key: Final = scenario.key(max_budget=5.0)
|
||||
first: Final = fresh_chat(rig.gateway, rig.paid, key, "guardrail-exhaust-" + uuid.uuid4().hex, tagged)
|
||||
assert first.status_code == 200, first.text
|
||||
settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED, seconds=90, extra=tagged)
|
||||
untagged: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "guardrail-untagged-" + uuid.uuid4().hex)
|
||||
assert untagged.status_code == 200, untagged.text
|
||||
with owned_proxy(rig.gateway, tmp_path, {}, config=config, workers=2) as candidate:
|
||||
settle_candidate(candidate, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=120)
|
||||
refused: Final = fresh_chat(candidate, rig.hidden_paid, key, "guardrail-refused-" + uuid.uuid4().hex)
|
||||
assert f"Tag={tag}" in refused.text, refused.text
|
||||
assert chat_statuses(candidate, rig.hidden_free, key, 8) == {200}
|
||||
served: Final = fresh_chat(candidate, rig.hidden_free, key, served_marker)
|
||||
assert served.status_code == 200, served.text
|
||||
row: Final = landed_once(key, served_marker)
|
||||
assert row["request_id"] == served.json()["id"], row
|
||||
assert_free_row(row, rig.hidden_free)
|
||||
|
||||
|
||||
def _assert_free_alias_served(rig: AliasRig, candidate: Gateway, alias: str, key: str, prefix: str) -> None:
|
||||
marker: Final = f"{prefix}-" + uuid.uuid4().hex
|
||||
response: Final = fresh_chat(candidate, alias, key, marker)
|
||||
assert response.status_code == 200, response.text
|
||||
assert_free_row(landed_once(key, marker), alias)
|
||||
|
||||
|
||||
def test_exhausted_key_reaches_visible_free_alias(rig: AliasRig) -> None:
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
_assert_free_alias_served(rig, rig.gateway, rig.visible_free, key, "visible")
|
||||
_assert_free_alias_served(rig, rig.peer, rig.visible_free, key, "visible")
|
||||
|
||||
|
||||
def _assert_hidden_paid_refused(rig: AliasRig, candidate: Gateway, key: str) -> None:
|
||||
marker: Final = "hidden-paid-" + uuid.uuid4().hex
|
||||
response: Final = fresh_chat(candidate, rig.hidden_paid, key, marker)
|
||||
assert response.status_code == BUDGET_EXCEEDED, response.text
|
||||
assert error_type(response) == "budget_exceeded", response.text
|
||||
assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0
|
||||
|
||||
|
||||
def test_exhausted_key_is_refused_on_hidden_paid_alias(rig: AliasRig) -> None:
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
_assert_hidden_paid_refused(rig, rig.gateway, key)
|
||||
_assert_hidden_paid_refused(rig, rig.peer, key)
|
||||
|
||||
|
||||
def test_exhausted_key_reaches_free_group_by_its_own_name(rig: AliasRig) -> None:
|
||||
marker: Final = "plain-free-" + uuid.uuid4().hex
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
response: Final = fresh_chat(rig.gateway, rig.free, key, marker)
|
||||
assert response.status_code == 200, response.text
|
||||
assert_free_row(landed_once(key, marker), rig.free)
|
||||
|
||||
|
||||
def test_key_with_headroom_is_billed_through_hidden_paid_alias(rig: AliasRig) -> None:
|
||||
paid_marker: Final = "headroom-paid-" + uuid.uuid4().hex
|
||||
free_marker: Final = "headroom-free-" + uuid.uuid4().hex
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = scenario.key(max_budget=5.0)
|
||||
paid: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, paid_marker)
|
||||
assert paid.status_code == 200, paid.text
|
||||
free: Final = fresh_chat(rig.gateway, rig.hidden_free, key, free_marker)
|
||||
assert free.status_code == 200, free.text
|
||||
billed: Final = landed_once(key, paid_marker)
|
||||
assert math.isclose(float(str(billed["spend"])), 20 * 0.001 + 20 * 0.002), billed
|
||||
assert billed["model_group"] == rig.hidden_paid, billed
|
||||
assert_free_row(landed_once(key, free_marker), rig.hidden_free)
|
||||
|
||||
|
||||
def test_key_restricted_to_the_free_group_reaches_its_hidden_alias(rig: AliasRig) -> None:
|
||||
marker: Final = "restricted-" + uuid.uuid4().hex
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = scenario.key(models=[rig.free], max_budget=5.0)
|
||||
response: Final = fresh_chat(rig.gateway, rig.hidden_free, key, marker)
|
||||
assert response.status_code == 200, response.text
|
||||
refused: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "restricted-paid-" + uuid.uuid4().hex)
|
||||
assert refused.status_code == 403, refused.text
|
||||
assert error_type(refused) == "key_model_access_denied", refused.text
|
||||
assert_free_row(landed_once(key, marker), rig.hidden_free)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("flag", ["false", "null"])
|
||||
def test_alias_with_a_non_hidden_flag_keeps_the_bypass(rig: AliasRig, flag: str) -> None:
|
||||
alias: Final = {"false": rig.shown_free, "null": rig.null_hidden_free}[flag]
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
_assert_free_alias_served(rig, rig.gateway, alias, key, f"flag-{flag}")
|
||||
_assert_free_alias_served(rig, rig.peer, alias, key, f"flag-{flag}")
|
||||
|
||||
|
||||
def test_hidden_alias_to_a_group_priced_by_the_cost_map_stays_budgeted(rig: AliasRig) -> None:
|
||||
marker: Final = "unpriced-" + uuid.uuid4().hex
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
response: Final = fresh_chat(rig.gateway, rig.hidden_unpriced, key, marker)
|
||||
assert response.status_code == BUDGET_EXCEEDED, response.text
|
||||
assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0
|
||||
|
||||
|
||||
def test_hidden_alias_to_a_missing_group_is_refused_and_the_proxy_stays_healthy(rig: AliasRig) -> None:
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
refused: Final = fresh_chat(rig.gateway, rig.hidden_missing, key, "missing-" + uuid.uuid4().hex)
|
||||
assert refused.status_code == BUDGET_EXCEEDED, refused.text
|
||||
unroutable: Final = fresh_chat(rig.gateway, rig.hidden_missing, rig.gateway.key, "missing-" + uuid.uuid4().hex)
|
||||
assert unroutable.status_code == 400, unroutable.text
|
||||
assert "no healthy deployments" in unroutable.text, unroutable.text
|
||||
for candidate in (rig.gateway, rig.peer):
|
||||
assert candidate.request("GET", "/health/liveliness").status_code == 200
|
||||
assert candidate.request("GET", "/model/info").status_code == 200
|
||||
assert candidate.request("GET", "/v1/models").status_code == 200
|
||||
served: Final = fresh_chat(rig.gateway, rig.hidden_free, key, "missing-control-" + uuid.uuid4().hex)
|
||||
assert served.status_code == 200, served.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("shape", "status"),
|
||||
[("int", BUDGET_EXCEEDED), ("list", 400), ("empty", BUDGET_EXCEEDED), ("oversized", BUDGET_EXCEEDED)],
|
||||
)
|
||||
def test_malformed_model_value_never_takes_the_bypass(rig: AliasRig, shape: str, status: int) -> None:
|
||||
marker: Final = f"malformed-{shape}-" + uuid.uuid4().hex
|
||||
models: Final[Mapping[str, JsonValue]] = {
|
||||
"int": 5,
|
||||
"list": [rig.hidden_free],
|
||||
"empty": "",
|
||||
"oversized": rig.hidden_free + "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
|
||||
served: Final = fresh_chat(rig.gateway, rig.hidden_free, key, "malformed-control-" + uuid.uuid4().hex)
|
||||
assert served.status_code == 200, served.text
|
||||
|
||||
|
||||
def test_unauthenticated_request_to_hidden_alias_is_rejected(rig: AliasRig) -> None:
|
||||
marker: Final = "unauthenticated-" + uuid.uuid4().hex
|
||||
response: Final = httpx.post(
|
||||
f"{base_url(rig.gateway)}/v1/chat/completions",
|
||||
json={"model": rig.hidden_free, "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 _assert_hidden_aliases_unlisted(rig: AliasRig, candidate: Gateway) -> None:
|
||||
models: Final = candidate.get("/v1/models")["data"]
|
||||
groups: Final = candidate.get("/model_group/info")["data"]
|
||||
assert isinstance(models, list) and isinstance(groups, list)
|
||||
listed: Final = frozenset(str(object_value(entry)["id"]) for entry in models)
|
||||
described: Final = frozenset(str(object_value(entry)["model_group"]) for entry in groups)
|
||||
assert rig.visible_free in listed and rig.visible_free in described
|
||||
assert rig.shown_free in listed and rig.shown_free in described
|
||||
for name in (rig.hidden_free, rig.hidden_paid, rig.hidden_responses, rig.hidden_missing):
|
||||
assert name not in listed and name not in described, name
|
||||
|
||||
|
||||
def test_hidden_alias_stays_out_of_model_listings(rig: AliasRig) -> None:
|
||||
_assert_hidden_aliases_unlisted(rig, rig.gateway)
|
||||
_assert_hidden_aliases_unlisted(rig, rig.peer)
|
||||
|
|
@ -0,0 +1,208 @@
|
|||
import os
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic import JsonValue
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.process import owned_proxy_process, owned_upstream
|
||||
from integration.authorization._hidden_alias_budget import (
|
||||
BUDGET_EXCEEDED,
|
||||
AliasRig,
|
||||
alias_rig,
|
||||
assert_free_row,
|
||||
chat_statuses,
|
||||
exhausted_key,
|
||||
fresh_chat,
|
||||
fresh_message,
|
||||
fresh_response,
|
||||
hidden,
|
||||
install_aliases,
|
||||
landed_all_once,
|
||||
remove_aliases,
|
||||
settle_candidate,
|
||||
settle_chat,
|
||||
upstream_hits,
|
||||
upstream_requests,
|
||||
)
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(240)
|
||||
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
_WAVE: Final = 8
|
||||
_STREAM: Final[Mapping[str, JsonValue]] = {"stream": True}
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig() -> Iterator[AliasRig]:
|
||||
with alias_rig() as built:
|
||||
yield built
|
||||
|
||||
|
||||
def _burst_markers(prefix: str, count: int) -> tuple[str, ...]:
|
||||
return tuple(f"{prefix}-{index}-{uuid.uuid4().hex}" for index in range(count))
|
||||
|
||||
|
||||
def _send_all(send: Callable[[str], httpx.Response], markers: tuple[str, ...]) -> tuple[httpx.Response, ...]:
|
||||
with ThreadPoolExecutor(max_workers=len(markers)) as pool:
|
||||
return tuple(pool.map(send, markers))
|
||||
|
||||
|
||||
def _mixed_call(rig: AliasRig, key: str, marker: str) -> httpx.Response:
|
||||
senders: Final[Mapping[str, Callable[[], httpx.Response]]] = {
|
||||
"chat": lambda: fresh_chat(rig.gateway, rig.hidden_free, key, marker),
|
||||
"chatstream": lambda: fresh_chat(rig.gateway, rig.hidden_free, key, marker, _STREAM),
|
||||
"messages": lambda: fresh_message(rig.gateway, rig.hidden_responses, key, marker),
|
||||
"messagesstream": lambda: fresh_message(rig.gateway, rig.hidden_responses_stream, key, marker, _STREAM),
|
||||
"responses": lambda: fresh_response(rig.gateway, rig.hidden_responses, key, marker),
|
||||
"responsesstream": lambda: fresh_response(rig.gateway, rig.hidden_responses_stream, key, marker, _STREAM),
|
||||
}
|
||||
return senders[marker.split("-")[1]]()
|
||||
|
||||
|
||||
_KINDS: Final = ("chat", "chatstream", "messages", "messagesstream", "responses", "responsesstream")
|
||||
|
||||
|
||||
def test_mixed_concurrent_burst_through_hidden_free_aliases_lands_each_call_once(rig: AliasRig) -> None:
|
||||
markers: Final = tuple(f"burst-{_KINDS[index % len(_KINDS)]}-{index}-{uuid.uuid4().hex}" for index in range(30))
|
||||
with rig.gateway.scenario() as scenario:
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
upstream_requests(rig.gateway.upstream_url)
|
||||
responses: Final = _send_all(lambda marker: _mixed_call(rig, key, marker), markers)
|
||||
assert [response.status_code for response in responses] == [200] * len(markers), [
|
||||
response.text for response in responses if response.status_code != 200
|
||||
]
|
||||
observed: Final = upstream_requests(rig.gateway.upstream_url)
|
||||
assert {marker: upstream_hits(observed, marker) for marker in markers} == dict.fromkeys(markers, 1)
|
||||
rows: Final = landed_all_once(key, frozenset(markers))
|
||||
assert len(rows) == len(markers), rows
|
||||
for row in rows:
|
||||
assert float(str(row["spend"])) == 0.0, row
|
||||
assert row["status"] == "success", row
|
||||
|
||||
|
||||
def test_alias_flipped_to_a_free_group_during_a_burst_only_ever_serves_or_refuses(rig: AliasRig) -> None:
|
||||
alias: Final = "flip-" + uuid.uuid4().hex
|
||||
seen: Final[SimpleQueue[tuple[str, int]]] = SimpleQueue()
|
||||
with rig.gateway.scenario() as scenario:
|
||||
install_aliases(rig.gateway, {alias: hidden(rig.paid)})
|
||||
scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({alias}))
|
||||
settle_chat(rig, alias, rig.gateway.key, 200)
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
settle_chat(rig, alias, key, BUDGET_EXCEEDED)
|
||||
upstream_requests(rig.gateway.upstream_url)
|
||||
|
||||
def wave(candidate: Gateway) -> frozenset[int]:
|
||||
markers: Final = _burst_markers("flip", _WAVE)
|
||||
responses: Final = _send_all(lambda marker: fresh_chat(candidate, alias, key, marker), markers)
|
||||
for marker, response in zip(markers, responses, strict=True):
|
||||
seen.put((marker, response.status_code))
|
||||
return frozenset(response.status_code for response in responses)
|
||||
|
||||
assert wave(rig.gateway) == frozenset({BUDGET_EXCEEDED})
|
||||
flip: Final = threading.Thread(target=install_aliases, args=(rig.gateway, {alias: hidden(rig.free)}))
|
||||
flip.start()
|
||||
eventually(lambda: wave(rig.gateway) | wave(rig.gateway), lambda found: found == frozenset({200}), seconds=60)
|
||||
flip.join(timeout=30)
|
||||
assert not flip.is_alive()
|
||||
settle_chat(rig, alias, key, 200)
|
||||
collected: Final = tuple(seen.get() for _ in range(seen.qsize()))
|
||||
assert {status for _, status in collected} <= {200, BUDGET_EXCEEDED}, collected
|
||||
served: Final = frozenset(marker for marker, status in collected if status == 200)
|
||||
refused: Final = frozenset(marker for marker, status in collected if status == BUDGET_EXCEEDED)
|
||||
assert served and refused, collected
|
||||
observed: Final = upstream_requests(rig.gateway.upstream_url)
|
||||
assert all(upstream_hits(observed, marker) == 1 for marker in served), collected
|
||||
assert all(upstream_hits(observed, marker) == 0 for marker in refused), collected
|
||||
for row in landed_all_once(key, served):
|
||||
assert_free_row(row, alias)
|
||||
|
||||
|
||||
def _tolerant_status(candidate: Gateway, model: str, key: str, marker: str) -> int | None:
|
||||
try:
|
||||
return fresh_chat(candidate, model, key, marker).status_code
|
||||
except httpx.TransportError:
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.timeout(480)
|
||||
def test_upstream_outage_behind_hidden_free_alias_is_a_provider_error_and_recovers(
|
||||
rig: AliasRig, tmp_path: Path
|
||||
) -> None:
|
||||
alias: Final = "hidden-outage-" + uuid.uuid4().hex
|
||||
with owned_upstream(tmp_path) as slot, rig.gateway.scenario() as scenario:
|
||||
group: Final = scenario.model(api_base=f"{slot.url}/v1", input_cost_per_token=0, output_cost_per_token=0)
|
||||
install_aliases(rig.gateway, {alias: hidden(group)})
|
||||
scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({alias}))
|
||||
settle_chat(rig, alias, rig.gateway.key, 200)
|
||||
key: Final = exhausted_key(rig, scenario)
|
||||
before: Final = _burst_markers("outage-before", 10)
|
||||
served_before: Final = _send_all(lambda marker: fresh_chat(rig.gateway, alias, key, marker), before)
|
||||
assert [response.status_code for response in served_before] == [200] * 10
|
||||
slot.stop()
|
||||
during: Final = _burst_markers("outage-during", 10)
|
||||
failed: Final = _send_all(lambda marker: fresh_chat(rig.gateway, alias, key, marker), during)
|
||||
for response in failed:
|
||||
assert response.status_code >= 500, response.text
|
||||
assert "budget" not in response.text.lower(), response.text
|
||||
assert rig.gateway.request("GET", "/health/liveliness").status_code == 200
|
||||
unrelated: Final = fresh_chat(rig.gateway, rig.hidden_free, key, "outage-unrelated-" + uuid.uuid4().hex)
|
||||
assert unrelated.status_code == 200, unrelated.text
|
||||
slot.start()
|
||||
settle_candidate(rig.gateway, alias, key, 200, seconds=90)
|
||||
after: Final = _burst_markers("outage-after", 10)
|
||||
served_after: Final = _send_all(lambda marker: fresh_chat(rig.gateway, alias, key, marker), after)
|
||||
assert [response.status_code for response in served_after] == [200] * 10
|
||||
observed: Final = upstream_requests(slot.url)
|
||||
assert {marker: upstream_hits(observed, marker) for marker in after} == dict.fromkeys(after, 1)
|
||||
for row in landed_all_once(key, frozenset(before + after)):
|
||||
assert_free_row(row, alias)
|
||||
|
||||
|
||||
def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]:
|
||||
text: Final = log.read_text()
|
||||
started: Final = tuple(int(found[1]) for found in _STARTED_WORKER.finditer(text))
|
||||
return started, text.count("Application startup complete.")
|
||||
|
||||
|
||||
@pytest.mark.timeout(480)
|
||||
def test_killed_worker_leaves_the_sibling_serving_hidden_free_aliases(rig: AliasRig, tmp_path: Path) -> None:
|
||||
with rig.gateway.scenario() as scenario:
|
||||
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, rig.hidden_free, key, 200, seconds=120)
|
||||
settle_candidate(candidate, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=120)
|
||||
os.kill(workers[0], signal.SIGKILL)
|
||||
eventually(
|
||||
lambda: _tolerant_status(candidate, rig.hidden_free, key, "kill-probe-" + uuid.uuid4().hex),
|
||||
lambda found: found == 200,
|
||||
seconds=60,
|
||||
)
|
||||
assert chat_statuses(candidate, rig.hidden_free, key, 8) == {200}
|
||||
assert chat_statuses(candidate, rig.hidden_paid, key, 8) == {BUDGET_EXCEEDED}
|
||||
eventually(
|
||||
lambda: _worker_startups(owned.log),
|
||||
lambda found: len(found[0]) == 3 and found[1] == 3,
|
||||
seconds=180,
|
||||
)
|
||||
settle_candidate(candidate, rig.hidden_free, key, 200, seconds=120)
|
||||
settle_candidate(candidate, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=120)
|
||||
markers: Final = _burst_markers("kill-after", 10)
|
||||
served: Final = _send_all(lambda marker: fresh_chat(candidate, rig.hidden_free, key, marker), markers)
|
||||
assert [response.status_code for response in served] == [200] * 10
|
||||
for row in landed_all_once(key, frozenset(markers)):
|
||||
assert_free_row(row, rig.hidden_free)
|
||||
|
|
@ -9,6 +9,8 @@ See: https://github.com/BerriAI/litellm/issues/24770
|
|||
|
||||
import copy
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
|
||||
from litellm.router import Router
|
||||
|
|
@ -39,10 +41,7 @@ class TestUnmappedModelBudgetEnforcement:
|
|||
]
|
||||
)
|
||||
result = _is_model_cost_zero(model="custom-model", llm_router=router)
|
||||
assert result is False, (
|
||||
"Unmapped model should enforce budget (return False), "
|
||||
"not bypass it (return True)"
|
||||
)
|
||||
assert result is False, "Unmapped model should enforce budget (return False), not bypass it (return True)"
|
||||
|
||||
def test_explicitly_free_model_bypasses_budget(self):
|
||||
"""A model with explicit cost=0 in model_info should bypass budget."""
|
||||
|
|
@ -65,9 +64,7 @@ class TestUnmappedModelBudgetEnforcement:
|
|||
]
|
||||
)
|
||||
result = _is_model_cost_zero(model="free-model", llm_router=router)
|
||||
assert (
|
||||
result is True
|
||||
), "Explicitly free model should bypass budget (return True)"
|
||||
assert result is True, "Explicitly free model should bypass budget (return True)"
|
||||
|
||||
def test_known_paid_model_enforces_budget(self):
|
||||
"""A model in the cost map with non-zero costs should enforce budget."""
|
||||
|
|
@ -101,9 +98,7 @@ class TestUnmappedModelBudgetEnforcement:
|
|||
]
|
||||
)
|
||||
result = _is_model_cost_zero(model="free-via-params", llm_router=router)
|
||||
assert (
|
||||
result is True
|
||||
), "Model with explicit cost=0 in litellm_params should bypass budget"
|
||||
assert result is True, "Model with explicit cost=0 in litellm_params should bypass budget"
|
||||
|
||||
def test_cache_invalidates_on_in_place_pricing_update(self):
|
||||
"""
|
||||
|
|
@ -285,9 +280,12 @@ class TestUnmappedModelBudgetEnforcement:
|
|||
"An aliased PTU group must not be read as free"
|
||||
)
|
||||
|
||||
def test_hidden_model_group_alias_enforces_budget(self):
|
||||
"""A hidden alias keeps budget enforced: get_model_group_info() returns None for it,
|
||||
so the cost is unknown before the configuration gate is reached."""
|
||||
def test_hidden_model_group_alias_to_free_model_bypasses_budget(self):
|
||||
"""A hidden alias to an explicitly free group bypasses budget, like the group itself.
|
||||
|
||||
``get_model_group_info`` returns None for hidden aliases, so the alias must be
|
||||
resolved to its target group before the cost lookup.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
|
|
@ -304,7 +302,22 @@ class TestUnmappedModelBudgetEnforcement:
|
|||
model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}},
|
||||
)
|
||||
|
||||
assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is False
|
||||
assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is True
|
||||
|
||||
def test_hidden_model_group_alias_to_paid_model_enforces_budget(self):
|
||||
"""A hidden alias to a priced group keeps budget enforced."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "paid-model",
|
||||
"litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-fake"},
|
||||
"model_info": {"id": "paid-model-id"},
|
||||
},
|
||||
],
|
||||
model_group_alias={"hidden-paid-alias": {"model": "paid-model", "hidden": True}},
|
||||
)
|
||||
|
||||
assert _is_model_cost_zero(model="hidden-paid-alias", llm_router=router) is False
|
||||
|
||||
def test_dangling_model_group_alias_enforces_budget(self):
|
||||
"""An alias pointing at a group that does not exist keeps budget enforced."""
|
||||
|
|
@ -326,6 +339,115 @@ class TestUnmappedModelBudgetEnforcement:
|
|||
|
||||
assert _is_model_cost_zero(model="dangling-alias", llm_router=router) is False
|
||||
|
||||
def test_repointed_hidden_alias_does_not_reuse_cached_free_result(self):
|
||||
"""Repointing a hidden alias from a free group to a paid group re-evaluates the cost.
|
||||
|
||||
``Router.update_settings`` is the one runtime path that rewrites the alias map (the
|
||||
proxy's config update applies ``router_settings`` through it), so the cached verdict
|
||||
has to drop there.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "free-model",
|
||||
"litellm_params": {
|
||||
"model": "ollama/llama2",
|
||||
"api_base": "http://localhost:11434",
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
},
|
||||
"model_info": {"id": "free-model-id"},
|
||||
},
|
||||
{
|
||||
"model_name": "paid-model",
|
||||
"litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-fake"},
|
||||
"model_info": {"id": "paid-model-id"},
|
||||
},
|
||||
],
|
||||
model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}},
|
||||
)
|
||||
|
||||
assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is True
|
||||
router.update_settings(model_group_alias={"hidden-alias": {"model": "paid-model", "hidden": True}})
|
||||
assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is False
|
||||
|
||||
@pytest.mark.parametrize("alias_name_first", [True, False])
|
||||
def test_alias_shadowing_a_real_group_gives_each_name_its_own_verdict(self, alias_name_first: bool):
|
||||
"""An alias whose name is also a real PTU-priced group never shares a verdict with its target.
|
||||
|
||||
The verdict is cached per requested name, so whichever name is asked first, the free target
|
||||
stays free and the shadowed PTU name stays enforced.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "free-model",
|
||||
"litellm_params": {
|
||||
"model": "ollama/llama2",
|
||||
"api_base": "http://localhost:11434",
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
},
|
||||
"model_info": {"id": "free-model-id"},
|
||||
},
|
||||
{
|
||||
"model_name": "ptu-model",
|
||||
"litellm_params": {
|
||||
"model": "azure/ptu-deployment",
|
||||
"api_base": "https://fake.openai.azure.com",
|
||||
"api_key": "sk-fake",
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
},
|
||||
"model_info": {"id": "ptu-model-id", "ptu_count": 100, "cost_per_ptu_per_hour": 2.0},
|
||||
},
|
||||
],
|
||||
model_group_alias={"ptu-model": "free-model"},
|
||||
)
|
||||
order = ("ptu-model", "free-model") if alias_name_first else ("free-model", "ptu-model")
|
||||
expected = {"ptu-model": False, "free-model": True}
|
||||
|
||||
assert [_is_model_cost_zero(model=name, llm_router=router) for name in order] == [
|
||||
expected[name] for name in order
|
||||
]
|
||||
assert [_is_model_cost_zero(model=name, llm_router=router) for name in order] == [
|
||||
expected[name] for name in order
|
||||
], "the cached verdicts must match the first evaluation"
|
||||
|
||||
def test_alias_chain_through_a_priced_group_enforces_budget(self):
|
||||
"""An alias to a group that is itself an alias key resolves one hop, like the router does.
|
||||
|
||||
The router serves ``chain-smart`` with the real ``chain-legacy`` deployment, which is priced,
|
||||
so following the second hop to the free group would waive the budget for a paid call.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "chain-legacy",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "sk-fake",
|
||||
"input_cost_per_token": 0.0000002,
|
||||
"output_cost_per_token": 0.0000012,
|
||||
},
|
||||
"model_info": {"id": "chain-legacy-id"},
|
||||
},
|
||||
{
|
||||
"model_name": "free-model",
|
||||
"litellm_params": {
|
||||
"model": "ollama/llama2",
|
||||
"api_base": "http://localhost:11434",
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
},
|
||||
"model_info": {"id": "free-model-id"},
|
||||
},
|
||||
],
|
||||
model_group_alias={"chain-smart": "chain-legacy", "chain-legacy": "free-model"},
|
||||
)
|
||||
|
||||
assert _is_model_cost_zero(model="chain-smart", llm_router=router) is False
|
||||
|
||||
def test_handles_router_without_zero_cost_cache_attribute(self):
|
||||
"""Tolerate router-like objects (e.g. ``MagicMock`` stand-ins) that
|
||||
do not expose ``_zero_cost_cache`` — the auth check must still
|
||||
|
|
|
|||
|
|
@ -2278,6 +2278,73 @@ def test_model_group_info_cost_none_for_unpriced_deployment_but_zero_when_declar
|
|||
assert priced.output_cost_per_token is not None and priced.output_cost_per_token > 0
|
||||
|
||||
|
||||
def _alias_cost_router() -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "vllm-free",
|
||||
"litellm_params": {
|
||||
"model": "openai/my-vllm-free",
|
||||
"api_key": "fake",
|
||||
"api_base": "http://localhost:8000/v1",
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": 0,
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-priced",
|
||||
"litellm_params": {"model": "gpt-4o", "api_key": "fake"},
|
||||
},
|
||||
],
|
||||
model_group_alias={"hidden-free": {"model": "vllm-free", "hidden": True}, "visible": "vllm-free"},
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_group_info_include_hidden_resolves_a_hidden_alias():
|
||||
router = _alias_cost_router()
|
||||
|
||||
assert router.get_model_group_info(model_group="hidden-free") is None
|
||||
|
||||
hidden: Final = router.get_model_group_info(model_group="hidden-free", include_hidden=True)
|
||||
assert hidden is not None
|
||||
assert hidden.model_group == "hidden-free"
|
||||
assert hidden.input_cost_per_token == 0
|
||||
assert hidden.output_cost_per_token == 0
|
||||
|
||||
|
||||
def test_update_settings_model_group_alias_drops_cached_group_info():
|
||||
router = _alias_cost_router()
|
||||
before: Final = router.cached_model_group_info("visible")
|
||||
assert before is not None and before.input_cost_per_token == 0
|
||||
|
||||
router.update_settings(model_group_alias={"visible": "gpt-priced"})
|
||||
|
||||
after: Final = router.cached_model_group_info("visible")
|
||||
assert after is not None
|
||||
assert after.input_cost_per_token is not None and after.input_cost_per_token > 0
|
||||
|
||||
|
||||
def test_switch_routing_strategy_installs_lar1_then_restores_the_default_selector():
|
||||
router = _alias_cost_router()
|
||||
|
||||
router._switch_routing_strategy(
|
||||
"lar1",
|
||||
{
|
||||
"routing_strategy_args": {
|
||||
"confidence_threshold_low": 0.1,
|
||||
"confidence_threshold_medium": 0.3,
|
||||
"confidence_threshold_high": 0.9,
|
||||
}
|
||||
},
|
||||
)
|
||||
assert router.routing_strategy == "lar1"
|
||||
assert "async_get_available_deployment" in router.__dict__
|
||||
|
||||
router._switch_routing_strategy("usage-based-routing-v2", {})
|
||||
assert router.lowesttpm_logger_v2 is not None
|
||||
assert "async_get_available_deployment" not in router.__dict__
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"value,expected",
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue