mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_invalid_index_after_migration_deadlock
This commit is contained in:
commit
84a5136f45
31 changed files with 1446 additions and 372 deletions
|
|
@ -1440,6 +1440,7 @@ jobs:
|
|||
TEST_FILES=$(printf "%s\n" \
|
||||
tests/local_testing/test_dual_cache.py \
|
||||
tests/local_testing/test_redis_batch_optimizations.py \
|
||||
tests/local_testing/test_redis_increment_with_floor.py \
|
||||
tests/local_testing/test_router_utils.py)
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ from contextvars import ContextVar
|
|||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -80,11 +82,29 @@ class _AsyncRedisCommands(Protocol):
|
|||
|
||||
def pipeline(self, transaction: bool = True) -> "Pipeline[bytes]": ...
|
||||
|
||||
def eval(self, script: str, numkeys: int, *keys_and_args: str | bytes | float) -> Awaitable[object]: ...
|
||||
|
||||
|
||||
_BREAKER_GUARD_FRAME_NAMES: Final = frozenset(
|
||||
{"<lambda>", "wrapper", "_run_under_circuit_breaker", "_run_under_circuit_breaker_sync"}
|
||||
)
|
||||
|
||||
_INCREMENT_WITH_FLOOR_LUA: Final = (
|
||||
"local count = redis.call('INCRBY', KEYS[1], ARGV[1]) "
|
||||
"if count < 0 then count = redis.call('INCRBY', KEYS[1], -count) end "
|
||||
"if redis.call('TTL', KEYS[1]) < 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]) end "
|
||||
"return count"
|
||||
)
|
||||
|
||||
_LUA_COUNT: Final = TypeAdapter(int)
|
||||
_OPTIONAL_COUNTS: Final = TypeAdapter(tuple[int | None, ...])
|
||||
|
||||
|
||||
def _decoded_counts(values: Sequence[bytes | str | None]) -> tuple[int | None, ...]:
|
||||
return _OPTIONAL_COUNTS.validate_python(
|
||||
tuple(value.decode("utf-8") if isinstance(value, bytes) else value for value in values)
|
||||
)
|
||||
|
||||
|
||||
def _get_call_stack_info(num_frames: int = 2) -> str:
|
||||
"""
|
||||
|
|
@ -736,6 +756,43 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
raise e
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
"""Add ``value`` to ``key``, clamp the result at zero, and give a new key ``ttl``, in one Lua call.
|
||||
|
||||
A counter whose key expired while a request was still in flight would otherwise be
|
||||
recreated negative by that request's decrement. Clamping inside the same call is what
|
||||
keeps it safe: a separate corrective write could land after another pod's increment and
|
||||
erase it.
|
||||
|
||||
The TTL is set only on a key that has none, so a counter expires ``ttl`` after it was
|
||||
created rather than ``ttl`` after it was last touched. Refreshing it on every touch
|
||||
would keep a count a dead worker never decremented alive for as long as the group
|
||||
takes traffic. Returns the resulting count.
|
||||
"""
|
||||
namespaced_key: Final = self.check_and_fix_namespace(key=key)
|
||||
count: Final[object] = self.redis_client.eval( # pyright: ignore[reportAttributeAccessIssue] # stubs omit eval
|
||||
_INCREMENT_WITH_FLOOR_LUA, 1, namespaced_key, value, ttl
|
||||
)
|
||||
return _LUA_COUNT.validate_python(count)
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
"""Read integer counters for ``key_list``, in order, raising when Redis cannot answer.
|
||||
|
||||
``batch_get_cache`` swallows every failure and returns an empty dict, which the caller
|
||||
cannot tell apart from "every counter is unset". A caller that has to fall back to its
|
||||
own numbers when Redis is unreachable needs the failure, not a dict of zeros.
|
||||
"""
|
||||
namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list]
|
||||
return _decoded_counts(self._run_redis_mget_operation(keys=namespaced_keys))
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
"""Async twin of ``batch_get_counts``, raising on failure the same way."""
|
||||
namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list]
|
||||
return _decoded_counts(await self._async_run_redis_mget_operation(keys=namespaced_keys))
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_scan_iter(self, pattern: str, count: int = 100) -> list:
|
||||
start_time: Final = time.time()
|
||||
|
|
@ -1241,6 +1298,14 @@ class RedisCache(BaseCache):
|
|||
result = result.decode()
|
||||
return float(result)
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
"""Async twin of ``increment_with_floor``, sharing its Lua script and its guarantees."""
|
||||
_redis_client: Final = self._async_commands()
|
||||
namespaced_key: Final = self.check_and_fix_namespace(key=key)
|
||||
count: Final = await _redis_client.eval(_INCREMENT_WITH_FLOOR_LUA, 1, namespaced_key, value, ttl)
|
||||
return _LUA_COUNT.validate_python(count)
|
||||
|
||||
async def flush_cache_buffer(self):
|
||||
print_verbose(f"flushing to redis....reached size of buffer {len(self.redis_batch_writing_buffer)}")
|
||||
await self.async_set_cache_pipeline(self.redis_batch_writing_buffer)
|
||||
|
|
|
|||
|
|
@ -73,6 +73,9 @@ DEFAULT_MAX_TOKENS: Final = int(os.getenv("DEFAULT_MAX_TOKENS", 4096))
|
|||
DEFAULT_ALLOWED_FAILS: Final = int(os.getenv("DEFAULT_ALLOWED_FAILS", 3))
|
||||
DEFAULT_REDIS_SYNC_INTERVAL: Final = int(os.getenv("DEFAULT_REDIS_SYNC_INTERVAL", 1))
|
||||
DEFAULT_COOLDOWN_TIME_SECONDS: Final = int(os.getenv("DEFAULT_COOLDOWN_TIME_SECONDS", 5))
|
||||
DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS: Final = float(
|
||||
os.getenv("DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS", "1")
|
||||
)
|
||||
DEFAULT_REPLICATE_POLLING_RETRIES: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_RETRIES", 5))
|
||||
DEFAULT_REPLICATE_POLLING_DELAY_SECONDS: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1))
|
||||
DEFAULT_IMAGE_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250))
|
||||
|
|
|
|||
|
|
@ -7,16 +7,19 @@ from litellm.types.utils import CredentialItem
|
|||
|
||||
|
||||
class CredentialAccessor:
|
||||
@staticmethod
|
||||
def find_credential(credential_name: str) -> CredentialItem | None:
|
||||
return next(
|
||||
(credential for credential in litellm.credential_list if credential.credential_name == credential_name),
|
||||
None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_credential_values(credential_name: str) -> dict:
|
||||
"""Safe accessor for credentials."""
|
||||
|
||||
if not litellm.credential_list:
|
||||
return {}
|
||||
for credential in litellm.credential_list:
|
||||
if credential.credential_name == credential_name:
|
||||
return credential.credential_values.copy()
|
||||
return {}
|
||||
credential: Final = CredentialAccessor.find_credential(credential_name)
|
||||
return {} if credential is None else credential.credential_values.copy()
|
||||
|
||||
@staticmethod
|
||||
def upsert_credentials(credentials: list[CredentialItem]):
|
||||
|
|
|
|||
|
|
@ -125,6 +125,7 @@ async def _resync_model_deployments(model_name: str) -> bool:
|
|||
)
|
||||
return proxy_server.llm_router is not None
|
||||
async with proxy_server.MODEL_RECONCILE_LOCK:
|
||||
await proxy_server.proxy_config.get_credentials(prisma_client=prisma_client)
|
||||
proxy_server.proxy_config._add_deployment(db_models=rows)
|
||||
proxy_server.llm_model_list = router.get_model_list()
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -46,17 +46,32 @@ def extract_sql_commands(diff_output: str) -> list[str]:
|
|||
|
||||
def check_prisma_schema_diff_helper(db_url: str) -> tuple[bool, list[str]]:
|
||||
"""Checks for differences between current database and Prisma schema.
|
||||
|
||||
Never raises: a diff that cannot be produced, because the runner is missing,
|
||||
because the command failed, or because it outlived its budget, is reported as
|
||||
"no diff" so boot continues.
|
||||
|
||||
Returns:
|
||||
A tuple containing:
|
||||
- A boolean indicating if differences were found (True) or not (False).
|
||||
- A string with the diff output or error message.
|
||||
Raises:
|
||||
subprocess.CalledProcessError: If the Prisma command fails.
|
||||
Exception: For any other errors during execution.
|
||||
- The SQL commands that would close the diff, empty when there is none.
|
||||
"""
|
||||
verbose_logger.debug("Checking for Prisma schema diff...")
|
||||
try:
|
||||
result: Final = subprocess.run(
|
||||
from litellm_proxy_extras.prisma_toolchain import (
|
||||
PRISMA_COMMAND_TIMEOUT_ENV_VAR,
|
||||
prisma_command_timeout,
|
||||
run_prisma,
|
||||
)
|
||||
except ImportError as e:
|
||||
print( # noqa: T201 # boot-time operator output, same channel as this helper's other messages
|
||||
f"Skipping the migration diff: litellm-proxy-extras has no Prisma runner. Error: {e}"
|
||||
)
|
||||
return False, []
|
||||
|
||||
verbose_logger.debug("Checking for Prisma schema diff...")
|
||||
timeout: Final = prisma_command_timeout()
|
||||
try:
|
||||
result: Final = run_prisma(
|
||||
[
|
||||
"prisma",
|
||||
"migrate",
|
||||
|
|
@ -67,12 +82,10 @@ def check_prisma_schema_diff_helper(db_url: str) -> tuple[bool, list[str]]:
|
|||
"./schema.prisma",
|
||||
"--script",
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=True,
|
||||
timeout=timeout,
|
||||
env=os.environ.copy(),
|
||||
)
|
||||
|
||||
# return True, "Migration diff generated successfully."
|
||||
sql_commands: Final = extract_sql_commands(result.stdout)
|
||||
|
||||
if sql_commands:
|
||||
|
|
@ -83,6 +96,12 @@ def check_prisma_schema_diff_helper(db_url: str) -> tuple[bool, list[str]]:
|
|||
return True, sql_commands
|
||||
else:
|
||||
return False, []
|
||||
except subprocess.TimeoutExpired:
|
||||
print( # noqa: T201 # boot-time operator output, same channel as this helper's other messages
|
||||
f"Timed out after {timeout}s generating the migration diff. "
|
||||
f"Raise {PRISMA_COMMAND_TIMEOUT_ENV_VAR} if this database needs longer."
|
||||
)
|
||||
return False, []
|
||||
except subprocess.CalledProcessError as e:
|
||||
error_message: Final = f"Failed to generate migration diff. Error: {e.stderr}"
|
||||
print(error_message) # noqa: T201
|
||||
|
|
|
|||
|
|
@ -937,9 +937,17 @@ class PrismaManager:
|
|||
use_v2_resolver=use_v2_resolver,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
from litellm_proxy_extras.prisma_toolchain import (
|
||||
prisma_command_timeout,
|
||||
run_prisma,
|
||||
)
|
||||
except ImportError as e:
|
||||
verbose_proxy_logger.error("\x1b[1;31mLiteLLM: Failed to import proxy extras. Got %s\x1b[0m", e)
|
||||
return False
|
||||
|
||||
PrismaManager._raise_if_partitioned_spend_logs()
|
||||
# Use prisma db push with increased timeout
|
||||
subprocess.run(
|
||||
run_prisma(
|
||||
[
|
||||
"prisma",
|
||||
"db",
|
||||
|
|
@ -947,13 +955,15 @@ class PrismaManager:
|
|||
"--accept-data-loss",
|
||||
"--skip-generate",
|
||||
],
|
||||
timeout=60,
|
||||
check=True,
|
||||
timeout=prisma_command_timeout(),
|
||||
env=os.environ.copy(),
|
||||
stdout=None,
|
||||
stderr=None,
|
||||
)
|
||||
PrismaManager._apply_replica_identity_full_if_requested()
|
||||
return True
|
||||
except subprocess.TimeoutExpired:
|
||||
verbose_proxy_logger.warning("Attempt %s timed out", attempt + 1)
|
||||
except subprocess.TimeoutExpired as e:
|
||||
verbose_proxy_logger.warning("Attempt %s timed out after %.0fs", attempt + 1, e.timeout)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
except subprocess.CalledProcessError as e:
|
||||
attempts_left = 3 - attempt
|
||||
|
|
|
|||
|
|
@ -7109,11 +7109,10 @@ class ProxyConfig:
|
|||
],
|
||||
)
|
||||
|
||||
# Only load models from DB if "models" is in supported_db_objects (or if supported_db_objects is not set)
|
||||
if self._should_load_db_object(object_type="models"):
|
||||
new_models: Final = await self._get_models_from_db(prisma_client=prisma_client)
|
||||
|
||||
# update llm router
|
||||
load_models: Final = self._should_load_db_object(object_type="models")
|
||||
new_models: Final = await self._get_models_from_db(prisma_client=prisma_client) if load_models else None
|
||||
await self.get_credentials(prisma_client=prisma_client)
|
||||
if load_models:
|
||||
still_desired_ids = await self._update_llm_router(
|
||||
new_models=new_models, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
|
|
@ -7153,12 +7152,9 @@ class ProxyConfig:
|
|||
async def _resync_config_from_db() -> None:
|
||||
await self.add_deployment(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj)
|
||||
|
||||
async def _resync_credentials_from_db() -> None:
|
||||
await self.get_credentials(prisma_client=prisma_client)
|
||||
|
||||
subscriber: Final = ConfigSyncSubscriber(
|
||||
redis_cache=redis_cache,
|
||||
resync_callbacks=(_resync_config_from_db, _resync_credentials_from_db),
|
||||
resync_callbacks=(_resync_config_from_db,),
|
||||
)
|
||||
self.config_sync_subscriber = subscriber
|
||||
subscriber.start()
|
||||
|
|
@ -8013,7 +8009,7 @@ class ProxyConfig:
|
|||
|
||||
async def get_credentials(self, prisma_client: PrismaClient):
|
||||
try:
|
||||
credentials = await CredentialsRepository(prisma_client).find_all()
|
||||
credentials = await CredentialsRepository(WriterPinnedClient(prisma_client.db)).find_all()
|
||||
credentials = [self.decrypt_credentials(cred) for cred in credentials]
|
||||
await self.delete_credentials(credentials) # delete credentials that are not in the all-up list
|
||||
CredentialAccessor.upsert_credentials(credentials) # upsert credentials that are in the all-up list
|
||||
|
|
@ -9597,19 +9593,6 @@ class ProxyStartupEvent:
|
|||
)
|
||||
|
||||
if store_model_in_db is True:
|
||||
### GET STORED CREDENTIALS ###
|
||||
scheduler.add_job(
|
||||
proxy_config.get_credentials,
|
||||
"interval",
|
||||
seconds=config_reload_interval_seconds,
|
||||
# REMOVED jitter parameter - major cause of memory leak
|
||||
args=[prisma_client],
|
||||
id="get_credentials_job",
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
await proxy_config.get_credentials(prisma_client=prisma_client)
|
||||
|
||||
# MEMORY LEAK FIX: Increase interval from 10s to 30s minimum
|
||||
# Frequent polling was causing excessive memory allocations
|
||||
scheduler.add_job(
|
||||
|
|
@ -9623,7 +9606,7 @@ class ProxyStartupEvent:
|
|||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
|
||||
# this will load all existing models on proxy startup
|
||||
# this will load all existing credentials and models on proxy startup
|
||||
await proxy_config.add_deployment(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj)
|
||||
|
||||
proxy_config.start_config_sync_subscriber(
|
||||
|
|
|
|||
|
|
@ -1248,7 +1248,7 @@ class Router:
|
|||
selector = LeastBusyLoggingHandler(router_cache=self.cache)
|
||||
if register_callbacks:
|
||||
if isinstance(litellm.input_callback, list):
|
||||
litellm.input_callback.append(selector)
|
||||
litellm.logging_callback_manager.add_litellm_input_callback(selector)
|
||||
else:
|
||||
litellm.input_callback = [selector]
|
||||
case RoutingStrategy.USAGE_BASED_ROUTING.value:
|
||||
|
|
@ -4214,10 +4214,12 @@ class Router:
|
|||
}
|
||||
)
|
||||
litellm_logging_object = cast(LiteLLMLogging, litellm_logging_object)
|
||||
prompt_management_deployment: Final = self.get_available_deployment(
|
||||
specific_deployment: Final = kwargs.pop("specific_deployment", None)
|
||||
prompt_management_deployment: Final = await self.async_get_available_deployment(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "prompt"}],
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
messages=cast(list[dict[str, str]], messages), # cast-ok: selection reads messages structurally
|
||||
specific_deployment=specific_deployment,
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
|
||||
self._update_kwargs_with_deployment(deployment=prompt_management_deployment, kwargs=kwargs)
|
||||
|
|
|
|||
|
|
@ -1,17 +1,103 @@
|
|||
#### What this does ####
|
||||
# identifies least busy deployment
|
||||
# How is this achieved?
|
||||
# - Before each call, have the router print the state of requests {"deployment": "requests_in_flight"}
|
||||
# - use litellm.input_callbacks to log when a request is just about to be made to a model - {"deployment-id": traffic}
|
||||
# - use litellm.success + failure callbacks to log when a request completed
|
||||
# - in get_available_deployment, for a given model group name -> pick based on traffic
|
||||
|
||||
import random
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
IN_FLIGHT_COUNT_TTL_SECONDS: Final = 60 * 60
|
||||
|
||||
|
||||
class _ModelInfo(TypedDict, total=False):
|
||||
id: ReadOnly[str | int | None]
|
||||
|
||||
|
||||
class _Metadata(TypedDict, total=False):
|
||||
model_group: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _LitellmParams(TypedDict, total=False):
|
||||
metadata: ReadOnly[_Metadata | None]
|
||||
model_info: ReadOnly[_ModelInfo | None]
|
||||
|
||||
|
||||
class _CallKwargs(TypedDict, total=False):
|
||||
litellm_params: ReadOnly[_LitellmParams | None]
|
||||
|
||||
|
||||
class _DeploymentModelInfo(TypedDict):
|
||||
id: ReadOnly[str | int]
|
||||
|
||||
|
||||
class _Deployment(TypedDict):
|
||||
model_info: ReadOnly[_DeploymentModelInfo]
|
||||
|
||||
|
||||
_CALL_KWARGS: Final = TypeAdapter(_CallKwargs)
|
||||
_DEPLOYMENTS: Final = TypeAdapter(list[_Deployment])
|
||||
_MEMORY_COUNTS: Final = TypeAdapter(tuple[float | None, ...] | None)
|
||||
|
||||
|
||||
def _request_count_key(model_group: str, deployment_id: str) -> str:
|
||||
return f"{model_group}_request_count:{deployment_id}"
|
||||
|
||||
|
||||
def _deployment_ref(kwargs: Mapping[str, object]) -> tuple[str, str] | None:
|
||||
try:
|
||||
call: Final = _CALL_KWARGS.validate_python(kwargs)
|
||||
except ValidationError:
|
||||
return None
|
||||
litellm_params: Final = call.get("litellm_params")
|
||||
metadata: Final = litellm_params.get("metadata") if litellm_params else None
|
||||
model_info: Final = litellm_params.get("model_info") if litellm_params else None
|
||||
model_group: Final = metadata.get("model_group") if metadata else None
|
||||
deployment_id: Final = model_info.get("id") if model_info else None
|
||||
if model_group is None or deployment_id is None:
|
||||
return None
|
||||
return model_group, str(deployment_id)
|
||||
|
||||
|
||||
def _request_count_keys(model_group: str, healthy_deployments: Sequence[Mapping[str, object]]) -> tuple[str, ...]:
|
||||
return tuple(
|
||||
_request_count_key(model_group, str(deployment["model_info"]["id"]))
|
||||
for deployment in _DEPLOYMENTS.validate_python(healthy_deployments)
|
||||
)
|
||||
|
||||
|
||||
def _as_counts(values: Sequence[float | None]) -> tuple[int, ...]:
|
||||
return tuple(0 if value is None else int(value) for value in values)
|
||||
|
||||
|
||||
def _local_counts(raw: object, keys: tuple[str, ...]) -> tuple[int, ...]:
|
||||
values: Final = _MEMORY_COUNTS.validate_python(raw)
|
||||
if values is None or len(values) != len(keys):
|
||||
return (0,) * len(keys)
|
||||
return _as_counts(values)
|
||||
|
||||
|
||||
def _least_busy(
|
||||
healthy_deployments: Sequence[Mapping[str, object]], counts: tuple[int, ...]
|
||||
) -> Mapping[str, object] | None:
|
||||
if not healthy_deployments:
|
||||
return None
|
||||
return healthy_deployments[min(range(len(healthy_deployments)), key=lambda index: counts[index])]
|
||||
|
||||
|
||||
def _warn_unreadable(model_group: str, error: Exception) -> None:
|
||||
verbose_router_logger.warning(
|
||||
"least-busy routing could not read the shared in-flight counts for %s, "
|
||||
"falling back to this worker's own counts: %s",
|
||||
model_group,
|
||||
error,
|
||||
)
|
||||
|
||||
|
||||
def _warn_unwritable(key: str, error: Exception) -> None:
|
||||
verbose_router_logger.warning("least-busy routing could not update the in-flight count under %s: %s", key, error)
|
||||
|
||||
|
||||
class LeastBusyLoggingHandler(CustomLogger):
|
||||
test_flag: bool = False
|
||||
|
|
@ -20,195 +106,101 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
|
||||
def __init__(self, router_cache: DualCache):
|
||||
self.router_cache = router_cache
|
||||
self.router_cache_id = str(id(router_cache))
|
||||
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
"""
|
||||
Log when a model is being used.
|
||||
def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None:
|
||||
self._increment(kwargs, 1)
|
||||
|
||||
Caching based on model group.
|
||||
"""
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
def log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
self._increment(kwargs, -1)
|
||||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# update cache
|
||||
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
|
||||
request_count_dict[id] = request_count_dict.get(id, 0) + 1
|
||||
def log_failure_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
self._increment(kwargs, -1)
|
||||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
|
||||
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
except Exception:
|
||||
pass
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
await self._async_increment(kwargs, -1)
|
||||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# decrement count in cache
|
||||
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
|
||||
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
|
||||
if request_count_value is None:
|
||||
return
|
||||
request_count_dict[id] = request_count_value - 1
|
||||
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
|
||||
### TESTING ###
|
||||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# decrement count in cache
|
||||
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
|
||||
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
|
||||
if request_count_value is None:
|
||||
return
|
||||
request_count_dict[id] = request_count_value - 1
|
||||
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
|
||||
### TESTING ###
|
||||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# decrement count in cache
|
||||
request_count_dict: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
|
||||
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
|
||||
if request_count_value is None:
|
||||
return
|
||||
request_count_dict[id] = request_count_value - 1
|
||||
await self.router_cache.async_set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
|
||||
### TESTING ###
|
||||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# decrement count in cache
|
||||
request_count_dict: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
|
||||
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
|
||||
if request_count_value is None:
|
||||
return
|
||||
request_count_dict[id] = request_count_value - 1
|
||||
await self.router_cache.async_set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
|
||||
### TESTING ###
|
||||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _get_available_deployments(
|
||||
self,
|
||||
healthy_deployments: list,
|
||||
all_deployments: dict,
|
||||
):
|
||||
"""
|
||||
Helper to get deployments using least busy strategy
|
||||
"""
|
||||
for d in healthy_deployments:
|
||||
## if healthy deployment not yet used
|
||||
if d["model_info"]["id"] not in all_deployments:
|
||||
all_deployments[d["model_info"]["id"]] = 0
|
||||
# map deployment to id
|
||||
# pick least busy deployment
|
||||
min_traffic = float("inf")
|
||||
min_deployment = None
|
||||
for k, v in all_deployments.items():
|
||||
if v < min_traffic:
|
||||
min_traffic = v
|
||||
min_deployment = k
|
||||
if min_deployment is not None:
|
||||
## check if min deployment is a string, if so, cast it to int
|
||||
for m in healthy_deployments:
|
||||
if m["model_info"]["id"] == min_deployment:
|
||||
return m
|
||||
min_deployment = random.choice(healthy_deployments)
|
||||
else:
|
||||
min_deployment = random.choice(healthy_deployments)
|
||||
return min_deployment
|
||||
async def async_log_failure_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
await self._async_increment(kwargs, -1)
|
||||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
|
||||
def get_available_deployments(
|
||||
self,
|
||||
model_group: str,
|
||||
healthy_deployments: list,
|
||||
):
|
||||
"""
|
||||
Sync helper to get deployments using least busy strategy
|
||||
"""
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
all_deployments: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
|
||||
return self._get_available_deployments(
|
||||
healthy_deployments=healthy_deployments,
|
||||
all_deployments=all_deployments,
|
||||
)
|
||||
self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]]
|
||||
) -> Mapping[str, object] | None:
|
||||
keys: Final = _request_count_keys(model_group, healthy_deployments)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
shared: Final = _as_counts(redis_cache.batch_get_counts(list(keys)))
|
||||
except Exception as e:
|
||||
_warn_unreadable(model_group, e)
|
||||
else:
|
||||
return _least_busy(healthy_deployments, shared)
|
||||
local: Final = _local_counts(self.router_cache.batch_get_cache(list(keys), local_only=True), keys)
|
||||
return _least_busy(healthy_deployments, local)
|
||||
|
||||
async def async_get_available_deployments(self, model_group: str, healthy_deployments: list):
|
||||
"""
|
||||
Async helper to get deployments using least busy strategy
|
||||
"""
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
all_deployments: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
|
||||
return self._get_available_deployments(
|
||||
healthy_deployments=healthy_deployments,
|
||||
all_deployments=all_deployments,
|
||||
)
|
||||
async def async_get_available_deployments(
|
||||
self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]]
|
||||
) -> Mapping[str, object] | None:
|
||||
keys: Final = _request_count_keys(model_group, healthy_deployments)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
shared: Final = _as_counts(await redis_cache.async_batch_get_counts(list(keys)))
|
||||
except Exception as e:
|
||||
_warn_unreadable(model_group, e)
|
||||
else:
|
||||
return _least_busy(healthy_deployments, shared)
|
||||
local: Final = _local_counts(await self.router_cache.async_batch_get_cache(list(keys), local_only=True), keys)
|
||||
return _least_busy(healthy_deployments, local)
|
||||
|
||||
def _increment(self, kwargs: Mapping[str, object], delta: int) -> None:
|
||||
ref: Final = _deployment_ref(kwargs)
|
||||
if ref is None:
|
||||
return
|
||||
key: Final = _request_count_key(*ref)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
try:
|
||||
local: Final = self.router_cache.increment_cache(
|
||||
key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS
|
||||
)
|
||||
if local < 0:
|
||||
self.router_cache.set_cache(key, 0, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
if redis_cache is None:
|
||||
return
|
||||
redis_cache.increment_with_floor(key, delta, IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
except Exception as e:
|
||||
_warn_unwritable(key, e)
|
||||
|
||||
async def _async_increment(self, kwargs: Mapping[str, object], delta: int) -> None:
|
||||
ref: Final = _deployment_ref(kwargs)
|
||||
if ref is None:
|
||||
return
|
||||
key: Final = _request_count_key(*ref)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
try:
|
||||
local: Final = await self.router_cache.async_increment_cache(
|
||||
key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS
|
||||
)
|
||||
if local is not None and local < 0:
|
||||
await self.router_cache.async_set_cache(key, 0, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
if redis_cache is None:
|
||||
return
|
||||
await redis_cache.async_increment_with_floor(key, delta, IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
except Exception as e:
|
||||
_warn_unwritable(key, e)
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing_extensions import TypedDict
|
|||
from litellm import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -36,10 +37,19 @@ _MAX_CORRECTED_IN_MEMORY_TTL_SECONDS: Final = 60.0
|
|||
|
||||
|
||||
class CooldownCache:
|
||||
def __init__(self, cache: DualCache, default_cooldown_time: float):
|
||||
def __init__(
|
||||
self,
|
||||
cache: DualCache,
|
||||
default_cooldown_time: float,
|
||||
redis_read_interval_seconds: float = DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS,
|
||||
):
|
||||
self.cache = cache
|
||||
self.default_cooldown_time = default_cooldown_time
|
||||
self.in_memory_cache = InMemoryCache()
|
||||
self._cooldown_store = DualCache(
|
||||
in_memory_cache=self.in_memory_cache,
|
||||
default_redis_batch_cache_expiry=redis_read_interval_seconds,
|
||||
)
|
||||
# Initialize the masker with custom settings for exception strings
|
||||
self.exception_masker = SensitiveDataMasker(
|
||||
visible_prefix=50, # Show first 50 characters
|
||||
|
|
@ -48,6 +58,21 @@ class CooldownCache:
|
|||
mask_short_values=False, # Truncate long messages only; keep short ones readable
|
||||
)
|
||||
|
||||
@property
|
||||
def cooldown_store(self) -> DualCache:
|
||||
"""
|
||||
The cache cooldown entries live in, with the router's Redis attached on first use.
|
||||
|
||||
It is kept separate from the router-wide cache so that a key missing from memory is
|
||||
re-read from Redis every `redis_read_interval_seconds` rather than on the router
|
||||
cache's much longer batch interval, which is what lets a sibling replica see a
|
||||
cooldown another replica wrote, and so that unrelated router keys cannot evict a
|
||||
cooldown from the in-memory tier before it expires. Redis is attached lazily because
|
||||
the router builds its cooldown cache before it wires up the shared Redis client.
|
||||
"""
|
||||
self._cooldown_store.attach_redis_cache(self.cache.redis_cache)
|
||||
return self._cooldown_store
|
||||
|
||||
def _common_add_cooldown_logic(
|
||||
self, model_id: str, original_exception, exception_status, cooldown_time: float
|
||||
) -> tuple[str, CooldownCacheValue]:
|
||||
|
|
@ -93,7 +118,7 @@ class CooldownCache:
|
|||
)
|
||||
|
||||
# Set the cache with a TTL equal to the cooldown time
|
||||
self.cache.set_cache(
|
||||
self.cooldown_store.set_cache(
|
||||
value=cooldown_data,
|
||||
key=cooldown_key,
|
||||
ttl=_cooldown_time,
|
||||
|
|
@ -122,13 +147,13 @@ class CooldownCache:
|
|||
cooldown_cache_value: Final = CooldownCacheValue(**result) # pyright: ignore[reportUnknownArgumentType] - result comes from an untyped cache read, not from our own code
|
||||
remaining: Final = (cooldown_cache_value["timestamp"] + cooldown_cache_value["cooldown_time"]) - current_time
|
||||
if remaining <= 0:
|
||||
self.cache.in_memory_cache.delete_cache(key)
|
||||
self.in_memory_cache.delete_cache(key)
|
||||
return None
|
||||
current_expiry: Final = self.cache.in_memory_cache.ttl_dict.get(key)
|
||||
current_expiry: Final = self.in_memory_cache.ttl_dict.get(key)
|
||||
if current_expiry is not None and current_expiry > current_time + remaining + 5:
|
||||
corrected_ttl: Final = min(remaining, _MAX_CORRECTED_IN_MEMORY_TTL_SECONDS)
|
||||
self.cache.in_memory_cache.delete_cache(key)
|
||||
self.cache.in_memory_cache.set_cache(key, result, ttl=corrected_ttl)
|
||||
self.in_memory_cache.delete_cache(key)
|
||||
self.in_memory_cache.set_cache(key, result, ttl=corrected_ttl)
|
||||
return cooldown_cache_value
|
||||
|
||||
async def async_get_active_cooldowns(
|
||||
|
|
@ -137,12 +162,7 @@ class CooldownCache:
|
|||
# Generate the keys for the deployments
|
||||
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
|
||||
|
||||
# Retrieve the values for the keys using mget
|
||||
## more likely to be none if no models ratelimited. So just check redis every 1s
|
||||
## each redis call adds ~100ms latency.
|
||||
|
||||
## check in memory cache first
|
||||
results: Final = await self.cache.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
|
||||
results: Final = await self.cooldown_store.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
|
||||
active_cooldowns: Final[list[tuple[str, CooldownCacheValue]]] = []
|
||||
|
||||
if results is None or all(v is None for v in results):
|
||||
|
|
@ -164,7 +184,7 @@ class CooldownCache:
|
|||
# Generate the keys for the deployments
|
||||
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
|
||||
# Retrieve the values for the keys using mget
|
||||
results: Final = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
|
||||
active_cooldowns: Final = []
|
||||
current_time: Final = time.time()
|
||||
|
|
@ -184,7 +204,7 @@ class CooldownCache:
|
|||
keys: Final = [f"deployment:{model_id}:cooldown" for model_id in model_ids]
|
||||
|
||||
# Retrieve the values for the keys using mget
|
||||
results: Final = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
|
||||
min_cooldown_time: float | None = None
|
||||
# Process the results
|
||||
|
|
|
|||
|
|
@ -672,11 +672,19 @@ def load_credentials_from_list(kwargs: dict):
|
|||
CredentialAccessor: Final = getattr(sys.modules[__name__], "CredentialAccessor")
|
||||
|
||||
credential_name: Final = kwargs.get("litellm_credential_name")
|
||||
if credential_name and litellm.credential_list:
|
||||
credential_accessor: Final[Mapping[str, object]] = CredentialAccessor.get_credential_values(credential_name)
|
||||
for key, value in credential_accessor.items():
|
||||
if key not in kwargs:
|
||||
kwargs[key] = value
|
||||
if not credential_name:
|
||||
return
|
||||
credential: Final = CredentialAccessor.find_credential(credential_name)
|
||||
if credential is None:
|
||||
verbose_logger.warning(
|
||||
"litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it",
|
||||
credential_name,
|
||||
len(litellm.credential_list),
|
||||
)
|
||||
return
|
||||
for key, value in credential.credential_values.items():
|
||||
if key not in kwargs:
|
||||
kwargs[key] = value
|
||||
|
||||
|
||||
def get_dynamic_callbacks(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 2956
|
||||
"limit": 2918
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
|
|
@ -9,10 +9,10 @@
|
|||
"limit": 806
|
||||
},
|
||||
"ANN201": {
|
||||
"limit": 1979
|
||||
"limit": 1965
|
||||
},
|
||||
"ANN202": {
|
||||
"limit": 831
|
||||
"limit": 829
|
||||
},
|
||||
"ANN204": {
|
||||
"limit": 683
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 3
|
||||
},
|
||||
"BLE001": {
|
||||
"limit": 2916
|
||||
"limit": 2914
|
||||
},
|
||||
"C401": {
|
||||
"limit": 8
|
||||
|
|
@ -189,7 +189,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"S110": {
|
||||
"limit": 217
|
||||
"limit": 207
|
||||
},
|
||||
"S112": {
|
||||
"limit": 22
|
||||
|
|
|
|||
|
|
@ -33,8 +33,8 @@ def test_model_added():
|
|||
}
|
||||
}
|
||||
least_busy_logger.log_pre_api_call(model="test", messages=[], kwargs=kwargs)
|
||||
request_count_api_key = f"gpt-3.5-turbo_request_count"
|
||||
assert test_cache.get_cache(key=request_count_api_key) is not None
|
||||
request_count_api_key = "gpt-3.5-turbo_request_count:1234"
|
||||
assert test_cache.get_cache(key=request_count_api_key) == 1
|
||||
|
||||
|
||||
def test_get_available_deployments():
|
||||
|
|
@ -52,8 +52,8 @@ def test_get_available_deployments():
|
|||
}
|
||||
}
|
||||
least_busy_logger.log_pre_api_call(model="test", messages=[], kwargs=kwargs)
|
||||
request_count_api_key = f"{model_group}_request_count"
|
||||
assert test_cache.get_cache(key=request_count_api_key) is not None
|
||||
request_count_api_key = f"{model_group}_request_count:1234"
|
||||
assert test_cache.get_cache(key=request_count_api_key) == 1
|
||||
|
||||
|
||||
# test_get_available_deployments()
|
||||
|
|
@ -104,15 +104,20 @@ async def test_router_get_available_deployments(async_test):
|
|||
router.leastbusy_logger.test_flag = True
|
||||
|
||||
model_group = "azure-model"
|
||||
request_count_dict = {1: 10, 2: 54, 3: 100}
|
||||
cache_key = f"{model_group}_request_count"
|
||||
request_count_dict = {"1": 10, "2": 54, "3": 100}
|
||||
cache_keys = {
|
||||
deployment_id: f"{model_group}_request_count:{deployment_id}"
|
||||
for deployment_id in request_count_dict
|
||||
}
|
||||
if async_test is True:
|
||||
await router.cache.async_set_cache(key=cache_key, value=request_count_dict)
|
||||
for deployment_id, count in request_count_dict.items():
|
||||
await router.cache.async_set_cache(key=cache_keys[deployment_id], value=count)
|
||||
deployment = await router.async_get_available_deployment(
|
||||
model=model_group, messages=None, request_kwargs={}
|
||||
)
|
||||
else:
|
||||
router.cache.set_cache(key=cache_key, value=request_count_dict)
|
||||
for deployment_id, count in request_count_dict.items():
|
||||
router.cache.set_cache(key=cache_keys[deployment_id], value=count)
|
||||
deployment = router.get_available_deployment(model=model_group, messages=None)
|
||||
print(f"deployment: {deployment}")
|
||||
assert deployment["model_info"]["id"] == "1"
|
||||
|
|
@ -124,15 +129,18 @@ async def test_router_get_available_deployments(async_test):
|
|||
messages=[{"role": "user", "content": "Hey, how's it going?"}],
|
||||
)
|
||||
|
||||
return_dict = router.cache.get_cache(key=cache_key)
|
||||
|
||||
# wait 2 seconds
|
||||
time.sleep(2)
|
||||
|
||||
return_dict = {
|
||||
deployment_id: router.cache.get_cache(key=cache_key)
|
||||
for deployment_id, cache_key in cache_keys.items()
|
||||
}
|
||||
|
||||
assert router.leastbusy_logger.logged_success == 1
|
||||
assert return_dict[1] == 10
|
||||
assert return_dict[2] == 54
|
||||
assert return_dict[3] == 100
|
||||
assert return_dict["1"] == 10
|
||||
assert return_dict["2"] == 54
|
||||
assert return_dict["3"] == 100
|
||||
|
||||
|
||||
## Test with Real calls ##
|
||||
|
|
@ -192,9 +200,11 @@ async def test_router_atext_completion_streaming():
|
|||
await asyncio.sleep(random.uniform(0, 2))
|
||||
await router.atext_completion(model=model, prompt=prompt, stream=True)
|
||||
|
||||
cache_key = f"{model}_request_count"
|
||||
## check if calls equally distributed
|
||||
cache_dict = router.cache.get_cache(key=cache_key)
|
||||
cache_dict = {
|
||||
deployment_id: router.cache.get_cache(key=f"{model}_request_count:{deployment_id}")
|
||||
for deployment_id in ("1", "2", "3")
|
||||
}
|
||||
for k, v in cache_dict.items():
|
||||
assert v == 1, f"Failed. K={k} called v={v} times, cache_dict={cache_dict}"
|
||||
|
||||
|
|
@ -259,8 +269,10 @@ async def test_router_completion_streaming():
|
|||
await asyncio.sleep(random.uniform(0, 2))
|
||||
await router.acompletion(model=model, messages=messages, stream=True)
|
||||
|
||||
cache_key = f"{model}_request_count"
|
||||
## check if calls equally distributed
|
||||
cache_dict = router.cache.get_cache(key=cache_key)
|
||||
cache_dict = {
|
||||
deployment_id: router.cache.get_cache(key=f"{model}_request_count:{deployment_id}")
|
||||
for deployment_id in ("1", "2", "3")
|
||||
}
|
||||
for k, v in cache_dict.items():
|
||||
assert v == 1, f"Failed. K={k} called v={v} times, cache_dict={cache_dict}"
|
||||
|
|
|
|||
80
tests/local_testing/test_redis_increment_with_floor.py
Normal file
80
tests/local_testing/test_redis_increment_with_floor.py
Normal file
|
|
@ -0,0 +1,80 @@
|
|||
"""Least-busy routing keeps its in-flight counters in Redis, and the clamp at zero plus the
|
||||
create-once TTL both live inside a Lua script. Nothing but a real Redis runs that script, so
|
||||
these are the only tests that fail when the script itself is wrong."""
|
||||
|
||||
import os
|
||||
import uuid
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
TTL: Final = 600
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def counter():
|
||||
cache: Final = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT"))
|
||||
key: Final = f"lit7039-{uuid.uuid4()}"
|
||||
yield cache, key, cache.check_and_fix_namespace(key=key)
|
||||
cache.delete_cache(key)
|
||||
|
||||
|
||||
def test_a_counter_adds_every_increment_and_reads_back_what_it_holds(counter):
|
||||
cache, key, _ = counter
|
||||
|
||||
assert cache.increment_with_floor(key, 3, TTL) == 3
|
||||
assert cache.increment_with_floor(key, 2, TTL) == 5
|
||||
assert cache.batch_get_counts([key]) == (5,)
|
||||
|
||||
|
||||
def test_a_decrement_past_zero_leaves_the_counter_at_zero(counter):
|
||||
"""A worker whose counter expired mid-request decrements a key that is no longer there.
|
||||
Without the clamp that deployment reads negative, and least-busy pins every later request
|
||||
on it until the count climbs back to zero."""
|
||||
cache, key, _ = counter
|
||||
|
||||
assert cache.increment_with_floor(key, 1, TTL) == 1
|
||||
assert cache.increment_with_floor(key, -5, TTL) == 0
|
||||
assert cache.batch_get_counts([key]) == (0,)
|
||||
|
||||
|
||||
def test_traffic_never_pushes_a_counters_expiry_back_out(counter):
|
||||
"""The TTL is what releases a count whose worker died mid-request. Rewriting it on every
|
||||
touch would keep that stuck count alive for as long as the group takes traffic."""
|
||||
cache, key, namespaced_key = counter
|
||||
|
||||
cache.increment_with_floor(key, 1, TTL)
|
||||
assert cache.redis_client.ttl(namespaced_key) > TTL - 60
|
||||
|
||||
cache.redis_client.expire(namespaced_key, 30)
|
||||
cache.increment_with_floor(key, 1, TTL)
|
||||
|
||||
assert cache.redis_client.ttl(namespaced_key) <= 30
|
||||
|
||||
|
||||
def test_clamping_to_zero_keeps_the_expiry_it_already_had(counter):
|
||||
cache, key, namespaced_key = counter
|
||||
|
||||
cache.increment_with_floor(key, 1, TTL)
|
||||
cache.redis_client.expire(namespaced_key, 30)
|
||||
|
||||
assert cache.increment_with_floor(key, -5, TTL) == 0
|
||||
assert cache.redis_client.ttl(namespaced_key) <= 30
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_counter_behaves_the_same_way(counter):
|
||||
cache, key, namespaced_key = counter
|
||||
|
||||
assert await cache.async_increment_with_floor(key, 2, TTL) == 2
|
||||
assert await cache.async_batch_get_counts([key]) == (2,)
|
||||
|
||||
cache.redis_client.expire(namespaced_key, 30)
|
||||
|
||||
assert await cache.async_increment_with_floor(key, -9, TTL) == 0
|
||||
assert cache.redis_client.ttl(namespaced_key) <= 30
|
||||
|
|
@ -246,12 +246,12 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time() - 120.0,
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
cc.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
|
||||
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
assert active == [], "Expired cooldown entry must not appear in active cooldowns"
|
||||
assert cc.cache.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
|
||||
assert cc.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
|
||||
|
||||
def test_active_entry_is_returned(self):
|
||||
"""
|
||||
|
|
@ -267,7 +267,7 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time(),
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60)
|
||||
cc.in_memory_cache.set_cache(key, active_value, ttl=60)
|
||||
|
||||
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
|
|
@ -290,14 +290,14 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time() - (60.0 - remaining),
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, value, ttl=600)
|
||||
cc.in_memory_cache.set_cache(key, value, ttl=600)
|
||||
|
||||
before_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
before_expiry = cc.in_memory_cache.ttl_dict.get(key)
|
||||
assert before_expiry is not None
|
||||
|
||||
cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
after_expiry = cc.in_memory_cache.ttl_dict.get(key)
|
||||
assert after_expiry is not None
|
||||
corrected_remaining = after_expiry - time.time()
|
||||
assert corrected_remaining <= 60.0, "Corrected TTL must not exceed 60s"
|
||||
|
|
@ -318,12 +318,12 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time() - 120.0,
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
cc.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
|
||||
active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
assert active == [], "Expired entry must not appear in async active cooldowns"
|
||||
assert cc.cache.in_memory_cache.get_cache(key) is None
|
||||
assert cc.in_memory_cache.get_cache(key) is None
|
||||
|
||||
|
||||
class TestFallbackDeploymentCooldown:
|
||||
|
|
|
|||
|
|
@ -525,6 +525,50 @@ def test_circuit_breaker_open_keeps_sync_batch_get_cache_as_a_miss(sync_batch_re
|
|||
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"]) == {}
|
||||
|
||||
|
||||
def test_batch_get_counts_raises_where_batch_get_cache_reports_a_miss(sync_batch_redis_cache):
|
||||
"""A caller that must fall back when Redis is unreachable needs the failure, not zeros.
|
||||
|
||||
The batch read answers a dead Redis with an empty dict, which a counting caller cannot tell
|
||||
apart from "every counter is unset". Least-busy routing read that as an idle deployment and
|
||||
kept sending traffic to it instead of falling back to this worker's own in-flight counts.
|
||||
"""
|
||||
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit7039"]) == {}
|
||||
|
||||
with pytest.raises(OSError, match="redis unavailable"):
|
||||
sync_batch_redis_cache.batch_get_counts(["lit7039"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_batch_get_counts_raises_where_async_batch_get_cache_reports_a_miss(redis_no_ping: None):
|
||||
"""Async twin: the async batch read hides the same failure behind an empty dict."""
|
||||
failing_client = AsyncMock()
|
||||
failing_client.mget.side_effect = OSError("redis unavailable")
|
||||
with patch( # test-quality-ok: RedisCache.__init__ builds its client eagerly, with no injection point
|
||||
"litellm._redis.get_redis_client", return_value=MagicMock()
|
||||
):
|
||||
cache = RedisCache(host="127.0.0.1", port=6379)
|
||||
|
||||
with patch.object(cache, "init_async_client", return_value=failing_client):
|
||||
assert await cache.async_batch_get_cache(key_list=["lit7039"]) == {}
|
||||
|
||||
with pytest.raises(OSError, match="redis unavailable"):
|
||||
await cache.async_batch_get_counts(["lit7039"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stored", [b"3", "3"])
|
||||
def test_batch_get_counts_reads_counters_in_order_and_keeps_unset_keys_apart(stored, redis_no_ping: None):
|
||||
"""Counters come back positionally, so an unset key has to stay a hole rather than shift the
|
||||
rest of the row onto the wrong deployments, and a count has to survive whether the client
|
||||
hands it back as bytes or as text."""
|
||||
with patch( # test-quality-ok: RedisCache.__init__ builds its client eagerly, with no injection point
|
||||
"litellm._redis.get_redis_client", return_value=MagicMock()
|
||||
):
|
||||
cache = RedisCache(host="127.0.0.1", port=6379)
|
||||
cache.redis_client.mget.return_value = [stored, None, b"0"]
|
||||
|
||||
assert cache.batch_get_counts(["dep-a", "dep-b", "dep-c"]) == (3, None, 0)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sync_batch_cache_with_service_logger(redis_no_ping: None) -> Iterator[tuple[RedisCache, ServiceLogging]]:
|
||||
service_logger = ServiceLogging(mock_testing=True)
|
||||
|
|
|
|||
|
|
@ -824,7 +824,7 @@ class _StopFailingSubscriber(ConfigSyncSubscriber):
|
|||
raise RuntimeError("stop failed")
|
||||
|
||||
|
||||
async def test_proxy_config_subscriber_resyncs_deployments_and_credentials() -> None:
|
||||
async def test_proxy_config_subscriber_resyncs_deployments_only() -> None:
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([_QueuePubSub()]))
|
||||
|
|
@ -852,10 +852,7 @@ async def test_proxy_config_subscriber_resyncs_deployments_and_credentials() ->
|
|||
await callback()
|
||||
await config.stop_config_sync_subscriber()
|
||||
|
||||
assert calls == [
|
||||
("add_deployment", prisma_client, proxy_logging_obj),
|
||||
("get_credentials", prisma_client, None),
|
||||
]
|
||||
assert calls == [("add_deployment", prisma_client, proxy_logging_obj)]
|
||||
assert config.config_sync_subscriber is None
|
||||
assert subscriber._task is None
|
||||
|
||||
|
|
|
|||
|
|
@ -393,6 +393,51 @@ async def test_resync_model_deployments_mutates_router_under_model_reconcile_loc
|
|||
assert not proxy_server.MODEL_RECONCILE_LOCK.locked()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resync_model_deployments_loads_db_credentials_before_reconciling_models(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import litellm
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.proxy.common_utils.registry_read_through import _resync_model_deployments
|
||||
from litellm.types.utils import CredentialItem
|
||||
|
||||
rows: Final = [MagicMock()]
|
||||
prisma_client: Final = MagicMock()
|
||||
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=rows)
|
||||
router: Final = MagicMock()
|
||||
router.get_model_list.return_value = []
|
||||
installed: Final = MagicMock()
|
||||
|
||||
async def load_credentials_from_db(prisma_client: object) -> None:
|
||||
CredentialAccessor.upsert_credentials(
|
||||
[
|
||||
CredentialItem(
|
||||
credential_name="openai-cred",
|
||||
credential_values={"api_key": "sk-from-db"},
|
||||
credential_info={},
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
def install_models(db_models: object) -> None:
|
||||
installed(db_models=db_models, credential=CredentialAccessor.get_credential_values("openai-cred"))
|
||||
|
||||
monkeypatch.setattr(litellm, "credential_list", [])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(proxy_server, "llm_model_list", None)
|
||||
monkeypatch.setattr(proxy_server.proxy_config, "get_credentials", load_credentials_from_db)
|
||||
monkeypatch.setattr(proxy_server.proxy_config, "_add_deployment", install_models)
|
||||
|
||||
assert await _resync_model_deployments("model-created-on-a-sibling-replica") is True
|
||||
installed.assert_called_once_with(db_models=rows, credential={"api_key": "sk-from-db"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resync_model_deployments_respects_supported_db_objects(monkeypatch):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ from litellm.proxy.common_utils.scheduled_job_stagger import (
|
|||
)
|
||||
|
||||
OPERATOR_CRON_JOB_ID = "spend_log_cleanup_job"
|
||||
SHARED_INTERVAL_JOB_IDS = ("periodic_reload_job", "get_credentials_job", "add_deployment_job")
|
||||
SHARED_INTERVAL_JOB_IDS = ("periodic_reload_job", "add_deployment_job")
|
||||
|
||||
|
||||
async def _noop() -> None: ...
|
||||
|
|
|
|||
|
|
@ -1,5 +1,11 @@
|
|||
import json
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Generator
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
|
|
@ -75,3 +81,76 @@ def reset_entra_token_provider_cache() -> Generator[None, None, None]:
|
|||
def unset_database_url(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("DATABASE_URL", "about-to-be-unset")
|
||||
monkeypatch.delenv("DATABASE_URL")
|
||||
|
||||
|
||||
FAKE_PRISMA_CLI = """#!{python}
|
||||
import json
|
||||
import os
|
||||
import pathlib
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
calls_file = pathlib.Path(os.environ["FAKE_PRISMA_CALLS"])
|
||||
earlier_calls = calls_file.read_text().splitlines() if calls_file.exists() else []
|
||||
with calls_file.open("a") as log:
|
||||
print(json.dumps(sys.argv[1:]), file=log)
|
||||
if not earlier_calls and os.environ.get("FAKE_PRISMA_HANG_FIRST"):
|
||||
grandchild = subprocess.Popen([sys.executable, "-c", "import time; time.sleep(600)"])
|
||||
pathlib.Path(os.environ["FAKE_PRISMA_GRANDCHILD_PIDFILE"]).write_text(str(grandchild.pid))
|
||||
time.sleep(600)
|
||||
sys.exit(0)
|
||||
"""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FakePrismaCli:
|
||||
"""A stand-in `prisma` on PATH, recording every invocation.
|
||||
|
||||
With FAKE_PRISMA_HANG_FIRST set it hangs on its first call from a process tree
|
||||
of its own, the way the real CLI wraps Node around a Rust schema engine, so a
|
||||
timeout that kills only the direct child leaves the rest of that tree running.
|
||||
"""
|
||||
|
||||
calls_file: Path
|
||||
grandchild_pidfile: Path
|
||||
|
||||
@property
|
||||
def calls(self) -> list[list[str]]:
|
||||
if not self.calls_file.exists():
|
||||
return []
|
||||
return [json.loads(line) for line in self.calls_file.read_text().splitlines()]
|
||||
|
||||
def grandchild_is_gone(self, within_seconds: float) -> bool:
|
||||
deadline = time.monotonic() + within_seconds
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
os.kill(int(self.grandchild_pidfile.read_text()), 0)
|
||||
except ProcessLookupError:
|
||||
return True
|
||||
time.sleep(0.05)
|
||||
return False
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_prisma_cli(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Generator[FakePrismaCli, None, None]:
|
||||
bin_dir = tmp_path / "fakebin"
|
||||
bin_dir.mkdir()
|
||||
script = bin_dir / "prisma"
|
||||
script.write_text(FAKE_PRISMA_CLI.format(python=sys.executable))
|
||||
script.chmod(0o755)
|
||||
cli = FakePrismaCli(
|
||||
calls_file=tmp_path / "calls.jsonl",
|
||||
grandchild_pidfile=tmp_path / "grandchild.pid",
|
||||
)
|
||||
monkeypatch.setenv("PATH", f"{bin_dir}{os.pathsep}{os.environ['PATH']}")
|
||||
monkeypatch.setenv("FAKE_PRISMA_CALLS", str(cli.calls_file))
|
||||
monkeypatch.setenv("FAKE_PRISMA_GRANDCHILD_PIDFILE", str(cli.grandchild_pidfile))
|
||||
monkeypatch.setenv("LITELLM_PRISMA_COMMAND_TIMEOUT", "1")
|
||||
monkeypatch.delenv("FAKE_PRISMA_HANG_FIRST", raising=False)
|
||||
yield cli
|
||||
if cli.grandchild_pidfile.exists():
|
||||
try:
|
||||
os.kill(int(cli.grandchild_pidfile.read_text()), signal.SIGKILL)
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -37,3 +37,35 @@ def test_check_migration_out_of_sync(mocker):
|
|||
check_migration.verbose_logger.exception.assert_called_once()
|
||||
actual_message = check_migration.verbose_logger.exception.call_args[0][0]
|
||||
assert "prisma schema out of sync with db" in actual_message
|
||||
|
||||
|
||||
@pytest.mark.timeout(30)
|
||||
def test_migrate_diff_stops_at_its_budget_and_takes_its_process_tree_with_it(fake_prisma_cli, monkeypatch):
|
||||
"""
|
||||
`prisma migrate diff` ran unbounded, so a database that never answers hung boot
|
||||
before uvicorn ever started, and interrupting the proxy orphaned the schema engine.
|
||||
"""
|
||||
from litellm.proxy.db.check_migration import check_prisma_schema_diff_helper
|
||||
|
||||
monkeypatch.setenv("FAKE_PRISMA_HANG_FIRST", "1")
|
||||
|
||||
assert check_prisma_schema_diff_helper("postgresql://u:p@localhost:9/x") == (False, [])
|
||||
assert fake_prisma_cli.calls == [
|
||||
["migrate", "diff", "--from-url", "postgresql://u:p@localhost:9/x",
|
||||
"--to-schema-datamodel", "./schema.prisma", "--script"]
|
||||
]
|
||||
assert fake_prisma_cli.grandchild_is_gone(within_seconds=5)
|
||||
|
||||
|
||||
def test_migrate_diff_without_the_prisma_runner_skips_instead_of_crashing_boot(monkeypatch):
|
||||
"""
|
||||
Boot calls this helper directly, so an ImportError here takes the proxy down before
|
||||
uvicorn starts. An install without the runner must lose the diagnostic, not the proxy.
|
||||
"""
|
||||
import sys
|
||||
|
||||
from litellm.proxy.db.check_migration import check_prisma_schema_diff_helper
|
||||
|
||||
monkeypatch.setitem(sys.modules, "litellm_proxy_extras.prisma_toolchain", None)
|
||||
|
||||
assert check_prisma_schema_diff_helper("postgresql://u:p@localhost:9/x") == (False, [])
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from fastapi.testclient import TestClient
|
|||
|
||||
|
||||
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper, should_update_prisma_schema
|
||||
from litellm.proxy.db.prisma_client import PrismaManager, PrismaWrapper, should_update_prisma_schema
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
@ -193,7 +193,10 @@ async def test_recreate_prisma_client_recovers_from_disconnected_client(
|
|||
mock_new_prisma.connect.assert_awaited_once()
|
||||
|
||||
|
||||
def test_db_push_applies_replica_identity_full_when_requested(monkeypatch):
|
||||
DB_PUSH_ARGV = ["db", "push", "--accept-data-loss", "--skip-generate"]
|
||||
|
||||
|
||||
def test_db_push_applies_replica_identity_full_when_requested(monkeypatch, fake_prisma_cli, unset_database_url):
|
||||
"""`prisma db push` bypasses litellm-proxy-extras, so it needs its own call
|
||||
into the opt-in REPLICA IDENTITY FULL step."""
|
||||
from litellm.proxy.db.prisma_client import PrismaManager
|
||||
|
|
@ -208,14 +211,13 @@ def test_db_push_applies_replica_identity_full_when_requested(monkeypatch):
|
|||
staticmethod(lambda: applied.append(True)),
|
||||
)
|
||||
|
||||
with patch("litellm.proxy.db.prisma_client.subprocess.run") as mock_run:
|
||||
assert PrismaManager.setup_database(use_migrate=False) is True
|
||||
assert PrismaManager.setup_database(use_migrate=False) is True
|
||||
|
||||
assert mock_run.call_args[0][0][:3] == ["prisma", "db", "push"]
|
||||
assert fake_prisma_cli.calls == [DB_PUSH_ARGV]
|
||||
assert applied == [True]
|
||||
|
||||
|
||||
def test_db_push_is_rejected_when_spend_logs_is_partitioned(monkeypatch):
|
||||
def test_db_push_is_rejected_when_spend_logs_is_partitioned(monkeypatch, fake_prisma_cli, unset_database_url):
|
||||
"""A doc-partitioned LiteLLM_SpendLogs makes `prisma db push` rewrite the
|
||||
primary key back to ("request_id"), which Postgres rejects; the guard must
|
||||
fail fast with guidance instead of running the push."""
|
||||
|
|
@ -228,29 +230,23 @@ def test_db_push_is_rejected_when_spend_logs_is_partitioned(monkeypatch):
|
|||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "spend_logs_is_partitioned", staticmethod(lambda: True)
|
||||
)
|
||||
with patch( # test-quality-ok: subprocess.run is the external prisma CLI boundary, asserted never reached
|
||||
"litellm.proxy.db.prisma_client.subprocess.run"
|
||||
) as mock_run:
|
||||
with pytest.raises(RuntimeError) as err:
|
||||
PrismaManager.setup_database(use_migrate=False)
|
||||
with pytest.raises(RuntimeError) as err:
|
||||
PrismaManager.setup_database(use_migrate=False)
|
||||
|
||||
assert str(err.value) == PARTITIONED_SPEND_LOGS_PUSH_ERROR
|
||||
mock_run.assert_not_called()
|
||||
assert fake_prisma_cli.calls == []
|
||||
|
||||
|
||||
def test_db_push_proceeds_when_spend_logs_is_not_partitioned(monkeypatch):
|
||||
def test_db_push_proceeds_when_spend_logs_is_not_partitioned(monkeypatch, fake_prisma_cli, unset_database_url):
|
||||
from litellm.proxy.db.prisma_client import PrismaManager
|
||||
from litellm_proxy_extras.utils import ProxyExtrasDBManager
|
||||
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "spend_logs_is_partitioned", staticmethod(lambda: False)
|
||||
)
|
||||
with patch( # test-quality-ok: subprocess.run is the external prisma CLI boundary, not SDK logic
|
||||
"litellm.proxy.db.prisma_client.subprocess.run"
|
||||
) as mock_run:
|
||||
assert PrismaManager.setup_database(use_migrate=False) is True
|
||||
assert PrismaManager.setup_database(use_migrate=False) is True
|
||||
|
||||
assert mock_run.call_args[0][0][:3] == ["prisma", "db", "push"]
|
||||
assert fake_prisma_cli.calls == [DB_PUSH_ARGV]
|
||||
|
||||
|
||||
def _entra_jwt(expires_in_seconds: int) -> str:
|
||||
|
|
@ -377,3 +373,30 @@ def test_minting_without_the_database_env_vars_names_them(azure_env, monkeypatch
|
|||
|
||||
with pytest.raises(RuntimeError, match="DATABASE_HOST"):
|
||||
wrapper.get_rds_iam_token()
|
||||
|
||||
|
||||
@pytest.mark.timeout(45)
|
||||
def test_db_push_timeout_takes_its_process_tree_with_it(fake_prisma_cli, unset_database_url, monkeypatch):
|
||||
"""
|
||||
A timed-out `db push` used to leave Node and the schema engine writing the schema,
|
||||
so the next attempt pushed into a database the abandoned one was still mutating.
|
||||
"""
|
||||
monkeypatch.delenv("LITELLM_SET_REPLICA_IDENTITY_FULL", raising=False)
|
||||
monkeypatch.setenv("FAKE_PRISMA_HANG_FIRST", "1")
|
||||
|
||||
assert PrismaManager.setup_database(use_migrate=False) is True
|
||||
assert fake_prisma_cli.calls == [DB_PUSH_ARGV, DB_PUSH_ARGV]
|
||||
assert fake_prisma_cli.grandchild_is_gone(within_seconds=5)
|
||||
|
||||
|
||||
def test_db_push_without_the_prisma_runner_fails_the_migration_instead_of_crashing_boot(
|
||||
fake_prisma_cli, unset_database_url, monkeypatch
|
||||
):
|
||||
"""
|
||||
An ImportError out of setup_database escapes the caller's RuntimeError handler and
|
||||
kills boot, bypassing the operator's enforce_prisma_migration_check choice.
|
||||
"""
|
||||
monkeypatch.setitem(sys.modules, "litellm_proxy_extras.prisma_toolchain", None)
|
||||
|
||||
assert PrismaManager.setup_database(use_migrate=False) is False
|
||||
assert fake_prisma_cli.calls == []
|
||||
|
|
|
|||
|
|
@ -2939,6 +2939,129 @@ async def test_ProxyConfig_add_deployment_applies_db_router_settings(monkeypatch
|
|||
fake_router.update_settings.assert_called_once_with(routing_strategy="latency-based-routing")
|
||||
|
||||
|
||||
def _stub_add_deployment_collaborators(
|
||||
monkeypatch: pytest.MonkeyPatch, pc: ProxyConfig, fake_prisma: MagicMock
|
||||
) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
fake_router = MagicMock()
|
||||
fake_router.get_model_list = MagicMock(return_value=[])
|
||||
|
||||
async def fake_get_config(*args: object, **kwargs: object) -> dict[str, object]:
|
||||
return {}
|
||||
|
||||
monkeypatch.setattr(litellm, "credential_list", [])
|
||||
monkeypatch.setattr(pc, "get_config", fake_get_config)
|
||||
monkeypatch.setattr(pc, "_init_non_llm_objects_in_db", AsyncMock())
|
||||
monkeypatch.setattr(proxy_server, "prefetch_config_params", AsyncMock())
|
||||
monkeypatch.setattr(proxy_server, "get_config_param", AsyncMock(return_value=None))
|
||||
monkeypatch.setattr(proxy_server, "llm_router", fake_router)
|
||||
monkeypatch.setattr(proxy_server, "master_key", "sk-master")
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma)
|
||||
monkeypatch.setattr(proxy_server, "proxy_config", pc)
|
||||
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
|
||||
|
||||
|
||||
def _encrypted_credential_row(credential_name: str, api_key: str) -> dict[str, object]:
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
|
||||
return {
|
||||
"credential_name": credential_name,
|
||||
"credential_values": {"api_key": encrypt_value_helper(api_key, new_encryption_key="sk-master")},
|
||||
"credential_info": {"custom_llm_provider": "openai"},
|
||||
}
|
||||
|
||||
|
||||
def _fake_prisma_with_encrypted_credential(credential_name: str, api_key: str) -> MagicMock:
|
||||
fake_prisma = MagicMock()
|
||||
fake_prisma.db.litellm_credentialstable.find_many = AsyncMock(
|
||||
return_value=[_encrypted_credential_row(credential_name, api_key)]
|
||||
)
|
||||
return fake_prisma
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_add_deployment_loads_db_credentials_before_reconciling_models(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.utils import load_credentials_from_list
|
||||
|
||||
pc = ProxyConfig()
|
||||
fake_prisma = MagicMock()
|
||||
fake_prisma.db.litellm_credentialstable.find_many = AsyncMock(return_value=[])
|
||||
_stub_add_deployment_collaborators(monkeypatch, pc, fake_prisma)
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
installed = MagicMock()
|
||||
|
||||
async def read_models_while_a_credential_lands(prisma_client: object) -> list[MagicMock]:
|
||||
fake_prisma.db.litellm_credentialstable.find_many.return_value = [
|
||||
_encrypted_credential_row("openai-cred", "sk-from-db")
|
||||
]
|
||||
return [MagicMock()]
|
||||
|
||||
async def install_models(new_models: object, proxy_logging_obj: object) -> None:
|
||||
installed(credential=CredentialAccessor.get_credential_values("openai-cred"))
|
||||
|
||||
monkeypatch.setattr(pc, "_get_models_from_db", read_models_while_a_credential_lands)
|
||||
monkeypatch.setattr(pc, "_update_llm_router", install_models)
|
||||
|
||||
await pc.add_deployment(prisma_client=fake_prisma, proxy_logging_obj=MagicMock())
|
||||
|
||||
installed.assert_called_once_with(credential={"api_key": "sk-from-db"})
|
||||
assert CredentialAccessor.get_credential_values("openai-cred") == {"api_key": "sk-from-db"}
|
||||
request_kwargs = {"litellm_credential_name": "openai-cred"}
|
||||
load_credentials_from_list(request_kwargs)
|
||||
assert request_kwargs == {"litellm_credential_name": "openai-cred", "api_key": "sk-from-db"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_add_deployment_loads_db_credentials_even_when_models_are_not_db_objects(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
pc = ProxyConfig()
|
||||
fake_prisma = _fake_prisma_with_encrypted_credential("openai-cred", "sk-from-db")
|
||||
_stub_add_deployment_collaborators(monkeypatch, pc, fake_prisma)
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": ["mcp"]})
|
||||
models_fetch = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr(pc, "_get_models_from_db", models_fetch)
|
||||
|
||||
await pc.add_deployment(prisma_client=fake_prisma, proxy_logging_obj=MagicMock())
|
||||
|
||||
models_fetch.assert_not_awaited()
|
||||
assert CredentialAccessor.get_credential_values("openai-cred") == {"api_key": "sk-from-db"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_get_credentials_reads_from_writer_not_replica(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
|
||||
|
||||
pc = ProxyConfig()
|
||||
writer_inner = MagicMock(name="writer_prisma")
|
||||
reader_inner = MagicMock(name="reader_prisma")
|
||||
writer_inner.litellm_credentialstable.find_many = AsyncMock(
|
||||
return_value=[_encrypted_credential_row("openai-cred", "sk-from-writer")]
|
||||
)
|
||||
reader_inner.litellm_credentialstable.find_many = AsyncMock(return_value=[])
|
||||
fake_prisma = MagicMock()
|
||||
fake_prisma.db = RoutingPrismaWrapper(
|
||||
writer=PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False),
|
||||
reader=PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False),
|
||||
)
|
||||
_stub_add_deployment_collaborators(monkeypatch, pc, fake_prisma)
|
||||
|
||||
await pc.get_credentials(prisma_client=fake_prisma)
|
||||
|
||||
assert CredentialAccessor.get_credential_values("openai-cred") == {"api_key": "sk-from-writer"}
|
||||
reader_inner.litellm_credentialstable.find_many.assert_not_awaited()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ProxyConfig._add_general_settings_from_db_config
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -822,16 +822,16 @@ def _mock_scheduled_proxy_config() -> MagicMock:
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_scheduled_jobs_credentials(monkeypatch):
|
||||
"""
|
||||
Test that get_credentials is only called when store_model_in_db is True
|
||||
"""
|
||||
async def test_initialize_scheduled_jobs_loads_credentials_only_through_add_deployment(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
|
||||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
# Mock dependencies
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||||
|
|
@ -841,25 +841,6 @@ async def test_initialize_scheduled_jobs_credentials(monkeypatch):
|
|||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||||
): # set store_model_in_db to False
|
||||
# Test when store_model_in_db is False
|
||||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||||
general_settings={},
|
||||
prisma_client=mock_prisma_client,
|
||||
proxy_budget_rescheduler_min_time=1,
|
||||
proxy_budget_rescheduler_max_time=2,
|
||||
proxy_batch_write_at=5,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
# Verify get_credentials was not called
|
||||
mock_proxy_config.get_credentials.assert_not_called()
|
||||
|
||||
# Now test with store_model_in_db = True
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||||
):
|
||||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||||
general_settings={},
|
||||
|
|
@ -870,12 +851,31 @@ async def test_initialize_scheduled_jobs_credentials(monkeypatch):
|
|||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
# Verify get_credentials was called both directly and scheduled
|
||||
assert mock_proxy_config.get_credentials.call_count == 1 # Direct call
|
||||
mock_proxy_config.get_credentials.assert_not_called()
|
||||
mock_proxy_config.add_deployment.assert_not_called()
|
||||
|
||||
# Verify a scheduled job was added for get_credentials
|
||||
mock_scheduler_calls = [call[0] for call in mock_proxy_config.get_credentials.mock_calls]
|
||||
assert len(mock_scheduler_calls) > 0
|
||||
scheduler = AsyncIOScheduler()
|
||||
try:
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||||
patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=scheduler),
|
||||
):
|
||||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||||
general_settings={},
|
||||
prisma_client=mock_prisma_client,
|
||||
proxy_budget_rescheduler_min_time=1,
|
||||
proxy_budget_rescheduler_max_time=2,
|
||||
proxy_batch_write_at=5,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
assert scheduler.get_job("get_credentials_job") is None
|
||||
assert scheduler.get_job("add_deployment_job") is not None
|
||||
mock_proxy_config.get_credentials.assert_not_called()
|
||||
assert mock_proxy_config.add_deployment.call_count == 1
|
||||
finally:
|
||||
scheduler.shutdown(wait=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -924,7 +924,7 @@ async def test_periodic_reload_job_scheduled_without_store_model_in_db(monkeypat
|
|||
@pytest.mark.asyncio
|
||||
async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval(monkeypatch):
|
||||
"""
|
||||
The DB config-reload jobs (add_deployment, get_credentials) that keep multi-pod
|
||||
The DB config-reload job (add_deployment) that keeps multi-pod
|
||||
deployments in sync must be scheduled at the configured
|
||||
proxy_config_reload_interval_seconds, not a hardcoded value.
|
||||
"""
|
||||
|
|
@ -967,7 +967,7 @@ async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval(
|
|||
if "id" in job_call.kwargs
|
||||
}
|
||||
assert scheduled_seconds["add_deployment_job"] == configured_interval
|
||||
assert scheduled_seconds["get_credentials_job"] == configured_interval
|
||||
assert "get_credentials_job" not in scheduled_seconds
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1011,7 +1011,7 @@ async def test_initialize_scheduled_jobs_rejects_non_positive_config_reload_inte
|
|||
if "id" in job_call.kwargs
|
||||
}
|
||||
assert scheduled_seconds["add_deployment_job"] == 30
|
||||
assert scheduled_seconds["get_credentials_job"] == 30
|
||||
assert "get_credentials_job" not in scheduled_seconds
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -7446,10 +7446,8 @@ async def test_store_model_in_db_db_override_when_config_false():
|
|||
# store_model_in_db should now be True (overridden by DB)
|
||||
assert ps.store_model_in_db is True
|
||||
|
||||
# add_deployment and get_credentials should have been called
|
||||
# since store_model_in_db is now True
|
||||
assert mock_proxy_config.add_deployment.call_count == 1
|
||||
assert mock_proxy_config.get_credentials.call_count == 1
|
||||
mock_proxy_config.get_credentials.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
187
tests/test_litellm/router_strategy/test_least_busy.py
Normal file
187
tests/test_litellm/router_strategy/test_least_busy.py
Normal file
|
|
@ -0,0 +1,187 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.router_strategy.least_busy import IN_FLIGHT_COUNT_TTL_SECONDS, LeastBusyLoggingHandler
|
||||
|
||||
GROUP: Final = "least-busy-group"
|
||||
DEPLOYMENT_A: Final[dict[str, object]] = {"model_info": {"id": "dep-a"}}
|
||||
DEPLOYMENT_B: Final[dict[str, object]] = {"model_info": {"id": "dep-b"}}
|
||||
HEALTHY: Final = [DEPLOYMENT_A, DEPLOYMENT_B]
|
||||
|
||||
|
||||
def _call_kwargs(deployment_id: str) -> dict[str, object]:
|
||||
return {"litellm_params": {"metadata": {"model_group": GROUP}, "model_info": {"id": deployment_id}}}
|
||||
|
||||
|
||||
class SharedRedisCounters:
|
||||
"""Mirrors what Redis gives the handler: increments clamped at zero, a TTL set once when
|
||||
the key is created, and ordered reads that raise rather than invent a value."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.counts: dict[str, int] = {}
|
||||
self.ttls: dict[str, int] = {}
|
||||
|
||||
def count(self, key: str) -> int | None:
|
||||
return self.counts.get(key)
|
||||
|
||||
def expire(self, key: str) -> None:
|
||||
self.counts.pop(key, None)
|
||||
self.ttls.pop(key, None)
|
||||
|
||||
def increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
incremented: Final = max(0, self.counts.get(key, 0) + value)
|
||||
self.counts[key] = incremented
|
||||
self.ttls.setdefault(key, ttl)
|
||||
return incremented
|
||||
|
||||
async def async_increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
return self.increment_with_floor(key, value, ttl)
|
||||
|
||||
def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
return tuple(self.counts.get(key) for key in key_list)
|
||||
|
||||
async def async_batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
return self.batch_get_counts(key_list)
|
||||
|
||||
|
||||
def _worker(shared: SharedRedisCounters | None) -> LeastBusyLoggingHandler:
|
||||
cache: Final = DualCache(in_memory_cache=InMemoryCache(), redis_cache=shared) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
return LeastBusyLoggingHandler(router_cache=cache)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_routes_around_a_request_another_worker_started() -> None:
|
||||
shared: Final = SharedRedisCounters()
|
||||
streaming_worker: Final = _worker(shared)
|
||||
picking_worker: Final = _worker(shared)
|
||||
|
||||
picking_worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
await picking_worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
streaming_worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert await picking_worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
|
||||
await streaming_worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert await picking_worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
|
||||
|
||||
|
||||
def test_sync_pick_reads_the_shared_counts() -> None:
|
||||
shared: Final = SharedRedisCounters()
|
||||
streaming_worker: Final = _worker(shared)
|
||||
picking_worker: Final = _worker(shared)
|
||||
|
||||
picking_worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
picking_worker.log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
streaming_worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert picking_worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
|
||||
streaming_worker.log_failure_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert picking_worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
|
||||
|
||||
|
||||
def test_the_handler_never_pushes_a_counters_ttl_forward() -> None:
|
||||
"""A worker that dies mid-request leaves a +1 nobody will ever decrement. Redis expires that
|
||||
stuck count an hour after the key was created, which only works while nothing writes the TTL
|
||||
again: a handler that refreshed it on every touch would keep the count alive for as long as
|
||||
the group takes traffic, and the deployment would read busier than it is forever."""
|
||||
shared: Final = SharedRedisCounters()
|
||||
worker: Final = _worker(shared)
|
||||
key: Final = f"{GROUP}_request_count:dep-a"
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert shared.ttls == {key: IN_FLIGHT_COUNT_TTL_SECONDS}
|
||||
|
||||
shared.ttls[key] = 5
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
worker.log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert shared.count(key) == 1
|
||||
assert shared.ttls == {key: 5}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_counts_stay_in_memory_without_redis() -> None:
|
||||
worker: Final = _worker(None)
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
|
||||
await worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
|
||||
assert worker.router_cache.get_cache(f"{GROUP}_request_count:dep-a") == 0
|
||||
|
||||
|
||||
class UnavailableRedis(SharedRedisCounters):
|
||||
def increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
raise ConnectionError("redis is down")
|
||||
|
||||
def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
raise ConnectionError("redis is down")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_redis_outage_falls_back_to_this_workers_own_counts() -> None:
|
||||
worker: Final = _worker(UnavailableRedis())
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
|
||||
await worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
|
||||
|
||||
|
||||
def test_a_shared_counter_that_expired_mid_request_cannot_go_negative() -> None:
|
||||
shared: Final = SharedRedisCounters()
|
||||
worker: Final = _worker(shared)
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
shared.expire(f"{GROUP}_request_count:dep-a")
|
||||
worker.log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert shared.count(f"{GROUP}_request_count:dep-a") == 0
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert shared.count(f"{GROUP}_request_count:dep-a") == 1
|
||||
assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_local_counter_that_expired_mid_request_cannot_go_negative() -> None:
|
||||
worker: Final = _worker(None)
|
||||
in_memory: Final = worker.router_cache.in_memory_cache
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
in_memory.delete_cache(f"{GROUP}_request_count:dep-a")
|
||||
await worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert worker.router_cache.get_cache(f"{GROUP}_request_count:dep-a") == 0
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
|
||||
|
||||
def test_calls_without_a_deployment_are_ignored() -> None:
|
||||
shared: Final = SharedRedisCounters()
|
||||
worker: Final = _worker(shared)
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs={"litellm_params": {"metadata": None}})
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs={})
|
||||
|
||||
assert shared.counts == {}
|
||||
|
|
@ -12,6 +12,7 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.router import RoutingGroup, RoutingStrategy
|
||||
|
||||
|
||||
|
|
@ -435,6 +436,81 @@ def test_update_settings_unregisters_group_selectors_when_groups_removed(monkeyp
|
|||
assert router._group_selectors == {}
|
||||
|
||||
|
||||
def test_two_least_busy_groups_count_a_request_once(monkeypatch):
|
||||
"""
|
||||
Least-busy counts a request up from the pre-call hooks on `litellm.input_callback` and
|
||||
back down from the success hooks on `litellm.callbacks`. The success list drops a second
|
||||
selector of the same class, so a pre-call list that kept both counted every request twice
|
||||
and released it once, and the deployment's in-flight count climbed until it looked pinned.
|
||||
"""
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "input_callback", [])
|
||||
|
||||
router = _build_router(
|
||||
routing_strategy="least-busy",
|
||||
routing_groups=[
|
||||
{
|
||||
"group_name": "fast",
|
||||
"models": ["filtered-model"],
|
||||
"routing_strategy": "least-busy",
|
||||
}
|
||||
],
|
||||
)
|
||||
kwargs = {
|
||||
"litellm_params": {
|
||||
"metadata": {"model_group": "filtered-model"},
|
||||
"model_info": {"id": "deploy-1"},
|
||||
}
|
||||
}
|
||||
|
||||
for callback in litellm.input_callback:
|
||||
if isinstance(callback, CustomLogger):
|
||||
callback.log_pre_api_call(model="filtered-model", messages=[], kwargs=kwargs)
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
callback.log_success_event(kwargs, None, None, None)
|
||||
|
||||
assert router.cache.get_cache("filtered-model_request_count:deploy-1") == 0
|
||||
|
||||
|
||||
def test_two_routers_in_one_process_each_count_their_own_requests(monkeypatch):
|
||||
"""
|
||||
Least-busy hangs its counting off litellm's global callback lists, and those lists keep one
|
||||
logger per class unless the instances differ in a plain attribute. Two routers in one process
|
||||
(a second Router, or a per-request `user_config` one) therefore have to register separately:
|
||||
a second router whose selector is dropped counts nothing, reads zero for every deployment,
|
||||
and sends every request to whichever one is listed first.
|
||||
"""
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "input_callback", [])
|
||||
|
||||
first = _build_router(routing_strategy="least-busy")
|
||||
second = _build_router(routing_strategy="least-busy")
|
||||
kwargs = {
|
||||
"litellm_params": {
|
||||
"metadata": {"model_group": "filtered-model"},
|
||||
"model_info": {"id": "deploy-1"},
|
||||
}
|
||||
}
|
||||
|
||||
for callback in litellm.input_callback:
|
||||
if isinstance(callback, CustomLogger):
|
||||
callback.log_pre_api_call(model="filtered-model", messages=[], kwargs=kwargs)
|
||||
|
||||
assert second.cache.get_cache("filtered-model_request_count:deploy-1") == 1
|
||||
assert (
|
||||
second.get_available_deployment(model="filtered-model", messages=[])["model_info"]["id"]
|
||||
== "deploy-2"
|
||||
)
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
callback.log_success_event(kwargs, None, None, None)
|
||||
|
||||
assert first.cache.get_cache("filtered-model_request_count:deploy-1") == 0
|
||||
assert second.cache.get_cache("filtered-model_request_count:deploy-1") == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Direct helper coverage
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -56,6 +56,18 @@ class BlockEverything:
|
|||
return context
|
||||
|
||||
|
||||
class MessageRecorder:
|
||||
"""Records what each plugin pass was handed, then blocks so the request stops there."""
|
||||
|
||||
def __init__(self):
|
||||
self.seen = []
|
||||
|
||||
async def run(self, context: RoutingContext) -> RoutingContext:
|
||||
self.seen.append(list(context.raw_messages))
|
||||
context.candidate_models = []
|
||||
return context
|
||||
|
||||
|
||||
def _smart_router_model_list():
|
||||
return [
|
||||
{
|
||||
|
|
@ -164,6 +176,71 @@ async def test_async_completion_with_unsupported_strategy_rejects_configured_plu
|
|||
await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_management_model_still_runs_the_plugin_pipeline():
|
||||
"""
|
||||
A prompt-management model routes through its own factory, which picked the deployment
|
||||
on the synchronous path. Plugins never run there, so the guard turned every such request
|
||||
into an error message about the caller's own API choice, on an async call the caller made
|
||||
correctly. It also read the in-flight counts with a blocking call inside the event loop.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "cached-claude",
|
||||
"litellm_params": {
|
||||
"model": "anthropic_cache_control_hook/claude-sonnet-5",
|
||||
"prompt_id": "cache-points",
|
||||
},
|
||||
}
|
||||
],
|
||||
routing_strategy="least-busy",
|
||||
plugins=[BlockEverything()],
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="No deployments left after routing-plugin filtering"):
|
||||
await router.acompletion(
|
||||
model="cached-claude",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
litellm_call_id="lit-7039",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_management_plugins_see_the_callers_own_messages():
|
||||
"""
|
||||
The prompt-management factory picks its deployment with a placeholder message, which was
|
||||
harmless while that pick ran on the synchronous path (plugins never ran there at all). Now
|
||||
that the pick runs the plugin pipeline, a plugin that classifies request content would score
|
||||
the placeholder instead of the conversation, and the narrowing it produces decides which
|
||||
deployments the real call is allowed to use.
|
||||
"""
|
||||
recorder = MessageRecorder()
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "cached-claude",
|
||||
"litellm_params": {
|
||||
"model": "anthropic_cache_control_hook/claude-sonnet-5",
|
||||
"prompt_id": "cache-points",
|
||||
},
|
||||
}
|
||||
],
|
||||
routing_strategy="least-busy",
|
||||
plugins=[recorder],
|
||||
)
|
||||
messages = [{"role": "user", "content": "wire me $40,000 to account 12345"}]
|
||||
|
||||
with pytest.raises(ValueError, match="No deployments left after routing-plugin filtering"):
|
||||
await router.acompletion(
|
||||
model="cached-claude",
|
||||
messages=messages,
|
||||
litellm_call_id="lit-7039",
|
||||
)
|
||||
|
||||
assert recorder.seen == [messages]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_without_plugins_is_unaffected():
|
||||
"""Regression guard: a Router with no `plugins` configured behaves exactly as before."""
|
||||
|
|
|
|||
|
|
@ -268,12 +268,12 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time() - 120.0,
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
cc.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
|
||||
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
assert active == [], "Expired cooldown entry must not appear in active cooldowns"
|
||||
assert cc.cache.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
|
||||
assert cc.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
|
||||
|
||||
def test_active_entry_is_returned(self):
|
||||
"""
|
||||
|
|
@ -289,7 +289,7 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time(),
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60)
|
||||
cc.in_memory_cache.set_cache(key, active_value, ttl=60)
|
||||
|
||||
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
|
|
@ -312,14 +312,14 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time() - (60.0 - remaining),
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, value, ttl=600)
|
||||
cc.in_memory_cache.set_cache(key, value, ttl=600)
|
||||
|
||||
before_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
before_expiry = cc.in_memory_cache.ttl_dict.get(key)
|
||||
assert before_expiry is not None
|
||||
|
||||
cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
after_expiry = cc.in_memory_cache.ttl_dict.get(key)
|
||||
assert after_expiry is not None
|
||||
corrected_remaining = after_expiry - time.time()
|
||||
assert corrected_remaining <= 60.0, "Corrected TTL must not exceed 60s"
|
||||
|
|
@ -340,12 +340,12 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time() - 120.0,
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
cc.in_memory_cache.set_cache(key, expired_value, ttl=600)
|
||||
|
||||
active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
assert active == [], "Expired entry must not appear in async active cooldowns"
|
||||
assert cc.cache.in_memory_cache.get_cache(key) is None
|
||||
assert cc.in_memory_cache.get_cache(key) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_active_entry_is_returned(self):
|
||||
|
|
@ -363,7 +363,7 @@ class TestCooldownCacheTTLCorrection:
|
|||
"timestamp": time.time(),
|
||||
"cooldown_time": 60.0,
|
||||
}
|
||||
cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60)
|
||||
cc.in_memory_cache.set_cache(key, active_value, ttl=60)
|
||||
|
||||
active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
|
||||
|
||||
|
|
@ -389,18 +389,18 @@ class TestCorrectedActiveCooldown:
|
|||
cc = self._make_cooldown_cache()
|
||||
key = "deployment:expired-dep:cooldown"
|
||||
entry = self._entry(timestamp=time.time() - 120.0, cooldown_time=60.0)
|
||||
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=600)
|
||||
cc.in_memory_cache.set_cache(key, dict(entry), ttl=600)
|
||||
|
||||
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
|
||||
|
||||
assert result is None
|
||||
assert cc.cache.in_memory_cache.get_cache(key) is None
|
||||
assert cc.in_memory_cache.get_cache(key) is None
|
||||
|
||||
def test_active_entry_within_window_returns_value(self):
|
||||
cc = self._make_cooldown_cache()
|
||||
key = "deployment:active-dep:cooldown"
|
||||
entry = self._entry(timestamp=time.time(), cooldown_time=60.0)
|
||||
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=60)
|
||||
cc.in_memory_cache.set_cache(key, dict(entry), ttl=60)
|
||||
|
||||
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
|
||||
|
||||
|
|
@ -412,12 +412,12 @@ class TestCorrectedActiveCooldown:
|
|||
key = "deployment:backfilled-dep:cooldown"
|
||||
remaining = 30.0
|
||||
entry = self._entry(timestamp=time.time() - (60.0 - remaining), cooldown_time=60.0)
|
||||
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=600)
|
||||
cc.in_memory_cache.set_cache(key, dict(entry), ttl=600)
|
||||
|
||||
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
|
||||
|
||||
assert result is not None
|
||||
corrected_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
corrected_expiry = cc.in_memory_cache.ttl_dict.get(key)
|
||||
assert corrected_expiry is not None
|
||||
assert corrected_expiry - time.time() <= 60.0
|
||||
|
||||
|
|
@ -425,10 +425,160 @@ class TestCorrectedActiveCooldown:
|
|||
cc = self._make_cooldown_cache()
|
||||
key = "deployment:normal-dep:cooldown"
|
||||
entry = self._entry(timestamp=time.time(), cooldown_time=60.0)
|
||||
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=60)
|
||||
original_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
cc.in_memory_cache.set_cache(key, dict(entry), ttl=60)
|
||||
original_expiry = cc.in_memory_cache.ttl_dict.get(key)
|
||||
|
||||
cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
|
||||
|
||||
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
|
||||
after_expiry = cc.in_memory_cache.ttl_dict.get(key)
|
||||
assert after_expiry == original_expiry
|
||||
|
||||
|
||||
class SharedRedisDouble:
|
||||
"""
|
||||
In-process stand-in for RedisCache, shared by several DualCache instances so that
|
||||
tests can model two proxy replicas talking to one Redis.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.store: dict = {} # mutable-ok: stands in for Redis' own mutable keyspace
|
||||
|
||||
def set_cache(self, key, value, **kwargs):
|
||||
self.store[key] = value
|
||||
|
||||
async def async_set_cache(self, key, value, **kwargs):
|
||||
self.store[key] = value
|
||||
|
||||
def batch_get_cache(self, key_list, parent_otel_span=None, **kwargs):
|
||||
return {key: self.store.get(key) for key in key_list}
|
||||
|
||||
async def async_batch_get_cache(self, key_list, parent_otel_span=None, **kwargs):
|
||||
return {key: self.store.get(key) for key in key_list}
|
||||
|
||||
|
||||
class TestCooldownPropagationBetweenReplicas:
|
||||
"""
|
||||
A cooldown written by one replica has to reach its siblings quickly. The router's own
|
||||
DualCache re-reads a key that is missing from memory only every 10s, so cooldown reads
|
||||
get their own cache with a much shorter Redis read interval.
|
||||
"""
|
||||
|
||||
def _make_replica(self, redis: SharedRedisDouble, read_interval: float | None = None) -> CooldownCache:
|
||||
router_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis)
|
||||
if read_interval is None:
|
||||
return CooldownCache(cache=router_cache, default_cooldown_time=60.0)
|
||||
return CooldownCache(
|
||||
cache=router_cache,
|
||||
default_cooldown_time=60.0,
|
||||
redis_read_interval_seconds=read_interval,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sibling_replica_sees_cooldown_within_configured_read_interval(self):
|
||||
redis = SharedRedisDouble()
|
||||
replica_a = self._make_replica(redis, read_interval=0.25)
|
||||
replica_b = self._make_replica(redis, read_interval=0.25)
|
||||
model_id = "shared-deployment"
|
||||
|
||||
assert await replica_b.async_get_active_cooldowns([model_id], parent_otel_span=None) == []
|
||||
|
||||
replica_a.add_deployment_to_cooldown(
|
||||
model_id=model_id,
|
||||
original_exception=Exception("Internal server error"),
|
||||
exception_status=500,
|
||||
cooldown_time=60.0,
|
||||
)
|
||||
|
||||
time.sleep(0.3)
|
||||
|
||||
active = await replica_b.async_get_active_cooldowns([model_id], parent_otel_span=None)
|
||||
assert [model_id] == [entry[0] for entry in active], (
|
||||
"sibling replica must pick up a cooldown written by another replica within the read interval"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sibling_replica_sees_cooldown_within_default_read_interval(self):
|
||||
redis = SharedRedisDouble()
|
||||
replica_a = self._make_replica(redis)
|
||||
replica_b = self._make_replica(redis)
|
||||
model_id = "default-interval-deployment"
|
||||
|
||||
assert await replica_b.async_get_active_cooldowns([model_id], parent_otel_span=None) == []
|
||||
|
||||
replica_a.add_deployment_to_cooldown(
|
||||
model_id=model_id,
|
||||
original_exception=Exception("Internal server error"),
|
||||
exception_status=500,
|
||||
cooldown_time=60.0,
|
||||
)
|
||||
|
||||
time.sleep(1.2)
|
||||
|
||||
active = await replica_b.async_get_active_cooldowns([model_id], parent_otel_span=None)
|
||||
assert [model_id] == [entry[0] for entry in active], (
|
||||
"the shipped default read interval must let a sibling replica see a cooldown about a second later"
|
||||
)
|
||||
|
||||
def test_sync_read_path_sees_sibling_cooldown_within_read_interval(self):
|
||||
redis = SharedRedisDouble()
|
||||
replica_a = self._make_replica(redis, read_interval=0.25)
|
||||
replica_b = self._make_replica(redis, read_interval=0.25)
|
||||
model_id = "sync-shared-deployment"
|
||||
|
||||
assert replica_b.get_active_cooldowns([model_id], parent_otel_span=None) == []
|
||||
|
||||
replica_a.add_deployment_to_cooldown(
|
||||
model_id=model_id,
|
||||
original_exception=Exception("Internal server error"),
|
||||
exception_status=500,
|
||||
cooldown_time=60.0,
|
||||
)
|
||||
|
||||
time.sleep(0.3)
|
||||
|
||||
active = replica_b.get_active_cooldowns([model_id], parent_otel_span=None)
|
||||
assert [model_id] == [entry[0] for entry in active]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_attached_after_construction_is_still_used(self):
|
||||
redis = SharedRedisDouble()
|
||||
router_cache = DualCache(in_memory_cache=InMemoryCache())
|
||||
writer = CooldownCache(cache=router_cache, default_cooldown_time=60.0, redis_read_interval_seconds=0.25)
|
||||
router_cache.attach_redis_cache(redis)
|
||||
reader = self._make_replica(redis, read_interval=0.25)
|
||||
model_id = "late-redis-deployment"
|
||||
|
||||
writer.add_deployment_to_cooldown(
|
||||
model_id=model_id,
|
||||
original_exception=Exception("Internal server error"),
|
||||
exception_status=500,
|
||||
cooldown_time=60.0,
|
||||
)
|
||||
|
||||
active = await reader.async_get_active_cooldowns([model_id], parent_otel_span=None)
|
||||
assert [model_id] == [entry[0] for entry in active], (
|
||||
"a router that wires Redis after building its cooldown cache must still publish cooldowns to it"
|
||||
)
|
||||
|
||||
|
||||
class TestCooldownSurvivesUnrelatedCacheTraffic:
|
||||
@pytest.mark.asyncio
|
||||
async def test_unrelated_router_cache_writes_do_not_evict_active_cooldown(self):
|
||||
router_cache = DualCache(in_memory_cache=InMemoryCache())
|
||||
cc = CooldownCache(cache=router_cache, default_cooldown_time=60.0)
|
||||
model_id = "busy-router-deployment"
|
||||
|
||||
cc.add_deployment_to_cooldown(
|
||||
model_id=model_id,
|
||||
original_exception=Exception("Internal server error"),
|
||||
exception_status=500,
|
||||
cooldown_time=30.0,
|
||||
)
|
||||
|
||||
for i in range(400):
|
||||
router_cache.set_cache(key=f"unrelated-router-key-{i}", value={"n": i})
|
||||
|
||||
active = await cc.async_get_active_cooldowns([model_id], parent_otel_span=None)
|
||||
assert [model_id] == [entry[0] for entry in active], (
|
||||
"unrelated router cache traffic must not evict a cooldown that is still running"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6092,3 +6092,47 @@ class TestFinalOptionalParamsLineRedaction:
|
|||
|
||||
assert "'max_tokens': 17" in printed
|
||||
assert "'temperature': 0.25" in printed
|
||||
|
||||
|
||||
def _credential_warnings(caplog: pytest.LogCaptureFixture) -> list[str]:
|
||||
return [record.getMessage() for record in caplog.records if "litellm_credential_name=" in record.getMessage()]
|
||||
|
||||
|
||||
def test_load_credentials_from_list_warns_when_the_named_credential_is_not_loaded(
|
||||
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
from litellm.utils import load_credentials_from_list
|
||||
|
||||
monkeypatch.setattr(litellm, "credential_list", [])
|
||||
request_kwargs = {"litellm_credential_name": "openai-cred", "model": "openai/gpt-5.4-mini"}
|
||||
with caplog.at_level(logging.WARNING, logger=verbose_logger.name):
|
||||
load_credentials_from_list(request_kwargs)
|
||||
|
||||
assert request_kwargs == {"litellm_credential_name": "openai-cred", "model": "openai/gpt-5.4-mini"}
|
||||
assert _credential_warnings(caplog) == [
|
||||
"litellm_credential_name=openai-cred matched none of the 0 loaded credentials; the request runs without it"
|
||||
]
|
||||
|
||||
|
||||
def test_load_credentials_from_list_fills_kwargs_from_the_loaded_credential_without_warning(
|
||||
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
from litellm.types.utils import CredentialItem
|
||||
from litellm.utils import load_credentials_from_list
|
||||
|
||||
loaded = CredentialItem(
|
||||
credential_name="openai-cred",
|
||||
credential_values={"api_key": "sk-from-db", "api_base": "https://credential.example"},
|
||||
credential_info={},
|
||||
)
|
||||
monkeypatch.setattr(litellm, "credential_list", [loaded])
|
||||
request_kwargs = {"litellm_credential_name": "openai-cred", "api_base": "https://request.example"}
|
||||
with caplog.at_level(logging.WARNING, logger=verbose_logger.name):
|
||||
load_credentials_from_list(request_kwargs)
|
||||
|
||||
assert request_kwargs == {
|
||||
"litellm_credential_name": "openai-cred",
|
||||
"api_base": "https://request.example",
|
||||
"api_key": "sk-from-db",
|
||||
}
|
||||
assert _credential_warnings(caplog) == []
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22180
|
||||
"limit": 22174
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26729
|
||||
"limit": 26715
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 261
|
||||
|
|
@ -27,10 +27,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16426
|
||||
"limit": 16398
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5506
|
||||
"limit": 5504
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4486
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue