fix(proxy): return HTTP 503 from readiness probe when DATABASE_URL is set but prisma_client is uninitialized

This commit is contained in:
NoySolvin 2026-06-21 16:19:36 +03:00
parent 84c1414aef
commit 663b1d4f91
4 changed files with 213 additions and 333 deletions

View file

@ -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()

View file

@ -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)

View file

@ -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"}

View file

@ -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."