mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
54d404ef2c
commit
7a7639993c
6 changed files with 164 additions and 38 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue