diff --git a/litellm/constants.py b/litellm/constants.py index 87d6fa1a744..c9d9ff155ff 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1478,6 +1478,10 @@ CLOUDZERO_MAX_FETCHED_DATA_RECORDS: Final = int(os.getenv("CLOUDZERO_MAX_FETCHED SPEND_LOG_CLEANUP_JOB_NAME: Final = "spend_log_cleanup" KEY_ROTATION_JOB_NAME: Final = "litellm_key_rotation_job" EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME: Final = "litellm_expired_ui_session_key_cleanup_job" +WEEKLY_SPEND_REPORT_JOB_ID: Final = "weekly_spend_report_job" +MONTHLY_SPEND_REPORT_JOB_ID: Final = "monthly_spend_report_job" +PROMETHEUS_FALLBACK_STATS_JOB_ID: Final = "prometheus_fallback_stats_job" +SLACK_DAILY_REPORT_LOCK_ID: Final = "slack_daily_report" SPEND_LOG_RUN_LOOPS: Final = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500)) SPEND_LOG_CLEANUP_BATCH_SIZE: Final = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000)) SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3)) diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 7edfb93e581..f3cd937599c 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -17,7 +17,7 @@ import litellm.litellm_core_utils.litellm_logging import litellm.types from litellm._logging import verbose_logger, verbose_proxy_logger from litellm.caching.caching import DualCache -from litellm.constants import HOURS_IN_A_DAY +from litellm.constants import HOURS_IN_A_DAY, SLACK_DAILY_REPORT_LOCK_ID from litellm.integrations.custom_batch_logger import CustomBatchLogger from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_type from litellm.integrations.SlackAlerting.hanging_request_check import ( @@ -51,6 +51,7 @@ from .batching_handler import send_to_webhook, squash_payloads from .utils import process_slack_alerting_variables if TYPE_CHECKING: + from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager from litellm.router import Router as _Router Router = _Router @@ -1576,7 +1577,11 @@ Model Info: except Exception: pass - async def _run_scheduler_helper(self, llm_router) -> bool: + async def _run_scheduler_helper( + self, + llm_router, + pod_lock_manager: "PodLockManager | None" = None, + ) -> bool: """ Returns: - True -> report sent @@ -1601,6 +1606,16 @@ Model Info: interval_seconds: Final = self.alerting_args.daily_report_frequency if current_time - report_sent >= interval_seconds: + if ( + pod_lock_manager is not None + and ( + await pod_lock_manager.acquire_lock( + cronjob_id=SLACK_DAILY_REPORT_LOCK_ID, ttl=interval_seconds, allow_reentrant=False + ) + ) + is False + ): + return False # Sneak in the reporting logic here await self.send_daily_reports(router=llm_router) # Also, don't forget to update the report_sent time after sending the report! @@ -1612,7 +1627,11 @@ Model Info: return report_sent_bool - async def _run_scheduled_daily_report(self, llm_router: Any | None = None): + async def _run_scheduled_daily_report( + self, + llm_router: Any | None = None, + pod_lock_manager: "PodLockManager | None" = None, + ): """ If 'daily_reports' enabled @@ -1625,7 +1644,7 @@ Model Info: if "daily_reports" in self.alert_types: while True: - await self._run_scheduler_helper(llm_router=llm_router) + await self._run_scheduler_helper(llm_router=llm_router, pod_lock_manager=pod_lock_manager) interval = random.randint( self.alerting_args.report_check_interval - 3, self.alerting_args.report_check_interval + 3, diff --git a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py index c74cb412c68..4be1331e955 100644 --- a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py +++ b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py @@ -43,6 +43,7 @@ end self, cronjob_id: str, ttl: int | None = None, + allow_reentrant: bool = True, ) -> bool | None: """ Attempt to acquire the lock for a specific cron job using Redis. @@ -53,6 +54,10 @@ end ttl: Optional custom TTL in seconds. Defaults to DEFAULT_CRON_JOB_LOCK_TTL_SECONDS. Use a longer TTL for jobs that may take longer than the default 60s (e.g. key rotation with many keys). + allow_reentrant: With the default True, a pod that already holds the lock + acquires it again (leader election semantics). Pass False when the live + lock marks work as already done for this window, so not even the holder + may redo it before the TTL expires. """ if self.redis_cache is None: verbose_proxy_logger.debug("redis_cache is None, skipping acquire_lock") @@ -88,7 +93,7 @@ end if current_value is not None: if isinstance(current_value, bytes): current_value = current_value.decode("utf-8") - if current_value == self.pod_id: + if current_value == self.pod_id and allow_reentrant: verbose_proxy_logger.info( "Pod %s already holds the Redis lock for cronjob_id=%s", self.pod_id, @@ -96,14 +101,12 @@ end ) self._emit_acquired_lock_event(cronjob_id, self.pod_id) return True - else: - verbose_proxy_logger.info( - "Spend tracking - pod %s could not acquire lock for cronjob_id=%s, " - "held by pod %s. Spend updates in Redis will wait for the leader pod to commit.", - self.pod_id, - cronjob_id, - current_value, - ) + verbose_proxy_logger.info( + "Pod %s could not acquire lock for cronjob_id=%s, held by pod %s.", + self.pod_id, + cronjob_id, + current_value, + ) return False except Exception as e: verbose_proxy_logger.error("Error acquiring Redis lock for %s: %s", cronjob_id, e) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f79819c76d3..19030d50110 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -233,6 +233,8 @@ from litellm.constants import ( GLOBAL_PROXY_SPEND_CACHE_KEY, LITELLM_PROXY_ADMIN_NAME, LITELLM_PROXY_BUDGET_NAME, + MONTHLY_SPEND_REPORT_JOB_ID, + PROMETHEUS_FALLBACK_STATS_JOB_ID, PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS, PROXY_BATCH_POLLING_ENABLED, PROXY_BATCH_POLLING_INTERVAL, @@ -240,6 +242,7 @@ from litellm.constants import ( PROXY_BUDGET_RESCHEDULER_MAX_TIME, PROXY_BUDGET_RESCHEDULER_MIN_TIME, PROXY_CONFIG_RELOAD_INTERVAL_SECONDS, + WEEKLY_SPEND_REPORT_JOB_ID, ) from litellm.exceptions import RejectedRequestError from litellm.integrations.custom_guardrail import ModifyResponseException @@ -9043,41 +9046,76 @@ class ProxyStartupEvent: spend_report_frequency: Final[str] = general_settings.get("spend_report_frequency", "7d") or "7d" days: Final = int(spend_report_frequency[:-1]) - if spend_report_frequency[-1].lower() != "d": - raise ValueError("spend_report_frequency must be specified in days, e.g., '1d', '7d'") + if spend_report_frequency[-1].lower() != "d" or days <= 0: + raise ValueError("spend_report_frequency must be a positive number of days, e.g., '1d', '7d'") + + pod_lock_manager: Final = proxy_logging_obj.db_spend_update_writer.pod_lock_manager + weekly_lock_ttl: Final = duration_in_seconds(spend_report_frequency) - 3600 + + async def _scheduled_weekly_spend_report() -> None: + # TTL spans the whole reporting window: each pod's interval anchor is its own + # boot time + jitter, so a shorter lock would let a later pod re-send the report. + # Minus an hour so the next window's first firer finds a free key + if ( + await pod_lock_manager.acquire_lock( + cronjob_id=WEEKLY_SPEND_REPORT_JOB_ID, ttl=weekly_lock_ttl, allow_reentrant=False + ) + is False + ): + return + await proxy_logging_obj.slack_alerting_instance.send_weekly_spend_report(spend_report_frequency) + + async def _scheduled_monthly_spend_report() -> None: + if ( + await pod_lock_manager.acquire_lock( + cronjob_id=MONTHLY_SPEND_REPORT_JOB_ID, ttl=3600, allow_reentrant=False + ) + is False + ): + return + await proxy_logging_obj.slack_alerting_instance.send_monthly_spend_report() scheduler.add_job( - proxy_logging_obj.slack_alerting_instance.send_weekly_spend_report, + _scheduled_weekly_spend_report, "interval", days=days, next_run_time=datetime.now() + timedelta(seconds=10 + random.randint(0, 300)), - args=[spend_report_frequency], - id="weekly_spend_report_job", + id=WEEKLY_SPEND_REPORT_JOB_ID, replace_existing=True, misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, ) scheduler.add_job( - proxy_logging_obj.slack_alerting_instance.send_monthly_spend_report, + _scheduled_monthly_spend_report, "cron", day=1, - id="monthly_spend_report_job", + id=MONTHLY_SPEND_REPORT_JOB_ID, replace_existing=True, ) if os.getenv("PROMETHEUS_URL"): from zoneinfo import ZoneInfo + async def _scheduled_fallback_stats() -> None: + if ( + await pod_lock_manager.acquire_lock( + cronjob_id=PROMETHEUS_FALLBACK_STATS_JOB_ID, ttl=3600, allow_reentrant=False + ) + is False + ): + return + await proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus() + scheduler.add_job( - proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus, + _scheduled_fallback_stats, "cron", hour=PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS, minute=0, timezone=ZoneInfo("America/Los_Angeles"), - id="prometheus_fallback_stats_job", + id=PROMETHEUS_FALLBACK_STATS_JOB_ID, replace_existing=True, ) - await proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus() + await _scheduled_fallback_stats() @classmethod async def _setup_prisma_client( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b5972935806..9653106f7e0 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -468,7 +468,10 @@ class ProxyLogging: and not self.daily_report_started ): asyncio.create_task( - self.slack_alerting_instance._run_scheduled_daily_report(llm_router=llm_router) + self.slack_alerting_instance._run_scheduled_daily_report( + llm_router=llm_router, + pod_lock_manager=self.db_spend_update_writer.pod_lock_manager, + ) ) # RUN DAILY REPORT (if scheduled) self.daily_report_started = True diff --git a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py b/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py index 1ea4795207d..23a35098697 100644 --- a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py +++ b/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py @@ -1,17 +1,21 @@ +import asyncio import datetime import json import os import sys +import time import unittest -from typing import List, Optional, Tuple +from typing import Final, List, Optional, Tuple from unittest.mock import ANY, AsyncMock, MagicMock, Mock, patch -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system-path +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system-path import litellm +from litellm.caching.caching import DualCache from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting from litellm.proxy._types import CallInfo, Litellm_EntityType +from litellm.types.integrations.slack_alerting import SlackAlertingCacheKeys class TestSlackAlerting(unittest.TestCase): @@ -20,37 +24,27 @@ class TestSlackAlerting(unittest.TestCase): def test_get_percent_of_max_budget_left(self): # Test case 1: When max_budget is None - user_info = CallInfo( - max_budget=None, spend=50.0, event_group=Litellm_EntityType.KEY - ) + user_info = CallInfo(max_budget=None, spend=50.0, event_group=Litellm_EntityType.KEY) result = self.slack_alerting._get_percent_of_max_budget_left(user_info) self.assertEqual(result, 0.0) # Test case 2: When max_budget is 0 - user_info = CallInfo( - max_budget=0.0, spend=50.0, event_group=Litellm_EntityType.KEY - ) + user_info = CallInfo(max_budget=0.0, spend=50.0, event_group=Litellm_EntityType.KEY) result = self.slack_alerting._get_percent_of_max_budget_left(user_info) self.assertEqual(result, 0.0) # Test case 3: When spend is less than max_budget - user_info = CallInfo( - max_budget=100.0, spend=75.0, event_group=Litellm_EntityType.KEY - ) + user_info = CallInfo(max_budget=100.0, spend=75.0, event_group=Litellm_EntityType.KEY) result = self.slack_alerting._get_percent_of_max_budget_left(user_info) self.assertEqual(result, 0.25) # Test case 4: When spend equals max_budget - user_info = CallInfo( - max_budget=100.0, spend=100.0, event_group=Litellm_EntityType.KEY - ) + user_info = CallInfo(max_budget=100.0, spend=100.0, event_group=Litellm_EntityType.KEY) result = self.slack_alerting._get_percent_of_max_budget_left(user_info) self.assertEqual(result, 0.0) # Test case 5: When spend exceeds max_budget - user_info = CallInfo( - max_budget=100.0, spend=120.0, event_group=Litellm_EntityType.KEY - ) + user_info = CallInfo(max_budget=100.0, spend=120.0, event_group=Litellm_EntityType.KEY) result = self.slack_alerting._get_percent_of_max_budget_left(user_info) self.assertEqual(result, -0.2) @@ -189,7 +183,9 @@ class TestSlackAlerting(unittest.TestCase): # Test the specific formatting logic we're interested in alert_type_formatted = f"Alert type: `{alert_type.name}`\n" - formatted_message = f"{alert_type_formatted}\n Level: `{level}`\nTimestamp: `{current_time}`\n\nMessage: {message}" + formatted_message = ( + f"{alert_type_formatted}\n Level: `{level}`\nTimestamp: `{current_time}`\n\nMessage: {message}" + ) # Verify alert_type is in the formatted message as expected self.assertIn("Alert type: `llm_exceptions`", formatted_message) @@ -214,9 +210,7 @@ class TestSlackAlerting(unittest.TestCase): json.dumps(outage_value) # Verify the specific error message - self.assertIn( - "Object of type set is not JSON serializable", str(context.exception) - ) + self.assertIn("Object of type set is not JSON serializable", str(context.exception)) def test_fixed_redis_serialization(self): """Test that our fix resolves the Redis serialization error.""" @@ -245,3 +239,133 @@ class TestSlackAlerting(unittest.TestCase): ) self.assertEqual(parsed_data["alerts"], [408]) self.assertEqual(parsed_data["provider_region_id"], "vertex_aius-east1") + + +_REPORT_SENT_KEY: Final = SlackAlertingCacheKeys.report_sent_key.value +_DAILY_REPORT_FREQUENCY: Final = 900 + + +async def _slack_alerting_with_due_daily_report() -> SlackAlerting: + slack_alerting: Final = SlackAlerting( + internal_usage_cache=DualCache(), + alerting_args={"daily_report_frequency": _DAILY_REPORT_FREQUENCY}, + ) + await slack_alerting.internal_usage_cache.async_set_cache( + key=_REPORT_SENT_KEY, + value=time.time() - _DAILY_REPORT_FREQUENCY - 1, + ) + slack_alerting.send_daily_reports = AsyncMock() + return slack_alerting + + +async def _read_report_sent(slack_alerting: SlackAlerting) -> float: + return await slack_alerting.internal_usage_cache.async_get_cache( + key=_REPORT_SENT_KEY, + parent_otel_span=None, + ) + + +@pytest.mark.asyncio +async def test_daily_report_skipped_when_another_pod_holds_the_lock(): + """regression: issue #14809 - every pod sent its own copy of the daily report. + + The losing pod must also leave report_sent untouched so the winner's window still counts. + """ + slack_alerting: Final = await _slack_alerting_with_due_daily_report() + report_sent_before: Final = await _read_report_sent(slack_alerting) + pod_lock_manager: Final = AsyncMock() + pod_lock_manager.acquire_lock.return_value = False + + result: Final = await slack_alerting._run_scheduler_helper( + llm_router=MagicMock(), + pod_lock_manager=pod_lock_manager, + ) + + assert result is False + slack_alerting.send_daily_reports.assert_not_awaited() + assert await _read_report_sent(slack_alerting) == report_sent_before + pod_lock_manager.acquire_lock.assert_awaited_once_with( + cronjob_id="slack_daily_report", + ttl=_DAILY_REPORT_FREQUENCY, + allow_reentrant=False, + ) + + +@pytest.mark.asyncio +async def test_daily_report_sent_by_the_pod_that_wins_the_lock(): + slack_alerting: Final = await _slack_alerting_with_due_daily_report() + report_sent_before: Final = await _read_report_sent(slack_alerting) + llm_router: Final = MagicMock() + pod_lock_manager: Final = AsyncMock() + pod_lock_manager.acquire_lock.return_value = True + + result: Final = await slack_alerting._run_scheduler_helper( + llm_router=llm_router, + pod_lock_manager=pod_lock_manager, + ) + + assert result is True + slack_alerting.send_daily_reports.assert_awaited_once_with(router=llm_router) + assert await _read_report_sent(slack_alerting) > report_sent_before + pod_lock_manager.acquire_lock.assert_awaited_once_with( + cronjob_id="slack_daily_report", + ttl=_DAILY_REPORT_FREQUENCY, + allow_reentrant=False, + ) + + +@pytest.mark.parametrize("lock_state", ["no_pod_lock_manager", "no_redis_configured"]) +@pytest.mark.asyncio +async def test_daily_report_still_sent_without_a_working_lock(lock_state: str): + """Single-pod parity: a missing lock manager, or one whose acquire_lock returns None + because redis isn't configured, must not suppress the report.""" + slack_alerting: Final = await _slack_alerting_with_due_daily_report() + report_sent_before: Final = await _read_report_sent(slack_alerting) + llm_router: Final = MagicMock() + pod_lock_manager: Final = ( + None if lock_state == "no_pod_lock_manager" else AsyncMock(acquire_lock=AsyncMock(return_value=None)) + ) + + result: Final = await slack_alerting._run_scheduler_helper( + llm_router=llm_router, + pod_lock_manager=pod_lock_manager, + ) + + assert result is True + slack_alerting.send_daily_reports.assert_awaited_once_with(router=llm_router) + assert await _read_report_sent(slack_alerting) > report_sent_before + + +@pytest.mark.asyncio +async def test_daily_report_lock_not_attempted_before_the_interval_elapses(): + """The lock is a per-window marker, so a pod must not burn it on a check that isn't due yet.""" + slack_alerting: Final = await _slack_alerting_with_due_daily_report() + await slack_alerting.internal_usage_cache.async_set_cache(key=_REPORT_SENT_KEY, value=time.time()) + pod_lock_manager: Final = AsyncMock() + pod_lock_manager.acquire_lock.return_value = True + + result: Final = await slack_alerting._run_scheduler_helper( + llm_router=MagicMock(), + pod_lock_manager=pod_lock_manager, + ) + + assert result is False + pod_lock_manager.acquire_lock.assert_not_awaited() + slack_alerting.send_daily_reports.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_scheduled_daily_report_threads_the_pod_lock_manager_through(): + """The loop in _run_scheduled_daily_report is where the lock manager reaches the gate.""" + slack_alerting: Final = SlackAlerting(alert_types=["daily_reports"]) + pod_lock_manager: Final = AsyncMock() + slack_alerting._run_scheduler_helper = AsyncMock(side_effect=asyncio.CancelledError) + + with pytest.raises(asyncio.CancelledError): + await slack_alerting._run_scheduled_daily_report( + llm_router=MagicMock(), + pod_lock_manager=pod_lock_manager, + ) + + _, kwargs = slack_alerting._run_scheduler_helper.await_args + assert kwargs["pod_lock_manager"] is pod_lock_manager diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py index f2745052faa..7a1ab60c547 100644 --- a/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py @@ -7,9 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi.testclient import TestClient -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path from litellm.constants import DEFAULT_CRON_JOB_LOCK_TTL_SECONDS from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager @@ -310,9 +308,7 @@ async def test_lock_takeover_race_condition(mock_redis): @pytest.mark.asyncio -async def test_release_lock_uses_atomic_compare_delete_script_when_available( - pod_lock_manager, mock_redis -): +async def test_release_lock_uses_atomic_compare_delete_script_when_available(pod_lock_manager, mock_redis): """ Test that release_lock prefers atomic compare-and-delete Lua script when redis cache exposes script registration. @@ -323,12 +319,8 @@ async def test_release_lock_uses_atomic_compare_delete_script_when_available( await pod_lock_manager.release_lock(cronjob_id="test_job") lock_key = pod_lock_manager.get_redis_lock_key(cronjob_id="test_job") - mock_redis.async_register_script.assert_called_once_with( - PodLockManager._COMPARE_AND_DELETE_LOCK_SCRIPT - ) - script_callable.assert_called_once_with( - keys=[lock_key], args=[json.dumps(pod_lock_manager.pod_id)] - ) + mock_redis.async_register_script.assert_called_once_with(PodLockManager._COMPARE_AND_DELETE_LOCK_SCRIPT) + script_callable.assert_called_once_with(keys=[lock_key], args=[json.dumps(pod_lock_manager.pod_id)]) mock_redis.async_get_cache.assert_not_called() mock_redis.async_delete_cache.assert_not_called() @@ -359,9 +351,7 @@ async def test_release_lock_lua_path_emits_released_event(pod_lock_manager, mock with patch.object(pod_lock_manager, "_emit_released_lock_event") as mock_emit: await pod_lock_manager.release_lock(cronjob_id="test_job") - mock_emit.assert_called_once_with( - cronjob_id="test_job", pod_id=pod_lock_manager.pod_id - ) + mock_emit.assert_called_once_with(cronjob_id="test_job", pod_id=pod_lock_manager.pod_id) class FakeRedisLockStore: @@ -437,9 +427,7 @@ async def test_release_lock_preserves_lock_held_by_other_pod(): @pytest.mark.asyncio -async def test_release_lock_falls_back_to_get_del_when_lua_execution_fails( - pod_lock_manager, mock_redis -): +async def test_release_lock_falls_back_to_get_del_when_lua_execution_fails(pod_lock_manager, mock_redis): """ Test that release_lock falls back to GET+DEL when Lua script execution raises (e.g. Redis restart cleared loaded scripts). @@ -457,3 +445,14 @@ async def test_release_lock_falls_back_to_get_del_when_lua_execution_fails( mock_redis.async_delete_cache.assert_called_once_with(lock_key) # Cached script handle should be reset so next call re-registers assert pod_lock_manager._release_lock_script is None + + +@pytest.mark.asyncio +async def test_acquire_lock_own_lock_not_reentrant(pod_lock_manager, mock_redis): + """With allow_reentrant=False a live lock means the window's work is done, so even + the holder gets False; the default stays reentrant for leader-election callers.""" + mock_redis.async_set_cache.return_value = False + mock_redis.async_get_cache.return_value = pod_lock_manager.pod_id + + assert await pod_lock_manager.acquire_lock(cronjob_id="test_job", allow_reentrant=False) is False + assert await pod_lock_manager.acquire_lock(cronjob_id="test_job") is True diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index 6ac1e15e7b5..40ca7e3a64e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -22,6 +22,7 @@ import inspect import json import logging import os +from collections.abc import Awaitable, Callable from typing import List, Optional, Union from unittest.mock import AsyncMock, MagicMock, patch @@ -397,9 +398,7 @@ def test__redact_worker_config_for_logging_masks_nested_secret_fields(): "database_url": nested_db_url, "database_extra_connection_params": {"password": nested_extra_pw}, "alert_to_webhook_url": {"budget_alerts": nested_webhook}, - "pass_through_endpoints": [ - {"path": "/up", "headers": {"Authorization": nested_bearer}} - ], + "pass_through_endpoints": [{"path": "/up", "headers": {"Authorization": nested_bearer}}], } } } @@ -451,16 +450,13 @@ def test_load_from_azure_key_vault_disabled_no_side_effect(monkeypatch): import litellm sentinel_secret_mgr = object() - monkeypatch.setattr( - litellm, "secret_manager_client", sentinel_secret_mgr, raising=False - ) + monkeypatch.setattr(litellm, "secret_manager_client", sentinel_secret_mgr, raising=False) result = load_from_azure_key_vault(use_azure_key_vault=False) observed = { "return_value": result, - "secret_manager_unchanged": litellm.secret_manager_client - is sentinel_secret_mgr, + "secret_manager_unchanged": litellm.secret_manager_client is sentinel_secret_mgr, "called_with": False, } assert normalize(observed) == { @@ -614,9 +610,7 @@ def test_get_litellm_model_info_uses_base_model_for_lookup(monkeypatch): observed = { "called_arg": ( - fake_get.call_args.args[0] - if fake_get.call_args.args - else fake_get.call_args.kwargs.get("model") + fake_get.call_args.args[0] if fake_get.call_args.args else fake_get.call_args.kwargs.get("model") ), "returned_max_tokens": result.get("max_tokens"), "returned_cost": result.get("input_cost_per_token"), @@ -663,9 +657,7 @@ def test_run_ollama_serve_invokes_subprocess_popen(monkeypatch): def test_run_ollama_serve_popen_failure_is_swallowed(monkeypatch): """Popen raising OSError must NOT propagate — function logs and returns.""" - monkeypatch.setattr( - ps.subprocess, "Popen", MagicMock(side_effect=OSError("no ollama binary")) - ) + monkeypatch.setattr(ps.subprocess, "Popen", MagicMock(side_effect=OSError("no ollama binary"))) result = run_ollama_serve() assert result is None @@ -685,8 +677,7 @@ async def test_proxy_startup_event_is_async_context_manager_with_expected_signat observed = { "param_count": len(sig.parameters), "has_app_param": "app" in sig.parameters, - "wrapped_is_async": inspect.iscoroutinefunction(wrapped) - or inspect.isasyncgenfunction(wrapped), + "wrapped_is_async": inspect.iscoroutinefunction(wrapped) or inspect.isasyncgenfunction(wrapped), "has_asynccontextmanager_wrapper": wrapped is not None, } assert normalize(observed) == { @@ -777,3 +768,164 @@ def test_proxy_startup_event_warns_for_global_budget_without_database(): assert budget_check_pos < warn_pos < next_startup_section_pos, ( "DB-less budget warning must run after Prisma setup and the DB-backed budget block" ) + + +# --------------------------------------------------------------------------- +# _initialize_slack_alerting_jobs — spend-report pod locking (issue #14809) +# --------------------------------------------------------------------------- + +SlackAlertingJobs = dict[str, Callable[[], Awaitable[None]]] + + +def _make_slack_alerting_proxy_logging(acquire_lock_result: bool | None) -> MagicMock: + proxy_logging_obj = MagicMock() + proxy_logging_obj.slack_alerting_instance.alerting = ["slack"] + proxy_logging_obj.slack_alerting_instance.send_weekly_spend_report = AsyncMock() + proxy_logging_obj.slack_alerting_instance.send_monthly_spend_report = AsyncMock() + proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus = AsyncMock() + pod_lock_manager = proxy_logging_obj.db_spend_update_writer.pod_lock_manager + pod_lock_manager.acquire_lock = AsyncMock(return_value=acquire_lock_result) + pod_lock_manager.release_lock = AsyncMock() + return proxy_logging_obj + + +async def _init_slack_alerting_jobs( + acquire_lock_result: bool | None, + spend_report_frequency: str = "7d", +) -> tuple[SlackAlertingJobs, MagicMock]: + scheduler = MagicMock() + proxy_logging_obj = _make_slack_alerting_proxy_logging(acquire_lock_result) + + await ProxyStartupEvent._initialize_slack_alerting_jobs( + scheduler=scheduler, + general_settings={"spend_report_frequency": spend_report_frequency}, + proxy_logging_obj=proxy_logging_obj, + prisma_client=MagicMock(), + ) + + jobs = {call.kwargs["id"]: call.args[0] for call in scheduler.add_job.call_args_list} + return jobs, proxy_logging_obj + + +@pytest.mark.parametrize("spend_report_frequency", ["0d", "-1d", "7h"]) +@pytest.mark.asyncio +async def test_initialize_slack_alerting_jobs_invalid_frequency_raises(spend_report_frequency: str): + """A non-positive window used to become an every-second APScheduler interval, and now also + computes a negative lock TTL that expires instantly and suppresses the report for good. + match= is load-bearing: drop the guard and "-1d" still raises, but from duration_in_seconds.""" + with pytest.raises(ValueError, match="positive number of days"): + await _init_slack_alerting_jobs( + acquire_lock_result=True, + spend_report_frequency=spend_report_frequency, + ) + + +@pytest.mark.asyncio +async def test_weekly_spend_report_skipped_when_another_pod_holds_the_lock(): + """regression: issue #14809 - every pod ran its own weekly spend report job.""" + jobs, proxy_logging_obj = await _init_slack_alerting_jobs(acquire_lock_result=False) + + await jobs["weekly_spend_report_job"]() + + proxy_logging_obj.slack_alerting_instance.send_weekly_spend_report.assert_not_awaited() + proxy_logging_obj.db_spend_update_writer.pod_lock_manager.acquire_lock.assert_awaited_once_with( + cronjob_id="weekly_spend_report_job", + ttl=7 * 86400 - 3600, + allow_reentrant=False, + ) + + +@pytest.mark.parametrize("acquire_lock_result", [True, None]) +@pytest.mark.asyncio +async def test_weekly_spend_report_sent_when_the_lock_is_free_or_absent(acquire_lock_result): + """None means redis isn't configured; a single-pod deploy must still report.""" + jobs, proxy_logging_obj = await _init_slack_alerting_jobs(acquire_lock_result=acquire_lock_result) + + await jobs["weekly_spend_report_job"]() + + proxy_logging_obj.slack_alerting_instance.send_weekly_spend_report.assert_awaited_once_with("7d") + + +@pytest.mark.asyncio +async def test_weekly_spend_report_lock_ttl_tracks_the_configured_window(): + """TTL is the window less an hour: long enough that no second pod re-sends inside the + window, short enough that the lock is gone before the next one opens. A fixed TTL would + break one end or the other as soon as spend_report_frequency changes.""" + jobs, proxy_logging_obj = await _init_slack_alerting_jobs(acquire_lock_result=True, spend_report_frequency="1d") + + await jobs["weekly_spend_report_job"]() + + proxy_logging_obj.db_spend_update_writer.pod_lock_manager.acquire_lock.assert_awaited_once_with( + cronjob_id="weekly_spend_report_job", + ttl=86400 - 3600, + allow_reentrant=False, + ) + proxy_logging_obj.slack_alerting_instance.send_weekly_spend_report.assert_awaited_once_with("1d") + + +@pytest.mark.asyncio +async def test_monthly_spend_report_skipped_when_another_pod_holds_the_lock(): + jobs, proxy_logging_obj = await _init_slack_alerting_jobs(acquire_lock_result=False) + + await jobs["monthly_spend_report_job"]() + + proxy_logging_obj.slack_alerting_instance.send_monthly_spend_report.assert_not_awaited() + proxy_logging_obj.db_spend_update_writer.pod_lock_manager.acquire_lock.assert_awaited_once_with( + cronjob_id="monthly_spend_report_job", + ttl=3600, + allow_reentrant=False, + ) + + +@pytest.mark.parametrize("acquire_lock_result", [True, None]) +@pytest.mark.asyncio +async def test_monthly_spend_report_sent_when_the_lock_is_free_or_absent(acquire_lock_result): + jobs, proxy_logging_obj = await _init_slack_alerting_jobs(acquire_lock_result=acquire_lock_result) + + await jobs["monthly_spend_report_job"]() + + proxy_logging_obj.slack_alerting_instance.send_monthly_spend_report.assert_awaited_once_with() + + +@pytest.mark.asyncio +async def test_spend_report_locks_are_never_released(): + """The lock is a per-window marker, not a mutex: releasing it lets the next pod re-send.""" + jobs, proxy_logging_obj = await _init_slack_alerting_jobs(acquire_lock_result=True) + + await jobs["weekly_spend_report_job"]() + await jobs["monthly_spend_report_job"]() + + proxy_logging_obj.db_spend_update_writer.pod_lock_manager.release_lock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_prometheus_fallback_stats_job_skipped_when_another_pod_holds_the_lock(monkeypatch): + """The boot-time send goes through the same gate, so a losing pod sends nothing at all: + startup and the scheduled job both stay at zero.""" + monkeypatch.setenv("PROMETHEUS_URL", "http://prometheus.invalid") + jobs, proxy_logging_obj = await _init_slack_alerting_jobs(acquire_lock_result=False) + send_fallback_stats = proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus + assert send_fallback_stats.await_count == 0 + + await jobs["prometheus_fallback_stats_job"]() + + assert send_fallback_stats.await_count == 0 + proxy_logging_obj.db_spend_update_writer.pod_lock_manager.acquire_lock.assert_awaited_with( + cronjob_id="prometheus_fallback_stats_job", + ttl=3600, + allow_reentrant=False, + ) + assert proxy_logging_obj.db_spend_update_writer.pod_lock_manager.acquire_lock.await_count == 2 + + +@pytest.mark.parametrize("acquire_lock_result", [True, None]) +@pytest.mark.asyncio +async def test_prometheus_fallback_stats_job_runs_when_the_lock_is_free_or_absent(monkeypatch, acquire_lock_result): + monkeypatch.setenv("PROMETHEUS_URL", "http://prometheus.invalid") + jobs, proxy_logging_obj = await _init_slack_alerting_jobs(acquire_lock_result=acquire_lock_result) + send_fallback_stats = proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus + assert send_fallback_stats.await_count == 1 + + await jobs["prometheus_fallback_stats_job"]() + + assert send_fallback_stats.await_count == 2 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index d9ac45c0531..3acb9fcafd3 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -19,9 +19,7 @@ from fastapi import FastAPI from fastapi.staticfiles import StaticFiles from fastapi.testclient import TestClient -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system-path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system-path import litellm import litellm.proxy.proxy_server as proxy_server_module @@ -179,9 +177,7 @@ def test_login_v2_returns_json_on_http_exception(monkeypatch): from fastapi import HTTPException mock_prisma_client = MagicMock() - mock_authenticate_user = AsyncMock( - side_effect=HTTPException(status_code=401, detail="Unauthorized") - ) + mock_authenticate_user = AsyncMock(side_effect=HTTPException(status_code=401, detail="Unauthorized")) monkeypatch.setattr( "litellm.proxy.auth.login_utils.authenticate_user", @@ -477,9 +473,7 @@ def test_fallback_login_has_no_deprecation_banner(client_no_auth): "relative/path/logo.png", ], ) -def test_get_logo_url_does_not_disclose_local_paths( - client_no_auth, monkeypatch, ui_logo_path -): +def test_get_logo_url_does_not_disclose_local_paths(client_no_auth, monkeypatch, ui_logo_path): # ``/get_logo_url`` is unauthenticated. Returning a local filesystem # path verbatim discloses admin-only config to any caller. Only # browser-loadable HTTP(S) URLs should be returned; for local paths @@ -579,9 +573,7 @@ def test_restructure_ui_html_files_handles_nested_routes(tmp_path): assert not (ui_root / "home.html").exists() assert (ui_root / "home" / "index.html").read_text() == "home" assert not (ui_root / "mcp" / "oauth" / "callback.html").exists() - assert ( - ui_root / "mcp" / "oauth" / "callback" / "index.html" - ).read_text() == "callback" + assert (ui_root / "mcp" / "oauth" / "callback" / "index.html").read_text() == "callback" assert (ui_root / "existing" / "index.html").read_text() == "keep" assert (ui_root / "_next" / "ignore.html").read_text() == "asset" assert (ui_root / "litellm-asset-prefix" / "ignore.html").read_text() == "asset" @@ -626,9 +618,7 @@ def test_admin_ui_export_serves_nested_extensionless_routes(): and "_next" not in path.parts and "litellm-asset-prefix" not in path.parts ] - assert not nested_html_offenders, ( - "Nested routes must be named index.html. Offenders: " f"{nested_html_offenders}" - ) + assert not nested_html_offenders, f"Nested routes must be named index.html. Offenders: {nested_html_offenders}" callback_index = out_dir / "mcp" / "oauth" / "callback" / "index.html" assert callback_index.is_file(), ( @@ -645,9 +635,7 @@ def test_admin_ui_export_serves_nested_extensionless_routes(): follow_redirects=False, ) assert redirect.status_code == 307 - assert redirect.headers["location"].endswith( - "/ui/mcp/oauth/callback/?code=abc&state=xyz" - ) + assert redirect.headers["location"].endswith("/ui/mcp/oauth/callback/?code=abc&state=xyz") landed = client.get("/ui/mcp/oauth/callback?code=abc&state=xyz") assert landed.status_code == 200 @@ -712,6 +700,7 @@ async def test_initialize_scheduled_jobs_credentials(monkeypatch): mock_prisma_client = MagicMock() mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() with ( @@ -750,9 +739,7 @@ async def test_initialize_scheduled_jobs_credentials(monkeypatch): assert mock_proxy_config.get_credentials.call_count == 1 # Direct call # Verify a scheduled job was added for get_credentials - mock_scheduler_calls = [ - call[0] for call in mock_proxy_config.get_credentials.mock_calls - ] + mock_scheduler_calls = [call[0] for call in mock_proxy_config.get_credentials.mock_calls] assert len(mock_scheduler_calls) > 0 @@ -773,6 +760,7 @@ async def test_periodic_reload_job_scheduled_without_store_model_in_db(monkeypat mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() scheduler = AsyncIOScheduler() @@ -813,6 +801,7 @@ async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval( mock_prisma_client = MagicMock() mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() mock_scheduler = MagicMock() @@ -861,6 +850,7 @@ async def test_initialize_scheduled_jobs_rejects_non_positive_config_reload_inte mock_prisma_client = MagicMock() mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() mock_scheduler = MagicMock() @@ -907,6 +897,7 @@ async def test_initialize_scheduled_jobs_hydrates_mcp_when_store_model_in_db_fal mock_prisma_client = MagicMock() mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() with ( @@ -1051,9 +1042,7 @@ def test_get_config_custom_callback_api_env_vars(monkeypatch): assert response.status_code == 200 callbacks = response.json()["callbacks"] - custom_cb = next( - (cb for cb in callbacks if cb["name"] == "custom_callback_api"), None - ) + custom_cb = next((cb for cb in callbacks if cb["name"] == "custom_callback_api"), None) assert custom_cb is not None assert custom_cb["variables"] == { @@ -1101,9 +1090,7 @@ def test_get_config_callbacks_fall_back_to_process_env(mock_env_vars, monkeypatc app.dependency_overrides = original_overrides assert response.status_code == 200 - langfuse_cb = next( - (cb for cb in response.json()["callbacks"] if cb["name"] == "langfuse"), None - ) + langfuse_cb = next((cb for cb in response.json()["callbacks"] if cb["name"] == "langfuse"), None) assert langfuse_cb is not None assert langfuse_cb["variables"] == { "LANGFUSE_PUBLIC_KEY": "pk-env-only", @@ -1150,9 +1137,7 @@ def test_get_config_callback_env_secrets_redacted_for_non_admin(mock_env_vars, m app.dependency_overrides = original_overrides assert response.status_code == 200 - langfuse_cb = next( - (cb for cb in response.json()["callbacks"] if cb["name"] == "langfuse"), None - ) + langfuse_cb = next((cb for cb in response.json()["callbacks"] if cb["name"] == "langfuse"), None) assert langfuse_cb is not None assert langfuse_cb["variables"]["LANGFUSE_SECRET_KEY"] == "REDACTED" assert langfuse_cb["variables"]["LANGFUSE_HOST"] == "https://cloud.langfuse.com" @@ -1202,9 +1187,7 @@ def test_get_config_returns_email_settings(monkeypatch): app.dependency_overrides = original_overrides assert response.status_code == 200 - email_alert = next( - (a for a in response.json()["alerts"] if a["name"] == "email"), None - ) + email_alert = next((a for a in response.json()["alerts"] if a["name"] == "email"), None) assert email_alert is not None variables = email_alert["variables"] @@ -1349,9 +1332,7 @@ def test_get_config_returns_slack_webhook(monkeypatch): mock_logging = MagicMock() mock_logging.slack_alerting_instance.alert_types = ["budget_alerts"] - mock_logging.slack_alerting_instance._all_possible_alert_types.return_value = [ - "budget_alerts" - ] + mock_logging.slack_alerting_instance._all_possible_alert_types.return_value = ["budget_alerts"] mock_logging.slack_alerting_instance.alert_to_webhook_url = {} monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_logging) monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data)) @@ -1368,9 +1349,7 @@ def test_get_config_returns_slack_webhook(monkeypatch): app.dependency_overrides = original_overrides assert response.status_code == 200 - slack_alert = next( - (a for a in response.json()["alerts"] if a["name"] == "slack"), None - ) + slack_alert = next((a for a in response.json()["alerts"] if a["name"] == "slack"), None) assert slack_alert is not None masked_url = slack_alert["variables"]["SLACK_WEBHOOK_URL"] @@ -1390,9 +1369,7 @@ def test_get_config_cleared_slack_webhook_not_overridden_by_os_env(monkeypatch): """ from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth - monkeypatch.setenv( - "SLACK_WEBHOOK_URL", "https://hooks.slack.com/services/STALE/OS/ENVVALUE" - ) + monkeypatch.setenv("SLACK_WEBHOOK_URL", "https://hooks.slack.com/services/STALE/OS/ENVVALUE") config_data = { "litellm_settings": {}, "general_settings": {"alerting": ["slack"]}, @@ -1405,9 +1382,7 @@ def test_get_config_cleared_slack_webhook_not_overridden_by_os_env(monkeypatch): mock_logging = MagicMock() mock_logging.slack_alerting_instance.alert_types = ["budget_alerts"] - mock_logging.slack_alerting_instance._all_possible_alert_types.return_value = [ - "budget_alerts" - ] + mock_logging.slack_alerting_instance._all_possible_alert_types.return_value = ["budget_alerts"] mock_logging.slack_alerting_instance.alert_to_webhook_url = {} monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_logging) monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data)) @@ -1424,9 +1399,7 @@ def test_get_config_cleared_slack_webhook_not_overridden_by_os_env(monkeypatch): app.dependency_overrides = original_overrides assert response.status_code == 200 - slack_alert = next( - (a for a in response.json()["alerts"] if a["name"] == "slack"), None - ) + slack_alert = next((a for a in response.json()["alerts"] if a["name"] == "slack"), None) assert slack_alert is not None assert slack_alert["variables"]["SLACK_WEBHOOK_URL"] == "" @@ -1505,9 +1478,7 @@ async def test_aaaproxy_startup_master_key(mock_prisma, monkeypatch, tmp_path): # Test Case 3: Master key with os.environ prefix test_resolved_key = "sk-resolved-key" - test_config_with_prefix = { - "general_settings": {"master_key": "os.environ/CUSTOM_MASTER_KEY"} - } + test_config_with_prefix = {"general_settings": {"master_key": "os.environ/CUSTOM_MASTER_KEY"}} # Create config with os.environ prefix with open(config_path, "w") as f: @@ -1659,9 +1630,7 @@ async def test_get_all_team_models(): ) # Verify find_many was called with where clause for specific teams - mock_litellm_teamtable.find_many.assert_called_with( - where={"team_id": {"in": ["team1"]}} - ) + mock_litellm_teamtable.find_many.assert_called_with(where={"team_id": {"in": ["team1"]}}) # Verify router.get_model_list was called only for team1 models expected_calls = [ @@ -1856,14 +1825,10 @@ async def test_apply_search_filter_scopes_byok_to_caller_teams(): prisma_client = MagicMock() prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=2) - prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock( - return_value=[db_caller_row, db_other_row] - ) + prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[db_caller_row, db_other_row]) caller_user_row = MagicMock() caller_user_row.teams = ["team-mine"] - prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=caller_user_row - ) + prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=caller_user_row) proxy_config = MagicMock() proxy_config.decrypt_model_list_from_db = lambda rows: [ @@ -1893,12 +1858,10 @@ async def test_apply_search_filter_scopes_byok_to_caller_teams(): assert "byok-db-mine" in filtered_ids assert "public-id" in filtered_ids assert "byok-other" not in filtered_ids, ( - "router-side BYOK from another team must be dropped from search " - "when caller doesn't belong to that team" + "router-side BYOK from another team must be dropped from search when caller doesn't belong to that team" ) assert "byok-db-other" not in filtered_ids, ( - "DB-only BYOK from another team must be dropped from search when " - "caller doesn't belong to that team" + "DB-only BYOK from another team must be dropped from search when caller doesn't belong to that team" ) # total_count is router_models_count (2: caller_team_byok + public_model, # other_team_byok dropped router-side) + DB count (2 from the mocked @@ -2049,9 +2012,7 @@ async def test_filter_models_by_team_id_excludes_viewer_direct_access(): assert "byok-team-111" in visible_ids, "team-111's own BYOK must always be visible" assert "byok-team-222" not in visible_ids, "must not leak other teams' BYOK" - assert ( - "public-id" not in visible_ids - ), "viewer's direct_access must not widen the team's visible set" + assert "public-id" not in visible_ids, "viewer's direct_access must not widen the team's visible set" @pytest.mark.asyncio @@ -2234,9 +2195,7 @@ async def test_add_access_group_models_to_team_models(): mock_ag_row.access_model_names = ["claude-3", "gemini"] mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock( - return_value=[mock_ag_row] - ) + mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_ag_row]) result = await _add_access_group_models_to_team_models( team_db_objects_typed=[ @@ -2312,9 +2271,7 @@ async def test_add_access_group_models_multiple_teams_shared_group(): mock_extra_row.access_model_names = ["gemini"] mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock( - return_value=[mock_shared_row, mock_extra_row] - ) + mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_shared_row, mock_extra_row]) result = await _add_access_group_models_to_team_models( team_db_objects_typed=[team_a, team_b], @@ -2507,24 +2464,14 @@ async def test_delete_deployment_type_mismatch(): # The two SHA-hash models have no corresponding entry in combined_id_list # and must be evicted. assert len(deleted_ids) == 2, f"Expected 2 deletions (SHA-hash models), got {deleted_ids}" - assert ( - "a96e12e76b36a57cfae57a41288eb41567629cac89b4828c6f7074afc3534695" - in deleted_ids - ) - assert ( - "a40186dd0fdb9b7282380277d7f57044d29de95bfbfcd7f4322b3493702d5cd3" - in deleted_ids - ) + assert "a96e12e76b36a57cfae57a41288eb41567629cac89b4828c6f7074afc3534695" in deleted_ids + assert "a40186dd0fdb9b7282380277d7f57044d29de95bfbfcd7f4322b3493702d5cd3" in deleted_ids # Models 12345678 and 12345679 exist in the config (as integers); str() # conversion in _delete_deployment makes them match the router's string IDs, # so they must NOT be evicted. - assert ( - "12345678" not in deleted_ids - ), f"Model 12345678 should NOT be deleted. Deleted IDs: {deleted_ids}" - assert ( - "12345679" not in deleted_ids - ), f"Model 12345679 should NOT be deleted. Deleted IDs: {deleted_ids}" + assert "12345678" not in deleted_ids, f"Model 12345678 should NOT be deleted. Deleted IDs: {deleted_ids}" + assert "12345679" not in deleted_ids, f"Model 12345679 should NOT be deleted. Deleted IDs: {deleted_ids}" assert still_desired is not None assert {"12345678", "12345679"} <= still_desired, ( @@ -2597,9 +2544,7 @@ async def test_get_config_from_file(tmp_path, monkeypatch): await proxy_config._get_config_from_file(str(empty_file)) # Test Case 5: Using global user_config_file_path when no config_file_path provided - monkeypatch.setattr( - "litellm.proxy.proxy_server.user_config_file_path", str(config_file) - ) + monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", str(config_file)) result = await proxy_config._get_config_from_file(None) assert result == test_config @@ -2718,9 +2663,7 @@ async def test_add_proxy_budget_to_db_only_creates_user_no_keys(): ) # Patch generate_key_helper_fn in proxy_server where it's being called from - with patch( - "litellm.proxy.proxy_server.generate_key_helper_fn", mock_generate_key_helper - ): + with patch("litellm.proxy.proxy_server.generate_key_helper_fn", mock_generate_key_helper): # Call the function under test ProxyStartupEvent._add_proxy_budget_to_db() @@ -2846,9 +2789,7 @@ async def test_custom_ui_sso_sign_in_handler_config_loading(): proxy_config = ProxyConfig() # Create a mock router since load_config requires it mock_router = MagicMock() - await proxy_config.load_config( - router=mock_router, config_file_path=config_file_path - ) + await proxy_config.load_config(router=mock_router, config_file_path=config_file_path) # Verify get_instance_fn was called with correct parameters mock_get_instance.assert_called_with( @@ -2888,9 +2829,7 @@ async def test_load_config_max_budget_env_var_coerced_to_float(tmp_path, monkeyp original_max_budget = litellm.max_budget try: proxy_config = ProxyConfig() - await proxy_config.load_config( - router=MagicMock(), config_file_path=str(config_file) - ) + await proxy_config.load_config(router=MagicMock(), config_file_path=str(config_file)) assert isinstance(litellm.max_budget, float) assert litellm.max_budget == 10.0 assert litellm.max_budget > 0 @@ -2925,9 +2864,7 @@ async def test_load_config_max_ui_session_budget_applied_and_coerced(tmp_path, m original_budget = litellm.max_ui_session_budget try: proxy_config = ProxyConfig() - await proxy_config.load_config( - router=MagicMock(), config_file_path=str(config_file) - ) + await proxy_config.load_config(router=MagicMock(), config_file_path=str(config_file)) assert isinstance(litellm.max_ui_session_budget, float) assert litellm.max_ui_session_budget == 2.5 finally: @@ -2953,9 +2890,7 @@ async def test_load_config_max_ui_session_budget_none_disables_cap(tmp_path): original_budget = litellm.max_ui_session_budget try: proxy_config = ProxyConfig() - await proxy_config.load_config( - router=MagicMock(), config_file_path=str(config_file) - ) + await proxy_config.load_config(router=MagicMock(), config_file_path=str(config_file)) assert litellm.max_ui_session_budget is None finally: litellm.max_ui_session_budget = original_budget @@ -3010,10 +2945,7 @@ async def test_load_config_default_internal_user_params_without_max_budget(tmp_p absent_config_file = tmp_path / "absent_config.yaml" absent_config_file.write_text( - "model_list: []\n" - "litellm_settings:\n" - " default_internal_user_params:\n" - " user_role: internal_user\n" + "model_list: []\nlitellm_settings:\n default_internal_user_params:\n user_role: internal_user\n" ) null_config_file = tmp_path / "null_config.yaml" @@ -3060,9 +2992,7 @@ async def test_load_config_user_url_validation_handles_null_and_string_false(tmp ) ) - await ProxyConfig().load_config( - router=MagicMock(), config_file_path=str(null_config_file) - ) + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(null_config_file)) assert litellm.user_url_validation is True assert litellm.user_url_allowed_hosts is None assert litellm.provider_url_destination_allowed_hosts is None @@ -3077,9 +3007,7 @@ async def test_load_config_user_url_validation_handles_null_and_string_false(tmp ) ) - await ProxyConfig().load_config( - router=MagicMock(), config_file_path=str(false_config_file) - ) + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(false_config_file)) assert litellm.user_url_validation is False @@ -3107,12 +3035,8 @@ async def test_load_environment_variables_direct_and_os_environ(): # Mock get_secret_str to return a resolved value mock_secret_value = "resolved_secret_value" - with patch( - "litellm.proxy.proxy_server.get_secret_str", return_value=mock_secret_value - ) as mock_get_secret: - with patch.dict( - os.environ, {}, clear=False - ): # Don't clear existing env vars, just track changes + with patch("litellm.proxy.proxy_server.get_secret_str", return_value=mock_secret_value) as mock_get_secret: + with patch.dict(os.environ, {}, clear=False): # Don't clear existing env vars, just track changes # Call the method under test proxy_config._load_environment_variables(test_config) @@ -3125,9 +3049,7 @@ async def test_load_environment_variables_direct_and_os_environ(): assert os.environ["SECRET_VAR"] == mock_secret_value # Verify get_secret_str was called with the correct value - mock_get_secret.assert_called_once_with( - secret_name="os.environ/ACTUAL_SECRET_VAR" - ) + mock_get_secret.assert_called_once_with(secret_name="os.environ/ACTUAL_SECRET_VAR") @pytest.mark.asyncio @@ -3180,9 +3102,7 @@ async def test_load_environment_variables_litellm_license_and_edge_cases(): assert result is None # Method returns None # Test Case 4: os.environ/ prefix but get_secret_str returns None - test_config_secret_none = { - "environment_variables": {"FAILED_SECRET": "os.environ/NONEXISTENT_SECRET"} - } + test_config_secret_none = {"environment_variables": {"FAILED_SECRET": "os.environ/NONEXISTENT_SECRET"}} with patch("litellm.proxy.proxy_server.get_secret_str", return_value=None): with patch.dict(os.environ, {}, clear=False): @@ -3221,9 +3141,7 @@ async def test_load_environment_variables_blocks_dangerous_keys(): # Blocked keys should not be set to the attacker value assert os.environ.get("PATH") != "/tmp/evil" - assert ( - "LD_PRELOAD" not in os.environ or os.environ["LD_PRELOAD"] != "/tmp/evil.so" - ) + assert "LD_PRELOAD" not in os.environ or os.environ["LD_PRELOAD"] != "/tmp/evil.so" assert os.environ.get("PYTHONPATH") != "/tmp/evil" # Safe keys should still be set @@ -3297,15 +3215,11 @@ async def test_write_config_to_file(monkeypatch): # Mock general_settings mock_general_settings = {"store_model_in_db": True} - monkeypatch.setattr( - "litellm.proxy.proxy_server.general_settings", mock_general_settings - ) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", mock_general_settings) # Mock user_config_file_path test_config_path = "/tmp/test_config.yaml" - monkeypatch.setattr( - "litellm.proxy.proxy_server.user_config_file_path", test_config_path - ) + monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", test_config_path) proxy_config = ProxyConfig() @@ -3326,9 +3240,7 @@ async def test_write_config_to_file(monkeypatch): # Verify the config passed to DB has model_list removed call_args = mock_prisma_client.insert_data.call_args - assert call_args.kwargs["data"] == { - "key": "value" - } # model_list should be popped + assert call_args.kwargs["data"] == {"key": "value"} # model_list should be popped assert call_args.kwargs["table_name"] == "config" @@ -3349,15 +3261,11 @@ async def test_write_config_to_file_when_store_model_in_db_false(monkeypatch): # Mock general_settings mock_general_settings = {"store_model_in_db": False} - monkeypatch.setattr( - "litellm.proxy.proxy_server.general_settings", mock_general_settings - ) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", mock_general_settings) # Mock user_config_file_path test_config_path = "/tmp/test_config.yaml" - monkeypatch.setattr( - "litellm.proxy.proxy_server.user_config_file_path", test_config_path - ) + monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", test_config_path) proxy_config = ProxyConfig() @@ -3412,22 +3320,20 @@ async def test_async_data_generator_midstream_error(): for chunk in mock_chunks: yield chunk - mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = ( - mock_streaming_iterator - ) + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = mock_streaming_iterator # Mock async_post_call_streaming_hook to return error on third chunk def mock_streaming_hook(*args, **kwargs): chunk = kwargs.get("response") # Return error message for the third chunk (simulating guardrail trigger) if chunk == mock_chunks[2]: - return 'data: {"error": {"error": "Azure Content Safety Guardrail: Hate crossed severity 2, Got severity: 2"}}' + return ( + 'data: {"error": {"error": "Azure Content Safety Guardrail: Hate crossed severity 2, Got severity: 2"}}' + ) # Return normal chunks for first two return chunk - mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock( - side_effect=mock_streaming_hook - ) + mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(side_effect=mock_streaming_hook) mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() # Mock the global proxy_logging_obj @@ -3438,26 +3344,18 @@ async def test_async_data_generator_midstream_error(): # Collect all yielded data from the generator yielded_data = [] try: - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) except Exception as e: # If there's an exception, that's also part of what we want to test pass # Verify the results - assert ( - len(yielded_data) >= 3 - ), f"Expected at least 3 chunks, got {len(yielded_data)}: {yielded_data}" + assert len(yielded_data) >= 3, f"Expected at least 3 chunks, got {len(yielded_data)}: {yielded_data}" # First two chunks should be normal data - assert yielded_data[0].startswith( - "data: " - ), f"First chunk should start with 'data: ', got: {yielded_data[0]}" - assert yielded_data[1].startswith( - "data: " - ), f"Second chunk should start with 'data: ', got: {yielded_data[1]}" + assert yielded_data[0].startswith("data: "), f"First chunk should start with 'data: ', got: {yielded_data[0]}" + assert yielded_data[1].startswith("data: "), f"Second chunk should start with 'data: ', got: {yielded_data[1]}" # The error message should be yielded error_found = False @@ -3469,15 +3367,11 @@ async def test_async_data_generator_midstream_error(): if "data: [DONE]" in data: done_found = True - assert ( - error_found - ), f"Error message should be found in yielded data. Got: {yielded_data}" + assert error_found, f"Error message should be found in yielded data. Got: {yielded_data}" assert done_found, f"[DONE] message should be found at the end. Got: {yielded_data}" # Verify that the streaming hook was called for each chunk - assert mock_proxy_logging_obj.async_post_call_streaming_hook.call_count == len( - mock_chunks - ) + assert mock_proxy_logging_obj.async_post_call_streaming_hook.call_count == len(mock_chunks) # Verify that post_call_failure_hook was NOT called (since this is not an exception case) mock_proxy_logging_obj.post_call_failure_hook.assert_not_called() @@ -3564,15 +3458,11 @@ async def test_chat_completion_result_no_nested_none_values(): # Verify the mock has None values before serialization raw_dict = mock_model_response.model_dump() none_paths_before = _has_nested_none_values(raw_dict) - assert ( - len(none_paths_before) > 0 - ), "Mock should have None values before exclude_none=True" + assert len(none_paths_before) > 0, "Mock should have None values before exclude_none=True" # Mock the request processing to return our mock response mock_base_processor = MagicMock() - mock_base_processor.base_process_llm_request = AsyncMock( - return_value=mock_model_response - ) + mock_base_processor.base_process_llm_request = AsyncMock(return_value=mock_model_response) # Mock other dependencies mock_request = MagicMock(spec=Request) @@ -3601,9 +3491,9 @@ async def test_chat_completion_result_no_nested_none_values(): # Check that there are no nested None values in the result none_paths_after = _has_nested_none_values(result) - assert ( - len(none_paths_after) == 0 - ), f"Result should not contain nested None values. Found None at: {none_paths_after}" + assert len(none_paths_after) == 0, ( + f"Result should not contain nested None values. Found None at: {none_paths_after}" + ) # Verify essential fields are present assert "id" in result @@ -3629,9 +3519,7 @@ async def test_chat_completion_result_no_nested_none_values(): "annotations", ] for field in excluded_fields: - assert ( - field not in message - ), f"Field '{field}' should be excluded when it's None" + assert field not in message, f"Field '{field}' should be excluded when it's None" # ============================================================================ @@ -3686,9 +3574,7 @@ class TestPriceDataReloadAPI: with patch( "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new=AsyncMock( - return_value=ModelCostMapReloaded( - model_cost_map={"gpt-3.5-turbo": {"input_cost_per_token": 0.001}} - ) + return_value=ModelCostMapReloaded(model_cost_map={"gpt-3.5-turbo": {"input_cost_per_token": 0.001}}) ), ): # Mock the database connection @@ -3706,10 +3592,7 @@ class TestPriceDataReloadAPI: assert "timestamp" in data assert "models_count" in data # The new implementation immediately reloads and returns the count - assert ( - "Price data reloaded successfully! 1 models updated." - in data["message"] - ) + assert "Price data reloaded successfully! 1 models updated." in data["message"] assert data["models_count"] == 1 finally: # Restore the full model cost map so subsequent tests are not affected @@ -3732,9 +3615,7 @@ class TestPriceDataReloadAPI: def test_get_model_cost_map_public_access(self, client_no_auth): """Test that the model cost map endpoint is publicly accessible""" - with patch( - "litellm.model_cost", {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}} - ): + with patch("litellm.model_cost", {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}}): response = client_no_auth.get("/public/litellm_model_cost_map") assert response.status_code == 200 @@ -3756,9 +3637,7 @@ class TestPriceDataReloadAPI: response = client_with_auth.post("/reload/model_cost_map") - assert ( - response.status_code == 500 - ) # An unexpected exception still maps to 500 + assert response.status_code == 500 # An unexpected exception still maps to 500 data = response.json() assert "Failed to reload model cost map" in data["detail"] @@ -3966,9 +3845,7 @@ class TestPriceDataReloadIntegration: try: with patch( "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", - new=AsyncMock( - return_value=ModelCostMapReloaded(model_cost_map=mock_cost_map) - ), + new=AsyncMock(return_value=ModelCostMapReloaded(model_cost_map=mock_cost_map)), ): # Mock the database connection with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: @@ -4036,10 +3913,14 @@ class TestPriceDataReloadIntegration: original_model_cost = litellm.model_cost.copy() try: with ( - patch("litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock) as mock_get_map, + patch( + "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock + ) as mock_get_map, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): - mock_get_map.return_value = ModelCostMapReloaded(model_cost_map={"gpt-3.5-turbo": {"input_cost_per_token": 0.001}}) + mock_get_map.return_value = ModelCostMapReloaded( + model_cost_map={"gpt-3.5-turbo": {"input_cost_per_token": 0.001}} + ) asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) @@ -4080,7 +3961,9 @@ class TestPriceDataReloadIntegration: original_model_cost = litellm.model_cost.copy() try: with ( - patch("litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock) as mock_get_map, + patch( + "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock + ) as mock_get_map, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) @@ -4115,10 +3998,14 @@ class TestPriceDataReloadIntegration: original_model_cost = litellm.model_cost.copy() try: with ( - patch("litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock) as mock_get_map, + patch( + "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock + ) as mock_get_map, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): - mock_get_map.return_value = ModelCostMapReloaded(model_cost_map={"gpt-4-test": {"input_cost_per_token": 0.5}}) + mock_get_map.return_value = ModelCostMapReloaded( + model_cost_map={"gpt-4-test": {"input_cost_per_token": 0.5}} + ) asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) @@ -4156,10 +4043,14 @@ class TestPriceDataReloadIntegration: original_model_cost = litellm.model_cost.copy() try: with ( - patch("litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock) as mock_get_map, + patch( + "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock + ) as mock_get_map, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): - mock_get_map.return_value = ModelCostMapReloaded(model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}) + mock_get_map.return_value = ModelCostMapReloaded( + model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}} + ) for _ in range(3): for pod in pods: @@ -4196,10 +4087,14 @@ class TestPriceDataReloadIntegration: original_model_cost = litellm.model_cost.copy() try: with ( - patch("litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock) as mock_get_map, + patch( + "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock + ) as mock_get_map, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): - mock_get_map.return_value = ModelCostMapReloaded(model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}) + mock_get_map.return_value = ModelCostMapReloaded( + model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}} + ) asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) @@ -4228,10 +4123,14 @@ class TestPriceDataReloadIntegration: original_model_cost = litellm.model_cost.copy() try: with ( - patch("litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock) as mock_get_map, + patch( + "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock + ) as mock_get_map, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): - mock_get_map.return_value = ModelCostMapReloaded(model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}) + mock_get_map.return_value = ModelCostMapReloaded( + model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}} + ) asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) @@ -4271,7 +4170,9 @@ class TestPriceDataReloadIntegration: ) as mock_get_map, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): - mock_get_map.return_value = ModelCostMapReloaded(model_cost_map={"gpt-4": {"input_cost_per_token": 0.1}}) + mock_get_map.return_value = ModelCostMapReloaded( + model_cost_map={"gpt-4": {"input_cost_per_token": 0.1}} + ) asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) @@ -4317,8 +4218,7 @@ class TestPriceDataReloadIntegration: asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) assert litellm.model_cost is original_model_cost, ( - "a failed reload must keep the currently loaded cost map, " - "not swap in the packaged backup" + "a failed reload must keep the currently loaded cost map, not swap in the packaged backup" ) assert proxy_config.model_cost_map_loaded_at == pod_data_loaded_at, ( "a failed reload must not stamp the pod's data age, otherwise the retry waits a full interval" @@ -4428,11 +4328,15 @@ class TestPriceDataReloadIntegration: original_model_cost = litellm.model_cost.copy() try: with ( - patch("litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock) as mock_get_map, + patch( + "litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock + ) as mock_get_map, patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): - mock_get_map.return_value = ModelCostMapReloaded(model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}) + mock_get_map.return_value = ModelCostMapReloaded( + model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}} + ) mock_prisma.db.litellm_config.upsert = AsyncMock( return_value=_reload_schedule_row({}, reload_revision=9) ) @@ -4480,14 +4384,10 @@ class TestPriceDataReloadIntegration: mock_prisma.get_generic_data = AsyncMock(return_value=mock_config) mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1)) - with patch( - "litellm.anthropic_beta_headers_manager.reload_beta_headers_config" - ) as mock_reload: + with patch("litellm.anthropic_beta_headers_manager.reload_beta_headers_config") as mock_reload: mock_reload.return_value = {"anthropic": {"beta_header": "test-value"}} - asyncio.run( - proxy_config._check_and_reload_anthropic_beta_headers(mock_prisma) - ) + asyncio.run(proxy_config._check_and_reload_anthropic_beta_headers(mock_prisma)) # Verify the upsert update branch preserves interval_hours mock_prisma.db.litellm_config.upsert.assert_called() @@ -4519,9 +4419,7 @@ class TestPriceDataReloadIntegration: app.dependency_overrides[user_api_key_auth] = lambda: mock_auth client = TestClient(app) - with patch( - "litellm.anthropic_beta_headers_manager.reload_beta_headers_config" - ) as mock_reload: + with patch("litellm.anthropic_beta_headers_manager.reload_beta_headers_config") as mock_reload: mock_reload.return_value = {"anthropic": {"beta_header": "test-value"}} with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: @@ -4619,9 +4517,7 @@ async def test_add_router_settings_from_db_config_merge_logic(): # Mock prisma client mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_config.find_first = AsyncMock( - return_value=mock_db_config - ) + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) # Call the method under test await proxy_config._add_router_settings_from_db_config( @@ -4631,9 +4527,7 @@ async def test_add_router_settings_from_db_config_merge_logic(): ) # Verify find_first was called with correct parameters - mock_prisma_client.db.litellm_config.find_first.assert_called_once_with( - where={"param_name": "router_settings"} - ) + mock_prisma_client.db.litellm_config.find_first.assert_called_once_with(where={"param_name": "router_settings"}) # Verify update_settings was called mock_router.update_settings.assert_called_once() @@ -4713,9 +4607,7 @@ async def test_add_router_settings_from_db_config_edge_cases(): # Test Case 4: Config has no router_settings mock_db_config = MagicMock() mock_db_config.param_value = {"db_setting": "db_value"} - mock_prisma_client.db.litellm_config.find_first = AsyncMock( - return_value=mock_db_config - ) + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) await proxy_config._add_router_settings_from_db_config( config_data={}, # No router_settings in config @@ -4740,9 +4632,7 @@ async def test_add_router_settings_from_db_config_edge_cases(): # Test Case 6: DB config exists but param_value is not a dict mock_db_config_invalid = MagicMock() mock_db_config_invalid.param_value = "not_a_dict" - mock_prisma_client.db.litellm_config.find_first = AsyncMock( - return_value=mock_db_config_invalid - ) + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config_invalid) config_data = {"router_settings": {"config_setting": "config_value"}} @@ -4794,9 +4684,7 @@ async def test_add_router_settings_shallow_merge_behavior(): } mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_config.find_first = AsyncMock( - return_value=mock_db_config - ) + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) await proxy_config._add_router_settings_from_db_config( config_data=config_data, @@ -4873,9 +4761,7 @@ async def test_model_info_v1_oci_secrets_not_leaked(): patch("litellm.proxy.proxy_server.user_model", None), ): # Call the model_info_v1 endpoint - result = await model_info_v1( - user_api_key_dict=mock_user_api_key_dict, litellm_model_id=None - ) + result = await model_info_v1(user_api_key_dict=mock_user_api_key_dict, litellm_model_id=None) # Verify the result structure assert "data" in result @@ -4886,40 +4772,24 @@ async def test_model_info_v1_oci_secrets_not_leaked(): # Verify that sensitive OCI fields are masked assert "****" in litellm_params["oci_key"], "oci_key should be masked" - assert ( - "****" in litellm_params["oci_fingerprint"] - ), "oci_fingerprint should be masked" + assert "****" in litellm_params["oci_fingerprint"], "oci_fingerprint should be masked" assert "****" in litellm_params["oci_tenancy"], "oci_tenancy should be masked" assert "****" in litellm_params["oci_key_file"], "oci_key_file should be masked" # Verify that non-sensitive fields are NOT masked - assert ( - litellm_params["model"] == "oci/xai.grok-4" - ), "model field should not be masked" - assert ( - litellm_params["oci_region"] == "us-phoenix-1" - ), "oci_region should not be masked" + assert litellm_params["model"] == "oci/xai.grok-4", "model field should not be masked" + assert litellm_params["oci_region"] == "us-phoenix-1", "oci_region should not be masked" assert litellm_params["drop_params"] is True, "drop_params should not be masked" # Verify the model field specifically is not masked (this was the original issue) - assert ( - "****" not in litellm_params["model"] - ), "model field should never be masked" - assert litellm_params["model"].startswith( - "oci/" - ), "model should retain its full value" + assert "****" not in litellm_params["model"], "model field should never be masked" + assert litellm_params["model"].startswith("oci/"), "model should retain its full value" # Verify that actual secret values are not present in the response result_str = str(result) - assert ( - "ocid1.api_key.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk" - not in result_str - ) + assert "ocid1.api_key.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk" not in result_str assert "aa:bb:cc:dd:ee:ff:11:22:33:44:55:66:77:88:99:00" not in result_str - assert ( - "ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk" - not in result_str - ) + assert "ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk" not in result_str assert "/path/to/oci_api_key.pem" not in result_str @@ -4949,9 +4819,7 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks(): event_types=["success"], existing_callbacks=mock_success_callbacks, ) - mock_callback_manager.add_litellm_success_callback.assert_called_once_with( - "prometheus" - ) + mock_callback_manager.add_litellm_success_callback.assert_called_once_with("prometheus") mock_callback_manager.reset_mock() # Test Case 2: Add failure callback @@ -4961,9 +4829,7 @@ def test_add_callback_from_db_to_in_memory_litellm_callbacks(): event_types=["failure"], existing_callbacks=mock_failure_callbacks, ) - mock_callback_manager.add_litellm_failure_callback.assert_called_once_with( - "langfuse" - ) + mock_callback_manager.add_litellm_failure_callback.assert_called_once_with("langfuse") mock_callback_manager.reset_mock() # Test Case 3: Add callback for both success and failure @@ -5064,10 +4930,7 @@ def test_should_load_db_object_with_supported_db_objects(): assert proxy_config._should_load_db_object(object_type="mcp") is True assert proxy_config._should_load_db_object(object_type="guardrails") is True assert proxy_config._should_load_db_object(object_type="vector_stores") is True - assert ( - proxy_config._should_load_db_object(object_type="pass_through_endpoints") - is True - ) + assert proxy_config._should_load_db_object(object_type="pass_through_endpoints") is True assert proxy_config._should_load_db_object(object_type="prompts") is True assert proxy_config._should_load_db_object(object_type="model_cost_map") is True @@ -5093,12 +4956,8 @@ async def test_tag_cache_update_called(): "spend": 10.0, } - with patch.object( - cache, "async_get_cache", new=AsyncMock(return_value=mock_tag_obj) - ) as mock_get_cache: - with patch.object( - cache, "async_set_cache_pipeline", new=AsyncMock() - ) as mock_set_cache: + with patch.object(cache, "async_get_cache", new=AsyncMock(return_value=mock_tag_obj)) as mock_get_cache: + with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( token=None, user_id=None, @@ -5152,9 +5011,7 @@ async def test_tag_cache_update_multiple_tags(): with patch.object( cache, "async_get_cache", new=AsyncMock(side_effect=mock_get_cache_side_effect) ) as mock_get_cache: - with patch.object( - cache, "async_set_cache_pipeline", new=AsyncMock() - ) as mock_set_cache: + with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( token=None, user_id=None, @@ -5175,9 +5032,7 @@ async def test_tag_cache_update_multiple_tags(): assert len(cache_list) == 2 - tag_updates = { - cache_key: cache_value for cache_key, cache_value in cache_list - } + tag_updates = {cache_key: cache_value for cache_key, cache_value in cache_list} assert "tag:tag1" in tag_updates assert "tag:tag2" in tag_updates assert tag_updates["tag:tag1"]["spend"] == 15.0 @@ -5203,9 +5058,7 @@ async def test_update_cache_pipeline_honors_user_api_key_cache_ttl(): "async_get_cache", new=AsyncMock(return_value={"tag_name": "active-tag", "spend": 1.0}), ): - with patch.object( - cache, "async_set_cache_pipeline", new=AsyncMock() - ) as mock_set_cache: + with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( token=None, user_id=None, @@ -5248,9 +5101,7 @@ async def test_spend_tracking_never_writes_the_auth_object_back(): model_type=UserAPIKeyAuth, ) with ( - patch.object( - cache, "async_set_cache_pipeline", new=AsyncMock() - ) as mock_pipeline, + patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_pipeline, patch.object(cache, "async_set_cache", new=AsyncMock()) as mock_set, ): await litellm.proxy.proxy_server.update_cache( @@ -5261,9 +5112,7 @@ async def test_spend_tracking_never_writes_the_auth_object_back(): response_cost=5.0, parent_otel_span=None, ) - pending = [ - t for t in asyncio.all_tasks() if t is not asyncio.current_task() - ] + pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()] if pending: await asyncio.wait(pending, timeout=5) @@ -5305,12 +5154,8 @@ async def test_update_cache_global_proxy_spend_scalar_stays_shared(): cache = DualCache(default_in_memory_ttl=300) setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache) try: - with patch.object( - cache, "async_get_cache", new=AsyncMock(side_effect=fake_get) - ): - with patch.object( - cache, "async_set_cache_pipeline", new=AsyncMock() - ) as mock_set_cache: + with patch.object(cache, "async_get_cache", new=AsyncMock(side_effect=fake_get)): + with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( token=None, user_id="user-lit", @@ -5320,24 +5165,14 @@ async def test_update_cache_global_proxy_spend_scalar_stays_shared(): parent_otel_span=None, ) - pending = [ - t for t in asyncio.all_tasks() if t is not asyncio.current_task() - ] + pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()] if pending: await asyncio.wait(pending, timeout=5) calls = mock_set_cache.await_args_list - local_keys = [ - k - for c in calls - if c.kwargs.get("local_only") is True - for k, _ in c.kwargs["cache_list"] - ] + local_keys = [k for c in calls if c.kwargs.get("local_only") is True for k, _ in c.kwargs["cache_list"]] shared_keys = [ - k - for c in calls - if c.kwargs.get("local_only") is not True - for k, _ in c.kwargs["cache_list"] + k for c in calls if c.kwargs.get("local_only") is not True for k, _ in c.kwargs["cache_list"] ] assert "user-lit" in local_keys assert global_key not in local_keys @@ -5368,20 +5203,14 @@ async def test_init_sso_settings_in_db(): } mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock( - return_value=mock_sso_config - ) + mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config) # Mock _decrypt_and_set_db_env_variables - with patch.object( - proxy_config, "_decrypt_and_set_db_env_variables" - ) as mock_decrypt_and_set: + with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called with correct parameters - mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with( - where={"id": "sso_config"} - ) + mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with(where={"id": "sso_config"}) # Verify _decrypt_and_set_db_env_variables was called with uppercased keys mock_decrypt_and_set.assert_called_once() @@ -5421,15 +5250,11 @@ async def test_init_sso_settings_in_db_no_settings(): mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) # Mock _decrypt_and_set_db_env_variables - with patch.object( - proxy_config, "_decrypt_and_set_db_env_variables" - ) as mock_decrypt_and_set: + with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called - mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with( - where={"id": "sso_config"} - ) + mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with(where={"id": "sso_config"}) # Verify _decrypt_and_set_db_env_variables was NOT called when no settings exist mock_decrypt_and_set.assert_not_called() @@ -5448,9 +5273,7 @@ async def test_init_sso_settings_in_db_error_handling(): # Mock prisma client to raise an exception mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock( - side_effect=Exception("Database connection error") - ) + mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=Exception("Database connection error")) # The method should not raise an exception, it should log it instead try: @@ -5459,9 +5282,7 @@ async def test_init_sso_settings_in_db_error_handling(): assert True except Exception as e: # The exception should be caught and logged, not propagated - pytest.fail( - f"Exception should have been caught and logged, but was raised: {e}" - ) + pytest.fail(f"Exception should have been caught and logged, but was raised: {e}") @pytest.mark.asyncio @@ -5480,20 +5301,14 @@ async def test_init_sso_settings_in_db_empty_settings(): mock_sso_config.sso_settings = {} mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock( - return_value=mock_sso_config - ) + mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config) # Mock _decrypt_and_set_db_env_variables - with patch.object( - proxy_config, "_decrypt_and_set_db_env_variables" - ) as mock_decrypt_and_set: + with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called - mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with( - where={"id": "sso_config"} - ) + mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with(where={"id": "sso_config"}) # Verify _decrypt_and_set_db_env_variables was called with empty dict mock_decrypt_and_set.assert_called_once() @@ -5526,16 +5341,12 @@ async def test_init_sso_settings_in_db_retries_on_transport_error(): return mock_sso_config mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock( - side_effect=_flaky_find_unique - ) + mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=_flaky_find_unique) mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 - with patch.object( - proxy_config, "_decrypt_and_set_db_env_variables" - ) as mock_decrypt: + with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) assert len(invocations) == 2 @@ -5556,9 +5367,7 @@ async def test_init_sso_settings_in_db_propagates_when_reconnect_fails(): proxy_config = ProxyConfig() mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock( - side_effect=prisma.errors.ClientNotConnectedError() - ) + mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=prisma.errors.ClientNotConnectedError()) mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=False) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 @@ -5589,24 +5398,17 @@ async def test_init_hashicorp_vault_config_override_retries_on_transport_error() return None # No config in DB → function returns early after retry. mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_configoverrides.find_unique = AsyncMock( - side_effect=_flaky_find_unique - ) + mock_prisma_client.db.litellm_configoverrides.find_unique = AsyncMock(side_effect=_flaky_find_unique) mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 - await proxy_config._init_hashicorp_vault_config_override( - prisma_client=mock_prisma_client - ) + await proxy_config._init_hashicorp_vault_config_override(prisma_client=mock_prisma_client) assert len(invocations) == 2 mock_prisma_client.attempt_db_reconnect.assert_awaited_once() reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs - assert ( - reconnect_kwargs["reason"] - == "init_hashicorp_vault_config_override_lookup_failure" - ) + assert reconnect_kwargs["reason"] == "init_hashicorp_vault_config_override_lookup_failure" def test_update_config_fields_uppercases_env_vars(monkeypatch): @@ -5656,37 +5458,20 @@ def test_encrypt_env_variables_for_db_is_idempotent(monkeypatch): plaintext = "pk-langfuse-secret-value" # First write: plaintext in -> single-encrypted out. - enc1 = proxy_config._encrypt_env_variables_for_db( - {"LANGFUSE_PUBLIC_KEY": plaintext} - ) + enc1 = proxy_config._encrypt_env_variables_for_db({"LANGFUSE_PUBLIC_KEY": plaintext}) assert enc1["LANGFUSE_PUBLIC_KEY"] != plaintext - assert ( - decrypt_value_helper( - value=enc1["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY" - ) - == plaintext - ) + assert decrypt_value_helper(value=enc1["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY") == plaintext # UI round-trip: feed the ciphertext back in. Must NOT double-encrypt. enc2 = proxy_config._encrypt_env_variables_for_db(enc1) - assert ( - decrypt_value_helper( - value=enc2["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY" - ) - == plaintext - ) + assert decrypt_value_helper(value=enc2["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY") == plaintext # And again, ×3 total ciphertext re-feeds — still exactly one layer, # never stacked, no matter how many times the UI re-saves. enc3 = proxy_config._encrypt_env_variables_for_db(enc2) enc4 = proxy_config._encrypt_env_variables_for_db(enc3) for stacked in (enc3, enc4): - assert ( - decrypt_value_helper( - value=stacked["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY" - ) - == plaintext - ) + assert decrypt_value_helper(value=stacked["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY") == plaintext # Write path must not leak the value into the process environment. assert os.environ.get("LANGFUSE_PUBLIC_KEY") is None @@ -5728,15 +5513,11 @@ def test_get_prompt_spec_for_db_prompt_with_versions(): } # Test version 1 - prompt_spec_v1 = proxy_config._get_prompt_spec_for_db_prompt( - db_prompt=mock_prompt_v1 - ) + prompt_spec_v1 = proxy_config._get_prompt_spec_for_db_prompt(db_prompt=mock_prompt_v1) assert prompt_spec_v1.prompt_id == "chat_prompt.v1" # Test version 2 - prompt_spec_v2 = proxy_config._get_prompt_spec_for_db_prompt( - db_prompt=mock_prompt_v2 - ) + prompt_spec_v2 = proxy_config._get_prompt_spec_for_db_prompt(db_prompt=mock_prompt_v2) assert prompt_spec_v2.prompt_id == "chat_prompt.v2" @@ -5804,9 +5585,7 @@ async def test_get_image_non_root_uses_var_lib_assets_dir(monkeypatch): with ( patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs, - patch( - "litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect - ), + patch("litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect), patch("litellm.proxy.proxy_server.os.access", return_value=True), patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv, patch("litellm.proxy.proxy_server.FileResponse") as mock_file_response, @@ -5856,9 +5635,7 @@ async def test_get_image_non_root_fallback_to_default_logo(monkeypatch): # Mock os.path operations with ( patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs, - patch( - "litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect - ), + patch("litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect), patch("litellm.proxy.proxy_server.os.access", return_value=True), patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv, patch("litellm.proxy.proxy_server.FileResponse") as mock_file_response, @@ -5881,9 +5658,7 @@ async def test_get_image_non_root_fallback_to_default_logo(monkeypatch): # Verify that exists was called to check /var/lib/litellm/assets/logo.jpg assets_logo_path = "/var/lib/litellm/assets/logo.jpg" - assert any( - assets_logo_path in str(call) for call in exists_calls - ), f"Should check if {assets_logo_path} exists" + assert any(assets_logo_path in str(call) for call in exists_calls), f"Should check if {assets_logo_path} exists" # Verify FileResponse was called (with fallback logo) assert mock_file_response.called, "FileResponse should be called" @@ -5923,14 +5698,8 @@ async def test_get_image_root_case_uses_current_dir(monkeypatch): await get_image() # Verify makedirs was NOT called with /var/lib/litellm/assets (should not create it for root case) - var_lib_assets_calls = [ - call - for call in mock_makedirs.call_args_list - if "/var/lib/litellm/assets" in str(call) - ] - assert ( - len(var_lib_assets_calls) == 0 - ), "Should not create /var/lib/litellm/assets for root case" + var_lib_assets_calls = [call for call in mock_makedirs.call_args_list if "/var/lib/litellm/assets" in str(call)] + assert len(var_lib_assets_calls) == 0, "Should not create /var/lib/litellm/assets for root case" # Verify FileResponse was called assert mock_file_response.called, "FileResponse should be called" @@ -5961,15 +5730,11 @@ async def test_get_image_custom_local_logo_bypasses_cache(monkeypatch, tmp_path) return MagicMock() with ( - patch( - "litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response - ), + patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response), ): await get_image() - assert ( - len(calls_to_file_response) == 1 - ), "FileResponse should be called exactly once" + assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once" assert calls_to_file_response[0] == str(custom_logo.resolve()), ( f"Expected custom logo path, got {calls_to_file_response[0]}. " "A stale cached_logo.jpg may have been returned instead." @@ -5999,24 +5764,18 @@ async def test_get_image_default_logo_ignores_stale_cache(monkeypatch, tmp_path) return MagicMock() with ( - patch( - "litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response - ), + patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response), ): await get_image() - assert ( - len(calls_to_file_response) == 1 - ), "FileResponse should be called exactly once" + assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once" served_path = calls_to_file_response[0] assert served_path != str(cache_path.resolve()) assert served_path.endswith("logo.jpg") @pytest.mark.asyncio -async def test_get_image_custom_logo_missing_falls_through_to_default( - monkeypatch, tmp_path -): +async def test_get_image_custom_logo_missing_falls_through_to_default(monkeypatch, tmp_path): """ Test that when UI_LOGO_PATH points to a non-existent local file, get_image falls through to the default logo instead of failing. @@ -6037,26 +5796,18 @@ async def test_get_image_custom_logo_missing_falls_through_to_default( return MagicMock() with ( - patch( - "litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response - ), + patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response), ): await get_image() - assert ( - len(calls_to_file_response) == 1 - ), "FileResponse should be called exactly once" + assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once" served_path = calls_to_file_response[0] - assert served_path != str( - custom_logo_path - ), "Should not attempt to serve a non-existent custom logo" + assert served_path != str(custom_logo_path), "Should not attempt to serve a non-existent custom logo" assert served_path.endswith("logo.jpg") @pytest.mark.asyncio -async def test_get_image_custom_logo_missing_no_cache_serves_default( - monkeypatch, tmp_path -): +async def test_get_image_custom_logo_missing_no_cache_serves_default(monkeypatch, tmp_path): """ Test that when UI_LOGO_PATH points to a non-existent file AND there is no cached_logo.jpg, get_image serves the default logo instead of the non-existent @@ -6078,22 +5829,14 @@ async def test_get_image_custom_logo_missing_no_cache_serves_default( return MagicMock() with ( - patch( - "litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response - ), + patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response), ): await get_image() - assert ( - len(calls_to_file_response) == 1 - ), "FileResponse should be called exactly once" + assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once" served_path = calls_to_file_response[0] - assert served_path != str( - custom_logo_path - ), "Should not attempt to serve a non-existent custom logo" - assert served_path.endswith( - "logo.jpg" - ), f"Expected fallback to default logo.jpg, got {served_path}" + assert served_path != str(custom_logo_path), "Should not attempt to serve a non-existent custom logo" + assert served_path.endswith("logo.jpg"), f"Expected fallback to default logo.jpg, got {served_path}" def test_get_config_normalizes_string_callbacks(monkeypatch): @@ -6133,9 +5876,7 @@ def test_get_config_normalizes_string_callbacks(monkeypatch): success_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "success"] failure_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "failure"] - success_and_failure_callbacks = [ - cb["name"] for cb in callbacks if cb.get("type") == "success_and_failure" - ] + success_and_failure_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "success_and_failure"] assert "langfuse" in success_callbacks assert len(failure_callbacks) == 0 @@ -6172,9 +5913,7 @@ def test_deep_merge_dicts_skips_none_and_empty_lists(monkeypatch): }, } - result = proxy_config._update_config_fields( - current_config, "general_settings", db_param_value - ) + result = proxy_config._update_config_fields(current_config, "general_settings", db_param_value) assert result["general_settings"]["max_parallel_requests"] == 10 assert result["general_settings"]["allowed_models"] == ["gpt-3.5-turbo", "gpt-4"] @@ -6241,9 +5980,7 @@ class TestInvitationEndpoints: ), ], ) - def test_invitation_endpoints_proxy_admin_success( - self, client_with_auth, endpoint, payload, mock_return - ): + def test_invitation_endpoints_proxy_admin_success(self, client_with_auth, endpoint, payload, mock_return): """Proxy admin can successfully create and delete invitations.""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_prisma.db.litellm_invitationlink = MagicMock() @@ -6258,9 +5995,7 @@ class TestInvitationEndpoints: mock_prisma.db.litellm_invitationlink.find_unique = AsyncMock( return_value={**mock_return, "created_by": "admin-user-id"} ) - mock_prisma.db.litellm_invitationlink.delete = AsyncMock( - return_value=mock_return - ) + mock_prisma.db.litellm_invitationlink.delete = AsyncMock(return_value=mock_return) response = client_with_auth.post(endpoint, json=payload) assert response.status_code == 200 @@ -6275,9 +6010,7 @@ class TestInvitationEndpoints: ("/invitation/delete", {"invitation_id": "inv-456"}), ], ) - def test_invitation_endpoints_non_admin_denied( - self, client_with_auth, endpoint, payload - ): + def test_invitation_endpoints_non_admin_denied(self, client_with_auth, endpoint, payload): """Non-admin users cannot access invitation endpoints.""" from litellm.proxy._types import LitellmUserRoles @@ -6332,9 +6065,7 @@ async def test_async_data_generator_cleanup_on_early_exit(): for chunk in mock_chunks: yield chunk - mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = ( - mock_streaming_iterator - ) + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = mock_streaming_iterator mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock( side_effect=lambda **kwargs: kwargs.get("response") ) @@ -6346,9 +6077,7 @@ async def test_async_data_generator_cleanup_on_early_exit(): with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): # Consume only the first chunk then abandon the generator (simulates client disconnect) - gen = async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ) + gen = async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data) first_chunk = await gen.__anext__() assert first_chunk.startswith("data: ") @@ -6401,19 +6130,12 @@ async def test_async_data_generator_uses_direct_stream_fast_path_without_callbac mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object( - ProxyLogging, "_fire_deferred_stream_logging" - ) as mock_deferred_logging: + with patch.object(ProxyLogging, "_fire_deferred_stream_logging") as mock_deferred_logging: yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) - yielded_text = [ - chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk - for chunk in yielded_data - ] + yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data] assert len([chunk for chunk in yielded_text if chunk.startswith("data: {")]) == 2 assert yielded_text[-1] == "data: [DONE]\n\n" mock_proxy_logging_obj.async_post_call_streaming_iterator_hook.assert_not_called() @@ -6466,18 +6188,13 @@ async def test_async_data_generator_preserves_non_raw_sse_like_bytes(): with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) - yielded_text = [ - chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk - for chunk in yielded_data - ] + yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data] assert yielded_text[0] == gemini_event.decode("utf-8") assert yielded_text[1] == gemini_event_without_terminator.decode("utf-8") + "\n\n" - assert yielded_text[2] == f'data: {raw_payload.decode("utf-8")}\n\n' + assert yielded_text[2] == f"data: {raw_payload.decode('utf-8')}\n\n" assert "b'data:" not in "".join(yielded_text) assert yielded_text[-1] == "data: [DONE]\n\n" @@ -6500,12 +6217,8 @@ async def test_async_data_generator_buffers_split_google_native_sse_json_frame() ) raw_chunks = [ payload[:2].encode("utf-8"), - payload[ - 2 : payload.index("thoughtSignature") + len('thoughtSignature": "abc') - ].encode("utf-8"), - payload[ - payload.index("thoughtSignature") + len('thoughtSignature": "abc') : - ].encode("utf-8"), + payload[2 : payload.index("thoughtSignature") + len('thoughtSignature": "abc')].encode("utf-8"), + payload[payload.index("thoughtSignature") + len('thoughtSignature": "abc') :].encode("utf-8"), ] class MockStream: @@ -6532,15 +6245,10 @@ async def test_async_data_generator_buffers_split_google_native_sse_json_frame() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) - yielded_text = [ - chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk - for chunk in yielded_data - ] + yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data] assert yielded_text == [payload] for chunk in yielded_text: @@ -6586,15 +6294,10 @@ async def test_async_data_generator_flushes_raw_sse_stream_without_trailing_deli patch.object(ProxyLogging, "_fire_deferred_stream_logging"), ): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) - yielded_text = [ - chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk - for chunk in yielded_data - ] + yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data] assert len(yielded_text) == 1 assert yielded_text[0] == 'data: {"candidates": [{"content": "unterminated"}]\n\n' assert "[DONE]" not in yielded_text[0] @@ -6641,15 +6344,10 @@ async def test_async_data_generator_errors_when_raw_sse_frame_exceeds_buffer_lim patch.object(ProxyLogging, "_fire_deferred_stream_logging"), ): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) - yielded_text = [ - chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk - for chunk in yielded_data - ] + yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data] assert len(yielded_text) == 1 assert "maximum buffered size" in yielded_text[0] assert "[DONE]" not in yielded_text[0] @@ -6702,15 +6400,10 @@ async def test_async_data_generator_checks_raw_sse_buffer_limit_after_complete_f patch.object(ProxyLogging, "_fire_deferred_stream_logging"), ): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) - yielded_text = [ - chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk - for chunk in yielded_data - ] + yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data] assert yielded_text[0] == complete_frame assert yielded_text[1] == partial_frame + "\n\n" assert "[DONE]" not in "".join(yielded_text) @@ -6731,9 +6424,7 @@ async def test_async_data_generator_google_genai_stream_omits_openai_done(): "model": "gemini-2.0-flash", "_litellm_skip_openai_stream_done": True, } - gemini_event = ( - b'data: {"candidates": [{"content": {"parts": [{"text": "Hi"}]}}]}\n\n' - ) + gemini_event = b'data: {"candidates": [{"content": {"parts": [{"text": "Hi"}]}}]}\n\n' class MockStream: def __aiter__(self): @@ -6758,15 +6449,10 @@ async def test_async_data_generator_google_genai_stream_omits_openai_done(): with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) - yielded_text = [ - chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk - for chunk in yielded_data - ] + yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data] assert yielded_text == [gemini_event.decode("utf-8")] assert "[DONE]" not in "".join(yielded_text) @@ -6855,15 +6541,10 @@ async def test_async_data_generator_google_genai_stream_forwards_error_without_d with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) - yielded_text = [ - chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk - for chunk in yielded_data - ] + yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data] assert yielded_text == [error_sse] assert "[DONE]" not in "".join(yielded_text) @@ -6893,9 +6574,7 @@ async def test_async_data_generator_cleanup_on_normal_completion(): for chunk in mock_chunks: yield chunk - mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = ( - mock_streaming_iterator - ) + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = mock_streaming_iterator mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock( side_effect=lambda **kwargs: kwargs.get("response") ) @@ -6906,9 +6585,7 @@ async def test_async_data_generator_cleanup_on_normal_completion(): with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) # Should have completed normally with [DONE] @@ -6939,9 +6616,7 @@ async def test_async_data_generator_cleanup_on_midstream_error(): yield {"choices": [{"delta": {"content": "Hello"}}]} raise RuntimeError("upstream connection reset") - mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = ( - mock_streaming_iterator_with_error - ) + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = mock_streaming_iterator_with_error mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock( side_effect=lambda **kwargs: kwargs.get("response") ) @@ -6952,9 +6627,7 @@ async def test_async_data_generator_cleanup_on_midstream_error(): with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): yielded_data = [] - async for data in async_data_generator( - mock_response, mock_user_api_key_dict, mock_request_data - ): + async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) # Should have yielded data chunk and then an error chunk @@ -7009,9 +6682,7 @@ async def test_update_general_settings_store_model_in_db_true(): patch("litellm.proxy.proxy_server.store_model_in_db", False) as mock_store, patch("litellm.proxy.proxy_server.general_settings", {}) as mock_gs, ): - await proxy_config._update_general_settings( - db_general_settings={"store_model_in_db": True} - ) + await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": True}) import litellm.proxy.proxy_server as ps @@ -7033,9 +6704,7 @@ async def test_update_general_settings_store_model_in_db_false(): patch("litellm.proxy.proxy_server.store_model_in_db", True), patch("litellm.proxy.proxy_server.general_settings", {}), ): - await proxy_config._update_general_settings( - db_general_settings={"store_model_in_db": False} - ) + await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": False}) import litellm.proxy.proxy_server as ps @@ -7116,9 +6785,7 @@ async def test_update_general_settings_store_model_in_db_string_normalization(): patch("litellm.proxy.proxy_server.store_model_in_db", False), patch("litellm.proxy.proxy_server.general_settings", {}), ): - await proxy_config._update_general_settings( - db_general_settings={"store_model_in_db": "true"} - ) + await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": "true"}) import litellm.proxy.proxy_server as ps assert ps.store_model_in_db is True @@ -7128,9 +6795,7 @@ async def test_update_general_settings_store_model_in_db_string_normalization(): patch("litellm.proxy.proxy_server.store_model_in_db", False), patch("litellm.proxy.proxy_server.general_settings", {}), ): - await proxy_config._update_general_settings( - db_general_settings={"store_model_in_db": "True"} - ) + await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": "True"}) import litellm.proxy.proxy_server as ps assert ps.store_model_in_db is True @@ -7140,9 +6805,7 @@ async def test_update_general_settings_store_model_in_db_string_normalization(): patch("litellm.proxy.proxy_server.store_model_in_db", True), patch("litellm.proxy.proxy_server.general_settings", {}), ): - await proxy_config._update_general_settings( - db_general_settings={"store_model_in_db": "false"} - ) + await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": "false"}) import litellm.proxy.proxy_server as ps assert ps.store_model_in_db is False @@ -7163,9 +6826,7 @@ async def test_update_general_settings_store_model_in_db_none_keeps_current(): patch("litellm.proxy.proxy_server.store_model_in_db", True), patch("litellm.proxy.proxy_server.general_settings", {}), ): - await proxy_config._update_general_settings( - db_general_settings={"store_model_in_db": None} - ) + await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": None}) import litellm.proxy.proxy_server as ps assert ps.store_model_in_db is True @@ -7175,9 +6836,7 @@ async def test_update_general_settings_store_model_in_db_none_keeps_current(): patch("litellm.proxy.proxy_server.store_model_in_db", False), patch("litellm.proxy.proxy_server.general_settings", {}), ): - await proxy_config._update_general_settings( - db_general_settings={"store_model_in_db": None} - ) + await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": None}) import litellm.proxy.proxy_server as ps assert ps.store_model_in_db is False @@ -7197,12 +6856,11 @@ async def test_store_model_in_db_db_override_when_config_false(): # Mock DB returning store_model_in_db=True in general_settings mock_db_record = MagicMock() mock_db_record.param_value = {"store_model_in_db": True} - mock_prisma_client.db.litellm_config.find_first = AsyncMock( - return_value=mock_db_record - ) + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_record) mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() with ( @@ -7245,6 +6903,7 @@ async def test_store_model_in_db_db_check_skipped_when_already_true(monkeypatch) mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() with ( @@ -7283,12 +6942,11 @@ async def test_store_model_in_db_db_failure_graceful(monkeypatch): mock_prisma_client = MagicMock() # Simulate DB failure - mock_prisma_client.db.litellm_config.find_first = AsyncMock( - side_effect=Exception("DB connection error") - ) + mock_prisma_client.db.litellm_config.find_first = AsyncMock(side_effect=Exception("DB connection error")) mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() with ( @@ -7423,9 +7081,7 @@ async def test_increment_spend_counters_initializes_and_increments(): ) # Counter should be: base(5.0) + increment(0.50) = 5.50 - counter = counter_cache.in_memory_cache.get_cache( - key=f"spend:key:{hashed_token}" - ) + counter = counter_cache.in_memory_cache.get_cache(key=f"spend:key:{hashed_token}") assert counter == 5.50 # Second increment — counter already exists, just increment @@ -7436,9 +7092,7 @@ async def test_increment_spend_counters_initializes_and_increments(): response_cost=0.25, ) - counter = counter_cache.in_memory_cache.get_cache( - key=f"spend:key:{hashed_token}" - ) + counter = counter_cache.in_memory_cache.get_cache(key=f"spend:key:{hashed_token}") assert counter == 5.75 finally: ps.user_api_key_cache = original_key_cache @@ -7484,9 +7138,7 @@ async def test_increment_spend_counters_team_and_member(): team_counter = counter_cache.in_memory_cache.get_cache(key="spend:team:team-1") assert team_counter == 2.30 - member_counter = counter_cache.in_memory_cache.get_cache( - key="spend:team_member:user-1:team-1" - ) + member_counter = counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") assert member_counter == 1.30 finally: ps.user_api_key_cache = original_key_cache @@ -7544,14 +7196,10 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss( increment=1.5, ) - fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with( - where={"team_id": "team-9"} - ) + fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-9"}) # Seed uses SET NX with db_spend (42) — cross-pod safe, no INCR of 42. # Only the per-request delta (1.5) goes through INCRBYFLOAT. - fake_redis.async_set_cache.assert_awaited_once_with( - key="spend:team:team-9", value=42.0, nx=True - ) + fake_redis.async_set_cache.assert_awaited_once_with(key="spend:team:team-9", value=42.0, nx=True) writes = [(c["key"], c["value"]) for c in recorded_increments] assert writes == [("spend:team:team-9", 1.5)] finally: @@ -7620,9 +7268,7 @@ async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed( return row fake_prisma = MagicMock() - fake_prisma.db.litellm_teamtable.find_unique = AsyncMock( - side_effect=slow_find_unique - ) + fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(side_effect=slow_find_unique) pod_a = DualCache() pod_a.redis_cache = fake_redis @@ -7655,11 +7301,7 @@ async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed( # (winner) and one was rejected (loser). assert db_read_count == 2 assert fake_redis.async_set_cache.await_count == 2 - nx_writes = [ - call - for call in fake_redis.async_set_cache.await_args_list - if call.kwargs.get("nx") is True - ] + nx_writes = [call for call in fake_redis.async_set_cache.await_args_list if call.kwargs.get("nx") is True] assert len(nx_writes) == 2 assert sorted(set_results) == [ False, @@ -7668,9 +7310,7 @@ async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed( # Loser path executed: after the winner's SET NX returned True, the # losing coalesced() call falls back to async_get_cache to read the # winner's value rather than re-seeding. - assert ( - get_after_set_count >= 1 - ), "loser branch (else: read back winner's value) was never exercised" + assert get_after_set_count >= 1, "loser branch (else: read back winner's value) was never exercised" @pytest.mark.asyncio @@ -7692,14 +7332,10 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): fake_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) fake_prisma.db.litellm_endusertable.find_unique = AsyncMock() fake_prisma.db.litellm_tagtable.find_unique = AsyncMock() - fake_prisma.db.litellm_organizationtable.find_unique = AsyncMock( - return_value=org_row - ) + fake_prisma.db.litellm_organizationtable.find_unique = AsyncMock(return_value=org_row) assert await SpendCounterReseed.from_db(fake_prisma, "spend:user:alice") == 17.0 - fake_prisma.db.litellm_usertable.find_unique.assert_awaited_once_with( - where={"user_id": "alice"} - ) + fake_prisma.db.litellm_usertable.find_unique.assert_awaited_once_with(where={"user_id": "alice"}) assert ( await SpendCounterReseed.from_db( @@ -7714,9 +7350,7 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): fake_prisma.db.litellm_tagtable.find_unique.assert_not_awaited() assert await SpendCounterReseed.from_db(fake_prisma, "spend:org:acme") == 305.0 - fake_prisma.db.litellm_organizationtable.find_unique.assert_awaited_once_with( - where={"organization_id": "acme"} - ) + fake_prisma.db.litellm_organizationtable.find_unique.assert_awaited_once_with(where={"organization_id": "acme"}) @pytest.mark.asyncio @@ -7730,14 +7364,8 @@ async def test_reseed_spend_from_db_skips_window_variant_keys(): fake_prisma.db.litellm_verificationtoken.find_unique = AsyncMock() fake_prisma.db.litellm_teamtable.find_unique = AsyncMock() - assert ( - await SpendCounterReseed.from_db(fake_prisma, "spend:key:sk-abc:window:1h") - is None - ) - assert ( - await SpendCounterReseed.from_db(fake_prisma, "spend:team:team-1:window:1d") - is None - ) + assert await SpendCounterReseed.from_db(fake_prisma, "spend:key:sk-abc:window:1h") is None + assert await SpendCounterReseed.from_db(fake_prisma, "spend:team:team-1:window:1d") is None fake_prisma.db.litellm_verificationtoken.find_unique.assert_not_awaited() fake_prisma.db.litellm_teamtable.find_unique.assert_not_awaited() @@ -7773,9 +7401,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss(): where={"api_key": "key-window", "startTime": {"gte": window_start}}, sum={"spend": True}, ) - assert counter_cache.in_memory_cache.get_cache( - key="spend:key:key-window:window:1h" - ) == pytest.approx(2.75) + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-window:window:1h") == pytest.approx(2.75) finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma @@ -7830,14 +7456,10 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory(): increment=1.5, ) - fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with( - where={"team_id": "team-stale-local"} - ) + fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-stale-local"}) # Seed via SET NX (42) + delta via INCRBYFLOAT (1.5) = 43.5. assert redis_store[counter_key] == pytest.approx(43.5) - assert counter_cache.in_memory_cache.get_cache( - key=counter_key - ) == pytest.approx(43.5) + assert counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(43.5) finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma @@ -7900,9 +7522,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): sum={"spend": True}, ) assert redis_store[counter_key] == pytest.approx(2.75) - assert counter_cache.in_memory_cache.get_cache( - key=counter_key - ) == pytest.approx(2.75) + assert counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(2.75) finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma @@ -7938,9 +7558,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed() fake_prisma = MagicMock() fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( - return_value=[ - {"api_key": "key-window-concurrent-seed", "_sum": {"spend": 2.25}} - ] + return_value=[{"api_key": "key-window-concurrent-seed", "_sum": {"spend": 2.25}}] ) import litellm.proxy.proxy_server as ps @@ -7963,9 +7581,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed() nx=True, ) assert redis_store[counter_key] == pytest.approx(3.25) - assert counter_cache.in_memory_cache.get_cache( - key=counter_key - ) == pytest.approx(3.25) + assert counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(3.25) finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma @@ -7991,12 +7607,7 @@ async def test_window_spend_counter_skips_invalid_window_start(): increment=0.5, ) - assert ( - counter_cache.in_memory_cache.get_cache( - key="spend:key:key-invalid-window:window:not-a-duration" - ) - is None - ) + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-invalid-window:window:not-a-duration") is None finally: ps.spend_counter_cache = orig_counter @@ -8078,9 +7689,9 @@ async def test_increment_spend_counters_finalizes_after_unreserved_increments(): assert incremented_counters == ["spend:team:team-finalize-after-increments"] assert budget_reservation["finalized"] is True - assert counter_cache.in_memory_cache.get_cache( - key="spend:key:key-finalize-after-increments" - ) == pytest.approx(0.25) + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-finalize-after-increments") == pytest.approx( + 0.25 + ) finally: ps.spend_counter_cache = orig_counter ps.user_api_key_cache = orig_user @@ -8124,9 +7735,7 @@ async def test_increment_spend_counters_finalizes_none_cost_reservation(): ) assert budget_reservation["finalized"] is True - assert counter_cache.in_memory_cache.get_cache( - key="spend:key:key-finalize-none-cost" - ) == pytest.approx(0.0) + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-finalize-none-cost") == pytest.approx(0.0) finally: ps.spend_counter_cache = orig_counter @@ -8176,9 +7785,7 @@ async def test_increment_spend_counters_reseeds_from_db_on_bad_reserved_counter( assert budget_reservation["finalized"] is True # counter reseeded to the authoritative DB value, not deleted/left None # and not double-counted via a direct increment - assert counter_cache.in_memory_cache.get_cache( - key="spend:key:key-bad-reserved-counter" - ) == pytest.approx(0.6) + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-bad-reserved-counter") == pytest.approx(0.6) finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma @@ -8207,12 +7814,8 @@ async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure( increment=0.5, ) - assert ( - counter_cache.in_memory_cache.get_cache(key="spend:team:redis-fail") is None - ) - fake_redis.async_delete_cache.assert_awaited_once_with( - key="spend:team:redis-fail" - ) + assert counter_cache.in_memory_cache.get_cache(key="spend:team:redis-fail") is None + fake_redis.async_delete_cache.assert_awaited_once_with(key="spend:team:redis-fail") finally: ps.spend_counter_cache = orig_counter @@ -8258,16 +7861,13 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing(): fallback_spend=30.0, ) assert spend == 362.0, ( - f"expected DB reseed to return 362.0, got {spend} " - f"(fallback would have returned 30.0 and caused bypass)" + f"expected DB reseed to return 362.0, got {spend} (fallback would have returned 30.0 and caused bypass)" ) # Counter warmed via SET NX so subsequent reads are fast. assert ("spend:team_member:user-1:team-1", 362.0, True) in [ (s["key"], s["value"], s["nx"]) for s in recorded_seeds ] - assert counter_cache.in_memory_cache.get_cache( - key="spend:team_member:user-1:team-1" - ) == pytest.approx(362.0) + assert counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == pytest.approx(362.0) finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma @@ -8352,9 +7952,7 @@ async def test_get_current_spend_coalesces_concurrent_reseeds(): counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() - fake_prisma.db.litellm_teammembership.find_unique = AsyncMock( - side_effect=slow_find_unique - ) + fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=slow_find_unique) import litellm.proxy.proxy_server as ps @@ -8363,15 +7961,10 @@ async def test_get_current_spend_coalesces_concurrent_reseeds(): ps.prisma_client = fake_prisma try: results = await _asyncio.gather( - *[ - get_current_spend(counter_key=counter_key, fallback_spend=0.0) - for _ in range(5) - ] + *[get_current_spend(counter_key=counter_key, fallback_spend=0.0) for _ in range(5)] ) assert results == [100.0] * 5, f"all callers should see DB value, got {results}" - assert ( - db_call_count == 1 - ), f"expected exactly 1 DB query for 5 concurrent reseeds, got {db_call_count}" + assert db_call_count == 1, f"expected exactly 1 DB query for 5 concurrent reseeds, got {db_call_count}" finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma @@ -8408,9 +8001,7 @@ async def test_get_current_spend_uses_db_zero_over_stale_fallback(): counter_key="spend:team_member:user-1:team-after-reset", fallback_spend=42.0, ) - assert ( - spend == 0.0 - ), f"DB authoritative 0 must override stale fallback 42, got {spend}" + assert spend == 0.0, f"DB authoritative 0 must override stale fallback 42, got {spend}" finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma @@ -8468,9 +8059,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query(): counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() - fake_prisma.db.litellm_teammembership.find_unique = AsyncMock( - side_effect=slow_find_unique - ) + fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=slow_find_unique) import litellm.proxy.proxy_server as ps @@ -8492,9 +8081,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query(): ), get_current_spend(counter_key=counter_key, fallback_spend=0.0), ) - assert ( - db_call_count == 1 - ), f"expected 1 DB query for concurrent read+write+read, got {db_call_count}" + assert db_call_count == 1, f"expected 1 DB query for concurrent read+write+read, got {db_call_count}" # Read-path callers see the warmed counter; the write path's # increment may or may not have landed by then, so accept either # the seeded value or seeded+increment. @@ -8530,9 +8117,7 @@ async def test_reseed_locks_dict_is_bounded(): try: for i in range(7): await SpendCounterReseed._get_lock(f"spend:key:test-key-{i}") - assert ( - len(SpendCounterReseed._locks) == 5 - ), f"got {len(SpendCounterReseed._locks)}" + assert len(SpendCounterReseed._locks) == 5, f"got {len(SpendCounterReseed._locks)}" # Oldest two evicted assert "spend:key:test-key-0" not in SpendCounterReseed._locks assert "spend:key:test-key-1" not in SpendCounterReseed._locks @@ -8589,9 +8174,7 @@ async def test_reseed_warms_cache_even_on_zero_db_spend(): return row fake_prisma = MagicMock() - fake_prisma.db.litellm_teammembership.find_unique = AsyncMock( - side_effect=find_unique - ) + fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=find_unique) import litellm.proxy.proxy_server as ps @@ -8604,9 +8187,7 @@ async def test_reseed_warms_cache_even_on_zero_db_spend(): # Second call: cache should be warmed at 0, no second DB query. spend2 = await get_current_spend(counter_key=counter_key, fallback_spend=0.0) assert spend1 == 0.0 and spend2 == 0.0 - assert ( - db_call_count == 1 - ), f"second read should hit warmed cache, got {db_call_count} DB queries" + assert db_call_count == 1, f"second read should hit warmed cache, got {db_call_count} DB queries" assert redis_store.get(counter_key) == 0.0, "cache must be warmed at 0" finally: ps.spend_counter_cache = orig_counter @@ -8669,9 +8250,7 @@ def _update_config_setup(monkeypatch): def _install(initial_rows=None, store_model_in_db=True): prisma = _FakePrismaClient(initial_rows=initial_rows) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) - monkeypatch.setattr( - "litellm.proxy.proxy_server.store_model_in_db", store_model_in_db - ) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", store_model_in_db) monkeypatch.setattr( "litellm.proxy.proxy_server.encrypt_value_helper", lambda value, **_: f"enc:{value}", @@ -8682,9 +8261,7 @@ def _update_config_setup(monkeypatch): ) from litellm.proxy.proxy_server import proxy_config as real_proxy_config - monkeypatch.setattr( - real_proxy_config, "add_deployment", AsyncMock(return_value=None) - ) + monkeypatch.setattr(real_proxy_config, "add_deployment", AsyncMock(return_value=None)) original_overrides = app.dependency_overrides.copy() app.dependency_overrides[auth_dep] = lambda: UserAPIKeyAuth( @@ -8719,19 +8296,13 @@ def test_update_config_writes_only_sent_section(_update_config_setup): assert resp.status_code == 200 written = {name for name, _ in prisma.db.litellm_config.upsert_calls} assert written == {"general_settings"} - assert prisma.db.litellm_config.rows["litellm_settings"] == { - "drop_params": True - } - assert prisma.db.litellm_config.rows["environment_variables"] == { - "FOO": "enc:bar" - } + assert prisma.db.litellm_config.rows["litellm_settings"] == {"drop_params": True} + assert prisma.db.litellm_config.rows["environment_variables"] == {"FOO": "enc:bar"} finally: restore() -def test_update_config_env_var_round_trip_not_double_encrypted( - _update_config_setup, monkeypatch -): +def test_update_config_env_var_round_trip_not_double_encrypted(_update_config_setup, monkeypatch): """Endpoint-level regression for the /config/update double-encryption bug. The Admin UI reads config back via /get/config/callbacks (which returns @@ -8744,16 +8315,12 @@ def test_update_config_env_var_round_trip_not_double_encrypted( code this stored "enc:enc:..."; the assertions below would fail there. """ - def _fake_decrypt( - value, key=None, exception_type="error", return_original_value=False - ): + def _fake_decrypt(value, key=None, exception_type="error", return_original_value=False): if isinstance(value, str) and value.startswith("enc:"): return value[len("enc:") :] return value if return_original_value else None - monkeypatch.setattr( - "litellm.proxy.proxy_server.decrypt_value_helper", _fake_decrypt - ) + monkeypatch.setattr("litellm.proxy.proxy_server.decrypt_value_helper", _fake_decrypt) client, prisma, restore = _update_config_setup( initial_rows={"environment_variables": {"PREEXISTING_KEY": "enc:keepme"}} @@ -8771,21 +8338,14 @@ def test_update_config_env_var_round_trip_not_double_encrypted( # UI round-trip: re-POST the stored ciphertext (no field change). resp = client.post( "/config/update", - json={ - "environment_variables": { - "LANGFUSE_SECRET_KEY": stored["LANGFUSE_SECRET_KEY"] - } - }, + json={"environment_variables": {"LANGFUSE_SECRET_KEY": stored["LANGFUSE_SECRET_KEY"]}}, ) assert resp.status_code == 200 stored = prisma.db.litellm_config.rows["environment_variables"] # The bug: this would be "enc:enc:sk-secret". The fix keeps it single. assert stored["LANGFUSE_SECRET_KEY"] == "enc:sk-secret" - assert ( - _fake_decrypt(stored["LANGFUSE_SECRET_KEY"], return_original_value=True) - == "sk-secret" - ) + assert _fake_decrypt(stored["LANGFUSE_SECRET_KEY"], return_original_value=True) == "sk-secret" # Untouched key preserved byte-for-byte (only sent keys rewritten). assert stored["PREEXISTING_KEY"] == "enc:keepme" @@ -8800,14 +8360,9 @@ def test_update_config_can_flip_store_model_in_db_when_currently_false( False, blocking the very request that would flip it to True.""" client, prisma, restore = _update_config_setup(store_model_in_db=False) try: - resp = client.post( - "/config/update", json={"general_settings": {"store_model_in_db": True}} - ) + resp = client.post("/config/update", json={"general_settings": {"store_model_in_db": True}}) assert resp.status_code == 200 - assert ( - prisma.db.litellm_config.rows["general_settings"]["store_model_in_db"] - is True - ) + assert prisma.db.litellm_config.rows["general_settings"]["store_model_in_db"] is True finally: restore() @@ -8840,9 +8395,7 @@ def test_update_config_litellm_settings_request_wins_for_non_callback_keys( } ) try: - resp = client.post( - "/config/update", json={"litellm_settings": {"drop_params": False}} - ) + resp = client.post("/config/update", json={"litellm_settings": {"drop_params": False}}) assert resp.status_code == 200 stored = prisma.db.litellm_config.rows["litellm_settings"] assert stored["drop_params"] is False @@ -8938,9 +8491,7 @@ class TestLazyFeaturesNotImportedAtStartup: from litellm.proxy._lazy_features import LAZY_FEATURES - proxy_server_src = ( - Path(__file__).resolve().parents[3] / "litellm/proxy/proxy_server.py" - ).read_text() + proxy_server_src = (Path(__file__).resolve().parents[3] / "litellm/proxy/proxy_server.py").read_text() leaks = [] for feat in LAZY_FEATURES: @@ -9045,9 +8596,7 @@ class TestLazyFeatureMiddleware: ("/api/v1", "/api/v1/unrelated", False, "unrelated path under root"), ], ) - async def test_root_path_handling( - self, monkeypatch, server_root_path, request_path, should_load, case - ): + async def test_root_path_handling(self, monkeypatch, server_root_path, request_path, should_load, case): """ The middleware must strip SERVER_ROOT_PATH before prefix-matching so lazy features load under deployments that set a server root path, @@ -9157,9 +8706,7 @@ class TestLazyFeatureMiddleware: ) await asyncio.gather(hit(), hit(), hit(), hit(), hit()) - assert loads == [ - "json" - ], f"expected one registration despite concurrent first hits, got {loads}" + assert loads == ["json"], f"expected one registration despite concurrent first hits, got {loads}" @pytest.mark.asyncio async def test_failing_import_does_not_loop(self): @@ -9209,9 +8756,9 @@ class TestLazyFeatureMiddleware: receive, send, ) - assert attempts == [ - "called" - ], f"failing register_fn should be invoked once, not on every request; got {attempts}" + assert attempts == ["called"], ( + f"failing register_fn should be invoked once, not on every request; got {attempts}" + ) @pytest.mark.asyncio @@ -9279,9 +8826,7 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory(): counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() - fake_prisma.db.litellm_teammembership.find_unique = AsyncMock( - return_value=MagicMock(spend=999.0) - ) + fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=MagicMock(spend=999.0)) import litellm.proxy.proxy_server as ps @@ -9291,8 +8836,7 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory(): try: spend = await get_current_spend(counter_key=counter_key, fallback_spend=0.0) assert spend == 42.0, ( - f"expected in-memory fallback 42.0 on Redis error, got {spend} " - f"(should not have hit DB when Redis errored)" + f"expected in-memory fallback 42.0 on Redis error, got {spend} (should not have hit DB when Redis errored)" ) # DB query should NOT have fired - in-memory short-circuits. fake_prisma.db.litellm_teammembership.find_unique.assert_not_awaited() @@ -9315,9 +8859,7 @@ def test_realtime_websocket_route_aliases_registered(): from litellm.proxy.proxy_server import app from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes - websocket_paths = { - route.path for route in app.routes if isinstance(route, WebSocketRoute) - } + websocket_paths = {route.path for route in app.routes if isinstance(route, WebSocketRoute)} openai_routes = LiteLLMRoutes.openai_routes.value for expected in ("/openai/v1/realtime", "/v1/realtime", "/realtime"): @@ -9329,9 +8871,7 @@ def test_realtime_websocket_route_aliases_registered(): f"{expected!r} missing from LiteLLMRoutes.openai_routes; " f"non-admin / team / key-scoped users will get 403 on this path." ) - assert tuple(API_ROUTE_TO_CALL_TYPES.get(expected) or ()) == ( - CallTypes.arealtime, - ), ( + assert tuple(API_ROUTE_TO_CALL_TYPES.get(expected) or ()) == (CallTypes.arealtime,), ( f"{expected!r} missing from API_ROUTE_TO_CALL_TYPES; call-type " f"resolution will return None and break call-type-aware features." ) @@ -9381,8 +8921,7 @@ class TestTransformRequestBannedParams: }, ) assert response.status_code == 400, ( - f"Expected 400 for banned param '{banned}', " - f"got {response.status_code}: {response.json()}" + f"Expected 400 for banned param '{banned}', got {response.status_code}: {response.json()}" ) @@ -9408,13 +8947,8 @@ class TestSortModelsByDisplayName: {"model_name": "gpt-4o", "model_info": {}}, ] - sorted_models = _sort_models( - all_models=models, sort_by="model_name", sort_order="asc" - ) - displayed_order = [ - m["model_info"].get("team_public_model_name") or m["model_name"] - for m in sorted_models - ] + sorted_models = _sort_models(all_models=models, sort_by="model_name", sort_order="asc") + displayed_order = [m["model_info"].get("team_public_model_name") or m["model_name"] for m in sorted_models] assert displayed_order == [ "anthropic/claude", "claude-haiku-4-5", @@ -9433,13 +8967,8 @@ class TestSortModelsByDisplayName: {"model_name": "gpt-4o", "model_info": {}}, ] - sorted_models = _sort_models( - all_models=models, sort_by="model_name", sort_order="desc" - ) - displayed_order = [ - m["model_info"].get("team_public_model_name") or m["model_name"] - for m in sorted_models - ] + sorted_models = _sort_models(all_models=models, sort_by="model_name", sort_order="desc") + displayed_order = [m["model_info"].get("team_public_model_name") or m["model_name"] for m in sorted_models] assert displayed_order == [ "zeta/model", "gpt-4o", @@ -9457,9 +8986,7 @@ class TestSortModelsByDisplayName: {"model_name": "beta", "model_info": {}}, ] - sorted_models = _sort_models( - all_models=models, sort_by="model_name", sort_order="asc" - ) + sorted_models = _sort_models(all_models=models, sort_by="model_name", sort_order="asc") assert [m["model_name"] for m in sorted_models] == ["alpha", "beta"] @@ -9481,9 +9008,7 @@ class TestDeleteDeploymentSync: mock_router.delete_deployment.return_value = MagicMock() with patch("litellm.proxy.proxy_server.llm_router", mock_router): - with patch.object( - proxy_config, "get_config", AsyncMock(return_value={"model_list": []}) - ): + with patch.object(proxy_config, "get_config", AsyncMock(return_value={"model_list": []})): still_desired = await proxy_config._delete_deployment(db_models=[]) mock_router.delete_deployment.assert_called_once_with(id="model-id-to-evict") @@ -9507,9 +9032,7 @@ class TestDeleteDeploymentSync: with patch("litellm.proxy.proxy_server.llm_router", mock_router): with patch.object(proxy_config, "get_config", AsyncMock(return_value={})): - await proxy_config._update_llm_router( - new_models=None, proxy_logging_obj=MagicMock() - ) + await proxy_config._update_llm_router(new_models=None, proxy_logging_obj=MagicMock()) mock_router.delete_deployment.assert_not_called() mock_router.upsert_deployment.assert_not_called() @@ -9526,15 +9049,11 @@ class TestDeleteDeploymentSync: proxy_config = ProxyConfig() mock_prisma = MagicMock() - mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock( - side_effect=Exception("DB connection lost") - ) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=Exception("DB connection lost")) result = await proxy_config._get_models_from_db(prisma_client=mock_prisma) - assert ( - result is None - ), f"Expected None on DB failure to signal fetch error, got {result!r}" + assert result is None, f"Expected None on DB failure to signal fetch error, got {result!r}" def test_get_config_list_includes_cancel_on_disconnect(monkeypatch): @@ -9816,9 +9335,18 @@ def test_general_settings_ui_defaults_unchanged_for_existing_fields(): _general_settings_ui_litellm_default, ) - assert _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["budget_exceeded_throttle_percentage"]) is None - assert _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["enable_anthropic_prompt_caching"]) is False - assert _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["anthropic_prompt_caching_ttl"]) is None + assert ( + _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["budget_exceeded_throttle_percentage"]) + is None + ) + assert ( + _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["enable_anthropic_prompt_caching"]) + is False + ) + assert ( + _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["anthropic_prompt_caching_ttl"]) + is None + ) @pytest.mark.parametrize( @@ -10084,16 +9612,10 @@ def test_preserve_redacted_plugin_keys_keeps_stored_credential(): existing = [{"name": "p1", "url": "https://p1", "plugin_key": "sk-real-1"}] - redacted = _preserve_redacted_plugin_keys( - [{"name": "p1", "url": "https://p1-new", "plugin_key": "***"}], existing - ) - assert redacted == [ - {"name": "p1", "url": "https://p1-new", "plugin_key": "sk-real-1"} - ] + redacted = _preserve_redacted_plugin_keys([{"name": "p1", "url": "https://p1-new", "plugin_key": "***"}], existing) + assert redacted == [{"name": "p1", "url": "https://p1-new", "plugin_key": "sk-real-1"}] - blanked = _preserve_redacted_plugin_keys( - [{"name": "p1", "url": "https://p1", "plugin_key": ""}], existing - ) + blanked = _preserve_redacted_plugin_keys([{"name": "p1", "url": "https://p1", "plugin_key": ""}], existing) assert blanked[0]["plugin_key"] == "sk-real-1" @@ -10103,14 +9625,10 @@ def test_preserve_redacted_plugin_keys_sets_new_and_drops_orphan_placeholder(): existing = [{"name": "p1", "url": "https://p1", "plugin_key": "sk-real-1"}] - rotated = _preserve_redacted_plugin_keys( - [{"name": "p1", "url": "https://p1", "plugin_key": "sk-new"}], existing - ) + rotated = _preserve_redacted_plugin_keys([{"name": "p1", "url": "https://p1", "plugin_key": "sk-new"}], existing) assert rotated[0]["plugin_key"] == "sk-new" - new_plugin = _preserve_redacted_plugin_keys( - [{"name": "p2", "url": "https://p2", "plugin_key": "***"}], existing - ) + new_plugin = _preserve_redacted_plugin_keys([{"name": "p2", "url": "https://p2", "plugin_key": "***"}], existing) assert "plugin_key" not in new_plugin[0] @@ -10143,9 +9661,7 @@ def _config_field_info_client(monkeypatch, user_role): mock_prisma = MagicMock() mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) monkeypatch.setattr(ps, "prisma_client", mock_prisma) - app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( - user_id="u", user_role=user_role - ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="u", user_role=user_role) return TestClient(app) @@ -10156,9 +9672,7 @@ def test_config_field_info_redacts_secrets_for_view_only_admin(monkeypatch): is not a FULL PROXY_ADMIN, while non-secret fields stay readable.""" from litellm.proxy._types import LitellmUserRoles - client = _config_field_info_client( - monkeypatch, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY - ) + client = _config_field_info_client(monkeypatch, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) try: for secret_field in ("master_key", "database_url", "pass_through_endpoints"): resp = client.get("/config/field/info", params={"field_name": secret_field}) @@ -10168,9 +9682,7 @@ def test_config_field_info_redacts_secrets_for_view_only_admin(monkeypatch): assert "secret" not in str(body["field_value"]) assert "p4ssw0rd" not in str(body["field_value"]) - resp = client.get( - "/config/field/info", params={"field_name": "max_parallel_requests"} - ) + resp = client.get("/config/field/info", params={"field_name": "max_parallel_requests"}) assert resp.status_code == 200, resp.text assert resp.json()["field_value"] == 100 finally: @@ -10188,14 +9700,9 @@ def test_config_field_info_returns_raw_secrets_for_full_admin(monkeypatch): assert resp.status_code == 200, resp.text assert resp.json()["field_value"] == "sk-super-secret-master" - resp = client.get( - "/config/field/info", params={"field_name": "pass_through_endpoints"} - ) + resp = client.get("/config/field/info", params={"field_name": "pass_through_endpoints"}) assert resp.status_code == 200, resp.text - assert ( - resp.json()["field_value"][0]["headers"]["Authorization"] - == "Bearer sk-upstream-secret" - ) + assert resp.json()["field_value"][0]["headers"]["Authorization"] == "Bearer sk-upstream-secret" finally: app.dependency_overrides.clear() @@ -10437,9 +9944,7 @@ async def test_delete_config_general_settings_emits_deleted_audit_log(monkeypatc user_role=LitellmUserRoles.PROXY_ADMIN, ) await delete_config_general_settings( - data=ConfigFieldDelete( - field_name="max_parallel_requests", config_type="general_settings" - ), + data=ConfigFieldDelete(field_name="max_parallel_requests", config_type="general_settings"), user_api_key_dict=admin, ) # Audit is scheduled via asyncio.create_task; yield so it runs. @@ -10462,9 +9967,7 @@ def test_update_config_audits_every_written_section(_update_config_setup, monkey is the row that holds default_internal_user_params ("default user settings").""" import litellm.proxy.proxy_server as proxy_server_module - client, prisma, restore = _update_config_setup( - initial_rows={"litellm_settings": {"drop_params": True}} - ) + client, prisma, restore = _update_config_setup(initial_rows={"litellm_settings": {"drop_params": True}}) audit_create = AsyncMock() prisma.db.litellm_auditlog.create = audit_create monkeypatch.setattr(proxy_server_module, "premium_user", True) @@ -10475,17 +9978,14 @@ def test_update_config_audits_every_written_section(_update_config_setup, monkey json={ "general_settings": {"store_prompts_in_spend_logs": True}, "environment_variables": {"FOO": "bar"}, - "litellm_settings": { - "default_internal_user_params": {"max_budget": 10} - }, + "litellm_settings": {"default_internal_user_params": {"max_budget": 10}}, "router_settings": {"routing_strategy": "latency-based-routing"}, }, ) assert resp.status_code == 200, resp.text audited = { - call.kwargs["data"]["object_id"]: call.kwargs["data"]["action"] - for call in audit_create.await_args_list + call.kwargs["data"]["object_id"]: call.kwargs["data"]["action"] for call in audit_create.await_args_list } assert audited == { "general_settings": "updated", @@ -10497,20 +9997,14 @@ def test_update_config_audits_every_written_section(_update_config_setup, monkey assert call.kwargs["data"]["table_name"] == "LiteLLM_Config" assert call.kwargs["data"]["changed_by"] == "test_admin" - ls_call = next( - c - for c in audit_create.await_args_list - if c.kwargs["data"]["object_id"] == "litellm_settings" - ) + ls_call = next(c for c in audit_create.await_args_list if c.kwargs["data"]["object_id"] == "litellm_settings") after = json.loads(ls_call.kwargs["data"]["updated_values"]) assert after["default_internal_user_params"] == {"max_budget": 10} finally: restore() -def test_delete_callback_audits_litellm_settings_deletion( - _update_config_setup, monkeypatch -): +def test_delete_callback_audits_litellm_settings_deletion(_update_config_setup, monkeypatch): """/config/callback/delete must emit a deleted audit row for litellm_settings capturing the success_callback list before and after removal.""" import litellm.proxy.proxy_server as proxy_server_module @@ -10526,19 +10020,11 @@ def test_delete_callback_audits_litellm_settings_deletion( monkeypatch.setattr( real_proxy_config, "get_config", - AsyncMock( - return_value={ - "litellm_settings": {"success_callback": ["langfuse", "datadog"]} - } - ), - ) - monkeypatch.setattr( - real_proxy_config, "save_config", AsyncMock(return_value=None) + AsyncMock(return_value={"litellm_settings": {"success_callback": ["langfuse", "datadog"]}}), ) + monkeypatch.setattr(real_proxy_config, "save_config", AsyncMock(return_value=None)) try: - resp = client.post( - "/config/callback/delete", json={"callback_name": "datadog"} - ) + resp = client.post("/config/callback/delete", json={"callback_name": "datadog"}) assert resp.status_code == 200, resp.text audit_create.assert_awaited_once() @@ -10567,24 +10053,16 @@ def test_delete_callback_audits_before_reload_failure(_update_config_setup, monk monkeypatch.setattr( real_proxy_config, "get_config", - AsyncMock( - return_value={ - "litellm_settings": {"success_callback": ["langfuse", "datadog"]} - } - ), - ) - monkeypatch.setattr( - real_proxy_config, "save_config", AsyncMock(return_value=None) + AsyncMock(return_value={"litellm_settings": {"success_callback": ["langfuse", "datadog"]}}), ) + monkeypatch.setattr(real_proxy_config, "save_config", AsyncMock(return_value=None)) monkeypatch.setattr( real_proxy_config, "add_deployment", AsyncMock(side_effect=RuntimeError("reload failed")), ) try: - resp = client.post( - "/config/callback/delete", json={"callback_name": "datadog"} - ) + resp = client.post("/config/callback/delete", json={"callback_name": "datadog"}) assert resp.status_code == 500, resp.text audit_create.assert_awaited_once() @@ -10595,9 +10073,7 @@ def test_delete_callback_audits_before_reload_failure(_update_config_setup, monk restore() -def test_update_config_redacts_all_environment_variable_values( - _update_config_setup, monkeypatch -): +def test_update_config_redacts_all_environment_variable_values(_update_config_setup, monkeypatch): """environment_variables hold credentials under arbitrary uppercase keys (DATABASE_URL) that key-name secret matching misses, so every value in the section must be redacted before the audit row is written; a plaintext @@ -10607,11 +10083,7 @@ def test_update_config_redacts_all_environment_variable_values( # DATABASE_URL is the bug class: an uppercase env key that key-name secret # matching does NOT flag, so only whole-section value redaction protects it. client, prisma, restore = _update_config_setup( - initial_rows={ - "environment_variables": { - "DATABASE_URL": "enc:postgresql://OLDsecret@old.host:5432/db" - } - } + initial_rows={"environment_variables": {"DATABASE_URL": "enc:postgresql://OLDsecret@old.host:5432/db"}} ) audit_create = AsyncMock() prisma.db.litellm_auditlog.create = audit_create @@ -10630,9 +10102,7 @@ def test_update_config_redacts_all_environment_variable_values( assert resp.status_code == 200, resp.text env_call = next( - c - for c in audit_create.await_args_list - if c.kwargs["data"]["object_id"] == "environment_variables" + c for c in audit_create.await_args_list if c.kwargs["data"]["object_id"] == "environment_variables" ) data = env_call.kwargs["data"] @@ -10796,11 +10266,7 @@ def test_init_coordination_redis_startup_nodes_builds_cluster_client(): """A coordination_redis block with startup_nodes must construct a cluster client, so cluster-aware consumers (v3 rate limiter) take the cluster path.""" usage_cache, _, _ = _run_init_coordination_redis( - config={ - "general_settings": { - "coordination_redis": {"startup_nodes": [{"host": "node-1", "port": 7000}]} - } - }, + config={"general_settings": {"coordination_redis": {"startup_nodes": [{"host": "node-1", "port": 7000}]}}}, ) assert isinstance(usage_cache, _EnvBuiltClusterCache) @@ -11050,17 +10516,13 @@ async def _collect_async_data_generator_frames(request_data: dict) -> list: with patch.object(proxy_server_module.ProxyLogging, "_fire_deferred_stream_logging"): return [ frame.decode("utf-8") if isinstance(frame, bytes) else frame - async for frame in async_data_generator( - MockStream(), MagicMock(spec=UserAPIKeyAuth), request_data - ) + async for frame in async_data_generator(MockStream(), MagicMock(spec=UserAPIKeyAuth), request_data) ] @pytest.mark.asyncio async def test_async_data_generator_strips_injected_usage_chunk(): - frames = await _collect_async_data_generator_frames( - {"model": "gpt-5.4-nano", "_litellm_strip_stream_usage": True} - ) + frames = await _collect_async_data_generator_frames({"model": "gpt-5.4-nano", "_litellm_strip_stream_usage": True}) data_frames = [frame for frame in frames if frame.startswith("data: {")] assert len(data_frames) == 2 @@ -11138,9 +10600,7 @@ def test_startup_warns_when_mock_testing_params_enabled(caplog): ) with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_if_mock_testing_params_enabled( - general_settings={MOCK_TESTING_CONFIG_KEY: True} - ) + ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings={MOCK_TESTING_CONFIG_KEY: True}) assert MOCK_TESTING_CONFIG_KEY in caplog.text for param_name in GATED_MOCK_PARAM_NAMES: @@ -11201,9 +10661,7 @@ async def test_setup_prisma_client_retains_connected_client_when_startup_health_ {"allow_requests_on_db_unavailable": True}, ) - mock_client = _mock_startup_prisma_client( - health_check_error=httpx.ReadTimeout("startup health check timed out") - ) + mock_client = _mock_startup_prisma_client(health_check_error=httpx.ReadTimeout("startup health check timed out")) result = await _run_setup_prisma_client(mock_client) assert mock_client.connect.await_count == 1 @@ -11227,9 +10685,7 @@ async def test_setup_prisma_client_arms_health_watchdog_before_startup_health_ch {"allow_requests_on_db_unavailable": True}, ) - mock_client = _mock_startup_prisma_client( - health_check_error=httpx.ReadTimeout("startup health check timed out") - ) + mock_client = _mock_startup_prisma_client(health_check_error=httpx.ReadTimeout("startup health check timed out")) call_order = MagicMock() call_order.attach_mock(mock_client.start_db_health_watchdog_task, "watchdog") call_order.attach_mock(mock_client.health_check, "health_check") @@ -11253,9 +10709,7 @@ async def test_setup_prisma_client_raises_when_db_unavailable_is_not_allowed(mon {"allow_requests_on_db_unavailable": False}, ) - mock_client = _mock_startup_prisma_client( - health_check_error=httpx.ReadTimeout("startup health check timed out") - ) + mock_client = _mock_startup_prisma_client(health_check_error=httpx.ReadTimeout("startup health check timed out")) with pytest.raises(httpx.ReadTimeout): await _run_setup_prisma_client(mock_client) @@ -11289,6 +10743,7 @@ async def _run_scheduled_background_jobs(): mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() mock_proxy_config = AsyncMock() with ( diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py index cf906259246..db842802435 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py @@ -8,6 +8,7 @@ because they are direct dependents on the lifecycle state. from __future__ import annotations +import asyncio from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock, patch @@ -138,6 +139,41 @@ def test_startup_event_propagates_init_callbacks_failure_raises(proxy_logging): proxy_logging.startup_event(llm_router=None, redis_usage_cache=None) +@pytest.mark.asyncio +async def test_startup_event_hands_the_daily_report_this_pods_lock_manager(proxy_logging): + """regression: issue #14809 - the daily report's dedupe lock only works if startup_event + passes the writer's pod_lock_manager down; dropping the argument silently restores the + every-pod-reports behavior.""" + proxy_logging.slack_alerting_instance = MagicMock() + proxy_logging.slack_alerting_instance.alert_types = ["daily_reports"] + proxy_logging.slack_alerting_instance._run_scheduled_daily_report = AsyncMock() + proxy_logging._init_litellm_callbacks = MagicMock() + proxy_logging.update_values = MagicMock() + llm_router = MagicMock() + + proxy_logging.startup_event(llm_router=llm_router, redis_usage_cache=None) + await asyncio.sleep(0) + + call = proxy_logging.slack_alerting_instance._run_scheduled_daily_report.call_args + assert proxy_logging.slack_alerting_instance._run_scheduled_daily_report.call_count == 1 + assert call.kwargs["pod_lock_manager"] is proxy_logging.db_spend_update_writer.pod_lock_manager + assert call.kwargs["llm_router"] is llm_router + + +@pytest.mark.asyncio +async def test_startup_event_skips_the_daily_report_when_it_is_not_an_alert_type(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock() + proxy_logging.slack_alerting_instance.alert_types = [] + proxy_logging.slack_alerting_instance._run_scheduled_daily_report = AsyncMock() + proxy_logging._init_litellm_callbacks = MagicMock() + proxy_logging.update_values = MagicMock() + + proxy_logging.startup_event(llm_router=None, redis_usage_cache=None) + await asyncio.sleep(0) + + proxy_logging.slack_alerting_instance._run_scheduled_daily_report.assert_not_called() + + # --------------------------------------------------------------------------- # _add_proxy_hooks # ---------------------------------------------------------------------------