This commit is contained in:
Charan Rathore 2026-09-30 10:30:02 -04:00 • committed by GitHub
commit a54e427a97
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 101 additions and 24 deletions

View file

@ -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:
"""

View file

@ -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"}}