diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index a886a754f5e..87eca63d6cd 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -1092,6 +1092,7 @@ router_settings: | SPEND_LOG_CLEANUP_BATCH_SIZE | Number of logs deleted per batch during cleanup. Default is 1000 | SPEND_LOG_QUEUE_POLL_INTERVAL | Polling interval in seconds for spend log queue. Default is 2.0 | SPEND_LOG_QUEUE_SIZE_THRESHOLD | Threshold for spend log queue size before processing. Default is 100 +| SPEND_COUNTER_REDIS_TTL_SECONDS | TTL in seconds for cross-pod Redis spend counters used by budget enforcement. Default is 300 (5 minutes); on counter expiry the value is reseeded from the authoritative DB spend. | COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY | Maximum size for CoroutineChecker in-memory cache. Default is 1000 | DEFAULT_SHARED_HEALTH_CHECK_TTL | Time-to-live in seconds for cached health check results in shared health check mode. Default is 300 (5 minutes) | DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL | Time-to-live in seconds for health check lock in shared health check mode. Default is 60 (1 minute) diff --git a/litellm/constants.py b/litellm/constants.py index 6c89cf5946d..0ccdc3dfa0b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -274,6 +274,9 @@ REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY = ( REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_agent_spend_update_buffer" REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_tag_spend_update_buffer" MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", 100)) +SPEND_COUNTER_REDIS_TTL_SECONDS = int( + os.getenv("SPEND_COUNTER_REDIS_TTL_SECONDS", 5 * 60) +) # Bounds asyncio.Queue() instances (log queues, spend update queues, etc.) to prevent unbounded memory growth LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000)) TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60)) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index d245ec53ece..0315c1bc5b0 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -899,6 +899,63 @@ async def get_default_end_user_budget( return None +@log_db_metrics +async def get_team_member_default_budget( + budget_id: str, + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, +) -> Optional[LiteLLM_BudgetTable]: + """ + Fetches the team-level default per-member budget referenced by team.metadata["team_member_budget_id"]. + + This budget is applied to team members whose TeamMembership row has no + linked budget. Results are cached for performance. + + Args: + budget_id: The budget_id pulled from team.metadata["team_member_budget_id"] + prisma_client: Database client instance + user_api_key_cache: Cache for storing/retrieving budget data + + Returns: + LiteLLM_BudgetTable if found, None otherwise + """ + if prisma_client is None: + return None + + cache_key = f"team_member_default_budget:{budget_id}" + + cached_budget = await user_api_key_cache.async_get_cache(key=cache_key) + if isinstance(cached_budget, LiteLLM_BudgetTable): + return cached_budget + if isinstance(cached_budget, dict): + return LiteLLM_BudgetTable(**cached_budget) + + try: + budget_record = await prisma_client.db.litellm_budgettable.find_unique( + where={"budget_id": budget_id} + ) + + if budget_record is None: + verbose_proxy_logger.warning( + f"Team-default member budget not found in database: {budget_id}" + ) + return None + + await user_api_key_cache.async_set_cache( + key=cache_key, + value=budget_record.dict(), + ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, + ) + + return LiteLLM_BudgetTable(**budget_record.dict()) + + except Exception: + verbose_proxy_logger.exception( + f"Error fetching team-default member budget {budget_id}" + ) + return None + + async def _apply_default_budget_to_end_user( end_user_obj: LiteLLM_EndUserTable, prisma_client: PrismaClient, @@ -3225,13 +3282,26 @@ async def _check_team_member_budget( proxy_logging_obj=proxy_logging_obj, ) - if ( - team_membership is not None - and team_membership.litellm_budget_table is not None - and team_membership.litellm_budget_table.max_budget is not None - ): + # Per-member override wins; otherwise fall back to the team-level + # default configured via team.metadata["team_member_budget_id"]. + team_member_budget: Optional[float] = None + if team_membership is not None and team_membership.litellm_budget_table is not None: team_member_budget = team_membership.litellm_budget_table.max_budget - team_member_spend = team_membership.spend or 0.0 + else: + default_budget_id = (team_object.metadata or {}).get("team_member_budget_id") + if isinstance(default_budget_id, str): + default_budget = await get_team_member_default_budget( + budget_id=default_budget_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + if default_budget is not None: + team_member_budget = default_budget.max_budget + + if team_member_budget is not None: + team_member_spend = ( + team_membership.spend if team_membership is not None else 0.0 + ) or 0.0 # Read from cross-pod counter (Redis-first) if available from litellm.proxy.proxy_server import get_current_spend diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 8e21b851857..45194896582 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -302,14 +302,15 @@ class TeamMemberBudgetHandler: prisma_client: PrismaClient, ) -> None: """ - Create team_memberships entries for existing members that don't have one. + Ensure every team member has a TeamMembership row linked to the + team_member_budget. - Called after team_member_budget is set/updated on a team to ensure - members who joined before the budget was configured also get budget - enforcement. - - Only creates missing entries — does not touch existing memberships - (which may carry individual per-member budgets). + Called after team_member_budget is set/updated on a team. Creates + rows for members who don't have one, and populates budget_id on + existing rows where it is NULL. Rows with a non-NULL budget_id + are left untouched, which preserves per-member overrides but also + means rows pointing to a prior team-default budget_id are not + migrated to the new one. """ if not members_with_roles: return @@ -347,6 +348,21 @@ class TeamMemberBudgetHandler: _sanitize_for_log(team_member_budget_id), ) + # Heal existing membership rows that predate the team_member_budget + # configuration: populate budget_id where it is currently NULL. + # Rows with an explicit budget_id (per-member override) are left alone. + updated = await prisma_client.db.litellm_teammembership.update_many( + where={"team_id": team_id, "budget_id": None}, + data={"budget_id": team_member_budget_id}, + ) + if updated: + verbose_proxy_logger.info( + "Populated budget_id on %d existing team_memberships for team %s with budget %s", + updated, + _sanitize_for_log(team_id), + _sanitize_for_log(team_member_budget_id), + ) + def _get_default_team_param(field: str) -> Any: """ diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d8354a798b1..a1a7669ef37 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -55,6 +55,7 @@ from litellm.constants import ( LITELLM_UI_ALLOW_HEADERS, LITELLM_UI_SESSION_DURATION, DAILY_TAG_SPEND_BATCH_MULTIPLIER, + SPEND_COUNTER_REDIS_TTL_SECONDS, ) from litellm.litellm_core_utils.litellm_logging import ( _init_custom_logger_compatible_class, @@ -1845,6 +1846,7 @@ async def increment_spend_counters( await spend_counter_cache.async_increment_cache( key=f"spend:key:{hashed_token}:window:{duration}", value=response_cost, + ttl=SPEND_COUNTER_REDIS_TTL_SECONDS, ) if team_id is not None: @@ -1872,6 +1874,7 @@ async def increment_spend_counters( await spend_counter_cache.async_increment_cache( key=f"spend:team:{team_id}:window:{duration}", value=response_cost, + ttl=SPEND_COUNTER_REDIS_TTL_SECONDS, ) if user_id is not None and team_id is not None: @@ -1882,40 +1885,95 @@ async def increment_spend_counters( ) +async def _reseed_spend_from_db(counter_key: str) -> float: + """ + Read the authoritative spend for a missing counter from the DB. The + counter_key prefix encodes the table to query: + + spend:key:{token} -> LiteLLM_VerificationToken.spend + spend:team:{team_id} -> LiteLLM_TeamTable.spend + spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend + + Returns 0.0 if prisma is unavailable, the row is missing, or the + key format is unrecognized. On failure, logs and returns 0.0 rather + than raising so the caller can still record the current increment. + """ + if prisma_client is None: + return 0.0 + try: + if counter_key.startswith("spend:key:"): + token = counter_key[len("spend:key:") :] + row = await prisma_client.db.litellm_verificationtoken.find_unique( + where={"token": token} + ) + elif counter_key.startswith("spend:team_member:"): + suffix = counter_key[len("spend:team_member:") :] + if ":" not in suffix: + return 0.0 + user_id, team_id = suffix.rsplit(":", 1) + row = await prisma_client.db.litellm_teammembership.find_unique( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}} + ) + elif counter_key.startswith("spend:team:"): + team_id = counter_key[len("spend:team:") :] + row = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) + else: + return 0.0 + except Exception: + verbose_proxy_logger.exception( + "Failed to reseed spend counter %s from DB", counter_key + ) + return 0.0 + if row is None: + return 0.0 + return float(getattr(row, "spend", 0.0) or 0.0) + + async def _init_and_increment_spend_counter( counter_key: str, source_cache_key: str, increment: float, ): """ - Initialize counter from cached object's DB-loaded spend if not yet set, - then atomically increment in both in-memory and Redis. + Initialize counter from the authoritative DB spend value if not yet + set, then atomically increment in both in-memory and Redis. On first access per pod: - 1. Check spend_counter_cache (in-memory -> Redis via DualCache for init check) - 2. If not found anywhere, read base spend from user_api_key_cache (DB-loaded object) + 1. Check spend_counter_cache (in-memory -> Redis via DualCache) + 2. If not found, reseed from the DB (`_reseed_spend_from_db`). Falls + back to the cached object's `.spend` via user_api_key_cache only + if prisma is unavailable, since that value can lag the flusher. 3. Seed counter via async_increment_cache (not async_set_cache) to avoid a check-then-set race: if two pods cold-start simultaneously, both may see - the counter as absent and seed it. Using increment instead of set means - the worst case is over-counting (conservative — blocks slightly early) - rather than under-counting (would allow overspend). + the counter as absent and seed it. Using increment means the worst case + is over-counting (conservative, blocks slightly early) rather than + under-counting (would allow overspend). 4. Increment atomically (both in-memory + Redis) """ current = await spend_counter_cache.async_get_cache(key=counter_key) if current is None: - source = await user_api_key_cache.async_get_cache(key=source_cache_key) - base_spend = 0.0 - if source is not None: - if isinstance(source, dict): - base_spend = source.get("spend", 0.0) or 0.0 - else: - base_spend = getattr(source, "spend", 0.0) or 0.0 + base_spend = await _reseed_spend_from_db(counter_key) + if prisma_client is None: + # Best-effort fallback when prisma is unavailable (tests or + # early-startup paths). May be stale but avoids resetting to 0. + source = await user_api_key_cache.async_get_cache(key=source_cache_key) + if source is not None: + if isinstance(source, dict): + base_spend = source.get("spend", 0.0) or 0.0 + else: + base_spend = getattr(source, "spend", 0.0) or 0.0 if base_spend > 0: await spend_counter_cache.async_increment_cache( - key=counter_key, value=base_spend + key=counter_key, + value=base_spend, + ttl=SPEND_COUNTER_REDIS_TTL_SECONDS, ) - await spend_counter_cache.async_increment_cache(key=counter_key, value=increment) + await spend_counter_cache.async_increment_cache( + key=counter_key, value=increment, ttl=SPEND_COUNTER_REDIS_TTL_SECONDS + ) async def update_cache( # noqa: PLR0915 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 19fffffc65b..8612d243c41 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -2126,3 +2126,192 @@ class TestGuardrailModificationCheck: """Unparseable strings should not trigger a 403 — they have no keys.""" self._call({"metadata": "not-json"}) self._call({"metadata": '"just a string"'}) + + +@pytest.mark.asyncio +async def test_team_member_budget_check_falls_back_to_team_default_budget_id(): + """When a member's TeamMembership has no linked budget row, the check + should fall back to team.metadata["team_member_budget_id"] and still + enforce the cap. Pre-fix, this path silently skipped enforcement.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.utils import ProxyLogging + + team_object = LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_budget_id": "budget-default"}, + ) + user_object = LiteLLM_UserTable(user_id="test-user") + valid_token = UserAPIKeyAuth( + token="test-token", + user_id="test-user", + team_id="test-team", + ) + + # Membership row without an attached budget. + team_membership = LiteLLM_TeamMembership( + user_id="test-user", + team_id="test-team", + spend=0.0, + budget_id=None, + litellm_budget_table=None, + ) + + proxy_logging_obj = ProxyLogging(user_api_key_cache=None) + + fake_budget_row = MagicMock() + fake_budget_row.max_budget = 50.0 + fake_budget_row.dict = MagicMock( + return_value={"budget_id": "budget-default", "max_budget": 50.0} + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_budgettable.find_unique = AsyncMock( + return_value=fake_budget_row + ) + + async def mock_get_current_spend(counter_key, fallback_spend): + if counter_key == "spend:team_member:test-user:test-team": + return 70.0 + return fallback_spend + + user_api_key_cache = DualCache() + + with ( + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new_callable=AsyncMock, + return_value=team_membership, + ), + ): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _check_team_member_budget( + team_object=team_object, + user_object=user_object, + valid_token=valid_token, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert exc_info.value.current_cost == 70.0 + assert exc_info.value.max_budget == 50.0 + + # First call did perform the fallback DB lookup. + prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once() + + # Second call hits the cached budget row, no additional prisma read. + prisma_client.db.litellm_budgettable.find_unique.reset_mock() + with ( + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new_callable=AsyncMock, + return_value=team_membership, + ), + ): + with pytest.raises(litellm.BudgetExceededError) as second_exc_info: + await _check_team_member_budget( + team_object=team_object, + user_object=user_object, + valid_token=valid_token, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + # The cached $50 cap is still being applied (not a coincidental skip) + assert second_exc_info.value.current_cost == 70.0 + assert second_exc_info.value.max_budget == 50.0 + prisma_client.db.litellm_budgettable.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_team_member_budget_check_per_member_override_wins_over_team_default(): + """If a member has a per-member budget AND the team carries a + team_member_budget_id default, the per-member value wins and the + fallback prisma lookup is never performed.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership + from litellm.proxy.utils import ProxyLogging + + team_object = LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_budget_id": "budget-default"}, + ) + user_object = LiteLLM_UserTable(user_id="test-user") + valid_token = UserAPIKeyAuth( + token="test-token", + user_id="test-user", + team_id="test-team", + ) + + team_membership = LiteLLM_TeamMembership( + user_id="test-user", + team_id="test-team", + spend=0.0, + budget_id="budget-override", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=200.0), + ) + + proxy_logging_obj = ProxyLogging(user_api_key_cache=None) + + # Team-default row resolves to $50. If the fallback fired (it must + # not here), spend $70 would exceed that $50 cap and raise. + fake_budget_row = MagicMock() + fake_budget_row.max_budget = 50.0 + + prisma_client = MagicMock() + prisma_client.db.litellm_budgettable.find_unique = AsyncMock( + return_value=fake_budget_row + ) + + mocked_spend = 70.0 + + async def mock_get_current_spend(counter_key, fallback_spend): + if counter_key == "spend:team_member:test-user:test-team": + return mocked_spend + return fallback_spend + + # 1. spend ($70) < per-member cap ($200) → no raise, no fallback lookup. + with ( + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new_callable=AsyncMock, + return_value=team_membership, + ), + ): + await _check_team_member_budget( + team_object=team_object, + user_object=user_object, + valid_token=valid_token, + prisma_client=prisma_client, + user_api_key_cache=DualCache(), + proxy_logging_obj=proxy_logging_obj, + ) + + prisma_client.db.litellm_budgettable.find_unique.assert_not_awaited() + + # 2. Now push spend above the per-member cap ($200). Must raise with + # max_budget=200 to prove the per-member cap is the value being + # enforced (not just that enforcement silently skipped). + mocked_spend = 250.0 + with ( + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new_callable=AsyncMock, + return_value=team_membership, + ), + ): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _check_team_member_budget( + team_object=team_object, + user_object=user_object, + valid_token=valid_token, + prisma_client=prisma_client, + user_api_key_cache=DualCache(), + proxy_logging_obj=proxy_logging_obj, + ) + assert exc_info.value.current_cost == 250.0 + assert exc_info.value.max_budget == 200.0 diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 8da0ef19f81..65187fb52dc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -1795,6 +1795,7 @@ async def test_backfill_team_member_budget_entries_creates_missing_memberships() return_value=[existing_membership] ) mock_prisma.db.litellm_teammembership.create_many = AsyncMock(return_value=None) + mock_prisma.db.litellm_teammembership.update_many = AsyncMock(return_value=0) # Test with Member instances members = [ @@ -1823,6 +1824,7 @@ async def test_backfill_team_member_budget_entries_creates_missing_memberships() # Also test with raw dicts (members_with_roles may be dicts when deserialized from DB) mock_prisma.db.litellm_teammembership.find_many.reset_mock() mock_prisma.db.litellm_teammembership.create_many.reset_mock() + mock_prisma.db.litellm_teammembership.update_many.reset_mock() members_as_dicts = [ {"user_id": "user-A", "role": "user"}, @@ -1868,6 +1870,7 @@ async def test_backfill_team_member_budget_entries_no_op_when_all_exist(): return_value=[existing_a, existing_b] ) mock_prisma.db.litellm_teammembership.create_many = AsyncMock(return_value=None) + mock_prisma.db.litellm_teammembership.update_many = AsyncMock(return_value=0) members = [ Member(user_id="user-A", role="user"), @@ -1884,6 +1887,55 @@ async def test_backfill_team_member_budget_entries_no_op_when_all_exist(): mock_prisma.db.litellm_teammembership.create_many.assert_not_awaited() +@pytest.mark.asyncio +async def test_backfill_team_member_budget_entries_populates_null_budget_id_on_existing_rows(): + """ + backfill_team_member_budget_entries should populate budget_id on + existing TeamMembership rows where it is currently NULL, so admins + can configure a team member budget after members have already joined + and have enforcement apply to those pre-existing members. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import Member + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) + + team_id = "team-abc" + budget_id = "budget-xyz" + + # Both members already have rows, so create_many must not fire; + # update_many must fire with the NULL-budget_id filter. + existing_a = MagicMock() + existing_a.user_id = "user-A" + existing_b = MagicMock() + existing_b.user_id = "user-B" + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teammembership.find_many = AsyncMock( + return_value=[existing_a, existing_b] + ) + mock_prisma.db.litellm_teammembership.create_many = AsyncMock(return_value=None) + mock_prisma.db.litellm_teammembership.update_many = AsyncMock(return_value=2) + + await TeamMemberBudgetHandler.backfill_team_member_budget_entries( + team_id=team_id, + members_with_roles=[ + Member(user_id="user-A", role="user"), + Member(user_id="user-B", role="user"), + ], + team_member_budget_id=budget_id, + prisma_client=mock_prisma, + ) + + mock_prisma.db.litellm_teammembership.create_many.assert_not_awaited() + mock_prisma.db.litellm_teammembership.update_many.assert_awaited_once_with( + where={"team_id": team_id, "budget_id": None}, + data={"budget_id": budget_id}, + ) + + @pytest.mark.asyncio async def test_backfill_team_member_budget_entries_empty_members(): """ diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 79eba81dc40..fe6809b0445 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -4965,3 +4965,160 @@ async def test_increment_spend_counters_team_and_member(): finally: ps.user_api_key_cache = original_key_cache ps.spend_counter_cache = original_counter_cache + + +@pytest.mark.asyncio +async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(): + """When the Redis counter is missing, the reseed path reads the + authoritative spend from the DB (not a stale cache), so the next + increment continues from the correct base value.""" + from litellm.caching.dual_cache import DualCache + + counter_cache = DualCache() + recorded_increments: list = [] + + async def record_increment(key, value, ttl=None, **kwargs): + recorded_increments.append({"key": key, "value": value, "ttl": ttl}) + return value + + fake_redis = AsyncMock() + fake_redis.async_increment = AsyncMock(side_effect=record_increment) + fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing + counter_cache.redis_cache = fake_redis + + # Prisma returns spend=42.0 (authoritative) while the stale cached + # value (would be read only if prisma is None) is 10.0. The counter + # must seed from 42, not 10. + db_row = MagicMock() + db_row.spend = 42.0 + fake_prisma = MagicMock() + fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row) + + stale_cache = DualCache() + stale_team = MagicMock() + stale_team.spend = 10.0 + stale_cache.in_memory_cache.set_cache(key="team_id:team-9", value=stale_team) + + import litellm.proxy.proxy_server as ps + from litellm.proxy.proxy_server import _init_and_increment_spend_counter + + orig_user, orig_counter, orig_prisma = ( + ps.user_api_key_cache, + ps.spend_counter_cache, + ps.prisma_client, + ) + ps.user_api_key_cache = stale_cache + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + try: + await _init_and_increment_spend_counter( + counter_key="spend:team:team-9", + source_cache_key="team_id:team-9", + increment=1.5, + ) + + fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with( + where={"team_id": "team-9"} + ) + # Two increments keyed on the counter: seed ($42) then request ($1.50). + writes = [(c["key"], c["value"]) for c in recorded_increments] + assert ("spend:team:team-9", 42.0) in writes + assert ("spend:team:team-9", 1.5) in writes + finally: + ps.user_api_key_cache = orig_user + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + +@pytest.mark.asyncio +async def test_increment_spend_counters_uses_long_ttl_on_redis_writes(): + """Every counter write (key, team, team_member) must carry the long + TTL. Redis's 60s default would expire counters mid-cycle and re-seed + from stale cached spend, causing budget bypass.""" + from litellm.caching.dual_cache import DualCache + from litellm.constants import SPEND_COUNTER_REDIS_TTL_SECONDS + from litellm.proxy._types import ( + LiteLLM_TeamTable, + LiteLLM_VerificationTokenView, + hash_token, + ) + + key_cache = DualCache() + counter_cache = DualCache() + + hashed_token = hash_token("sk-ttl-test-token") + # Seed budget_limits so the per-window counter paths also fire. + key_cache.in_memory_cache.set_cache( + key=hashed_token, + value=LiteLLM_VerificationTokenView( + token=hashed_token, + spend=0.0, + max_budget=10.0, + budget_limits=[{"budget_duration": "1h", "max_budget": 1.0}], + ), + ) + # Seed team + membership so increment_spend_counters exercises + # all counter paths (key, team, team_member, plus window variants). + key_cache.in_memory_cache.set_cache( + key="team_id:team-1", + value=LiteLLM_TeamTable( + team_id="team-1", + spend=0.0, + budget_limits=[{"budget_duration": "1d", "max_budget": 5.0}], + ), + ) + key_cache.in_memory_cache.set_cache( + key="team_membership:user-1:team-1", + value={"user_id": "user-1", "team_id": "team-1", "spend": 0.0}, + ) + + recorded_writes: list = [] + + async def record_increment(key, value, ttl=None, **kwargs): + recorded_writes.append({"op": "increment", "key": key, "ttl": ttl}) + return value + + async def record_set(key, value, **kwargs): + recorded_writes.append({"op": "set", "key": key, "ttl": kwargs.get("ttl")}) + + fake_redis = AsyncMock() + fake_redis.async_increment = AsyncMock(side_effect=record_increment) + fake_redis.async_set_cache = AsyncMock(side_effect=record_set) + fake_redis.async_get_cache = AsyncMock(return_value=None) + counter_cache.redis_cache = fake_redis + + import litellm.proxy.proxy_server as ps + + original_key_cache = ps.user_api_key_cache + original_counter_cache = ps.spend_counter_cache + ps.user_api_key_cache = key_cache + ps.spend_counter_cache = counter_cache + + try: + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=hashed_token, + team_id="team-1", + user_id="user-1", + response_cost=0.05, + ) + + # Every write carries the long TTL. + for call in recorded_writes: + assert call["ttl"] == SPEND_COUNTER_REDIS_TTL_SECONDS, ( + f"expected ttl={SPEND_COUNTER_REDIS_TTL_SECONDS} on every " + f"counter write, got ttl={call['ttl']} " + f"on op={call['op']} key={call['key']}" + ) + + # All three primary counter keys plus the per-window counters were touched. + keys_written = {c["key"] for c in recorded_writes} + assert f"spend:key:{hashed_token}" in keys_written + assert "spend:team:team-1" in keys_written + assert "spend:team_member:user-1:team-1" in keys_written + assert f"spend:key:{hashed_token}:window:1h" in keys_written + assert "spend:team:team-1:window:1d" in keys_written + finally: + ps.user_api_key_cache = original_key_cache + ps.spend_counter_cache = original_counter_cache