mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix(proxy): price one cost estimate at one moment
The totals, the per-token-type lines and the reported rates each resolved off-peak pricing on their own clock read, so a quote taken as a window opened could bill on one side of the boundary and report rates from the other. /cost/estimate now pins a billing moment for the whole quote, and every rate lookup answers for the pinned moment instead of reading the clock again Claude-Session: https://claude.ai/code/session_011Tn3657NkV6ojLqewL64Kb
This commit is contained in:
parent
43a02e2dbc
commit
ba91588b15
5 changed files with 126 additions and 22 deletions
|
|
@ -6,9 +6,33 @@ be settable from user input. Context variables are scoped to the current
|
|||
asyncio task and cannot be injected via HTTP request bodies.
|
||||
"""
|
||||
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final
|
||||
|
||||
# When True, suppresses async logging and billing for internal sub-calls
|
||||
# (e.g., emulated file-search steps that make nested LLM calls).
|
||||
is_internal_call: Final[ContextVar[bool]] = ContextVar("is_internal_call", default=False)
|
||||
|
||||
# One request prices its totals, its per-token-type lines and the rates it reports on
|
||||
# separate code paths. Each reads the clock for off-peak pricing, so without a pinned
|
||||
# moment they can land on either side of a window boundary and disagree with each other.
|
||||
_billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", default=None)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def pinned_billing_time(moment: datetime) -> Generator[None]:
|
||||
"""Price every rate lookup inside this block at ``moment`` rather than at each one's own clock read."""
|
||||
token: Final = _billing_time.set(moment)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_billing_time.reset(token)
|
||||
|
||||
|
||||
def current_billing_time() -> datetime:
|
||||
"""The pinned billing moment, or now in UTC outside a pinned block."""
|
||||
pinned: Final = _billing_time.get()
|
||||
return pinned if pinned is not None else datetime.now(timezone.utc)
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from typing import Any, Final, Literal, TypedDict, cast
|
|||
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import current_billing_time
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import (
|
||||
select_tier_for_input,
|
||||
|
|
@ -306,7 +307,7 @@ def _is_within_off_peak_window(off_peak_hours_utc: str | Sequence[str], current_
|
|||
than being localised, so callers must pass datetime.now(timezone.utc), never datetime.now(),
|
||||
or every window shifts by the host's offset.
|
||||
"""
|
||||
reference: Final = current_time if current_time is not None else datetime.now(timezone.utc)
|
||||
reference: Final = current_time if current_time is not None else current_billing_time()
|
||||
now: Final = (reference.astimezone(timezone.utc) if reference.tzinfo is not None else reference).time()
|
||||
windows: Final = (off_peak_hours_utc,) if isinstance(off_peak_hours_utc, str) else off_peak_hours_utc
|
||||
for window in windows:
|
||||
|
|
@ -393,7 +394,7 @@ def _is_off_peak(off_peak: Mapping[str, object], current_time: datetime | None =
|
|||
rules: the flat hours_utc windows, which apply every day, or any entry in windows, whose
|
||||
hours apply only on its weekdays.
|
||||
"""
|
||||
reference: Final = current_time if current_time is not None else datetime.now(timezone.utc)
|
||||
reference: Final = current_time if current_time is not None else current_billing_time()
|
||||
reference_utc: Final = (
|
||||
reference.astimezone(timezone.utc) if reference.tzinfo is not None else reference.replace(tzinfo=timezone.utc)
|
||||
)
|
||||
|
|
@ -1187,7 +1188,7 @@ def generic_cost_per_token(
|
|||
usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens, 0
|
||||
)
|
||||
|
||||
billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
|
||||
billing_time: Final = current_time if current_time is not None else current_billing_time()
|
||||
(
|
||||
prompt_base_cost,
|
||||
completion_base_cost,
|
||||
|
|
@ -1379,7 +1380,7 @@ def _cost_map_billed_rates(
|
|||
vertex_location: str | None,
|
||||
current_time: datetime | None,
|
||||
) -> BilledTokenRates:
|
||||
billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
|
||||
billing_time: Final = current_time if current_time is not None else current_billing_time()
|
||||
(
|
||||
prompt_base_cost,
|
||||
completion_base_cost,
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from fastapi import APIRouter, Depends, HTTPException
|
|||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import current_billing_time, pinned_billing_time
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import get_billed_token_rates
|
||||
|
|
@ -631,33 +632,39 @@ async def estimate_cost(
|
|||
function_id="cost-estimate",
|
||||
)
|
||||
|
||||
# Use completion_cost which handles all the logic including margins/discounts
|
||||
try:
|
||||
cost_per_request: Final = completion_cost(
|
||||
completion_response=mock_response,
|
||||
# The totals, the per-token-type lines and the reported rates each resolve pricing on their
|
||||
# own path. Pinning one moment keeps an off-peak window that opens mid-quote from splitting them.
|
||||
billed_at: Final = current_billing_time()
|
||||
with pinned_billing_time(billed_at):
|
||||
# Use completion_cost which handles all the logic including margins/discounts
|
||||
try:
|
||||
cost_per_request: Final = completion_cost(
|
||||
completion_response=mock_response,
|
||||
model=resolved_model,
|
||||
custom_llm_provider=resolved_provider,
|
||||
custom_cost_per_token=resolved.custom_cost_per_token,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={
|
||||
"error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e}"
|
||||
},
|
||||
)
|
||||
|
||||
rates: Final = get_billed_token_rates(
|
||||
model=resolved_model,
|
||||
custom_llm_provider=resolved_provider,
|
||||
usage=usage,
|
||||
custom_cost_per_token=resolved.custom_cost_per_token,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={
|
||||
"error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e}"
|
||||
},
|
||||
current_time=billed_at,
|
||||
)
|
||||
|
||||
per_request: Final = _cost_lines(cost_per_request, litellm_logging_obj.cost_breakdown)
|
||||
daily: Final = per_request.times(request.num_requests_per_day)
|
||||
monthly: Final = per_request.times(request.num_requests_per_month)
|
||||
|
||||
rates: Final = get_billed_token_rates(
|
||||
model=resolved_model,
|
||||
custom_llm_provider=resolved_provider,
|
||||
usage=usage,
|
||||
custom_cost_per_token=resolved.custom_cost_per_token,
|
||||
)
|
||||
model_info: Final = _lookup_model_info(resolved_model)
|
||||
mapped_provider: Final = model_info.get("litellm_provider") if model_info is not None else None
|
||||
custom_llm_provider: Final = mapped_provider if mapped_provider is not None else resolved_provider
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
import json
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import pinned_billing_time
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
|
||||
StandardBuiltInToolCostTracking,
|
||||
)
|
||||
|
|
@ -3981,6 +3983,46 @@ def test_billed_token_rates_follow_the_token_tier_the_breakdown_bills_at(monkeyp
|
|||
assert breakdown.reasoning_cost == pytest.approx(200 * rates.output_cost_per_reasoning_token)
|
||||
|
||||
|
||||
def test_a_pinned_billing_time_prices_the_totals_and_the_reported_rates_at_one_moment(monkeypatch):
|
||||
"""Totals and reported rates resolve off-peak pricing on separate paths that each read the
|
||||
clock, so a window opening between the two reads used to leave them describing one request
|
||||
at two different prices. Pinned, both must answer for the pinned moment."""
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"off-peak-model",
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"off_peak_pricing": {
|
||||
"hours_utc": "02:00-03:00",
|
||||
"input_cost_per_token": 1e-6,
|
||||
"output_cost_per_token": 5e-6,
|
||||
},
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
|
||||
|
||||
with pinned_billing_time(datetime(2026, 1, 1, 2, 30, tzinfo=timezone.utc)):
|
||||
off_peak_prompt_cost, off_peak_completion_cost = generic_cost_per_token(
|
||||
model="off-peak-model", usage=usage, custom_llm_provider="openai"
|
||||
)
|
||||
off_peak_rates = get_billed_token_rates(model="off-peak-model", custom_llm_provider="openai", usage=usage)
|
||||
with pinned_billing_time(datetime(2026, 1, 1, 12, 30, tzinfo=timezone.utc)):
|
||||
peak_prompt_cost, peak_completion_cost = generic_cost_per_token(
|
||||
model="off-peak-model", usage=usage, custom_llm_provider="openai"
|
||||
)
|
||||
peak_rates = get_billed_token_rates(model="off-peak-model", custom_llm_provider="openai", usage=usage)
|
||||
|
||||
assert off_peak_rates.input_cost_per_token == pytest.approx(1e-6)
|
||||
assert peak_rates.input_cost_per_token == pytest.approx(3e-6)
|
||||
assert off_peak_prompt_cost == pytest.approx(1000 * off_peak_rates.input_cost_per_token)
|
||||
assert off_peak_completion_cost == pytest.approx(500 * off_peak_rates.output_cost_per_token)
|
||||
assert peak_prompt_cost == pytest.approx(1000 * peak_rates.input_cost_per_token)
|
||||
assert peak_completion_cost == pytest.approx(500 * peak_rates.output_cost_per_token)
|
||||
|
||||
|
||||
def test_billed_token_rates_are_none_for_an_unpriced_model():
|
||||
usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ Tests for cost tracking settings management endpoints.
|
|||
Tests the GET and PATCH endpoints for managing cost discount configuration.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -12,6 +13,7 @@ from pydantic import ValidationError
|
|||
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import pinned_billing_time
|
||||
from litellm.proxy._types import CostEstimateRequest
|
||||
from litellm.proxy.management_endpoints.cost_tracking_settings import router
|
||||
from litellm.proxy.proxy_server import app
|
||||
|
|
@ -1108,6 +1110,34 @@ class TestEstimateCostCacheAndReasoningTokens:
|
|||
)
|
||||
assert response.output_cost_per_request == pytest.approx(1_000 * response.output_cost_per_token)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_quote_prices_its_totals_and_its_rates_at_the_same_moment(self, monkeypatch):
|
||||
"""The totals and the reported rates resolve off-peak pricing on separate paths. A quote
|
||||
taken as a window opens must not bill on one side of it and report rates from the other."""
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
A_MAPPED_MODEL,
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"off_peak_pricing": {
|
||||
"hours_utc": "02:00-03:00",
|
||||
"input_cost_per_token": 1e-6,
|
||||
"output_cost_per_token": 5e-6,
|
||||
},
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
|
||||
with pinned_billing_time(datetime(2026, 1, 1, 2, 30, tzinfo=timezone.utc)):
|
||||
response = await _estimate(None, model=A_MAPPED_MODEL)
|
||||
|
||||
assert response.input_cost_per_token == pytest.approx(1e-6)
|
||||
assert response.output_cost_per_token == pytest.approx(5e-6)
|
||||
assert response.input_cost_per_request == pytest.approx(INPUT_TOKENS * response.input_cost_per_token)
|
||||
assert response.output_cost_per_request == pytest.approx(OUTPUT_TOKENS * response.output_cost_per_token)
|
||||
|
||||
|
||||
class TestCostEstimateRequestTokenSubsets:
|
||||
def test_cache_tokens_beyond_the_input_tokens_are_rejected(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue