From 4652c73259e2e3b1d3cadea360f2bc5f4b27d1eb Mon Sep 17 00:00:00 2001 From: Sean Marsh Glover Date: Tue, 24 Feb 2026 11:16:59 -0500 Subject: [PATCH] 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 Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- deploy/charts/litellm-helm/README.md | 4 + .../templates/configmap-litellm.yaml | 2 +- .../litellm-helm/templates/deployment.yaml | 25 +- .../litellm-helm/tests/deployment_tests.yaml | 148 ++++++++++- deploy/charts/litellm-helm/values.yaml | 25 ++ litellm/proxy/_types.py | 7 + litellm/proxy/health_check.py | 194 ++++++++++++-- .../shared_health_check_manager.py | 96 +++---- .../health_endpoints/_health_endpoints.py | 114 ++++---- litellm/proxy/proxy_server.py | 247 ++++++++++++++---- .../litellm_utils_tests/test_health_check.py | 119 ++++++++- tests/proxy_unit_tests/test_proxy_server.py | 8 +- 12 files changed, 809 insertions(+), 180 deletions(-) diff --git a/deploy/charts/litellm-helm/README.md b/deploy/charts/litellm-helm/README.md index 2fa856843f..74e70f4aeb 100644 --- a/deploy/charts/litellm-helm/README.md +++ b/deploy/charts/litellm-helm/README.md @@ -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` | | `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` | +| `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`) | `""` | | `ingress.labels` | Additional labels for the Ingress resource | `{}` | | `ingress.*` | See [values.yaml](./values.yaml) for example settings | N/A | diff --git a/deploy/charts/litellm-helm/templates/configmap-litellm.yaml b/deploy/charts/litellm-helm/templates/configmap-litellm.yaml index cf35917da0..acbe4e3a4b 100644 --- a/deploy/charts/litellm-helm/templates/configmap-litellm.yaml +++ b/deploy/charts/litellm-helm/templates/configmap-litellm.yaml @@ -6,4 +6,4 @@ metadata: data: config.yaml: | {{ .Values.proxy_config | toYaml | indent 6 }} -{{- end }} \ No newline at end of file +{{- end }} diff --git a/deploy/charts/litellm-helm/templates/deployment.yaml b/deploy/charts/litellm-helm/templates/deployment.yaml index 4ac5582d06..df483ab927 100644 --- a/deploy/charts/litellm-helm/templates/deployment.yaml +++ b/deploy/charts/litellm-helm/templates/deployment.yaml @@ -158,18 +158,31 @@ spec: {{- end }} livenessProbe: httpGet: - path: /health/liveliness + path: {{ .Values.livenessProbe.path | quote }} 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: httpGet: - path: /health/readiness + path: {{ .Values.readinessProbe.path | quote }} 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: httpGet: - path: /health/readiness + path: {{ .Values.startupProbe.path | quote }} port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }} - failureThreshold: 30 - periodSeconds: 10 + initialDelaySeconds: {{ .Values.startupProbe.initialDelaySeconds }} + periodSeconds: {{ .Values.startupProbe.periodSeconds }} + timeoutSeconds: {{ .Values.startupProbe.timeoutSeconds }} + successThreshold: {{ .Values.startupProbe.successThreshold }} + failureThreshold: {{ .Values.startupProbe.failureThreshold }} resources: {{- toYaml .Values.resources | nindent 12 }} volumeMounts: @@ -235,4 +248,4 @@ spec: {{- if .Values.topologySpreadConstraints }} topologySpreadConstraints: {{- toYaml .Values.topologySpreadConstraints | nindent 8 }} - {{- end }} \ No newline at end of file + {{- end }} diff --git a/deploy/charts/litellm-helm/tests/deployment_tests.yaml b/deploy/charts/litellm-helm/tests/deployment_tests.yaml index f1229e1023..2e9c48043d 100644 --- a/deploy/charts/litellm-helm/tests/deployment_tests.yaml +++ b/deploy/charts/litellm-helm/tests/deployment_tests.yaml @@ -159,4 +159,150 @@ tests: value: -c - equal: path: spec.template.spec.containers[0].lifecycle.preStop.exec.command[2] - value: echo "Container stopping" \ No newline at end of file + 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: {} diff --git a/deploy/charts/litellm-helm/values.yaml b/deploy/charts/litellm-helm/values.yaml index cea25974bb..d62f5b29c2 100644 --- a/deploy/charts/litellm-helm/values.yaml +++ b/deploy/charts/litellm-helm/values.yaml @@ -84,6 +84,31 @@ service: separateHealthApp: false 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: enabled: false className: "nginx" diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 95739834a9..f354e28acd 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2079,6 +2079,13 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): health_check_interval: int = Field( 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( None, description="List of alerting integrations. Today, just slack - `alerting: ['slack']`", diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 427a16a980..4777f64405 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -3,6 +3,9 @@ import asyncio import logging import random +import sys +import threading +import time from typing import List, Optional import litellm @@ -23,6 +26,29 @@ ILLEGAL_DISPLAY_PARAMS = [ 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(): """ Get a random message from the LLM. @@ -67,26 +93,29 @@ async def run_with_timeout(task, timeout): try: return await asyncio.wait_for(task, timeout) except asyncio.TimeoutError: - task.cancel() - # Only cancel child tasks of the current task - 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 + # `asyncio.wait_for()` already cancels only the awaited task on timeout. + # Do not cancel unrelated sibling health check tasks. 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. + + max_concurrency: Optional limit on concurrent health check requests. """ - tasks = [] - for model in model_list: + 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") + + async def _run_model_health_check(model: dict): litellm_params = model["litellm_params"] model_info = model.get("model_info", {}) 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 - task = run_with_timeout( + return await run_with_timeout( litellm.ahealth_check( - model["litellm_params"], + litellm_params, mode=mode, prompt=DEFAULT_HEALTH_CHECK_PROMPT, input=["test from litellm"], @@ -105,9 +134,73 @@ async def _perform_health_check(model_list: list, details: Optional[bool] = True 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 = [] unhealthy_endpoints = [] @@ -190,6 +283,8 @@ async def perform_health_check( model: Optional[str] = None, cli_model: Optional[str] = None, details: Optional[bool] = True, + max_concurrency: Optional[int] = None, + instrumentation_context: Optional[dict] = None, ): """ Perform a health check on the system. @@ -197,14 +292,28 @@ async def perform_health_check( Returns: (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 cli_model: model_list = [ {"model_name": cli_model, "litellm_params": {"model": cli_model}} ] else: + if instrumentation_enabled: + logger.debug( + "health_check_cycle_skipped source=%s cycle_id=%s reason=no_models", + source, + cycle_id, + ) return [], [] + cycle_start_time = time.monotonic() + requested_model_count = len(model_list) + if model is not None: _new_model_list = [ 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] model_list = _new_model_list + post_filter_model_count = len(model_list) model_list = filter_deployments_by_id( model_list=model_list ) # filter duplicate deployments (e.g. when model alias'es are used) - healthy_endpoints, unhealthy_endpoints = await _perform_health_check( - model_list, details - ) + deduped_model_count = len(model_list) + + 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 diff --git a/litellm/proxy/health_check_utils/shared_health_check_manager.py b/litellm/proxy/health_check_utils/shared_health_check_manager.py index d0c99d84e9..ae18a42c02 100644 --- a/litellm/proxy/health_check_utils/shared_health_check_manager.py +++ b/litellm/proxy/health_check_utils/shared_health_check_manager.py @@ -16,7 +16,7 @@ from litellm.proxy.health_check import perform_health_check class SharedHealthCheckManager: """ Manager for coordinating health checks across multiple pods using Redis. - + This class implements a shared health check state mechanism that: - Prevents duplicate health checks across pods - Caches health check results with configurable TTL @@ -58,7 +58,7 @@ class SharedHealthCheckManager: async def acquire_health_check_lock(self) -> bool: """ Attempt to acquire the global health check lock. - + Returns: bool: True if lock was acquired, False otherwise """ @@ -74,7 +74,7 @@ class SharedHealthCheckManager: nx=True, # Only set if key doesn't exist ttl=self.lock_ttl, ) - + if acquired: verbose_proxy_logger.info( "Pod %s acquired health check lock", self.pod_id @@ -83,12 +83,10 @@ class SharedHealthCheckManager: verbose_proxy_logger.debug( "Pod %s failed to acquire health check lock", self.pod_id ) - + return acquired except Exception as e: - verbose_proxy_logger.error( - "Error acquiring health check lock: %s", str(e) - ) + verbose_proxy_logger.error("Error acquiring health check lock: %s", str(e)) return False async def release_health_check_lock(self) -> None: @@ -106,14 +104,12 @@ class SharedHealthCheckManager: "Pod %s released health check lock", self.pod_id ) except Exception as e: - verbose_proxy_logger.error( - "Error releasing health check lock: %s", str(e) - ) + verbose_proxy_logger.error("Error releasing health check lock: %s", str(e)) async def get_cached_health_check_results(self) -> Optional[Dict[str, Any]]: """ Get cached health check results from Redis. - + Returns: Optional[Dict]: Cached health check results or None if not found/expired """ @@ -123,7 +119,7 @@ class SharedHealthCheckManager: try: cache_key = self.get_health_check_cache_key() cached_data = await self.redis_cache.async_get_cache(cache_key) - + if cached_data is None: return None @@ -136,7 +132,7 @@ class SharedHealthCheckManager: # Check if the cache is still valid cache_timestamp = cached_results.get("timestamp", 0) current_time = time.time() - + if current_time - cache_timestamp > self.health_check_ttl: verbose_proxy_logger.debug("Cached health check results expired") return None @@ -151,13 +147,13 @@ class SharedHealthCheckManager: return None async def cache_health_check_results( - self, - healthy_endpoints: List[Dict[str, Any]], - unhealthy_endpoints: List[Dict[str, Any]] + self, + healthy_endpoints: List[Dict[str, Any]], + unhealthy_endpoints: List[Dict[str, Any]], ) -> None: """ Cache health check results in Redis. - + Args: healthy_endpoints: List of healthy endpoints unhealthy_endpoints: List of unhealthy endpoints @@ -181,7 +177,7 @@ class SharedHealthCheckManager: safe_dumps(cache_data), ttl=self.health_check_ttl, ) - + verbose_proxy_logger.info( "Cached health check results for %d healthy and %d unhealthy endpoints", len(healthy_endpoints), @@ -189,29 +185,29 @@ class SharedHealthCheckManager: ) except Exception as e: - verbose_proxy_logger.error( - "Error caching health check results: %s", str(e) - ) + verbose_proxy_logger.error("Error caching health check results: %s", str(e)) async def perform_shared_health_check( - self, - model_list: List[Dict[str, Any]], - details: bool = True + self, + model_list: List[Dict[str, Any]], + details: bool = True, + max_concurrency: Optional[int] = None, ) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]: """ Perform health check with shared state coordination. - + This method: 1. First checks if there are recent cached results 2. If no recent cache, tries to acquire lock to run health check 3. If lock acquired, runs health check and caches results 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 - + Args: model_list: List of models to check details: Whether to include detailed information - + max_concurrency: Optional limit on concurrent health check requests + Returns: Tuple of (healthy_endpoints, unhealthy_endpoints) """ @@ -225,27 +221,29 @@ class SharedHealthCheckManager: # No recent cache, try to acquire lock lock_acquired = await self.acquire_health_check_lock() - + if lock_acquired: try: # We have the lock, run health check verbose_proxy_logger.info( - "Pod %s running health check for %d models", - self.pod_id, - len(model_list) + "Pod %s running health check for %d models", + self.pod_id, + len(model_list), ) - + 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 await self.cache_health_check_results( healthy_endpoints, unhealthy_endpoints ) - + return healthy_endpoints, unhealthy_endpoints - + finally: # Always release the lock await self.release_health_check_lock() @@ -254,10 +252,10 @@ class SharedHealthCheckManager: verbose_proxy_logger.debug( "Pod %s waiting for other pod to complete health check", self.pod_id ) - + # Wait a bit for the other pod to complete await asyncio.sleep(2) - + # Try to get cached results again cached_results = await self.get_cached_health_check_results() if cached_results is not None: @@ -265,19 +263,23 @@ class SharedHealthCheckManager: cached_results.get("healthy_endpoints", []), cached_results.get("unhealthy_endpoints", []), ) - + # Still no cache, fall back to local health check verbose_proxy_logger.warning( - "Pod %s falling back to local health check (no cache available)", - self.pod_id + "Pod %s falling back to local health check (no cache available)", + 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: """ Check if a health check is currently in progress by another pod. - + Returns: 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]: """ Get the current status of health check coordination. - + Returns: Dict containing status information """ @@ -320,7 +322,9 @@ class SharedHealthCheckManager: cached_results = await self.get_cached_health_check_results() status["cache_available"] = cached_results is not None 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") except Exception as e: diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index da90696ec2..3570b2dd6a 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -110,26 +110,31 @@ def _resolve_os_environ_variables(params: dict) -> dict: def get_callback_identifier(callback): """ Get the callback identifier string, handling both strings and objects. - + This function extracts a string identifier from a callback, which can be: - A string (returned as-is) - An object with a callback_name attribute - An object registered in CustomLoggerRegistry - Falls back to callback_name() helper function - + Args: callback: The callback to identify (can be str or object) - + Returns: str: The callback identifier string """ if isinstance(callback, str): return callback - if hasattr(callback, 'callback_name') and callback.callback_name: + if hasattr(callback, "callback_name") and callback.callback_name: return callback.callback_name - if hasattr(callback, '__class__'): - callback_strs = CustomLoggerRegistry.get_all_callback_strs_from_class_type(callback.__class__) - if hasattr(callback, 'callback_name') and callback.callback_name in callback_strs: + if hasattr(callback, "__class__"): + callback_strs = CustomLoggerRegistry.get_all_callback_strs_from_class_type( + callback.__class__ + ) + if ( + hasattr(callback, "callback_name") + and callback.callback_name in callback_strs + ): return callback.callback_name if callback_strs: return callback_strs[0] @@ -151,7 +156,7 @@ services = Union[ "datadog_llm_observability", "generic_api", "arize", - "sqs" + "sqs", ], str, ] @@ -224,7 +229,7 @@ async def health_services_endpoint( # noqa: PLR0915 "datadog_llm_observability", "generic_api", "arize", - "sqs" + "sqs", ]: raise HTTPException( status_code=400, @@ -238,14 +243,14 @@ async def health_services_endpoint( # noqa: PLR0915 service_in_success_callbacks = True else: 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 break cb_id = get_callback_identifier(cb) if cb_id == service: service_in_success_callbacks = True break - + if ( service == "openmeter" or service == "braintrust" @@ -320,6 +325,7 @@ async def health_services_endpoint( # noqa: PLR0915 ) elif service == "sqs": from litellm.integrations.sqs import SQSLogger + sqs_logger = SQSLogger() response = await sqs_logger.async_health_check() return { @@ -518,12 +524,12 @@ async def _save_health_check_to_db( def _build_model_param_to_info_mapping(model_list: list) -> dict: """ 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. - + Args: model_list: List of model configurations - + Returns: 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") litellm_params = model.get("litellm_params", {}) model_param = litellm_params.get("model") - + if model_param and model_name: if model_param not in model_param_to_info: model_param_to_info[model_param] = [] - model_param_to_info[model_param].append({ - "model_name": model_name, - "model_id": model_id, - }) + model_param_to_info[model_param].append( + { + "model_name": model_name, + "model_id": model_id, + } + ) return model_param_to_info @@ -552,19 +560,19 @@ def _aggregate_health_check_results( ) -> dict: """ Aggregate health check results per unique model. - + Uses (model_id, model_name) as key, or (None, model_name) if model_id is None. - + Args: model_param_to_info: Mapping from model parameter to model info healthy_endpoints: List of healthy endpoint results unhealthy_endpoints: List of unhealthy endpoint results - + Returns: Dictionary mapping (model_id, model_name) to aggregated health check results """ model_results = {} - + # Process healthy endpoints for endpoint in healthy_endpoints: model_param = endpoint.get("model") @@ -580,7 +588,7 @@ def _aggregate_health_check_results( "error_message": None, } model_results[key]["healthy_count"] += 1 - + # Process unhealthy endpoints for endpoint in unhealthy_endpoints: model_param = endpoint.get("model") @@ -600,7 +608,7 @@ def _aggregate_health_check_results( # Use the first error message encountered if not model_results[key]["error_message"] and error_message: model_results[key]["error_message"] = str(error_message)[:500] - + 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. - + 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. - + - Stable systems: ~1 write/hour per model (instead of 12 writes/hour with 5-min intervals) - Status changes: Immediate write (no delay) - Result: ~92% reduction in DB writes for stable systems, while maintaining real-time updates on changes - + Args: prisma_client: Database client 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(): new_status = "healthy" if result["healthy_count"] > 0 else "unhealthy" - + # Check if we should save this result should_save = True 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) if last_check.checked_at: from datetime import datetime, timezone + time_since_last_check = ( datetime.now(timezone.utc) - last_check.checked_at ).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 if time_since_last_check < 3600: # 1 hour threshold should_save = False - + if should_save: asyncio.create_task( 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. - + 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). - + 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. """ if prisma_client is None: return - + try: # Step 1: Build mapping from model parameter to model info model_param_to_info = _build_model_param_to_info_mapping(model_list) - + # Step 2: Aggregate health check results per unique model model_results = _aggregate_health_check_results( model_param_to_info, healthy_endpoints, unhealthy_endpoints, ) - + # 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_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 if key not in latest_checks_map: latest_checks_map[key] = check - + # Step 4: Save aggregated results, but only if status changed await _save_health_check_results_if_changed( prisma_client, @@ -729,10 +738,15 @@ async def _perform_health_check_and_save( start_time, user_id, model_id=None, + max_concurrency=None, ): """Helper function to perform health check and save results to database""" 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) @@ -789,6 +803,7 @@ async def health_endpoint( import time from litellm.proxy.proxy_server import ( + health_check_concurrency, health_check_details, health_check_results, llm_model_list, @@ -841,6 +856,7 @@ async def health_endpoint( start_time=start_time, user_id=user_api_key_dict.user_id, model_id=None, # CLI model doesn't have model_id + max_concurrency=health_check_concurrency, ) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -864,6 +880,7 @@ async def health_endpoint( start_time=start_time, user_id=user_api_key_dict.user_id, model_id=model_id, + max_concurrency=health_check_concurrency, ) except Exception as e: verbose_proxy_logger.error( @@ -1420,11 +1437,11 @@ async def test_model_connection( status_code=500, detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - + # Get model name from litellm_params request_litellm_params = litellm_params or {} model_name = request_litellm_params.get("model") - + # Look up model configuration from router if model name is provided # This gets the litellm_params from proxy config (with resolved env vars) config_litellm_params: dict = {} @@ -1432,34 +1449,39 @@ async def test_model_connection( try: # First try to find by proxy model_name (e.g., "gpt-4o") 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 deployments or len(deployments) == 0: all_deployments = llm_router.get_model_list(model_name=None) if 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] break - + if deployments and len(deployments) > 0: # Use the first deployment's litellm_params as base 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: verbose_proxy_logger.debug( f"Could not find model {model_name} in router: {e}. " "Proceeding with request params only." ) - + # Merge: config params (from proxy config) as base, request params override # This allows users to override specific params while using config for credentials merged_litellm_params = {**config_litellm_params, **request_litellm_params} - + # Resolve os.environ/ environment variables in any remaining request params # This handles cases where user explicitly passes os.environ/ values to override config litellm_params = _resolve_os_environ_variables(merged_litellm_params) - + ## Auth check await ModelManagementAuthChecks.can_user_make_model_call( model_params=Deployment( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 35a75d4caf..ddab560aaa 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -9,6 +9,7 @@ import secrets import shutil import subprocess import sys +import threading import time import traceback import warnings @@ -658,7 +659,7 @@ _description = ( 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 master_key = None @@ -672,6 +673,7 @@ def cleanup_router_config_variables(): use_background_health_checks = None use_shared_health_check = None health_check_interval = None + health_check_concurrency = 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") _config = proxy_config.get_config_state() _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( llm_router=llm_router, litellm_settings=_litellm_settings, @@ -1468,7 +1472,9 @@ redis_usage_cache: Optional[ RedisCache ] = None # redis cache used for tracking spend, tpm/rpm limits 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 user_custom_auth = None user_custom_key_generate = None @@ -1478,8 +1484,11 @@ use_background_health_checks = None use_shared_health_check = None use_queue = False health_check_interval = None +health_check_concurrency = None health_check_details = None 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 = [] litellm_proxy_budget_name = "litellm-proxy-budget" 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(): """ 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. 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 ( health_check_interval is None @@ -1943,6 +1978,24 @@ async def _run_background_health_check(): ): 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 shared_health_manager = 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") 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 _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 _llm_model_list = [ @@ -1967,6 +2025,52 @@ async def _run_background_health_check(): for m in _llm_model_list 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 # 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, unhealthy_endpoints, ) = 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: verbose_proxy_logger.error( "Error in shared health check, falling back to direct health check: %s", str(e), ) - healthy_endpoints, unhealthy_endpoints = await perform_health_check( - model_list=_llm_model_list, details=health_check_details + healthy_endpoints, unhealthy_endpoints = ( + await _run_direct_health_check_with_instrumentation() ) else: - healthy_endpoints, unhealthy_endpoints = await perform_health_check( - model_list=_llm_model_list, details=health_check_details + healthy_endpoints, unhealthy_endpoints = ( + await _run_direct_health_check_with_instrumentation() ) # 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["healthy_count"] = len(healthy_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) if prisma_client is not None: @@ -2480,7 +2605,7 @@ class ProxyConfig: """ 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) @@ -2905,7 +3030,18 @@ class ProxyConfig: health_check_interval = general_settings.get( "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) + 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_role_permissions = general_settings.get("role_permissions", None) @@ -2999,7 +3135,7 @@ class ProxyConfig: for k, v in router_settings.items(): if k in available_args: router_params[k] = v - elif k == "health_check_interval": + elif k in {"health_check_interval", "health_check_concurrency"}: raise ValueError( 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"): - await self._init_semantic_filter_settings_in_db( - prisma_client=prisma_client - ) + await self._init_semantic_filter_settings_in_db(prisma_client=prisma_client) 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""" 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 - if not mcp_semantic_filter_config or not mcp_semantic_filter_config.get("enabled", False): - verbose_proxy_logger.debug("Semantic tool filter not configured or not enabled, skipping initialization") + if not mcp_semantic_filter_config or not mcp_semantic_filter_config.get( + "enabled", False + ): + verbose_proxy_logger.debug( + "Semantic tool filter not configured or not enabled, " + "skipping initialization" + ) return - + verbose_proxy_logger.debug( f"Initializing semantic tool filter: llm_router={llm_router is not None}, " f"config={mcp_semantic_filter_config}" ) - hook = await SemanticToolFilterHook.initialize_from_config( config=mcp_semantic_filter_config, llm_router=llm_router, ) - + if hook: verbose_proxy_logger.debug("Semantic tool filter hook registered") litellm.logging_callback_manager.add_litellm_callback(hook) else: # 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 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 if router_models_count < models_needed_for_page: 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: @@ -8742,21 +8885,21 @@ async def _apply_search_filter_to_models( def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]: """ Normalize a datetime value to a timezone-aware UTC datetime for sorting. - + This function handles: - None values: returns None - 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 - + Args: dt: Datetime value (None, str, or datetime object) - + Returns: UTC-aware datetime object, or None if input is None or cannot be parsed """ if dt is None: return None - + if isinstance(dt, str): try: # Handle ISO format strings, including 'Z' suffix @@ -8770,14 +8913,14 @@ def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]: return parsed_dt except (ValueError, AttributeError): return None - + if isinstance(dt, datetime): # If naive, assume UTC and make it aware if dt.tzinfo is None: return dt.replace(tzinfo=timezone.utc) # If aware, convert to UTC return dt.astimezone(timezone.utc) - + return None @@ -8797,46 +8940,60 @@ def _sort_models( Returns: 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 reverse = sort_order.lower() == "desc" def get_sort_key(model: Dict[str, Any]) -> Any: model_info = model.get("model_info", {}) - + if sort_by == "model_name": return model.get("model_name", "").lower() - + elif sort_by == "created_at": created_at = model_info.get("created_at") normalized_dt = _normalize_datetime_for_sorting(created_at) if normalized_dt is None: # 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 - + elif sort_by == "updated_at": updated_at = model_info.get("updated_at") normalized_dt = _normalize_datetime_for_sorting(updated_at) 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 - + elif sort_by == "costs": input_cost = model_info.get("input_cost_per_token", 0) or 0 output_cost = model_info.get("output_cost_per_token", 0) or 0 total_cost = input_cost + output_cost # Put 0 or None costs at the end for asc, at the start for desc 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 - + elif sort_by == "status": # False (config) comes before True (db) for asc db_model = model_info.get("db_model", False) return db_model - + return None try: @@ -9032,9 +9189,7 @@ async def _find_model_by_id( ) if db_model: # Convert database model to router format - decrypted_models = proxy_config.decrypt_model_list_from_db( - [db_model] - ) + decrypted_models = proxy_config.decrypt_model_list_from_db([db_model]) if decrypted_models: found_model = decrypted_models[0] except Exception as e: @@ -9208,13 +9363,13 @@ async def model_info_v2( ) verbose_proxy_logger.debug("all_models: %s", all_models) - + # Append A2A agents to models list all_models = await append_agents_to_model_info( models=all_models, user_api_key_dict=user_api_key_dict, ) - + # Update total count to include agents search_total_count = len(all_models) @@ -10057,7 +10212,7 @@ async def 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 ) - + # Append A2A agents to model groups model_groups = await append_agents_to_model_group( model_groups=model_groups, diff --git a/tests/litellm_utils_tests/test_health_check.py b/tests/litellm_utils_tests/test_health_check.py index 19882bbe4b..7ea4574bab 100644 --- a/tests/litellm_utils_tests/test_health_check.py +++ b/tests/litellm_utils_tests/test_health_check.py @@ -92,7 +92,7 @@ async def test_azure_img_gen_health_check(): litellm._turn_on_debug() max_retries = 3 retry_delay = 1 # Start with 1 second delay - + for attempt in range(max_retries): response = await litellm.ahealth_check( model_params={ @@ -103,11 +103,11 @@ async def test_azure_img_gen_health_check(): mode="image_generation", prompt="cute baby sea otter", ) - + # Check if response is successful (no error) if isinstance(response, dict) and "error" not in response: return response - + # Check if error is a transient Azure internal server error error_str = str(response.get("error", "")).lower() is_transient_error = ( @@ -116,16 +116,18 @@ async def test_azure_img_gen_health_check(): or "internalfailure" in error_str or "internal failure" in error_str ) - + # If it's the last attempt or not a transient error, fail the test 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 - + # Wait before retrying with exponential backoff await asyncio.sleep(retry_delay) retry_delay *= 2 # Exponential backoff - + # Should not reach here, but just in case 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" +@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 async def test_ahealth_check_ocr(): 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(): """ 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 failed with "LLM Provider NOT provided" error for OpenAI-compatible self-hosted providers, even when a provider was selected in the dropdown. - + 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. """ from unittest.mock import MagicMock - + # Mock the completion call to avoid making real API calls mock_response = MagicMock() mock_response._hidden_params = {"headers": {"x-ratelimit-remaining-tokens": "1000"}} - + with patch("litellm.acompletion", return_value=mock_response): # Test with a custom model name that wouldn't be recognized without custom_llm_provider response = await litellm.ahealth_check( @@ -668,7 +763,7 @@ async def test_health_check_with_custom_llm_provider(): }, mode="chat", ) - + # Should succeed without "LLM Provider NOT provided" error assert "error" not in response assert isinstance(response, dict) diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 14b49901c9..0cbba7b5cc 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -2330,7 +2330,9 @@ async def test_run_background_health_check_reflects_llm_model_list(monkeypatch): test_model_list_2 = [{"model_name": "model-b"}] 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)) return (["healthy"], ["unhealthy"]) @@ -2378,7 +2380,9 @@ async def test_background_health_check_skip_disabled_models(monkeypatch): ] 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)) return (["healthy"], [])