diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 8a432eb2f42..d90974afaba 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -1660,8 +1660,13 @@ async def _resolve_public_readiness_db(response: Response) -> str: "Not connected" (no DB configured), "connected", "disconnected". """ from litellm.proxy.proxy_server import prisma_client + from litellm import get_secret + _db_url = get_secret("DATABASE_URL", None) if prisma_client is None: + if _db_url is not None: + response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE + return "disconnected" return "Not connected" db_health_status = await _db_health_readiness_check() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3f8b01cc865..4b93067a63c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -419,10 +419,6 @@ from litellm.proxy.management_endpoints.workflow_management_endpoints import ( ) from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update from litellm.proxy.memory.memory_endpoints import router as memory_router -from litellm.proxy.plugin_routes import ( - router as plugin_router, - register_plugins_from_config, -) from litellm.proxy.middleware.in_flight_requests_middleware import ( InFlightRequestsMiddleware, ) @@ -430,9 +426,6 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi from litellm.proxy.middleware.request_size_limit_middleware import ( RequestSizeLimitMiddleware, ) -from litellm.proxy.middleware.security_headers_middleware import ( - SecurityHeadersMiddleware, -) from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, @@ -847,6 +840,37 @@ async def proxy_startup_event(app: FastAPI): if isinstance(worker_config, dict): await initialize(**worker_config) + ## V2 OTEL: now that config (and therefore the callbacks) is loaded, publish + ## the chosen V2 logger's TracerProvider as the OTel global. The FastAPI + ## instrumentation mounted at app-creation binds to the global provider, so + ## this is what makes server spans and gen-ai spans share one provider and + ## land in the same trace. Prefer an already-registered preset logger + ## (arize, langfuse, …) so server spans export to that backend too; otherwise + ## build a generic one from OTEL_* envs. ``set_tracer_provider`` only takes + ## effect once, so the first configured logger wins. + try: + from litellm.integrations.otel.model.config import is_otel_v2_enabled + + if is_otel_v2_enabled(): + from opentelemetry import trace as _otel_trace + + from litellm.integrations.otel.logger import OpenTelemetryV2 + + _otel_v2_logger = ( + next( + ( + cb + for cb in litellm.service_callback + if isinstance(cb, OpenTelemetryV2) + ), + None, + ) + or OpenTelemetryV2() + ) + _otel_trace.set_tracer_provider(_otel_v2_logger._tracer_provider) + except Exception as e: + verbose_proxy_logger.debug("Skipping OTel V2 provider setup: %s", e) + # check if DATABASE_URL in environment - load from there if prisma_client is None: _db_url: Optional[str] = get_secret("DATABASE_URL", None) # type: ignore @@ -883,42 +907,6 @@ async def proxy_startup_event(app: FastAPI): redis_usage_cache=transaction_buffer_redis_cache, ) - ## V2 OTEL: publish the chosen V2 logger's TracerProvider as the OTel global. - ## This MUST run after callback initialization above: a preset (arize, langfuse, - ## …) builds its logger there, folding the OTEL_* base exporter and its own - ## exporter into one logger. The FastAPI instrumentation mounted at app-creation - ## binds to the global provider, so reusing that one logger is what makes the - ## server span and the gen-ai spans share one provider and land in the same - ## trace, exporting to every configured backend. Running before callback init - ## (when no logger exists yet) would build a second, generic logger whose - ## provider became the global, orphaning the gen-ai spans onto a different - ## backend than the server span. A generic logger is built only when none was - ## configured. - try: - from litellm.integrations.otel.model.config import is_otel_v2_enabled - - if is_otel_v2_enabled(): - from opentelemetry import trace as _otel_trace - - from litellm.litellm_core_utils.litellm_logging import _in_memory_loggers - from litellm.integrations.otel.logger import ( - OpenTelemetryV2, - publish_global_otel_v2_provider, - ) - - registered = ( - open_telemetry_logger - if isinstance(open_telemetry_logger, OpenTelemetryV2) - else None - ) - publish_global_otel_v2_provider( - _in_memory_loggers, # any-ok: pre-existing untyped List[Any] global - _otel_trace.set_tracer_provider, - registered=registered, - ) - except Exception as e: - verbose_proxy_logger.debug("Skipping OTel V2 provider setup: %s", e) - ## Validate use_redis_transaction_buffer requires Redis cache ## ProxyStartupEvent._validate_redis_transaction_buffer_config( general_settings=general_settings, @@ -1769,7 +1757,6 @@ app.add_middleware( app.add_middleware(PrometheusAuthMiddleware) app.add_middleware(InFlightRequestsMiddleware) -app.add_middleware(SecurityHeadersMiddleware) def mount_swagger_ui(): @@ -2039,43 +2026,7 @@ def cost_tracking(): ) -# Bounds authoritative DB re-reads when enforcing a budget against a -# stale-low spend counter: at most one DB read per counter per window. -SPEND_DB_FLOOR_CACHE_TTL_SECONDS = 5 - - -def _fail_closed_budget_enforcement() -> bool: - return general_settings.get("fail_closed_budget_enforcement") is True - - -def _raise_budget_unverifiable(counter_key: str) -> None: - verbose_proxy_logger.warning( - "fail_closed_budget_enforcement: rejecting request — spend for %s could " - "not be verified against Redis or the database", - counter_key, - ) - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail={ - "error": ( - "Budget enforcement unavailable: current spend could not be " - "verified against Redis or the database, and " - "fail_closed_budget_enforcement is enabled, so the request was " - "rejected to avoid exceeding the configured budget. Retry shortly." - ) - }, - ) - - -async def get_current_spend( - counter_key: str, - fallback_spend: float, - max_budget: float | None = None, - window_entity_type: str | None = None, - window_entity_id: str | None = None, - window_start: datetime | None = None, - fallback_authoritative: bool = False, -) -> float: +async def get_current_spend(counter_key: str, fallback_spend: float) -> float: """ Read current spend from the cross-pod spend counter. @@ -2089,168 +2040,7 @@ async def get_current_spend( 2. In-memory counter (single-instance or Redis failure) 3. Reseed from authoritative DB spend (counter expired, cross-pod stale) 4. Caller-supplied fallback (DB unavailable, cold start) - - When ``max_budget`` is supplied, the counter is re-checked against the - authoritative recorded spend before a request is admitted. A Redis counter - that survived a Redis restart can return a stale-low value loaded from an - older RDB snapshot; that read is a hit (not a clean miss), so step 3 never - runs and a key can leak spend past ``max_budget`` indefinitely. The - authoritative source depends on the counter: primary key/team/user/org - counters read the DB row; per-window counters (``window_start`` supplied) - aggregate spend logs; end-user/tag counters have no DB row, so the caller's - ``fallback_spend`` (loaded fresh in auth) is authoritative. The DB read is - skipped for healthy primary counters (counter at or above recorded spend) - and cached in-process for a few seconds, so a persistently stale counter - drives at most one read per counter per window rather than one per request. """ - current, verified = await _read_spend_counter_estimate( - counter_key=counter_key, fallback_spend=fallback_spend - ) - if fallback_authoritative: - verified = True - - if max_budget is None or current >= max_budget: - return current - - # Cheap staleness signal for primary counters: the counter reads below the - # spend this caller already knows about. Window counters have no such signal - # (fallback is 0), so they always re-check, bounded by the cache. Strict mode - # (fail_closed_budget_enforcement) always re-checks against the authoritative - # source too, so a counter that is stale-low at the same time as the caller's - # cached spend cannot slip through; the 5s cache keeps that bounded. - is_window = window_start is not None - if fallback_spend > current or is_window or _fail_closed_budget_enforcement(): - authoritative = await _authoritative_floor_spend( - counter_key=counter_key, - window_entity_type=window_entity_type, - window_entity_id=window_entity_id, - window_start=window_start, - ) - if authoritative is not None: - verified = True - if authoritative > current: - await _repair_stale_spend_counter( - counter_key=counter_key, db_spend=authoritative - ) - return authoritative - elif fallback_spend > current: - # end-user / tag counters have no DB row; fallback_spend is the - # authoritative recorded value loaded in auth. - return fallback_spend - - # Opt-in hard guarantee: when the spend backing this admit decision came - # only from a per-pod cache (Redis and DB both unreadable), reject rather - # than admit on an unverifiable budget. No-op unless the flag is set, so - # default behavior is unchanged. - if not verified and _fail_closed_budget_enforcement(): - _raise_budget_unverifiable(counter_key) - - return current - - -async def _repair_stale_spend_counter(counter_key: str, db_spend: float) -> None: - """Raise a counter that has fallen below the authoritative DB spend (e.g. - Redis restarted and reloaded an older snapshot) so every worker reads the - corrected value directly instead of re-deriving it per request, and so a - worker whose own cached spend is also stale still sees the true total. - - The write is monotonic: it only ever raises the counter, so a repair that - carries a slightly-stale DB total cannot clobber a concurrent increment that - already pushed the counter higher (which would let racing requests - under-count). Redis enforces this atomically via async_set_max; the - in-memory copy is guarded by a read-compare-write with no await in between, - so it is atomic within the worker. - """ - cached = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) - needs_update = True - if cached is not None: - try: - needs_update = float(cached) < db_spend - except (TypeError, ValueError): - needs_update = True - if needs_update: - spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=db_spend) - if spend_counter_cache.redis_cache is not None: - try: - await spend_counter_cache.redis_cache.async_set_max( - key=counter_key, value=db_spend - ) - except Exception: - verbose_proxy_logger.debug( - "Unable to repair stale spend counter %s in Redis", - counter_key, - exc_info=True, - ) - - -async def reseed_spend_counter_from_db(counter_key: str) -> None: - """Recover a counter that the reservation reconcile found in an inconsistent - state (missing, or where applying the reconcile delta would drive it - negative) by reseeding it from the DB instead of deleting it. - - The DB row is a LAGGING authoritative floor, not post-request truth: the - entity .spend column is flushed in batches (every PROXY_BATCH_WRITE_AT), so - it can exclude this request's just-recorded cost and other buffered spend. - That is fine here: the monotonic set-max can only RAISE a stale-low counter - toward that floor (never lowers it or clobbers a concurrent increment), and - the read-time floor (_authoritative_floor_spend) converges to the true total - as the buffer flushes. The point is to restore enforcement to a real floor - rather than leave the counter deleted and unenforced (the prior fail-open). - Counters with no DB row (window/end-user/tag) are left untouched rather than - deleted, so enforcement keeps reading whatever value they hold. - """ - db_spend = await SpendCounterReseed.from_db( - prisma_client=prisma_client, counter_key=counter_key - ) - if db_spend is None: - return - await _repair_stale_spend_counter(counter_key=counter_key, db_spend=db_spend) - - -async def _authoritative_floor_spend( - counter_key: str, - window_entity_type: str | None = None, - window_entity_id: str | None = None, - window_start: datetime | None = None, -) -> float | None: - marker_key = f"spend_db_floor:{counter_key}" - cached = spend_counter_cache.in_memory_cache.get_cache(key=marker_key) - if cached is not None: - return float(cached) - - db_spend = await SpendCounterReseed.from_db( - prisma_client=prisma_client, counter_key=counter_key - ) - if ( - db_spend is None - and window_entity_type is not None - and window_entity_id is not None - and window_start is not None - ): - db_spend = await SpendCounterReseed.window_from_spend_logs( - prisma_client=prisma_client, - entity_type=window_entity_type, - entity_id=window_entity_id, - window_start=window_start, - ) - if db_spend is None: - return None - - spend_counter_cache.in_memory_cache.set_cache( - key=marker_key, - value=db_spend, - ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS, - ) - return db_spend - - -async def _read_spend_counter_estimate( - counter_key: str, fallback_spend: float -) -> tuple[float, bool]: - """Return (spend, authoritative). ``authoritative`` is True when the value - came from Redis or a fresh DB read (cross-pod truth), False when it came - from the per-pod in-memory copy or the caller's fallback. Only the - fail-closed path reads the flag; normal callers ignore it.""" # 1. Redis first (cross-pod authoritative). On clean miss, skip # in-memory: per-pod in-memory only has this pod's writes, so it # would mask cross-pod increments. @@ -2259,7 +2049,7 @@ async def _read_spend_counter_estimate( try: val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key) if val is not None: - return float(val), True + return float(val) redis_clean_miss = True except Exception as e: verbose_proxy_logger.debug( @@ -2272,7 +2062,7 @@ async def _read_spend_counter_estimate( if not redis_clean_miss: val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) if val is not None: - return float(val), False + return float(val) # 3. Reseed from DB - fallback_spend lags cross-pod, would allow bypass. db_spend = await SpendCounterReseed.coalesced( @@ -2281,10 +2071,10 @@ async def _read_spend_counter_estimate( counter_key=counter_key, ) if db_spend is not None: - return db_spend, True + return db_spend # 4. Caller-supplied fallback (DB unavailable). - return fallback_spend, False + return fallback_spend async def increment_spend_counters( @@ -4506,8 +4296,6 @@ class ProxyConfig: load_from_azure_key_vault(use_azure_key_vault=use_azure_key_vault) ### ALERTING ### self._load_alerting_settings(general_settings=general_settings) - ### PLUGINS ### - register_plugins_from_config(general_settings) ### CONNECT TO DATABASE ### database_url = general_settings.get("database_url", None) if database_url and database_url.startswith("os.environ/"): @@ -4779,11 +4567,6 @@ class ProxyConfig: config ) - ## SANDBOX TOOLS SETTINGS - from litellm.sandbox.sandbox_tools import register_sandbox_tools - - register_sandbox_tools(config.get("sandbox_tools") or []) - ## /fine_tuning/jobs endpoints config finetuning_config = config.get("finetune_settings", None) set_fine_tuning_config(config=finetuning_config) @@ -5674,10 +5457,6 @@ class ProxyConfig: llm_router=llm_router, ) - if _general_settings is not None and "plugins" in _general_settings: - general_settings["plugins"] = _general_settings["plugins"] - register_plugins_from_config(general_settings) - async def _reschedule_spend_log_cleanup_job(self): """ Reschedule the spend log cleanup job based on current general_settings. @@ -7950,6 +7729,52 @@ class ProxyStartupEvent: "Failed to check DB for store_model_in_db: %s", str(e) ) + async def _retry_db_init(): + import asyncio + while True: + try: + await asyncio.sleep(5) + _retry_db_gs = await ConfigRepository(prisma_client).table.find_first( + where={"param_name": "general_settings"} + ) + if _retry_db_gs is not None and isinstance(_retry_db_gs.param_value, dict): + _retry_db_val = _retry_db_gs.param_value.get("store_model_in_db") + if _retry_db_val is True or ( + isinstance(_retry_db_val, str) and _retry_db_val.lower() == "true" + ): + global store_model_in_db + store_model_in_db = True + verbose_proxy_logger.info( + "store_model_in_db=True loaded from DB on retry" + ) + scheduler.add_job( + proxy_config.add_deployment, + "interval", + seconds=30, + args=[prisma_client, proxy_logging_obj], + id="add_deployment_job", + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + ) + await proxy_config.add_deployment( + prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj + ) + scheduler.add_job( + proxy_config.get_credentials, + "interval", + seconds=30, + args=[prisma_client], + id="get_credentials_job", + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + ) + await proxy_config.get_credentials(prisma_client=prisma_client) + break + except Exception: + pass + import asyncio + asyncio.create_task(_retry_db_init()) + if store_model_in_db is True: # MEMORY LEAK FIX: Increase interval from 10s to 30s minimum # Frequent polling was causing excessive memory allocations @@ -8979,7 +8804,7 @@ async def chat_completion( completion_stream=_iterator, model=e.model, custom_llm_provider="cached_response", - logging_obj=_data.get("litellm_logging_obj", None), + logging_obj=data.get("litellm_logging_obj", None), ) selected_data_generator = select_data_generator( response=_streaming_response, @@ -9014,7 +8839,7 @@ async def chat_completion( completion_stream=_iterator, model=data.get("model", ""), custom_llm_provider="cached_response", - logging_obj=_data.get("litellm_logging_obj", None), + logging_obj=data.get("litellm_logging_obj", None), ) selected_data_generator = select_data_generator( response=_streaming_response, @@ -13161,9 +12986,6 @@ async def model_info_v1( # use internal routing keys (model_name_{team_id}_{uuid}) and were omitted # when v1 resolved models only via public model_name strings. all_models: List[dict] = copy.deepcopy(llm_router.model_list) - alias_models = copy.deepcopy(llm_router.get_model_list_from_model_alias()) - all_models.extend(alias_models) - allowed_model_names = _get_v1_model_info_allowed_model_names( user_api_key_dict=user_api_key_dict, llm_router=llm_router, @@ -13731,24 +13553,26 @@ async def fallback_login(request: Request): # get url from request redirect_url = get_custom_url(str(request.base_url)) + ui_username = os.getenv("UI_USERNAME") if redirect_url.endswith("/"): redirect_url += "sso/callback" else: redirect_url += "/sso/callback" - from fastapi.responses import HTMLResponse + if ui_username is not None: + # No Google, Microsoft SSO + # Use UI Credentials set in .env + from fastapi.responses import HTMLResponse - hide_default_credentials_hint = ( - os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true" - or general_settings.get("hide_default_credentials_hint", False) is True - ) - return HTMLResponse( - content=build_ui_login_form( - show_deprecation_banner=False, - hide_default_credentials_hint=hide_default_credentials_hint, - ), - status_code=200, - ) + return HTMLResponse( + content=build_ui_login_form(show_deprecation_banner=False), status_code=200 + ) + else: + from fastapi.responses import HTMLResponse + + return HTMLResponse( + content=build_ui_login_form(show_deprecation_banner=False), status_code=200 + ) @router.post( @@ -14946,41 +14770,6 @@ async def update_config( Keep it more precise, to prevent overwrite other values unintentially """ -_PLUGIN_KEY_REDACTED = "***" - - -def _preserve_redacted_plugin_keys(incoming: object, existing: object) -> object: - """Restore real plugin_key values the client never sees. - - /config/field/info redacts every plugin_key to ``"***"``, so an admin - editing a plugin posts that placeholder (or a blank, when the UI clears the - field) straight back. Treat a blank or redacted plugin_key as "keep the - stored credential" by sourcing it from the existing config; only a real, - non-redacted value replaces it, and a blank with no stored key drops the - field entirely instead of persisting the placeholder. - """ - if not isinstance(incoming, list): - return incoming - - stored_keys = { - p["name"]: p["plugin_key"] - for p in (existing if isinstance(existing, list) else []) - if isinstance(p, dict) and p.get("name") and p.get("plugin_key") - } - - def resolve(plugin: object) -> object: - if not isinstance(plugin, dict): - return plugin - key = plugin.get("plugin_key") - if key not in (None, "", _PLUGIN_KEY_REDACTED): - return plugin - name = plugin.get("name") - if name in stored_keys: - return {**plugin, "plugin_key": stored_keys[name]} - return {k: v for k, v in plugin.items() if k != "plugin_key"} - - return [resolve(p) for p in incoming] - @router.post( "/config/field/update", @@ -15047,13 +14836,7 @@ async def update_config_general_settings( ## update db - field_value = data.field_value - if data.field_name == "plugins": - field_value = _preserve_redacted_plugin_keys( - field_value, general_settings.get("plugins") - ) - - general_settings[data.field_name] = field_value + general_settings[data.field_name] = data.field_value response = await ConfigRepository(prisma_client).table.upsert( where={"param_name": "general_settings"}, @@ -15064,9 +14847,6 @@ async def update_config_general_settings( ) await invalidate_config_param("general_settings") - if data.field_name == "plugins": - register_plugins_from_config(general_settings) - return response @@ -15122,19 +14902,9 @@ async def get_config_general_settings( general_settings = dict(db_general_settings.param_value) if field_name in general_settings: - field_value = general_settings[field_name] - # Redact plugin_key from plugin configs so the shared credential - # is never returned even to admin-viewer callers. - if field_name == "plugins" and isinstance(field_value, list): - field_value = [ - ( - {k: ("***" if k == "plugin_key" else v) for k, v in p.items()} - if isinstance(p, dict) - else p - ) - for p in field_value - ] - return ConfigFieldInfo(field_name=field_name, field_value=field_value) + return ConfigFieldInfo( + field_name=field_name, field_value=general_settings[field_name] + ) else: raise HTTPException( status_code=400, @@ -16456,7 +16226,6 @@ app.include_router(model_access_group_management_router) app.include_router(tag_management_router) app.include_router(workflow_management_router) app.include_router(memory_router) -app.include_router(plugin_router) app.include_router(cost_tracking_settings_router) app.include_router(router_settings_router) app.include_router(fallback_management_router) diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index a04ad5598df..9857f375d2f 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -2061,3 +2061,21 @@ def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): assert "aws_access_key_id" not in cleaned assert cleaned.get("api_base") == "https://example.test/v1" assert cleaned.get("api_version") == "2024-10-21" + + +@pytest.mark.asyncio +async def test_health_readiness_returns_503_when_db_configured_but_not_initialized(): + """ + readiness probe should return 503 when a DB is configured (DATABASE_URL is set) but prisma_client is None. + """ + from unittest.mock import patch + from fastapi import Response + from litellm.proxy.health_endpoints._health_endpoints import health_readiness + + response = Response() + with patch("litellm.proxy.proxy_server.prisma_client", None): + with patch("litellm.get_secret", return_value="postgresql://fake:fake@localhost:5432/fake"): + result = await health_readiness(response=response) + + assert response.status_code == 503 + assert result == {"status": "healthy", "db": "disconnected"} diff --git a/tests/test_litellm/proxy/test_proxy_startup_db_failure.py b/tests/test_litellm/proxy/test_proxy_startup_db_failure.py new file mode 100644 index 00000000000..a8263a69fda --- /dev/null +++ b/tests/test_litellm/proxy/test_proxy_startup_db_failure.py @@ -0,0 +1,88 @@ +import pytest +from unittest.mock import AsyncMock, MagicMock, patch +import asyncio + +@pytest.mark.asyncio +async def test_proxy_startup_db_failure_reproduction(monkeypatch): + """ + Test for litellm DB connection failure during startup. + When Litellm starts up and the database is unstable/down, + the query for general_settings in initialize_scheduled_background_jobs + throws an exception, causing store_model_in_db to remain False. + This starts a background retry task. Upon DB recovery, the retry task + registers the add_deployment and get_credentials jobs. + + The test asserts that both 'add_deployment_job' and 'get_credentials_job' + are registered after the background retry task runs and succeeds. + """ + # Delete environment overrides so we rely on DB check + monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) + + from litellm.proxy.proxy_server import ProxyStartupEvent + from litellm.proxy.utils import ProxyLogging + + # 1. Mock DB connection to throw an exception on initial call, but succeed on retry + mock_prisma_client = MagicMock() + mock_db_record = MagicMock() + mock_db_record.param_value = {"store_model_in_db": True} + + # Side effect: first call (startup check) fails; second call (retry loop) succeeds + mock_prisma_client.db.litellm_config.find_first = AsyncMock( + side_effect=[ + Exception("Database Connection Failed / Connection Timed Out"), + mock_db_record + ] + ) + + # 2. Mock proxy logging and proxy config + mock_proxy_logging = MagicMock(spec=ProxyLogging) + mock_proxy_logging.slack_alerting_instance = MagicMock() + + mock_proxy_config = AsyncMock() + + # 3. Patch AsyncIOScheduler to spy/mock scheduler calls + with patch("litellm.proxy.proxy_server.AsyncIOScheduler") as mock_scheduler_class: + mock_scheduler_instance = MagicMock() + mock_scheduler_class.return_value = mock_scheduler_instance + + # Patch proxy_config and initialize store_model_in_db to False (default YAML config behavior) + with ( + patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", False), + patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False), + ): + # 4. Mock asyncio.sleep to fast-forward the 5 seconds sleep in retry loop + original_sleep = asyncio.sleep + async def mock_sleep(seconds, *args, **kwargs): + if seconds == 5: + await original_sleep(0.01) + else: + await original_sleep(seconds) + + with patch("asyncio.sleep", side_effect=mock_sleep): + # 5. Invoke initialize_scheduled_background_jobs + await ProxyStartupEvent.initialize_scheduled_background_jobs( + general_settings={"disable_spend_logs": True}, + prisma_client=mock_prisma_client, + proxy_budget_rescheduler_min_time=1, + proxy_budget_rescheduler_max_time=2, + proxy_batch_write_at=5, + proxy_logging_obj=mock_proxy_logging, + ) + + # 6. Yield execution to allow the background retry task to run and complete + await asyncio.sleep(0.1) + + # 7. Assert that the add_deployment and get_credentials background jobs ARE registered. + add_deployment_job_registered = False + get_credentials_job_registered = False + + for call_args in mock_scheduler_instance.add_job.call_args_list: + job_id = call_args.kwargs.get("id") + if job_id == "add_deployment_job": + add_deployment_job_registered = True + elif job_id == "get_credentials_job": + get_credentials_job_registered = True + + assert add_deployment_job_registered, "add_deployment_job was not registered after database recovery." + assert get_credentials_job_registered, "get_credentials_job was not registered after database recovery."