mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge pull request #40886 from BerriAI/litellm_backport_redis_chaos_rc_1_101_0
fix(redis): backport Redis chaos fixes and load gate to rc/1.101.0
This commit is contained in:
commit
18243cd7af
50 changed files with 2865 additions and 530 deletions
103
.github/workflows/test-e2e-redis-chaos.yml
vendored
Normal file
103
.github/workflows/test-e2e-redis-chaos.yml
vendored
Normal file
|
|
@ -0,0 +1,103 @@
|
|||
name: "Redis Chaos E2E"
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
workflow_call:
|
||||
inputs:
|
||||
ref:
|
||||
description: "Commit SHA or ref to test. Defaults to the ref the workflow was triggered on"
|
||||
required: false
|
||||
type: string
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
redis-chaos-e2e:
|
||||
runs-on: ubuntu-latest-16-cores
|
||||
timeout-minutes: 30
|
||||
services:
|
||||
postgres:
|
||||
image: postgres:16.6@sha256:557fea37a744d5f4c8faab304b0a90858b53ab119735a88c131fd19dab802f36
|
||||
env:
|
||||
POSTGRES_USER: llmproxy
|
||||
POSTGRES_PASSWORD: dbpassword9090
|
||||
POSTGRES_DB: litellm
|
||||
ports:
|
||||
- 5432:5432
|
||||
options: >-
|
||||
--health-cmd "pg_isready -U llmproxy"
|
||||
--health-interval 5s
|
||||
--health-timeout 5s
|
||||
--health-retries 10
|
||||
valkey:
|
||||
image: valkey/valkey:8.1.4@sha256:81db6d39e1bba3b3ff32bd3a1b19a6d69690f94a3954ec131277b9a26b95b3aa
|
||||
ports:
|
||||
- 6379:6379
|
||||
options: >-
|
||||
--health-cmd "valkey-cli ping"
|
||||
--health-interval 5s
|
||||
--health-timeout 5s
|
||||
--health-retries 10
|
||||
env:
|
||||
DATABASE_URL: postgresql://llmproxy:dbpassword9090@localhost:5432/litellm
|
||||
LITELLM_MASTER_KEY: sk-redis-chaos-e2e
|
||||
LITELLM_LOG: WARNING
|
||||
JSON_LOGS: "true"
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
ref: ${{ inputs.ref || github.sha }}
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Cache the Rust build
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --group e2e-dev --extra proxy
|
||||
|
||||
- name: Cache Prisma binaries
|
||||
uses: ./.github/actions/cache-prisma-binaries
|
||||
|
||||
- name: Generate Prisma client
|
||||
run: |
|
||||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Start a multi-worker proxy on the chaos config
|
||||
run: |
|
||||
nohup uv run --no-sync litellm --config tests/e2e/gateway/redis_chaos_ci_config.yml --port 4000 --num_workers 4 > proxy.log 2>&1 &
|
||||
echo "E2E_PROXY_PID=$!" >> "$GITHUB_ENV"
|
||||
echo "E2E_PROXY_LOG=$(pwd)/proxy.log" >> "$GITHUB_ENV"
|
||||
for _ in $(seq 1 90); do
|
||||
if curl -fs http://localhost:4000/health/liveliness > /dev/null; then
|
||||
exit 0
|
||||
fi
|
||||
sleep 2
|
||||
done
|
||||
echo "proxy never became live"
|
||||
tail -n 100 proxy.log
|
||||
exit 1
|
||||
|
||||
- name: Run the Redis chaos load test
|
||||
env:
|
||||
E2E_REDIS_CHAOS: "1"
|
||||
LITELLM_PROXY_URL: http://localhost:4000
|
||||
REDIS_HOST: 127.0.0.1
|
||||
REDIS_PORT: "6379"
|
||||
run: |
|
||||
uv run --no-sync pytest tests/e2e/load/test_redis_chaos_e2e.py -v --tb=short -rA -s
|
||||
|
||||
- name: Show proxy log on failure
|
||||
if: failure()
|
||||
run: tail -n 300 proxy.log
|
||||
|
|
@ -10,6 +10,7 @@
|
|||
import ast
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
|
|
@ -32,7 +33,7 @@ from .dual_cache import DualCache # noqa: F401
|
|||
from .gcs_cache import GCSCache
|
||||
from .in_memory_cache import InMemoryCache
|
||||
from .qdrant_semantic_cache import QdrantSemanticCache
|
||||
from .redis_cache import RedisCache
|
||||
from .redis_cache import RedisCache, log_redis_failure
|
||||
from .redis_cluster_cache import RedisClusterCache
|
||||
from .redis_semantic_cache import RedisSemanticCache
|
||||
from .s3_cache import S3Cache
|
||||
|
|
@ -678,7 +679,7 @@ class Cache:
|
|||
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
|
||||
self.cache.set_cache(cache_key, cached_data, **kwargs)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("LiteLLM Cache: Excepton add_cache: %s", e)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Cache: exception in add_cache", e)
|
||||
|
||||
async def async_add_cache(self, result, dynamic_cache_object: BaseCache | None = None, **kwargs):
|
||||
"""
|
||||
|
|
@ -697,7 +698,7 @@ class Cache:
|
|||
else:
|
||||
await self.cache.async_set_cache(cache_key, cached_data, **kwargs)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("LiteLLM Cache: Excepton add_cache: %s", e)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Cache: exception in add_cache", e)
|
||||
|
||||
def _convert_to_cached_embedding(
|
||||
self,
|
||||
|
|
@ -876,7 +877,7 @@ class Cache:
|
|||
else:
|
||||
await self.cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("LiteLLM Cache: Excepton add_cache: %s", e)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Cache: exception in add_cache", e)
|
||||
|
||||
def should_use_cache(self, **kwargs):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -8,8 +8,8 @@ Has 4 primary methods:
|
|||
- async_get_cache
|
||||
"""
|
||||
|
||||
import logging
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Sequence
|
||||
from threading import Lock
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
|
@ -23,7 +23,7 @@ from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE
|
|||
|
||||
from .base_cache import BaseCache
|
||||
from .in_memory_cache import InMemoryCache
|
||||
from .redis_cache import RedisCache
|
||||
from .redis_cache import RedisCache, RedisCircuitBreakerOpenError, log_redis_failure
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
|
@ -177,8 +177,10 @@ class DualCache(BaseCache):
|
|||
|
||||
print_verbose(f"get cache: cache result: {result}")
|
||||
return result
|
||||
except Exception:
|
||||
verbose_logger.error(traceback.format_exc())
|
||||
except Exception as e:
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in get_cache", e, with_traceback=True
|
||||
)
|
||||
|
||||
def batch_get_cache(
|
||||
self,
|
||||
|
|
@ -204,9 +206,12 @@ class DualCache(BaseCache):
|
|||
redis_result: Final = self.redis_cache.batch_get_cache(
|
||||
key_list=sublist_keys, parent_otel_span=parent_otel_span
|
||||
)
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
# Do not throttle subsequent callers if the Redis read fails.
|
||||
self._rollback_redis_batch_key_reservations(previous_access_times)
|
||||
if isinstance(e, RedisCircuitBreakerOpenError):
|
||||
verbose_logger.debug("LiteLLM Cache: batch_get_cache served from memory only: %s", e)
|
||||
return result
|
||||
raise
|
||||
|
||||
if self.in_memory_cache is not None:
|
||||
|
|
@ -217,8 +222,10 @@ class DualCache(BaseCache):
|
|||
return list( # mutable-ok: public list contract
|
||||
redis_result.get(key) if value is None else value for key, value in zip(keys, result)
|
||||
)
|
||||
except Exception:
|
||||
verbose_logger.error(traceback.format_exc())
|
||||
except Exception as e:
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in batch_get_cache", e, with_traceback=True
|
||||
)
|
||||
|
||||
async def async_get_cache(
|
||||
self,
|
||||
|
|
@ -250,8 +257,10 @@ class DualCache(BaseCache):
|
|||
|
||||
print_verbose(f"get cache: cache result: {result}")
|
||||
return result
|
||||
except Exception:
|
||||
verbose_logger.error(traceback.format_exc())
|
||||
except Exception as e:
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in async_get_cache", e, with_traceback=True
|
||||
)
|
||||
|
||||
def _reserve_redis_batch_keys(
|
||||
self,
|
||||
|
|
@ -319,9 +328,12 @@ class DualCache(BaseCache):
|
|||
redis_result: Final = await self.redis_cache.async_batch_get_cache(
|
||||
sublist_keys, parent_otel_span=parent_otel_span
|
||||
)
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
# Do not throttle subsequent callers if the Redis read fails.
|
||||
self._rollback_redis_batch_key_reservations(previous_access_times)
|
||||
if isinstance(e, RedisCircuitBreakerOpenError):
|
||||
verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e)
|
||||
return result
|
||||
raise
|
||||
|
||||
# Short-circuit if redis_result is None or contains only None values
|
||||
|
|
@ -339,8 +351,14 @@ class DualCache(BaseCache):
|
|||
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
|
||||
|
||||
return result
|
||||
except Exception:
|
||||
verbose_logger.error(traceback.format_exc())
|
||||
except Exception as e:
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.ERROR,
|
||||
"LiteLLM Cache: exception in async_batch_get_cache",
|
||||
e,
|
||||
with_traceback=True,
|
||||
)
|
||||
|
||||
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs):
|
||||
print_verbose(f"async set cache: cache key: {key}; local_only: {local_only}; value: {value}")
|
||||
|
|
@ -353,7 +371,9 @@ class DualCache(BaseCache):
|
|||
if self.redis_cache is not None and local_only is False:
|
||||
await self.redis_cache.async_set_cache(key, value, **kwargs)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("LiteLLM Cache: Excepton async add_cache: %s", e)
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in async add_cache", e, with_traceback=True
|
||||
)
|
||||
|
||||
# async_batch_set_cache
|
||||
async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs):
|
||||
|
|
@ -372,7 +392,9 @@ class DualCache(BaseCache):
|
|||
cache_list=cache_list, ttl=kwargs.pop("ttl", None), **kwargs
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("LiteLLM Cache: Excepton async add_cache: %s", e)
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in async add_cache", e, with_traceback=True
|
||||
)
|
||||
|
||||
async def async_increment_cache(
|
||||
self,
|
||||
|
|
@ -410,8 +432,10 @@ class DualCache(BaseCache):
|
|||
|
||||
return result
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"Redis async_increment_cache failed, falling back to in-memory result: %s",
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.WARNING,
|
||||
"Redis async_increment_cache failed, falling back to in-memory result",
|
||||
e,
|
||||
)
|
||||
return result
|
||||
|
|
@ -439,8 +463,10 @@ class DualCache(BaseCache):
|
|||
|
||||
return result
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"Redis async_increment_cache_pipeline failed, falling back to in-memory result: %s",
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.WARNING,
|
||||
"Redis async_increment_cache_pipeline failed, falling back to in-memory result",
|
||||
e,
|
||||
)
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -14,9 +14,11 @@ import functools
|
|||
import hashlib
|
||||
import inspect
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
|
||||
|
||||
|
|
@ -175,8 +177,14 @@ class RedisCircuitBreaker:
|
|||
self._timeout_streak_started_at: float | None = None
|
||||
self._opened_at: float | None = None
|
||||
self._state = self.CLOSED
|
||||
self._generation = 0
|
||||
_breaker_metrics().record_state_change(None, self._state)
|
||||
|
||||
@property
|
||||
def generation(self) -> int:
|
||||
"""Counts state transitions, so a call can tell whether the breaker moved while it ran."""
|
||||
return self._generation
|
||||
|
||||
def is_open(self) -> bool:
|
||||
"""Returns True if Redis calls should be skipped."""
|
||||
if not self.enabled:
|
||||
|
|
@ -229,7 +237,7 @@ class RedisCircuitBreaker:
|
|||
self._set_state(self.OPEN)
|
||||
|
||||
def record_success(self) -> None:
|
||||
if not self.enabled:
|
||||
if not self.enabled or self._state == self.OPEN:
|
||||
return
|
||||
if self._state == self.HALF_OPEN:
|
||||
verbose_logger.info("Redis circuit breaker CLOSED — Redis recovered")
|
||||
|
|
@ -245,6 +253,7 @@ class RedisCircuitBreaker:
|
|||
_breaker_metrics().record_transition(state)
|
||||
_breaker_metrics().record_state_change(self._state, state)
|
||||
self._state = state
|
||||
self._generation += 1
|
||||
|
||||
|
||||
_RedisCallResult = TypeVar("_RedisCallResult")
|
||||
|
|
@ -371,21 +380,46 @@ def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseExcep
|
|||
_swallowed_redis_failures.set(_swallowed_redis_failures.get() + 1)
|
||||
|
||||
|
||||
def _enter_circuit_breaker(breaker: RedisCircuitBreaker, name: str) -> int:
|
||||
"""Reject the call if the breaker is open, else return the swallowed-failure count to compare against."""
|
||||
class RedisCircuitBreakerOpenError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def log_redis_failure(
|
||||
logger: logging.Logger, level: int, message: str, exc: BaseException, with_traceback: bool = False
|
||||
) -> None:
|
||||
if isinstance(exc, RedisCircuitBreakerOpenError):
|
||||
logger.debug("%s: %s", message, exc)
|
||||
return
|
||||
logger.log(level, "%s: %s", message, exc, exc_info=exc if with_traceback else None)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BreakerAdmission:
|
||||
swallowed_before: int
|
||||
generation: int
|
||||
|
||||
|
||||
def _enter_circuit_breaker(breaker: RedisCircuitBreaker, name: str) -> _BreakerAdmission:
|
||||
"""Reject the call if the breaker is open, else record what its success may later prove."""
|
||||
if breaker.is_open():
|
||||
raise Exception(f"Redis circuit breaker is open — skipping {name}")
|
||||
return _swallowed_redis_failures.get()
|
||||
raise RedisCircuitBreakerOpenError(f"Redis circuit breaker is open — skipping {name}")
|
||||
return _BreakerAdmission(swallowed_before=_swallowed_redis_failures.get(), generation=breaker.generation)
|
||||
|
||||
|
||||
def _exit_circuit_breaker(breaker: RedisCircuitBreaker, swallowed_before: int) -> None:
|
||||
"""Record success only when nothing failed while the call ran.
|
||||
def _exit_circuit_breaker(breaker: RedisCircuitBreaker, admission: _BreakerAdmission) -> None:
|
||||
"""Record success only when nothing failed while the call ran and the breaker has not moved since.
|
||||
|
||||
Several Redis methods catch their own connection errors and return a default, so a
|
||||
method that returned is not on its own proof of a healthy Redis.
|
||||
method that returned is not on its own proof of a healthy Redis. A success also vouches
|
||||
only for the breaker state that admitted the call: a call admitted before the breaker
|
||||
opened, or a probe admitted before a later failure reopened it, finishes knowing nothing
|
||||
about whether Redis has recovered since, so only the current probe may close the breaker.
|
||||
"""
|
||||
if _swallowed_redis_failures.get() == swallowed_before:
|
||||
breaker.record_success()
|
||||
if _swallowed_redis_failures.get() != admission.swallowed_before:
|
||||
return
|
||||
if breaker.generation != admission.generation:
|
||||
return
|
||||
breaker.record_success()
|
||||
|
||||
|
||||
async def _run_under_circuit_breaker(
|
||||
|
|
@ -398,14 +432,14 @@ async def _run_under_circuit_breaker(
|
|||
Shared by the method decorator and the Lua script executor so both feed the same
|
||||
health signal.
|
||||
"""
|
||||
swallowed_before: Final = _enter_circuit_breaker(breaker, name)
|
||||
admission: Final = _enter_circuit_breaker(breaker, name)
|
||||
try:
|
||||
result: Final = await call()
|
||||
except Exception as e:
|
||||
if _is_redis_health_failure(e):
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(e))
|
||||
raise
|
||||
_exit_circuit_breaker(breaker, swallowed_before)
|
||||
_exit_circuit_breaker(breaker, admission)
|
||||
return result
|
||||
|
||||
|
||||
|
|
@ -415,14 +449,14 @@ def _run_under_circuit_breaker_sync(
|
|||
call: Callable[[], _RedisCallResult],
|
||||
) -> _RedisCallResult:
|
||||
"""Run one blocking Redis call under a circuit breaker, feeding the same health signal as the async path."""
|
||||
swallowed_before: Final = _enter_circuit_breaker(breaker, name)
|
||||
admission: Final = _enter_circuit_breaker(breaker, name)
|
||||
try:
|
||||
result: Final = call()
|
||||
except Exception as e:
|
||||
if _is_redis_health_failure(e):
|
||||
breaker.record_failure()
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(e))
|
||||
raise
|
||||
_exit_circuit_breaker(breaker, swallowed_before)
|
||||
_exit_circuit_breaker(breaker, admission)
|
||||
return result
|
||||
|
||||
|
||||
|
|
@ -1258,6 +1292,7 @@ class RedisCache(BaseCache):
|
|||
except Exception:
|
||||
return ast.literal_eval(decoded)
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def get_cache(self, key, parent_otel_span: Span | None = None, **kwargs):
|
||||
try:
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
|
|
@ -1277,8 +1312,8 @@ class RedisCache(BaseCache):
|
|||
print_verbose(f"Got Redis Cache: key: {key}, cached_response {cached_response}")
|
||||
return self._get_cache_logic(cached_response=cached_response)
|
||||
except Exception as e:
|
||||
# NON blocking - notify users Redis is throwing an exception
|
||||
verbose_logger.error("litellm.caching.caching: get() - Got exception from REDIS: ", e)
|
||||
verbose_logger.error("litellm.caching.caching: get() - Got exception from REDIS: %s", e)
|
||||
_record_swallowed_redis_failure(self._circuit_breaker, e)
|
||||
|
||||
def _run_redis_mget_operation(self, keys: list[str]) -> Sequence[bytes | str | None]:
|
||||
"""
|
||||
|
|
@ -1315,12 +1350,12 @@ class RedisCache(BaseCache):
|
|||
key_value_dict = {}
|
||||
_key_list: Final = [key for key in key_list if key is not None]
|
||||
start_time: Final = time.time()
|
||||
admission: Final = _enter_circuit_breaker(self._circuit_breaker, "batch_get_cache")
|
||||
|
||||
try:
|
||||
swallowed_before: Final = _enter_circuit_breaker(self._circuit_breaker, "batch_get_cache")
|
||||
_keys: Final = [self.check_and_fix_namespace(key=cache_key or "") for cache_key in _key_list]
|
||||
results: Final = self._run_redis_mget_operation(keys=_keys)
|
||||
_exit_circuit_breaker(self._circuit_breaker, swallowed_before)
|
||||
_exit_circuit_breaker(self._circuit_breaker, admission)
|
||||
end_time: Final = time.time()
|
||||
_duration: Final = end_time - start_time
|
||||
self.service_logger_obj.service_success_hook(
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.caching.redis_cache import RedisCache, log_redis_failure
|
||||
from litellm.constants import DEFAULT_CRON_JOB_LOCK_TTL_SECONDS
|
||||
from litellm.proxy.db.db_transaction_queue.base_update_queue import service_logger_obj
|
||||
from litellm.types.services import ServiceTypes
|
||||
|
|
@ -109,7 +110,7 @@ end
|
|||
)
|
||||
return False
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error acquiring Redis lock for %s: %s", cronjob_id, e)
|
||||
log_redis_failure(verbose_proxy_logger, logging.ERROR, f"Error acquiring Redis lock for {cronjob_id}", e)
|
||||
return False
|
||||
|
||||
async def release_lock(
|
||||
|
|
@ -151,7 +152,7 @@ end
|
|||
cronjob_id,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error releasing Redis lock for %s: %s", cronjob_id, e)
|
||||
log_redis_failure(verbose_proxy_logger, logging.ERROR, f"Error releasing Redis lock for {cronjob_id}", e)
|
||||
|
||||
async def _compare_and_delete_lock(self, lock_key: str) -> int:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ the reservation is refunded when the batch reaches a terminal state
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import math
|
||||
import time
|
||||
import uuid
|
||||
|
|
@ -19,6 +20,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias
|
|||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
from litellm.constants import BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, BATCH_ENQUEUED_TOKEN_TTL_SECONDS
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
|
@ -233,8 +235,11 @@ class BatchEnqueuedTokenStore:
|
|||
try:
|
||||
return await self._reserve_via_redis(reserve_script, refund_script, tokens=tokens, scopes=scopes)
|
||||
except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters
|
||||
verbose_proxy_logger.warning(
|
||||
"Redis enqueued-token reserve failed, falling back to in-memory: %s", str(e)
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"Redis enqueued-token reserve failed, falling back to in-memory",
|
||||
e,
|
||||
)
|
||||
return await self._reserve_in_memory(tokens=tokens, scopes=scopes, span=litellm_parent_otel_span)
|
||||
|
||||
|
|
@ -374,8 +379,11 @@ class BatchEnqueuedTokenStore:
|
|||
(serialized, ttl),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record
|
||||
verbose_proxy_logger.warning(
|
||||
"Redis enqueued-token reservation save failed, falling back to in-memory: %s", str(e)
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"Redis enqueued-token reservation save failed, falling back to in-memory",
|
||||
e,
|
||||
)
|
||||
else:
|
||||
return
|
||||
|
|
@ -421,8 +429,11 @@ class BatchEnqueuedTokenStore:
|
|||
await pop_script((self._record_key(batch_id),), (BATCH_ENQUEUED_TOKEN_TTL_SECONDS,))
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record
|
||||
verbose_proxy_logger.warning(
|
||||
"Redis enqueued-token reservation pop failed, falling back to in-memory: %s", str(e)
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"Redis enqueued-token reservation pop failed, falling back to in-memory",
|
||||
e,
|
||||
)
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -14,11 +14,13 @@ Works across multiple proxy instances via DualCache (in-memory + Redis).
|
|||
Follows the same pattern as max_iterations_limiter.py.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
from litellm.exceptions import RateLimitType
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -215,9 +217,11 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
return float(result)
|
||||
return 0.0
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"MaxBudgetPerSessionHandler: Redis GET failed, falling back to in-memory: %s",
|
||||
str(e),
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"MaxBudgetPerSessionHandler: Redis GET failed, falling back to in-memory",
|
||||
e,
|
||||
)
|
||||
|
||||
result = await self.internal_usage_cache.async_get_cache(
|
||||
|
|
@ -239,9 +243,11 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
)
|
||||
return float(result)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"MaxBudgetPerSessionHandler: Redis INCRBYFLOAT failed, falling back to in-memory: %s",
|
||||
str(e),
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"MaxBudgetPerSessionHandler: Redis INCRBYFLOAT failed, falling back to in-memory",
|
||||
e,
|
||||
)
|
||||
|
||||
return await self._in_memory_increment_spend(cache_key, amount)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ This is currently in development and not yet ready for production.
|
|||
|
||||
import asyncio
|
||||
import binascii
|
||||
import logging
|
||||
import os
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence, Set
|
||||
|
|
@ -26,6 +27,7 @@ from typing_extensions import NotRequired, ReadOnly
|
|||
|
||||
from litellm import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
|
|
@ -1223,7 +1225,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
all_cache_values.extend(group_cache_values)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("Redis Lua script failed for hash tag %s: %s", hash_tag, e)
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger, logging.WARNING, f"Redis Lua script failed for hash tag {hash_tag}", e
|
||||
)
|
||||
# Fallback to in-memory cache for this group
|
||||
group_cache_values = await self.in_memory_cache_sliding_window(
|
||||
keys=group_keys,
|
||||
|
|
@ -1470,7 +1474,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
counts = [max(0, int(value)) for value in raw_counts]
|
||||
except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the local mirror, never a 500
|
||||
verbose_proxy_logger.warning("parallel_count_script failed, using local mirror: %s", e)
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger, logging.WARNING, "parallel_count_script failed, using local mirror", e
|
||||
)
|
||||
counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span)
|
||||
else:
|
||||
counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span)
|
||||
|
|
@ -1500,7 +1506,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
],
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to in-memory enforcement, never a 500
|
||||
verbose_proxy_logger.warning("parallel_acquire_script failed, falling back to in-memory gauge: %s", e)
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"parallel_acquire_script failed, falling back to in-memory gauge",
|
||||
e,
|
||||
)
|
||||
async with self._check_and_increment_lock:
|
||||
return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span)
|
||||
if int(raw[0]) == 1:
|
||||
|
|
@ -1626,7 +1637,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
return
|
||||
except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the in-memory release, never a 500
|
||||
verbose_proxy_logger.warning("parallel_release_script failed, falling back to in-memory release: %s", e)
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"parallel_release_script failed, falling back to in-memory release",
|
||||
e,
|
||||
)
|
||||
|
||||
async with self._check_and_increment_lock:
|
||||
for counter_key in counter_keys:
|
||||
|
|
@ -1809,12 +1825,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
# state ambiguous. Refund any prior groups so Redis returns
|
||||
# to its pre-call state, then fall back to in-memory for the
|
||||
# whole call (counters there are independent of Redis).
|
||||
verbose_proxy_logger.error(
|
||||
"atomic_check_and_increment_by_n: Redis Lua execution failed (%s: %s). Refunding %s prior descriptors and falling back to in-memory enforcement — counters will diverge from Redis until window expires (window_size=%ss).",
|
||||
type(e).__name__,
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.ERROR,
|
||||
f"atomic_check_and_increment_by_n: Redis Lua execution failed ({type(e).__name__}). Refunding "
|
||||
f"{len(applied)} prior descriptors and falling back to in-memory enforcement, counters will "
|
||||
f"diverge from Redis until window expires (window_size={self.window_size}s)",
|
||||
e,
|
||||
len(applied),
|
||||
self.window_size,
|
||||
)
|
||||
await self._refund_applied_descriptor_groups(applied)
|
||||
flat_meta: list[AtomicCounterMeta] = [m for _k, _a, group_meta in descriptor_groups for m in group_meta]
|
||||
|
|
@ -1861,8 +1878,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
value=-entry["increment"],
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to refund %s on cross-descriptor rollback: %s", entry["counter_key"], e
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
f"Failed to refund {entry['counter_key']} on cross-descriptor rollback",
|
||||
e,
|
||||
)
|
||||
|
||||
def _build_atomic_response(
|
||||
|
|
@ -3851,7 +3871,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("TTL preservation failed, falling back to regular pipeline: %s", e)
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger, logging.WARNING, "TTL preservation failed, falling back to regular pipeline", e
|
||||
)
|
||||
# Fallback to regular pipeline on error
|
||||
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
|
||||
increment_list=pipeline_operations,
|
||||
|
|
@ -3917,9 +3939,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
continue
|
||||
except Exception as e: # noqa: BLE001 # Redis failures use the plain increment fallback
|
||||
verbose_proxy_logger.warning(
|
||||
"Window-guarded token adjustment failed for %s: %s",
|
||||
operation["key"],
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
f"Window-guarded token adjustment failed for {operation['key']}",
|
||||
e,
|
||||
)
|
||||
if operation["increment_value"] > 0:
|
||||
|
|
|
|||
|
|
@ -10,11 +10,13 @@ this hook manages:
|
|||
Works across multiple proxy instances via DualCache (in-memory + Redis).
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
from litellm.integrations.custom_guardrail import get_session_id_from_request_data
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -96,9 +98,11 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
|
|||
)
|
||||
return routed_model
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"SensitiveDataRoutingHandler: Redis GET failed, falling back to in-memory: %s",
|
||||
str(e),
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"SensitiveDataRoutingHandler: Redis GET failed, falling back to in-memory",
|
||||
e,
|
||||
)
|
||||
|
||||
result = await self.internal_usage_cache.async_get_cache(
|
||||
|
|
@ -142,9 +146,11 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
|
|||
ttl=self.ttl,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"SensitiveDataRoutingHandler: Redis SET failed, falling back to in-memory: %s",
|
||||
str(e),
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"SensitiveDataRoutingHandler: Redis SET failed, falling back to in-memory",
|
||||
e,
|
||||
)
|
||||
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.constants import (
|
||||
CLI_SSO_CLAIM_MAP,
|
||||
CLI_SSO_CLAIM_MAX_SCALAR_LENGTH,
|
||||
|
|
@ -353,6 +354,16 @@ def _check_cli_sso_start_rate_limit(
|
|||
)
|
||||
|
||||
|
||||
def _read_cli_sso_flow(cache: DualCache, cache_key: str) -> object:
|
||||
redis_cache: Final = cache.redis_cache
|
||||
if redis_cache is None:
|
||||
return cache.get_cache(key=cache_key)
|
||||
try:
|
||||
return redis_cache.get_cache(key=cache_key)
|
||||
except RedisCircuitBreakerOpenError:
|
||||
return None
|
||||
|
||||
|
||||
def _get_cli_sso_flow_or_raise(login_id: str | None, cache: DualCache) -> dict:
|
||||
if isinstance(login_id, str) and login_id.startswith("sk-"):
|
||||
raise HTTPException(
|
||||
|
|
@ -365,12 +376,7 @@ def _get_cli_sso_flow_or_raise(login_id: str | None, cache: DualCache) -> dict:
|
|||
if not _is_valid_cli_sso_login_id(login_id):
|
||||
raise HTTPException(status_code=400, detail="Invalid CLI login session id")
|
||||
|
||||
cache_key: Final = _get_cli_sso_flow_cache_key(cast(str, login_id))
|
||||
redis_cache: Final = cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
flow = redis_cache.get_cache(key=cache_key)
|
||||
else:
|
||||
flow = cache.get_cache(key=cache_key)
|
||||
flow = _read_cli_sso_flow(cache, _get_cli_sso_flow_cache_key(cast(str, login_id)))
|
||||
if isinstance(flow, str):
|
||||
try:
|
||||
flow = _as_object(json.loads(flow))
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import time
|
|||
import traceback
|
||||
import warnings
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Collection, Mapping, MutableMapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType, UnionType
|
||||
from typing import (
|
||||
|
|
@ -125,6 +126,7 @@ from litellm.router_utils.auto_router_model_naming import (
|
|||
count_capability_routers,
|
||||
validate_complexity_router_config_placement,
|
||||
)
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.utils import (
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
|
|
@ -234,6 +236,7 @@ import litellm._redis
|
|||
from litellm import Router
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger
|
||||
from litellm.caching.caching import DualCache, RedisCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||||
from litellm.constants import (
|
||||
_REALTIME_BODY_CACHE_SIZE,
|
||||
|
|
@ -2682,6 +2685,12 @@ async def _read_spend_counter_estimate(counter_key: str, fallback_spend: float)
|
|||
return fallback_spend, False
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PendingSpendIncrement:
|
||||
counter_key: str
|
||||
increment: float
|
||||
|
||||
|
||||
async def increment_spend_counters(
|
||||
token: str | None,
|
||||
team_id: str | None,
|
||||
|
|
@ -2716,7 +2725,7 @@ async def increment_spend_counters(
|
|||
|
||||
cost: Final[float] = response_cost
|
||||
|
||||
async def _key_scope(key_token: str) -> None:
|
||||
async def _key_scope(key_token: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
# key_token arrives pre-hashed from metadata["user_api_key"] (auth flow
|
||||
# hashes raw "sk-..." keys before they reach the callback). The
|
||||
# startswith("sk-") check is a safety net matching update_cache —
|
||||
|
|
@ -2727,30 +2736,29 @@ async def increment_spend_counters(
|
|||
hash_token(token=key_token) if isinstance(key_token, str) and key_token.startswith("sk-") else key_token
|
||||
)
|
||||
key_counter_key: Final = f"spend:key:{hashed_token}"
|
||||
if key_counter_key not in reserved_counter_keys:
|
||||
await _init_and_increment_spend_counter(
|
||||
counter_key=key_counter_key,
|
||||
source_cache_key=hashed_token,
|
||||
increment=cost,
|
||||
key_pending: Final[tuple[_PendingSpendIncrement, ...]] = (
|
||||
()
|
||||
if key_counter_key in reserved_counter_keys
|
||||
else (
|
||||
await _prepare_spend_counter_increment(
|
||||
counter_key=key_counter_key,
|
||||
source_cache_key=hashed_token,
|
||||
increment=cost,
|
||||
),
|
||||
)
|
||||
|
||||
key_obj: Final[object] = await user_api_key_cache.async_get_cache(key=hashed_token)
|
||||
if key_obj is None:
|
||||
return
|
||||
key_budget_limits = getattr(key_obj, "budget_limits", None) or (
|
||||
key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
|
||||
)
|
||||
if isinstance(key_budget_limits, str):
|
||||
key_budget_limits = json.loads(key_budget_limits)
|
||||
if not isinstance(key_budget_limits, list):
|
||||
return
|
||||
for window in key_budget_limits:
|
||||
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
|
||||
key_window_reset_at = window.get("reset_at") if isinstance(window, dict) else window.reset_at
|
||||
key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
|
||||
|
||||
async def _key_window_increment(window: object) -> _PendingSpendIncrement | None:
|
||||
duration = (
|
||||
window["budget_duration"] if isinstance(window, dict) else getattr(window, "budget_duration", None)
|
||||
)
|
||||
key_window_reset_at = (
|
||||
window.get("reset_at") if isinstance(window, dict) else getattr(window, "reset_at", None)
|
||||
)
|
||||
key_window_counter: Final = f"spend:key:{hashed_token}:window:{duration}"
|
||||
key_window_start = get_budget_window_start(window)
|
||||
if key_window_counter not in reserved_counter_keys:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending_window: Final = (
|
||||
await _prepare_window_spend_counter_increment(
|
||||
counter_key=key_window_counter,
|
||||
entity_type="Key",
|
||||
entity_id=hashed_token,
|
||||
|
|
@ -2758,6 +2766,9 @@ async def increment_spend_counters(
|
|||
window_start=key_window_start,
|
||||
increment=cost,
|
||||
)
|
||||
if key_window_counter not in reserved_counter_keys
|
||||
else None
|
||||
)
|
||||
await _enqueue_window_spend_row_update(
|
||||
entity_type=Litellm_EntityType.KEY,
|
||||
entity_id=hashed_token,
|
||||
|
|
@ -2767,33 +2778,48 @@ async def increment_spend_counters(
|
|||
increment=cost,
|
||||
request_started_at=request_started_at,
|
||||
)
|
||||
return pending_window
|
||||
|
||||
async def _team_scope(scope_team_id: str) -> None:
|
||||
team_counter_key: Final = f"spend:team:{scope_team_id}"
|
||||
if team_counter_key not in reserved_counter_keys:
|
||||
await _init_and_increment_spend_counter(
|
||||
counter_key=team_counter_key,
|
||||
source_cache_key=f"team_id:{scope_team_id}",
|
||||
increment=cost,
|
||||
)
|
||||
|
||||
team_obj: Final[object] = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
|
||||
if team_obj is None:
|
||||
return
|
||||
team_budget_limits = getattr(team_obj, "budget_limits", None) or (
|
||||
team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
|
||||
key_obj: Final[object] = await user_api_key_cache.async_get_cache(key=hashed_token)
|
||||
if key_obj is None:
|
||||
return key_pending
|
||||
key_budget_limits = getattr(key_obj, "budget_limits", None) or (
|
||||
key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
|
||||
)
|
||||
if isinstance(team_budget_limits, str):
|
||||
team_budget_limits = json.loads(team_budget_limits)
|
||||
if not isinstance(team_budget_limits, list):
|
||||
return
|
||||
for window in team_budget_limits:
|
||||
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
|
||||
team_window_reset_at = window.get("reset_at") if isinstance(window, dict) else window.reset_at
|
||||
team_window_counter = f"spend:team:{scope_team_id}:window:{duration}"
|
||||
if isinstance(key_budget_limits, str):
|
||||
key_budget_limits = json.loads(key_budget_limits)
|
||||
if not isinstance(key_budget_limits, list):
|
||||
return key_pending
|
||||
window_pending: Final = await asyncio.gather(
|
||||
*(_key_window_increment(window) for window in key_budget_limits), return_exceptions=True
|
||||
)
|
||||
return key_pending + tuple(item for item in window_pending if item is not None)
|
||||
|
||||
async def _team_scope(scope_team_id: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
team_counter_key: Final = f"spend:team:{scope_team_id}"
|
||||
team_pending: Final[tuple[_PendingSpendIncrement, ...]] = (
|
||||
()
|
||||
if team_counter_key in reserved_counter_keys
|
||||
else (
|
||||
await _prepare_spend_counter_increment(
|
||||
counter_key=team_counter_key,
|
||||
source_cache_key=f"team_id:{scope_team_id}",
|
||||
increment=cost,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
async def _team_window_increment(window: object) -> _PendingSpendIncrement | None:
|
||||
duration = (
|
||||
window["budget_duration"] if isinstance(window, dict) else getattr(window, "budget_duration", None)
|
||||
)
|
||||
team_window_reset_at = (
|
||||
window.get("reset_at") if isinstance(window, dict) else getattr(window, "reset_at", None)
|
||||
)
|
||||
team_window_counter: Final = f"spend:team:{scope_team_id}:window:{duration}"
|
||||
team_window_start = get_budget_window_start(window)
|
||||
if team_window_counter not in reserved_counter_keys:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending_window: Final = (
|
||||
await _prepare_window_spend_counter_increment(
|
||||
counter_key=team_window_counter,
|
||||
entity_type="Team",
|
||||
entity_id=scope_team_id,
|
||||
|
|
@ -2801,6 +2827,9 @@ async def increment_spend_counters(
|
|||
window_start=team_window_start,
|
||||
increment=cost,
|
||||
)
|
||||
if team_window_counter not in reserved_counter_keys
|
||||
else None
|
||||
)
|
||||
await _enqueue_window_spend_row_update(
|
||||
entity_type=Litellm_EntityType.TEAM,
|
||||
entity_id=scope_team_id,
|
||||
|
|
@ -2810,25 +2839,47 @@ async def increment_spend_counters(
|
|||
increment=cost,
|
||||
request_started_at=request_started_at,
|
||||
)
|
||||
return pending_window
|
||||
|
||||
async def _team_member_scope(scope_user_id: str, scope_team_id: str) -> None:
|
||||
team_obj: Final[object] = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
|
||||
if team_obj is None:
|
||||
return team_pending
|
||||
team_budget_limits = getattr(team_obj, "budget_limits", None) or (
|
||||
team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
|
||||
)
|
||||
if isinstance(team_budget_limits, str):
|
||||
team_budget_limits = json.loads(team_budget_limits)
|
||||
if not isinstance(team_budget_limits, list):
|
||||
return team_pending
|
||||
window_pending: Final = await asyncio.gather(
|
||||
*(_team_window_increment(window) for window in team_budget_limits), return_exceptions=True
|
||||
)
|
||||
return team_pending + tuple(item for item in window_pending if item is not None)
|
||||
|
||||
async def _team_member_scope(
|
||||
scope_user_id: str, scope_team_id: str
|
||||
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
team_member_counter_key: Final = f"spend:team_member:{scope_user_id}:{scope_team_id}"
|
||||
if team_member_counter_key in reserved_counter_keys:
|
||||
return
|
||||
await _init_and_increment_spend_counter(
|
||||
counter_key=team_member_counter_key,
|
||||
source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
|
||||
increment=cost,
|
||||
return ()
|
||||
return (
|
||||
await _prepare_spend_counter_increment(
|
||||
counter_key=team_member_counter_key,
|
||||
source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
|
||||
increment=cost,
|
||||
),
|
||||
)
|
||||
|
||||
async def _user_scope(scope_user_id: str) -> None:
|
||||
async def _user_scope(scope_user_id: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
user_counter_key: Final = f"spend:user:{scope_user_id}"
|
||||
if user_counter_key in reserved_counter_keys:
|
||||
return
|
||||
await _init_and_increment_spend_counter(
|
||||
counter_key=user_counter_key,
|
||||
source_cache_key=scope_user_id,
|
||||
increment=cost,
|
||||
return ()
|
||||
return (
|
||||
await _prepare_spend_counter_increment(
|
||||
counter_key=user_counter_key,
|
||||
source_cache_key=scope_user_id,
|
||||
increment=cost,
|
||||
),
|
||||
)
|
||||
|
||||
scope_coros: Final = tuple(
|
||||
|
|
@ -2838,7 +2889,7 @@ async def increment_spend_counters(
|
|||
_team_scope(team_id) if team_id is not None else None,
|
||||
_team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None,
|
||||
_user_scope(user_id) if user_id is not None else None,
|
||||
_increment_end_user_and_tag_spend_counters(
|
||||
_prepare_end_user_and_tag_spend_increments(
|
||||
end_user_id=end_user_id,
|
||||
tags=tags,
|
||||
response_cost=cost,
|
||||
|
|
@ -2846,14 +2897,14 @@ async def increment_spend_counters(
|
|||
)
|
||||
if end_user_id is not None or tags is not None
|
||||
else None,
|
||||
_increment_model_access_group_spend_counters(
|
||||
_prepare_model_access_group_spend_increments(
|
||||
model_access_groups=model_access_groups,
|
||||
response_cost=cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
if model_access_groups
|
||||
else None,
|
||||
_increment_org_spend_counter(
|
||||
_prepare_org_spend_increment(
|
||||
org_id=org_id,
|
||||
response_cost=cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
|
|
@ -2868,7 +2919,20 @@ async def increment_spend_counters(
|
|||
# as orphaned tasks that race the caller's reservation-counter invalidation;
|
||||
# all scopes settle, then the first error propagates as before.
|
||||
scope_results: Final = await asyncio.gather(*scope_coros, return_exceptions=True)
|
||||
scope_errors: Final = [r for r in scope_results if isinstance(r, BaseException)]
|
||||
scope_errors: Final = tuple(
|
||||
item
|
||||
for scope in scope_results
|
||||
for item in (scope if isinstance(scope, tuple) else (scope,))
|
||||
if isinstance(item, BaseException)
|
||||
)
|
||||
pending: Final = tuple(
|
||||
item
|
||||
for scope in scope_results
|
||||
if not isinstance(scope, BaseException)
|
||||
for item in scope
|
||||
if not isinstance(item, BaseException)
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=pending)
|
||||
if scope_errors:
|
||||
raise scope_errors[0]
|
||||
|
||||
|
|
@ -2911,41 +2975,49 @@ async def _reconcile_budget_reservation_for_counter_update(
|
|||
return reserved_counter_keys
|
||||
|
||||
|
||||
async def _increment_end_user_and_tag_spend_counters(
|
||||
async def _prepare_end_user_and_tag_spend_increments(
|
||||
end_user_id: str | None,
|
||||
tags: list[str] | None,
|
||||
response_cost: float,
|
||||
reserved_counter_keys: set[str],
|
||||
) -> None:
|
||||
if end_user_id is not None:
|
||||
await _init_and_increment_unreserved_spend_counter(
|
||||
counter_key=f"spend:end_user:{end_user_id}",
|
||||
source_cache_key=end_user_cache_key(end_user_id),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
|
||||
if tags is None:
|
||||
return
|
||||
|
||||
seen_tags: Final[set[str]] = set()
|
||||
for tag_name in tags:
|
||||
if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags:
|
||||
continue
|
||||
seen_tags.add(tag_name)
|
||||
await _init_and_increment_unreserved_spend_counter(
|
||||
counter_key=f"spend:tag:{tag_name}",
|
||||
source_cache_key=tag_cache_key(tag_name),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
unique_tags: Final = (
|
||||
tuple(dict.fromkeys(tag for tag in tags if tag and isinstance(tag, str))) if tags is not None else ()
|
||||
)
|
||||
results: Final = await asyncio.gather(
|
||||
*(
|
||||
coro
|
||||
for coro in (
|
||||
_prepare_unreserved_spend_counter_increment(
|
||||
counter_key=f"spend:end_user:{end_user_id}",
|
||||
source_cache_key=end_user_cache_key(end_user_id),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
if end_user_id is not None
|
||||
else None,
|
||||
*(
|
||||
_prepare_unreserved_spend_counter_increment(
|
||||
counter_key=f"spend:tag:{tag_name}",
|
||||
source_cache_key=tag_cache_key(tag_name),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
for tag_name in unique_tags
|
||||
),
|
||||
)
|
||||
if coro is not None
|
||||
),
|
||||
return_exceptions=True,
|
||||
)
|
||||
return tuple(item for item in results if item is not None)
|
||||
|
||||
|
||||
async def _increment_model_access_group_spend_counters(
|
||||
async def _prepare_model_access_group_spend_increments(
|
||||
model_access_groups: Sequence[object],
|
||||
response_cost: float,
|
||||
reserved_counter_keys: set[str],
|
||||
) -> None:
|
||||
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
"""Charge the model access groups that authorized this request.
|
||||
|
||||
Without this the counter auth reads is written only by the reservation path, so
|
||||
|
|
@ -2959,55 +3031,63 @@ async def _increment_model_access_group_spend_counters(
|
|||
unique_groups: Final = tuple(
|
||||
dict.fromkeys(group for group in model_access_groups if group and isinstance(group, str))
|
||||
)
|
||||
for group in unique_groups:
|
||||
await _init_and_increment_unreserved_spend_counter(
|
||||
counter_key=model_access_group_spend_counter_key(group),
|
||||
source_cache_key=model_access_group_cache_key(group),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
results: Final = await asyncio.gather(
|
||||
*(
|
||||
_prepare_unreserved_spend_counter_increment(
|
||||
counter_key=model_access_group_spend_counter_key(group),
|
||||
source_cache_key=model_access_group_cache_key(group),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
for group in unique_groups
|
||||
),
|
||||
return_exceptions=True,
|
||||
)
|
||||
return tuple(item for item in results if item is not None)
|
||||
|
||||
|
||||
async def _increment_org_spend_counter(
|
||||
async def _prepare_org_spend_increment(
|
||||
org_id: str | None,
|
||||
response_cost: float,
|
||||
reserved_counter_keys: set[str],
|
||||
) -> None:
|
||||
) -> tuple[_PendingSpendIncrement, ...]:
|
||||
if org_id is None:
|
||||
return
|
||||
return ()
|
||||
|
||||
await _init_and_increment_unreserved_spend_counter(
|
||||
pending: Final = await _prepare_unreserved_spend_counter_increment(
|
||||
counter_key=f"spend:org:{org_id}",
|
||||
source_cache_key=[f"org_id:{org_id}:with_budget", f"org_id:{org_id}"],
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
return (pending,) if pending is not None else ()
|
||||
|
||||
|
||||
async def _init_and_increment_unreserved_spend_counter(
|
||||
async def _prepare_unreserved_spend_counter_increment(
|
||||
counter_key: str,
|
||||
source_cache_key: str | list[str],
|
||||
increment: float,
|
||||
reserved_counter_keys: set[str],
|
||||
) -> None:
|
||||
) -> _PendingSpendIncrement | None:
|
||||
if counter_key in reserved_counter_keys:
|
||||
return
|
||||
return None
|
||||
|
||||
await _init_and_increment_spend_counter(
|
||||
return await _prepare_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
source_cache_key=source_cache_key,
|
||||
increment=increment,
|
||||
)
|
||||
|
||||
|
||||
async def _init_and_increment_spend_counter(
|
||||
async def _prepare_spend_counter_increment(
|
||||
counter_key: str,
|
||||
source_cache_key: str | list[str],
|
||||
increment: float,
|
||||
):
|
||||
) -> _PendingSpendIncrement:
|
||||
"""
|
||||
Initialize counter from the authoritative DB spend value if not yet
|
||||
set, then atomically increment in both in-memory and Redis.
|
||||
set, then return the pending increment for the caller to apply in one
|
||||
pipelined Redis call.
|
||||
|
||||
On first access per pod:
|
||||
1. Check spend_counter_cache (in-memory -> Redis via DualCache)
|
||||
|
|
@ -3019,13 +3099,13 @@ async def _init_and_increment_spend_counter(
|
|||
the counter as absent and seed it. Using increment means the worst case
|
||||
is over-counting (conservative, blocks slightly early) rather than
|
||||
under-counting (would allow overspend).
|
||||
4. Increment atomically (both in-memory + Redis)
|
||||
4. Increment is returned for the caller to apply via pipeline
|
||||
"""
|
||||
await _ensure_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
source_cache_key=source_cache_key,
|
||||
)
|
||||
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
|
||||
return _PendingSpendIncrement(counter_key=counter_key, increment=increment)
|
||||
|
||||
|
||||
async def _enqueue_window_spend_row_update(
|
||||
|
|
@ -3077,20 +3157,20 @@ async def _enqueue_window_spend_row_update(
|
|||
)
|
||||
|
||||
|
||||
async def _init_and_increment_window_spend_counter(
|
||||
async def _prepare_window_spend_counter_increment(
|
||||
counter_key: str,
|
||||
entity_type: str,
|
||||
entity_id: str,
|
||||
window_duration: str | None,
|
||||
window_start: datetime | None,
|
||||
increment: float,
|
||||
):
|
||||
) -> _PendingSpendIncrement | None:
|
||||
if window_start is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Skipping spend counter increment for invalid budget window %s",
|
||||
counter_key,
|
||||
)
|
||||
return
|
||||
return None
|
||||
|
||||
initialized: Final = await _ensure_window_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
|
|
@ -3100,8 +3180,8 @@ async def _init_and_increment_window_spend_counter(
|
|||
window_start=window_start,
|
||||
)
|
||||
if initialized is False:
|
||||
return
|
||||
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
|
||||
return None
|
||||
return _PendingSpendIncrement(counter_key=counter_key, increment=increment)
|
||||
|
||||
|
||||
async def _ensure_spend_counter_initialized(
|
||||
|
|
@ -3234,6 +3314,34 @@ async def _invalidate_spend_counter(counter_key: str):
|
|||
)
|
||||
|
||||
|
||||
async def _apply_spend_counter_increments(pending: Sequence[_PendingSpendIncrement]) -> None:
|
||||
if not pending:
|
||||
return
|
||||
redis_cache: Final = spend_counter_cache.redis_cache
|
||||
if redis_cache is None:
|
||||
for item in pending:
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=item.counter_key,
|
||||
value=item.increment,
|
||||
refresh_ttl=True,
|
||||
)
|
||||
return
|
||||
ttl: Final = redis_cache.get_ttl()
|
||||
increment_list: Final = [ # mutable-ok: async_increment_pipeline signature requires list[RedisPipelineIncrementOperation]
|
||||
RedisPipelineIncrementOperation(key=item.counter_key, increment_value=item.increment, ttl=ttl)
|
||||
for item in pending
|
||||
]
|
||||
try:
|
||||
results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list)
|
||||
except Exception as e:
|
||||
await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending))
|
||||
if isinstance(e, RedisCircuitBreakerOpenError):
|
||||
return
|
||||
raise
|
||||
for item, current_value in zip(pending, results or ()):
|
||||
spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value)
|
||||
|
||||
|
||||
async def update_cache(
|
||||
token: str | None,
|
||||
user_id: str | None,
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ from litellm.caching.caching import (
|
|||
RedisCache,
|
||||
RedisClusterCache,
|
||||
)
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
from litellm.constants import (
|
||||
CONSUMED_REQUEST_TAGS_METADATA_KEY,
|
||||
DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS,
|
||||
|
|
@ -12935,8 +12936,10 @@ class Router:
|
|||
return await session_cache.async_get_cache(key=cache_key)
|
||||
return await session_cache.redis_cache.async_get_cache(key=cache_key)
|
||||
except Exception as e: # noqa: BLE001 # an optional binding must not make routing depend on Redis
|
||||
verbose_router_logger.warning(
|
||||
"Failed to read Claude Code session router binding; using the requested model: %s",
|
||||
log_redis_failure(
|
||||
verbose_router_logger,
|
||||
logging.WARNING,
|
||||
"Failed to read Claude Code session router binding; using the requested model",
|
||||
e,
|
||||
)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -3,12 +3,13 @@ Base class across routing strategies to abstract commmon functions like batch in
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from abc import ABC
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisPipelineIncrementOperation
|
||||
from litellm.caching.redis_cache import RedisPipelineIncrementOperation, log_redis_failure
|
||||
from litellm.constants import DEFAULT_REDIS_SYNC_INTERVAL
|
||||
|
||||
|
||||
|
|
@ -147,7 +148,7 @@ class BaseRoutingStrategy(ABC):
|
|||
return return_result
|
||||
|
||||
except Exception as e:
|
||||
verbose_router_logger.error("Error syncing in-memory cache with Redis: %s", e)
|
||||
log_redis_failure(verbose_router_logger, logging.ERROR, "Error syncing in-memory cache with Redis", e)
|
||||
self.redis_increment_operation_queue = []
|
||||
|
||||
def add_to_in_memory_keys_to_update(self, key: str):
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ anthropic:
|
|||
|
||||
import asyncio
|
||||
import builtins
|
||||
import logging
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Final
|
||||
|
|
@ -27,7 +28,7 @@ from typing import Any, Final
|
|||
import litellm
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisPipelineIncrementOperation
|
||||
from litellm.caching.redis_cache import RedisCache, RedisPipelineIncrementOperation, log_redis_failure
|
||||
from litellm.integrations.custom_logger import CustomLogger, Span
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
|
|
@ -92,6 +93,13 @@ class _LiteLLMParamsDictView:
|
|||
return dict(self._params)
|
||||
|
||||
|
||||
async def _push_increments_to_redis(redis_cache: RedisCache, queued: list[RedisPipelineIncrementOperation]) -> None:
|
||||
try:
|
||||
await redis_cache.async_increment_pipeline(increment_list=queued)
|
||||
except Exception as e:
|
||||
log_redis_failure(verbose_router_logger, logging.ERROR, "Error syncing in-memory cache with Redis", e)
|
||||
|
||||
|
||||
class RouterBudgetLimiting(CustomLogger):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -536,17 +544,13 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
"Pushing Redis Increment Pipeline for queue: %s",
|
||||
self.redis_increment_operation_queue,
|
||||
)
|
||||
if len(self.redis_increment_operation_queue) > 0:
|
||||
asyncio.create_task(
|
||||
self.dual_cache.redis_cache.async_increment_pipeline(
|
||||
increment_list=self.redis_increment_operation_queue,
|
||||
)
|
||||
)
|
||||
|
||||
queued: Final = self.redis_increment_operation_queue
|
||||
self.redis_increment_operation_queue = []
|
||||
if queued:
|
||||
asyncio.create_task(_push_increments_to_redis(self.dual_cache.redis_cache, queued))
|
||||
|
||||
except Exception as e:
|
||||
verbose_router_logger.error("Error syncing in-memory cache with Redis: %s", e)
|
||||
log_redis_failure(verbose_router_logger, logging.ERROR, "Error syncing in-memory cache with Redis", e)
|
||||
|
||||
async def _sync_in_memory_spend_with_redis(self):
|
||||
"""
|
||||
|
|
@ -601,7 +605,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
verbose_router_logger.debug("Updated in-memory cache for %s: %s", key, value)
|
||||
|
||||
except Exception as e:
|
||||
verbose_router_logger.error("Error syncing in-memory cache with Redis: %s", e)
|
||||
log_redis_failure(verbose_router_logger, logging.ERROR, "Error syncing in-memory cache with Redis", e)
|
||||
|
||||
def _get_budget_config_for_deployment(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing_extensions import TypedDict
|
|||
|
||||
from litellm import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
|
@ -27,6 +28,16 @@ class DeploymentHealthStateValue(TypedDict):
|
|||
reason: str
|
||||
|
||||
|
||||
def _read_shared_health_snapshot(cache: DualCache, key: str) -> object:
|
||||
redis_cache: Final = cache.redis_cache
|
||||
if redis_cache is None:
|
||||
return None
|
||||
try:
|
||||
return redis_cache.get_cache(key)
|
||||
except RedisCircuitBreakerOpenError:
|
||||
return None
|
||||
|
||||
|
||||
class DeploymentHealthCache:
|
||||
"""
|
||||
Cache for deployment health states produced by background health checks.
|
||||
|
|
@ -50,13 +61,12 @@ class DeploymentHealthCache:
|
|||
coexist on the one shared entry without erasing each other's results.
|
||||
The snapshot is read from Redis when available, since a pod-local read
|
||||
would only ever see this writer's own previous merge. When the Redis
|
||||
read comes back empty (a miss, or a swallowed connection error), the
|
||||
pod-local copy of the last merge is used so peers are not erased.
|
||||
read comes back empty (a miss, a swallowed connection error, or a read
|
||||
refused by the open circuit breaker), the pod-local copy of the last
|
||||
merge is used so peers are not erased.
|
||||
"""
|
||||
try:
|
||||
redis_raw: Final = (
|
||||
self.cache.redis_cache.get_cache(self.CACHE_KEY) if self.cache.redis_cache is not None else None
|
||||
)
|
||||
redis_raw: Final = _read_shared_health_snapshot(self.cache, self.CACHE_KEY)
|
||||
raw: Final = redis_raw if isinstance(redis_raw, dict) else self.cache.get_cache(key=self.CACHE_KEY)
|
||||
existing: Final = raw if isinstance(raw, dict) else {}
|
||||
expiry_seconds: Final = self.staleness_threshold * 1.5
|
||||
|
|
|
|||
|
|
@ -220,6 +220,7 @@ e2e-dev = [
|
|||
"playwright==1.61.0",
|
||||
"websockets>=15.0.1,<16.0",
|
||||
"locust==2.45.0",
|
||||
"psutil==7.2.2",
|
||||
"mcp>=1.28.1,<2.0",
|
||||
]
|
||||
proxy-dev = [
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
|
|||
- `logging/` - logging-integration delivery (datadog and friends)
|
||||
- `security/` - secret handling and log-leak protection
|
||||
- `router/` - routing and reliability behavior (fallbacks, cooldowns)
|
||||
- `load/` - performance-category tests, kept OUT of the main suite: throughput/load SLO tests are a different testing category from functional e2e (variance-driven, historically flaky) and live outside this suite until re-implemented as their own pipeline (LIT-5163); do not add a live load test that runs in the default collection. What remains here: the weekly session-anomaly test (`test_weekly_session_anomaly_e2e.py`, Claude Code-shaped multi-turn sessions against real providers with ceilings on error rate, cache read/write, turn time, and spend; marked `weekly` and deselected unless `E2E_WEEKLY_ANOMALY` is set, driven by `.github/workflows/weekly_load_anomaly.yml`) and markerless harness unit tests for the Locust/session-anomaly aggregation logic
|
||||
- `load/` - performance-category tests, kept OUT of the main suite: throughput/load SLO tests are a different testing category from functional e2e (variance-driven, historically flaky) and live outside this suite until re-implemented as their own pipeline (LIT-5163); do not add a live load test that runs in the default collection. What lives here: the weekly session-anomaly test (`test_weekly_session_anomaly_e2e.py`, Claude Code-shaped multi-turn sessions against real providers with ceilings on error rate, cache read/write, turn time, and spend; marked `weekly` and deselected unless `E2E_WEEKLY_ANOMALY` is set, driven by `.github/workflows/weekly_load_anomaly.yml`), the Redis chaos test (`test_redis_chaos_e2e.py`, locust load against mock deployments split round robin over `/chat/completions` and `/v1/messages`, one endpoint per simulated user, with `CLIENT PAUSE ALL` on the proxy's Redis mid-run to simulate it being down outright, asserting zero failed requests on every endpoint, budgeting RSS and CPU-per-request as ratios against the same run's healthy phase, and holding p50/p90/p99 latency and log-bytes-per-request to flat ceilings (a ratio cannot bound those two: an open breaker skips Redis instead of waiting on it, so the chaos phase can measure cheaper than baseline while still being far slower than a user should see); needs a proxy booted from `gateway/redis_chaos_ci_config.yml` on the same host with `E2E_PROXY_PID` and `E2E_PROXY_LOG` set, marked `redis_chaos`, deselected unless `E2E_REDIS_CHAOS` is set and excluded from the per-PR selector like the rest of `load/`, driven by `.github/workflows/test-e2e-redis-chaos.yml` and by the Buildkite `e2e-redis-chaos` step in project-releaser, which runs the proxy, Postgres and Valkey co-located with pytest in one pod and sets the opt-in), and markerless harness unit tests for the locust, process-usage, and session-anomaly aggregation logic
|
||||
- `other/` - the holding-pen suite for the `other.*` registry cluster with no home of its own yet: the master-key auth gate and the process-lifecycle health probes (liveness, public readiness, authenticated readiness diagnostics). Promote a cluster out once it is large/stable enough for its own suite
|
||||
- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests
|
||||
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher plus `_*_unit_tests/` trees. The HTTP probes ride the shared transport (`ProxyClient.count_tokens` / `ProxyClient.messages`); the CLI-driving path stays bespoke
|
||||
|
|
|
|||
|
|
@ -54,6 +54,10 @@ Some suites need extra services the bare proxy does not start. The `logging/` OT
|
|||
|
||||
A couple of logging destinations are configured on the proxy rather than by the test. The Weave tests scope their callback to the key they create, but litellm builds the `weave_otel` logger from `WANDB_API_KEY` and `WANDB_PROJECT_ID` before it applies the per-key vars, so the proxy needs both in its own environment or the key-scoped callback never initializes and nothing ships
|
||||
|
||||
### Redis chaos load test
|
||||
|
||||
The Redis chaos test under `load/` needs a proxy and its own Redis on the same host, using `gateway/redis_chaos_ci_config.yml`. `.github/workflows/test-e2e-redis-chaos.yml` boots that stack, and the Buildkite `e2e-redis-chaos` step in project-releaser runs the proxy, Postgres and Valkey together in an isolated pod. The test is deselected unless `E2E_REDIS_CHAOS` is set
|
||||
|
||||
### Record and replay
|
||||
|
||||
Record/replay scopes to the proxy's provider-bound traffic only. In `E2E_FIXTURE_MODE=record` the harness boots a local provider-edge server, edge-wired tests register their deployments with an `api_base` pointing at it, and every provider call the proxy makes is forwarded verbatim and written to a fixture bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`). `E2E_FIXTURE_MODE=replay` runs the same tests against the same live proxy and database, but the edge answers the proxy's provider calls from the bundle instead of the provider, so the run makes zero provider calls and spends nothing while key auth, routing, cost calculation, and spend-log writes all still execute for real. Unset (or `live`) behaves exactly as before the knob existed. Both record and replay need the proxy up; only the provider is taken out of the loop
|
||||
|
|
|
|||
|
|
@ -20,16 +20,14 @@ from datetime import datetime, timezone
|
|||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
from e2e_config import CONTROL_PLANE_BASE_URL, FIXTURE_DIR, FIXTURE_MODE_RAW, PROXY_BASE_URL
|
||||
from e2e_db import RESET_OPT_IN_ENV, reset_spend_logs, run_spend_log_cleanup
|
||||
from fixture_mode import fixture_mode_collection_error, fixture_report_lines
|
||||
from provider_edge import replay_leftover_error
|
||||
from junit_properties import attach_result_properties
|
||||
from lifecycle import ProxyClientProvider, ResourceManager
|
||||
from provider_edge import replay_leftover_error
|
||||
from proxy_client import ProxyClient, build_proxy_client
|
||||
|
||||
|
||||
_E2E_TEST_RAN = pytest.StashKey[bool]()
|
||||
_CALL_PASSED = pytest.StashKey[bool]()
|
||||
|
||||
|
|
@ -60,6 +58,11 @@ def pytest_configure(config: pytest.Config) -> None:
|
|||
"markers",
|
||||
"managed_files: needs a proxy running with require_managed_files enabled; deselected unless E2E_MANAGED_FILES_STACK is set",
|
||||
)
|
||||
config.addinivalue_line(
|
||||
"markers",
|
||||
"redis_chaos: load test that pauses the proxy's Redis outright mid-run; needs a proxy booted from "
|
||||
"gateway/redis_chaos_ci_config.yml on the same host, and is deselected unless E2E_REDIS_CHAOS is set",
|
||||
)
|
||||
|
||||
|
||||
def pytest_sessionstart(session: pytest.Session) -> None:
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@
|
|||
- {id: reliability.cache.exact.returns_cached, module: reliability, tier: P1, behavior: cache, variant: exact, assertions: [returns_cached], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/caching.py", rationale: "Response cache returns cached on exact match"}
|
||||
- {id: reliability.cache.prompt_caching_model_select.returns_cached, module: reliability, tier: P1, behavior: cache, variant: prompt_caching_model_select, assertions: [returns_cached], exercised_on: [chat_completions], source: "router_utils/prompt_caching_cache.py", rationale: "Selects model supporting prompt caching for cacheable prefix"}
|
||||
- {id: reliability.circuit_breaker.redis.trips_then_recovers, module: reliability, tier: P0, behavior: circuit_breaker, variant: redis, assertions: [trips_then_recovers], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/redis_cache.py:99", rationale: "Redis breaker CLOSED->OPEN->HALF_OPEN; guards all cache/rate-limit ops"}
|
||||
- {id: reliability.circuit_breaker.redis_timeout.stays_responsive, module: reliability, tier: P1, behavior: circuit_breaker, variant: redis_timeout, assertions: [stays_responsive], exercised_on: [chat_completions, messages], source: "litellm/proxy/hooks/proxy_track_cost_callback.py:386", fail_before_fix: proven, rationale: "Under locust load split round robin over /chat/completions and /v1/messages with every request retrying through failing mock deployments, holding Redis in CLIENT PAUSE ALL for the phase trips the breaker and every request still succeeds, with latency, RSS, and CPU reported as p50/p90/p99 against the pre-pause baseline; on v1.100.0 the failed-tracking alert body doubled per request until the worker OOMed (LIT-6780)"}
|
||||
- {id: reliability.timeout.request_timeout.exceeds_deadline, module: reliability, tier: P1, behavior: timeout, variant: request_timeout, assertions: [exceeds_deadline], exercised_on: [chat_completions, messages], source: "litellm/router.py:545-551", rationale: "Per-request timeout raises Timeout"}
|
||||
- {id: reliability.timeout.stream_timeout.exceeds_deadline, module: reliability, tier: P1, behavior: timeout, variant: stream_timeout, assertions: [exceeds_deadline], exercised_on: [chat_completions], source: "litellm/router.py:551", rationale: "Streaming chunk-delivery timeout"}
|
||||
- {id: reliability.perf.throughput.under_slo, module: reliability, tier: P1, behavior: perf, variant: throughput, assertions: [under_slo], exercised_on: [chat_completions, messages], source: grammar, rationale: "Throughput SLO under load"}
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ from pathlib import Path
|
|||
from typing import Final
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from fixture_mode import deterministic_marker, parse_fixture_mode
|
||||
from provider_edge import provider_edge_api_base
|
||||
|
||||
|
|
@ -135,6 +134,7 @@ LOAD_MIN_CONCURRENCY_EFFICIENCY = float(os.environ.get("E2E_LOAD_MIN_CONCURRENCY
|
|||
|
||||
WEEKLY_ANOMALY_OPT_IN_ENV = "E2E_WEEKLY_ANOMALY"
|
||||
MANAGED_FILES_OPT_IN_ENV = "E2E_MANAGED_FILES_STACK"
|
||||
REDIS_CHAOS_OPT_IN_ENV = "E2E_REDIS_CHAOS"
|
||||
ANOMALY_SESSIONS = int(os.environ.get("E2E_ANOMALY_SESSIONS", "6"))
|
||||
ANOMALY_TURNS_PER_SESSION = int(os.environ.get("E2E_ANOMALY_TURNS_PER_SESSION", "6"))
|
||||
ANOMALY_TURN_ATTEMPTS = int(os.environ.get("E2E_ANOMALY_TURN_ATTEMPTS", "3"))
|
||||
|
|
@ -172,8 +172,7 @@ def datadog_mcp_url(*, toolsets: str = "core") -> str:
|
|||
site = (
|
||||
os.environ.get("DD_SITE", DD_SITE) or "datadoghq.com"
|
||||
).strip().removeprefix("https://").removeprefix("http://").rstrip("/")
|
||||
if site.startswith("app."):
|
||||
site = site[len("app.") :]
|
||||
site = site.removeprefix("app.")
|
||||
host = "mcp.datadoghq.com" if site in ("", "datadoghq.com") else f"mcp.{site}"
|
||||
base = f"https://{host}/v1/mcp"
|
||||
return f"{base}?toolsets={toolsets}" if toolsets else base
|
||||
|
|
|
|||
19
tests/e2e/gateway/redis_chaos_ci_config.yml
Normal file
19
tests/e2e/gateway/redis_chaos_ci_config.yml
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
general_settings:
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
store_model_in_db: true
|
||||
use_redis_transaction_buffer: true
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["prometheus"]
|
||||
require_auth_for_metrics_endpoint: false
|
||||
enable_redis_auth_cache: true
|
||||
cache: true
|
||||
cache_params:
|
||||
type: redis
|
||||
host: 127.0.0.1
|
||||
port: 6379
|
||||
socket_timeout: 0.1
|
||||
|
||||
router_settings:
|
||||
num_retries: 2
|
||||
disable_cooldowns: true
|
||||
|
|
@ -3,26 +3,23 @@ from __future__ import annotations
|
|||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_config import WEEKLY_ANOMALY_OPT_IN_ENV
|
||||
from e2e_config import REDIS_CHAOS_OPT_IN_ENV, WEEKLY_ANOMALY_OPT_IN_ENV
|
||||
from load_client import LoadClient, build_client
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
_OPT_IN_MARKERS = (
|
||||
("weekly", WEEKLY_ANOMALY_OPT_IN_ENV),
|
||||
("redis_chaos", REDIS_CHAOS_OPT_IN_ENV),
|
||||
)
|
||||
|
||||
def pytest_collection_modifyitems(
|
||||
config: pytest.Config, items: list[pytest.Item]
|
||||
) -> None:
|
||||
if os.environ.get(WEEKLY_ANOMALY_OPT_IN_ENV):
|
||||
return
|
||||
deselected = [
|
||||
item for item in items if item.get_closest_marker("weekly") is not None
|
||||
]
|
||||
|
||||
def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item]) -> None:
|
||||
opted_out = {marker for marker, opt_in_env in _OPT_IN_MARKERS if not os.environ.get(opt_in_env)}
|
||||
deselected = [item for item in items if any(item.get_closest_marker(marker) is not None for marker in opted_out)]
|
||||
if not deselected:
|
||||
return
|
||||
config.hook.pytest_deselected(items=deselected)
|
||||
items[:] = [
|
||||
item for item in items if item.get_closest_marker("weekly") is None
|
||||
]
|
||||
items[:] = [item for item in items if item not in deselected]
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
|
|
|
|||
|
|
@ -1,17 +1,26 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from itertools import accumulate
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
_LOCUSTFILE = Path(__file__).with_name("locustfile.py")
|
||||
_CSV_PREFIX = "locust"
|
||||
_GENERATOR_SATURATION_MARKER = "CPU usage above"
|
||||
_MAX_REPORTED_ERRORS = 5
|
||||
|
||||
|
||||
class LocustStatEntry(BaseModel):
|
||||
name: str
|
||||
num_requests: int
|
||||
num_failures: int
|
||||
start_time: float
|
||||
|
|
@ -29,12 +38,25 @@ class LoadError:
|
|||
occurrences: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EndpointLoad:
|
||||
"""One route's share of a phase, so a run that silently drove only one of them is visible."""
|
||||
|
||||
name: str
|
||||
requests: int
|
||||
failures: int
|
||||
p50_seconds: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LoadResult:
|
||||
requests: int
|
||||
failures: int
|
||||
requests_per_second: float
|
||||
median_response_seconds: float
|
||||
p50_seconds: float
|
||||
p90_seconds: float
|
||||
p99_seconds: float
|
||||
endpoints: tuple[EndpointLoad, ...]
|
||||
errors: tuple[LoadError, ...]
|
||||
generator_warnings: tuple[str, ...]
|
||||
|
||||
|
|
@ -53,33 +75,65 @@ class LoadResult:
|
|||
lines.append("locust recorded no error breakdown")
|
||||
return "; ".join((*lines, *self.generator_warnings))
|
||||
|
||||
def latency_summary(self) -> str:
|
||||
return f"p50 {self.p50_seconds:.3f}s, p90 {self.p90_seconds:.3f}s, p99 {self.p99_seconds:.3f}s"
|
||||
|
||||
def median_seconds(entries: list[LocustStatEntry]) -> float:
|
||||
samples = sorted(
|
||||
(milliseconds, count) for entry in entries for milliseconds, count in entry.response_times.items()
|
||||
)
|
||||
def endpoint_summary(self) -> str:
|
||||
return ", ".join(
|
||||
f"{endpoint.name} {endpoint.requests} requests, {endpoint.failures} failures, "
|
||||
f"p50 {endpoint.p50_seconds:.3f}s"
|
||||
for endpoint in self.endpoints
|
||||
)
|
||||
|
||||
|
||||
def percentile_seconds(entries: Sequence[LocustStatEntry], fraction: float) -> float:
|
||||
"""The response time at `fraction` of the merged histograms, in seconds.
|
||||
|
||||
Locust buckets response times by millisecond, so this reads the first bucket whose
|
||||
running count reaches the rank, the same lower-sample convention locust's own
|
||||
percentiles use.
|
||||
"""
|
||||
samples = sorted((milliseconds, count) for entry in entries for milliseconds, count in entry.response_times.items())
|
||||
total = sum(count for _, count in samples)
|
||||
if total == 0:
|
||||
return 0.0
|
||||
running = accumulate(count for _, count in samples)
|
||||
return next(
|
||||
milliseconds for (milliseconds, _), seen in zip(samples, running) if seen >= total / 2
|
||||
) / 1000.0
|
||||
rank: Final = total * fraction
|
||||
return next(milliseconds for (milliseconds, _), seen in zip(samples, running) if seen >= rank) / 1000.0
|
||||
|
||||
|
||||
def per_endpoint(entries: Sequence[LocustStatEntry]) -> tuple[EndpointLoad, ...]:
|
||||
"""Each locust request name's own totals, in the order the names first appear."""
|
||||
names: Final = tuple(dict.fromkeys(entry.name for entry in entries))
|
||||
grouped: Final = ((name, tuple(entry for entry in entries if entry.name == name)) for name in names)
|
||||
return tuple(
|
||||
EndpointLoad(
|
||||
name=name,
|
||||
requests=sum(entry.num_requests for entry in group),
|
||||
failures=sum(entry.num_failures for entry in group),
|
||||
p50_seconds=percentile_seconds(group, 0.5),
|
||||
)
|
||||
for name, group in grouped
|
||||
)
|
||||
|
||||
|
||||
def aggregate_stats(
|
||||
entries: list[LocustStatEntry],
|
||||
entries: Sequence[LocustStatEntry],
|
||||
errors: tuple[LoadError, ...],
|
||||
generator_warnings: tuple[str, ...],
|
||||
) -> LoadResult:
|
||||
requests = sum(entry.num_requests for entry in entries)
|
||||
failures = sum(entry.num_failures for entry in entries)
|
||||
endpoints = per_endpoint(entries)
|
||||
if not entries or requests == 0:
|
||||
return LoadResult(
|
||||
requests=requests,
|
||||
failures=failures,
|
||||
requests_per_second=0.0,
|
||||
median_response_seconds=0.0,
|
||||
p50_seconds=0.0,
|
||||
p90_seconds=0.0,
|
||||
p99_seconds=0.0,
|
||||
endpoints=endpoints,
|
||||
errors=errors,
|
||||
generator_warnings=generator_warnings,
|
||||
)
|
||||
|
|
@ -88,7 +142,10 @@ def aggregate_stats(
|
|||
requests=requests,
|
||||
failures=failures,
|
||||
requests_per_second=requests / elapsed if elapsed > 0 else 0.0,
|
||||
median_response_seconds=median_seconds(entries),
|
||||
p50_seconds=percentile_seconds(entries, 0.5),
|
||||
p90_seconds=percentile_seconds(entries, 0.9),
|
||||
p99_seconds=percentile_seconds(entries, 0.99),
|
||||
endpoints=endpoints,
|
||||
errors=errors,
|
||||
generator_warnings=generator_warnings,
|
||||
)
|
||||
|
|
@ -121,3 +178,74 @@ def read_generator_warnings(stderr: str) -> tuple[str, ...]:
|
|||
if _GENERATOR_SATURATION_MARKER in line
|
||||
)
|
||||
return tuple(dict.fromkeys(saturated))
|
||||
|
||||
|
||||
def run_gateway_load(
|
||||
*,
|
||||
base_url: str,
|
||||
api_keys: tuple[str, ...],
|
||||
model: str,
|
||||
endpoints: tuple[str, ...],
|
||||
users: int,
|
||||
spawn_rate: float,
|
||||
duration_seconds: float,
|
||||
) -> LoadResult:
|
||||
"""Drive `endpoints` from headless locust and aggregate what it reported.
|
||||
|
||||
Each simulated user picks one of `api_keys`, so auth and budget lookups spread over a
|
||||
pool of virtual keys instead of keeping one key's cache entry permanently warm, and one
|
||||
of `endpoints` round robin, so the run covers every route the caller asked for.
|
||||
"""
|
||||
with tempfile.TemporaryDirectory(prefix="e2e-load-") as report_dir:
|
||||
csv_prefix = Path(report_dir) / _CSV_PREFIX
|
||||
completed = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-m",
|
||||
"locust",
|
||||
"--headless",
|
||||
"--json",
|
||||
"--csv",
|
||||
str(csv_prefix),
|
||||
"--locustfile",
|
||||
str(_LOCUSTFILE),
|
||||
"--host",
|
||||
base_url,
|
||||
"--users",
|
||||
str(users),
|
||||
"--spawn-rate",
|
||||
str(spawn_rate),
|
||||
"--run-time",
|
||||
f"{int(duration_seconds)}s",
|
||||
"--exit-code-on-error",
|
||||
"0",
|
||||
],
|
||||
env={
|
||||
**os.environ,
|
||||
"LOAD_API_KEYS": ",".join(api_keys),
|
||||
"LOAD_MODEL": model,
|
||||
"LOAD_ENDPOINTS": ",".join(endpoints),
|
||||
},
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=duration_seconds + 120,
|
||||
check=False,
|
||||
)
|
||||
if completed.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"locust exited {completed.returncode} before it could report throughput "
|
||||
f"(a startup failure, not request failures, which are folded into the JSON summary via "
|
||||
f"--exit-code-on-error 0):\n{completed.stderr}"
|
||||
)
|
||||
try:
|
||||
entries = _STATS_ADAPTER.validate_json(completed.stdout)
|
||||
except ValueError as exc:
|
||||
raise RuntimeError(
|
||||
f"locust exited 0 but did not print a parseable --json throughput summary on stdout; "
|
||||
f"got stdout={completed.stdout!r}, stderr={completed.stderr!r}"
|
||||
) from exc
|
||||
return aggregate_stats(
|
||||
entries,
|
||||
read_errors(csv_prefix.with_name(f"{_CSV_PREFIX}_failures.csv")),
|
||||
read_generator_warnings(completed.stderr),
|
||||
)
|
||||
|
|
|
|||
52
tests/e2e/load/locustfile.py
Normal file
52
tests/e2e/load/locustfile.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import random
|
||||
import uuid
|
||||
from itertools import cycle
|
||||
from typing import Final
|
||||
|
||||
from locust import FastHttpUser, constant, task
|
||||
|
||||
_MODEL: Final = os.environ["LOAD_MODEL"]
|
||||
_API_KEYS: Final = tuple(os.environ["LOAD_API_KEYS"].split(","))
|
||||
_NEXT_ENDPOINT: Final = cycle(os.environ["LOAD_ENDPOINTS"].split(","))
|
||||
_FILLER: Final = "x" * 40_000
|
||||
|
||||
|
||||
def _payload() -> dict[str, object]:
|
||||
"""A prompt no other request sent, so the response cache never answers for the deployment.
|
||||
|
||||
Both endpoints take the same body: /v1/messages requires max_tokens, which /chat/completions
|
||||
also accepts, so one payload serves the whole round robin. Padded to tens of KB so a
|
||||
per-request bookkeeping cost that scales with body size (string formatting, hashing) shows
|
||||
up in the CPU and log-size budgets instead of hiding behind a 40-byte prompt.
|
||||
"""
|
||||
return {
|
||||
"model": _MODEL,
|
||||
"messages": [{"role": "user", "content": f"load test ping {uuid.uuid4().hex} {_FILLER}"}],
|
||||
"max_tokens": 16,
|
||||
}
|
||||
|
||||
|
||||
class GatewayUser(FastHttpUser):
|
||||
"""One simulated user, pinned to one endpoint for its lifetime.
|
||||
|
||||
Endpoints are handed out round robin as users spawn, so a run spreads evenly over them
|
||||
while each user's traffic stays on a single route, the way a real client behaves.
|
||||
"""
|
||||
|
||||
wait_time = constant(0)
|
||||
|
||||
def on_start(self) -> None:
|
||||
self.headers = {"Authorization": f"Bearer {random.choice(_API_KEYS)}"}
|
||||
self.endpoint = next(_NEXT_ENDPOINT)
|
||||
|
||||
@task
|
||||
def call(self) -> None:
|
||||
self.client.post( # pyright: ignore[reportUnknownMemberType] # locust FastHttpSession.post types json/**kwargs as Any
|
||||
self.endpoint,
|
||||
json=_payload(),
|
||||
headers=self.headers,
|
||||
name=self.endpoint,
|
||||
)
|
||||
81
tests/e2e/load/phase_budget.py
Normal file
81
tests/e2e/load/phase_budget.py
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
"""Comparing one load phase against another, for tests that degrade a dependency mid-run.
|
||||
|
||||
Two shapes of ceiling, because the metrics divide into two kinds. RSS and CPU are
|
||||
machine-shaped: RSS scales with worker count and CPU with core count, so an absolute number
|
||||
calibrated on one runner means nothing on the next, and what travels is the ratio against a
|
||||
healthy phase measured on the same machine in the same run. Latency and log volume are not:
|
||||
a ratio there is actively misleading, because a dependency that fails fast once its breaker
|
||||
opens can make the degraded phase look cheaper than the healthy one while still being far
|
||||
slower or noisier than a user should ever see. Those get a flat ceiling, which is the promise
|
||||
the test is actually making.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
|
||||
def _rendered(value: float, unit: str, decimals: int) -> str:
|
||||
return f"{value:.{decimals}f}{unit}"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RatioBudget:
|
||||
"""One metric's healthy value, its degraded value, and how much growth is allowed."""
|
||||
|
||||
name: str
|
||||
baseline: float
|
||||
degraded: float
|
||||
ratio_ceiling: float
|
||||
unit: str
|
||||
decimals: int = 1
|
||||
|
||||
@property
|
||||
def ratio(self) -> float | None:
|
||||
"""How many times the baseline the degraded value is, or None if there is no baseline."""
|
||||
return self.degraded / self.baseline if self.baseline > 0 else None
|
||||
|
||||
def violation(self) -> str | None:
|
||||
"""Why this metric fails its budget, or None if it passes."""
|
||||
ratio: Final = self.ratio
|
||||
if ratio is None:
|
||||
return (
|
||||
f"{self.name} measured {_rendered(self.baseline, self.unit, self.decimals)} in the healthy phase, "
|
||||
f"so there is nothing to compare the degraded phase against; the measurement did not happen"
|
||||
)
|
||||
if ratio > self.ratio_ceiling:
|
||||
return (
|
||||
f"{self.name} went from {_rendered(self.baseline, self.unit, self.decimals)} healthy to "
|
||||
f"{_rendered(self.degraded, self.unit, self.decimals)} degraded, {ratio:.1f}x the baseline and past "
|
||||
f"the {self.ratio_ceiling:.1f}x allowed"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AbsoluteBudget:
|
||||
"""One metric's degraded value against a flat ceiling, for metrics a ratio cannot bound."""
|
||||
|
||||
name: str
|
||||
measured: float
|
||||
ceiling: float
|
||||
unit: str
|
||||
decimals: int = 1
|
||||
|
||||
def violation(self) -> str | None:
|
||||
"""Why this metric fails its budget, or None if it passes."""
|
||||
if self.measured > self.ceiling:
|
||||
return (
|
||||
f"{self.name} measured {_rendered(self.measured, self.unit, self.decimals)} in the degraded phase, "
|
||||
f"past the {_rendered(self.ceiling, self.unit, self.decimals)} allowed"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
Budget: TypeAlias = RatioBudget | AbsoluteBudget
|
||||
|
||||
|
||||
def violations(budgets: tuple[Budget, ...]) -> tuple[str, ...]:
|
||||
"""Every budget the run blew, so one failure reports all of them instead of the first."""
|
||||
return tuple(violation for budget in budgets if (violation := budget.violation()) is not None)
|
||||
164
tests/e2e/load/proxy_usage.py
Normal file
164
tests/e2e/load/proxy_usage.py
Normal file
|
|
@ -0,0 +1,164 @@
|
|||
"""Resident memory and CPU of the proxy process tree, sampled on a background thread.
|
||||
|
||||
The proxy under load runs several worker processes, and `/metrics` cannot report their
|
||||
memory: litellm sets PROMETHEUS_MULTIPROC_DIR when num_workers > 1, and the multiprocess
|
||||
collector drops the process collector's `process_resident_memory_bytes` /
|
||||
`process_cpu_seconds_total` entirely. So the test measures the tree itself through psutil,
|
||||
which needs the proxy to run on the same host as the test.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import psutil
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
|
||||
class _MemoryInfo(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
rss: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class UsageSample:
|
||||
elapsed_seconds: float
|
||||
rss_bytes: int
|
||||
cpu_seconds: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class UsageWindow:
|
||||
"""The samples taken across one phase, plus what they say about that phase."""
|
||||
|
||||
samples: tuple[UsageSample, ...]
|
||||
|
||||
def rss_percentile(self, fraction: float) -> int:
|
||||
if not self.samples:
|
||||
return 0
|
||||
ordered: Final = sorted(sample.rss_bytes for sample in self.samples)
|
||||
return ordered[_rank(len(ordered), fraction)]
|
||||
|
||||
def cpu_seconds_consumed(self) -> float:
|
||||
"""CPU seconds the tree burned across the window, from its monotonic counter."""
|
||||
if len(self.samples) < 2:
|
||||
return 0.0
|
||||
return self.samples[-1].cpu_seconds - self.samples[0].cpu_seconds
|
||||
|
||||
def cpu_seconds_per_request(self, requests: int) -> float:
|
||||
"""CPU seconds the tree spent per request served.
|
||||
|
||||
The portable cost figure: cores-busy saturates at the worker count under enough load,
|
||||
so it reads the same whether a request costs 10 ms of CPU or 40 ms. This does not.
|
||||
"""
|
||||
return self.cpu_seconds_consumed() / requests if requests else 0.0
|
||||
|
||||
def cpu_utilization_percentiles(self) -> tuple[float, float, float]:
|
||||
"""Per-interval CPU utilization (cores busy) at p50, p90 and p99.
|
||||
|
||||
Derived from consecutive samples of the cumulative counter rather than
|
||||
psutil's own cpu_percent, so it covers every process in the tree including
|
||||
workers that came and went between samples.
|
||||
"""
|
||||
rates: Final = sorted(
|
||||
(later.cpu_seconds - earlier.cpu_seconds) / (later.elapsed_seconds - earlier.elapsed_seconds)
|
||||
for earlier, later in zip(self.samples, self.samples[1:])
|
||||
if later.elapsed_seconds > earlier.elapsed_seconds
|
||||
)
|
||||
if not rates:
|
||||
return 0.0, 0.0, 0.0
|
||||
return (
|
||||
rates[_rank(len(rates), 0.5)],
|
||||
rates[_rank(len(rates), 0.9)],
|
||||
rates[_rank(len(rates), 0.99)],
|
||||
)
|
||||
|
||||
def summary(self) -> str:
|
||||
p50_cpu, p90_cpu, p99_cpu = self.cpu_utilization_percentiles()
|
||||
return (
|
||||
f"RSS p50 {self.rss_percentile(0.5) / 2**20:.0f} MB, "
|
||||
f"p90 {self.rss_percentile(0.9) / 2**20:.0f} MB, "
|
||||
f"p99 {self.rss_percentile(0.99) / 2**20:.0f} MB; "
|
||||
f"CPU cores busy p50 {p50_cpu:.2f}, p90 {p90_cpu:.2f}, p99 {p99_cpu:.2f}; "
|
||||
f"{self.cpu_seconds_consumed():.1f} CPU seconds consumed"
|
||||
)
|
||||
|
||||
|
||||
def _rank(count: int, fraction: float) -> int:
|
||||
"""Index of the sample at `fraction`, the same lower-sample convention as locust's percentiles."""
|
||||
return min(count - 1, max(0, math.ceil(count * fraction) - 1))
|
||||
|
||||
|
||||
def _read_process(process: psutil.Process) -> tuple[int, float] | None:
|
||||
try:
|
||||
with process.oneshot():
|
||||
memory: Final = _MemoryInfo.model_validate(process.memory_info())
|
||||
times: Final = process.cpu_times()
|
||||
return memory.rss, times.user + times.system
|
||||
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||
return None
|
||||
|
||||
|
||||
class ProxyUsageSampler:
|
||||
"""Samples the proxy process tree every `interval_seconds` until stopped.
|
||||
|
||||
`split()` returns the samples taken so far and starts a new window, so one sampler
|
||||
covers a baseline phase and a chaos phase without a gap between them.
|
||||
"""
|
||||
|
||||
def __init__(self, pid: int, interval_seconds: float = 1.0) -> None:
|
||||
self._process: Final = psutil.Process(pid)
|
||||
self._interval: Final = interval_seconds
|
||||
self._stop: Final = threading.Event()
|
||||
self._lock: Final = threading.Lock()
|
||||
self._samples: list[UsageSample] = [] # mutable-ok: a sampling buffer the reader drains under a lock
|
||||
self._started: Final = time.monotonic()
|
||||
self._thread: Final = threading.Thread(target=self._run, name="proxy-usage-sampler", daemon=True)
|
||||
|
||||
def __enter__(self) -> ProxyUsageSampler:
|
||||
self._thread.start()
|
||||
return self
|
||||
|
||||
def __exit__(self, *_: object) -> None:
|
||||
self._stop.set()
|
||||
self._thread.join(timeout=self._interval * 5)
|
||||
|
||||
def _tree(self) -> tuple[psutil.Process, ...]:
|
||||
try:
|
||||
return (self._process, *self._process.children(recursive=True))
|
||||
except psutil.NoSuchProcess:
|
||||
return ()
|
||||
|
||||
def _sample(self) -> UsageSample | None:
|
||||
readings: Final = tuple(reading for process in self._tree() if (reading := _read_process(process)) is not None)
|
||||
if not readings:
|
||||
return None
|
||||
return UsageSample(
|
||||
elapsed_seconds=time.monotonic() - self._started,
|
||||
rss_bytes=sum(rss for rss, _ in readings),
|
||||
cpu_seconds=sum(cpu for _, cpu in readings),
|
||||
)
|
||||
|
||||
def _run(self) -> None:
|
||||
while not self._stop.is_set():
|
||||
sample = self._sample()
|
||||
if sample is not None:
|
||||
with self._lock:
|
||||
self._samples.append(sample)
|
||||
self._stop.wait(self._interval)
|
||||
|
||||
def split(self) -> UsageWindow:
|
||||
"""The window that ends now; the next one starts from this window's last sample.
|
||||
|
||||
The boundary sample is carried into the next window so its CPU counter has a
|
||||
starting point, which is what makes the two windows' utilization comparable.
|
||||
"""
|
||||
with self._lock:
|
||||
taken = tuple(self._samples)
|
||||
self._samples = [taken[-1]] if taken else [] # rebind-ok: drains the buffer under the lock
|
||||
return UsageWindow(samples=taken)
|
||||
|
|
@ -1,13 +1,14 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
from locust_load import (
|
||||
LoadError,
|
||||
LoadResult,
|
||||
LocustStatEntry,
|
||||
aggregate_stats,
|
||||
median_seconds,
|
||||
percentile_seconds,
|
||||
read_errors,
|
||||
read_generator_warnings,
|
||||
)
|
||||
|
|
@ -18,12 +19,14 @@ _FAILURES_HEADER = "Method,Name,Error,Occurrences,First Seen,Last Seen\n"
|
|||
def _entry(
|
||||
*,
|
||||
num_requests: int,
|
||||
name: str = "/chat/completions",
|
||||
num_failures: int = 0,
|
||||
start_time: float = 1000.0,
|
||||
last_request_timestamp: float = 1010.0,
|
||||
response_times: dict[int, int] | None = None,
|
||||
) -> LocustStatEntry:
|
||||
return LocustStatEntry(
|
||||
name=name,
|
||||
num_requests=num_requests,
|
||||
num_failures=num_failures,
|
||||
start_time=start_time,
|
||||
|
|
@ -41,35 +44,48 @@ def _result(
|
|||
requests=10,
|
||||
failures=10,
|
||||
requests_per_second=1.0,
|
||||
median_response_seconds=0.05,
|
||||
p50_seconds=0.05,
|
||||
p90_seconds=0.08,
|
||||
p99_seconds=0.1,
|
||||
endpoints=(),
|
||||
errors=errors,
|
||||
generator_warnings=generator_warnings,
|
||||
)
|
||||
|
||||
|
||||
class TestSerialLatency:
|
||||
class TestPercentiles:
|
||||
def test_median_is_the_middle_sample_not_the_mean_a_slow_tail_would_drag(self) -> None:
|
||||
# Nine fast requests and one very slow one: the mean is 1.99s, the median is 20ms.
|
||||
entry = _entry(num_requests=10, response_times={20: 9, 20000: 1})
|
||||
|
||||
assert median_seconds([entry]) == 0.02
|
||||
assert percentile_seconds([entry], 0.5) == 0.02
|
||||
|
||||
def test_median_merges_the_histograms_of_every_stats_entry(self) -> None:
|
||||
def test_the_tail_percentiles_reach_the_slow_samples_the_median_hides(self) -> None:
|
||||
# 100 samples: 89 fast, 10 slow, 1 very slow. p50 sits in the fast bucket, p90 in the
|
||||
# slow one, and p99 lands on the single very slow sample.
|
||||
entry = _entry(num_requests=100, response_times={20: 89, 500: 10, 20000: 1})
|
||||
|
||||
assert percentile_seconds([entry], 0.5) == 0.02
|
||||
assert percentile_seconds([entry], 0.9) == 0.5
|
||||
assert percentile_seconds([entry], 0.99) == 0.5
|
||||
assert percentile_seconds([entry], 1.0) == 20.0
|
||||
|
||||
def test_percentiles_merge_the_histograms_of_every_stats_entry(self) -> None:
|
||||
# Per entry the median would be 10ms and 90ms; merged, the middle of the five samples is 90ms.
|
||||
entries = [
|
||||
_entry(num_requests=2, response_times={10: 2}),
|
||||
_entry(num_requests=3, response_times={90: 3}),
|
||||
]
|
||||
|
||||
assert median_seconds(entries) == 0.09
|
||||
assert percentile_seconds(entries, 0.5) == 0.09
|
||||
|
||||
def test_an_even_split_takes_the_lower_middle_sample_as_locust_itself_does(self) -> None:
|
||||
entry = _entry(num_requests=4, response_times={10: 2, 90: 2})
|
||||
|
||||
assert median_seconds([entry]) == 0.01
|
||||
assert percentile_seconds([entry], 0.5) == 0.01
|
||||
|
||||
def test_no_samples_reports_zero_rather_than_dividing_by_an_empty_histogram(self) -> None:
|
||||
assert median_seconds([]) == 0.0
|
||||
assert percentile_seconds([], 0.5) == 0.0
|
||||
|
||||
|
||||
class TestAggregate:
|
||||
|
|
@ -84,9 +100,20 @@ class TestAggregate:
|
|||
result = aggregate_stats([entry], (), ())
|
||||
|
||||
assert result.requests_per_second == 3.0
|
||||
assert result.median_response_seconds == 0.057
|
||||
assert result.p50_seconds == 0.057
|
||||
assert result.p99_seconds == 0.057
|
||||
assert result.failure_ratio == 0.0
|
||||
|
||||
def test_tail_percentiles_come_from_the_slow_end_of_the_histogram(self) -> None:
|
||||
entry = _entry(num_requests=100, response_times={20: 89, 500: 10, 3000: 1})
|
||||
|
||||
result = aggregate_stats([entry], (), ())
|
||||
|
||||
assert result.p50_seconds == 0.02
|
||||
assert result.p90_seconds == 0.5
|
||||
assert result.p99_seconds == 0.5
|
||||
assert result.latency_summary() == "p50 0.020s, p90 0.500s, p99 0.500s"
|
||||
|
||||
def test_throughput_spans_from_the_earliest_start_when_locust_reports_several_entries(self) -> None:
|
||||
entries = [
|
||||
_entry(num_requests=60, start_time=1000.0, last_request_timestamp=1030.0),
|
||||
|
|
@ -103,6 +130,49 @@ class TestAggregate:
|
|||
assert result.requests == 0
|
||||
assert result.requests_per_second == 0.0
|
||||
assert result.failure_ratio == 1.0
|
||||
assert result.endpoints == ()
|
||||
|
||||
|
||||
class TestPerEndpoint:
|
||||
def test_each_route_keeps_its_own_requests_failures_and_median(self) -> None:
|
||||
entries: Final = (
|
||||
_entry(name="/chat/completions", num_requests=100, response_times={20: 100}),
|
||||
_entry(name="/v1/messages", num_requests=40, num_failures=3, response_times={900: 40}),
|
||||
)
|
||||
|
||||
result: Final = aggregate_stats(entries, (), ())
|
||||
|
||||
assert tuple((one.name, one.requests, one.failures, one.p50_seconds) for one in result.endpoints) == (
|
||||
("/chat/completions", 100, 0, 0.02),
|
||||
("/v1/messages", 40, 3, 0.9),
|
||||
)
|
||||
|
||||
def test_several_stats_entries_for_one_route_fold_into_a_single_row(self) -> None:
|
||||
entries: Final = (
|
||||
_entry(name="/v1/messages", num_requests=10, response_times={30: 10}),
|
||||
_entry(name="/v1/messages", num_requests=30, num_failures=1, response_times={30: 30}),
|
||||
)
|
||||
|
||||
result: Final = aggregate_stats(entries, (), ())
|
||||
|
||||
assert tuple((one.name, one.requests, one.failures) for one in result.endpoints) == (("/v1/messages", 40, 1),)
|
||||
|
||||
def test_a_route_that_never_ran_is_absent_so_a_one_sided_run_cannot_pass_unnoticed(self) -> None:
|
||||
result: Final = aggregate_stats((_entry(name="/chat/completions", num_requests=10),), (), ())
|
||||
|
||||
assert tuple(one.name for one in result.endpoints) == ("/chat/completions",)
|
||||
|
||||
def test_the_summary_names_every_route_with_its_counts(self) -> None:
|
||||
entries: Final = (
|
||||
_entry(name="/chat/completions", num_requests=2, response_times={20: 2}),
|
||||
_entry(name="/v1/messages", num_requests=1, num_failures=1, response_times={500: 1}),
|
||||
)
|
||||
|
||||
result: Final = aggregate_stats(entries, (), ())
|
||||
|
||||
assert result.endpoint_summary() == (
|
||||
"/chat/completions 2 requests, 0 failures, p50 0.020s, /v1/messages 1 requests, 1 failures, p50 0.500s"
|
||||
)
|
||||
|
||||
|
||||
class TestErrorBreakdown:
|
||||
|
|
@ -133,8 +203,7 @@ class TestErrorBreakdown:
|
|||
def test_diagnosis_caps_the_list_and_says_how_many_it_left_out(self) -> None:
|
||||
result = _result(
|
||||
errors=tuple(
|
||||
LoadError(name="/chat/completions", error=f"error-{index}", occurrences=index)
|
||||
for index in range(1, 9)
|
||||
LoadError(name="/chat/completions", error=f"error-{index}", occurrences=index) for index in range(1, 9)
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
105
tests/e2e/load/test_phase_budget.py
Normal file
105
tests/e2e/load/test_phase_budget.py
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
from phase_budget import AbsoluteBudget, RatioBudget, violations
|
||||
|
||||
|
||||
def _budget(*, baseline: float, degraded: float, ceiling: float = 2.0) -> RatioBudget:
|
||||
return RatioBudget(
|
||||
name="p99 RSS", baseline=baseline, degraded=degraded, ratio_ceiling=ceiling, unit=" MB", decimals=0
|
||||
)
|
||||
|
||||
|
||||
class TestRatioBudget:
|
||||
def test_growth_within_the_ceiling_is_not_a_violation(self) -> None:
|
||||
assert _budget(baseline=100, degraded=199).violation() is None
|
||||
|
||||
def test_growth_exactly_at_the_ceiling_is_allowed(self) -> None:
|
||||
assert _budget(baseline=100, degraded=200).violation() is None
|
||||
|
||||
def test_growth_past_the_ceiling_reports_both_values_and_the_ratio(self) -> None:
|
||||
violation: Final = _budget(baseline=100, degraded=250).violation()
|
||||
|
||||
assert violation is not None
|
||||
assert "100 MB" in violation
|
||||
assert "250 MB" in violation
|
||||
assert "2.5x" in violation
|
||||
assert "2.0x allowed" in violation
|
||||
|
||||
def test_shrinking_is_never_a_violation(self) -> None:
|
||||
assert _budget(baseline=100, degraded=10).violation() is None
|
||||
|
||||
def test_a_missing_baseline_is_a_violation_rather_than_a_silent_pass(self) -> None:
|
||||
# The trap this guards: 0 as a baseline would make every ratio a division by zero, and
|
||||
# treating it as "no growth" would pass a run that measured nothing at all.
|
||||
violation: Final = _budget(baseline=0, degraded=4000).violation()
|
||||
|
||||
assert violation is not None
|
||||
assert "nothing to compare" in violation
|
||||
|
||||
def test_the_unit_and_decimals_carry_into_the_message(self) -> None:
|
||||
violation: Final = RatioBudget(
|
||||
name="p99 latency", baseline=0.16, degraded=9.5, ratio_ceiling=8.0, unit="s", decimals=3
|
||||
).violation()
|
||||
|
||||
assert violation is not None
|
||||
assert "0.160s" in violation
|
||||
assert "9.500s" in violation
|
||||
|
||||
|
||||
class TestAbsoluteBudget:
|
||||
def test_a_value_under_the_ceiling_is_not_a_violation(self) -> None:
|
||||
assert AbsoluteBudget(name="p99 latency", measured=1.2, ceiling=5.0, unit="s", decimals=3).violation() is None
|
||||
|
||||
def test_a_value_exactly_at_the_ceiling_is_allowed(self) -> None:
|
||||
assert AbsoluteBudget(name="p99 latency", measured=5.0, ceiling=5.0, unit="s", decimals=3).violation() is None
|
||||
|
||||
def test_a_value_past_the_ceiling_reports_the_measurement_and_the_ceiling(self) -> None:
|
||||
violation: Final = AbsoluteBudget(
|
||||
name="p99 latency", measured=9.5, ceiling=5.0, unit="s", decimals=3
|
||||
).violation()
|
||||
|
||||
assert violation is not None
|
||||
assert "9.500s" in violation
|
||||
assert "5.000s allowed" in violation
|
||||
|
||||
def test_a_flat_ceiling_fails_a_degraded_phase_that_is_cheaper_than_its_baseline(self) -> None:
|
||||
# The whole reason this shape exists: once the breaker opens, requests skip Redis instead
|
||||
# of waiting on its socket timeout, so the chaos phase can measure faster than the healthy
|
||||
# one. A ratio against that baseline passes; the user still waited 9.5s.
|
||||
assert _budget(baseline=20.0, degraded=9.5, ceiling=2.0).violation() is None
|
||||
assert AbsoluteBudget(name="p99 latency", measured=9.5, ceiling=5.0, unit="s").violation() is not None
|
||||
|
||||
def test_a_zero_measurement_is_not_a_violation(self) -> None:
|
||||
assert AbsoluteBudget(name="log bytes per request", measured=0, ceiling=12_000, unit=" B").violation() is None
|
||||
|
||||
|
||||
class TestViolations:
|
||||
def test_every_blown_budget_is_reported_not_just_the_first(self) -> None:
|
||||
blown: Final = violations(
|
||||
(
|
||||
_budget(baseline=100, degraded=500),
|
||||
_budget(baseline=100, degraded=120),
|
||||
RatioBudget(name="CPU per request", baseline=10, degraded=90, ratio_ceiling=6.0, unit=" ms"),
|
||||
)
|
||||
)
|
||||
|
||||
assert len(blown) == 2
|
||||
assert blown[0].startswith("p99 RSS")
|
||||
assert blown[1].startswith("CPU per request")
|
||||
|
||||
def test_both_budget_shapes_report_together(self) -> None:
|
||||
blown: Final = violations(
|
||||
(
|
||||
_budget(baseline=100, degraded=500),
|
||||
AbsoluteBudget(name="p99 latency", measured=9.5, ceiling=5.0, unit="s", decimals=3),
|
||||
)
|
||||
)
|
||||
|
||||
assert len(blown) == 2
|
||||
assert blown[0].startswith("p99 RSS")
|
||||
assert blown[1].startswith("p99 latency")
|
||||
|
||||
def test_a_run_inside_every_budget_reports_nothing(self) -> None:
|
||||
assert violations((_budget(baseline=100, degraded=150),)) == ()
|
||||
71
tests/e2e/load/test_proxy_usage.py
Normal file
71
tests/e2e/load/test_proxy_usage.py
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
from proxy_usage import UsageSample, UsageWindow
|
||||
|
||||
_MB: Final = 2**20
|
||||
|
||||
|
||||
def _window(*points: tuple[float, int, float]) -> UsageWindow:
|
||||
return UsageWindow(
|
||||
samples=tuple(
|
||||
UsageSample(elapsed_seconds=elapsed, rss_bytes=rss, cpu_seconds=cpu) for elapsed, rss, cpu in points
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class TestRssPercentiles:
|
||||
def test_the_tail_percentiles_reach_the_peak_the_median_hides(self) -> None:
|
||||
# 100 one-second samples: 89 flat, 10 elevated, 1 spike. The median stays flat, p90 sees the
|
||||
# elevated plateau, and only the max reaches the spike.
|
||||
window: Final = _window(
|
||||
*((float(i), 100 * _MB, float(i)) for i in range(89)),
|
||||
*((float(89 + i), 300 * _MB, float(89 + i)) for i in range(10)),
|
||||
(99.0, 900 * _MB, 99.0),
|
||||
)
|
||||
|
||||
assert window.rss_percentile(0.5) == 100 * _MB
|
||||
assert window.rss_percentile(0.9) == 300 * _MB
|
||||
assert window.rss_percentile(0.99) == 300 * _MB
|
||||
assert window.rss_percentile(1.0) == 900 * _MB
|
||||
|
||||
def test_an_empty_window_reports_zero_rather_than_indexing_nothing(self) -> None:
|
||||
assert _window().rss_percentile(0.5) == 0
|
||||
|
||||
|
||||
class TestCpuUtilization:
|
||||
def test_utilization_is_the_counter_delta_over_the_interval_not_the_counter_itself(self) -> None:
|
||||
# The counter climbs 0.5 CPU seconds per second, then 4.0 per second: half a core, then four.
|
||||
window: Final = _window((0.0, _MB, 0.0), (1.0, _MB, 0.5), (2.0, _MB, 1.0), (3.0, _MB, 5.0))
|
||||
|
||||
p50, p90, p99 = window.cpu_utilization_percentiles()
|
||||
|
||||
assert (p50, p90, p99) == (0.5, 4.0, 4.0)
|
||||
assert window.cpu_seconds_consumed() == 5.0
|
||||
|
||||
def test_a_single_sample_has_no_interval_and_reports_zero(self) -> None:
|
||||
window: Final = _window((0.0, _MB, 3.0))
|
||||
|
||||
assert window.cpu_utilization_percentiles() == (0.0, 0.0, 0.0)
|
||||
assert window.cpu_seconds_consumed() == 0.0
|
||||
|
||||
def test_cost_per_request_separates_runs_that_cores_busy_reports_identically(self) -> None:
|
||||
# Both windows pin 4 cores for 10 seconds, so utilization cannot tell them apart. The
|
||||
# second one served a tenth of the traffic for the same CPU, which is the regression shape.
|
||||
window: Final = _window(*((float(i), _MB, 4.0 * i) for i in range(11)))
|
||||
|
||||
assert window.cpu_utilization_percentiles()[0] == 4.0
|
||||
assert window.cpu_seconds_per_request(4000) == 0.01
|
||||
assert window.cpu_seconds_per_request(400) == 0.1
|
||||
|
||||
def test_no_requests_reports_zero_cost_rather_than_dividing_by_zero(self) -> None:
|
||||
assert _window((0.0, _MB, 0.0), (1.0, _MB, 1.0)).cpu_seconds_per_request(0) == 0.0
|
||||
|
||||
def test_summary_reports_every_percentile_in_human_units(self) -> None:
|
||||
window: Final = _window((0.0, 200 * _MB, 0.0), (1.0, 200 * _MB, 1.5), (2.0, 200 * _MB, 3.0))
|
||||
|
||||
assert window.summary() == (
|
||||
"RSS p50 200 MB, p90 200 MB, p99 200 MB; "
|
||||
"CPU cores busy p50 1.50, p90 1.50, p99 1.50; 3.0 CPU seconds consumed"
|
||||
)
|
||||
454
tests/e2e/load/test_redis_chaos_e2e.py
Normal file
454
tests/e2e/load/test_redis_chaos_e2e.py
Normal file
|
|
@ -0,0 +1,454 @@
|
|||
"""Live e2e: the proxy under load keeps serving every request while Redis is down entirely.
|
||||
|
||||
Runs against a proxy booted from tests/e2e/gateway/redis_chaos_ci_config.yml, which points
|
||||
cache_params at a real Redis with litellm's default socket_timeout. That one client backs all
|
||||
three Redis touchpoints on the request path: the virtual-key auth cache, the response cache,
|
||||
and the cross-pod spend counter the cost-tracking callback awaits.
|
||||
|
||||
The load runs in two phases against one model group of three mock deployments. The two at
|
||||
order 1 raise InternalServerError and the one at order 2 serves, so every request burns its
|
||||
retries on the failing pair (a 500 is retryable, so retries keep re-picking inside the lowest
|
||||
order) and the router's order-based fallback then re-targets order 2. Every request is expected
|
||||
to succeed, and each one carries retry breadcrumbs into cost tracking.
|
||||
|
||||
Traffic is split round robin between /chat/completions and /v1/messages, one endpoint per
|
||||
simulated user: the Redis touchpoints and the cost-tracking callback are shared by both, but
|
||||
the Anthropic Messages route reaches them through its own request path, so a regression that
|
||||
only shows up there would not surface from chat completions alone.
|
||||
|
||||
Phase A is a baseline with Redis healthy; phase B holds Redis in CLIENT PAUSE ALL for the
|
||||
length of the phase, simulating Redis being down outright rather than merely slow to write.
|
||||
Every touchpoint times out: the auth cache read falls back to Postgres, the response cache
|
||||
read and write both fail, and the spend counter increment times out and the callback
|
||||
stringifies the request metadata, breadcrumbs included, into a failed-tracking alert. On
|
||||
v1.100.0 that string doubled per request until the worker hung (LIT-6780), which is what the
|
||||
per-phase RSS, CPU, and log-bytes budgets are here to catch.
|
||||
|
||||
Needs the proxy on the same host, since RSS and CPU come from psutil on its process tree:
|
||||
a multi-worker proxy serves /metrics from the prometheus multiprocess collector, which drops
|
||||
the process collector's memory and CPU series. Log bytes are read from the file the proxy's
|
||||
stdout/stderr was redirected to, so the same host requirement covers that too. Deselected
|
||||
unless E2E_REDIS_CHAOS is set.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from itertools import pairwise
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import redis
|
||||
from e2e_config import PROXY_BASE_URL, unique_marker
|
||||
from e2e_http import NoBody
|
||||
from lifecycle import ResourceManager
|
||||
from load_client import LoadClient
|
||||
from locust_load import LoadResult, run_gateway_load
|
||||
from models import KeyGenerateBody, LiteLLMParamsBody
|
||||
from phase_budget import AbsoluteBudget, Budget, RatioBudget, violations
|
||||
from proxy_client import ProxyClient
|
||||
from proxy_usage import ProxyUsageSampler, UsageWindow
|
||||
|
||||
pytestmark: Final = pytest.mark.e2e
|
||||
|
||||
MODEL_GROUP: Final = f"redis-chaos-fable-{unique_marker()}"
|
||||
MOCK_MODEL: Final = "anthropic/claude-fable-5-1"
|
||||
FAILING_DEPLOYMENTS: Final = 2
|
||||
SERVING_DEPLOYMENTS: Final = 1
|
||||
FAILING_ORDER: Final = 1
|
||||
SERVING_ORDER: Final = 2
|
||||
KEY_POOL_SIZE: Final = 8
|
||||
LOAD_ENDPOINTS: Final = ("/chat/completions", "/v1/messages")
|
||||
LOCUST_USERS: Final = 50
|
||||
LOCUST_SPAWN_RATE: Final = 50.0
|
||||
BASELINE_SECONDS: Final = 60.0
|
||||
CHAOS_SECONDS: Final = 90.0
|
||||
REDIS_PAUSE_MS: Final = int(CHAOS_SECONDS * 1000)
|
||||
|
||||
# RSS and CPU are budgeted as a multiple of the same metric in the baseline phase, because both
|
||||
# are machine-shaped: RSS scales with worker count and CPU with core count, so a number
|
||||
# calibrated on one runner means nothing on another. RSS moved 0.91x-1.40x across three otherwise
|
||||
# identical local runs, so it stays loose; CPU per request held steady at 1.33x-1.36x across the
|
||||
# same runs, so it sits close to what is actually measured. That makes CPU the likeliest of these
|
||||
# to flake first on a runner whose core count shifts how much of baseline CPU is fixed per-request
|
||||
# work: loosen it rather than widening the others if a CI run trips it without a real cause.
|
||||
CHAOS_RSS_RATIO_CEILING: Final = 2.0
|
||||
CHAOS_CPU_PER_REQUEST_RATIO_CEILING: Final = 2.0
|
||||
|
||||
# Latency and log volume get flat ceilings instead, because a ratio cannot bound either one. Once
|
||||
# the breaker opens, a request skips Redis rather than waiting on its socket timeout, so the chaos
|
||||
# phase can come in faster than baseline (local runs measured p90 at 0.61x) and a ratio passes on a
|
||||
# phase that was never slow. What a user actually cares about is the wall-clock number, which these
|
||||
# hold directly. Calibrated from local runs whose worst chaos phase was p50 0.19s, p90 0.23s, p99
|
||||
# 0.69s and 3.5 KB of log per request, with several times that left as slack for a shared CI runner.
|
||||
CHAOS_P50_LATENCY_CEILING_SECONDS: Final = 1.0
|
||||
CHAOS_P90_LATENCY_CEILING_SECONDS: Final = 2.0
|
||||
CHAOS_P99_LATENCY_CEILING_SECONDS: Final = 3.0
|
||||
CHAOS_LOG_BYTES_PER_REQUEST_CEILING: Final = 10_000.0
|
||||
|
||||
DRAIN_TIMEOUT_SECONDS: Final = 30.0
|
||||
DRAIN_POLL_SECONDS: Final = 1.0
|
||||
|
||||
TIMEOUT_FAILURES_RE: Final = re.compile(
|
||||
r'^litellm_redis_circuit_breaker_failures_total\{failure_class="timeout"\} ([0-9.e+]+)$', re.M
|
||||
)
|
||||
# The state gauge carries a pid label under the multiprocess collector, one series per worker,
|
||||
# so this matches any label order rather than a bare {state="open"} that never appears.
|
||||
BREAKER_OPEN_RE: Final = re.compile(
|
||||
r'^litellm_redis_circuit_breaker_state\{[^}]*state="open"[^}]*\} ([0-9.e+]+)$', re.M
|
||||
)
|
||||
BREAKER_TRANSITIONS_RE: Final = re.compile(
|
||||
r'^litellm_redis_circuit_breaker_transitions_total\{state="[a-z_]+"\} ([0-9.e+]+)$', re.M
|
||||
)
|
||||
|
||||
|
||||
def _deployment_metric_re(name: str, model_ids: tuple[str, ...]) -> re.Pattern[str]:
|
||||
"""A per-deployment counter, narrowed to the deployments one run registered, so traffic
|
||||
anything else sends the same proxy during the run cannot pad the retry count."""
|
||||
ids: Final = "|".join(re.escape(model_id) for model_id in model_ids)
|
||||
return re.compile(rf'^litellm_{name}\{{[^}}]*model_id="(?:{ids})"[^}}]*\}} ([0-9.e+]+)$', re.M)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Phase:
|
||||
"""One load phase's traffic and what the proxy's process tree did during it."""
|
||||
|
||||
name: str
|
||||
load: LoadResult
|
||||
usage: UsageWindow
|
||||
redis_timeouts: float
|
||||
log_bytes: int
|
||||
|
||||
@property
|
||||
def timeouts_per_request(self) -> float:
|
||||
return self.redis_timeouts / self.load.requests if self.load.requests else 0.0
|
||||
|
||||
@property
|
||||
def cpu_seconds_per_request(self) -> float:
|
||||
return self.usage.cpu_seconds_per_request(self.load.requests)
|
||||
|
||||
@property
|
||||
def log_bytes_per_request(self) -> float:
|
||||
return self.log_bytes / self.load.requests if self.load.requests else 0.0
|
||||
|
||||
def report(self) -> str:
|
||||
return (
|
||||
f"{self.name}: {self.load.requests} requests, {self.load.failures} failures, "
|
||||
f"{self.load.requests_per_second:.0f} rps, {self.load.latency_summary()}; {self.usage.summary()}; "
|
||||
f"{self.cpu_seconds_per_request * 1000:.1f} ms CPU per request; "
|
||||
f"{self.log_bytes_per_request:.0f} log bytes per request; "
|
||||
f"{self.timeouts_per_request:.2f} Redis timeouts per request; "
|
||||
f"by endpoint: {self.load.endpoint_summary()}"
|
||||
)
|
||||
|
||||
|
||||
def _failing_params() -> LiteLLMParamsBody:
|
||||
return LiteLLMParamsBody(
|
||||
model=MOCK_MODEL,
|
||||
api_key="sk-redis-chaos-not-used",
|
||||
mock_response="litellm.InternalServerError",
|
||||
order=FAILING_ORDER,
|
||||
)
|
||||
|
||||
|
||||
def _serving_params() -> LiteLLMParamsBody:
|
||||
return LiteLLMParamsBody(
|
||||
model=MOCK_MODEL,
|
||||
api_key="sk-redis-chaos-not-used",
|
||||
mock_response="redis chaos ok",
|
||||
order=SERVING_ORDER,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def proxy_pid() -> int:
|
||||
"""The proxy's PID, which the workflow exports after starting it.
|
||||
|
||||
Required rather than discovered: picking a process out of the table by name would be
|
||||
ambiguous on a developer machine running more than one proxy.
|
||||
"""
|
||||
pid: Final = os.environ.get("E2E_PROXY_PID")
|
||||
assert pid and pid.isdigit(), (
|
||||
"E2E_PROXY_PID must hold the PID of the proxy under test; RSS and CPU are read from "
|
||||
"its process tree because a multi-worker proxy does not report them on /metrics"
|
||||
)
|
||||
return int(pid)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def proxy_log() -> Path:
|
||||
"""Path to the proxy's stdout/stderr log, which the workflow captures to a file.
|
||||
|
||||
Required rather than discovered for the same reason as proxy_pid: a developer machine may
|
||||
have more than one proxy log around.
|
||||
"""
|
||||
path: Final = os.environ.get("E2E_PROXY_LOG")
|
||||
assert path, "E2E_PROXY_LOG must hold the path the proxy's stdout/stderr was redirected to"
|
||||
return Path(path)
|
||||
|
||||
|
||||
def _log_bytes(path: Path) -> int:
|
||||
return path.stat().st_size
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def redis_control() -> Iterator[redis.Redis[bytes]]:
|
||||
"""A control connection to the proxy's Redis, which unpauses it in teardown as a safety net.
|
||||
|
||||
CLIENT PAUSE ALL freezes every connection including this one, so REDIS_PAUSE_MS is sized
|
||||
to the chaos phase: by the time teardown runs, the pause has
|
||||
already lapsed on its own and CLIENT UNPAUSE here returns immediately. It only actually
|
||||
waits out a lapsed pause if the chaos phase itself overran that duration.
|
||||
"""
|
||||
host: Final = os.environ.get("REDIS_HOST")
|
||||
port: Final = os.environ.get("REDIS_PORT")
|
||||
assert host and port, "REDIS_HOST and REDIS_PORT must name the Redis the proxy under test uses"
|
||||
control: Final = redis.Redis(host=host, port=int(port), socket_timeout=5)
|
||||
try:
|
||||
yield control
|
||||
finally:
|
||||
control.client_unpause() # pyright: ignore[reportUnknownMemberType] # redis-py stubs return Any
|
||||
control.close()
|
||||
|
||||
|
||||
def _scrape(proxy: ProxyClient) -> str:
|
||||
"""One /metrics body, read once per checkpoint so every counter comes from the same instant."""
|
||||
scrape: Final = proxy.probe("/metrics", params=NoBody())
|
||||
assert scrape.status_code == 200, (
|
||||
f"/metrics did not answer ({scrape.status_code}: {scrape.body[:200]}), so no counter can be read; "
|
||||
f"a silent 0 here would turn every before-and-after difference negative"
|
||||
)
|
||||
return scrape.body
|
||||
|
||||
|
||||
def _metric(scrape: str, pattern: re.Pattern[str]) -> float:
|
||||
return sum(float(match.group(1)) for match in pattern.finditer(scrape))
|
||||
|
||||
|
||||
def _scrape_after_drain(proxy: ProxyClient, pattern: re.Pattern[str]) -> str:
|
||||
"""A /metrics body taken once `pattern`'s count has stopped moving.
|
||||
|
||||
`set_llm_deployment_failure_metrics` runs from the async logging callback queue, so a load
|
||||
generator that just stopped sending traffic can still have thousands of failure increments
|
||||
in flight, and a scrape taken the instant load stops undercounts them. Settling on the
|
||||
counter rather than sleeping a fixed duration keeps the wait proportional to how backed up
|
||||
the queue actually is.
|
||||
"""
|
||||
deadline: Final = time.monotonic() + DRAIN_TIMEOUT_SECONDS
|
||||
|
||||
def scrapes() -> Iterator[str]:
|
||||
yield _scrape(proxy)
|
||||
while time.monotonic() < deadline:
|
||||
time.sleep(DRAIN_POLL_SECONDS)
|
||||
yield _scrape(proxy)
|
||||
|
||||
settled: Final = next(
|
||||
(later for earlier, later in pairwise(scrapes()) if _metric(earlier, pattern) == _metric(later, pattern)),
|
||||
None,
|
||||
)
|
||||
return settled if settled is not None else _scrape(proxy)
|
||||
|
||||
|
||||
def _register_deployments(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, ...]:
|
||||
"""The model ids this run registered, which scope its per-deployment metric reads."""
|
||||
params: Final = (
|
||||
*(_failing_params() for _ in range(FAILING_DEPLOYMENTS)),
|
||||
*(_serving_params() for _ in range(SERVING_DEPLOYMENTS)),
|
||||
)
|
||||
model_ids: Final = tuple(proxy.create_model(MODEL_GROUP, one) for one in params)
|
||||
for model_id in model_ids:
|
||||
resources.defer(lambda doomed=model_id: proxy.delete_model(doomed))
|
||||
return model_ids
|
||||
|
||||
|
||||
def _generate_key_pool(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, ...]:
|
||||
"""A pool of virtual keys so auth and budget lookups are not one permanently warm
|
||||
cache entry; each locust user picks one, so Redis auth reads actually happen."""
|
||||
keys: Final = tuple(
|
||||
proxy.generate_key(
|
||||
KeyGenerateBody(models=[MODEL_GROUP], key_alias=f"e2e-redis-chaos-{unique_marker()}-{index}")
|
||||
)
|
||||
for index in range(KEY_POOL_SIZE)
|
||||
)
|
||||
for key in keys:
|
||||
resources.defer(lambda doomed=key: proxy.delete_key(doomed))
|
||||
return keys
|
||||
|
||||
|
||||
def _drive(keys: tuple[str, ...], seconds: float) -> LoadResult:
|
||||
return run_gateway_load(
|
||||
base_url=PROXY_BASE_URL,
|
||||
api_keys=keys,
|
||||
model=MODEL_GROUP,
|
||||
endpoints=LOAD_ENDPOINTS,
|
||||
users=LOCUST_USERS,
|
||||
spawn_rate=LOCUST_SPAWN_RATE,
|
||||
duration_seconds=seconds,
|
||||
)
|
||||
|
||||
|
||||
def _latency_budget(percentile: str, measured: float, ceiling: float) -> Budget:
|
||||
return AbsoluteBudget(name=f"{percentile} latency", measured=measured, ceiling=ceiling, unit="s", decimals=3)
|
||||
|
||||
|
||||
def _rss_budget(percentile: str, baseline: UsageWindow, degraded: UsageWindow, fraction: float) -> Budget:
|
||||
return RatioBudget(
|
||||
name=f"{percentile} RSS",
|
||||
baseline=baseline.rss_percentile(fraction) / 2**20,
|
||||
degraded=degraded.rss_percentile(fraction) / 2**20,
|
||||
ratio_ceiling=CHAOS_RSS_RATIO_CEILING,
|
||||
unit=" MB",
|
||||
decimals=0,
|
||||
)
|
||||
|
||||
|
||||
def _chaos_budgets(baseline: Phase, chaos: Phase) -> tuple[Budget, ...]:
|
||||
"""What a Redis outage is allowed to cost.
|
||||
|
||||
Every request still succeeding is the headline assertion, but a proxy can answer every
|
||||
request while leaking: the v1.100.0 regression (LIT-6780) served traffic the whole way up
|
||||
to a 61 GB worker. These bound the cost of serving it. RSS and CPU are bounded against the
|
||||
same run's healthy phase, latency and log bytes against a flat ceiling; see phase_budget
|
||||
for why the two kinds of metric cannot share one shape.
|
||||
|
||||
Latency and RSS are budgeted at p50, p90 and p99 so a regression that only shows up in the
|
||||
tail (or only in the median) cannot hide behind the other. RSS gets the tightest bound: the
|
||||
failure path has no business allocating more per request. CPU and log bytes are each budgeted
|
||||
once, as an amount per request rather than per percentile: cores-busy saturates at the worker
|
||||
count under load, so its percentiles read the same whether a request costs 10 ms of CPU or
|
||||
40, and cannot budget anything; per-request is the figure that actually moves. Log bytes
|
||||
isolates the cost of the failed-tracking alert's own noisy error handling from the CPU it
|
||||
burns doing useful retry work, since the two would otherwise be indistinguishable in one
|
||||
CPU number.
|
||||
"""
|
||||
return (
|
||||
_latency_budget("p50", chaos.load.p50_seconds, CHAOS_P50_LATENCY_CEILING_SECONDS),
|
||||
_latency_budget("p90", chaos.load.p90_seconds, CHAOS_P90_LATENCY_CEILING_SECONDS),
|
||||
_latency_budget("p99", chaos.load.p99_seconds, CHAOS_P99_LATENCY_CEILING_SECONDS),
|
||||
_rss_budget("p50", baseline.usage, chaos.usage, 0.5),
|
||||
_rss_budget("p90", baseline.usage, chaos.usage, 0.9),
|
||||
_rss_budget("p99", baseline.usage, chaos.usage, 0.99),
|
||||
RatioBudget(
|
||||
name="CPU per request",
|
||||
baseline=baseline.cpu_seconds_per_request * 1000,
|
||||
degraded=chaos.cpu_seconds_per_request * 1000,
|
||||
ratio_ceiling=CHAOS_CPU_PER_REQUEST_RATIO_CEILING,
|
||||
unit=" ms",
|
||||
),
|
||||
AbsoluteBudget(
|
||||
name="log bytes per request",
|
||||
measured=chaos.log_bytes_per_request,
|
||||
ceiling=CHAOS_LOG_BYTES_PER_REQUEST_CEILING,
|
||||
unit=" B",
|
||||
decimals=0,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.redis_chaos
|
||||
class TestRedisChaos:
|
||||
@pytest.mark.covers(
|
||||
"reliability.circuit_breaker.redis_timeout.stays_responsive",
|
||||
exercised_on=("chat_completions", "messages"),
|
||||
)
|
||||
def test_load_survives_redis_being_down(
|
||||
self,
|
||||
client: LoadClient,
|
||||
resources: ResourceManager,
|
||||
proxy_pid: int,
|
||||
proxy_log: Path,
|
||||
redis_control: redis.Redis[bytes],
|
||||
) -> None:
|
||||
proxy: Final = client.proxy
|
||||
model_ids: Final = _register_deployments(proxy, resources)
|
||||
keys: Final = _generate_key_pool(proxy, resources)
|
||||
|
||||
retries_re: Final = _deployment_metric_re("deployment_failure_responses_total", model_ids)
|
||||
cooldown_re: Final = _deployment_metric_re("deployment_cooled_down_total", model_ids)
|
||||
|
||||
at_start: Final = _scrape(proxy)
|
||||
log_at_start: Final = _log_bytes(proxy_log)
|
||||
|
||||
with ProxyUsageSampler(proxy_pid) as sampler:
|
||||
baseline_load: Final = _drive(keys, BASELINE_SECONDS)
|
||||
baseline_usage: Final = sampler.split()
|
||||
after_baseline: Final = _scrape(proxy)
|
||||
log_after_baseline: Final = _log_bytes(proxy_log)
|
||||
|
||||
redis_control.client_pause(REDIS_PAUSE_MS, all=True) # pyright: ignore[reportUnknownMemberType] # redis-py stubs return Any
|
||||
chaos_load: Final = _drive(keys, CHAOS_SECONDS)
|
||||
chaos_usage: Final = sampler.split()
|
||||
at_end: Final = _scrape_after_drain(proxy, retries_re)
|
||||
log_at_end: Final = _log_bytes(proxy_log)
|
||||
|
||||
baseline: Final = Phase(
|
||||
name="baseline",
|
||||
load=baseline_load,
|
||||
usage=baseline_usage,
|
||||
redis_timeouts=_metric(after_baseline, TIMEOUT_FAILURES_RE) - _metric(at_start, TIMEOUT_FAILURES_RE),
|
||||
log_bytes=log_after_baseline - log_at_start,
|
||||
)
|
||||
chaos: Final = Phase(
|
||||
name="chaos",
|
||||
load=chaos_load,
|
||||
usage=chaos_usage,
|
||||
redis_timeouts=_metric(at_end, TIMEOUT_FAILURES_RE) - _metric(after_baseline, TIMEOUT_FAILURES_RE),
|
||||
log_bytes=log_at_end - log_after_baseline,
|
||||
)
|
||||
report: Final = f"{baseline.report()} | {chaos.report()}"
|
||||
|
||||
for phase in (baseline, chaos):
|
||||
assert phase.load.requests > 0, (
|
||||
f"{phase.name} drove no traffic at all, so it proved nothing: {phase.load.diagnosis()}. {report}"
|
||||
)
|
||||
assert frozenset(endpoint.name for endpoint in phase.load.endpoints) == frozenset(LOAD_ENDPOINTS), (
|
||||
f"{phase.name} drove {tuple(endpoint.name for endpoint in phase.load.endpoints)} rather than every "
|
||||
f"endpoint in {LOAD_ENDPOINTS}; the round robin hands one endpoint to each simulated user, so a "
|
||||
f"missing one means a route never ran and its request path was never exercised. {report}"
|
||||
)
|
||||
assert phase.load.failures == 0, (
|
||||
f"{phase.name} had {phase.load.failures} of {phase.load.requests} requests fail. Every request "
|
||||
f"must succeed: the failing deployments sit at order {FAILING_ORDER} and the serving one at order "
|
||||
f"{SERVING_ORDER}, so once the retries on order {FAILING_ORDER} are spent the order-based fallback "
|
||||
f"lands on the serving deployment. Failures mean it was cooled down, the fallback did not run, or "
|
||||
f"a Redis failure reached the response path. {phase.load.diagnosis()}. {report}"
|
||||
)
|
||||
|
||||
cooldowns: Final = _metric(at_end, cooldown_re) - _metric(at_start, cooldown_re)
|
||||
assert cooldowns == 0, (
|
||||
f"{cooldowns:.0f} deployments were cooled down during the run; the failing deployments are supposed "
|
||||
f"to stay in rotation so every request keeps exercising the retry path. {report}"
|
||||
)
|
||||
|
||||
retries: Final = _metric(at_end, retries_re) - _metric(at_start, retries_re)
|
||||
assert retries >= baseline.load.requests + chaos.load.requests, (
|
||||
f"only {retries:.0f} deployment failures were counted across "
|
||||
f"{baseline.load.requests + chaos.load.requests} requests; the mock deployments did not fail, so no "
|
||||
f"request carried retry breadcrumbs into cost tracking and the regression path was never entered. "
|
||||
f"{report}"
|
||||
)
|
||||
|
||||
transitions: Final = _metric(at_end, BREAKER_TRANSITIONS_RE) - _metric(after_baseline, BREAKER_TRANSITIONS_RE)
|
||||
breaker_open: Final = _metric(at_end, BREAKER_OPEN_RE) >= 1
|
||||
assert transitions >= 1 or breaker_open, (
|
||||
f"pausing Redis produced no circuit breaker state transitions and it ended closed; nothing on the "
|
||||
f"request path ever saw Redis fail, so this run proved nothing. {report}"
|
||||
)
|
||||
|
||||
blown: Final = violations(_chaos_budgets(baseline, chaos))
|
||||
assert not blown, (
|
||||
f"pausing Redis cost the proxy more than a Redis outage is allowed to: {'; '.join(blown)}. {report}"
|
||||
)
|
||||
|
||||
rows: Final = proxy.poll_logs_for_key(keys[0], min_rows=1)
|
||||
assert rows, (
|
||||
f"no spend rows landed for the first key in the pool; a Redis outage must not cost the proxy its "
|
||||
f"spend logs, which are written to Postgres through a queue rather than through Redis. {report}"
|
||||
)
|
||||
|
||||
print(f"\nredis chaos load: {report}") # noqa: T201 # the numbers this test exists to report, read off the CI log
|
||||
|
|
@ -816,10 +816,11 @@ class LiteLLMParamsBody(BaseModel):
|
|||
auto_router_default_model: str | None = None
|
||||
auto_router_embedding_model: str | None = None
|
||||
tags: list[str] | None = None
|
||||
mock_response: str | None = None
|
||||
mock_response: str | list[float] | None = None
|
||||
timeout: float | None = None
|
||||
tpm: int | None = None
|
||||
weight: int | None = None
|
||||
order: int | None = None
|
||||
|
||||
|
||||
ModelMode = Literal["batch", "realtime", "image_generation"]
|
||||
|
|
|
|||
|
|
@ -9,3 +9,4 @@ markers =
|
|||
load: heavy throughput/load test; collected last so it never perturbs latency-sensitive suites
|
||||
weekly: real-provider anomaly load test that spends real money; deselected unless E2E_WEEKLY_ANOMALY is set
|
||||
managed_files: needs a proxy running with require_managed_files enabled; deselected unless E2E_MANAGED_FILES_STACK is set
|
||||
redis_chaos: load test that pauses the proxy's Redis outright mid-run; needs a proxy booted from gateway/redis_chaos_ci_config.yml on the same host, and is deselected unless E2E_REDIS_CHAOS is set
|
||||
|
|
|
|||
|
|
@ -13,10 +13,7 @@ from __future__ import annotations
|
|||
from collections.abc import Iterator
|
||||
|
||||
import pytest
|
||||
from requests import RequestException
|
||||
|
||||
from complexity_router_client import ComplexityRouterClient, build_client
|
||||
from proxy_client import ProxyClient
|
||||
from e2e_http import NoBody, Success
|
||||
from lifecycle import ResourceManager
|
||||
from models import (
|
||||
|
|
@ -26,6 +23,8 @@ from models import (
|
|||
LiteLLMParamsBody,
|
||||
ModelsListResponse,
|
||||
)
|
||||
from proxy_client import ProxyClient
|
||||
from requests import RequestException
|
||||
|
||||
ROUTER_MODEL = "complexity-smart-router"
|
||||
ROUTER_PARAMS = LiteLLMParamsBody(
|
||||
|
|
@ -120,8 +119,6 @@ def _ensure_complexity_smart_router( # pyright: ignore[reportUnusedFunction] #
|
|||
@pytest.fixture
|
||||
def complexity_key(resources: ResourceManager, client: ComplexityRouterClient) -> str:
|
||||
"""Per-test key allowed to call the complexity router and its tier backends."""
|
||||
key = client.proxy.generate_key(
|
||||
KeyGenerateBody(models=ROUTER_KEY_MODELS, user_id="e2e-complexity-router")
|
||||
)
|
||||
key = client.proxy.generate_key(KeyGenerateBody(models=ROUTER_KEY_MODELS, user_id="e2e-complexity-router"))
|
||||
resources.defer(lambda: client.proxy.delete_key(key))
|
||||
return key
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -7,7 +8,8 @@ import pytest
|
|||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.caching.redis_cache import RedisCache, _redis_circuit_breaker_guard, _redis_circuit_breaker_guard_sync
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -576,3 +578,129 @@ async def test_dual_cache_late_attach_redis_wires_writes_and_ttl_async():
|
|||
assert mock_redis.async_set_cache.call_args[0][:2] == (key_after, val_after)
|
||||
|
||||
assert in_memory.get_cache(key_after) == val_after
|
||||
|
||||
|
||||
class _OpenBreakerRedis:
|
||||
def __init__(self) -> None:
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker
|
||||
|
||||
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
||||
for _ in range(3):
|
||||
self._circuit_breaker.record_failure()
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_batch_get_cache(self, key_list, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_set_cache(self, key, value, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_set_cache_pipeline(self, cache_list, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_increment_pipeline(self, increment_list, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_increment(self, key, value, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def get_cache(self, key, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def batch_get_cache(self, key_list, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"call",
|
||||
[
|
||||
lambda cache: cache.async_get_cache("k"),
|
||||
lambda cache: cache.async_batch_get_cache(["k1", "k2"]),
|
||||
lambda cache: cache.async_set_cache("k", "v"),
|
||||
lambda cache: cache.async_set_cache_pipeline([("k", "v")]),
|
||||
lambda cache: cache.async_increment_cache_pipeline(
|
||||
increment_list=[RedisPipelineIncrementOperation(key="k", increment_value=1.0, ttl=60)]
|
||||
),
|
||||
lambda cache: cache.async_increment_cache("k", 1.0),
|
||||
],
|
||||
ids=["get", "batch_get", "set", "set_pipeline", "increment_pipeline", "increment"],
|
||||
)
|
||||
async def test_an_open_circuit_breaker_is_not_an_error_per_request(caplog, call):
|
||||
cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
caplog.clear()
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
await call(cache)
|
||||
|
||||
assert [record.levelno for record in caplog.records if record.levelno >= logging.WARNING] == []
|
||||
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"call",
|
||||
[lambda cache: cache.get_cache("k"), lambda cache: cache.batch_get_cache(["k1", "k2"])],
|
||||
ids=["get", "batch_get"],
|
||||
)
|
||||
def test_an_open_circuit_breaker_is_not_an_error_per_sync_request(caplog, call):
|
||||
cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
caplog.clear()
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
call(cache)
|
||||
|
||||
assert [record.levelno for record in caplog.records if record.levelno >= logging.WARNING] == []
|
||||
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_real_redis_failure_still_logs_an_error(caplog):
|
||||
class _BrokenRedis:
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
raise ConnectionError("redis is down")
|
||||
|
||||
cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=_BrokenRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
assert await cache.async_get_cache("k") is None
|
||||
|
||||
errors = [record for record in caplog.records if record.levelno == logging.ERROR]
|
||||
assert [record.getMessage() for record in errors] == ["LiteLLM Cache: exception in async_get_cache: redis is down"]
|
||||
assert errors[0].exc_info is not None
|
||||
|
||||
|
||||
def _dual_cache_with_open_breaker_and_a_memory_hit() -> DualCache:
|
||||
in_memory = InMemoryCache()
|
||||
in_memory.set_cache("k1", "v1")
|
||||
return DualCache(in_memory_cache=in_memory, redis_cache=_OpenBreakerRedis(), default_redis_batch_cache_expiry=10) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
|
||||
|
||||
def test_open_breaker_keeps_sync_batch_read_memory_hits_and_releases_reservations():
|
||||
"""A refused Redis batch read must still answer with the in-memory hits and hold no reservation.
|
||||
|
||||
The refusal was logged and turned into a bare None, so a caller lost its in-memory hits
|
||||
for as long as the breaker stayed open, and the reserved keys stayed throttled until
|
||||
the batch expiry passed even though nothing was ever read for them.
|
||||
"""
|
||||
cache = _dual_cache_with_open_breaker_and_a_memory_hit()
|
||||
|
||||
assert list(cache.batch_get_cache(["k1", "k2"])) == ["v1", None]
|
||||
assert "k2" not in cache.last_redis_batch_access_time
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_open_breaker_keeps_async_batch_read_memory_hits_and_releases_reservations():
|
||||
cache = _dual_cache_with_open_breaker_and_a_memory_hit()
|
||||
|
||||
assert list(await cache.async_batch_get_cache(["k1", "k2"])) == ["v1", None]
|
||||
assert "k2" not in cache.last_redis_batch_access_time
|
||||
|
|
|
|||
|
|
@ -1,11 +1,12 @@
|
|||
import asyncio
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm._service_logger import ServiceLogging
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.caching.redis_cache import RedisCache, RedisCircuitBreakerOpenError
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -515,14 +516,46 @@ async def test_circuit_breaker_opens_when_method_swallows_redis_failure(call_met
|
|||
await call_method(cache)
|
||||
|
||||
|
||||
def test_circuit_breaker_open_keeps_sync_batch_get_cache_as_a_miss(sync_batch_redis_cache):
|
||||
"""An open breaker must preserve the sync batch read's dictionary fallback."""
|
||||
def test_circuit_breaker_open_makes_sync_batch_get_cache_fast_fail(sync_batch_redis_cache, caplog):
|
||||
"""Once the breaker is open the sync batch read refuses with the typed error instead of a miss.
|
||||
|
||||
Swallowing the refusal into `{}` made every sync batch read on an open breaker emit an ERROR
|
||||
log and a service failure event per call, and the DualCache caller could not tell the
|
||||
refusal from a dead Redis, so it dropped its in-memory hits too.
|
||||
"""
|
||||
from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
|
||||
|
||||
for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD):
|
||||
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"]) == {}
|
||||
|
||||
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"]) == {}
|
||||
caplog.clear()
|
||||
with caplog.at_level("INFO"):
|
||||
with pytest.raises(RedisCircuitBreakerOpenError):
|
||||
sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"])
|
||||
sync_batch_redis_cache.redis_client.mget.assert_called()
|
||||
assert caplog.records == []
|
||||
|
||||
|
||||
def test_sync_get_cache_failure_feeds_the_breaker_and_logs_a_well_formed_record(sync_batch_redis_cache, caplog):
|
||||
"""The sync get path swallowed its Redis error without recording it, and its log call was malformed.
|
||||
|
||||
`verbose_logger.error("...: ", e)` passes the exception as a format argument to a message
|
||||
with no placeholder, so the record carried no error text. Nothing fed the breaker either,
|
||||
so a dead Redis read through this path never opened it.
|
||||
"""
|
||||
from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
|
||||
|
||||
sync_batch_redis_cache.redis_client.get.side_effect = OSError("redis unavailable")
|
||||
|
||||
with caplog.at_level("ERROR"):
|
||||
for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD):
|
||||
assert sync_batch_redis_cache.get_cache("lit7468") is None
|
||||
|
||||
assert all("redis unavailable" in record.getMessage() for record in caplog.records)
|
||||
assert len(caplog.records) == REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
|
||||
assert sync_batch_redis_cache._circuit_breaker.is_open() is True
|
||||
with pytest.raises(RedisCircuitBreakerOpenError):
|
||||
sync_batch_redis_cache.get_cache("lit7468")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -617,7 +650,8 @@ def test_sync_batch_get_cache_survives_a_service_callback_that_raises(
|
|||
with ThreadPoolExecutor(max_workers=1) as pool:
|
||||
assert pool.submit(cache.batch_get_cache, key_list=["lit6729"]).result() == {}
|
||||
|
||||
assert cache.batch_get_cache(key_list=["lit6729"]) == {}
|
||||
with pytest.raises(RedisCircuitBreakerOpenError):
|
||||
cache.batch_get_cache(key_list=["lit6729"])
|
||||
|
||||
|
||||
def test_call_stack_info_skips_breaker_guard_frames():
|
||||
|
|
@ -966,6 +1000,161 @@ async def test_breaker_metrics_track_state_and_failure_class():
|
|||
assert sample("litellm_redis_circuit_breaker_state", {"state": "open"}) == open_gauge_before + 1
|
||||
assert sample("litellm_redis_circuit_breaker_state", {"state": "closed"}) == closed_gauge_before
|
||||
|
||||
breaker._opened_at = time.time() - 9999
|
||||
assert breaker.is_open() is False
|
||||
breaker.record_success()
|
||||
assert sample("litellm_redis_circuit_breaker_state", {"state": "open"}) == open_gauge_before
|
||||
assert sample("litellm_redis_circuit_breaker_state", {"state": "closed"}) == closed_gauge_before + 1
|
||||
|
||||
|
||||
def test_sync_guard_counts_a_timeout_as_a_timeout():
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, _run_under_circuit_breaker_sync
|
||||
|
||||
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, timeout_min_duration=5.0)
|
||||
|
||||
def timing_out_call() -> str:
|
||||
raise RedisTimeoutError("read timed out")
|
||||
|
||||
for _ in range(6):
|
||||
with pytest.raises(RedisTimeoutError):
|
||||
_run_under_circuit_breaker_sync(breaker, "op", timing_out_call)
|
||||
|
||||
assert breaker.is_open() is False
|
||||
|
||||
|
||||
def test_success_admitted_before_the_breaker_opened_cannot_close_it():
|
||||
"""A stale in-flight success must not close a breaker that opened while it ran.
|
||||
|
||||
Calls admitted while the breaker was still closed finish after later failures opened it.
|
||||
Recording their success unconditionally closed the breaker again, skipping the recovery
|
||||
timeout and the single half-open probe, so the breaker flapped between open and closed
|
||||
on every straggler while Redis was still down.
|
||||
"""
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker
|
||||
|
||||
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
||||
for _ in range(3):
|
||||
breaker.record_failure()
|
||||
assert breaker._state == breaker.OPEN
|
||||
|
||||
breaker.record_success()
|
||||
|
||||
assert breaker._state == breaker.OPEN
|
||||
assert breaker.is_open() is True
|
||||
|
||||
|
||||
def test_recovery_probe_still_closes_the_breaker():
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker
|
||||
|
||||
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
||||
for _ in range(3):
|
||||
breaker.record_failure()
|
||||
breaker._opened_at = time.time() - 9999
|
||||
assert breaker.is_open() is False
|
||||
assert breaker._state == breaker.HALF_OPEN
|
||||
|
||||
breaker.record_success()
|
||||
|
||||
assert breaker._state == breaker.CLOSED
|
||||
assert breaker.is_open() is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_success_during_the_recovery_probe_leaves_the_breaker_to_the_probe():
|
||||
"""A call admitted before the trip that finishes while HALF_OPEN must not close the breaker.
|
||||
|
||||
Only the one call designated as the recovery probe has actually reached Redis after the
|
||||
outage, so closing on the straggler's success resumed full Redis traffic before the probe
|
||||
had proven anything.
|
||||
"""
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, _run_under_circuit_breaker
|
||||
|
||||
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
||||
stale_admitted = asyncio.Event()
|
||||
stale_release = asyncio.Event()
|
||||
probe_admitted = asyncio.Event()
|
||||
probe_release = asyncio.Event()
|
||||
|
||||
async def stale_call() -> str:
|
||||
stale_admitted.set()
|
||||
await stale_release.wait()
|
||||
return "stale"
|
||||
|
||||
async def probe_call() -> str:
|
||||
probe_admitted.set()
|
||||
await probe_release.wait()
|
||||
return "probe"
|
||||
|
||||
stale = asyncio.ensure_future(_run_under_circuit_breaker(breaker, "op", stale_call))
|
||||
await stale_admitted.wait()
|
||||
for _ in range(3):
|
||||
breaker.record_failure()
|
||||
assert breaker._state == breaker.OPEN
|
||||
breaker._opened_at = time.time() - 9999
|
||||
probe = asyncio.ensure_future(_run_under_circuit_breaker(breaker, "op", probe_call))
|
||||
await probe_admitted.wait()
|
||||
assert breaker._state == breaker.HALF_OPEN
|
||||
|
||||
stale_release.set()
|
||||
assert await stale == "stale"
|
||||
|
||||
assert breaker._state == breaker.HALF_OPEN, "the straggler must not close the breaker for the probe"
|
||||
assert breaker.is_open() is True
|
||||
|
||||
probe_release.set()
|
||||
assert await probe == "probe"
|
||||
|
||||
assert breaker._state == breaker.CLOSED
|
||||
assert breaker.is_open() is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_probe_overtaken_by_a_later_outage_leaves_the_breaker_to_the_new_probe():
|
||||
"""A probe still in flight when a late failure reopens the breaker must not close it for the next probe.
|
||||
|
||||
Once the breaker has reopened, only the probe admitted after that outage has reached
|
||||
Redis, so the older probe's success no longer says anything about whether Redis recovered.
|
||||
"""
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, _run_under_circuit_breaker
|
||||
|
||||
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
||||
old_probe_admitted = asyncio.Event()
|
||||
old_probe_release = asyncio.Event()
|
||||
new_probe_admitted = asyncio.Event()
|
||||
new_probe_release = asyncio.Event()
|
||||
|
||||
async def old_probe_call() -> str:
|
||||
old_probe_admitted.set()
|
||||
await old_probe_release.wait()
|
||||
return "old probe"
|
||||
|
||||
async def new_probe_call() -> str:
|
||||
new_probe_admitted.set()
|
||||
await new_probe_release.wait()
|
||||
return "new probe"
|
||||
|
||||
for _ in range(3):
|
||||
breaker.record_failure()
|
||||
breaker._opened_at = time.time() - 9999
|
||||
old_probe = asyncio.ensure_future(_run_under_circuit_breaker(breaker, "op", old_probe_call))
|
||||
await old_probe_admitted.wait()
|
||||
assert breaker._state == breaker.HALF_OPEN
|
||||
|
||||
breaker.record_failure()
|
||||
assert breaker._state == breaker.OPEN
|
||||
breaker._opened_at = time.time() - 9999
|
||||
new_probe = asyncio.ensure_future(_run_under_circuit_breaker(breaker, "op", new_probe_call))
|
||||
await new_probe_admitted.wait()
|
||||
assert breaker._state == breaker.HALF_OPEN
|
||||
|
||||
old_probe_release.set()
|
||||
assert await old_probe == "old probe"
|
||||
|
||||
assert breaker._state == breaker.HALF_OPEN, "the overtaken probe must not close the breaker for the new probe"
|
||||
assert breaker.is_open() is True
|
||||
|
||||
new_probe_release.set()
|
||||
assert await new_probe == "new probe"
|
||||
assert breaker._state == breaker.CLOSED
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
import logging
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -6,6 +7,7 @@ import pytest
|
|||
from fastapi.testclient import TestClient
|
||||
|
||||
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.constants import DEFAULT_CRON_JOB_LOCK_TTL_SECONDS
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
|
||||
|
||||
|
|
@ -215,6 +217,22 @@ async def test_redis_error_handling(pod_lock_manager, mock_redis):
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lock_refused_by_the_open_circuit_breaker_is_not_logged_as_an_error(pod_lock_manager, mock_redis, caplog):
|
||||
"""Every cron job retries its lock on a timer, so an open breaker must not add an error line per cycle."""
|
||||
refused = RedisCircuitBreakerOpenError("Redis circuit breaker is open - skipping async_set_cache")
|
||||
mock_redis.async_set_cache.side_effect = refused
|
||||
mock_redis.async_get_cache.return_value = pod_lock_manager.pod_id
|
||||
mock_redis.async_delete_cache.side_effect = refused
|
||||
|
||||
with caplog.at_level(logging.ERROR):
|
||||
acquired = await pod_lock_manager.acquire_lock(cronjob_id="test_job")
|
||||
await pod_lock_manager.release_lock(cronjob_id="test_job")
|
||||
|
||||
assert acquired is False
|
||||
assert caplog.records == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bytes_handling(pod_lock_manager, mock_redis):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ response-shape helpers the v3 limiter's post-call hooks rely on.
|
|||
"""
|
||||
|
||||
import base64
|
||||
import logging
|
||||
import socket
|
||||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
|
|
@ -16,6 +17,7 @@ from typing import Final
|
|||
import pytest
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.constants import BATCH_ENQUEUED_TOKEN_TTL_SECONDS
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.batch_enqueued_tokens import (
|
||||
|
|
@ -439,3 +441,29 @@ async def test_redis_lua_path_full_lifecycle():
|
|||
refill = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
|
||||
assert isinstance(refill, BatchEnqueuedTokenReservation)
|
||||
await store.refund(refill)
|
||||
|
||||
|
||||
class _OpenBreakerRedis:
|
||||
def async_register_script(self, script: str):
|
||||
async def refused(keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> object:
|
||||
raise RedisCircuitBreakerOpenError("Redis circuit breaker is open")
|
||||
|
||||
return refused
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_open_circuit_breaker_keeps_reservations_in_memory_without_a_warning(caplog):
|
||||
scope = _scope(limit=100)
|
||||
store = BatchEnqueuedTokenStore(
|
||||
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=_OpenBreakerRedis(), default_in_memory_ttl=60)) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
reservation = await store.reserve(tokens=60, scopes=(scope,))
|
||||
assert isinstance(reservation, BatchEnqueuedTokenReservation)
|
||||
assert reservation.backend == "memory"
|
||||
await store.save_reservation("batch_quiet", reservation)
|
||||
assert await store.pop_reservation("batch_quiet") == reservation
|
||||
|
||||
assert not [record for record in caplog.records if record.levelno >= logging.WARNING]
|
||||
assert sum("circuit breaker is open" in record.getMessage() for record in caplog.records) == 3
|
||||
|
|
|
|||
|
|
@ -10,10 +10,12 @@ Tests that session-scoped budget tracking works correctly:
|
|||
|
||||
from unittest.mock import patch
|
||||
|
||||
import logging
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import _redis_circuit_breaker_guard
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.max_budget_per_session_limiter import (
|
||||
_PROXY_MaxBudgetPerSessionHandler,
|
||||
|
|
@ -163,3 +165,37 @@ async def test_no_agent_id_passes():
|
|||
call_type="",
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
class _OpenBreakerRedis:
|
||||
def __init__(self) -> None:
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker
|
||||
|
||||
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
||||
for _ in range(3):
|
||||
self._circuit_breaker.record_failure()
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
def async_register_script(self, script):
|
||||
@_redis_circuit_breaker_guard
|
||||
async def refused(_self, keys, args):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
return lambda keys, args: refused(self, keys, args)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_open_circuit_breaker_reads_session_spend_locally_without_a_warning(caplog):
|
||||
cache = DualCache(redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
handler = _PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
caplog.clear()
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
spend = await handler._get_current_spend("{session_budget:quiet}:spend")
|
||||
|
||||
assert spend == 0.0
|
||||
assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == []
|
||||
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
|
||||
|
|
|
|||
|
|
@ -6171,3 +6171,51 @@ async def test_success_hook_leaves_stash_untouched_for_non_batch_responses():
|
|||
data={}, user_api_key_dict=user, response=ModelResponse(usage=Usage(total_tokens=5))
|
||||
)
|
||||
assert get_request_stash().batch_enqueued_reservation == reservation
|
||||
|
||||
|
||||
class _OpenBreakerRedis:
|
||||
def async_register_script(self, script: str):
|
||||
async def refused(keys, args):
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
|
||||
raise RedisCircuitBreakerOpenError("Redis circuit breaker is open")
|
||||
|
||||
return refused
|
||||
|
||||
async def async_increment_pipeline(self, increment_list, **kwargs):
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
|
||||
raise RedisCircuitBreakerOpenError("Redis circuit breaker is open")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_open_circuit_breaker_falls_back_to_the_pipeline_without_a_warning(caplog):
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=_OpenBreakerRedis())) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
await handler.async_increment_tokens_with_ttl_preservation(
|
||||
pipeline_operations=[RedisPipelineIncrementOperation(key="quiet_key", increment_value=10.0, ttl=60)]
|
||||
)
|
||||
|
||||
assert await handler.internal_usage_cache.dual_cache.async_get_cache("quiet_key") == 10.0
|
||||
assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_open_circuit_breaker_reads_the_sliding_window_locally_without_a_warning(caplog):
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=_OpenBreakerRedis())) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
values = await handler._execute_redis_batch_rate_limiter_script(
|
||||
["{quiet}:window", "{quiet}:counter"], now_int=int(time.time())
|
||||
)
|
||||
|
||||
assert isinstance(values, list)
|
||||
assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == []
|
||||
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ This feature allows guardrails to route requests to a different model
|
|||
All subsequent requests in the same session are routed to the same model.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import asyncio
|
||||
from typing import Any, Dict, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -13,12 +14,14 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import pytest
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import _redis_circuit_breaker_guard
|
||||
from litellm.exceptions import SensitiveDataRouteException
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
get_session_id_from_request_data,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import InternalUsageCache
|
||||
from litellm.proxy.hooks.sensitive_data_routing import (
|
||||
_PROXY_SensitiveDataRoutingHandler,
|
||||
SENSITIVE_ROUTING_CACHE_PREFIX,
|
||||
|
|
@ -1034,3 +1037,35 @@ class TestPreCallHookDeferredRouting:
|
|||
metrics_kwargs = prom._record_guardrail_metrics.call_args.kwargs
|
||||
assert metrics_kwargs["status"] == "intervened"
|
||||
assert metrics_kwargs["error_type"] is None
|
||||
|
||||
|
||||
class _OpenBreakerRedis:
|
||||
def __init__(self) -> None:
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker
|
||||
|
||||
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
||||
for _ in range(3):
|
||||
self._circuit_breaker.record_failure()
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_set_cache(self, key, value, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_open_circuit_breaker_keeps_session_routing_in_memory_without_a_warning(caplog):
|
||||
cache = DualCache(redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
caplog.clear()
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
|
||||
await handler.set_session_routing("quiet-session", "safe-model")
|
||||
routed = await handler._get_routed_model("quiet-session", None)
|
||||
|
||||
assert routed == "safe-model"
|
||||
assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == []
|
||||
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
|
||||
|
|
|
|||
|
|
@ -2616,6 +2616,28 @@ class TestCLIKeyRegenerationFlow:
|
|||
)
|
||||
cache.set_cache.assert_not_called()
|
||||
|
||||
def test_cli_sso_flow_lookup_treats_an_open_redis_breaker_as_a_miss(self):
|
||||
"""A Redis read refused by the open circuit breaker is a missing session, not a server error.
|
||||
|
||||
The direct Redis read is what keeps the flow authoritative across workers, so the
|
||||
refusal must not fall back to a possibly stale in-memory copy either.
|
||||
"""
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.proxy.management_endpoints.ui_sso import _get_cli_sso_flow_or_raise
|
||||
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.get_cache.side_effect = RedisCircuitBreakerOpenError("Redis circuit breaker is open")
|
||||
cache = MagicMock()
|
||||
cache.redis_cache = redis_cache
|
||||
cache.get_cache.return_value = {"poll_secret_hash": "stale", "sso_complete": False}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_get_cli_sso_flow_or_raise(login_id="cli-breaker_open_1234567890", cache=cache)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "not found or expired" in exc_info.value.detail
|
||||
cache.get_cache.assert_not_called()
|
||||
|
||||
def test_cli_sso_flow_with_enum_survives_redis_round_trip(self):
|
||||
"""
|
||||
RedisCache stores values via str(value) and reads them back through
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -41,6 +41,15 @@ class _FlakyRedisCache:
|
|||
self._store[key] = float(value)
|
||||
return True
|
||||
|
||||
async def async_increment_pipeline(self, increment_list, **kwargs):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await self.async_increment(op["key"], op["increment_value"]))
|
||||
return results
|
||||
|
||||
def get_ttl(self, **kwargs):
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_direct_increment_runs_when_reservation_reconcile_hits_redis_failure(
|
||||
|
|
|
|||
|
|
@ -7710,7 +7710,7 @@ async def test_increment_spend_counters_team_and_member():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss():
|
||||
async def test_prepare_spend_counter_increment_reseeds_from_db_on_counter_miss():
|
||||
"""When the Redis counter is missing, the reseed path reads the
|
||||
authoritative spend from the DB (not a stale cache), so the next
|
||||
increment continues from the correct base value."""
|
||||
|
|
@ -7723,8 +7723,17 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
|
|||
recorded_increments.append({"key": key, "value": value, "ttl": ttl})
|
||||
return value
|
||||
|
||||
async def record_pipeline(increment_list, **kwargs):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
await record_increment(key=op["key"], value=op["increment_value"], ttl=op["ttl"])
|
||||
results.append(op["increment_value"])
|
||||
return results
|
||||
|
||||
fake_redis = AsyncMock()
|
||||
fake_redis.async_increment = AsyncMock(side_effect=record_increment)
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=record_pipeline)
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing
|
||||
fake_redis.async_set_cache = AsyncMock(return_value=True) # SET NX wins
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
|
@ -7743,7 +7752,10 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
|
|||
stale_cache.in_memory_cache.set_cache(key="team_id:team-9", value=stale_team)
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.proxy.proxy_server import _init_and_increment_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_spend_counter_increment,
|
||||
)
|
||||
|
||||
orig_user, orig_counter, orig_prisma = (
|
||||
ps.user_api_key_cache,
|
||||
|
|
@ -7754,11 +7766,12 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
|
|||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_spend_counter(
|
||||
pending = await _prepare_spend_counter_increment(
|
||||
counter_key="spend:team:team-9",
|
||||
source_cache_key="team_id:team-9",
|
||||
increment=1.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,))
|
||||
|
||||
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-9"})
|
||||
# Seed uses SET NX with db_spend (42) — cross-pod safe, no INCR of 42.
|
||||
|
|
@ -7937,7 +7950,10 @@ async def test_reseed_spend_from_db_skips_window_variant_keys():
|
|||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_window_spend_counter_increment,
|
||||
)
|
||||
|
||||
counter_cache = DualCache()
|
||||
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
|
|
@ -7953,7 +7969,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
|||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending = await _prepare_window_spend_counter_increment(
|
||||
counter_key="spend:key:key-window:window:1h",
|
||||
entity_type="Key",
|
||||
entity_id="key-window",
|
||||
|
|
@ -7961,6 +7977,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
|||
window_start=window_start,
|
||||
increment=0.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
|
||||
|
||||
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
|
||||
by=["api_key"],
|
||||
|
|
@ -7976,7 +7993,10 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
|||
@pytest.mark.asyncio
|
||||
async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_spend_counter_increment,
|
||||
)
|
||||
|
||||
counter_cache = DualCache()
|
||||
counter_key = "spend:team:team-stale-local"
|
||||
|
|
@ -7998,6 +8018,15 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||||
|
||||
async def redis_increment_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
|
||||
return results
|
||||
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
||||
db_row = MagicMock()
|
||||
|
|
@ -8016,11 +8045,12 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
ps.prisma_client = fake_prisma
|
||||
ps.user_api_key_cache = DualCache()
|
||||
try:
|
||||
await _init_and_increment_spend_counter(
|
||||
pending = await _prepare_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
source_cache_key="team_id:team-stale-local",
|
||||
increment=1.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,))
|
||||
|
||||
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-stale-local"})
|
||||
# Seed via SET NX (42) + delta via INCRBYFLOAT (1.5) = 43.5.
|
||||
|
|
@ -8035,7 +8065,10 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_window_spend_counter_increment,
|
||||
)
|
||||
|
||||
counter_cache = DualCache()
|
||||
counter_key = "spend:key:key-window-stale-local:window:1h"
|
||||
|
|
@ -8058,6 +8091,15 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||||
|
||||
async def redis_increment_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
|
||||
return results
|
||||
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
||||
fake_prisma = MagicMock()
|
||||
|
|
@ -8072,7 +8114,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending = await _prepare_window_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
entity_type="Key",
|
||||
entity_id="key-window-stale-local",
|
||||
|
|
@ -8080,6 +8122,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
window_start=window_start,
|
||||
increment=0.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
|
||||
|
||||
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
|
||||
by=["api_key"],
|
||||
|
|
@ -8099,7 +8142,10 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_window_spend_counter_increment,
|
||||
)
|
||||
|
||||
counter_cache = DualCache()
|
||||
counter_key = "spend:key:key-window-concurrent-seed:window:1h"
|
||||
|
|
@ -8122,6 +8168,15 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
|
|||
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache)
|
||||
fake_redis.async_set_cache = AsyncMock(return_value=False)
|
||||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||||
|
||||
async def redis_increment_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
|
||||
return results
|
||||
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
||||
fake_prisma = MagicMock()
|
||||
|
|
@ -8136,7 +8191,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
|
|||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending = await _prepare_window_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
entity_type="Key",
|
||||
entity_id="key-window-concurrent-seed",
|
||||
|
|
@ -8144,6 +8199,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
|
|||
window_start=window_start,
|
||||
increment=0.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
|
||||
|
||||
fake_redis.async_set_cache.assert_awaited_once_with(
|
||||
key=counter_key,
|
||||
|
|
@ -8160,7 +8216,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
|
|||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_skips_invalid_window_start():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
from litellm.proxy.proxy_server import _prepare_window_spend_counter_increment
|
||||
|
||||
counter_cache = DualCache()
|
||||
|
||||
|
|
@ -8169,7 +8225,7 @@ async def test_window_spend_counter_skips_invalid_window_start():
|
|||
orig_counter = ps.spend_counter_cache
|
||||
ps.spend_counter_cache = counter_cache
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending = await _prepare_window_spend_counter_increment(
|
||||
counter_key="spend:key:key-invalid-window:window:not-a-duration",
|
||||
entity_type="Key",
|
||||
entity_id="key-invalid-window",
|
||||
|
|
@ -8177,6 +8233,7 @@ async def test_window_spend_counter_skips_invalid_window_start():
|
|||
window_start=None,
|
||||
increment=0.5,
|
||||
)
|
||||
assert pending is None
|
||||
|
||||
assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-invalid-window:window:not-a-duration") is None
|
||||
finally:
|
||||
|
|
@ -8240,6 +8297,9 @@ async def test_increment_spend_counters_finalizes_after_unreserved_increments():
|
|||
async def assert_reservation_not_finalized_yet(**kwargs):
|
||||
assert budget_reservation["finalized"] is False
|
||||
incremented_counters.append(kwargs["counter_key"])
|
||||
return ps._PendingSpendIncrement(
|
||||
counter_key=kwargs["counter_key"], increment=kwargs["increment"]
|
||||
)
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
|
|
@ -8248,7 +8308,7 @@ async def test_increment_spend_counters_finalizes_after_unreserved_increments():
|
|||
ps.user_api_key_cache = DualCache()
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server._init_and_increment_spend_counter",
|
||||
"litellm.proxy.proxy_server._prepare_spend_counter_increment",
|
||||
new=AsyncMock(side_effect=assert_reservation_not_finalized_yet),
|
||||
):
|
||||
await increment_spend_counters(
|
||||
|
|
@ -8583,7 +8643,7 @@ async def test_get_current_spend_uses_db_zero_over_stale_fallback():
|
|||
async def test_concurrent_read_and_write_paths_share_one_db_query():
|
||||
"""
|
||||
The read path (`get_current_spend`) and the write path
|
||||
(`_init_and_increment_spend_counter`) both reseed cold counters from
|
||||
(`_prepare_spend_counter_increment`) both reseed cold counters from
|
||||
the DB. They must share the per-counter lock so a concurrent pre-call
|
||||
enforcement read and post-call increment for the same counter collapse
|
||||
to one DB query, not two.
|
||||
|
|
@ -8592,7 +8652,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query():
|
|||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import (
|
||||
_init_and_increment_spend_counter,
|
||||
_prepare_spend_counter_increment,
|
||||
get_current_spend,
|
||||
)
|
||||
|
||||
|
|
@ -8646,7 +8706,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query():
|
|||
try:
|
||||
results = await _asyncio.gather(
|
||||
get_current_spend(counter_key=counter_key, fallback_spend=0.0),
|
||||
_init_and_increment_spend_counter(
|
||||
_prepare_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
source_cache_key="ignored",
|
||||
increment=1.5,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
import logging
|
||||
from typing import Any, Dict, List, Optional, Set, Union
|
||||
|
||||
import pytest
|
||||
|
|
@ -9,7 +10,7 @@ from unittest.mock import MagicMock, patch
|
|||
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisPipelineIncrementOperation
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError, RedisPipelineIncrementOperation
|
||||
from litellm.router_strategy.base_routing_strategy import BaseRoutingStrategy
|
||||
|
||||
|
||||
|
|
@ -146,3 +147,18 @@ async def test_cache_keys_management(base_strategy):
|
|||
# Test resetting cache keys
|
||||
base_strategy.reset_in_memory_keys_to_update()
|
||||
assert len(base_strategy.get_in_memory_keys_to_update()) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_push_refused_by_the_open_circuit_breaker_is_not_logged_as_an_error(base_strategy, mock_dual_cache, caplog):
|
||||
"""The sync loop pushes every 100 ms under usage-based routing, so an open breaker must not add an error line per cycle."""
|
||||
mock_dual_cache.redis_cache.async_increment_pipeline.side_effect = RedisCircuitBreakerOpenError(
|
||||
"Redis circuit breaker is open - skipping async_increment_pipeline"
|
||||
)
|
||||
base_strategy.redis_increment_operation_queue = [{"key": "k", "increment_value": 1.0, "ttl": 60}]
|
||||
|
||||
with caplog.at_level(logging.ERROR):
|
||||
await base_strategy._push_in_memory_increments_to_redis()
|
||||
|
||||
assert caplog.records == []
|
||||
assert base_strategy.redis_increment_operation_queue == []
|
||||
|
|
|
|||
|
|
@ -1,7 +1,13 @@
|
|||
import asyncio
|
||||
import gc
|
||||
import logging
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisCache, RedisCircuitBreakerOpenError
|
||||
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
|
||||
from litellm.types.router import LiteLLM_Params
|
||||
from litellm.types.utils import BudgetConfig
|
||||
|
|
@ -303,3 +309,90 @@ def test_router_add_deployment_registers_deployment_budget(
|
|||
)
|
||||
assert config is not None
|
||||
assert config.max_budget == 0.000000000001
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_refused_by_the_open_circuit_breaker_is_quiet_and_leaks_no_task(disable_budget_sync, caplog):
|
||||
"""The budget sync runs every second, so an open breaker must not add an error line or an unretrieved task exception per cycle."""
|
||||
refused = RedisCircuitBreakerOpenError("Redis circuit breaker is open - skipping async_increment_pipeline")
|
||||
redis_cache = MagicMock(spec=RedisCache)
|
||||
redis_cache.async_increment_pipeline = AsyncMock(side_effect=refused)
|
||||
redis_cache.async_batch_get_cache = AsyncMock(side_effect=refused)
|
||||
limiter = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(redis_cache=redis_cache),
|
||||
provider_budget_config={"openai": BudgetConfig(max_budget=1.0, budget_duration="1d")},
|
||||
)
|
||||
await asyncio.gather(*(task for task in asyncio.all_tasks() if task is not asyncio.current_task()))
|
||||
limiter.redis_increment_operation_queue = [{"key": "provider_spend:openai:1d", "increment_value": 0.5, "ttl": 60}]
|
||||
loop = asyncio.get_running_loop()
|
||||
unretrieved = MagicMock()
|
||||
loop.set_exception_handler(unretrieved)
|
||||
|
||||
try:
|
||||
with caplog.at_level(logging.ERROR):
|
||||
await limiter._sync_in_memory_spend_with_redis()
|
||||
await asyncio.sleep(0)
|
||||
gc.collect()
|
||||
finally:
|
||||
loop.set_exception_handler(None)
|
||||
|
||||
assert caplog.records == []
|
||||
unretrieved.assert_not_called()
|
||||
assert limiter.redis_increment_operation_queue == []
|
||||
assert redis_cache.async_increment_pipeline.await_count == 1
|
||||
|
||||
|
||||
async def _limiter_with_redis(redis_cache: MagicMock) -> RouterBudgetLimiting:
|
||||
limiter = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(redis_cache=redis_cache),
|
||||
provider_budget_config={"openai": BudgetConfig(max_budget=1.0, budget_duration="1d")},
|
||||
)
|
||||
await asyncio.gather(*(task for task in asyncio.all_tasks() if task is not asyncio.current_task()))
|
||||
limiter.redis_increment_operation_queue = [{"key": "provider_spend:openai:1d", "increment_value": 0.5, "ttl": 60}]
|
||||
return limiter
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_push_returns_before_redis_answers(disable_budget_sync):
|
||||
"""The push runs inside the request success callback, so it must hand the Redis round trip to a task instead of waiting on it."""
|
||||
redis_answered = asyncio.Event()
|
||||
|
||||
async def wait_for_redis(**_: object) -> None:
|
||||
await redis_answered.wait()
|
||||
|
||||
redis_cache = MagicMock(spec=RedisCache)
|
||||
redis_cache.async_increment_pipeline = AsyncMock(side_effect=wait_for_redis)
|
||||
limiter = await _limiter_with_redis(redis_cache)
|
||||
|
||||
await asyncio.wait_for(limiter._push_in_memory_increments_to_redis(), timeout=1)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert not redis_answered.is_set()
|
||||
assert redis_cache.async_increment_pipeline.await_count == 1
|
||||
assert limiter.redis_increment_operation_queue == []
|
||||
redis_answered.set()
|
||||
await asyncio.gather(*(task for task in asyncio.all_tasks() if task is not asyncio.current_task()))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_push_task_failure_is_logged_once_and_not_leaked(disable_budget_sync, caplog):
|
||||
"""A real Redis failure on the background push must surface as one error line, never as an unretrieved task exception."""
|
||||
redis_cache = MagicMock(spec=RedisCache)
|
||||
redis_cache.async_increment_pipeline = AsyncMock(side_effect=ConnectionError("Error 61 connecting to 127.0.0.1:6379"))
|
||||
limiter = await _limiter_with_redis(redis_cache)
|
||||
loop = asyncio.get_running_loop()
|
||||
unretrieved = MagicMock()
|
||||
loop.set_exception_handler(unretrieved)
|
||||
|
||||
try:
|
||||
with caplog.at_level(logging.ERROR):
|
||||
await limiter._push_in_memory_increments_to_redis()
|
||||
await asyncio.sleep(0)
|
||||
gc.collect()
|
||||
finally:
|
||||
loop.set_exception_handler(None)
|
||||
|
||||
assert [record.getMessage() for record in caplog.records] == [
|
||||
"Error syncing in-memory cache with Redis: Error 61 connecting to 127.0.0.1:6379"
|
||||
]
|
||||
unretrieved.assert_not_called()
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ import time
|
|||
import pytest
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.router_utils.health_state_cache import DeploymentHealthCache
|
||||
|
||||
|
||||
|
|
@ -145,8 +146,11 @@ class _SharedRedisFake:
|
|||
def __init__(self):
|
||||
self.store = {}
|
||||
self.fail_get = False
|
||||
self.breaker_open = False
|
||||
|
||||
def get_cache(self, key, parent_otel_span=None, **kwargs):
|
||||
if self.breaker_open:
|
||||
raise RedisCircuitBreakerOpenError("Redis circuit breaker is open - skipping get_cache")
|
||||
if self.fail_get:
|
||||
return None # RedisCache.get_cache swallows connection errors and returns None
|
||||
return self.store.get(key)
|
||||
|
|
@ -192,3 +196,26 @@ def test_failed_redis_read_falls_back_to_local_copy():
|
|||
{"prod-bad": {"is_healthy": False, "timestamp": time.time(), "reason": "check_failed"}}
|
||||
)
|
||||
assert set(redis_fake.store[DeploymentHealthCache.CACHE_KEY]) == {"prod-bad", "internal-bad"}
|
||||
|
||||
|
||||
def test_open_circuit_breaker_read_still_merges_into_local_copy(caplog):
|
||||
"""A read refused by the open breaker is a miss, so the merge and local write still happen quietly."""
|
||||
redis_fake = _SharedRedisFake()
|
||||
pod_a = DeploymentHealthCache(cache=DualCache(redis_cache=redis_fake), staleness_threshold=60.0)
|
||||
pod_b = DeploymentHealthCache(cache=DualCache(redis_cache=redis_fake), staleness_threshold=60.0)
|
||||
pod_a.set_deployment_health_states(
|
||||
{"prod-bad": {"is_healthy": False, "timestamp": time.time(), "reason": "check_failed"}}
|
||||
)
|
||||
pod_b.set_deployment_health_states(
|
||||
{"internal-bad": {"is_healthy": False, "timestamp": time.time(), "reason": "timeout"}}
|
||||
)
|
||||
pod_a.set_deployment_health_states(
|
||||
{"prod-bad": {"is_healthy": False, "timestamp": time.time(), "reason": "check_failed"}}
|
||||
)
|
||||
redis_fake.breaker_open = True
|
||||
with caplog.at_level("ERROR"):
|
||||
pod_a.set_deployment_health_states(
|
||||
{"prod-new-bad": {"is_healthy": False, "timestamp": time.time(), "reason": "check_failed"}}
|
||||
)
|
||||
assert caplog.records == []
|
||||
assert pod_a.get_unhealthy_deployment_ids() == {"prod-bad", "internal-bad", "prod-new-bad"}
|
||||
|
|
|
|||
|
|
@ -18,6 +18,8 @@ import respx
|
|||
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import _redis_circuit_breaker_guard
|
||||
from litellm import Router
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -13421,3 +13423,32 @@ async def test_router_max_parallel_requests_slot_released_when_stream_closed_ear
|
|||
|
||||
assert tracker.peak == 1
|
||||
assert tracker.current == 0
|
||||
|
||||
|
||||
class _OpenBreakerRedis:
|
||||
def __init__(self) -> None:
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker
|
||||
|
||||
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
||||
for _ in range(3):
|
||||
self._circuit_breaker.record_failure()
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
raise AssertionError("never reached")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_open_circuit_breaker_skips_the_session_binding_without_a_warning(caplog):
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "haiku", "litellm_params": {"model": "anthropic/claude-haiku-4-5", "api_key": "k"}}]
|
||||
)
|
||||
router._claude_code_session_router_cache = DualCache(redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
caplog.clear()
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM Router"):
|
||||
binding = await router._get_claude_code_session_router_binding("quiet-session")
|
||||
|
||||
assert binding is None
|
||||
assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == []
|
||||
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
|
||||
|
|
|
|||
2
uv.lock
generated
2
uv.lock
generated
|
|
@ -4559,6 +4559,7 @@ e2e-dev = [
|
|||
{ name = "locust" },
|
||||
{ name = "mcp" },
|
||||
{ name = "playwright" },
|
||||
{ name = "psutil" },
|
||||
{ name = "websockets" },
|
||||
]
|
||||
healthcheck = [
|
||||
|
|
@ -4747,6 +4748,7 @@ e2e-dev = [
|
|||
{ name = "locust", specifier = "==2.45.0" },
|
||||
{ name = "mcp", specifier = ">=1.28.1,<2.0" },
|
||||
{ name = "playwright", specifier = "==1.61.0" },
|
||||
{ name = "psutil", specifier = "==7.2.2" },
|
||||
{ name = "websockets", specifier = ">=15.0.1,<16.0" },
|
||||
]
|
||||
healthcheck = [
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue