diff --git a/litellm/proxy/auth/resolvers/divergence.py b/litellm/proxy/auth/resolvers/divergence.py new file mode 100644 index 00000000000..f0c69098607 --- /dev/null +++ b/litellm/proxy/auth/resolvers/divergence.py @@ -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, + ) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 8fc9d009e67..4dd66e06607 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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: diff --git a/tests/test_litellm/proxy/auth/test_resolvers_divergence.py b/tests/test_litellm/proxy/auth/test_resolvers_divergence.py new file mode 100644 index 00000000000..635597f41e3 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_resolvers_divergence.py @@ -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