mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(cost): price rule-only model names at the deployment's rate (#44144)
* fix(cost): price streamed aliases that only match a capability rule from the deployment model Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(cost): satisfy basedpyright delta after the cost-candidate sort Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * perf(cost): skip the capability-rule check for exact cost-map keys Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(spend): price streamed alias rows from the deployment's own rates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(spend): isolate streamed alias deployments with per-run model names Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(streaming): keep the unpriceable-stamp case on a truly unmapped model Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): treat capability-rule matches as unmapped in every cost lookup Route every model-info lookup on cost paths through get_priced_model_info and _cached_get_priced_model_info_helper, which raise ModelNotMappedError when the only match is a pricing-free fallback-generalizations rule. An alias that matches a capability rule now falls through to the deployment's real model instead of billing 0, and the earlier candidate-sorting fix is reverted since the priced helper is the single choke point. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(spend): assert the rule alias bills the same as the plain alias Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): treat router-registered rule-only model_cost entries as unmapped Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): count string rates and check pricing before the capability-rule match Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(cost): repoint cost-path patches at get_priced_model_info Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(cost): give the model_cost row cast a cast-ok reason Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): type the lazy get_priced_model_info export Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(cost): narrow the fix to ordering rule-only cost candidates last Drops get_priced_model_info and its call-site swaps, the lazy-import entry, and the cost-path code check. Only the candidate sort in completion_cost and pricing_entry_for_cost_calc stays, with its tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(spend): cover rule-only base_model billing on every endpoint, client and failure mode Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry <kerry@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ca1994e403
commit
f9a32ffcb5
4 changed files with 547 additions and 10 deletions
|
|
@ -137,6 +137,8 @@ from litellm.utils import (
|
|||
TextCompletionResponse,
|
||||
TranscriptionResponse,
|
||||
_cached_get_model_info_helper,
|
||||
_get_model_info_from_generalization,
|
||||
_get_potential_model_names,
|
||||
token_counter,
|
||||
)
|
||||
|
||||
|
|
@ -924,6 +926,22 @@ def _get_response_model(completion_response: object) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def _prices_only_via_capability_rule(model: str | None, custom_llm_provider: str | None) -> bool:
|
||||
if model is None or model in litellm.model_cost or f"{custom_llm_provider}/{model}" in litellm.model_cost:
|
||||
return False
|
||||
try:
|
||||
return (
|
||||
_get_model_info_from_generalization(
|
||||
model=model,
|
||||
potential_model_names=_get_potential_model_names(model=model, custom_llm_provider=custom_llm_provider),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
is not None
|
||||
)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
_GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER: Final[dict] = {
|
||||
# ON_DEMAND_PRIORITY maps to "priority" — selects input_cost_per_token_priority, etc.
|
||||
"ON_DEMAND_PRIORITY": "priority",
|
||||
|
|
@ -1440,12 +1458,10 @@ def completion_cost(
|
|||
region_name=region_name,
|
||||
)
|
||||
|
||||
potential_model_names: Final = [
|
||||
selected_model,
|
||||
_get_response_model(completion_response),
|
||||
]
|
||||
if model is not None:
|
||||
potential_model_names.append(model)
|
||||
potential_model_names: Final = sorted(
|
||||
(selected_model, _get_response_model(completion_response), *((model,) if model is not None else ())),
|
||||
key=lambda candidate: _prices_only_via_capability_rule(candidate, cast(str | None, custom_llm_provider)),
|
||||
)
|
||||
|
||||
for idx, model in enumerate(potential_model_names):
|
||||
try:
|
||||
|
|
@ -1460,7 +1476,7 @@ def completion_cost(
|
|||
else:
|
||||
usage_obj = getattr(completion_response, "usage", {})
|
||||
if isinstance(usage_obj, BaseModel) and not _is_known_usage_objects(usage_obj=usage_obj):
|
||||
_usage_for_dump = cast(BaseModel, usage_obj)
|
||||
_usage_for_dump = usage_obj
|
||||
setattr(
|
||||
completion_response,
|
||||
"usage",
|
||||
|
|
@ -1469,7 +1485,7 @@ def completion_cost(
|
|||
if usage_obj is None:
|
||||
_usage = {}
|
||||
elif isinstance(usage_obj, BaseModel):
|
||||
_usage = cast(BaseModel, usage_obj).model_dump()
|
||||
_usage = usage_obj.model_dump()
|
||||
else:
|
||||
_usage = usage_obj
|
||||
|
||||
|
|
@ -2124,7 +2140,10 @@ def pricing_entry_for_cost_calc(
|
|||
router_model_id=router_model_id,
|
||||
region_name=region_name,
|
||||
)
|
||||
candidates: Final = (selected_model, _get_response_model(completion_response), model)
|
||||
candidates: Final = sorted(
|
||||
(selected_model, _get_response_model(completion_response), model),
|
||||
key=lambda candidate: _prices_only_via_capability_rule(candidate, custom_llm_provider),
|
||||
)
|
||||
resolved: Final = next(
|
||||
(info for info in (_cost_map_model_info(name, custom_llm_provider) for name in candidates if name) if info),
|
||||
None,
|
||||
|
|
|
|||
464
tests/integration/spend/test_stream_alias_billing.py
Normal file
464
tests/integration/spend/test_stream_alias_billing.py
Normal file
|
|
@ -0,0 +1,464 @@
|
|||
"""A model name that only matches a capability rule never zeroes the deployment's price (LIT-9065).
|
||||
|
||||
The proxy restamps every streamed chunk with the client's alias, so end-of-stream cost calculation can see
|
||||
"claude-opus-4.8-<digits>" before the deployment's model. That name is no cost-map key but matches the claude
|
||||
capability generalization rules, whose model info carries no prices, so the dotted alias must bill exactly what
|
||||
the plain alias "integration-<hex>" bills at the same deployment rates. The same holds for a deployment whose
|
||||
model_info.base_model only matches a rule, on every endpoint and client
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from hashlib import sha256
|
||||
from typing import Final
|
||||
from uuid import uuid4
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
|
||||
def _sse_event(name: str, payload: dict[str, JsonValue]) -> bytes:
|
||||
return f"event: {name}\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n".encode()
|
||||
|
||||
|
||||
def _anthropic_stream(request: Request) -> Reply:
|
||||
assert request.target.endswith("/v1/messages"), request.target
|
||||
body: Final = json.loads(request.body)
|
||||
assert body["model"] == "claude-opus-4-8" and body["stream"] is True, body
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(
|
||||
_sse_event(
|
||||
"message_start",
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": f"msg_{uuid4().hex[:12]}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-opus-4-8",
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 30, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
),
|
||||
_sse_event(
|
||||
"content_block_start",
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
),
|
||||
_sse_event(
|
||||
"content_block_delta",
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hi"}},
|
||||
),
|
||||
_sse_event("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
_sse_event(
|
||||
"message_delta",
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
|
||||
"usage": {"output_tokens": 40},
|
||||
},
|
||||
),
|
||||
_sse_event("message_stop", {"type": "message_stop"}),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _deployment(
|
||||
scenario: Scenario,
|
||||
model_name: str,
|
||||
litellm_params: dict[str, JsonValue],
|
||||
model_info: dict[str, JsonValue] | None = None,
|
||||
) -> str:
|
||||
created: Final = scenario.gateway.post(
|
||||
"/model/new", {"model_name": model_name, "litellm_params": litellm_params, "model_info": model_info or {}}
|
||||
)
|
||||
identity: Final = string_value(object_value(created["model_info"])["id"])
|
||||
scenario.cleanups.callback(scenario.delete_model, identity)
|
||||
return model_name
|
||||
|
||||
|
||||
def _streamed_spend(gateway: Gateway, scenario: Scenario, model: str, content: str) -> dict[str, JsonValue]:
|
||||
key: Final = scenario.key(models=[model])
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": content}],
|
||||
"stream": True,
|
||||
"stream_options": {"include_usage": True},
|
||||
},
|
||||
key=key,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE api_key=%s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return rows[0]
|
||||
|
||||
|
||||
def _listed_deployments(gateway: Gateway, model_name: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
entries: Final = gateway.get("/model/info")["data"]
|
||||
assert isinstance(entries, list)
|
||||
return tuple(object_value(entry) for entry in entries if object_value(entry)["model_name"] == model_name)
|
||||
|
||||
|
||||
def _deployment_pricing(gateway: Gateway, model_name: str) -> dict[str, JsonValue]:
|
||||
listed: Final = eventually(lambda: _listed_deployments(gateway, model_name), lambda found: len(found) == 1)
|
||||
return object_value(listed[0]["model_info"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"litellm_params",
|
||||
(
|
||||
pytest.param(
|
||||
lambda _: {"model": "vertex_ai/claude-opus-4-8@default", "mock_response": "hi"},
|
||||
id="vertex-mock-response",
|
||||
),
|
||||
pytest.param(
|
||||
lambda wire_url: {
|
||||
"model": "anthropic/claude-opus-4-8",
|
||||
"api_key": "integration-provider-key",
|
||||
"api_base": wire_url,
|
||||
},
|
||||
id="anthropic-upstream",
|
||||
),
|
||||
),
|
||||
)
|
||||
@pytest.mark.timeout(180)
|
||||
def test_streamed_alias_matching_a_capability_rule_bills_the_deployment_price(
|
||||
gateway: Gateway, litellm_params: Callable[[str], dict[str, JsonValue]]
|
||||
) -> None:
|
||||
with wire_server(_anthropic_stream) as wire, gateway.scenario() as scenario:
|
||||
content: Final = f"alias billing {uuid4().hex}"
|
||||
plain_alias: Final = f"integration-{uuid4().hex}"
|
||||
rule_alias: Final = f"claude-opus-4.8-{uuid4().int % 10**8:08d}"
|
||||
exact_row: Final = _streamed_spend(
|
||||
gateway, scenario, _deployment(scenario, plain_alias, litellm_params(wire.url)), content
|
||||
)
|
||||
alias_row: Final = _streamed_spend(
|
||||
gateway, scenario, _deployment(scenario, rule_alias, litellm_params(wire.url)), content
|
||||
)
|
||||
|
||||
for model_name, row in ((plain_alias, exact_row), (rule_alias, alias_row)):
|
||||
pricing: Final = _deployment_pricing(gateway, model_name)
|
||||
input_rate: Final = float(str(pricing["input_cost_per_token"]))
|
||||
output_rate: Final = float(str(pricing["output_cost_per_token"]))
|
||||
uplift: Final = float(str(pricing["regional_endpoint_uplift_multiplier"] or 1))
|
||||
assert input_rate > 0 and output_rate > 0, pricing
|
||||
assert float(str(row["spend"])) == pytest.approx(
|
||||
uplift
|
||||
* (float(str(row["prompt_tokens"])) * input_rate + float(str(row["completion_tokens"])) * output_rate)
|
||||
), (model_name, row, pricing)
|
||||
|
||||
|
||||
INPUT_TOKENS: Final = 30
|
||||
OUTPUT_TOKENS: Final = 40
|
||||
DEPLOYMENT_MODEL: Final = "anthropic/claude-opus-4-8"
|
||||
ENDPOINTS: Final = ("/v1/chat/completions", "/v1/messages", "/v1/responses")
|
||||
|
||||
|
||||
def _rule_only_name() -> str:
|
||||
return f"claude-opus-4.8-{uuid4().int % 10**8:08d}"
|
||||
|
||||
|
||||
def _anthropic_reply(request: Request) -> Reply:
|
||||
body: Final = json.loads(request.body)
|
||||
if body.get("stream") is True:
|
||||
return _anthropic_stream(request)
|
||||
assert request.target.endswith("/v1/messages") and body["model"] == "claude-opus-4-8", (request.target, body)
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": f"msg_{uuid4().hex[:12]}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-opus-4-8",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": INPUT_TOKENS, "output_tokens": OUTPUT_TOKENS},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_params(wire_url: str) -> dict[str, JsonValue]:
|
||||
return {"model": DEPLOYMENT_MODEL, "api_key": "integration-provider-key", "api_base": wire_url}
|
||||
|
||||
|
||||
def _listed_rates(scenario: Scenario, model: str, wire_url: str) -> tuple[float, float]:
|
||||
model_name: Final = _deployment(
|
||||
scenario, f"integration-{uuid4().hex}", {**_anthropic_params(wire_url), "model": model}
|
||||
)
|
||||
pricing: Final = _deployment_pricing(scenario.gateway, model_name)
|
||||
uplift: Final = float(str(pricing.get("regional_endpoint_uplift_multiplier") or 1))
|
||||
rates: Final = (
|
||||
uplift * float(str(pricing["input_cost_per_token"])),
|
||||
uplift * float(str(pricing["output_cost_per_token"])),
|
||||
)
|
||||
assert rates[0] > 0 and rates[1] > 0, pricing
|
||||
return rates
|
||||
|
||||
|
||||
def _body(path: str, model: str, content: str, stream: bool) -> dict[str, JsonValue]:
|
||||
match path:
|
||||
case "/v1/chat/completions":
|
||||
return {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": content}],
|
||||
"stream": stream,
|
||||
**({"stream_options": {"include_usage": True}} if stream else {}),
|
||||
}
|
||||
case "/v1/messages":
|
||||
return {
|
||||
"model": model,
|
||||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": content}],
|
||||
"stream": stream,
|
||||
}
|
||||
case _:
|
||||
return {"model": model, "input": content, "stream": stream}
|
||||
|
||||
|
||||
def _spend_rows(key: str, count: int) -> list[dict[str, JsonValue]]:
|
||||
return eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, spend, prompt_tokens, completion_tokens, status, cache_hit FROM "LiteLLM_SpendLogs"'
|
||||
' WHERE api_key=%s ORDER BY "startTime"',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
),
|
||||
lambda values: len(values) == count,
|
||||
seconds=90,
|
||||
)
|
||||
|
||||
|
||||
def _rule_only_base_model_deployment(scenario: Scenario, wire_url: str) -> str:
|
||||
return _deployment(
|
||||
scenario, f"integration-{uuid4().hex}", _anthropic_params(wire_url), {"base_model": _rule_only_name()}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", (False, True), ids=("non-streaming", "streaming"))
|
||||
@pytest.mark.parametrize("path", ENDPOINTS)
|
||||
@pytest.mark.timeout(180)
|
||||
def test_rule_only_base_model_bills_the_deployment_price(gateway: Gateway, path: str, stream: bool) -> None:
|
||||
with wire_server(_anthropic_reply) as wire, gateway.scenario() as scenario:
|
||||
input_rate, output_rate = _listed_rates(scenario, DEPLOYMENT_MODEL, wire.url)
|
||||
model: Final = _rule_only_base_model_deployment(scenario, wire.url)
|
||||
key: Final = scenario.key(models=[model])
|
||||
|
||||
response: Final = gateway.request(
|
||||
"POST", path, _body(path, model, f"base model {uuid4().hex}", stream), key=key
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_rows(key, 1)[0]
|
||||
assert (row["prompt_tokens"], row["completion_tokens"]) == (INPUT_TOKENS, OUTPUT_TOKENS), row
|
||||
assert float(str(row["spend"])) == pytest.approx(INPUT_TOKENS * input_rate + OUTPUT_TOKENS * output_rate), row
|
||||
|
||||
|
||||
def _openai_sync_chat_stream(base_url: str, key: str, model: str, content: str) -> None:
|
||||
with openai.OpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0) as client:
|
||||
chunks: Final = tuple(
|
||||
client.chat.completions.create(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": content}],
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
)
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == "hi", chunks
|
||||
|
||||
|
||||
def _openai_async_responses(base_url: str, key: str, model: str, content: str) -> None:
|
||||
async def call() -> str:
|
||||
async with openai.AsyncOpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0) as client:
|
||||
return (await client.responses.create(model=model, input=content)).output_text
|
||||
|
||||
assert asyncio.run(call()) == "hi"
|
||||
|
||||
|
||||
def _anthropic_async_messages_stream(base_url: str, key: str, model: str, content: str) -> None:
|
||||
async def call() -> int:
|
||||
async with anthropic.AsyncAnthropic(base_url=base_url, api_key=key, max_retries=0) as client:
|
||||
async with client.messages.stream(
|
||||
model=model, max_tokens=64, messages=[{"role": "user", "content": content}]
|
||||
) as stream:
|
||||
return (await stream.get_final_message()).usage.output_tokens
|
||||
|
||||
assert asyncio.run(call()) == OUTPUT_TOKENS
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"client_call",
|
||||
(
|
||||
pytest.param(_openai_sync_chat_stream, id="openai-sync-chat-stream"),
|
||||
pytest.param(_openai_async_responses, id="openai-async-responses"),
|
||||
pytest.param(_anthropic_async_messages_stream, id="anthropic-async-messages-stream"),
|
||||
),
|
||||
)
|
||||
@pytest.mark.timeout(180)
|
||||
def test_rule_only_base_model_bills_the_deployment_price_through_the_sdks(
|
||||
gateway: Gateway, client_call: Callable[[str, str, str, str], None]
|
||||
) -> None:
|
||||
with wire_server(_anthropic_reply) as wire, gateway.scenario() as scenario:
|
||||
input_rate, output_rate = _listed_rates(scenario, DEPLOYMENT_MODEL, wire.url)
|
||||
model: Final = _rule_only_base_model_deployment(scenario, wire.url)
|
||||
key: Final = scenario.key(models=[model])
|
||||
|
||||
client_call(str(gateway.client.base_url).rstrip("/"), key, model, f"sdk {uuid4().hex}")
|
||||
|
||||
row: Final = _spend_rows(key, 1)[0]
|
||||
assert float(str(row["spend"])) == pytest.approx(INPUT_TOKENS * input_rate + OUTPUT_TOKENS * output_rate), row
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", (False, True), ids=("non-streaming", "streaming"))
|
||||
@pytest.mark.timeout(180)
|
||||
def test_custom_pricing_still_beats_a_rule_only_base_model(gateway: Gateway, stream: bool) -> None:
|
||||
with wire_server(_anthropic_reply) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(
|
||||
scenario,
|
||||
f"integration-{uuid4().hex}",
|
||||
{**_anthropic_params(wire.url), "input_cost_per_token": 0.001, "output_cost_per_token": 0.002},
|
||||
{"base_model": _rule_only_name()},
|
||||
)
|
||||
key: Final = scenario.key(models=[model])
|
||||
path: Final = "/v1/chat/completions"
|
||||
|
||||
response: Final = gateway.request("POST", path, _body(path, model, f"custom {uuid4().hex}", stream), key=key)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert float(str(_spend_rows(key, 1)[0]["spend"])) == pytest.approx(30 * 0.001 + 40 * 0.002)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", (False, True), ids=("non-streaming", "streaming"))
|
||||
@pytest.mark.timeout(180)
|
||||
def test_priced_base_model_still_bills_its_own_price(gateway: Gateway, stream: bool) -> None:
|
||||
with wire_server(_anthropic_reply) as wire, gateway.scenario() as scenario:
|
||||
input_rate, output_rate = _listed_rates(scenario, "anthropic/claude-haiku-4-5", wire.url)
|
||||
model: Final = _deployment(
|
||||
scenario, f"integration-{uuid4().hex}", _anthropic_params(wire.url), {"base_model": "claude-haiku-4-5"}
|
||||
)
|
||||
key: Final = scenario.key(models=[model])
|
||||
path: Final = "/v1/chat/completions"
|
||||
|
||||
response: Final = gateway.request("POST", path, _body(path, model, f"priced {uuid4().hex}", stream), key=key)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_rows(key, 1)[0]
|
||||
assert float(str(row["spend"])) == pytest.approx(INPUT_TOKENS * input_rate + OUTPUT_TOKENS * output_rate), row
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"base_model",
|
||||
(
|
||||
pytest.param("", id="empty"),
|
||||
pytest.param(f"claude-opus-4.8-{'9' * 5000}", id="5kb-rule-only"),
|
||||
pytest.param(f"integration-unmapped-{uuid4().hex}", id="unmapped-no-rule"),
|
||||
),
|
||||
)
|
||||
@pytest.mark.timeout(180)
|
||||
def test_odd_base_model_values_bill_the_deployment_price(gateway: Gateway, base_model: str) -> None:
|
||||
with wire_server(_anthropic_reply) as wire, gateway.scenario() as scenario:
|
||||
input_rate, output_rate = _listed_rates(scenario, DEPLOYMENT_MODEL, wire.url)
|
||||
model: Final = _deployment(
|
||||
scenario, f"integration-{uuid4().hex}", _anthropic_params(wire.url), {"base_model": base_model}
|
||||
)
|
||||
key: Final = scenario.key(models=[model])
|
||||
path: Final = "/v1/chat/completions"
|
||||
|
||||
response: Final = gateway.request("POST", path, _body(path, model, f"odd {uuid4().hex}", True), key=key)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
row: Final = _spend_rows(key, 1)[0]
|
||||
assert float(str(row["spend"])) == pytest.approx(INPUT_TOKENS * input_rate + OUTPUT_TOKENS * output_rate), row
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ENDPOINTS)
|
||||
@pytest.mark.timeout(180)
|
||||
def test_upstream_failure_on_a_rule_only_base_model_logs_a_zero_spend_failure(gateway: Gateway, path: str) -> None:
|
||||
failure: Final = Reply(
|
||||
status=500, body=b'{"type":"error","error":{"type":"api_error","message":"integration upstream down"}}'
|
||||
)
|
||||
with wire_server(lambda _: failure) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _rule_only_base_model_deployment(scenario, wire.url)
|
||||
key: Final = scenario.key(models=[model])
|
||||
|
||||
response: Final = gateway.request("POST", path, _body(path, model, f"down {uuid4().hex}", False), key=key)
|
||||
|
||||
assert response.status_code == 500, response.text
|
||||
row: Final = _spend_rows(key, 1)[0]
|
||||
assert (row["status"], float(str(row["spend"]))) == ("failure", 0.0), row
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
def test_cache_hit_on_a_rule_only_base_model_bills_only_the_first_call(gateway: Gateway) -> None:
|
||||
with wire_server(_anthropic_reply) as wire, gateway.scenario() as scenario:
|
||||
input_rate, output_rate = _listed_rates(scenario, DEPLOYMENT_MODEL, wire.url)
|
||||
model: Final = _rule_only_base_model_deployment(scenario, wire.url)
|
||||
key: Final = scenario.key(models=[model])
|
||||
path: Final = "/v1/chat/completions"
|
||||
body: Final = _body(path, model, f"cached {uuid4().hex}", False)
|
||||
|
||||
responses: Final = tuple(gateway.request("POST", path, body, key=key) for _ in range(2))
|
||||
|
||||
assert [response.status_code for response in responses] == [200, 200], [r.text for r in responses]
|
||||
rows: Final = _spend_rows(key, 2)
|
||||
assert [(row["cache_hit"], float(str(row["spend"]))) for row in rows] == [
|
||||
("None", pytest.approx(INPUT_TOKENS * input_rate + OUTPUT_TOKENS * output_rate)),
|
||||
("True", 0.0),
|
||||
], rows
|
||||
assert len([request for request in wire.drain() if request.target.endswith("/v1/messages")]) == 1
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_burst_through_an_upstream_outage_bills_every_recovered_request_once(gateway: Gateway) -> None:
|
||||
outage: Final = threading.Event()
|
||||
overloaded: Final = Reply(status=529, body=b'{"type":"error","error":{"type":"overloaded_error","message":"x"}}')
|
||||
burst: Final = tuple((path, stream) for path in ENDPOINTS for stream in (False, True)) * 4
|
||||
with (
|
||||
wire_server(lambda request: overloaded if outage.is_set() else _anthropic_reply(request)) as wire,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
input_rate, output_rate = _listed_rates(scenario, DEPLOYMENT_MODEL, wire.url)
|
||||
model: Final = _rule_only_base_model_deployment(scenario, wire.url)
|
||||
key: Final = scenario.key(models=[model])
|
||||
|
||||
def send(cell: tuple[str, bool]) -> httpx.Response:
|
||||
return gateway.request("POST", cell[0], _body(cell[0], model, f"burst {uuid4().hex}", cell[1]), key=key)
|
||||
|
||||
outage.set()
|
||||
with ThreadPoolExecutor(max_workers=len(burst)) as pool:
|
||||
during: Final = tuple(pool.map(send, burst))
|
||||
outage.clear()
|
||||
with ThreadPoolExecutor(max_workers=len(burst)) as pool:
|
||||
after: Final = tuple(pool.map(send, burst))
|
||||
|
||||
assert all(response.status_code != 200 for response in during), [r.status_code for r in during]
|
||||
assert [response.status_code for response in after] == [200] * len(burst), [r.text for r in after]
|
||||
rows: Final = _spend_rows(key, 2 * len(burst))
|
||||
succeeded: Final = tuple(row for row in rows if row["status"] == "success")
|
||||
assert len({row["request_id"] for row in succeeded}) == len(succeeded) == len(burst), rows
|
||||
assert {float(str(row["spend"])) for row in rows if row["status"] != "success"} == {0.0}, rows
|
||||
assert all(
|
||||
float(str(row["spend"])) == pytest.approx(INPUT_TOKENS * input_rate + OUTPUT_TOKENS * output_rate)
|
||||
for row in succeeded
|
||||
), succeeded
|
||||
|
|
@ -15,6 +15,7 @@ from litellm.cost_calculator import (
|
|||
completion_cost,
|
||||
cost_per_token,
|
||||
handle_realtime_stream_cost_calculation,
|
||||
pricing_entry_for_cost_calc,
|
||||
response_cost_calculator,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
|
@ -5639,3 +5640,56 @@ def test_completion_cost_bills_base_when_gemini_serves_on_demand(
|
|||
)
|
||||
|
||||
assert cost == pytest.approx(100 * 0.001 + 50 * 0.002)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider,deployment_model,cost_map_key",
|
||||
[
|
||||
("vertex_ai", "claude-opus-4-8@default", "vertex_ai/claude-opus-4-8@default"),
|
||||
("anthropic", "claude-opus-4-8", "claude-opus-4-8"),
|
||||
],
|
||||
)
|
||||
def test_completion_cost_prices_capability_rule_alias_from_the_deployment(
|
||||
_local_model_cost_map: None, custom_llm_provider: str, deployment_model: str, cost_map_key: str
|
||||
) -> None:
|
||||
"""Streamed proxy chunks carry the client's alias, so the first cost candidate is the
|
||||
provider-prefixed alias. That name matches a claude capability generalization rule (unpriced)
|
||||
and must fall through to the deployment's priced model instead of stopping at $0."""
|
||||
response: Final = ModelResponse(
|
||||
id="chatcmpl_x",
|
||||
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
|
||||
model="claude-opus-4.8",
|
||||
usage=Usage(prompt_tokens=30, completion_tokens=40, total_tokens=70),
|
||||
)
|
||||
row: Final = litellm.model_cost[cost_map_key]
|
||||
expected: Final = 30 * row["input_cost_per_token"] + 40 * row["output_cost_per_token"]
|
||||
assert expected > 0
|
||||
|
||||
assert completion_cost(
|
||||
completion_response=response,
|
||||
model=deployment_model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
) == pytest.approx(expected)
|
||||
|
||||
|
||||
def test_pricing_entry_for_cost_calc_skips_capability_rule_alias(_local_model_cost_map: None) -> None:
|
||||
response: Final = ModelResponse(
|
||||
id="chatcmpl_x",
|
||||
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
|
||||
model="claude-opus-4.8",
|
||||
usage=Usage(prompt_tokens=30, completion_tokens=40, total_tokens=70),
|
||||
)
|
||||
|
||||
resolved: Final = pricing_entry_for_cost_calc(
|
||||
model="claude-opus-4-8@default",
|
||||
completion_response=response,
|
||||
custom_llm_provider="vertex_ai",
|
||||
custom_pricing=None,
|
||||
base_model=None,
|
||||
router_model_id=None,
|
||||
region_name=None,
|
||||
litellm_logging_obj=None,
|
||||
)
|
||||
|
||||
assert resolved is not None
|
||||
assert resolved[0] == "vertex_ai/claude-opus-4-8@default"
|
||||
|
|
|
|||
|
|
@ -4061,7 +4061,7 @@ def test_stream_chunk_builder_skips_stamp_when_cost_is_unpriceable():
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
|
||||
logging_obj: Final = LiteLLMLogging(
|
||||
model="us.anthropic.claude-opus-5",
|
||||
model="unmapped-deployment-without-cost-map-entry",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue