feat(proxy): limit concurrent health checks with health_check_concurrency (#20584)

* staged first pass

* black

* Update litellm/proxy/health_check.py

Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>

* simpler

* restore cached logo

* fix tests for perform_health_check max_concurrency arg

* implement pr suggestion

* and the helm chart

* add configureable resources and probes to the deployment in the helm chart

* more helm chart unittests

* move some background healthcheck loggin to debug

---------

Co-authored-by: Sean Glover <sglover@athenahealth.com>
Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>
This commit is contained in:
Sean Marsh Glover 2026-02-24 11:16:59 -05:00 committed by GitHub
parent 1fa0aad3f2
commit 4652c73259
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
12 changed files with 809 additions and 180 deletions

View File

@ -36,6 +36,10 @@ If `db.useStackgresOperator` is used (not yet implemented):
| `serviceAccount.create` | Whether or not to create a Kubernetes Service Account for this deployment. The default is `false` because LiteLLM has no need to access the Kubernetes API. | `false` | | `serviceAccount.create` | Whether or not to create a Kubernetes Service Account for this deployment. The default is `false` because LiteLLM has no need to access the Kubernetes API. | `false` |
| `service.type` | Kubernetes Service type (e.g. `LoadBalancer`, `ClusterIP`, etc.) | `ClusterIP` | | `service.type` | Kubernetes Service type (e.g. `LoadBalancer`, `ClusterIP`, etc.) | `ClusterIP` |
| `service.port` | TCP port that the Kubernetes Service will listen on. Also the TCP port within the Pod that the proxy will listen on. | `4000` | | `service.port` | TCP port that the Kubernetes Service will listen on. Also the TCP port within the Pod that the proxy will listen on. | `4000` |
| `livenessProbe.*` | Liveness probe settings for the LiteLLM container (`path`, `periodSeconds`, `timeoutSeconds`, thresholds, and initial delay). | See `values.yaml` |
| `readinessProbe.*` | Readiness probe settings for the LiteLLM container (`path`, `periodSeconds`, `timeoutSeconds`, thresholds, and initial delay). | See `values.yaml` |
| `startupProbe.*` | Startup probe settings for the LiteLLM container (`path`, `periodSeconds`, `timeoutSeconds`, thresholds, and initial delay). | See `values.yaml` |
| `resources.*` | CPU/memory requests and limits for the LiteLLM container. | `{}` |
| `service.loadBalancerClass` | Optional LoadBalancer implementation class (only used when `service.type` is `LoadBalancer`) | `""` | | `service.loadBalancerClass` | Optional LoadBalancer implementation class (only used when `service.type` is `LoadBalancer`) | `""` |
| `ingress.labels` | Additional labels for the Ingress resource | `{}` | | `ingress.labels` | Additional labels for the Ingress resource | `{}` |
| `ingress.*` | See [values.yaml](./values.yaml) for example settings | N/A | | `ingress.*` | See [values.yaml](./values.yaml) for example settings | N/A |

View File

@ -6,4 +6,4 @@ metadata:
data: data:
config.yaml: | config.yaml: |
{{ .Values.proxy_config | toYaml | indent 6 }} {{ .Values.proxy_config | toYaml | indent 6 }}
{{- end }} {{- end }}

View File

@ -158,18 +158,31 @@ spec:
{{- end }} {{- end }}
livenessProbe: livenessProbe:
httpGet: httpGet:
path: /health/liveliness path: {{ .Values.livenessProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }} port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
initialDelaySeconds: {{ .Values.livenessProbe.initialDelaySeconds }}
periodSeconds: {{ .Values.livenessProbe.periodSeconds }}
timeoutSeconds: {{ .Values.livenessProbe.timeoutSeconds }}
successThreshold: {{ .Values.livenessProbe.successThreshold }}
failureThreshold: {{ .Values.livenessProbe.failureThreshold }}
readinessProbe: readinessProbe:
httpGet: httpGet:
path: /health/readiness path: {{ .Values.readinessProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }} port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
initialDelaySeconds: {{ .Values.readinessProbe.initialDelaySeconds }}
periodSeconds: {{ .Values.readinessProbe.periodSeconds }}
timeoutSeconds: {{ .Values.readinessProbe.timeoutSeconds }}
successThreshold: {{ .Values.readinessProbe.successThreshold }}
failureThreshold: {{ .Values.readinessProbe.failureThreshold }}
startupProbe: startupProbe:
httpGet: httpGet:
path: /health/readiness path: {{ .Values.startupProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }} port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
failureThreshold: 30 initialDelaySeconds: {{ .Values.startupProbe.initialDelaySeconds }}
periodSeconds: 10 periodSeconds: {{ .Values.startupProbe.periodSeconds }}
timeoutSeconds: {{ .Values.startupProbe.timeoutSeconds }}
successThreshold: {{ .Values.startupProbe.successThreshold }}
failureThreshold: {{ .Values.startupProbe.failureThreshold }}
resources: resources:
{{- toYaml .Values.resources | nindent 12 }} {{- toYaml .Values.resources | nindent 12 }}
volumeMounts: volumeMounts:
@ -235,4 +248,4 @@ spec:
{{- if .Values.topologySpreadConstraints }} {{- if .Values.topologySpreadConstraints }}
topologySpreadConstraints: topologySpreadConstraints:
{{- toYaml .Values.topologySpreadConstraints | nindent 8 }} {{- toYaml .Values.topologySpreadConstraints | nindent 8 }}
{{- end }} {{- end }}

View File

@ -159,4 +159,150 @@ tests:
value: -c value: -c
- equal: - equal:
path: spec.template.spec.containers[0].lifecycle.preStop.exec.command[2] path: spec.template.spec.containers[0].lifecycle.preStop.exec.command[2]
value: echo "Container stopping" value: echo "Container stopping"
- it: should render background health check settings from proxy_config.general_settings
template: configmap-litellm.yaml
set:
proxy_config.general_settings.background_health_checks: true
proxy_config.general_settings.health_check_interval: 240
proxy_config.general_settings.health_check_concurrency: 16
proxy_config.general_settings.health_check_details: false
asserts:
- matchRegex:
path: data["config.yaml"]
pattern: '(?m)^\s*background_health_checks:\s*true$'
- matchRegex:
path: data["config.yaml"]
pattern: '(?m)^\s*health_check_interval:\s*240$'
- matchRegex:
path: data["config.yaml"]
pattern: '(?m)^\s*health_check_concurrency:\s*16$'
- matchRegex:
path: data["config.yaml"]
pattern: '(?m)^\s*health_check_details:\s*false$'
- it: should allow overriding liveness, readiness, and startup probes
template: deployment.yaml
set:
livenessProbe:
path: /custom/livez
initialDelaySeconds: 5
periodSeconds: 15
timeoutSeconds: 5
successThreshold: 1
failureThreshold: 5
readinessProbe:
path: /custom/readyz
initialDelaySeconds: 10
periodSeconds: 20
timeoutSeconds: 6
successThreshold: 1
failureThreshold: 6
startupProbe:
path: /custom/startupz
initialDelaySeconds: 15
periodSeconds: 25
timeoutSeconds: 7
successThreshold: 1
failureThreshold: 40
asserts:
- equal:
path: spec.template.spec.containers[0].livenessProbe.httpGet.path
value: /custom/livez
- equal:
path: spec.template.spec.containers[0].livenessProbe.timeoutSeconds
value: 5
- equal:
path: spec.template.spec.containers[0].readinessProbe.httpGet.path
value: /custom/readyz
- equal:
path: spec.template.spec.containers[0].readinessProbe.timeoutSeconds
value: 6
- equal:
path: spec.template.spec.containers[0].startupProbe.httpGet.path
value: /custom/startupz
- equal:
path: spec.template.spec.containers[0].startupProbe.failureThreshold
value: 40
- it: should render container resources from values
template: deployment.yaml
set:
resources:
limits:
cpu: 500m
memory: 2Gi
requests:
cpu: 250m
memory: 1Gi
asserts:
- equal:
path: spec.template.spec.containers[0].resources.limits.cpu
value: 500m
- equal:
path: spec.template.spec.containers[0].resources.limits.memory
value: 2Gi
- equal:
path: spec.template.spec.containers[0].resources.requests.cpu
value: 250m
- equal:
path: spec.template.spec.containers[0].resources.requests.memory
value: 1Gi
- it: should keep default probes and empty resources unchanged
template: deployment.yaml
asserts:
- equal:
path: spec.template.spec.containers[0].livenessProbe.httpGet.path
value: /health/liveliness
- equal:
path: spec.template.spec.containers[0].livenessProbe.initialDelaySeconds
value: 0
- equal:
path: spec.template.spec.containers[0].livenessProbe.periodSeconds
value: 10
- equal:
path: spec.template.spec.containers[0].livenessProbe.timeoutSeconds
value: 1
- equal:
path: spec.template.spec.containers[0].livenessProbe.successThreshold
value: 1
- equal:
path: spec.template.spec.containers[0].livenessProbe.failureThreshold
value: 3
- equal:
path: spec.template.spec.containers[0].readinessProbe.httpGet.path
value: /health/readiness
- equal:
path: spec.template.spec.containers[0].readinessProbe.initialDelaySeconds
value: 0
- equal:
path: spec.template.spec.containers[0].readinessProbe.periodSeconds
value: 10
- equal:
path: spec.template.spec.containers[0].readinessProbe.timeoutSeconds
value: 1
- equal:
path: spec.template.spec.containers[0].readinessProbe.successThreshold
value: 1
- equal:
path: spec.template.spec.containers[0].readinessProbe.failureThreshold
value: 3
- equal:
path: spec.template.spec.containers[0].startupProbe.httpGet.path
value: /health/readiness
- equal:
path: spec.template.spec.containers[0].startupProbe.initialDelaySeconds
value: 0
- equal:
path: spec.template.spec.containers[0].startupProbe.periodSeconds
value: 10
- equal:
path: spec.template.spec.containers[0].startupProbe.timeoutSeconds
value: 1
- equal:
path: spec.template.spec.containers[0].startupProbe.successThreshold
value: 1
- equal:
path: spec.template.spec.containers[0].startupProbe.failureThreshold
value: 30
- equal:
path: spec.template.spec.containers[0].resources
value: {}

View File

@ -84,6 +84,31 @@ service:
separateHealthApp: false separateHealthApp: false
separateHealthPort: 8081 separateHealthPort: 8081
# Probe tuning for proxy container
livenessProbe:
path: /health/liveliness
initialDelaySeconds: 0
periodSeconds: 10
timeoutSeconds: 1
successThreshold: 1
failureThreshold: 3
readinessProbe:
path: /health/readiness
initialDelaySeconds: 0
periodSeconds: 10
timeoutSeconds: 1
successThreshold: 1
failureThreshold: 3
startupProbe:
path: /health/readiness
initialDelaySeconds: 0
periodSeconds: 10
timeoutSeconds: 1
successThreshold: 1
failureThreshold: 30
ingress: ingress:
enabled: false enabled: false
className: "nginx" className: "nginx"

View File

@ -2079,6 +2079,13 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
health_check_interval: int = Field( health_check_interval: int = Field(
300, description="background health check interval in seconds" 300, description="background health check interval in seconds"
) )
health_check_concurrency: Optional[int] = Field(
None,
description=(
"limit concurrent health checks per cycle; when unset, "
"health checks run without a concurrency cap"
),
)
alerting: Optional[List] = Field( alerting: Optional[List] = Field(
None, None,
description="List of alerting integrations. Today, just slack - `alerting: ['slack']`", description="List of alerting integrations. Today, just slack - `alerting: ['slack']`",

View File

@ -3,6 +3,9 @@
import asyncio import asyncio
import logging import logging
import random import random
import sys
import threading
import time
from typing import List, Optional from typing import List, Optional
import litellm import litellm
@ -23,6 +26,29 @@ ILLEGAL_DISPLAY_PARAMS = [
MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"] MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"]
def _get_process_rss_mb() -> Optional[float]:
"""
Get process RSS memory in MB.
On Linux, ru_maxrss is in KB. On macOS, ru_maxrss is in bytes.
"""
try:
import resource
ru_maxrss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
if sys.platform == "darwin":
return float(ru_maxrss) / (1024 * 1024)
return float(ru_maxrss) / 1024
except Exception:
return None
def _rss_mb_for_log() -> str:
rss_mb = _get_process_rss_mb()
if rss_mb is None:
return "unknown"
return f"{rss_mb:.2f}"
def _get_random_llm_message(): def _get_random_llm_message():
""" """
Get a random message from the LLM. Get a random message from the LLM.
@ -67,26 +93,29 @@ async def run_with_timeout(task, timeout):
try: try:
return await asyncio.wait_for(task, timeout) return await asyncio.wait_for(task, timeout)
except asyncio.TimeoutError: except asyncio.TimeoutError:
task.cancel() # `asyncio.wait_for()` already cancels only the awaited task on timeout.
# Only cancel child tasks of the current task # Do not cancel unrelated sibling health check tasks.
current_task = asyncio.current_task()
for t in asyncio.all_tasks():
if t != current_task:
t.cancel()
try:
await asyncio.wait_for(task, 0.1) # Give 100ms for cleanup
except (asyncio.TimeoutError, asyncio.CancelledError, Exception):
pass
return {"error": "Timeout exceeded"} return {"error": "Timeout exceeded"}
async def _perform_health_check(model_list: list, details: Optional[bool] = True): async def _perform_health_check(
model_list: list,
details: Optional[bool] = True,
max_concurrency: Optional[int] = None,
instrumentation_context: Optional[dict] = None,
):
""" """
Perform a health check for each model in the list. Perform a health check for each model in the list.
max_concurrency: Optional limit on concurrent health check requests.
""" """
tasks = [] instrumentation_context = instrumentation_context or {}
for model in model_list: instrumentation_enabled = bool(instrumentation_context.get("enabled", False))
cycle_id = instrumentation_context.get("cycle_id", "unknown")
source = instrumentation_context.get("source", "unknown")
async def _run_model_health_check(model: dict):
litellm_params = model["litellm_params"] litellm_params = model["litellm_params"]
model_info = model.get("model_info", {}) model_info = model.get("model_info", {})
mode = model_info.get("mode", None) mode = model_info.get("mode", None)
@ -95,9 +124,9 @@ async def _perform_health_check(model_list: list, details: Optional[bool] = True
) )
timeout = model_info.get("health_check_timeout") or HEALTH_CHECK_TIMEOUT_SECONDS timeout = model_info.get("health_check_timeout") or HEALTH_CHECK_TIMEOUT_SECONDS
task = run_with_timeout( return await run_with_timeout(
litellm.ahealth_check( litellm.ahealth_check(
model["litellm_params"], litellm_params,
mode=mode, mode=mode,
prompt=DEFAULT_HEALTH_CHECK_PROMPT, prompt=DEFAULT_HEALTH_CHECK_PROMPT,
input=["test from litellm"], input=["test from litellm"],
@ -105,9 +134,73 @@ async def _perform_health_check(model_list: list, details: Optional[bool] = True
timeout, timeout,
) )
tasks.append(task) async def _run_health_checks_with_bounded_concurrency(
models: list, concurrency_limit: int
) -> tuple[list, int]:
"""
Run health checks with at most `concurrency_limit` active tasks.
Preserves result ordering to match `models`.
"""
results: list = [None] * len(models)
tasks_to_index: dict[asyncio.Task, int] = {}
model_iter = iter(enumerate(models))
peak_in_flight = 0
results = await asyncio.gather(*tasks, return_exceptions=True) def _schedule_next() -> bool:
nonlocal peak_in_flight
try:
idx, next_model = next(model_iter)
except StopIteration:
return False
task = asyncio.create_task(_run_model_health_check(next_model))
tasks_to_index[task] = idx
peak_in_flight = max(peak_in_flight, len(tasks_to_index))
return True
for _ in range(min(concurrency_limit, len(models))):
_schedule_next()
while tasks_to_index:
done, _ = await asyncio.wait(
set(tasks_to_index.keys()),
return_when=asyncio.FIRST_COMPLETED,
)
for task in done:
idx = tasks_to_index.pop(task)
try:
results[idx] = task.result()
except Exception as e:
results[idx] = e
_schedule_next()
return results, peak_in_flight
dispatch_mode = "unbounded"
peak_in_flight = 0
if isinstance(max_concurrency, int) and max_concurrency > 0:
dispatch_mode = "bounded"
results, peak_in_flight = await _run_health_checks_with_bounded_concurrency(
model_list, max_concurrency
)
else:
tasks = [
asyncio.create_task(_run_model_health_check(model)) for model in model_list
]
peak_in_flight = len(tasks)
results = await asyncio.gather(*tasks, return_exceptions=True)
if instrumentation_enabled:
logger.debug(
"health_check_dispatch_summary source=%s cycle_id=%s mode=%s model_count=%d max_concurrency=%s peak_in_flight=%d thread_count=%d rss_mb=%s",
source,
cycle_id,
dispatch_mode,
len(model_list),
max_concurrency,
peak_in_flight,
threading.active_count(),
_rss_mb_for_log(),
)
healthy_endpoints = [] healthy_endpoints = []
unhealthy_endpoints = [] unhealthy_endpoints = []
@ -190,6 +283,8 @@ async def perform_health_check(
model: Optional[str] = None, model: Optional[str] = None,
cli_model: Optional[str] = None, cli_model: Optional[str] = None,
details: Optional[bool] = True, details: Optional[bool] = True,
max_concurrency: Optional[int] = None,
instrumentation_context: Optional[dict] = None,
): ):
""" """
Perform a health check on the system. Perform a health check on the system.
@ -197,14 +292,28 @@ async def perform_health_check(
Returns: Returns:
(bool): True if the health check passes, False otherwise. (bool): True if the health check passes, False otherwise.
""" """
instrumentation_context = instrumentation_context or {}
instrumentation_enabled = bool(instrumentation_context.get("enabled", False))
cycle_id = instrumentation_context.get("cycle_id", "unknown")
source = instrumentation_context.get("source", "unknown")
if not model_list: if not model_list:
if cli_model: if cli_model:
model_list = [ model_list = [
{"model_name": cli_model, "litellm_params": {"model": cli_model}} {"model_name": cli_model, "litellm_params": {"model": cli_model}}
] ]
else: else:
if instrumentation_enabled:
logger.debug(
"health_check_cycle_skipped source=%s cycle_id=%s reason=no_models",
source,
cycle_id,
)
return [], [] return [], []
cycle_start_time = time.monotonic()
requested_model_count = len(model_list)
if model is not None: if model is not None:
_new_model_list = [ _new_model_list = [
x for x in model_list if x["litellm_params"]["model"] == model x for x in model_list if x["litellm_params"]["model"] == model
@ -213,11 +322,56 @@ async def perform_health_check(
_new_model_list = [x for x in model_list if x["model_name"] == model] _new_model_list = [x for x in model_list if x["model_name"] == model]
model_list = _new_model_list model_list = _new_model_list
post_filter_model_count = len(model_list)
model_list = filter_deployments_by_id( model_list = filter_deployments_by_id(
model_list=model_list model_list=model_list
) # filter duplicate deployments (e.g. when model alias'es are used) ) # filter duplicate deployments (e.g. when model alias'es are used)
healthy_endpoints, unhealthy_endpoints = await _perform_health_check( deduped_model_count = len(model_list)
model_list, details
) if instrumentation_enabled:
logger.debug(
"health_check_cycle_start source=%s cycle_id=%s requested_model_count=%d post_model_filter_count=%d deduped_model_count=%d max_concurrency=%s thread_count=%d rss_mb=%s",
source,
cycle_id,
requested_model_count,
post_filter_model_count,
deduped_model_count,
max_concurrency,
threading.active_count(),
_rss_mb_for_log(),
)
try:
healthy_endpoints, unhealthy_endpoints = await _perform_health_check(
model_list,
details,
max_concurrency=max_concurrency,
instrumentation_context=instrumentation_context,
)
except Exception:
if instrumentation_enabled:
logger.exception(
"health_check_cycle_failed source=%s cycle_id=%s model_count=%d duration_ms=%.2f thread_count=%d rss_mb=%s",
source,
cycle_id,
deduped_model_count,
(time.monotonic() - cycle_start_time) * 1000,
threading.active_count(),
_rss_mb_for_log(),
)
raise
if instrumentation_enabled:
logger.debug(
"health_check_cycle_complete source=%s cycle_id=%s model_count=%d healthy_count=%d unhealthy_count=%d duration_ms=%.2f thread_count=%d rss_mb=%s",
source,
cycle_id,
deduped_model_count,
len(healthy_endpoints),
len(unhealthy_endpoints),
(time.monotonic() - cycle_start_time) * 1000,
threading.active_count(),
_rss_mb_for_log(),
)
return healthy_endpoints, unhealthy_endpoints return healthy_endpoints, unhealthy_endpoints

View File

@ -16,7 +16,7 @@ from litellm.proxy.health_check import perform_health_check
class SharedHealthCheckManager: class SharedHealthCheckManager:
""" """
Manager for coordinating health checks across multiple pods using Redis. Manager for coordinating health checks across multiple pods using Redis.
This class implements a shared health check state mechanism that: This class implements a shared health check state mechanism that:
- Prevents duplicate health checks across pods - Prevents duplicate health checks across pods
- Caches health check results with configurable TTL - Caches health check results with configurable TTL
@ -58,7 +58,7 @@ class SharedHealthCheckManager:
async def acquire_health_check_lock(self) -> bool: async def acquire_health_check_lock(self) -> bool:
""" """
Attempt to acquire the global health check lock. Attempt to acquire the global health check lock.
Returns: Returns:
bool: True if lock was acquired, False otherwise bool: True if lock was acquired, False otherwise
""" """
@ -74,7 +74,7 @@ class SharedHealthCheckManager:
nx=True, # Only set if key doesn't exist nx=True, # Only set if key doesn't exist
ttl=self.lock_ttl, ttl=self.lock_ttl,
) )
if acquired: if acquired:
verbose_proxy_logger.info( verbose_proxy_logger.info(
"Pod %s acquired health check lock", self.pod_id "Pod %s acquired health check lock", self.pod_id
@ -83,12 +83,10 @@ class SharedHealthCheckManager:
verbose_proxy_logger.debug( verbose_proxy_logger.debug(
"Pod %s failed to acquire health check lock", self.pod_id "Pod %s failed to acquire health check lock", self.pod_id
) )
return acquired return acquired
except Exception as e: except Exception as e:
verbose_proxy_logger.error( verbose_proxy_logger.error("Error acquiring health check lock: %s", str(e))
"Error acquiring health check lock: %s", str(e)
)
return False return False
async def release_health_check_lock(self) -> None: async def release_health_check_lock(self) -> None:
@ -106,14 +104,12 @@ class SharedHealthCheckManager:
"Pod %s released health check lock", self.pod_id "Pod %s released health check lock", self.pod_id
) )
except Exception as e: except Exception as e:
verbose_proxy_logger.error( verbose_proxy_logger.error("Error releasing health check lock: %s", str(e))
"Error releasing health check lock: %s", str(e)
)
async def get_cached_health_check_results(self) -> Optional[Dict[str, Any]]: async def get_cached_health_check_results(self) -> Optional[Dict[str, Any]]:
""" """
Get cached health check results from Redis. Get cached health check results from Redis.
Returns: Returns:
Optional[Dict]: Cached health check results or None if not found/expired Optional[Dict]: Cached health check results or None if not found/expired
""" """
@ -123,7 +119,7 @@ class SharedHealthCheckManager:
try: try:
cache_key = self.get_health_check_cache_key() cache_key = self.get_health_check_cache_key()
cached_data = await self.redis_cache.async_get_cache(cache_key) cached_data = await self.redis_cache.async_get_cache(cache_key)
if cached_data is None: if cached_data is None:
return None return None
@ -136,7 +132,7 @@ class SharedHealthCheckManager:
# Check if the cache is still valid # Check if the cache is still valid
cache_timestamp = cached_results.get("timestamp", 0) cache_timestamp = cached_results.get("timestamp", 0)
current_time = time.time() current_time = time.time()
if current_time - cache_timestamp > self.health_check_ttl: if current_time - cache_timestamp > self.health_check_ttl:
verbose_proxy_logger.debug("Cached health check results expired") verbose_proxy_logger.debug("Cached health check results expired")
return None return None
@ -151,13 +147,13 @@ class SharedHealthCheckManager:
return None return None
async def cache_health_check_results( async def cache_health_check_results(
self, self,
healthy_endpoints: List[Dict[str, Any]], healthy_endpoints: List[Dict[str, Any]],
unhealthy_endpoints: List[Dict[str, Any]] unhealthy_endpoints: List[Dict[str, Any]],
) -> None: ) -> None:
""" """
Cache health check results in Redis. Cache health check results in Redis.
Args: Args:
healthy_endpoints: List of healthy endpoints healthy_endpoints: List of healthy endpoints
unhealthy_endpoints: List of unhealthy endpoints unhealthy_endpoints: List of unhealthy endpoints
@ -181,7 +177,7 @@ class SharedHealthCheckManager:
safe_dumps(cache_data), safe_dumps(cache_data),
ttl=self.health_check_ttl, ttl=self.health_check_ttl,
) )
verbose_proxy_logger.info( verbose_proxy_logger.info(
"Cached health check results for %d healthy and %d unhealthy endpoints", "Cached health check results for %d healthy and %d unhealthy endpoints",
len(healthy_endpoints), len(healthy_endpoints),
@ -189,29 +185,29 @@ class SharedHealthCheckManager:
) )
except Exception as e: except Exception as e:
verbose_proxy_logger.error( verbose_proxy_logger.error("Error caching health check results: %s", str(e))
"Error caching health check results: %s", str(e)
)
async def perform_shared_health_check( async def perform_shared_health_check(
self, self,
model_list: List[Dict[str, Any]], model_list: List[Dict[str, Any]],
details: bool = True details: bool = True,
max_concurrency: Optional[int] = None,
) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]: ) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]:
""" """
Perform health check with shared state coordination. Perform health check with shared state coordination.
This method: This method:
1. First checks if there are recent cached results 1. First checks if there are recent cached results
2. If no recent cache, tries to acquire lock to run health check 2. If no recent cache, tries to acquire lock to run health check
3. If lock acquired, runs health check and caches results 3. If lock acquired, runs health check and caches results
4. If lock not acquired, waits briefly and tries to get cached results again 4. If lock not acquired, waits briefly and tries to get cached results again
5. Falls back to running health check locally if no cache available 5. Falls back to running health check locally if no cache available
Args: Args:
model_list: List of models to check model_list: List of models to check
details: Whether to include detailed information details: Whether to include detailed information
max_concurrency: Optional limit on concurrent health check requests
Returns: Returns:
Tuple of (healthy_endpoints, unhealthy_endpoints) Tuple of (healthy_endpoints, unhealthy_endpoints)
""" """
@ -225,27 +221,29 @@ class SharedHealthCheckManager:
# No recent cache, try to acquire lock # No recent cache, try to acquire lock
lock_acquired = await self.acquire_health_check_lock() lock_acquired = await self.acquire_health_check_lock()
if lock_acquired: if lock_acquired:
try: try:
# We have the lock, run health check # We have the lock, run health check
verbose_proxy_logger.info( verbose_proxy_logger.info(
"Pod %s running health check for %d models", "Pod %s running health check for %d models",
self.pod_id, self.pod_id,
len(model_list) len(model_list),
) )
healthy_endpoints, unhealthy_endpoints = await perform_health_check( healthy_endpoints, unhealthy_endpoints = await perform_health_check(
model_list=model_list, details=details model_list=model_list,
details=details,
max_concurrency=max_concurrency,
) )
# Cache the results # Cache the results
await self.cache_health_check_results( await self.cache_health_check_results(
healthy_endpoints, unhealthy_endpoints healthy_endpoints, unhealthy_endpoints
) )
return healthy_endpoints, unhealthy_endpoints return healthy_endpoints, unhealthy_endpoints
finally: finally:
# Always release the lock # Always release the lock
await self.release_health_check_lock() await self.release_health_check_lock()
@ -254,10 +252,10 @@ class SharedHealthCheckManager:
verbose_proxy_logger.debug( verbose_proxy_logger.debug(
"Pod %s waiting for other pod to complete health check", self.pod_id "Pod %s waiting for other pod to complete health check", self.pod_id
) )
# Wait a bit for the other pod to complete # Wait a bit for the other pod to complete
await asyncio.sleep(2) await asyncio.sleep(2)
# Try to get cached results again # Try to get cached results again
cached_results = await self.get_cached_health_check_results() cached_results = await self.get_cached_health_check_results()
if cached_results is not None: if cached_results is not None:
@ -265,19 +263,23 @@ class SharedHealthCheckManager:
cached_results.get("healthy_endpoints", []), cached_results.get("healthy_endpoints", []),
cached_results.get("unhealthy_endpoints", []), cached_results.get("unhealthy_endpoints", []),
) )
# Still no cache, fall back to local health check # Still no cache, fall back to local health check
verbose_proxy_logger.warning( verbose_proxy_logger.warning(
"Pod %s falling back to local health check (no cache available)", "Pod %s falling back to local health check (no cache available)",
self.pod_id self.pod_id,
)
return await perform_health_check(
model_list=model_list,
details=details,
max_concurrency=max_concurrency,
) )
return await perform_health_check(model_list=model_list, details=details)
async def is_health_check_in_progress(self) -> bool: async def is_health_check_in_progress(self) -> bool:
""" """
Check if a health check is currently in progress by another pod. Check if a health check is currently in progress by another pod.
Returns: Returns:
bool: True if health check is in progress, False otherwise bool: True if health check is in progress, False otherwise
""" """
@ -297,7 +299,7 @@ class SharedHealthCheckManager:
async def get_health_check_status(self) -> Dict[str, Any]: async def get_health_check_status(self) -> Dict[str, Any]:
""" """
Get the current status of health check coordination. Get the current status of health check coordination.
Returns: Returns:
Dict containing status information Dict containing status information
""" """
@ -320,7 +322,9 @@ class SharedHealthCheckManager:
cached_results = await self.get_cached_health_check_results() cached_results = await self.get_cached_health_check_results()
status["cache_available"] = cached_results is not None status["cache_available"] = cached_results is not None
if cached_results: if cached_results:
status["cache_age_seconds"] = time.time() - cached_results.get("timestamp", 0) status["cache_age_seconds"] = time.time() - cached_results.get(
"timestamp", 0
)
status["last_checked_by"] = cached_results.get("checked_by") status["last_checked_by"] = cached_results.get("checked_by")
except Exception as e: except Exception as e:

View File

@ -110,26 +110,31 @@ def _resolve_os_environ_variables(params: dict) -> dict:
def get_callback_identifier(callback): def get_callback_identifier(callback):
""" """
Get the callback identifier string, handling both strings and objects. Get the callback identifier string, handling both strings and objects.
This function extracts a string identifier from a callback, which can be: This function extracts a string identifier from a callback, which can be:
- A string (returned as-is) - A string (returned as-is)
- An object with a callback_name attribute - An object with a callback_name attribute
- An object registered in CustomLoggerRegistry - An object registered in CustomLoggerRegistry
- Falls back to callback_name() helper function - Falls back to callback_name() helper function
Args: Args:
callback: The callback to identify (can be str or object) callback: The callback to identify (can be str or object)
Returns: Returns:
str: The callback identifier string str: The callback identifier string
""" """
if isinstance(callback, str): if isinstance(callback, str):
return callback return callback
if hasattr(callback, 'callback_name') and callback.callback_name: if hasattr(callback, "callback_name") and callback.callback_name:
return callback.callback_name return callback.callback_name
if hasattr(callback, '__class__'): if hasattr(callback, "__class__"):
callback_strs = CustomLoggerRegistry.get_all_callback_strs_from_class_type(callback.__class__) callback_strs = CustomLoggerRegistry.get_all_callback_strs_from_class_type(
if hasattr(callback, 'callback_name') and callback.callback_name in callback_strs: callback.__class__
)
if (
hasattr(callback, "callback_name")
and callback.callback_name in callback_strs
):
return callback.callback_name return callback.callback_name
if callback_strs: if callback_strs:
return callback_strs[0] return callback_strs[0]
@ -151,7 +156,7 @@ services = Union[
"datadog_llm_observability", "datadog_llm_observability",
"generic_api", "generic_api",
"arize", "arize",
"sqs" "sqs",
], ],
str, str,
] ]
@ -224,7 +229,7 @@ async def health_services_endpoint( # noqa: PLR0915
"datadog_llm_observability", "datadog_llm_observability",
"generic_api", "generic_api",
"arize", "arize",
"sqs" "sqs",
]: ]:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
@ -238,14 +243,14 @@ async def health_services_endpoint( # noqa: PLR0915
service_in_success_callbacks = True service_in_success_callbacks = True
else: else:
for cb in litellm.success_callback: for cb in litellm.success_callback:
if hasattr(cb, 'callback_name') and cb.callback_name == service: if hasattr(cb, "callback_name") and cb.callback_name == service:
service_in_success_callbacks = True service_in_success_callbacks = True
break break
cb_id = get_callback_identifier(cb) cb_id = get_callback_identifier(cb)
if cb_id == service: if cb_id == service:
service_in_success_callbacks = True service_in_success_callbacks = True
break break
if ( if (
service == "openmeter" service == "openmeter"
or service == "braintrust" or service == "braintrust"
@ -320,6 +325,7 @@ async def health_services_endpoint( # noqa: PLR0915
) )
elif service == "sqs": elif service == "sqs":
from litellm.integrations.sqs import SQSLogger from litellm.integrations.sqs import SQSLogger
sqs_logger = SQSLogger() sqs_logger = SQSLogger()
response = await sqs_logger.async_health_check() response = await sqs_logger.async_health_check()
return { return {
@ -518,12 +524,12 @@ async def _save_health_check_to_db(
def _build_model_param_to_info_mapping(model_list: list) -> dict: def _build_model_param_to_info_mapping(model_list: list) -> dict:
""" """
Build a mapping from model parameter to model info (model_name, model_id). Build a mapping from model parameter to model info (model_name, model_id).
Multiple models might share the same model parameter, so we use a list. Multiple models might share the same model parameter, so we use a list.
Args: Args:
model_list: List of model configurations model_list: List of model configurations
Returns: Returns:
Dictionary mapping model parameter to list of model info dicts Dictionary mapping model parameter to list of model info dicts
""" """
@ -534,14 +540,16 @@ def _build_model_param_to_info_mapping(model_list: list) -> dict:
model_id = model_info.get("id") model_id = model_info.get("id")
litellm_params = model.get("litellm_params", {}) litellm_params = model.get("litellm_params", {})
model_param = litellm_params.get("model") model_param = litellm_params.get("model")
if model_param and model_name: if model_param and model_name:
if model_param not in model_param_to_info: if model_param not in model_param_to_info:
model_param_to_info[model_param] = [] model_param_to_info[model_param] = []
model_param_to_info[model_param].append({ model_param_to_info[model_param].append(
"model_name": model_name, {
"model_id": model_id, "model_name": model_name,
}) "model_id": model_id,
}
)
return model_param_to_info return model_param_to_info
@ -552,19 +560,19 @@ def _aggregate_health_check_results(
) -> dict: ) -> dict:
""" """
Aggregate health check results per unique model. Aggregate health check results per unique model.
Uses (model_id, model_name) as key, or (None, model_name) if model_id is None. Uses (model_id, model_name) as key, or (None, model_name) if model_id is None.
Args: Args:
model_param_to_info: Mapping from model parameter to model info model_param_to_info: Mapping from model parameter to model info
healthy_endpoints: List of healthy endpoint results healthy_endpoints: List of healthy endpoint results
unhealthy_endpoints: List of unhealthy endpoint results unhealthy_endpoints: List of unhealthy endpoint results
Returns: Returns:
Dictionary mapping (model_id, model_name) to aggregated health check results Dictionary mapping (model_id, model_name) to aggregated health check results
""" """
model_results = {} model_results = {}
# Process healthy endpoints # Process healthy endpoints
for endpoint in healthy_endpoints: for endpoint in healthy_endpoints:
model_param = endpoint.get("model") model_param = endpoint.get("model")
@ -580,7 +588,7 @@ def _aggregate_health_check_results(
"error_message": None, "error_message": None,
} }
model_results[key]["healthy_count"] += 1 model_results[key]["healthy_count"] += 1
# Process unhealthy endpoints # Process unhealthy endpoints
for endpoint in unhealthy_endpoints: for endpoint in unhealthy_endpoints:
model_param = endpoint.get("model") model_param = endpoint.get("model")
@ -600,7 +608,7 @@ def _aggregate_health_check_results(
# Use the first error message encountered # Use the first error message encountered
if not model_results[key]["error_message"] and error_message: if not model_results[key]["error_message"] and error_message:
model_results[key]["error_message"] = str(error_message)[:500] model_results[key]["error_message"] = str(error_message)[:500]
return model_results return model_results
@ -613,14 +621,14 @@ async def _save_health_check_results_if_changed(
): ):
""" """
Save health check results to database, but only if status changed or >1 hour since last save. Save health check results to database, but only if status changed or >1 hour since last save.
OPTIMIZATION: Only saves to database if the status has changed from the last saved check. OPTIMIZATION: Only saves to database if the status has changed from the last saved check.
This dramatically reduces database writes when health status remains stable. This dramatically reduces database writes when health status remains stable.
- Stable systems: ~1 write/hour per model (instead of 12 writes/hour with 5-min intervals) - Stable systems: ~1 write/hour per model (instead of 12 writes/hour with 5-min intervals)
- Status changes: Immediate write (no delay) - Status changes: Immediate write (no delay)
- Result: ~92% reduction in DB writes for stable systems, while maintaining real-time updates on changes - Result: ~92% reduction in DB writes for stable systems, while maintaining real-time updates on changes
Args: Args:
prisma_client: Database client prisma_client: Database client
model_results: Dictionary of aggregated health check results per model model_results: Dictionary of aggregated health check results per model
@ -630,7 +638,7 @@ async def _save_health_check_results_if_changed(
""" """
for result in model_results.values(): for result in model_results.values():
new_status = "healthy" if result["healthy_count"] > 0 else "unhealthy" new_status = "healthy" if result["healthy_count"] > 0 else "unhealthy"
# Check if we should save this result # Check if we should save this result
should_save = True should_save = True
lookup_key = result["model_id"] if result["model_id"] else result["model_name"] lookup_key = result["model_id"] if result["model_id"] else result["model_name"]
@ -641,6 +649,7 @@ async def _save_health_check_results_if_changed(
# Check if last check was recent (within 1 hour) # Check if last check was recent (within 1 hour)
if last_check.checked_at: if last_check.checked_at:
from datetime import datetime, timezone from datetime import datetime, timezone
time_since_last_check = ( time_since_last_check = (
datetime.now(timezone.utc) - last_check.checked_at datetime.now(timezone.utc) - last_check.checked_at
).total_seconds() ).total_seconds()
@ -648,7 +657,7 @@ async def _save_health_check_results_if_changed(
# This ensures we still get periodic updates even if status is stable # This ensures we still get periodic updates even if status is stable
if time_since_last_check < 3600: # 1 hour threshold if time_since_last_check < 3600: # 1 hour threshold
should_save = False should_save = False
if should_save: if should_save:
asyncio.create_task( asyncio.create_task(
prisma_client.save_health_check_result( prisma_client.save_health_check_result(
@ -675,27 +684,27 @@ async def _save_background_health_checks_to_db(
): ):
""" """
Save background health check results to database for each model. Save background health check results to database for each model.
Maps health check endpoints back to their original models to get model_name and model_id. Maps health check endpoints back to their original models to get model_name and model_id.
Aggregates results per unique model (by model_id if available, otherwise model_name). Aggregates results per unique model (by model_id if available, otherwise model_name).
OPTIMIZATION: Only saves to database if the status has changed from the last saved check. OPTIMIZATION: Only saves to database if the status has changed from the last saved check.
This dramatically reduces database writes when health status remains stable. This dramatically reduces database writes when health status remains stable.
""" """
if prisma_client is None: if prisma_client is None:
return return
try: try:
# Step 1: Build mapping from model parameter to model info # Step 1: Build mapping from model parameter to model info
model_param_to_info = _build_model_param_to_info_mapping(model_list) model_param_to_info = _build_model_param_to_info_mapping(model_list)
# Step 2: Aggregate health check results per unique model # Step 2: Aggregate health check results per unique model
model_results = _aggregate_health_check_results( model_results = _aggregate_health_check_results(
model_param_to_info, model_param_to_info,
healthy_endpoints, healthy_endpoints,
unhealthy_endpoints, unhealthy_endpoints,
) )
# Step 3: Get latest health checks for all models in one query to compare status # Step 3: Get latest health checks for all models in one query to compare status
latest_checks = await prisma_client.get_all_latest_health_checks() latest_checks = await prisma_client.get_all_latest_health_checks()
latest_checks_map = {} latest_checks_map = {}
@ -704,7 +713,7 @@ async def _save_background_health_checks_to_db(
key = check.model_id if check.model_id else check.model_name key = check.model_id if check.model_id else check.model_name
if key not in latest_checks_map: if key not in latest_checks_map:
latest_checks_map[key] = check latest_checks_map[key] = check
# Step 4: Save aggregated results, but only if status changed # Step 4: Save aggregated results, but only if status changed
await _save_health_check_results_if_changed( await _save_health_check_results_if_changed(
prisma_client, prisma_client,
@ -729,10 +738,15 @@ async def _perform_health_check_and_save(
start_time, start_time,
user_id, user_id,
model_id=None, model_id=None,
max_concurrency=None,
): ):
"""Helper function to perform health check and save results to database""" """Helper function to perform health check and save results to database"""
healthy_endpoints, unhealthy_endpoints = await perform_health_check( healthy_endpoints, unhealthy_endpoints = await perform_health_check(
model_list=model_list, cli_model=cli_model, model=target_model, details=details model_list=model_list,
cli_model=cli_model,
model=target_model,
details=details,
max_concurrency=max_concurrency,
) )
# Optionally save health check result to database (non-blocking) # Optionally save health check result to database (non-blocking)
@ -789,6 +803,7 @@ async def health_endpoint(
import time import time
from litellm.proxy.proxy_server import ( from litellm.proxy.proxy_server import (
health_check_concurrency,
health_check_details, health_check_details,
health_check_results, health_check_results,
llm_model_list, llm_model_list,
@ -841,6 +856,7 @@ async def health_endpoint(
start_time=start_time, start_time=start_time,
user_id=user_api_key_dict.user_id, user_id=user_api_key_dict.user_id,
model_id=None, # CLI model doesn't have model_id model_id=None, # CLI model doesn't have model_id
max_concurrency=health_check_concurrency,
) )
raise HTTPException( raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
@ -864,6 +880,7 @@ async def health_endpoint(
start_time=start_time, start_time=start_time,
user_id=user_api_key_dict.user_id, user_id=user_api_key_dict.user_id,
model_id=model_id, model_id=model_id,
max_concurrency=health_check_concurrency,
) )
except Exception as e: except Exception as e:
verbose_proxy_logger.error( verbose_proxy_logger.error(
@ -1420,11 +1437,11 @@ async def test_model_connection(
status_code=500, status_code=500,
detail={"error": CommonProxyErrors.db_not_connected_error.value}, detail={"error": CommonProxyErrors.db_not_connected_error.value},
) )
# Get model name from litellm_params # Get model name from litellm_params
request_litellm_params = litellm_params or {} request_litellm_params = litellm_params or {}
model_name = request_litellm_params.get("model") model_name = request_litellm_params.get("model")
# Look up model configuration from router if model name is provided # Look up model configuration from router if model name is provided
# This gets the litellm_params from proxy config (with resolved env vars) # This gets the litellm_params from proxy config (with resolved env vars)
config_litellm_params: dict = {} config_litellm_params: dict = {}
@ -1432,34 +1449,39 @@ async def test_model_connection(
try: try:
# First try to find by proxy model_name (e.g., "gpt-4o") # First try to find by proxy model_name (e.g., "gpt-4o")
deployments = llm_router.get_model_list(model_name=model_name) deployments = llm_router.get_model_list(model_name=model_name)
# If not found, try to find by litellm model name (e.g., "azure/gpt-4o") # If not found, try to find by litellm model name (e.g., "azure/gpt-4o")
if not deployments or len(deployments) == 0: if not deployments or len(deployments) == 0:
all_deployments = llm_router.get_model_list(model_name=None) all_deployments = llm_router.get_model_list(model_name=None)
if all_deployments: if all_deployments:
for deployment in all_deployments: for deployment in all_deployments:
if deployment.get("litellm_params", {}).get("model") == model_name: if (
deployment.get("litellm_params", {}).get("model")
== model_name
):
deployments = [deployment] deployments = [deployment]
break break
if deployments and len(deployments) > 0: if deployments and len(deployments) > 0:
# Use the first deployment's litellm_params as base config # Use the first deployment's litellm_params as base config
# These already have resolved environment variables from proxy config # These already have resolved environment variables from proxy config
config_litellm_params = dict(deployments[0].get("litellm_params", {})) config_litellm_params = dict(
deployments[0].get("litellm_params", {})
)
except Exception as e: except Exception as e:
verbose_proxy_logger.debug( verbose_proxy_logger.debug(
f"Could not find model {model_name} in router: {e}. " f"Could not find model {model_name} in router: {e}. "
"Proceeding with request params only." "Proceeding with request params only."
) )
# Merge: config params (from proxy config) as base, request params override # Merge: config params (from proxy config) as base, request params override
# This allows users to override specific params while using config for credentials # This allows users to override specific params while using config for credentials
merged_litellm_params = {**config_litellm_params, **request_litellm_params} merged_litellm_params = {**config_litellm_params, **request_litellm_params}
# Resolve os.environ/ environment variables in any remaining request params # Resolve os.environ/ environment variables in any remaining request params
# This handles cases where user explicitly passes os.environ/ values to override config # This handles cases where user explicitly passes os.environ/ values to override config
litellm_params = _resolve_os_environ_variables(merged_litellm_params) litellm_params = _resolve_os_environ_variables(merged_litellm_params)
## Auth check ## Auth check
await ModelManagementAuthChecks.can_user_make_model_call( await ModelManagementAuthChecks.can_user_make_model_call(
model_params=Deployment( model_params=Deployment(

View File

@ -9,6 +9,7 @@ import secrets
import shutil import shutil
import subprocess import subprocess
import sys import sys
import threading
import time import time
import traceback import traceback
import warnings import warnings
@ -658,7 +659,7 @@ _description = (
def cleanup_router_config_variables(): def cleanup_router_config_variables():
global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, prisma_client global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, health_check_concurrency, prisma_client
# Set all variables to None # Set all variables to None
master_key = None master_key = None
@ -672,6 +673,7 @@ def cleanup_router_config_variables():
use_background_health_checks = None use_background_health_checks = None
use_shared_health_check = None use_shared_health_check = None
health_check_interval = None health_check_interval = None
health_check_concurrency = None
prisma_client = None prisma_client = None
@ -822,7 +824,9 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
verbose_proxy_logger.debug("About to initialize semantic tool filter") verbose_proxy_logger.debug("About to initialize semantic tool filter")
_config = proxy_config.get_config_state() _config = proxy_config.get_config_state()
_litellm_settings = _config.get("litellm_settings", {}) _litellm_settings = _config.get("litellm_settings", {})
verbose_proxy_logger.debug(f"litellm_settings keys = {list(_litellm_settings.keys())}") verbose_proxy_logger.debug(
f"litellm_settings keys = {list(_litellm_settings.keys())}"
)
await ProxyStartupEvent._initialize_semantic_tool_filter( await ProxyStartupEvent._initialize_semantic_tool_filter(
llm_router=llm_router, llm_router=llm_router,
litellm_settings=_litellm_settings, litellm_settings=_litellm_settings,
@ -1468,7 +1472,9 @@ redis_usage_cache: Optional[
RedisCache RedisCache
] = None # redis cache used for tracking spend, tpm/rpm limits ] = None # redis cache used for tracking spend, tpm/rpm limits
polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False
native_background_mode: List[str] = [] # Models that should use native provider background mode instead of polling native_background_mode: List[
str
] = [] # Models that should use native provider background mode instead of polling
polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache
user_custom_auth = None user_custom_auth = None
user_custom_key_generate = None user_custom_key_generate = None
@ -1478,8 +1484,11 @@ use_background_health_checks = None
use_shared_health_check = None use_shared_health_check = None
use_queue = False use_queue = False
health_check_interval = None health_check_interval = None
health_check_concurrency = None
health_check_details = None health_check_details = None
health_check_results: Dict[str, Union[int, List[Dict[str, Any]]]] = {} health_check_results: Dict[str, Union[int, List[Dict[str, Any]]]] = {}
background_health_check_loop_active = False
background_health_check_cycle_seq = 0
queue: List = [] queue: List = []
litellm_proxy_budget_name = "litellm-proxy-budget" litellm_proxy_budget_name = "litellm-proxy-budget"
litellm_proxy_admin_name = LITELLM_PROXY_ADMIN_NAME litellm_proxy_admin_name = LITELLM_PROXY_ADMIN_NAME
@ -1927,6 +1936,29 @@ def run_ollama_serve():
) )
def _get_process_rss_mb() -> Optional[float]:
"""
Get process RSS memory in MB.
On Linux, ru_maxrss is in KB. On macOS, ru_maxrss is in bytes.
"""
try:
import resource
ru_maxrss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
if sys.platform == "darwin":
return float(ru_maxrss) / (1024 * 1024)
return float(ru_maxrss) / 1024
except Exception:
return None
def _rss_mb_for_log() -> str:
rss_mb = _get_process_rss_mb()
if rss_mb is None:
return "unknown"
return f"{rss_mb:.2f}"
async def _run_background_health_check(): async def _run_background_health_check():
""" """
Periodically run health checks in the background on the endpoints. Periodically run health checks in the background on the endpoints.
@ -1934,7 +1966,10 @@ async def _run_background_health_check():
Update health_check_results, based on this. Update health_check_results, based on this.
Uses shared health check state when Redis is available to coordinate across pods. Uses shared health check state when Redis is available to coordinate across pods.
""" """
global health_check_results, llm_model_list, health_check_interval, health_check_details, use_shared_health_check, redis_usage_cache, prisma_client global health_check_results, llm_model_list, health_check_interval
global health_check_concurrency, health_check_details, use_shared_health_check
global redis_usage_cache, prisma_client
global background_health_check_loop_active, background_health_check_cycle_seq
if ( if (
health_check_interval is None health_check_interval is None
@ -1943,6 +1978,24 @@ async def _run_background_health_check():
): ):
return return
if background_health_check_loop_active:
verbose_proxy_logger.warning(
"background_health_check_loop_overlap_detected existing_loop_active=true interval_seconds=%s max_concurrency=%s shared=%s",
health_check_interval,
health_check_concurrency,
use_shared_health_check,
)
background_health_check_loop_active = True
verbose_proxy_logger.info(
"background_health_check_loop_started interval_seconds=%s max_concurrency=%s shared=%s details=%s thread_count=%d rss_mb=%s",
health_check_interval,
health_check_concurrency,
use_shared_health_check,
health_check_details,
threading.active_count(),
_rss_mb_for_log(),
)
# Initialize shared health check manager if Redis is available and feature is enabled # Initialize shared health check manager if Redis is available and feature is enabled
shared_health_manager = None shared_health_manager = None
if use_shared_health_check and redis_usage_cache is not None: if use_shared_health_check and redis_usage_cache is not None:
@ -1958,8 +2011,13 @@ async def _run_background_health_check():
verbose_proxy_logger.info("Initialized shared health check manager") verbose_proxy_logger.info("Initialized shared health check manager")
while True: while True:
background_health_check_cycle_seq += 1
cycle_id = f"bg-{background_health_check_cycle_seq}"
cycle_start_time = time.monotonic()
# make 1 deep copy of llm_model_list on every health check iteration # make 1 deep copy of llm_model_list on every health check iteration
_llm_model_list = copy.deepcopy(llm_model_list) or [] _llm_model_list = copy.deepcopy(llm_model_list) or []
model_count_total = len(_llm_model_list)
# filter out models that have disabled background health checks # filter out models that have disabled background health checks
_llm_model_list = [ _llm_model_list = [
@ -1967,6 +2025,52 @@ async def _run_background_health_check():
for m in _llm_model_list for m in _llm_model_list
if not m.get("model_info", {}).get("disable_background_health_check", False) if not m.get("model_info", {}).get("disable_background_health_check", False)
] ]
model_count_enabled = len(_llm_model_list)
expected_peak_in_flight = model_count_enabled
if (
isinstance(health_check_concurrency, int)
and health_check_concurrency > 0
and model_count_enabled > 0
):
expected_peak_in_flight = min(model_count_enabled, health_check_concurrency)
verbose_proxy_logger.debug(
"background_health_check_cycle_start cycle_id=%s model_count_total=%d model_count_enabled=%d interval_seconds=%s max_concurrency=%s expected_peak_in_flight=%d shared=%s thread_count=%d rss_mb=%s",
cycle_id,
model_count_total,
model_count_enabled,
health_check_interval,
health_check_concurrency,
expected_peak_in_flight,
shared_health_manager is not None,
threading.active_count(),
_rss_mb_for_log(),
)
instrumentation_context = {
"enabled": True,
"source": "proxy_background_loop",
"cycle_id": cycle_id,
}
async def _run_direct_health_check_with_instrumentation():
try:
return await perform_health_check(
model_list=_llm_model_list,
details=health_check_details,
max_concurrency=health_check_concurrency,
instrumentation_context=instrumentation_context,
)
except TypeError as e:
if "instrumentation_context" not in str(e):
raise
# Backward compatibility for monkeypatched or wrapped callables
# that do not accept instrumentation_context.
return await perform_health_check(
model_list=_llm_model_list,
details=health_check_details,
max_concurrency=health_check_concurrency,
)
# Use shared health check if available, otherwise fall back to direct health check # Use shared health check if available, otherwise fall back to direct health check
# Convert health_check_details to bool for perform_shared_health_check (defaults to True if None) # Convert health_check_details to bool for perform_shared_health_check (defaults to True if None)
@ -1980,19 +2084,21 @@ async def _run_background_health_check():
healthy_endpoints, healthy_endpoints,
unhealthy_endpoints, unhealthy_endpoints,
) = await shared_health_manager.perform_shared_health_check( ) = await shared_health_manager.perform_shared_health_check(
model_list=_llm_model_list, details=details_bool model_list=_llm_model_list,
details=details_bool,
max_concurrency=health_check_concurrency,
) )
except Exception as e: except Exception as e:
verbose_proxy_logger.error( verbose_proxy_logger.error(
"Error in shared health check, falling back to direct health check: %s", "Error in shared health check, falling back to direct health check: %s",
str(e), str(e),
) )
healthy_endpoints, unhealthy_endpoints = await perform_health_check( healthy_endpoints, unhealthy_endpoints = (
model_list=_llm_model_list, details=health_check_details await _run_direct_health_check_with_instrumentation()
) )
else: else:
healthy_endpoints, unhealthy_endpoints = await perform_health_check( healthy_endpoints, unhealthy_endpoints = (
model_list=_llm_model_list, details=health_check_details await _run_direct_health_check_with_instrumentation()
) )
# Update the global variable with the health check results # Update the global variable with the health check results
@ -2000,6 +2106,25 @@ async def _run_background_health_check():
health_check_results["unhealthy_endpoints"] = unhealthy_endpoints health_check_results["unhealthy_endpoints"] = unhealthy_endpoints
health_check_results["healthy_count"] = len(healthy_endpoints) health_check_results["healthy_count"] = len(healthy_endpoints)
health_check_results["unhealthy_count"] = len(unhealthy_endpoints) health_check_results["unhealthy_count"] = len(unhealthy_endpoints)
cycle_duration_ms = (time.monotonic() - cycle_start_time) * 1000
verbose_proxy_logger.debug(
"background_health_check_cycle_complete cycle_id=%s model_count_enabled=%d healthy_count=%d unhealthy_count=%d duration_ms=%.2f interval_seconds=%s thread_count=%d rss_mb=%s",
cycle_id,
model_count_enabled,
len(healthy_endpoints),
len(unhealthy_endpoints),
cycle_duration_ms,
health_check_interval,
threading.active_count(),
_rss_mb_for_log(),
)
if cycle_duration_ms > (health_check_interval * 1000):
verbose_proxy_logger.warning(
"background_health_check_cycle_duration_exceeded_interval cycle_id=%s duration_ms=%.2f interval_seconds=%s",
cycle_id,
cycle_duration_ms,
health_check_interval,
)
# Save background health checks to database (non-blocking) # Save background health checks to database (non-blocking)
if prisma_client is not None: if prisma_client is not None:
@ -2480,7 +2605,7 @@ class ProxyConfig:
""" """
Load config values into proxy global state Load config values into proxy global state
""" """
global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, use_queue, proxy_budget_rescheduler_max_time, proxy_budget_rescheduler_min_time, ui_access_mode, litellm_master_key_hash, proxy_batch_write_at, disable_spend_logs, prompt_injection_detection_obj, redis_usage_cache, store_model_in_db, premium_user, open_telemetry_logger, health_check_details, proxy_batch_polling_interval, config_passthrough_endpoints global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, health_check_concurrency, use_queue, proxy_budget_rescheduler_max_time, proxy_budget_rescheduler_min_time, ui_access_mode, litellm_master_key_hash, proxy_batch_write_at, disable_spend_logs, prompt_injection_detection_obj, redis_usage_cache, store_model_in_db, premium_user, open_telemetry_logger, health_check_details, proxy_batch_polling_interval, config_passthrough_endpoints
config: dict = await self.get_config(config_file_path=config_file_path) config: dict = await self.get_config(config_file_path=config_file_path)
@ -2905,7 +3030,18 @@ class ProxyConfig:
health_check_interval = general_settings.get( health_check_interval = general_settings.get(
"health_check_interval", DEFAULT_HEALTH_CHECK_INTERVAL "health_check_interval", DEFAULT_HEALTH_CHECK_INTERVAL
) )
health_check_concurrency = general_settings.get(
"health_check_concurrency", None
)
health_check_details = general_settings.get("health_check_details", True) health_check_details = general_settings.get("health_check_details", True)
verbose_proxy_logger.info(
"background_health_check_config enabled=%s shared=%s interval_seconds=%s max_concurrency=%s details=%s",
use_background_health_checks,
use_shared_health_check,
health_check_interval,
health_check_concurrency,
health_check_details,
)
### RBAC ### ### RBAC ###
rbac_role_permissions = general_settings.get("role_permissions", None) rbac_role_permissions = general_settings.get("role_permissions", None)
@ -2999,7 +3135,7 @@ class ProxyConfig:
for k, v in router_settings.items(): for k, v in router_settings.items():
if k in available_args: if k in available_args:
router_params[k] = v router_params[k] = v
elif k == "health_check_interval": elif k in {"health_check_interval", "health_check_concurrency"}:
raise ValueError( raise ValueError(
f"'{k}' is NOT a valid router_settings parameter. Please move it to 'general_settings'." f"'{k}' is NOT a valid router_settings parameter. Please move it to 'general_settings'."
) )
@ -4201,9 +4337,7 @@ class ProxyConfig:
) )
if self._should_load_db_object(object_type="semantic_filter_settings"): if self._should_load_db_object(object_type="semantic_filter_settings"):
await self._init_semantic_filter_settings_in_db( await self._init_semantic_filter_settings_in_db(prisma_client=prisma_client)
prisma_client=prisma_client
)
async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient): async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient):
""" """
@ -5259,30 +5393,38 @@ class ProxyStartupEvent:
): ):
"""Initialize MCP semantic tool filter if configured""" """Initialize MCP semantic tool filter if configured"""
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
mcp_semantic_filter_config = litellm_settings.get("mcp_semantic_tool_filter", None) mcp_semantic_filter_config = litellm_settings.get(
"mcp_semantic_tool_filter", None
)
# Only proceed if the feature is configured and enabled # Only proceed if the feature is configured and enabled
if not mcp_semantic_filter_config or not mcp_semantic_filter_config.get("enabled", False): if not mcp_semantic_filter_config or not mcp_semantic_filter_config.get(
verbose_proxy_logger.debug("Semantic tool filter not configured or not enabled, skipping initialization") "enabled", False
):
verbose_proxy_logger.debug(
"Semantic tool filter not configured or not enabled, "
"skipping initialization"
)
return return
verbose_proxy_logger.debug( verbose_proxy_logger.debug(
f"Initializing semantic tool filter: llm_router={llm_router is not None}, " f"Initializing semantic tool filter: llm_router={llm_router is not None}, "
f"config={mcp_semantic_filter_config}" f"config={mcp_semantic_filter_config}"
) )
hook = await SemanticToolFilterHook.initialize_from_config( hook = await SemanticToolFilterHook.initialize_from_config(
config=mcp_semantic_filter_config, config=mcp_semantic_filter_config,
llm_router=llm_router, llm_router=llm_router,
) )
if hook: if hook:
verbose_proxy_logger.debug("Semantic tool filter hook registered") verbose_proxy_logger.debug("Semantic tool filter hook registered")
litellm.logging_callback_manager.add_litellm_callback(hook) litellm.logging_callback_manager.add_litellm_callback(hook)
else: else:
# Only warn if the feature was configured but failed to initialize # Only warn if the feature was configured but failed to initialize
verbose_proxy_logger.warning("Semantic tool filter hook was configured but failed to initialize") verbose_proxy_logger.warning(
"Semantic tool filter hook was configured but failed to initialize"
)
@classmethod @classmethod
def _initialize_jwt_auth( def _initialize_jwt_auth(
@ -8706,7 +8848,8 @@ async def _apply_search_filter_to_models(
# Fetch database models if we need more for the current page # Fetch database models if we need more for the current page
if router_models_count < models_needed_for_page: if router_models_count < models_needed_for_page:
models_to_fetch = min( models_to_fetch = min(
models_needed_for_page - router_models_count, db_models_total_count models_needed_for_page - router_models_count,
db_models_total_count,
) )
if models_to_fetch > 0: if models_to_fetch > 0:
@ -8742,21 +8885,21 @@ async def _apply_search_filter_to_models(
def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]: def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]:
""" """
Normalize a datetime value to a timezone-aware UTC datetime for sorting. Normalize a datetime value to a timezone-aware UTC datetime for sorting.
This function handles: This function handles:
- None values: returns None - None values: returns None
- String values: parses ISO format strings and converts to UTC-aware datetime - String values: parses ISO format strings and converts to UTC-aware datetime
- Datetime objects: converts naive datetimes to UTC-aware, and aware datetimes to UTC - Datetime objects: converts naive datetimes to UTC-aware, and aware datetimes to UTC
Args: Args:
dt: Datetime value (None, str, or datetime object) dt: Datetime value (None, str, or datetime object)
Returns: Returns:
UTC-aware datetime object, or None if input is None or cannot be parsed UTC-aware datetime object, or None if input is None or cannot be parsed
""" """
if dt is None: if dt is None:
return None return None
if isinstance(dt, str): if isinstance(dt, str):
try: try:
# Handle ISO format strings, including 'Z' suffix # Handle ISO format strings, including 'Z' suffix
@ -8770,14 +8913,14 @@ def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]:
return parsed_dt return parsed_dt
except (ValueError, AttributeError): except (ValueError, AttributeError):
return None return None
if isinstance(dt, datetime): if isinstance(dt, datetime):
# If naive, assume UTC and make it aware # If naive, assume UTC and make it aware
if dt.tzinfo is None: if dt.tzinfo is None:
return dt.replace(tzinfo=timezone.utc) return dt.replace(tzinfo=timezone.utc)
# If aware, convert to UTC # If aware, convert to UTC
return dt.astimezone(timezone.utc) return dt.astimezone(timezone.utc)
return None return None
@ -8797,46 +8940,60 @@ def _sort_models(
Returns: Returns:
Sorted list of models Sorted list of models
""" """
if not sort_by or sort_by not in ["model_name", "created_at", "updated_at", "costs", "status"]: if not sort_by or sort_by not in [
"model_name",
"created_at",
"updated_at",
"costs",
"status",
]:
return all_models return all_models
reverse = sort_order.lower() == "desc" reverse = sort_order.lower() == "desc"
def get_sort_key(model: Dict[str, Any]) -> Any: def get_sort_key(model: Dict[str, Any]) -> Any:
model_info = model.get("model_info", {}) model_info = model.get("model_info", {})
if sort_by == "model_name": if sort_by == "model_name":
return model.get("model_name", "").lower() return model.get("model_name", "").lower()
elif sort_by == "created_at": elif sort_by == "created_at":
created_at = model_info.get("created_at") created_at = model_info.get("created_at")
normalized_dt = _normalize_datetime_for_sorting(created_at) normalized_dt = _normalize_datetime_for_sorting(created_at)
if normalized_dt is None: if normalized_dt is None:
# Put None values at the end for asc, at the start for desc # Put None values at the end for asc, at the start for desc
return (datetime.max.replace(tzinfo=timezone.utc) if not reverse else datetime.min.replace(tzinfo=timezone.utc)) return (
datetime.max.replace(tzinfo=timezone.utc)
if not reverse
else datetime.min.replace(tzinfo=timezone.utc)
)
return normalized_dt return normalized_dt
elif sort_by == "updated_at": elif sort_by == "updated_at":
updated_at = model_info.get("updated_at") updated_at = model_info.get("updated_at")
normalized_dt = _normalize_datetime_for_sorting(updated_at) normalized_dt = _normalize_datetime_for_sorting(updated_at)
if normalized_dt is None: if normalized_dt is None:
return (datetime.max.replace(tzinfo=timezone.utc) if not reverse else datetime.min.replace(tzinfo=timezone.utc)) return (
datetime.max.replace(tzinfo=timezone.utc)
if not reverse
else datetime.min.replace(tzinfo=timezone.utc)
)
return normalized_dt return normalized_dt
elif sort_by == "costs": elif sort_by == "costs":
input_cost = model_info.get("input_cost_per_token", 0) or 0 input_cost = model_info.get("input_cost_per_token", 0) or 0
output_cost = model_info.get("output_cost_per_token", 0) or 0 output_cost = model_info.get("output_cost_per_token", 0) or 0
total_cost = input_cost + output_cost total_cost = input_cost + output_cost
# Put 0 or None costs at the end for asc, at the start for desc # Put 0 or None costs at the end for asc, at the start for desc
if total_cost == 0: if total_cost == 0:
return (float("inf") if not reverse else float("-inf")) return float("inf") if not reverse else float("-inf")
return total_cost return total_cost
elif sort_by == "status": elif sort_by == "status":
# False (config) comes before True (db) for asc # False (config) comes before True (db) for asc
db_model = model_info.get("db_model", False) db_model = model_info.get("db_model", False)
return db_model return db_model
return None return None
try: try:
@ -9032,9 +9189,7 @@ async def _find_model_by_id(
) )
if db_model: if db_model:
# Convert database model to router format # Convert database model to router format
decrypted_models = proxy_config.decrypt_model_list_from_db( decrypted_models = proxy_config.decrypt_model_list_from_db([db_model])
[db_model]
)
if decrypted_models: if decrypted_models:
found_model = decrypted_models[0] found_model = decrypted_models[0]
except Exception as e: except Exception as e:
@ -9208,13 +9363,13 @@ async def model_info_v2(
) )
verbose_proxy_logger.debug("all_models: %s", all_models) verbose_proxy_logger.debug("all_models: %s", all_models)
# Append A2A agents to models list # Append A2A agents to models list
all_models = await append_agents_to_model_info( all_models = await append_agents_to_model_info(
models=all_models, models=all_models,
user_api_key_dict=user_api_key_dict, user_api_key_dict=user_api_key_dict,
) )
# Update total count to include agents # Update total count to include agents
search_total_count = len(all_models) search_total_count = len(all_models)
@ -10057,7 +10212,7 @@ async def model_group_info(
model_groups: List[ModelGroupInfoProxy] = _get_model_group_info( model_groups: List[ModelGroupInfoProxy] = _get_model_group_info(
llm_router=llm_router, all_models_str=all_models_str, model_group=model_group llm_router=llm_router, all_models_str=all_models_str, model_group=model_group
) )
# Append A2A agents to model groups # Append A2A agents to model groups
model_groups = await append_agents_to_model_group( model_groups = await append_agents_to_model_group(
model_groups=model_groups, model_groups=model_groups,

View File

@ -92,7 +92,7 @@ async def test_azure_img_gen_health_check():
litellm._turn_on_debug() litellm._turn_on_debug()
max_retries = 3 max_retries = 3
retry_delay = 1 # Start with 1 second delay retry_delay = 1 # Start with 1 second delay
for attempt in range(max_retries): for attempt in range(max_retries):
response = await litellm.ahealth_check( response = await litellm.ahealth_check(
model_params={ model_params={
@ -103,11 +103,11 @@ async def test_azure_img_gen_health_check():
mode="image_generation", mode="image_generation",
prompt="cute baby sea otter", prompt="cute baby sea otter",
) )
# Check if response is successful (no error) # Check if response is successful (no error)
if isinstance(response, dict) and "error" not in response: if isinstance(response, dict) and "error" not in response:
return response return response
# Check if error is a transient Azure internal server error # Check if error is a transient Azure internal server error
error_str = str(response.get("error", "")).lower() error_str = str(response.get("error", "")).lower()
is_transient_error = ( is_transient_error = (
@ -116,16 +116,18 @@ async def test_azure_img_gen_health_check():
or "internalfailure" in error_str or "internalfailure" in error_str
or "internal failure" in error_str or "internal failure" in error_str
) )
# If it's the last attempt or not a transient error, fail the test # If it's the last attempt or not a transient error, fail the test
if attempt == max_retries - 1 or not is_transient_error: if attempt == max_retries - 1 or not is_transient_error:
assert isinstance(response, dict) and "error" not in response, f"Health check failed: {response.get('error', 'Unknown error')}" assert (
isinstance(response, dict) and "error" not in response
), f"Health check failed: {response.get('error', 'Unknown error')}"
return response return response
# Wait before retrying with exponential backoff # Wait before retrying with exponential backoff
await asyncio.sleep(retry_delay) await asyncio.sleep(retry_delay)
retry_delay *= 2 # Exponential backoff retry_delay *= 2 # Exponential backoff
# Should not reach here, but just in case # Should not reach here, but just in case
assert False, "Health check failed after all retries" assert False, "Health check failed after all retries"
@ -562,6 +564,99 @@ async def test_health_check_bad_model():
), "Health check took longer than health_check_timeout" ), "Health check took longer than health_check_timeout"
@pytest.mark.asyncio
async def test_health_check_respects_concurrency_limit():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}}
for i in range(6)
]
active = 0
max_active = 0
async def mock_health_check(litellm_params, **kwargs):
nonlocal active, max_active
active += 1
max_active = max(max_active, active)
await asyncio.sleep(0.05)
active -= 1
return {"status": "healthy"}
with patch("litellm.ahealth_check", side_effect=mock_health_check):
await _perform_health_check(model_list, max_concurrency=2)
assert max_active <= 2
@pytest.mark.asyncio
async def test_health_check_creates_only_bounded_initial_tasks():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}}
for i in range(10)
]
release_event = asyncio.Event()
create_task_call_count = 0
real_create_task = asyncio.create_task
async def mock_health_check(litellm_params, **kwargs):
await release_event.wait()
return {"status": "healthy"}
def tracked_create_task(coro):
nonlocal create_task_call_count
create_task_call_count += 1
return real_create_task(coro)
with patch("litellm.ahealth_check", side_effect=mock_health_check), patch(
"litellm.proxy.health_check.asyncio.create_task", side_effect=tracked_create_task
):
perform_task = real_create_task(
_perform_health_check(model_list, max_concurrency=2)
)
await asyncio.sleep(0.05)
assert create_task_call_count == 2
release_event.set()
await perform_task
@pytest.mark.asyncio
async def test_timeout_does_not_cancel_other_health_checks():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{
"litellm_params": {"model": "openai/slow-model", "api_key": "fake-key"},
"model_info": {"health_check_timeout": 0.05},
},
{
"litellm_params": {"model": "openai/fast-model", "api_key": "fake-key"},
"model_info": {"health_check_timeout": 1},
},
]
async def mock_health_check(litellm_params, **kwargs):
if litellm_params["model"] == "openai/slow-model":
await asyncio.sleep(0.2)
return {"status": "healthy"}
await asyncio.sleep(0.01)
return {"status": "healthy"}
with patch("litellm.ahealth_check", side_effect=mock_health_check):
healthy_endpoints, unhealthy_endpoints = await _perform_health_check(
model_list, max_concurrency=1
)
healthy_models = {endpoint["model"] for endpoint in healthy_endpoints}
unhealthy_models = {endpoint["model"] for endpoint in unhealthy_endpoints}
assert "openai/fast-model" in healthy_models
assert "openai/slow-model" in unhealthy_models
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ahealth_check_ocr(): async def test_ahealth_check_ocr():
litellm._turn_on_debug() litellm._turn_on_debug()
@ -643,20 +738,20 @@ async def test_image_generation_health_check_prompt(monkeypatch):
async def test_health_check_with_custom_llm_provider(): async def test_health_check_with_custom_llm_provider():
""" """
Test that ahealth_check correctly uses custom_llm_provider from model_params. Test that ahealth_check correctly uses custom_llm_provider from model_params.
This test verifies the fix for the issue where the UI's "Test connect" button This test verifies the fix for the issue where the UI's "Test connect" button
failed with "LLM Provider NOT provided" error for OpenAI-compatible self-hosted failed with "LLM Provider NOT provided" error for OpenAI-compatible self-hosted
providers, even when a provider was selected in the dropdown. providers, even when a provider was selected in the dropdown.
The fix ensures that when custom_llm_provider is passed in model_params, The fix ensures that when custom_llm_provider is passed in model_params,
it's properly forwarded to get_llm_provider() to identify the correct provider. it's properly forwarded to get_llm_provider() to identify the correct provider.
""" """
from unittest.mock import MagicMock from unittest.mock import MagicMock
# Mock the completion call to avoid making real API calls # Mock the completion call to avoid making real API calls
mock_response = MagicMock() mock_response = MagicMock()
mock_response._hidden_params = {"headers": {"x-ratelimit-remaining-tokens": "1000"}} mock_response._hidden_params = {"headers": {"x-ratelimit-remaining-tokens": "1000"}}
with patch("litellm.acompletion", return_value=mock_response): with patch("litellm.acompletion", return_value=mock_response):
# Test with a custom model name that wouldn't be recognized without custom_llm_provider # Test with a custom model name that wouldn't be recognized without custom_llm_provider
response = await litellm.ahealth_check( response = await litellm.ahealth_check(
@ -668,7 +763,7 @@ async def test_health_check_with_custom_llm_provider():
}, },
mode="chat", mode="chat",
) )
# Should succeed without "LLM Provider NOT provided" error # Should succeed without "LLM Provider NOT provided" error
assert "error" not in response assert "error" not in response
assert isinstance(response, dict) assert isinstance(response, dict)

View File

@ -2330,7 +2330,9 @@ async def test_run_background_health_check_reflects_llm_model_list(monkeypatch):
test_model_list_2 = [{"model_name": "model-b"}] test_model_list_2 = [{"model_name": "model-b"}]
called_model_lists = [] called_model_lists = []
async def fake_perform_health_check(model_list, details): async def fake_perform_health_check(
model_list, details, max_concurrency=None
):
called_model_lists.append(copy.deepcopy(model_list)) called_model_lists.append(copy.deepcopy(model_list))
return (["healthy"], ["unhealthy"]) return (["healthy"], ["unhealthy"])
@ -2378,7 +2380,9 @@ async def test_background_health_check_skip_disabled_models(monkeypatch):
] ]
called_model_lists = [] called_model_lists = []
async def fake_perform_health_check(model_list, details): async def fake_perform_health_check(
model_list, details, max_concurrency=None
):
called_model_lists.append(copy.deepcopy(model_list)) called_model_lists.append(copy.deepcopy(model_list))
return (["healthy"], []) return (["healthy"], [])