test: add custom pricing tests

This commit is contained in:
mubashir1osmani 2026-06-19 01:34:33 -07:00
parent ff7c3c55c7
commit 61e259cc9d
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
4 changed files with 291 additions and 1 deletions

View file

@ -32,6 +32,8 @@ from models import (
KeyInfo,
KeyInfoParams,
KeyInfoResponse,
ModelInfoEntry,
ModelInfoResponse,
SpendLogRow,
SpendLogs,
SpendLogsParams,
@ -94,6 +96,18 @@ class Gateway:
)
).info
def model_info(self) -> list[ModelInfoEntry]:
"""Every configured deployment with the price the proxy resolved for it
(config override merged over cost-map defaults)."""
return unwrap(
self.transport.get(
"/model/info",
headers=self.transport.master,
params=NoBody(),
response_type=ModelInfoResponse,
)
).data
# ---- LLM calls ------------------------------------------------------
def chat(self, key: str, body: ChatBody) -> Result[ChatResponse]:

View file

@ -103,6 +103,17 @@ model_list:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
# Custom per-token pricing exercised by llm_translation/test_custom_pricing_e2e.py.
# Rates deliberately exceed canonical gemini-2.5-flash (input 3e-7 / output 2.5e-6)
# so an override that is ignored or under-applied reports spend at the base rate
# and fails that test. The test reads these same rates back from this file.
- model_name: custom-priced-flash
litellm_params:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
input_cost_per_token: 0.00005
output_cost_per_token: 0.0001
# embedding models
- model_name: openai-text-embedding-3-small
litellm_params:

View file

@ -0,0 +1,212 @@
"""Live e2e: a model's custom per-token pricing is loaded, billed, and isolated.
The gateway config declares ``custom-priced-flash`` (gemini-2.5-flash underneath)
with input/output rates deliberately far above the canonical gemini price, read
back here from the same config file. Three behaviors are checked independently:
- billing: a real call's logged cost breakdown charges input and output tokens at
the custom rates, each component checked separately (a base-rate bill lands
~100x lower; a swapped input/output rate passes a total-only check but not this)
- reporting: /model/info surfaces those rates for the model
- isolation: gemini-2.5-flash shares the same underlying gemini/gemini-2.5-flash
but sets no override, so it must keep its own price; an override that leaks into
the shared cost map misprices it. This currently fails on a real gap and is left
failing rather than weakened to pass.
"""
import time
from dataclasses import dataclass
from pathlib import Path
import pytest
import yaml
from pydantic import BaseModel, RootModel
from e2e_config import unique_marker
from e2e_http import Success, unwrap
from models import ChatBody, ChatMessage, CustomPricing, ModelInfoEntry, SpendLogsParams
from passthrough_client import PassthroughClient
pytestmark = pytest.mark.e2e
CUSTOM_MODEL = "custom-priced-flash"
BASE_MODEL = "gemini-2.5-flash"
CONFIG_PATH = Path(__file__).resolve().parents[1] / "gateway" / "litellm-config.yml"
@dataclass(frozen=True, slots=True)
class _Rates:
input_per_token: float
output_per_token: float
class _ConfiguredModel(BaseModel):
model_name: str
litellm_params: CustomPricing
class _GatewayConfig(BaseModel):
model_list: list[_ConfiguredModel]
class _CostBreakdown(BaseModel):
input_cost: float | None = None
output_cost: float | None = None
class _RowMetadata(BaseModel):
cost_breakdown: _CostBreakdown | None = None
class _SpendRow(BaseModel):
request_id: str | None = None
prompt_tokens: int | None = None
completion_tokens: int | None = None
metadata: _RowMetadata | None = None
class _SpendRows(RootModel[list[_SpendRow]]):
pass
def _approx_equal(actual: float, expected: float) -> bool:
"""Within 1% or 1e-9 absolute - spend math, not exact float identity."""
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
def _configured_pricing(model_name: str) -> _Rates:
"""The custom rates declared for `model_name` in the gateway config the proxy
runs with - the source of truth the billed and reported prices are checked
against."""
config = _GatewayConfig.model_validate(yaml.safe_load(CONFIG_PATH.read_text()))
for entry in config.model_list:
if entry.model_name == model_name:
pricing = entry.litellm_params
assert pricing.input_cost_per_token and pricing.output_cost_per_token, (
f"{model_name} declares no custom per-token rates in {CONFIG_PATH.name}"
)
return _Rates(pricing.input_cost_per_token, pricing.output_cost_per_token)
pytest.fail(f"{model_name} not found in {CONFIG_PATH.name}")
def _model_info_entry(
entries: list[ModelInfoEntry], model_name: str
) -> ModelInfoEntry:
for entry in entries:
if entry.model_name == model_name:
return entry
pytest.fail(f"{model_name} absent from /model/info; the override did not load")
def _poll_breakdown_row(
client: PassthroughClient, key: str, response_id: str | None
) -> _SpendRow:
"""Poll /spend/logs until the call's row lands with a cost breakdown (rows
flush ~60s behind the call via proxy_batch_write_at)."""
deadline = time.monotonic() + client.gateway.poll_timeout
while time.monotonic() < deadline:
result = client.gateway.transport.get(
"/spend/logs",
headers=client.gateway.transport.master,
params=SpendLogsParams(api_key=key),
response_type=_SpendRows,
)
match result:
case Success(data=data):
rows = data.root
case _:
rows = []
priced = [
row
for row in rows
if row.metadata
and row.metadata.cost_breakdown
and row.metadata.cost_breakdown.input_cost is not None
]
for row in priced:
if response_id and row.request_id == response_id:
return row
if priced:
return priced[0]
time.sleep(client.gateway.poll_interval)
pytest.fail("no spend row with a cost breakdown landed before the deadline")
def test_custom_pricing_is_billed_at_configured_rate(
client: PassthroughClient, scoped_key: str
) -> None:
rates = _configured_pricing(CUSTOM_MODEL)
chat = unwrap(
client.gateway.chat(
scoped_key,
ChatBody(
model=CUSTOM_MODEL,
messages=[
ChatMessage(
role="user", content=f"reply with one word {unique_marker()}"
)
],
max_tokens=16,
),
)
)
row = _poll_breakdown_row(client, scoped_key, chat.id)
assert row.metadata and row.metadata.cost_breakdown # guaranteed by the poll
breakdown = row.metadata.cost_breakdown
prompt = row.prompt_tokens or 0
completion = row.completion_tokens or 0
assert prompt > 0 and completion > 0, f"call tokens not logged on the row: {row}"
input_cost = breakdown.input_cost
output_cost = breakdown.output_cost
assert input_cost is not None and output_cost is not None, (
f"row cost breakdown missing input/output cost: {breakdown}"
)
assert _approx_equal(input_cost, prompt * rates.input_per_token), (
f"input_cost {input_cost} != {prompt} tokens * {rates.input_per_token} "
f"= {prompt * rates.input_per_token}"
)
assert _approx_equal(output_cost, completion * rates.output_per_token), (
f"output_cost {output_cost} != {completion} tokens * {rates.output_per_token} "
f"= {completion * rates.output_per_token}"
)
def test_model_info_reports_custom_pricing(client: PassthroughClient) -> None:
rates = _configured_pricing(CUSTOM_MODEL)
entry = _model_info_entry(client.gateway.model_info(), CUSTOM_MODEL)
assert entry.litellm_params.input_cost_per_token == rates.input_per_token, (
f"/model/info litellm_params input rate "
f"{entry.litellm_params.input_cost_per_token} != configured "
f"{rates.input_per_token}"
)
assert entry.litellm_params.output_cost_per_token == rates.output_per_token, (
f"/model/info litellm_params output rate "
f"{entry.litellm_params.output_cost_per_token} != configured "
f"{rates.output_per_token}"
)
def test_custom_pricing_is_isolated_from_sibling_deployment(
client: PassthroughClient,
) -> None:
entries = {entry.model_name: entry for entry in client.gateway.model_info()}
custom = entries.get(CUSTOM_MODEL)
base = entries.get(BASE_MODEL)
assert custom is not None, f"{CUSTOM_MODEL} absent from /model/info"
assert base is not None, f"{BASE_MODEL} absent from /model/info"
# custom-priced-flash overrides pricing; gemini-2.5-flash shares the same
# underlying gemini/gemini-2.5-flash but sets no override, so it must keep its
# own price. Equal rates mean the override leaked into the shared cost map.
assert (
base.model_info.input_cost_per_token != custom.model_info.input_cost_per_token
), (
f"{BASE_MODEL} input rate {base.model_info.input_cost_per_token} matches "
f"{CUSTOM_MODEL}'s override {custom.model_info.input_cost_per_token}; "
f"per-deployment custom pricing is not isolated"
)

