mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
351a54e849
6 changed files with 272 additions and 9 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue