mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 343360d7af into b781d157d7
This commit is contained in:
commit
a54e427a97
2 changed files with 101 additions and 24 deletions
|
|
@ -1,8 +1,8 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
|
@ -38,6 +38,7 @@ class SharedHealthCheckManager:
|
|||
self.health_check_ttl = health_check_ttl
|
||||
self.lock_ttl = lock_ttl
|
||||
self.pod_id = f"pod_{int(time.time() * 1000)}"
|
||||
self._release_lock_script: Callable[..., Awaitable[int]] | None = None
|
||||
|
||||
@staticmethod
|
||||
def get_health_check_lock_key() -> str:
|
||||
|
|
@ -89,20 +90,36 @@ class SharedHealthCheckManager:
|
|||
verbose_proxy_logger.error("Error acquiring health check lock: %s", str(e))
|
||||
return False
|
||||
|
||||
_COMPARE_AND_DELETE_LOCK_SCRIPT = """
|
||||
if redis.call("get", KEYS[1]) == ARGV[1] then
|
||||
return redis.call("del", KEYS[1])
|
||||
else
|
||||
return 0
|
||||
end
|
||||
"""
|
||||
|
||||
async def release_health_check_lock(self) -> None:
|
||||
"""Release the global health check lock."""
|
||||
"""Release only this pod's lock using an atomic compare-and-delete."""
|
||||
if self.redis_cache is None:
|
||||
return
|
||||
|
||||
lock_key: Final = self.get_health_check_lock_key()
|
||||
script_register: Final = getattr(self.redis_cache, "async_register_script", None)
|
||||
if not callable(script_register):
|
||||
verbose_proxy_logger.warning("Cannot atomically release health check lock; leaving it to expire")
|
||||
return
|
||||
|
||||
try:
|
||||
lock_key: Final = self.get_health_check_lock_key()
|
||||
# Only release if we own the lock
|
||||
current_owner: Final = await self.redis_cache.async_get_cache(lock_key)
|
||||
if current_owner == self.pod_id:
|
||||
await self.redis_cache.async_delete_cache(lock_key)
|
||||
if self._release_lock_script is None:
|
||||
self._release_lock_script = cast(
|
||||
Callable[..., Awaitable[int]], script_register(self._COMPARE_AND_DELETE_LOCK_SCRIPT)
|
||||
)
|
||||
result: Final = await self._release_lock_script(keys=[lock_key], args=[json.dumps(self.pod_id)])
|
||||
if int(result or 0) == 1:
|
||||
verbose_proxy_logger.info("Pod %s released health check lock", self.pod_id)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error releasing health check lock: %s", str(e))
|
||||
except Exception as exc:
|
||||
self._release_lock_script = None
|
||||
verbose_proxy_logger.warning("Atomic health check lock release failed; leaving it to expire: %s", exc)
|
||||
|
||||
async def get_cached_health_check_results(self) -> dict[str, Any] | None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -121,29 +121,88 @@ class TestSharedHealthCheckManager:
|
|||
assert result is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_health_check_lock_success(
|
||||
self, shared_health_manager, mock_redis_cache
|
||||
):
|
||||
"""Test successful lock release"""
|
||||
mock_redis_cache.async_get_cache.return_value = shared_health_manager.pod_id
|
||||
async def test_release_health_check_lock_success(self, shared_health_manager, mock_redis_cache):
|
||||
"""The Lua script compares with the JSON-encoded owner written by RedisCache."""
|
||||
script = AsyncMock(return_value=1)
|
||||
mock_redis_cache.async_register_script = MagicMock(return_value=script)
|
||||
|
||||
await shared_health_manager.release_health_check_lock()
|
||||
|
||||
mock_redis_cache.async_get_cache.assert_called_once_with("health_check_lock")
|
||||
mock_redis_cache.async_delete_cache.assert_called_once_with("health_check_lock")
|
||||
mock_redis_cache.async_register_script.assert_called_once_with(
|
||||
SharedHealthCheckManager._COMPARE_AND_DELETE_LOCK_SCRIPT
|
||||
)
|
||||
script.assert_awaited_once_with(keys=["health_check_lock"], args=[json.dumps(shared_health_manager.pod_id)])
|
||||
mock_redis_cache.async_get_cache.assert_not_called()
|
||||
mock_redis_cache.async_delete_cache.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_health_check_lock_wrong_owner(
|
||||
self, shared_health_manager, mock_redis_cache
|
||||
):
|
||||
"""Test lock release when not the owner"""
|
||||
mock_redis_cache.async_get_cache.return_value = "other_pod_id"
|
||||
async def test_release_health_check_lock_wrong_owner(self, shared_health_manager, mock_redis_cache):
|
||||
script = AsyncMock(return_value=0)
|
||||
mock_redis_cache.async_register_script = MagicMock(return_value=script)
|
||||
|
||||
await shared_health_manager.release_health_check_lock()
|
||||
|
||||
mock_redis_cache.async_get_cache.assert_called_once_with("health_check_lock")
|
||||
script.assert_awaited_once()
|
||||
mock_redis_cache.async_delete_cache.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_health_check_lock_script_failure_does_not_delete(
|
||||
self, shared_health_manager, mock_redis_cache
|
||||
):
|
||||
script = AsyncMock(side_effect=RuntimeError("NOSCRIPT"))
|
||||
mock_redis_cache.async_register_script = MagicMock(return_value=script)
|
||||
|
||||
await shared_health_manager.release_health_check_lock()
|
||||
|
||||
assert shared_health_manager._release_lock_script is None
|
||||
mock_redis_cache.async_get_cache.assert_not_called()
|
||||
mock_redis_cache.async_delete_cache.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_health_check_lock_without_script_preserves_new_owner(self):
|
||||
class ReplacedLockCache:
|
||||
def __init__(self):
|
||||
self.owner = None
|
||||
self.reads = 0
|
||||
self.deletes = 0
|
||||
|
||||
async def async_set_cache(self, key, value, nx=False, ttl=None):
|
||||
self.owner = value
|
||||
return True
|
||||
|
||||
async def async_get_cache(self, key):
|
||||
self.reads += 1
|
||||
old_owner = self.owner
|
||||
self.owner = "new-pod"
|
||||
return old_owner
|
||||
|
||||
async def async_delete_cache(self, key):
|
||||
self.deletes += 1
|
||||
self.owner = None
|
||||
return 1
|
||||
|
||||
cache = ReplacedLockCache()
|
||||
manager = SharedHealthCheckManager(redis_cache=cache)
|
||||
assert await manager.acquire_health_check_lock() is True
|
||||
await manager.release_health_check_lock()
|
||||
|
||||
assert cache.owner == manager.pod_id
|
||||
assert cache.reads == 0
|
||||
assert cache.deletes == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_health_check_lock_reuses_script(self, shared_health_manager, mock_redis_cache):
|
||||
script = AsyncMock(side_effect=[1, 0])
|
||||
mock_redis_cache.async_register_script = MagicMock(return_value=script)
|
||||
|
||||
await shared_health_manager.release_health_check_lock()
|
||||
await shared_health_manager.release_health_check_lock()
|
||||
|
||||
mock_redis_cache.async_register_script.assert_called_once()
|
||||
assert script.await_count == 2
|
||||
mock_redis_cache.async_delete_cache.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_health_check_lock_no_redis(self):
|
||||
"""Test lock release without Redis"""
|
||||
|
|
@ -288,8 +347,9 @@ class TestSharedHealthCheckManager:
|
|||
"""Test performing shared health check when acquiring lock"""
|
||||
# No cached results
|
||||
mock_redis_cache.async_get_cache.return_value = None
|
||||
# Lock acquisition succeeds
|
||||
# Lock acquisition succeeds; registration returns an awaitable script callable.
|
||||
mock_redis_cache.async_set_cache.return_value = True
|
||||
mock_redis_cache.async_register_script = MagicMock(return_value=AsyncMock(return_value=1))
|
||||
|
||||
model_list = [
|
||||
{"model_name": "test-model", "litellm_params": {"model": "test-model"}}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue