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:
parent
1fa0aad3f2
commit
4652c73259
@ -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 |
|
||||||
|
|||||||
@ -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 }}
|
||||||
|
|||||||
@ -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 }}
|
||||||
|
|||||||
@ -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: {}
|
||||||
|
|||||||
@ -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"
|
||||||
|
|||||||
@ -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']`",
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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(
|
||||||
|
|||||||
@ -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,
|
||||||
|
|||||||
@ -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)
|
||||||
|
|||||||
@ -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"], [])
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user