Merge pull request #41507 from BerriAI/litellm_attribute_router_rejected_spend_provider

fix(spend_tracking): attribute router-rejected requests to the model group provider
This commit is contained in:
Yassin Kortam 2026-09-16 18:21:44 -07:00 • committed by GitHub
commit 351a54e849
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 272 additions and 9 deletions

View file

@ -160,7 +160,7 @@ def get_llm_provider(
if model is None:
raise ValueError("model parameter is required but was None. Please provide a valid model name.")
if litellm.LiteLLMProxyChatConfig._should_use_litellm_proxy_by_default(
if litellm.LiteLLMProxyChatConfig.should_use_litellm_proxy_by_default(
litellm_params=cast(LiteLLM_Params | None, litellm_params)
):
return litellm.LiteLLMProxyChatConfig.litellm_proxy_get_custom_llm_provider_info(

View file

@ -54,7 +54,7 @@ class LiteLLMProxyChatConfig(OpenAIGPTConfig):
return api_key or get_secret_str("LITELLM_PROXY_API_KEY")
@staticmethod
def _should_use_litellm_proxy_by_default(
def should_use_litellm_proxy_by_default(
litellm_params: LiteLLM_Params | None = None,
):
"""

View file

@ -257,8 +257,8 @@ class DBSpendUpdateWriter:
# Completion object fields
kwargs: dict | None,
completion_response: object,
start_time: datetime | None,
end_time: datetime | None,
start_time: datetime,
end_time: datetime,
response_cost: float | None,
) -> bool:
"""Record the request's spend, answering whether its cost still needs charging.
@ -299,6 +299,7 @@ class DBSpendUpdateWriter:
response_obj=completion_response,
start_time=start_time,
end_time=end_time,
llm_router=get_llm_router(),
)
payload["spend"] = response_cost or 0.0
if isinstance(payload["startTime"], datetime):

View file

@ -5,7 +5,7 @@ from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from datetime import datetime as dt
from types import MappingProxyType
from typing import Final, Literal, Protocol, cast, runtime_checkable
from typing import TYPE_CHECKING, Final, Literal, Protocol, cast, runtime_checkable
from pydantic import BaseModel
@ -32,6 +32,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
reconstruct_model_name,
)
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call
from litellm.litellm_core_utils.litellm_logging import (
coerce_model_access_groups,
@ -43,10 +44,12 @@ from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsR
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.proxy.utils import PrismaClient, hash_token
from litellm.types.router import DeploymentTypedDict, LiteLLM_Params
from litellm.types.utils import (
PROMPT_CARRYING_GUARDRAIL_FIELDS,
CallTypes,
CostBreakdown,
LlmProviders,
StandardLoggingGuardrailInformation,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
@ -57,6 +60,9 @@ from litellm.types.utils import (
)
from litellm.utils import get_end_user_id_for_cost_tracking
if TYPE_CHECKING:
from litellm.router import Router
def _get_max_string_length_prompt_in_db() -> int:
"""
@ -339,12 +345,45 @@ def _sl_attribution_fallback(
return standard_logging_payload.get(field) or ""
def _deployment_provider(deployment: DeploymentTypedDict) -> str | None:
litellm_params: Final = LiteLLM_Params.model_validate(deployment["litellm_params"])
if litellm.LiteLLMProxyChatConfig.should_use_litellm_proxy_by_default(litellm_params=litellm_params):
return LlmProviders.LITELLM_PROXY.value
declared: Final = declared_authenticating_provider(litellm_params.model, litellm_params.custom_llm_provider)
if declared is not None:
return declared
try:
_, provider, _, _ = litellm.get_llm_provider(
model=litellm_params.model, custom_llm_provider=litellm_params.custom_llm_provider
)
except litellm.exceptions.BadRequestError:
return None
return provider or None
def _model_group_provider(model_group: str, llm_router: "Router | None") -> str | None:
if llm_router is None or not model_group:
return None
providers: Final = frozenset(
provider
for deployment in llm_router.get_model_list(model_name=model_group) or ()
if (provider := _deployment_provider(deployment)) is not None
)
return next(iter(providers)) if len(providers) == 1 else None
def _looks_like_model_name(model: str) -> bool:
candidate: Final = model.removeprefix(MCP_SPEND_LOG_MODEL_PREFIX)
return len(candidate) <= MAX_SPEND_LOG_MODEL_NAME_LENGTH and not any(char.isspace() for char in candidate)
def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogsPayload:
def get_logging_payload(
kwargs: dict | None,
response_obj: object,
start_time: datetime,
end_time: datetime,
llm_router: "Router | None" = None,
) -> SpendLogsPayload:
if kwargs is None:
kwargs = {}
@ -440,15 +479,16 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
hidden_params: Final = standard_logging_payload.get("hidden_params", {})
litellm_overhead_time_ms = hidden_params.get("litellm_overhead_time_ms")
custom_llm_provider: Final = (
logged_provider: Final = (
kwargs.get("custom_llm_provider")
or _sl_attribution_fallback(standard_logging_payload, "custom_llm_provider")
or None
)
custom_llm_provider: Final = logged_provider or _model_group_provider(_model_group, llm_router)
raw_model: Final = cast(str, kwargs.get("model") or "")
resolved_model: Final = (
standard_logging_payload.get("model") if standard_logging_payload is not None else None
) or reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})
) or reconstruct_model_name(raw_model, logged_provider, metadata or {})
failed_with_prompt_shaped_model: Final = (
_get_status_for_spend_log(metadata=metadata) == "failure"
and not _model_group

View file

@ -76,6 +76,49 @@ async def test_daily_spend_tracking_with_disabled_spend_logs():
assert call_args["payload"]["custom_llm_provider"] == "openai"
@pytest.mark.asyncio
async def test_update_database_attributes_router_rejected_failure_to_model_group_provider():
db_writer = DBSpendUpdateWriter()
db_writer._insert_spend_log_to_db = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_user_transaction = AsyncMock()
llm_router: Final = litellm.Router(
model_list=[
{"model_name": "openai-outage", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-a"}},
{"model_name": "openai-outage", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-b"}},
]
)
with (
patch("litellm.proxy.proxy_server.disable_spend_logs", True), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.llm_router", llm_router), # test-quality-ok: get_llm_router reads this proxy_server module global at call time; no injection seam
):
await db_writer.update_database(
token="test-token",
user_id="test-user",
end_user_id=None,
team_id=None,
org_id=None,
kwargs={
"model": "openai-outage",
"litellm_params": {
"metadata": {"user_api_key": "test-token", "model_group": "openai-outage", "status": "failure"}
},
},
completion_response={},
start_time=datetime.now(timezone.utc),
end_time=datetime.now(timezone.utc),
response_cost=0.0,
)
await asyncio.sleep(0)
payload: Final = db_writer.add_spend_log_transaction_to_daily_user_transaction.call_args[1]["payload"]
assert payload["model_group"] == "openai-outage"
assert payload["custom_llm_provider"] == "openai"
def _tool_call_response(*names: str) -> object:
from types import SimpleNamespace

View file

@ -1,7 +1,7 @@
import asyncio
import datetime
import json
from collections.abc import Mapping
from collections.abc import Callable, Mapping
from datetime import timezone
from typing import Any, Final, cast
from unittest.mock import AsyncMock, MagicMock, patch
@ -44,6 +44,7 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import (
should_store_prompts_and_responses_in_spend_logs,
)
from litellm.proxy.utils import hash_token
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
StandardLoggingHiddenParams,
StandardLoggingMetadata,
@ -4003,6 +4004,184 @@ def test_get_logging_payload_failed_request_without_standard_logging_payload_lea
assert payload["custom_llm_provider"] == ""
def _router_rejected_failure_payload(model_group: str, llm_router: litellm.Router | None) -> SpendLogsPayload:
return get_logging_payload(
kwargs={
"model": model_group,
"litellm_params": {
"metadata": {"user_api_key": "test-key", "model_group": model_group, "status": "failure"}
},
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
llm_router=llm_router,
)
_ProviderResolution = tuple[str, str, str | None, str | None]
def _router_init_provider_stub(
model: str,
custom_llm_provider: str | None = None,
api_base: str | None = None,
api_key: str | None = None,
litellm_params: GenericLiteLLMParams | None = None,
) -> _ProviderResolution:
prefix, _, suffix = model.partition("/")
return (suffix or model, custom_llm_provider or (prefix if suffix else "openai"), api_base, api_key)
def _oauth_tripwire(resolution_attempts: list[str]) -> Callable[..., _ProviderResolution]:
def _trip(
model: str,
custom_llm_provider: str | None = None,
api_base: str | None = None,
api_key: str | None = None,
litellm_params: GenericLiteLLMParams | None = None,
) -> _ProviderResolution:
resolution_attempts.append(model)
raise AssertionError("get_llm_provider would run the OAuth device flow")
return _trip
def _openai_and_anthropic_router() -> litellm.Router:
return litellm.Router(
model_list=[
{"model_name": "openai-group", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-a"}},
{"model_name": "openai-group", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-b"}},
{"model_name": "mixed-group", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-a"}},
{
"model_name": "mixed-group",
"litellm_params": {"model": "anthropic/claude-haiku-4-5", "api_key": "sk-c"},
},
]
)
@pytest.mark.parametrize(
"model_group,expected_provider",
[("openai-group", "openai"), ("mixed-group", ""), ("not-in-router", "")],
)
def test_get_logging_payload_router_rejected_request_takes_provider_from_model_group(
model_group: str, expected_provider: str
):
payload = _router_rejected_failure_payload(model_group, _openai_and_anthropic_router())
assert payload["model_group"] == model_group
assert payload["custom_llm_provider"] == expected_provider
def test_get_logging_payload_router_rejected_request_without_router_leaves_provider_empty():
assert _router_rejected_failure_payload("openai-group", None)["custom_llm_provider"] == ""
@pytest.mark.parametrize(
"litellm_params,expected_provider",
[
({"model": "github_copilot/gpt-4o"}, "github_copilot"),
({"model": "gpt-5", "custom_llm_provider": "chatgpt"}, "chatgpt"),
],
)
def test_get_logging_payload_inferred_provider_never_resolves_declared_authenticating_providers(
monkeypatch: pytest.MonkeyPatch, litellm_params: dict[str, str], expected_provider: str
):
resolution_attempts: list[str] = []
monkeypatch.setattr(litellm, "get_llm_provider", _router_init_provider_stub)
llm_router = litellm.Router(model_list=[{"model_name": "oauth-group", "litellm_params": litellm_params}])
monkeypatch.setattr(litellm, "get_llm_provider", _oauth_tripwire(resolution_attempts))
payload = _router_rejected_failure_payload("oauth-group", llm_router)
assert payload["custom_llm_provider"] == expected_provider
assert resolution_attempts == []
@pytest.mark.parametrize(
"litellm_params",
[
{"model": "github_copilot/gpt-4o"},
{"model": "gpt-5", "custom_llm_provider": "chatgpt"},
{"model": "openai/gpt-4o-mini", "api_key": "sk-a"},
],
)
def test_get_logging_payload_inferred_provider_honours_global_litellm_proxy_override(
monkeypatch: pytest.MonkeyPatch, litellm_params: dict[str, str]
):
monkeypatch.setattr(litellm, "get_llm_provider", _router_init_provider_stub)
llm_router = litellm.Router(model_list=[{"model_name": "proxied-group", "litellm_params": litellm_params}])
monkeypatch.setattr(litellm, "get_llm_provider", _oauth_tripwire([]))
monkeypatch.setattr(litellm, "use_litellm_proxy", True)
payload = _router_rejected_failure_payload("proxied-group", llm_router)
assert payload["custom_llm_provider"] == "litellm_proxy"
def test_get_logging_payload_router_rejected_request_for_unresolvable_deployment_leaves_provider_empty(
monkeypatch: pytest.MonkeyPatch,
):
with monkeypatch.context() as router_init:
router_init.setattr(litellm, "get_llm_provider", _router_init_provider_stub)
llm_router = litellm.Router(
model_list=[{"model_name": "opaque-group", "litellm_params": {"model": "my-unprefixed-model"}}]
)
payload = _router_rejected_failure_payload("opaque-group", llm_router)
assert payload["model_group"] == "opaque-group"
assert payload["custom_llm_provider"] == ""
def test_get_logging_payload_inferred_provider_does_not_rewrite_spend_log_model():
llm_router = litellm.Router(
model_list=[
{
"model_name": "bedrock-group",
"litellm_params": {
"model": "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
"aws_region_name": "us-east-1",
},
},
{
"model_name": "bedrock-group",
"litellm_params": {
"model": "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
"aws_region_name": "us-west-2",
},
},
]
)
payload = _router_rejected_failure_payload("bedrock-group", llm_router)
assert payload["custom_llm_provider"] == "bedrock"
assert payload["model"] == "bedrock-group"
def test_get_logging_payload_logged_provider_wins_over_model_group_provider():
payload = get_logging_payload(
kwargs={
"model": "openai-group",
"litellm_params": {"metadata": {"user_api_key": "test-key", "model_group": "openai-group"}},
"standard_logging_object": {
**_make_failed_request_standard_logging_payload(),
"model_group": "openai-group",
"custom_llm_provider": "azure",
},
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
llm_router=_openai_and_anthropic_router(),
)
assert payload["custom_llm_provider"] == "azure"
class _ModelRouterSpendLogKwargs(TypedDict):
model: ReadOnly[str]
litellm_params: ReadOnly[dict[str, dict[str, str]]]