mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
feat(spend): warn when spend-attribution metadata diverges from the resolved key
This commit is contained in:
parent
6437b812be
commit
9933c5031d
3 changed files with 177 additions and 0 deletions
68
litellm/proxy/auth/resolvers/divergence.py
Normal file
68
litellm/proxy/auth/resolvers/divergence.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
class SpendIdentity(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
user_id: str | None = None
|
||||
team_id: str | None = None
|
||||
org_id: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FieldDivergence:
|
||||
field: str
|
||||
resolved_value: str | None
|
||||
consumed_value: str | None
|
||||
|
||||
|
||||
def resolved_identity_from_key(key: UserAPIKeyAuth) -> SpendIdentity:
|
||||
return SpendIdentity(
|
||||
user_id=key.user_id,
|
||||
team_id=key.team_id,
|
||||
org_id=key.org_id,
|
||||
)
|
||||
|
||||
|
||||
def spend_identity_divergence(
|
||||
resolved: SpendIdentity, consumed: SpendIdentity
|
||||
) -> tuple[FieldDivergence, ...]:
|
||||
pairs = (
|
||||
("user_id", resolved.user_id, consumed.user_id),
|
||||
("team_id", resolved.team_id, consumed.team_id),
|
||||
("org_id", resolved.org_id, consumed.org_id),
|
||||
)
|
||||
return tuple(
|
||||
FieldDivergence(
|
||||
field=field, resolved_value=resolved_value, consumed_value=consumed_value
|
||||
)
|
||||
for field, resolved_value, consumed_value in pairs
|
||||
if resolved_value != consumed_value
|
||||
)
|
||||
|
||||
|
||||
def log_identity_divergence(
|
||||
resolved_key: UserAPIKeyAuth, consumed: SpendIdentity
|
||||
) -> None:
|
||||
divergences = spend_identity_divergence(
|
||||
resolved_identity_from_key(resolved_key), consumed
|
||||
)
|
||||
if not divergences:
|
||||
return
|
||||
fields = ", ".join(
|
||||
f"{d.field} (resolved={d.resolved_value!r}, consumed={d.consumed_value!r})"
|
||||
for d in divergences
|
||||
)
|
||||
verbose_proxy_logger.warning(
|
||||
"Spend attribution metadata diverges from the resolved key identity "
|
||||
"[credential_ref=%s]: %s",
|
||||
resolved_key.token,
|
||||
fields,
|
||||
)
|
||||
|
|
@ -17,6 +17,10 @@ from litellm.proxy.auth.auth_checks import (
|
|||
get_team_object,
|
||||
log_db_metrics,
|
||||
)
|
||||
from litellm.proxy.auth.resolvers.divergence import (
|
||||
SpendIdentity,
|
||||
log_identity_divergence,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.spend_tracking.spend_log_error_logger import (
|
||||
|
|
@ -377,6 +381,20 @@ class _ProxyDBLogger(CustomLogger):
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
try:
|
||||
log_identity_divergence(
|
||||
key_obj,
|
||||
SpendIdentity(
|
||||
user_id=metadata.get("user_api_key_user_id"),
|
||||
team_id=metadata.get("user_api_key_team_id"),
|
||||
org_id=metadata.get("user_api_key_org_id"),
|
||||
),
|
||||
)
|
||||
except Exception as divergence_error:
|
||||
verbose_proxy_logger.debug(
|
||||
"Spend identity divergence guard failed (non-fatal): %s",
|
||||
divergence_error,
|
||||
)
|
||||
if metadata.get("user_api_key_alias") is None:
|
||||
metadata["user_api_key_alias"] = key_obj.key_alias
|
||||
if metadata.get("user_api_key_user_id") is None:
|
||||
|
|
|
|||
91
tests/test_litellm/proxy/auth/test_resolvers_divergence.py
Normal file
91
tests/test_litellm/proxy/auth/test_resolvers_divergence.py
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.resolvers.divergence import (
|
||||
SpendIdentity,
|
||||
log_identity_divergence,
|
||||
spend_identity_divergence,
|
||||
)
|
||||
|
||||
|
||||
def _key(**overrides: str | None) -> UserAPIKeyAuth:
|
||||
base: dict[str, str | None] = {
|
||||
"token": "hashed-token",
|
||||
"user_id": "u-1",
|
||||
"team_id": "t-1",
|
||||
"org_id": "o-1",
|
||||
}
|
||||
base.update(overrides)
|
||||
return UserAPIKeyAuth(**base)
|
||||
|
||||
|
||||
def _consumed(**overrides: str | None) -> SpendIdentity:
|
||||
base: dict[str, str | None] = {
|
||||
"user_id": "u-1",
|
||||
"team_id": "t-1",
|
||||
"org_id": "o-1",
|
||||
}
|
||||
base.update(overrides)
|
||||
return SpendIdentity(**base)
|
||||
|
||||
|
||||
def test_consumed_matches_resolved_is_silent(caplog):
|
||||
key = _key()
|
||||
consumed = _consumed()
|
||||
|
||||
assert (
|
||||
spend_identity_divergence(
|
||||
SpendIdentity(user_id="u-1", team_id="t-1", org_id="o-1"),
|
||||
consumed,
|
||||
)
|
||||
== ()
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
log_identity_divergence(key, consumed)
|
||||
assert caplog.records == []
|
||||
|
||||
|
||||
def test_empty_consumed_metadata_against_populated_key_warns_naming_fields(caplog):
|
||||
key = _key()
|
||||
consumed = _consumed(user_id=None, team_id=None, org_id=None)
|
||||
|
||||
diverged = spend_identity_divergence(
|
||||
SpendIdentity(user_id="u-1", team_id="t-1", org_id="o-1"),
|
||||
consumed,
|
||||
)
|
||||
by_field = {d.field: (d.resolved_value, d.consumed_value) for d in diverged}
|
||||
|
||||
assert set(by_field) == {"user_id", "team_id", "org_id"}
|
||||
assert by_field["user_id"] == ("u-1", None)
|
||||
assert by_field["team_id"] == ("t-1", None)
|
||||
assert by_field["org_id"] == ("o-1", None)
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
log_identity_divergence(key, consumed)
|
||||
|
||||
assert len(caplog.records) == 1
|
||||
message = caplog.records[0].getMessage()
|
||||
assert "user_id" in message and "team_id" in message and "org_id" in message
|
||||
|
||||
|
||||
def test_different_consumed_user_id_warns(caplog):
|
||||
key = _key()
|
||||
consumed = _consumed(user_id="u-OTHER")
|
||||
|
||||
diverged = spend_identity_divergence(
|
||||
SpendIdentity(user_id="u-1", team_id="t-1", org_id="o-1"),
|
||||
consumed,
|
||||
)
|
||||
by_field = {d.field: (d.resolved_value, d.consumed_value) for d in diverged}
|
||||
|
||||
assert by_field == {"user_id": ("u-1", "u-OTHER")}
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
log_identity_divergence(key, consumed)
|
||||
|
||||
assert len(caplog.records) == 1
|
||||
message = caplog.records[0].getMessage()
|
||||
assert "user_id" in message and "u-1" in message and "u-OTHER" in message
|
||||
Loading…
Add table
Reference in a new issue