This commit is contained in:
Milan Zolota 2026-08-26 23:39:20 -07:00 • committed by GitHub
commit b985eb824f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1591 additions and 214 deletions

View file

@ -2381,6 +2381,28 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
)
use_google_kms: bool | None = Field(None, description="decrypt keys with google kms")
use_azure_key_vault: bool | None = Field(None, description="load keys from azure key vault")
track_spend_across_all_user_teams: bool | None = Field(
None,
description=(
"attribute each request's spend to EVERY team the calling user belongs to "
"(and every organization reached through those teams), not only the team "
"stamped on the virtual key. Budget gates expand to match, so one "
"over-budget team blocks the caller everywhere. Summing team spend then "
"exceeds real spend by design, because one request is charged to several "
"teams; key, user, and org totals stay single-counted. Default off."
),
)
enforce_rate_limits_across_all_user_teams: bool | None = Field(
None,
description=(
"apply the RPM/TPM limits of EVERY team the calling user belongs to, not "
"only the team stamped on the virtual key. The caller's effective limit "
"becomes the minimum across their memberships, so a busy team can throttle "
"someone who is mostly working for a different team. Separate from "
"track_spend_across_all_user_teams so spend attribution can be adopted "
"without this. Default off."
),
)
master_key: str | None = Field(None, description="require a key for all calls to proxy")
coordination_redis: CoordinationRedisParams | None = Field(
None,
@ -2782,6 +2804,20 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken):
organization_rpm_limit: int | None = None
organization_metadata: dict | None = None
# Membership attribution params (see litellm/proxy/auth/membership_attribution.py).
# Populated during auth ONLY when track_spend_across_all_user_teams or
# enforce_rate_limits_across_all_user_teams is on; None otherwise, so the
# default single-team path is untouched. Server-only and stripped from
# validated input for the same reason as mcp_source_team_rpm_limits below:
# a forged entry would let a caller choose which teams they are charged
# against, or name a team with generous limits to escape their real ones.
attributed_team_ids: tuple[str, ...] | None = Field(default=None, exclude=True)
attributed_org_ids: tuple[str, ...] | None = Field(default=None, exclude=True)
# team_id -> {"rpm": int | None, "tpm": int | None}. Precomputed during auth
# because rate-limit descriptor construction is synchronous and cannot await
# a per-team lookup.
attributed_team_limits: Mapping[str, Mapping[str, int | None]] | None = Field(default=None, exclude=True)
# Project Params
project_alias: str | None = None
project_metadata: dict | None = None
@ -2892,6 +2928,9 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
values.pop("mcp_source_team_rpm_limits", None)
values.pop("mcp_session_resource_server_id", None)
values.pop("via_virtual_key", None)
values.pop("attributed_team_ids", None)
values.pop("attributed_org_ids", None)
values.pop("attributed_team_limits", None)
if values.get("api_key") is not None:
values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))})
if isinstance(values.get("api_key"), str):

View file

