mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): return HTTP 503 from readiness probe when DATABASE_URL is set but prisma_client is uninitialized
This commit is contained in:
parent
84c1414aef
commit
663b1d4f91
4 changed files with 213 additions and 333 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
88
tests/test_litellm/proxy/test_proxy_startup_db_failure.py
Normal file
88
tests/test_litellm/proxy/test_proxy_startup_db_failure.py
Normal 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."
|
||||
Loading…
Add table
Reference in a new issue