From be71a8fdbf46f0335a8ba71daa0ffb9e39568c02 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 11 Aug 2026 12:41:11 -0700 Subject: [PATCH] fix(alerting): dedupe scheduled Slack spend reports across pods (#36489) * fix(alerting): dedupe scheduled Slack spend reports across pods Every pod ran its own weekly/monthly spend report jobs, prometheus fallback stats cron, and daily report loop, so deployments with multiple replicas or uvicorn workers received one copy per pod. Gate each scheduled send behind the shared PodLockManager redis lock. The lock is never released: its TTL (the full reporting window for the weekly interval job, whose per-pod anchors drift by boot time and jitter) doubles as a sent-this-window marker. acquire_lock returning None (no redis wired) proceeds, preserving single-pod behavior. Also generalize the pod lock could-not-acquire log line, which claimed to be about spend tracking for every consumer. Fixes #14809 * fix(alerting): harden spend report locks after adversarial review Weekly lock TTL gets an hour haircut: with ttl equal to the interval, the winner re-fires just before its own key expires, reacquires without a TTL refresh, and the key then lapses in time for a trailing pod to re-send. Job/lock ids move to litellm/constants.py per convention, and spend_report_frequency now rejects non-positive day counts, which previously coerced to an every-second schedule and would now compute a negative lock TTL that silently never sends. Adds the missing test coverage the review flagged: startup_event's pod_lock_manager wiring (identity-asserted), the prometheus closure's positive path, and the ungated immediate prometheus send pinned to exactly one await. * test(alerting): consolidate spend_report_frequency validator coverage Drops a duplicate non-positive-days test and parametrizes the survivor over the suffix half of the validator too * fix(alerting): route the startup prometheus fallback send through the pod lock Greptile caught that the boot-time send still ran once per pod when PROMETHEUS_URL is set, the same duplication class this PR removes * fix(alerting): make report lock acquisition non-reentrant Greptile caught that a pod booting within an hour of the fallback stats cron sent twice: the startup send takes the lock, then the cron fire hits acquire_lock's reacquire branch, which returns True for the holder. Window-marker gates now pass allow_reentrant=False so a live lock blocks everyone including its holder; leader-election consumers keep the reentrant default * test(proxy): give spec'd ProxyLogging mocks a db_spend_update_writer _initialize_slack_alerting_jobs now reads it for the pod lock manager, and spec=ProxyLogging blocks instance-only attributes --- litellm/constants.py | 4 + .../SlackAlerting/slack_alerting.py | 27 +- .../db_transaction_queue/pod_lock_manager.py | 21 +- litellm/proxy/proxy_server.py | 58 +- litellm/proxy/utils.py | 5 +- .../SlackAlerting/test_slack_alerting.py | 170 ++- .../test_pod_lock_manager.py | 35 +- .../proxy/proxy_server/test_lifecycle.py | 184 ++- tests/test_litellm/proxy/test_proxy_server.py | 1241 +++++------------ .../utils/proxy_logging/test_lifecycle.py | 36 + 10 files changed, 807 insertions(+), 974 deletions(-) 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 # ---------------------------------------------------------------------------