feat(spend): persist effective service_tier in spend log metadata

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-10 18:55:36 +00:00
parent 54d404ef2c
commit 7a7639993c
6 changed files with 164 additions and 38 deletions

View file

@ -29,6 +29,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import (
_parse_prompt_tokens_details,
calculate_cost_component,
generic_cost_per_token,
get_effective_service_tier,
get_token_type_cost_breakdown,
get_billable_input_tokens,
select_cost_metric_for_model,
@ -98,7 +99,6 @@ from litellm.types.utils import (
LlmProviders,
LlmProvidersSet,
ModelInfo,
ServiceTier,
StandardBuiltInToolsParams,
TranscriptionUsageDurationObject,
TranscriptionUsageTokensObject,
@ -852,20 +852,6 @@ def _map_traffic_type_to_service_tier(traffic_type: Optional[str]) -> Optional[s
return service_tier
def _normalize_service_tier(service_tier: object) -> str | None:
"""
Reduce a service_tier value to a concrete billable tier string or None.
"auto" is a routing preference and any non-string value is not a billable
tier, so both defer to standard pricing (or to the tier the provider reports
on the response usage) instead of crashing the downstream cost-key lookup,
which calls service_tier.lower()
"""
if not isinstance(service_tier, str) or service_tier.lower() == ServiceTier.AUTO.value:
return None
return service_tier
def _get_usage_object(
completion_response: Any,
) -> Optional[Usage]:
@ -1181,29 +1167,12 @@ def completion_cost(
cost_per_token_usage_object: Optional[Usage] = _get_usage_object(completion_response=completion_response)
rerank_billed_units: Optional[RerankBilledUnits] = None
# Extract service_tier from optional_params if not provided directly
if service_tier is None and optional_params is not None:
service_tier = optional_params.get("service_tier")
service_tier = _normalize_service_tier(service_tier)
# Extract service_tier from completion_response if not provided
if service_tier is None and completion_response is not None:
if isinstance(completion_response, BaseModel):
service_tier = getattr(completion_response, "service_tier", None)
elif isinstance(completion_response, dict):
service_tier = completion_response.get("service_tier")
service_tier = _normalize_service_tier(service_tier)
# Extract service_tier from usage object if not provided
if service_tier is None and cost_per_token_usage_object is not None:
if isinstance(cost_per_token_usage_object, BaseModel):
service_tier = getattr(cost_per_token_usage_object, "service_tier", None)
elif isinstance(cost_per_token_usage_object, dict):
service_tier = cost_per_token_usage_object.get("service_tier")
service_tier = _normalize_service_tier(service_tier)
service_tier = get_effective_service_tier(
service_tier=service_tier,
optional_params=optional_params,
completion_response=completion_response,
usage_object=cost_per_token_usage_object,
)
selected_model = _select_model_name_for_cost_calc(
model=model,

View file

@ -4,6 +4,8 @@
from dataclasses import dataclass
from typing import Any, Literal, Optional, Tuple, TypedDict, cast
from pydantic import BaseModel
import litellm
from litellm._logging import verbose_logger
from litellm.types.utils import (
@ -193,6 +195,59 @@ def _get_service_tier_cost_key(base_key: str, service_tier: Optional[str]) -> st
return base_key
def _normalize_service_tier(service_tier: object) -> Optional[str]:
"""
Reduce a service_tier value to a concrete billable tier string or None.
"auto" is a routing preference and any non-string value is not a billable
tier, so both defer to standard pricing (or to the tier the provider reports
on the response usage) instead of crashing the downstream cost-key lookup,
which calls service_tier.lower()
"""
if not isinstance(service_tier, str) or service_tier.lower() == ServiceTier.AUTO.value:
return None
return service_tier
def _extract_service_tier(obj: object) -> object:
if isinstance(obj, BaseModel):
return getattr(obj, "service_tier", None) # pyright: ignore[reportAny] # service_tier is a dynamic provider-set field, not declared on the model
if isinstance(obj, dict):
return obj.get("service_tier")
return None
def get_effective_service_tier(
service_tier: object = None,
optional_params: Optional[dict] = None,
completion_response: object = None,
usage_object: object = None,
) -> Optional[str]:
"""
Resolve the effective (billed) service tier for a request.
Precedence: an explicitly passed tier, then the request `optional_params`,
then the tier the provider reports on the response, then the tier on the
usage object. At each step "auto" and any non-string value normalize to
None so a routing preference like "auto" falls through to the concrete tier
the provider actually used
"""
normalized = _normalize_service_tier(service_tier)
if normalized is not None:
return normalized
if optional_params is not None:
normalized = _normalize_service_tier(optional_params.get("service_tier"))
if normalized is not None:
return normalized
normalized = _normalize_service_tier(_extract_service_tier(completion_response))
if normalized is not None:
return normalized
return _normalize_service_tier(_extract_service_tier(usage_object))
def _parse_above_token_threshold(key: str) -> float:
threshold_str = key.split("_above_")[1].split("_tokens")[0]
return float(threshold_str.replace("k", "")) * (1000 if "k" in threshold_str else 1)

View file

@ -3142,6 +3142,7 @@ class SpendLogsMetadata(TypedDict):
attempted_retries: Optional[int] # Number of retries attempted (0 = first attempt succeeded)
max_retries: Optional[int] # Max retries configured for this request
cost_breakdown: Optional[CostBreakdown] # Detailed cost breakdown (input_cost, output_cost, margin, discount, etc.)
service_tier: Optional[str] # effective (billed) service tier resolved from the request/response
class SpendLogsPayload(TypedDict):

View file

@ -24,6 +24,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
reconstruct_model_name,
)
from litellm.litellm_core_utils.llm_cost_calc.utils import get_effective_service_tier
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
@ -85,6 +86,7 @@ def _get_spend_logs_metadata(
litellm_overhead_time_ms: Optional[float] = None,
cost_breakdown: Optional[CostBreakdown] = None,
litellm_call_id: Optional[str] = None,
service_tier: Optional[str] = None,
) -> SpendLogsMetadata:
if metadata is None:
return SpendLogsMetadata(
@ -116,6 +118,7 @@ def _get_spend_logs_metadata(
max_retries=None,
cost_breakdown=None,
litellm_call_id=litellm_call_id,
service_tier=service_tier,
)
verbose_proxy_logger.debug(
"getting payload for SpendLogs, available keys in metadata: " + str(list(metadata.keys()))
@ -143,6 +146,7 @@ def _get_spend_logs_metadata(
clean_metadata["litellm_overhead_time_ms"] = litellm_overhead_time_ms
clean_metadata["cost_breakdown"] = cost_breakdown
clean_metadata["litellm_call_id"] = litellm_call_id
clean_metadata["service_tier"] = service_tier
return clean_metadata
@ -364,6 +368,11 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
Optional[str],
kwargs.get("litellm_call_id") or litellm_params.get("litellm_call_id"),
),
service_tier=get_effective_service_tier(
optional_params=kwargs.get("optional_params"),
completion_response=response_obj,
usage_object=usage,
),
)
special_usage_fields = ["completion_tokens", "prompt_tokens", "total_tokens"]

View file

@ -2137,3 +2137,33 @@ def test_token_type_cost_breakdown_applies_regional_uplift():
text_input_cost = 600 * model_info["input_cost_per_token"] * uplift
assert text_output_cost + eu.reasoning_cost == pytest.approx(completion_cost)
assert text_input_cost + eu.cache_read_cost == pytest.approx(prompt_cost)
from litellm.litellm_core_utils.llm_cost_calc.utils import get_effective_service_tier
@pytest.mark.parametrize(
"service_tier, optional_params, completion_response, usage_object, expected",
[
("priority", None, None, None, "priority"),
(None, {"service_tier": "flex"}, None, None, "flex"),
(None, {"service_tier": "auto"}, {"service_tier": "priority"}, None, "priority"),
(None, None, {"service_tier": "flex"}, None, "flex"),
(None, None, None, {"service_tier": "priority"}, "priority"),
("auto", {"service_tier": "auto"}, {"service_tier": "auto"}, None, None),
(None, {"service_tier": 123}, None, None, None),
(None, None, None, None, None),
],
)
def test_get_effective_service_tier_precedence(
service_tier, optional_params, completion_response, usage_object, expected
):
assert (
get_effective_service_tier(
service_tier=service_tier,
optional_params=optional_params,
completion_response=completion_response,
usage_object=usage_object,
)
== expected
)

View file

@ -2571,3 +2571,65 @@ def test_get_logging_payload_hashes_bearer_prefixed_api_key():
assert not metadata_dict["user_api_key"].startswith("sk-"), (
f"metadata user_api_key contains unhashed key: {metadata_dict['user_api_key']}"
)
def _service_tier_kwargs(request_service_tier):
litellm_params = {"metadata": {"user_api_key": "sk-test-key"}}
optional_params = {} if request_service_tier is None else {"service_tier": request_service_tier}
return {
"model": "gpt-4.1",
"custom_llm_provider": "openai",
"call_type": "acompletion",
"optional_params": optional_params,
"litellm_params": litellm_params,
}
def _run_service_tier_payload(request_service_tier, response_service_tier):
kwargs = _service_tier_kwargs(request_service_tier)
response_obj = {
"id": "test-response-123",
"choices": [{"message": {"content": "Hello!"}}],
"usage": {"total_tokens": 100, "prompt_tokens": 50, "completion_tokens": 50},
}
if response_service_tier is not None:
response_obj["service_tier"] = response_service_tier
payload = get_logging_payload(
kwargs=kwargs,
response_obj=response_obj,
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
return json.loads(payload["metadata"])
@patch("litellm.proxy.proxy_server.master_key", None)
@patch("litellm.proxy.proxy_server.general_settings", {})
def test_get_logging_payload_persists_service_tier_from_request():
"""Effective service_tier from the request lands in spend log metadata."""
metadata = _run_service_tier_payload(request_service_tier="priority", response_service_tier=None)
assert metadata["service_tier"] == "priority"
@patch("litellm.proxy.proxy_server.master_key", None)
@patch("litellm.proxy.proxy_server.general_settings", {})
def test_get_logging_payload_persists_effective_service_tier_when_request_is_auto():
"""A request 'auto' must resolve to the concrete tier the provider actually used."""
metadata = _run_service_tier_payload(request_service_tier="auto", response_service_tier="flex")
assert metadata["service_tier"] == "flex"
@patch("litellm.proxy.proxy_server.master_key", None)
@patch("litellm.proxy.proxy_server.general_settings", {})
def test_get_logging_payload_persists_service_tier_from_response_only():
"""When the request omits service_tier, the response-reported tier is persisted."""
metadata = _run_service_tier_payload(request_service_tier=None, response_service_tier="priority")
assert metadata["service_tier"] == "priority"
@patch("litellm.proxy.proxy_server.master_key", None)
@patch("litellm.proxy.proxy_server.general_settings", {})
def test_get_logging_payload_service_tier_none_when_absent():
"""No service_tier anywhere resolves to None rather than a bogus value."""
metadata = _run_service_tier_payload(request_service_tier=None, response_service_tier=None)
assert metadata["service_tier"] is None