@ -997,6 +997,23 @@ async def common_checks(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
),
# No-ops unless track_spend_across_all_user_teams is on: with the
# setting off, attributed_team_ids/attributed_org_ids are None and
# both return before touching the cache.
_attributed_teams_max_budget_check(
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
),
_attributed_orgs_max_budget_check(
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
),
_tag_max_budget_check(
request_body=request_body,
prisma_client=prisma_client,
@ -4818,6 +4835,193 @@ async def _team_max_budget_check(
)
async def _attributed_teams_max_budget_check(
valid_token: UserAPIKeyAuth | None,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
general_settings: Mapping[str, object],
) -> None:
"""Enforce the max budget of every OTHER team the caller belongs to.
Only active when ``track_spend_across_all_user_teams`` is on -- otherwise
``attributed_team_ids`` is None and this returns immediately.
The stamped team is skipped because ``_team_max_budget_check`` already
covers it; checking it twice would fire a duplicate budget alert. Teams are
evaluated concurrently (each lookup is a cache hit in steady state) but the
raised error is chosen by membership order, not by which coroutine finished
first, so the same request always reports the same offending team.
Consequence worth stating plainly: with attribution on, ONE exhausted team
blocks the caller on every team, because their usage counts against all of
them.
"""
from litellm.proxy.auth.membership_attribution import (
attributed_team_ids,
spend_attribution_enabled,
)
# Gated on the SPEND setting specifically. The resolver populates
# attributed_team_ids when either setting is on, so an operator who enabled
# only rate-limit attribution must not silently get all-team budget gates.
if not spend_attribution_enabled(general_settings):
return
if valid_token is None or prisma_client is None:
return
other_team_ids: Final = tuple(
team_id for team_id in attributed_team_ids(valid_token) if team_id != valid_token.team_id
)
if not other_team_ids:
return
from litellm.proxy.proxy_server import get_current_spend
async def _check(team_id: str) -> None:
try:
team_object: Final = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=valid_token.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception: # noqa: BLE001 # a team that cannot be loaded contributes no ceiling
return
if team_object.max_budget is None or not math.isfinite(team_object.max_budget):
return
spend: Final = await get_current_spend(
counter_key=f"spend:team:{team_id}",
fallback_spend=team_object.spend or 0.0,
max_budget=team_object.max_budget,
)
if spend <= team_object.max_budget:
return
call_info: Final = CallInfo(
token=valid_token.token,
spend=spend,
max_budget=team_object.max_budget,
user_id=valid_token.user_id,
team_id=team_id,
team_alias=team_object.team_alias,
organization_id=team_object.organization_id,
event_group=Litellm_EntityType.TEAM,
)
asyncio.create_task(
proxy_logging_obj.budget_alerts(
type="team_budget",
user_info=call_info,
)
)
raise litellm.BudgetExceededError(
current_cost=spend,
max_budget=team_object.max_budget,
message=(
f"Budget has been exceeded! Team={team_id} Current cost: {spend}, Max budget: {team_object.max_budget}"
),
entity_type=Litellm_EntityType.TEAM.value,
entity_id=team_id,
)
results: Final = await asyncio.gather(*(_check(team_id) for team_id in other_team_ids), return_exceptions=True)
first_error: Final = next((r for r in results if isinstance(r, BaseException)), None)
if first_error is not None:
raise first_error
async def _attributed_orgs_max_budget_check(
valid_token: UserAPIKeyAuth | None,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
general_settings: Mapping[str, object],
) -> None:
"""Enforce the max budget of every OTHER organization the caller belongs to.
Skips ``valid_token.org_id`` -- ``_organization_max_budget_check`` owns that
one. Only reaches past it when the caller's teams span several
organizations; in a single-org deployment this is always a no-op.
"""
from litellm.proxy.auth.membership_attribution import (
attributed_org_ids,
spend_attribution_enabled,
)
if not spend_attribution_enabled(general_settings):
return
if valid_token is None or prisma_client is None:
return
other_org_ids: Final = tuple(org_id for org_id in attributed_org_ids(valid_token) if org_id != valid_token.org_id)
if not other_org_ids:
return
from litellm.proxy.proxy_server import get_current_spend
async def _check(org_id: str) -> None:
try:
org_table: Final = await get_org_object(
org_id=org_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=valid_token.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
include_budget_table=True,
)
except Exception: # noqa: BLE001 # an org that cannot be loaded contributes no ceiling
return
if org_table is None or org_table.litellm_budget_table is None:
return
org_max_budget: Final = org_table.litellm_budget_table.max_budget
if org_max_budget is None or org_max_budget <= 0 or not math.isfinite(org_max_budget):
return
org_spend: Final = await get_current_spend(
counter_key=f"spend:org:{org_id}",
fallback_spend=org_table.spend or 0.0,
max_budget=org_max_budget,
)
if org_spend <= org_max_budget:
return
asyncio.create_task(
proxy_logging_obj.budget_alerts(
type="team_budget",
user_info=CallInfo(
token=valid_token.token,
spend=org_spend,
max_budget=org_max_budget,
user_id=valid_token.user_id,
team_id=valid_token.team_id,
organization_id=org_id,
event_group=Litellm_EntityType.ORGANIZATION,
),
)
)
raise litellm.BudgetExceededError(
current_cost=org_spend,
max_budget=org_max_budget,
message=(
f"Budget has been exceeded! Organization={org_id} Current cost: {org_spend}, "
f"Max budget: {org_max_budget}"
),
entity_type=Litellm_EntityType.ORGANIZATION.value,
entity_id=org_id,
)
results: Final = await asyncio.gather(*(_check(org_id) for org_id in other_org_ids), return_exceptions=True)
first_error: Final = next((r for r in results if isinstance(r, BaseException)), None)
if first_error is not None:
raise first_error
async def _team_multi_budget_check(
team_object: LiteLLM_TeamTable | None,
):

View file

@ -0,0 +1,273 @@
"""Membership-based usage attribution.
By default LiteLLM attributes a request to the single team stamped on the
virtual key (or the single team a JWT claim resolved to), and to that team's
organization. A user who belongs to many teams therefore contributes spend to
whichever team the key happens to name and nothing to the rest, and an operator
who wants "what did this team consume?" only gets an answer for keys that
happen to name it.
Two opt-in settings change that. Both default to off, so an existing deployment
keeps the single-team behavior byte for byte:
``track_spend_across_all_user_teams``
Spend increments, daily rollups, and budget gates apply to every team the
caller belongs to, and to every organization reached through those teams.
``enforce_rate_limits_across_all_user_teams``
The same expansion for the RPM/TPM limiter, so a request must fit inside
every membership's limit rather than only the stamped team's.
They are separate settings because they carry different costs. Spend
attribution is additive bookkeeping: the only surprise is that summing team
spend now exceeds real spend, because one request is deliberately charged to
several teams. Rate-limit expansion is a live behavior change: a caller's
effective limit becomes the MINIMUM across their memberships, so a busy team
can throttle someone who is mostly working for a different team. Operators
should be able to adopt the first without the second.
Nothing here builds a team tree. LiteLLM teams are a flat set, each optionally
belonging to one organization -- there is no ``parent_team_id`` in the schema.
"All memberships" therefore means the caller's teams plus the organizations
those teams (and the caller's own user row) belong to, never a recursive walk.
"""
import asyncio
from collections.abc import Iterable, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LiteLLM_UserTable, UserAPIKeyAuth
if TYPE_CHECKING:
from opentelemetry.trace import Span
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import PrismaClient, ProxyLogging
SPEND_ATTRIBUTION_SETTING: Final = "track_spend_across_all_user_teams"
RATE_LIMIT_ATTRIBUTION_SETTING: Final = "enforce_rate_limits_across_all_user_teams"
# One resolved team: its id, and the team row if it could be loaded.
TeamResolution: TypeAlias = tuple[str, "LiteLLM_TeamTableCachedObj | None"]
def spend_attribution_enabled(general_settings: Mapping[str, object] | None) -> bool:
"""Whether spend, rollups, and budget gates fan out across memberships."""
if not general_settings:
return False
return general_settings.get(SPEND_ATTRIBUTION_SETTING) is True
def rate_limit_attribution_enabled(general_settings: Mapping[str, object] | None) -> bool:
"""Whether the RPM/TPM limiter fans out across memberships."""
if not general_settings:
return False
return general_settings.get(RATE_LIMIT_ATTRIBUTION_SETTING) is True
def _attribution_enabled(general_settings: Mapping[str, object] | None) -> bool:
return spend_attribution_enabled(general_settings) or rate_limit_attribution_enabled(general_settings)
def attributed_team_ids(valid_token: UserAPIKeyAuth | None) -> tuple[str, ...]:
"""Every team this request is attributed to, stamped team first.
Falls back to the single stamped team whenever attribution is off or
resolved nothing, so a call site can use this unconditionally and keep
identical behavior with the settings disabled.
"""
if valid_token is None:
return ()
resolved: Final = valid_token.attributed_team_ids
if resolved:
return tuple(resolved)
return (valid_token.team_id,) if valid_token.team_id else ()
def attributed_org_ids(valid_token: UserAPIKeyAuth | None) -> tuple[str, ...]:
"""Every organization this request is attributed to.
Same fallback contract as :func:`attributed_team_ids`.
"""
if valid_token is None:
return ()
resolved: Final = valid_token.attributed_org_ids
if resolved:
return tuple(resolved)
return (valid_token.org_id,) if valid_token.org_id else ()
def attribution_targets(attributed_ids: Sequence[str] | None, stamped_id: str | None) -> tuple[str, ...]:
"""The entity ids one request should be charged against.
The attributed set when membership attribution resolved one, otherwise the
single stamped id -- so with both settings off this returns exactly
``(stamped_id,)`` and every caller keeps its historical behavior.
Order-preserving dedupe: the stamped team normally also appears in the
caller's membership list, and charging it twice would double-count.
"""
if attributed_ids:
return _ordered_unique(attributed_ids)
return (stamped_id,) if stamped_id else ()
def _ordered_unique(values: Iterable[str | None]) -> tuple[str, ...]:
"""Dedupe while preserving order, dropping empties.
Order is preserved so the stamped team stays first. Budget and rate-limit
errors report the first offending entity, and a caller reading that error
is best served by hearing about the team their key actually names.
"""
return tuple(dict.fromkeys(v for v in values if v))
async def resolve_membership_attribution(
*,
user_api_key_auth_obj: UserAPIKeyAuth,
user_object: LiteLLM_UserTable | None,
general_settings: Mapping[str, object] | None,
prisma_client: "PrismaClient | None",
user_api_key_cache: "UserApiKeyCache",
proxy_logging_obj: "ProxyLogging | None" = None,
) -> None:
"""Populate the attributed-membership fields on ``user_api_key_auth_obj``.
No-op unless one of the attribution settings is on, so the default path
pays nothing -- not even a cache read.
Fails open. A team that cannot be resolved (deleted row, transient DB
error) is skipped rather than raised: attribution is bookkeeping layered on
top of an authorization decision that has already been made, and losing one
team's attribution must never turn an authorized request into a 500. The
stamped team is seeded first and is never dropped, so a failure degrades to
exactly the default single-team behavior.
"""
if not _attribution_enabled(general_settings):
return
if prisma_client is None:
return
# The stamped team leads, then the caller's SCIM/IdP-maintained membership
# list. Seeding the stamped team explicitly matters: LiteLLM_UserTable.teams
# is maintained by SCIM and JWT sync, so it can lag a key that was just
# pointed at a new team, and losing that team would silently under-charge
# the one team the operator explicitly named.
candidate_team_ids: Final = _ordered_unique(
(user_api_key_auth_obj.team_id, *(user_object.teams if user_object and user_object.teams else ()))
)
if not candidate_team_ids:
_apply_org_attribution(
user_api_key_auth_obj=user_api_key_auth_obj,
user_object=user_object,
team_objects=(),
)
return
team_objects: Final = await _load_team_objects(
team_ids=candidate_team_ids,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth_obj.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
resolved_ids: Final = tuple(team_id for team_id, team_object in team_objects if team_object is not None)
if resolved_ids:
team_limits: Final = MappingProxyType(
{
team_id: MappingProxyType(
{
"rpm": getattr(team_object, "rpm_limit", None),
"tpm": getattr(team_object, "tpm_limit", None),
}
)
for team_id, team_object in team_objects
if team_object is not None
}
)
# The resolved context belongs on the auth object every later stage
# already reads -- the key/team org fallback immediately above this
# call does exactly the same. Returning a new object instead would
# mean rebuilding every consumer of user_api_key_auth.
user_api_key_auth_obj.attributed_team_ids = resolved_ids # rebind-ok: stamping resolved auth context
user_api_key_auth_obj.attributed_team_limits = team_limits # rebind-ok: stamping resolved auth context
_apply_org_attribution(
user_api_key_auth_obj=user_api_key_auth_obj,
user_object=user_object,
team_objects=team_objects,
)
def _apply_org_attribution(
*,
user_api_key_auth_obj: UserAPIKeyAuth,
user_object: LiteLLM_UserTable | None,
team_objects: Sequence[TeamResolution],
) -> None:
"""Derive the attributed organizations from what is already loaded.
Sources, in order: the org already resolved onto the token, the caller's own
``LiteLLM_UserTable.organization_id``, then the org of each attributed team.
Deliberately no ``LiteLLM_OrganizationMembership`` query -- that table would
add a fresh round trip to the hot path, and these three sources already
cover both the common single-org deployment and a caller whose teams span
several organizations.
"""
org_ids: Final = _ordered_unique(
(
user_api_key_auth_obj.org_id,
user_object.organization_id if user_object else None,
*(
getattr(team_object, "organization_id", None)
for _team_id, team_object in team_objects
if team_object is not None
),
)
)
if org_ids:
user_api_key_auth_obj.attributed_org_ids = org_ids # rebind-ok: stamping resolved auth context
async def _load_team_objects(
*,
team_ids: Sequence[str],
prisma_client: "PrismaClient",
user_api_key_cache: "UserApiKeyCache",
parent_otel_span: "Span | None",
proxy_logging_obj: "ProxyLogging | None",
) -> tuple[TeamResolution, ...]:
"""Resolve every candidate team, preserving input order.
Lookups run concurrently: ``get_team_object`` is cache-first, so the steady
state is N in-memory hits, but a cold pod pays a DB read per team and those
must not serialize.
"""
from litellm.proxy.auth.auth_checks import get_team_object
async def _safe_get(team_id: str) -> TeamResolution:
try:
team_object: Final = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception as e: # noqa: BLE001 # attribution must never fail an authorized request
verbose_proxy_logger.debug(
"membership attribution: skipping team_id=%s, could not resolve: %s",
team_id,
e,
)
return team_id, None
else:
return team_id, team_object
return tuple(await asyncio.gather(*(_safe_get(team_id) for team_id in team_ids)))

View file

@ -71,6 +71,7 @@ from litellm.proxy.auth.auth_utils import (
route_in_additonal_public_routes,
)
from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler
from litellm.proxy.auth.membership_attribution import resolve_membership_attribution
from litellm.proxy.auth.network import TrustedProxyConfig, resolve_network_context
from litellm.proxy.auth.oauth2_check import Oauth2Handler
from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
@ -2581,6 +2582,20 @@ async def _run_centralized_common_checks(
if user_api_key_auth_obj.org_id is None and team_object is not None and team_object.organization_id is not None:
user_api_key_auth_obj.org_id = team_object.organization_id
# Expand the single stamped team/org into every membership the caller has,
# when the operator opted in. Must run BEFORE common_checks so the
# all-teams budget gate sees the expansion, and before the request metadata
# is stamped so the cost callback can attribute spend to the same set.
# No-op (not even a cache read) when both settings are off.
await resolve_membership_attribution(
user_api_key_auth_obj=user_api_key_auth_obj,
user_object=user_object,
general_settings=general_settings,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# common_checks identifies admin via user_object, not the token
# (non_proxy_admin_allowed_routes_check). JWT admin shortcut and
# master_key tokens get admin from the token; the DB row for the

View file

@ -12,6 +12,7 @@ import os
import random
import time
import traceback
from collections.abc import Sequence
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload
@ -160,6 +161,11 @@ class DBSpendUpdateWriter:
start_time: datetime | None,
end_time: datetime | None,
response_cost: float | None,
# Every team/org this request is attributed to, when
# track_spend_across_all_user_teams is on. None keeps the historical
# single-team/single-org behavior.
attributed_team_ids: Sequence[str] | None = None,
attributed_org_ids: Sequence[str] | None = None,
):
from litellm.proxy.proxy_server import (
disable_spend_logs,
@ -238,6 +244,8 @@ class DBSpendUpdateWriter:
prisma_client=prisma_client,
litellm_proxy_budget_name=litellm_proxy_budget_name,
payload=payload,
attributed_team_ids=attributed_team_ids,
attributed_org_ids=attributed_org_ids,
)
)
@ -429,6 +437,8 @@ class DBSpendUpdateWriter:
prisma_client: PrismaClient | None,
litellm_proxy_budget_name: str | None,
payload: SpendLogsPayload,
attributed_team_ids: Sequence[str] | None = None,
attributed_org_ids: Sequence[str] | None = None,
):
"""
Runs all 11 spend-update helpers sequentially inside a single asyncio task.
@ -472,6 +482,7 @@ class DBSpendUpdateWriter:
team_id=team_id,
user_id=user_id,
prisma_client=prisma_client,
attributed_team_ids=attributed_team_ids,
)
except Exception:
verbose_proxy_logger.debug(
@ -484,6 +495,7 @@ class DBSpendUpdateWriter:
response_cost=response_cost,
org_id=org_id,
prisma_client=prisma_client,
attributed_org_ids=attributed_org_ids,
)
except Exception:
verbose_proxy_logger.debug(
@ -553,6 +565,7 @@ class DBSpendUpdateWriter:
await self.add_spend_log_transaction_to_daily_team_transaction(
payload=payload_copy,
prisma_client=prisma_client,
attributed_team_ids=attributed_team_ids,
)
except Exception:
verbose_proxy_logger.debug(
@ -565,6 +578,7 @@ class DBSpendUpdateWriter:
payload=payload_copy,
org_id=org_id,
prisma_client=prisma_client,
attributed_org_ids=attributed_org_ids,
)
except Exception:
verbose_proxy_logger.debug(
@ -651,54 +665,76 @@ class DBSpendUpdateWriter:
exc=e,
)
@staticmethod
def _attribution_targets(attributed_ids: Sequence[str] | None, stamped_id: str | None) -> tuple[str, ...]:
"""The entity ids to charge for one request. See ``attribution_targets``."""
from litellm.proxy.auth.membership_attribution import attribution_targets
return attribution_targets(attributed_ids, stamped_id)
async def _update_team_db(
self,
response_cost: float | None,
team_id: str | None,
user_id: str | None,
prisma_client: PrismaClient | None,
attributed_team_ids: Sequence[str] | None = None,
):
"""Charge the request to every team it is attributed to.
With ``track_spend_across_all_user_teams`` off, ``attributed_team_ids``
is None and this charges the single stamped team, exactly as before.
Note the deliberate consequence when it is on: summing team spend across
the deployment then exceeds real spend, because one request is charged to
each of the caller's teams on purpose. The key, user, and org rows stay
single-counted, so those remain true totals.
"""
target_team_ids: Final = self._attribution_targets(attributed_team_ids, team_id)
try:
if team_id is None or prisma_client is None:
if not target_team_ids or prisma_client is None:
verbose_proxy_logger.debug(
"track_cost_callback: team_id is None or prisma_client is None. Not tracking spend for team"
"track_cost_callback: no attributed team or prisma_client is None. Not tracking spend for team"
)
return
await self.spend_update_queue.add_update(
update=SpendUpdateQueueItem(
entity_type=Litellm_EntityType.TEAM,
entity_id=team_id,
response_cost=response_cost,
)
)
try:
# Track spend of the team member within this team
if user_id is not None:
# key is "team_id::<value>::user_id::<value>"
team_member_key: Final = f"team_id::{team_id}::user_id::{user_id}"
await self.spend_update_queue.add_update(
update=SpendUpdateQueueItem(
entity_type=Litellm_EntityType.TEAM_MEMBER,
entity_id=team_member_key,
response_cost=response_cost,
)
for target_team_id in target_team_ids:
await self.spend_update_queue.add_update(
update=SpendUpdateQueueItem(
entity_type=Litellm_EntityType.TEAM,
entity_id=target_team_id,
response_cost=response_cost,
)
except Exception as e:
spend_log_error(
"Spend tracking - failed to enqueue team member spend update. "
"team_id=%s, user_id=%s, response_cost=%s - %s",
team_id,
user_id,
response_cost,
str(e),
exc=e,
)
try:
# Track spend of the team member within this team
if user_id is not None:
# key is "team_id::<value>::user_id::<value>"
team_member_key = (
f"team_id::{target_team_id}::user_id::{user_id}" # rebind-ok: one per attributed team
)
await self.spend_update_queue.add_update(
update=SpendUpdateQueueItem(
entity_type=Litellm_EntityType.TEAM_MEMBER,
entity_id=team_member_key,
response_cost=response_cost,
)
)
except Exception as e:
spend_log_error(
"Spend tracking - failed to enqueue team member spend update. "
"team_id=%s, user_id=%s, response_cost=%s - %s",
target_team_id,
user_id,
response_cost,
str(e),
exc=e,
)
except Exception as e:
spend_log_error(
"Spend tracking - failed to enqueue team spend update. team_id=%s, response_cost=%s - %s",
team_id,
"Spend tracking - failed to enqueue team spend update. team_ids=%s, response_cost=%s - %s",
target_team_ids,
response_cost,
str(e),
exc=e,
@ -710,25 +746,33 @@ class DBSpendUpdateWriter:
response_cost: float | None,
org_id: str | None,
prisma_client: PrismaClient | None,
attributed_org_ids: Sequence[str] | None = None,
):
"""Charge the request to every organization it is attributed to.
Normally a single org: a caller belongs to one, and extra ids only
appear when their teams span several organizations.
"""
target_org_ids: Final = self._attribution_targets(attributed_org_ids, org_id)
try:
if org_id is None or prisma_client is None:
if not target_org_ids or prisma_client is None:
verbose_proxy_logger.debug(
"track_cost_callback: org_id is None or prisma_client is None. Not tracking spend for org"
"track_cost_callback: no attributed org or prisma_client is None. Not tracking spend for org"
)
return
await self.spend_update_queue.add_update(
update=SpendUpdateQueueItem(
entity_type=Litellm_EntityType.ORGANIZATION,
entity_id=org_id,
response_cost=response_cost,
for target_org_id in target_org_ids:
await self.spend_update_queue.add_update(
update=SpendUpdateQueueItem(
entity_type=Litellm_EntityType.ORGANIZATION,
entity_id=target_org_id,
response_cost=response_cost,
)
)
)
except Exception as e:
spend_log_error(
"Spend tracking - failed to enqueue org spend update. org_id=%s, response_cost=%s - %s",
org_id,
"Spend tracking - failed to enqueue org spend update. org_ids=%s, response_cost=%s - %s",
target_org_ids,
response_cost,
str(e),
exc=e,
@ -1948,7 +1992,17 @@ class DBSpendUpdateWriter:
self,
payload: SpendLogsPayload,
prisma_client: PrismaClient | None = None,
attributed_team_ids: Sequence[str] | None = None,
) -> None:
"""Enqueue one daily rollup row per attributed team.
``LiteLLM_DailyTeamSpend`` is already unique on
``(team_id, date, api_key, model, custom_llm_provider,
mcp_namespaced_tool_name, endpoint)``, so N teams means N distinct rows
and no migration is needed. The base transaction is computed once and
reused: it holds only per-request facts (tokens, cost, date, endpoint),
none of which depend on which team the row is keyed to.
"""
if prisma_client is None:
verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.")
return
@ -1958,47 +2012,56 @@ class DBSpendUpdateWriter:
)
if base_daily_transaction is None:
return
if payload["team_id"] is None:
verbose_proxy_logger.debug("team_id is None for request. Skipping incrementing team spend.")
target_team_ids: Final = self._attribution_targets(attributed_team_ids, payload["team_id"])
if not target_team_ids:
verbose_proxy_logger.debug("no attributed team for request. Skipping incrementing team spend.")
return
endpoint_str: Final = base_daily_transaction.get("endpoint") or ""
daily_transaction_key = f"{payload['team_id']}_{base_daily_transaction['date']}_{payload['api_key']}_{payload['model']}_{payload['custom_llm_provider']}_{endpoint_str}"
daily_transaction: Final = DailyTeamSpendTransaction(team_id=payload["team_id"], **base_daily_transaction)
await self.daily_team_spend_update_queue.add_update(update={daily_transaction_key: daily_transaction})
for target_team_id in target_team_ids:
daily_transaction_key = f"{target_team_id}_{base_daily_transaction['date']}_{payload['api_key']}_{payload['model']}_{payload['custom_llm_provider']}_{endpoint_str}"
daily_transaction = DailyTeamSpendTransaction(team_id=target_team_id, **base_daily_transaction)
await self.daily_team_spend_update_queue.add_update(update={daily_transaction_key: daily_transaction})
async def add_spend_log_transaction_to_daily_org_transaction(
self,
payload: SpendLogsPayload,
prisma_client: PrismaClient | None = None,
org_id: str | None = None,
attributed_org_ids: Sequence[str] | None = None,
) -> None:
"""Enqueue one daily rollup row per attributed organization."""
if prisma_client is None:
verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.")
return
if org_id is None:
target_org_ids: Final = self._attribution_targets(attributed_org_ids, org_id)
if not target_org_ids:
verbose_proxy_logger.debug("organization_id is None for request. Skipping incrementing organization spend.")
return
payload_with_org: Final = cast(
SpendLogsPayload,
{
**payload,
"organization_id": org_id,
},
)
for target_org_id in target_org_ids:
payload_with_org = cast(
SpendLogsPayload,
{
**payload,
"organization_id": target_org_id,
},
)
base_daily_transaction: Final = await self._common_add_spend_log_transaction_to_daily_transaction(
payload_with_org, prisma_client, "org"
)
if base_daily_transaction is None:
return
base_daily_transaction = await self._common_add_spend_log_transaction_to_daily_transaction(
payload_with_org, prisma_client, "org"
)
if base_daily_transaction is None:
continue
endpoint_str: Final = base_daily_transaction.get("endpoint") or ""
daily_transaction_key = f"{org_id}_{base_daily_transaction['date']}_{payload_with_org['api_key']}_{payload_with_org['model']}_{payload_with_org['custom_llm_provider']}_{endpoint_str}"
daily_transaction: Final = DailyOrganizationSpendTransaction(organization_id=org_id, **base_daily_transaction)
await self.daily_org_spend_update_queue.add_update(update={daily_transaction_key: daily_transaction})
endpoint_str = base_daily_transaction.get("endpoint") or ""
daily_transaction_key = f"{target_org_id}_{base_daily_transaction['date']}_{payload_with_org['api_key']}_{payload_with_org['model']}_{payload_with_org['custom_llm_provider']}_{endpoint_str}"
daily_transaction = DailyOrganizationSpendTransaction(
organization_id=target_org_id, **base_daily_transaction
)
await self.daily_org_spend_update_queue.add_update(update={daily_transaction_key: daily_transaction})
async def add_spend_log_transaction_to_daily_end_user_transaction(
self,

View file

@ -1198,25 +1198,44 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if self.batch_rate_limiter_script is None:
return []
# Bound here, not read off self inside the closure: the None check above
# narrows the attribute, but that narrowing does not reach into a nested
# function, so reading it there would be an optional call.
run_script: Final = self.batch_rate_limiter_script
key_groups: Final = self._group_keys_by_hash_tag(keys_to_fetch)
all_cache_values: Final[list[CacheCounterValue | None]] = []
for hash_tag, group_keys in key_groups.items():
# Both callees below (the Lua script and the in-memory fallback) declare
# `keys: list[str]`, and _group_keys_by_hash_tag hands back exactly that,
# so a read-only view here would not be assignable.
async def _read_group(tag: str, group_keys: list[str]) -> CacheCounterValues: # mutable-ok: callee needs list
try:
group_cache_values: CacheCounterValues = await self.batch_rate_limiter_script(
return await run_script(
keys=group_keys,
args=[now_int, self.window_size], # Use integer timestamp
)
all_cache_values.extend(group_cache_values)
except Exception as e:
verbose_proxy_logger.warning("Redis Lua script failed for hash tag %s: %s", hash_tag, e)
verbose_proxy_logger.warning("Redis Lua script failed for hash tag %s: %s", tag, e)
# Fallback to in-memory cache for this group
group_cache_values = await self.in_memory_cache_sliding_window(
return await self.in_memory_cache_sliding_window(
keys=group_keys,
now_int=now_int,
window_size=self.window_size,
)
all_cache_values.extend(group_cache_values)
# One Lua round trip per hash-tag group, issued concurrently. On
# non-cluster Redis there is exactly one group, so this stays a single
# call no matter how many descriptors the request carries. On Redis
# Cluster the groups are per-slot, and membership-attributed team
# descriptors can produce one group per team -- running them
# concurrently keeps the added latency at roughly one round trip
# instead of N. Results are consumed in group order so the returned
# list still lines up with ``keys_to_fetch``.
group_results: Final = await asyncio.gather(
*(_read_group(hash_tag, group_keys) for hash_tag, group_keys in key_groups.items())
)
for group_cache_values in group_results:
all_cache_values.extend(group_cache_values)
return all_cache_values
@ -2350,6 +2369,53 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
)
def _attributed_team_rate_limit_descriptors(
self,
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[RateLimitDescriptor, ...]:
"""Add a team descriptor for every OTHER team the caller belongs to.
Only active when ``enforce_rate_limits_across_all_user_teams`` is on;
otherwise ``attributed_team_limits`` is None and this returns at once.
``key="team"`` is reused deliberately: team T must share one bucket
whether T is the team stamped on the key or one the caller merely
belongs to, or the same team would be limited twice over under two
different counters. The stamped team is skipped because the block above
already emitted its descriptor.
Because ``should_rate_limit`` rejects when ANY descriptor is over, the
caller's effective limit becomes the minimum across their memberships.
That is the intended semantic, and it is why this rides a separate
setting from spend attribution: a busy team can now throttle someone
who is mostly working for a different team.
"""
from litellm.proxy.auth.membership_attribution import rate_limit_attribution_enabled
from litellm.proxy.proxy_server import general_settings
if not rate_limit_attribution_enabled(general_settings):
return ()
attributed_limits: Final = user_api_key_dict.attributed_team_limits
if not attributed_limits:
return ()
return tuple(
RateLimitDescriptor(
key="team",
value=team_id,
rate_limit=RateLimitDescriptorRateLimitObject(
requests_per_unit=team_limits.get("rpm"),
tokens_per_unit=team_limits.get("tpm"),
window_size=self.window_size,
),
)
for team_id, team_limits in attributed_limits.items()
if team_id
and team_id != user_api_key_dict.team_id
and not (team_limits.get("rpm") is None and team_limits.get("tpm") is None)
)
def _add_tag_per_key_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -2701,6 +2767,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
)
descriptors.extend(self._attributed_team_rate_limit_descriptors(user_api_key_dict=user_api_key_dict))
# Team Member rate limits
if user_api_key_dict.user_id and (
user_api_key_dict.team_member_rpm_limit is not None or user_api_key_dict.team_member_tpm_limit is not None

View file

@ -1,5 +1,6 @@
import asyncio
import traceback
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any, Final, cast
@ -206,6 +207,8 @@ class _ProxyDBLogger(CustomLogger):
start_time=actual_start_time,
end_time=datetime.now(),
org_id=user_api_key_dict.org_id,
attributed_team_ids=_attributed_team_ids(user_api_key_dict),
attributed_org_ids=_attributed_org_ids(user_api_key_dict),
)
@log_db_metrics
@ -245,6 +248,10 @@ class _ProxyDBLogger(CustomLogger):
user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None))
team_id: Final = cast(str | None, metadata.get("user_api_key_team_id", None))
org_id: Final = cast(str | None, metadata.get("user_api_key_org_id", None))
# Present only when track_spend_across_all_user_teams is on; None
# otherwise, which keeps the single-team write path unchanged.
attributed_team_ids: Final = _metadata_id_list(metadata, "user_api_key_attributed_team_ids")
attributed_org_ids: Final = _metadata_id_list(metadata, "user_api_key_attributed_org_ids")
key_alias: Final = cast(str | None, metadata.get("user_api_key_alias", None))
end_user_max_budget: Final = metadata.get("user_api_end_user_max_budget", None)
sl_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
@ -293,6 +300,8 @@ class _ProxyDBLogger(CustomLogger):
response_cost=response_cost,
budget_reservation=budget_reservation,
request_tags=tags,
attributed_team_ids=attributed_team_ids,
attributed_org_ids=attributed_org_ids,
)
# update cache (fire-and-forget for backward compat:
@ -538,6 +547,33 @@ def _get_budget_reservation_from_metadata(metadata: dict) -> dict | None:
return getattr(user_api_key_auth_obj, "budget_reservation", None)
def _metadata_id_list(metadata: Mapping[str, object], key: str) -> tuple[str, ...] | None:
"""Read a stamped id list out of request metadata, or None if absent.
Returns None (not an empty tuple) when the key is missing, so downstream
writers can tell "attribution off" apart from "attribution on, nothing
resolved".
"""
value: Final = metadata.get(key)
# tuple as stamped, list after any JSON round trip
if not isinstance(value, (list, tuple)):
return None
ids: Final = tuple(v for v in value if isinstance(v, str) and v)
return ids or None
def _attributed_team_ids(user_api_key_dict: UserAPIKeyAuth) -> tuple[str, ...] | None:
from litellm.proxy.auth.membership_attribution import attributed_team_ids
return attributed_team_ids(user_api_key_dict) if user_api_key_dict.attributed_team_ids else None
def _attributed_org_ids(user_api_key_dict: UserAPIKeyAuth) -> tuple[str, ...] | None:
from litellm.proxy.auth.membership_attribution import attributed_org_ids
return attributed_org_ids(user_api_key_dict) if user_api_key_dict.attributed_org_ids else None
def _get_request_tags_for_cost_tracking(
sl_object: StandardLoggingPayload | None,
metadata: dict,
@ -569,6 +605,8 @@ async def _update_database_and_spend_counters(
response_cost: float,
budget_reservation: dict | None,
request_tags: list[str] | None = None,
attributed_team_ids: Sequence[str] | None = None,
attributed_org_ids: Sequence[str] | None = None,
) -> None:
try:
await proxy_logging_obj.db_spend_update_writer.update_database(
@ -582,6 +620,8 @@ async def _update_database_and_spend_counters(
start_time=start_time,
end_time=end_time,
org_id=org_id,
attributed_team_ids=attributed_team_ids,
attributed_org_ids=attributed_org_ids,
)
except Exception:
if budget_reservation is not None:
@ -607,6 +647,8 @@ async def _update_database_and_spend_counters(
budget_reservation=budget_reservation,
end_user_id=end_user_id,
tags=request_tags,
attributed_team_ids=attributed_team_ids,
attributed_org_ids=attributed_org_ids,
)
except Exception:
if budget_reservation is not None:

View file

@ -1355,6 +1355,19 @@ class LiteLLMProxyRequestSetup:
)
if user_api_key_dict.budget_reservation is not None:
data[_metadata_variable_name]["user_api_key_budget_reservation"] = user_api_key_dict.budget_reservation
# Every team/org this request is attributed to, stamped as plain metadata
# so the cost callback reads it the same way it reads
# user_api_key_team_id. Written only when track_spend_across_all_user_teams
# is on, so the default path adds no metadata keys at all. Deliberately
# not folded into StandardLoggingUserAPIKeyMetadata: that TypedDict is
# consumed by ~20 logging integrations and widening it would make this a
# breaking payload change rather than an additive proxy feature.
LiteLLMProxyRequestSetup._add_attributed_membership_metadata(
data=data,
user_api_key_dict=user_api_key_dict,
_metadata_variable_name=_metadata_variable_name,
)
# UserAPIKeyAuth object for MCP server access control
data[_metadata_variable_name]["user_api_key_auth"] = user_api_key_dict.model_copy(
update={
@ -1366,6 +1379,39 @@ class LiteLLMProxyRequestSetup:
)
return data
@staticmethod
def _add_attributed_membership_metadata(
data: dict, # mutable-ok: every sibling stamper in this class writes into the caller's request dict
user_api_key_dict: UserAPIKeyAuth,
_metadata_variable_name: str,
) -> None:
"""Stamp the attributed team/org ids for the cost callback.
No-op unless ``track_spend_across_all_user_teams`` is on AND auth
actually resolved a membership set, so a deployment with the setting off
carries exactly the metadata it carries today.
"""
from litellm.proxy.auth.membership_attribution import (
attributed_org_ids,
attributed_team_ids,
spend_attribution_enabled,
)
from litellm.proxy.proxy_server import general_settings
if not spend_attribution_enabled(general_settings):
return
if user_api_key_dict.attributed_team_ids is None and user_api_key_dict.attributed_org_ids is None:
return
# Stored as tuples. This dict is serialized into the spend-log payload,
# and json.dumps renders a tuple as an array exactly like a list; the
# reader (_metadata_id_list) accepts either, so a JSON round trip that
# turns these back into lists is also fine.
team_ids: Final = attributed_team_ids(user_api_key_dict)
org_ids: Final = attributed_org_ids(user_api_key_dict)
data[_metadata_variable_name]["user_api_key_attributed_team_ids"] = team_ids # rebind-ok: metadata stamp
data[_metadata_variable_name]["user_api_key_attributed_org_ids"] = org_ids # rebind-ok: metadata stamp
@staticmethod
def add_management_endpoint_metadata_to_request_metadata(
data: dict,

View file

@ -292,6 +292,7 @@ from litellm.proxy.auth.auth_utils import (
)
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.litellm_license import LicenseCheck
from litellm.proxy.auth.membership_attribution import attribution_targets
from litellm.proxy.auth.model_checks import (
expand_wildcard_deployments_for_model_info,
get_all_fallbacks,
@ -2633,6 +2634,8 @@ async def increment_spend_counters(
budget_reservation: dict | None = None,
end_user_id: str | None = None,
tags: list[str] | None = None,
attributed_team_ids: Sequence[str] | None = None,
attributed_org_ids: Sequence[str] | None = None,
):
"""
Atomically increment spend counters for budget enforcement.
@ -2747,12 +2750,21 @@ async def increment_spend_counters(
increment=cost,
)
# With track_spend_across_all_user_teams on, one request increments the
# counter of every team the caller belongs to, and of every org reached
# through those teams. These are the counters the budget gates read
# Redis-first, so enforcement across memberships is immediate rather than
# waiting for the batched DB flush. With the setting off, both lists
# collapse to the single stamped id and the scopes are unchanged.
target_team_ids: Final = attribution_targets(attributed_team_ids, team_id)
target_org_ids: Final = attribution_targets(attributed_org_ids, org_id)
scope_coros: Final = tuple(
coro
for coro in (
_key_scope(token) if token is not None else None,
_team_scope(team_id) if team_id is not None else None,
_team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None,
*(_team_scope(scope_team_id) for scope_team_id in target_team_ids),
*(_team_member_scope(user_id, scope_team_id) for scope_team_id in target_team_ids if user_id is not None),
_user_scope(user_id) if user_id is not None else None,
_increment_end_user_and_tag_spend_counters(
end_user_id=end_user_id,
@ -2762,13 +2774,14 @@ async def increment_spend_counters(
)
if end_user_id is not None or tags is not None
else None,
_increment_org_spend_counter(
org_id=org_id,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
)
if org_id is not None
else None,
*(
_increment_org_spend_counter(
org_id=scope_org_id,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
)
for scope_org_id in target_org_ids
),
)
if coro is not None
)

View file

@ -0,0 +1,611 @@
"""Tests for membership-based usage attribution.
Covers the two opt-in settings introduced alongside
``litellm/proxy/auth/membership_attribution.py``:
- ``track_spend_across_all_user_teams``
- ``enforce_rate_limits_across_all_user_teams``
The most important cases here are the OFF cases. Both settings default to off,
and every one of those tests is a regression guard proving the default path is
byte-for-byte what it was before the feature existed.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm.proxy._types import Litellm_EntityType, LiteLLM_UserTable, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
_attributed_orgs_max_budget_check,
_attributed_teams_max_budget_check,
)
from litellm.proxy.auth.membership_attribution import (
attributed_org_ids,
attributed_team_ids,
attribution_targets,
rate_limit_attribution_enabled,
resolve_membership_attribution,
spend_attribution_enabled,
)
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
# --------------------------------------------------------------------------- #
# settings gates
# --------------------------------------------------------------------------- #
@pytest.mark.parametrize("general_settings", [None, {}, {"track_spend_across_all_user_teams": False}])
def test_spend_attribution_defaults_off(general_settings):
assert spend_attribution_enabled(general_settings) is False
@pytest.mark.parametrize("general_settings", [None, {}, {"enforce_rate_limits_across_all_user_teams": False}])
def test_rate_limit_attribution_defaults_off(general_settings):
assert rate_limit_attribution_enabled(general_settings) is False
def test_settings_are_independent():
"""An operator must be able to take spend attribution without rate limits."""
spend_only = {"track_spend_across_all_user_teams": True}
assert spend_attribution_enabled(spend_only) is True
assert rate_limit_attribution_enabled(spend_only) is False
# --------------------------------------------------------------------------- #
# target resolution
# --------------------------------------------------------------------------- #
def test_attribution_targets_falls_back_to_stamped_id():
assert attribution_targets(None, "team-a") == ("team-a",)
assert attribution_targets([], "team-a") == ("team-a",)
assert attribution_targets(None, None) == ()
def test_attribution_targets_dedupes_and_preserves_order():
"""The stamped team is normally also in the membership list; charge it once."""
assert attribution_targets(["team-a", "team-b", "team-a", ""], "team-a") == ("team-a", "team-b")
def test_read_helpers_fall_back_to_stamped_values():
token = UserAPIKeyAuth(api_key="sk-1", team_id="team-a", org_id="org-1")
assert attributed_team_ids(token) == ("team-a",)
assert attributed_org_ids(token) == ("org-1",)
def test_attributed_fields_cannot_be_forged_from_input():
"""Server-only fields. A caller who could set these would pick their own
budget and rate-limit buckets."""
forged = UserAPIKeyAuth(
api_key="sk-1",
team_id="team-a",
attributed_team_ids=["team-with-huge-budget"],
attributed_org_ids=["org-with-huge-budget"],
attributed_team_limits={"team-with-huge-budget": {"rpm": 10**9}},
)
assert forged.attributed_team_ids is None
assert forged.attributed_org_ids is None
assert forged.attributed_team_limits is None
# --------------------------------------------------------------------------- #
# resolver
# --------------------------------------------------------------------------- #
@pytest.mark.asyncio
async def test_resolver_is_noop_when_both_settings_off():
"""Not even a cache read on the default path."""
token = UserAPIKeyAuth(api_key="sk-1", team_id="team-a")
with patch("litellm.proxy.auth.auth_checks.get_team_object", new=AsyncMock()) as mock_get_team:
await resolve_membership_attribution(
user_api_key_auth_obj=token,
user_object=LiteLLM_UserTable(user_id="u1", teams=["team-a", "team-b"]),
general_settings={},
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
)
mock_get_team.assert_not_called()
assert token.attributed_team_ids is None
@pytest.mark.asyncio
async def test_resolver_collects_every_membership_with_stamped_team_first():
token = UserAPIKeyAuth(api_key="sk-1", team_id="team-a")
def _team(team_id: str, rpm=None, tpm=None, org="org-1"):
obj = MagicMock()
obj.team_id = team_id
obj.rpm_limit = rpm
obj.tpm_limit = tpm
obj.organization_id = org
return obj
async def _fake_get_team_object(team_id, **kwargs):
return _team(team_id, rpm=10 if team_id == "team-b" else None)
with patch("litellm.proxy.auth.auth_checks.get_team_object", new=AsyncMock(side_effect=_fake_get_team_object)):
await resolve_membership_attribution(
user_api_key_auth_obj=token,
user_object=LiteLLM_UserTable(user_id="u1", teams=["team-b", "team-a", "team-c"]),
general_settings={"track_spend_across_all_user_teams": True},
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
)
# stamped team leads, membership list follows, no duplicate of team-a
assert token.attributed_team_ids == ("team-a", "team-b", "team-c")
assert token.attributed_team_limits["team-b"] == {"rpm": 10, "tpm": None}
assert token.attributed_org_ids == ("org-1",)
@pytest.mark.asyncio
async def test_resolver_fails_open_on_unresolvable_team():
"""A deleted or erroring team is skipped, never raised: attribution must not
turn an already-authorized request into a 500."""
token = UserAPIKeyAuth(api_key="sk-1", team_id="team-a")
async def _fake_get_team_object(team_id, **kwargs):
if team_id == "team-broken":
raise Exception("team row is gone")
obj = MagicMock()
obj.rpm_limit = None
obj.tpm_limit = None
obj.organization_id = "org-1"
return obj
with patch("litellm.proxy.auth.auth_checks.get_team_object", new=AsyncMock(side_effect=_fake_get_team_object)):
await resolve_membership_attribution(
user_api_key_auth_obj=token,
user_object=LiteLLM_UserTable(user_id="u1", teams=["team-broken", "team-a"]),
general_settings={"track_spend_across_all_user_teams": True},
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
)
assert token.attributed_team_ids == ("team-a",)
# --------------------------------------------------------------------------- #
# spend fan-out
# --------------------------------------------------------------------------- #
def _enqueued(mock_add_update, entity_type):
return [
c.kwargs["update"]["entity_id"]
for c in mock_add_update.call_args_list
if c.kwargs["update"]["entity_type"] == entity_type
]
@pytest.mark.asyncio
async def test_update_team_db_single_team_when_attribution_off():
"""Regression guard for the default path."""
writer = DBSpendUpdateWriter()
writer.spend_update_queue.add_update = AsyncMock()
await writer._update_team_db(
response_cost=1.0,
team_id="team-a",
user_id="u1",
prisma_client=MagicMock(),
)
assert _enqueued(writer.spend_update_queue.add_update, Litellm_EntityType.TEAM) == ["team-a"]
assert _enqueued(writer.spend_update_queue.add_update, Litellm_EntityType.TEAM_MEMBER) == [
"team_id::team-a::user_id::u1"
]
@pytest.mark.asyncio
async def test_update_team_db_charges_every_attributed_team():
writer = DBSpendUpdateWriter()
writer.spend_update_queue.add_update = AsyncMock()
await writer._update_team_db(
response_cost=1.0,
team_id="team-a",
user_id="u1",
prisma_client=MagicMock(),
attributed_team_ids=["team-a", "team-b", "team-c"],
)
assert _enqueued(writer.spend_update_queue.add_update, Litellm_EntityType.TEAM) == [
"team-a",
"team-b",
"team-c",
]
assert _enqueued(writer.spend_update_queue.add_update, Litellm_EntityType.TEAM_MEMBER) == [
"team_id::team-a::user_id::u1",
"team_id::team-b::user_id::u1",
"team_id::team-c::user_id::u1",
]
@pytest.mark.asyncio
async def test_update_org_db_charges_every_attributed_org():
writer = DBSpendUpdateWriter()
writer.spend_update_queue.add_update = AsyncMock()
await writer._update_org_db(
response_cost=1.0,
org_id="org-1",
prisma_client=MagicMock(),
attributed_org_ids=["org-1", "org-2"],
)
assert _enqueued(writer.spend_update_queue.add_update, Litellm_EntityType.ORGANIZATION) == [
"org-1",
"org-2",
]
@pytest.mark.asyncio
async def test_daily_team_rollup_emits_one_row_per_team():
"""LiteLLM_DailyTeamSpend is already unique per team+date+key+model, so N
teams means N distinct transaction keys and no migration."""
writer = DBSpendUpdateWriter()
writer.daily_team_spend_update_queue.add_update = AsyncMock()
writer._common_add_spend_log_transaction_to_daily_transaction = AsyncMock(
return_value={"date": "2026-08-21", "spend": 1.0, "api_requests": 1, "endpoint": "/chat/completions"}
)
payload = {
"team_id": "team-a",
"api_key": "hashed-key",
"model": "gpt-4o",
"custom_llm_provider": "openai",
}
await writer.add_spend_log_transaction_to_daily_team_transaction(
payload=payload,
prisma_client=MagicMock(),
attributed_team_ids=["team-a", "team-b"],
)
keys = [list(c.kwargs["update"].keys())[0] for c in writer.daily_team_spend_update_queue.add_update.call_args_list]
assert len(keys) == 2
assert keys[0].startswith("team-a_2026-08-21_")
assert keys[1].startswith("team-b_2026-08-21_")
team_ids = [
list(c.kwargs["update"].values())[0]["team_id"]
for c in writer.daily_team_spend_update_queue.add_update.call_args_list
]
assert team_ids == ["team-a", "team-b"]
@pytest.mark.asyncio
async def test_daily_team_rollup_single_row_when_attribution_off():
writer = DBSpendUpdateWriter()
writer.daily_team_spend_update_queue.add_update = AsyncMock()
writer._common_add_spend_log_transaction_to_daily_transaction = AsyncMock(
return_value={"date": "2026-08-21", "spend": 1.0, "api_requests": 1, "endpoint": ""}
)
await writer.add_spend_log_transaction_to_daily_team_transaction(
payload={
"team_id": "team-a",
"api_key": "hashed-key",
"model": "gpt-4o",
"custom_llm_provider": "openai",
},
prisma_client=MagicMock(),
)
assert writer.daily_team_spend_update_queue.add_update.call_count == 1
# --------------------------------------------------------------------------- #
# rate-limit fan-out
# --------------------------------------------------------------------------- #
def _limiter():
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
)
return _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock())
def test_rate_limit_descriptors_not_added_when_setting_off():
limiter = _limiter()
token = UserAPIKeyAuth(api_key="sk-1", team_id="team-a")
token.attributed_team_limits = {"team-b": {"rpm": 10, "tpm": None}}
with patch("litellm.proxy.proxy_server.general_settings", {}):
descriptors = limiter._attributed_team_rate_limit_descriptors(user_api_key_dict=token)
assert descriptors == ()
def test_rate_limit_descriptors_added_for_other_teams_only():
limiter = _limiter()
token = UserAPIKeyAuth(api_key="sk-1", team_id="team-a")
token.attributed_team_limits = {
"team-a": {"rpm": 5, "tpm": None}, # stamped team: already emitted elsewhere
"team-b": {"rpm": 10, "tpm": None},
"team-c": {"rpm": None, "tpm": None}, # no limits: nothing to enforce
"team-d": {"rpm": None, "tpm": 900},
}
with patch(
"litellm.proxy.proxy_server.general_settings",
{"enforce_rate_limits_across_all_user_teams": True},
):
descriptors = limiter._attributed_team_rate_limit_descriptors(user_api_key_dict=token)
assert [d["value"] for d in descriptors] == ["team-b", "team-d"]
# same namespace as the stamped-team descriptor, so a team shares one bucket
assert {d["key"] for d in descriptors} == {"team"}
assert descriptors[0]["rate_limit"]["requests_per_unit"] == 10
assert descriptors[1]["rate_limit"]["tokens_per_unit"] == 900
# --------------------------------------------------------------------------- #
# budget enforcement across memberships
# --------------------------------------------------------------------------- #
SPEND_ON = {"track_spend_across_all_user_teams": True}
def _proxy_logging():
from litellm.proxy.utils import ProxyLogging
obj = ProxyLogging(user_api_key_cache=None)
obj.budget_alerts = AsyncMock()
return obj
def _team(team_id, max_budget=None, spend=0.0, org="org-1"):
obj = MagicMock()
obj.team_id = team_id
obj.team_alias = f"{team_id}-alias"
obj.max_budget = max_budget
obj.spend = spend
obj.organization_id = org
return obj
def _token_with_memberships():
token = UserAPIKeyAuth(api_key="sk-1", team_id="team-a", org_id="org-1")
token.attributed_team_ids = ["team-a", "team-b"]
token.attributed_org_ids = ["org-1", "org-2"]
return token
async def _team_budget_error(token, general_settings):
"""The budget error the team gate raised, or None if the caller got through.
Asserting on this instead of on which internals ran keeps these tests about
the only thing a caller can observe: blocked, or not blocked.
"""
try:
await _attributed_teams_max_budget_check(
valid_token=token,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=_proxy_logging(),
general_settings=general_settings,
)
except litellm.BudgetExceededError as e:
return e
return None
async def _org_budget_error(token, general_settings):
"""The budget error the org gate raised, or None if the caller got through."""
try:
await _attributed_orgs_max_budget_check(
valid_token=token,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=_proxy_logging(),
general_settings=general_settings,
)
except litellm.BudgetExceededError as e:
return e
return None
def _over_budget_team(team_id, **kwargs):
"""A team whose budget is already blown, for proving a gate is inert."""
return _team(team_id, max_budget=1.0, spend=999.0)
async def _spend_always_over(counter_key, fallback_spend, max_budget=None, **kwargs):
return 999.0
@pytest.mark.asyncio
async def test_attributed_team_budget_check_noop_when_setting_off():
"""Populated memberships must not gate budgets unless the SPEND setting is on.
Enabling only rate-limit attribution must not silently add budget gates, so
a caller whose other team is wildly over budget still gets through.
"""
token = _token_with_memberships()
with (
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new=AsyncMock(side_effect=lambda team_id, **kwargs: _over_budget_team(team_id)),
),
patch("litellm.proxy.proxy_server.get_current_spend", _spend_always_over),
):
error = await _team_budget_error(token, {"enforce_rate_limits_across_all_user_teams": True})
assert error is None
@pytest.mark.asyncio
async def test_attributed_team_budget_check_raises_for_other_team():
"""A team the caller merely belongs to can now block them."""
token = _token_with_memberships()
async def _get_team(team_id, **kwargs):
return _team(team_id, max_budget=10.0)
async def _spend(counter_key, fallback_spend, max_budget=None, **kwargs):
return 25.0 if counter_key == "spend:team:team-b" else 0.0
with (
patch("litellm.proxy.auth.auth_checks.get_team_object", new=AsyncMock(side_effect=_get_team)),
patch("litellm.proxy.proxy_server.get_current_spend", _spend),
):
with pytest.raises(litellm.BudgetExceededError) as exc:
await _attributed_teams_max_budget_check(
valid_token=token,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=_proxy_logging(),
general_settings=SPEND_ON,
)
assert exc.value.entity_id == "team-b"
assert exc.value.entity_type == Litellm_EntityType.TEAM.value
assert exc.value.current_cost == 25.0
@pytest.mark.asyncio
async def test_attributed_team_budget_check_skips_stamped_team():
"""The stamped team is _team_max_budget_check's job. Checking it here too
would raise twice and fire a duplicate budget alert."""
token = _token_with_memberships()
seen = []
async def _get_team(team_id, **kwargs):
seen.append(team_id)
return _team(team_id, max_budget=100.0)
async def _spend(counter_key, fallback_spend, max_budget=None, **kwargs):
return 0.0
with (
patch("litellm.proxy.auth.auth_checks.get_team_object", new=AsyncMock(side_effect=_get_team)),
patch("litellm.proxy.proxy_server.get_current_spend", _spend),
):
await _attributed_teams_max_budget_check(
valid_token=token,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=_proxy_logging(),
general_settings=SPEND_ON,
)
assert seen == ["team-b"]
@pytest.mark.asyncio
async def test_attributed_team_budget_check_passes_under_budget():
"""A membership with room left does not block the caller."""
token = _token_with_memberships()
async def _get_team(team_id, **kwargs):
return _team(team_id, max_budget=100.0)
async def _spend(counter_key, fallback_spend, max_budget=None, **kwargs):
return 5.0
with (
patch("litellm.proxy.auth.auth_checks.get_team_object", new=AsyncMock(side_effect=_get_team)),
patch("litellm.proxy.proxy_server.get_current_spend", _spend),
):
error = await _team_budget_error(token, SPEND_ON)
assert error is None
@pytest.mark.asyncio
async def test_attributed_team_budget_check_ignores_unloadable_team():
"""A team that cannot be loaded contributes no ceiling rather than a 500.
The caller is neither blocked nor served an internal error.
"""
token = _token_with_memberships()
async def _get_team(team_id, **kwargs):
raise Exception("team row is gone")
with patch("litellm.proxy.auth.auth_checks.get_team_object", new=AsyncMock(side_effect=_get_team)):
error = await _team_budget_error(token, SPEND_ON)
assert error is None
@pytest.mark.asyncio
async def test_attributed_org_budget_check_raises_for_other_org():
"""Only reachable when the caller's teams span several organizations."""
token = _token_with_memberships()
org = MagicMock()
org.spend = 0.0
org.litellm_budget_table = MagicMock()
org.litellm_budget_table.max_budget = 50.0
async def _spend(counter_key, fallback_spend, max_budget=None, **kwargs):
return 90.0 if counter_key == "spend:org:org-2" else 0.0
with (
patch("litellm.proxy.auth.auth_checks.get_org_object", new=AsyncMock(return_value=org)),
patch("litellm.proxy.proxy_server.get_current_spend", _spend),
):
with pytest.raises(litellm.BudgetExceededError) as exc:
await _attributed_orgs_max_budget_check(
valid_token=token,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=_proxy_logging(),
general_settings=SPEND_ON,
)
assert exc.value.entity_id == "org-2"
assert exc.value.entity_type == Litellm_EntityType.ORGANIZATION.value
@pytest.mark.asyncio
async def test_attributed_org_budget_check_noop_when_setting_off():
"""An over-budget non-stamped org does not block while the setting is off."""
token = _token_with_memberships()
org = MagicMock()
org.spend = 999.0
org.litellm_budget_table = MagicMock()
org.litellm_budget_table.max_budget = 1.0
with (
patch("litellm.proxy.auth.auth_checks.get_org_object", new=AsyncMock(return_value=org)),
patch("litellm.proxy.proxy_server.get_current_spend", _spend_always_over),
):
error = await _org_budget_error(token, {})
assert error is None
# --------------------------------------------------------------------------- #
# request-metadata stamping
# --------------------------------------------------------------------------- #
def test_metadata_stamped_only_when_setting_on():
from litellm.proxy.hooks.proxy_track_cost_callback import _metadata_id_list
token = _token_with_memberships()
data_off = {"metadata": {}}
with patch("litellm.proxy.proxy_server.general_settings", {}):
LiteLLMProxyRequestSetup._add_attributed_membership_metadata(
data=data_off, user_api_key_dict=token, _metadata_variable_name="metadata"
)
assert data_off["metadata"] == {}
data_on = {"metadata": {}}
with patch("litellm.proxy.proxy_server.general_settings", SPEND_ON):
LiteLLMProxyRequestSetup._add_attributed_membership_metadata(
data=data_on, user_api_key_dict=token, _metadata_variable_name="metadata"
)
assert data_on["metadata"]["user_api_key_attributed_team_ids"] == ("team-a", "team-b")
assert data_on["metadata"]["user_api_key_attributed_org_ids"] == ("org-1", "org-2")
# and the cost callback reads back exactly what was stamped
assert _metadata_id_list(data_on["metadata"], "user_api_key_attributed_team_ids") == ("team-a", "team-b")
def test_metadata_id_list_returns_none_when_absent():
"""None, not [], so a writer can tell "attribution off" from "on, nothing
resolved"."""
from litellm.proxy.hooks.proxy_track_cost_callback import _metadata_id_list
assert _metadata_id_list({}, "user_api_key_attributed_team_ids") is None
assert _metadata_id_list({"user_api_key_attributed_team_ids": []}, "user_api_key_attributed_team_ids") is None
assert _metadata_id_list({"user_api_key_attributed_team_ids": "nope"}, "user_api_key_attributed_team_ids") is None

View file

@ -598,6 +598,10 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
budget_reservation=budget_reservation,
end_user_id="test_end_user_id",
tags=["tag-a"],
# None unless track_spend_across_all_user_teams is on, which keeps this
# call identical to the pre-attribution behavior.
attributed_team_ids=None,
attributed_org_ids=None,
)

View file

@ -47,9 +47,7 @@ def _make_spend_counter_cache(
cache.in_memory_cache.delete_cache = MagicMock()
if with_redis:
cache.redis_cache = MagicMock()
cache.redis_cache.async_get_cache = AsyncMock(
return_value=redis_get_value, side_effect=redis_get_side_effect
)
cache.redis_cache.async_get_cache = AsyncMock(return_value=redis_get_value, side_effect=redis_get_side_effect)
cache.redis_cache.async_increment = AsyncMock(
return_value=redis_increment_value,
side_effect=redis_increment_side_effect,
@ -69,9 +67,7 @@ def _make_spend_counter_cache(
def _make_user_api_key_cache(get_value=None, get_side_effect=None):
cache = MagicMock()
cache.async_get_cache = AsyncMock(
return_value=get_value, side_effect=get_side_effect
)
cache.async_get_cache = AsyncMock(return_value=get_value, side_effect=get_side_effect)
cache.async_set_cache_pipeline = AsyncMock()
return cache
@ -108,9 +104,7 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory(monkeypatch
)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
result = await ps.get_current_spend(
counter_key="spend:key:abc", fallback_spend=99.0
)
result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=99.0)
assert result == 17.0
@ -135,9 +129,7 @@ async def test_get_current_spend_floors_stale_low_counter_against_db(monkeypatch
# the stale counter is repaired up to the authoritative DB value via a
# monotonic set-max so other workers read the corrected total, and a
# concurrent increment cannot be clobbered
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(
key="spend:key:abc", value=12.0
)
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key="spend:key:abc", value=12.0)
@pytest.mark.asyncio
@ -168,9 +160,7 @@ async def test_get_current_spend_no_floor_without_max_budget(monkeypatch):
from_db = AsyncMock(return_value=12.0)
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
result = await ps.get_current_spend(
counter_key="spend:key:abc", fallback_spend=12.0
)
result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0)
assert result == 2.0
assert from_db.await_count == 0
@ -209,12 +199,8 @@ async def test_get_current_spend_floor_caches_db_read(monkeypatch):
from_db = AsyncMock(return_value=12.0)
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
first = await ps.get_current_spend(
counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0
)
second = await ps.get_current_spend(
counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0
)
first = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0)
second = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0)
assert first == 12.0
assert second == 12.0
@ -266,9 +252,7 @@ async def test_get_current_spend_floors_window_against_spend_logs(monkeypatch):
assert result == 15.0
assert wfsl.await_count == 1
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(
key=counter_key, value=15.0
)
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key=counter_key, value=15.0)
@pytest.mark.asyncio
@ -278,21 +262,13 @@ async def test_get_current_spend_fail_closed_rejects_when_unverifiable(monkeypat
rather than admitted on an unverifiable budget."""
from fastapi import HTTPException
fake_cache = _make_spend_counter_cache(
redis_get_side_effect=RuntimeError("redis down")
)
fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(
ps, "general_settings", {"fail_closed_budget_enforcement": True}
)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
with pytest.raises(HTTPException) as exc:
await ps.get_current_spend(
counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
)
await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
assert exc.value.status_code == 503
@ -300,18 +276,12 @@ async def test_get_current_spend_fail_closed_rejects_when_unverifiable(monkeypat
async def test_get_current_spend_fail_closed_off_admits_when_unverifiable(monkeypatch):
"""Default (flag off): an unverifiable read keeps the existing behavior and
admits using the cached fallback — no new rejection."""
fake_cache = _make_spend_counter_cache(
redis_get_side_effect=RuntimeError("redis down")
)
fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "general_settings", {})
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
result = await ps.get_current_spend(
counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
)
result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
assert result == 1.0
@ -321,13 +291,9 @@ async def test_get_current_spend_fail_closed_admits_when_redis_verified(monkeypa
authoritative, so an under-budget request is admitted normally."""
fake_cache = _make_spend_counter_cache(redis_get_value=1.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(
ps, "general_settings", {"fail_closed_budget_enforcement": True}
)
monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
result = await ps.get_current_spend(
counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
)
result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
assert result == 1.0
@ -336,16 +302,10 @@ async def test_get_current_spend_fail_closed_allows_authoritative_fallback(monke
"""End-user/tag callers pass fallback_authoritative=True (their spend is
loaded fresh from the DB in auth), so fail-closed does not reject them even
when the counter path is unreadable."""
fake_cache = _make_spend_counter_cache(
redis_get_side_effect=RuntimeError("redis down")
)
fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(
ps, "general_settings", {"fail_closed_budget_enforcement": True}
)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
result = await ps.get_current_spend(
counter_key="spend:end_user:e1",
@ -363,9 +323,7 @@ async def test_get_current_spend_strict_floors_when_fallback_also_stale(monkeypa
re-checks the authoritative DB and enforces against it."""
fake_cache = _make_spend_counter_cache(redis_get_value=0.00001)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(
ps, "general_settings", {"fail_closed_budget_enforcement": True}
)
monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
from_db = AsyncMock(return_value=0.5)
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
@ -387,9 +345,7 @@ async def test_get_current_spend_strict_floors_when_fallback_also_stale(monkeypa
@pytest.mark.asyncio
async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
fake_cache = _make_spend_counter_cache(
redis_get_value=None, redis_increment_value=5.0
)
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=5.0)
fake_user_cache = _make_user_api_key_cache(get_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
@ -398,9 +354,7 @@ async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
async def _fake_coalesced(**kwargs):
return None
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced))
await ps.increment_spend_counters(
token="hashed-tok",
@ -421,6 +375,83 @@ async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
}
def _incremented_keys(fake_cache) -> set[str]:
"""Counter keys passed to the Redis increment, however they were spelled."""
keys = set()
for call in fake_cache.redis_cache.async_increment.call_args_list:
key = call.kwargs.get("key")
if key is None and call.args:
key = call.args[0]
if isinstance(key, str):
keys.add(key)
return keys
@pytest.mark.asyncio
async def test_increment_spend_counters_fans_out_across_attributed_teams(monkeypatch):
"""With track_spend_across_all_user_teams on, every attributed team gets its
own team and team-member counter, so budget gates reading Redis see the
spend against all of them immediately.
Baseline (single team) is 4 increments: key, team, team_member, user. Three
attributed teams makes it 8: key, 3 x team, 3 x team_member, user.
"""
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=5.0)
fake_user_cache = _make_user_api_key_cache(get_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
async def _fake_coalesced(**kwargs):
return None
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced))
await ps.increment_spend_counters(
token="hashed-tok",
team_id="t1",
user_id="u1",
response_cost=5.0,
attributed_team_ids=["t1", "t2", "t3"],
)
incremented_keys = _incremented_keys(fake_cache)
assert {"spend:team:t1", "spend:team:t2", "spend:team:t3"} <= incremented_keys
assert {
"spend:team_member:u1:t1",
"spend:team_member:u1:t2",
"spend:team_member:u1:t3",
} <= incremented_keys
assert fake_cache.redis_cache.async_increment.call_count == 8
@pytest.mark.asyncio
async def test_increment_spend_counters_single_team_when_attribution_off(monkeypatch):
"""Regression guard: without the setting, only the stamped team is charged."""
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=5.0)
fake_user_cache = _make_user_api_key_cache(get_value=None)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
async def _fake_coalesced(**kwargs):
return None
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced))
await ps.increment_spend_counters(
token="hashed-tok",
team_id="t1",
user_id="u1",
response_cost=5.0,
)
incremented_keys = _incremented_keys(fake_cache)
assert "spend:team:t1" in incremented_keys
assert not any(k.startswith("spend:team:t2") for k in incremented_keys)
assert fake_cache.redis_cache.async_increment.call_count == 4
class _ConcurrencyProbe:
"""Stand-in for redis_cache.async_increment that pins concurrency.
@ -463,9 +494,7 @@ async def test_increment_spend_counters_runs_scopes_concurrently(monkeypatch):
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
task = asyncio.create_task(
ps.increment_spend_counters(
@ -514,9 +543,7 @@ async def test_increment_spend_counters_skips_reserved_counter_keys(monkeypatch)
import litellm.proxy.spend_tracking.budget_reservation as br
reserved = {"spend:key:hashed-tok", "spend:org:org1"}
monkeypatch.setattr(
br, "get_reserved_counter_keys", MagicMock(return_value=set(reserved))
)
monkeypatch.setattr(br, "get_reserved_counter_keys", MagicMock(return_value=set(reserved)))
monkeypatch.setattr(br, "reconcile_budget_reservation", AsyncMock())
recorded: dict[str, float] = {}
@ -531,9 +558,7 @@ async def test_increment_spend_counters_skips_reserved_counter_keys(monkeypatch)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
reservation = {"finalized": False}
await ps.increment_spend_counters(
@ -578,9 +603,7 @@ async def test_increment_spend_counters_failing_scope_propagates_after_siblings_
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
reservation = {"finalized": False}
with pytest.raises(RuntimeError, match="redis increment failed"):
@ -636,9 +659,7 @@ async def test_increment_spend_counters_zero_cost_is_noop_finalizes_reservation(
@pytest.mark.asyncio
async def test_reconcile_budget_reservation_for_counter_update_returns_empty_set_when_none():
result = await ps._reconcile_budget_reservation_for_counter_update(
budget_reservation=None, response_cost=1.0
)
result = await ps._reconcile_budget_reservation_for_counter_update(budget_reservation=None, response_cost=1.0)
assert result == set()
@ -681,16 +702,12 @@ async def test_reconcile_budget_reservation_for_counter_update_failure_invalidat
async def test_increment_end_user_and_tag_spend_counters_increments_each_unique_tag(
monkeypatch,
):
fake_cache = _make_spend_counter_cache(
redis_get_value=None, redis_increment_value=3.0
)
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=3.0)
fake_user_cache = _make_user_api_key_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
await ps._increment_end_user_and_tag_spend_counters(
end_user_id="eu1",
@ -735,16 +752,12 @@ async def test_increment_end_user_and_tag_spend_counters_no_end_user_no_tags_inv
@pytest.mark.asyncio
async def test_increment_org_spend_counter_increments_when_org_present(monkeypatch):
fake_cache = _make_spend_counter_cache(
redis_get_value=None, redis_increment_value=10.0
)
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=10.0)
fake_user_cache = _make_user_api_key_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
await ps._increment_org_spend_counter(
org_id="org-1",
@ -755,9 +768,7 @@ async def test_increment_org_spend_counter_increments_when_org_present(monkeypat
observed = {
"increment_called": fake_cache.redis_cache.async_increment.called,
"increment_calls": fake_cache.redis_cache.async_increment.call_count,
"counter_key_arg": fake_cache.redis_cache.async_increment.call_args.kwargs[
"key"
],
"counter_key_arg": fake_cache.redis_cache.async_increment.call_args.kwargs["key"],
}
assert normalize(observed) == {
"increment_called": True,
@ -806,16 +817,12 @@ async def test_init_and_increment_unreserved_spend_counter_skips_reserved_keys(
async def test_init_and_increment_unreserved_spend_counter_proceeds_when_not_reserved(
monkeypatch,
):
fake_cache = _make_spend_counter_cache(
redis_get_value=None, redis_increment_value=2.0
)
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=2.0)
fake_user_cache = _make_user_api_key_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
await ps._init_and_increment_unreserved_spend_counter(
counter_key="spend:tag:y",
@ -843,9 +850,7 @@ async def test_init_and_increment_unreserved_spend_counter_proceeds_when_not_res
@pytest.mark.asyncio
async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypatch):
fake_cache = _make_spend_counter_cache(
redis_get_value=11.0, redis_increment_value=14.0
)
fake_cache = _make_spend_counter_cache(redis_get_value=11.0, redis_increment_value=14.0)
fake_user_cache = _make_user_api_key_cache()
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
@ -880,9 +885,7 @@ async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypa
async def test_init_and_increment_window_spend_counter_increments_when_initialized(
monkeypatch,
):
fake_cache = _make_spend_counter_cache(
redis_get_value=0.0, redis_increment_value=5.0
)
fake_cache = _make_spend_counter_cache(redis_get_value=0.0, redis_increment_value=5.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
@ -967,16 +970,12 @@ async def test_ensure_spend_counter_initialized_warm_skips_reseed_and_source(
async def test_ensure_spend_counter_initialized_cold_seeds_from_source_cache(
monkeypatch,
):
fake_cache = _make_spend_counter_cache(
redis_get_value=None, redis_increment_value=7.0
)
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=7.0)
fake_user_cache = _make_user_api_key_cache(get_value={"spend": 7.0})
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
)
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
await ps._ensure_spend_counter_initialized(
counter_key="spend:user:u",
@ -1016,9 +1015,7 @@ async def test_get_source_cache_base_spend_reads_first_hit_from_list(monkeypatch
fake_user_cache.async_get_cache = AsyncMock(side_effect=_get)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
result = await ps._get_source_cache_base_spend(
source_cache_key=["miss", "hit-obj", "miss2"]
)
result = await ps._get_source_cache_base_spend(source_cache_key=["miss", "hit-obj", "miss2"])
observed = {
"result": result,
@ -1145,9 +1142,7 @@ async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkey
fake_cache = _make_spend_counter_cache(redis_increment_value=44.0)
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
result = await ps._increment_spend_counter_cache(
counter_key="spend:key:k", increment=4.0
)
result = await ps._increment_spend_counter_cache(counter_key="spend:key:k", increment=4.0)
observed = {
"result": result,
@ -1165,15 +1160,11 @@ async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkey
async def test_increment_spend_counter_cache_redis_error_raises_and_invalidates(
monkeypatch,
):
fake_cache = _make_spend_counter_cache(
redis_increment_side_effect=RuntimeError("incr fail")
)
fake_cache = _make_spend_counter_cache(redis_increment_side_effect=RuntimeError("incr fail"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
with pytest.raises(RuntimeError):
await ps._increment_spend_counter_cache(
counter_key="spend:key:k", increment=1.0
)
await ps._increment_spend_counter_cache(counter_key="spend:key:k", increment=1.0)
assert fake_cache.in_memory_cache.delete_cache.called is True
assert fake_cache.redis_cache.async_delete_cache.called is True
@ -1194,9 +1185,7 @@ async def test_invalidate_spend_counter_deletes_in_memory_and_redis(monkeypatch)
observed = {
"in_memory_delete_called": fake_cache.in_memory_cache.delete_cache.called,
"redis_delete_called": fake_cache.redis_cache.async_delete_cache.called,
"delete_args_key": fake_cache.redis_cache.async_delete_cache.call_args.kwargs[
"key"
],
"delete_args_key": fake_cache.redis_cache.async_delete_cache.call_args.kwargs["key"],
}
assert normalize(observed) == {
"in_memory_delete_called": True,
@ -1208,9 +1197,7 @@ async def test_invalidate_spend_counter_deletes_in_memory_and_redis(monkeypatch)
@pytest.mark.asyncio
async def test_invalidate_spend_counter_swallows_redis_failure_no_raise(monkeypatch):
fake_cache = _make_spend_counter_cache()
fake_cache.redis_cache.async_delete_cache = AsyncMock(
side_effect=RuntimeError("redis down")
)
fake_cache.redis_cache.async_delete_cache = AsyncMock(side_effect=RuntimeError("redis down"))
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
await ps._invalidate_spend_counter(counter_key="spend:key:k")

View file

@ -924,6 +924,8 @@ async def test_api_key_preserved_through_failure_hook_to_database():
start_time,
end_time,
org_id,
attributed_team_ids=None,
attributed_org_ids=None,
):
"""Mock update_database and capture the payload it creates"""
from litellm.proxy.spend_tracking.spend_tracking_utils import (

View file

@ -24177,6 +24177,11 @@ export interface components {
* @default false
*/
enable_public_model_hub: boolean;
/**
* Enforce Rate Limits Across All User Teams
* @description apply the RPM/TPM limits of EVERY team the calling user belongs to, not only the team stamped on the virtual key. The caller's effective limit becomes the minimum across their memberships, so a busy team can throttle someone who is mostly working for a different team. Separate from track_spend_across_all_user_teams so spend attribution can be adopted without this. Default off.
*/
enforce_rate_limits_across_all_user_teams?: boolean | null;
/**
* Forward Client Headers To Llm Api
* @description If True, forwards client headers (e.g. Authorization) to the LLM API. Required for Claude Code with Max subscription.
@ -24349,6 +24354,11 @@ export interface components {
* @description Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map', 'tools', 'config_overrides'. If not set, all objects are loaded (default behavior).
*/
supported_db_objects?: components["schemas"]["SupportedDBObjectType"][] | null;
/**
* Track Spend Across All User Teams
* @description attribute each request's spend to EVERY team the calling user belongs to (and every organization reached through those teams), not only the team stamped on the virtual key. Budget gates expand to match, so one over-budget team blocks the caller everywhere. Summing team spend then exceeds real spend by design, because one request is charged to several teams; key, user, and org totals stay single-counted. Default off.
*/
track_spend_across_all_user_teams?: boolean | null;
/**
* Trusted Proxy Ranges
* @description CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler.