feat(spend): warn when spend-attribution metadata diverges from the resolved key

This commit is contained in:
Yassin Kortam 2026-06-22 12:56:42 -07:00
parent 6437b812be
commit 9933c5031d
3 changed files with 177 additions and 0 deletions

View 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,
)

View file

@ -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:

View 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