View file

@ -6,7 +6,7 @@ response validates without mirroring every proxy field. No untyped dicts.
from __future__ import annotations
from pydantic import BaseModel, RootModel
from pydantic import BaseModel, ConfigDict, RootModel
# ---------- keys ----------
@ -198,3 +198,56 @@ class RouteSpec(RootModel[dict[str, object]]):
class OpenAPISchema(BaseModel):
paths: dict[str, RouteSpec] = {}
# ---------- model info / custom pricing ----------
class CustomPricing(BaseModel):
"""The per-token custom-pricing fields a deployment can override in
litellm_params - the token-cost subset of litellm's CustomPricingLiteLLMParams
the proxy applies to chat spend. All optional: a config sets only what it
overrides, and /model/info echoes the rates the proxy resolved."""
model_config = ConfigDict(extra="ignore")
input_cost_per_token: float | None = None
output_cost_per_token: float | None = None
cache_read_input_token_cost: float | None = None
cache_creation_input_token_cost: float | None = None
def overrides(self) -> dict[str, float]:
"""The rates actually declared (non-null) - e.g. those a config.yml sets."""
declared = {
"input_cost_per_token": self.input_cost_per_token,
"output_cost_per_token": self.output_cost_per_token,
"cache_read_input_token_cost": self.cache_read_input_token_cost,
"cache_creation_input_token_cost": self.cache_creation_input_token_cost,
}
return {field: rate for field, rate in declared.items() if rate is not None}
def token_cost(self, prompt_tokens: int, completion_tokens: int) -> float:
"""Spend for a fresh (uncached) call under these rates: the proxy's
custom-pricing formula (prompt * input + completion * output)."""
assert (
self.input_cost_per_token is not None
and self.output_cost_per_token is not None
), "custom pricing has no per-token rates"
return (
prompt_tokens * self.input_cost_per_token
+ completion_tokens * self.output_cost_per_token
)
class ModelInfoEntry(BaseModel):
"""One /model/info row. `litellm_params` is the configured deployment (carries
any custom-pricing override); `model_info` is the price the proxy resolved for
it - the override merged over the cost-map defaults."""
model_config = ConfigDict(protected_namespaces=())
model_name: str
litellm_params: CustomPricing = CustomPricing()
model_info: CustomPricing = CustomPricing()
class ModelInfoResponse(BaseModel):
data: list[ModelInfoEntry] = []