mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): stop queued registry read-throughs spending the resync budget (#44277)
* fix(proxy): stop queued registry read-throughs spending the resync budget RegistryReadThrough.attempt serializes misses behind one lock, but every request that was queued behind the first one still spent a unit of the 20-per-5s resync budget and re-ran the DB resync, even though the first request had already loaded the object. A burst of more than 20 requests for a model created on another worker therefore exhausted the budget and the rest got 400 "Invalid model name". attempt now checks whether the key is already loaded once it holds the lock and returns early without touching the budget. Models check the router's model names and deployment ids, guardrails and agents reuse their existing registry lookups. * test(proxy): gate the queued read-through test on events and cover each registry's loaded check The queued-requests test now holds the first resync on an asyncio.Event instead of a timed sleep and records calls in a recorder with tuple and frozenset state. New tests show the model, guardrail and agent read-throughs each answer an object that is already loaded without reading the database, so rewiring any registry's loaded check now fails a test * test(proxy): keep the queued read-through recorder inside its test and type the agent registry fixture
This commit is contained in:
parent
dd86ca5175
commit
b877a38e5f
3 changed files with 184 additions and 14 deletions
|
|
@ -36,6 +36,7 @@ READ_THROUGH_MAX_RESYNCS_PER_WINDOW: Final = 20
|
|||
|
||||
class RegistryReadThrough:
|
||||
__slots__ = (
|
||||
"_is_loaded",
|
||||
"_lock",
|
||||
"_max_resyncs_per_window",
|
||||
"_miss_ttl_seconds",
|
||||
|
|
@ -49,11 +50,13 @@ class RegistryReadThrough:
|
|||
def __init__(
|
||||
self,
|
||||
resync: Callable[[str], Awaitable[bool]],
|
||||
is_loaded: Callable[[str], bool],
|
||||
miss_ttl_seconds: float = READ_THROUGH_MISS_TTL_SECONDS,
|
||||
max_resyncs_per_window: int = READ_THROUGH_MAX_RESYNCS_PER_WINDOW,
|
||||
resync_window_seconds: float = READ_THROUGH_RESYNC_WINDOW_SECONDS,
|
||||
) -> None:
|
||||
self._resync = resync
|
||||
self._is_loaded = is_loaded
|
||||
self._miss_ttl_seconds = miss_ttl_seconds
|
||||
self._max_resyncs_per_window = max_resyncs_per_window
|
||||
self._resync_window_seconds = resync_window_seconds
|
||||
|
|
@ -78,6 +81,8 @@ class RegistryReadThrough:
|
|||
async with self._lock:
|
||||
if self._recent_misses.get_cache(key) is not None:
|
||||
return False
|
||||
if self._is_loaded(key):
|
||||
return True
|
||||
if not self._consume_resync_budget():
|
||||
verbose_proxy_logger.warning(
|
||||
"registry read-through for %r skipped: resync budget of %s per %ss exhausted",
|
||||
|
|
@ -190,9 +195,26 @@ async def _resync_agents(agent_id_or_name: str) -> bool:
|
|||
return True
|
||||
|
||||
|
||||
model_registry_read_through: Final = RegistryReadThrough(resync=_resync_model_deployments)
|
||||
guardrail_registry_read_through: Final = RegistryReadThrough(resync=_resync_guardrails)
|
||||
agent_registry_read_through: Final = RegistryReadThrough(resync=_resync_agents)
|
||||
def _model_is_loaded(model_name_or_id: str) -> bool:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router: Final = proxy_server.llm_router
|
||||
if router is None:
|
||||
return False
|
||||
return model_name_or_id in router.model_names or router.has_model_id(model_name_or_id)
|
||||
|
||||
|
||||
def _guardrail_is_loaded(guardrail_name: str) -> bool:
|
||||
return _initialized_guardrail(guardrail_name) is not None
|
||||
|
||||
|
||||
def _agent_is_loaded(agent_id_or_name: str) -> bool:
|
||||
return _agent_from_registry(agent_id_or_name) is not None
|
||||
|
||||
|
||||
model_registry_read_through: Final = RegistryReadThrough(resync=_resync_model_deployments, is_loaded=_model_is_loaded)
|
||||
guardrail_registry_read_through: Final = RegistryReadThrough(resync=_resync_guardrails, is_loaded=_guardrail_is_loaded)
|
||||
agent_registry_read_through: Final = RegistryReadThrough(resync=_resync_agents, is_loaded=_agent_is_loaded)
|
||||
|
||||
|
||||
def _agent_from_registry(agent_id_or_name: str) -> "AgentResponse | None":
|
||||
|
|
|
|||
|
|
@ -1,10 +1,17 @@
|
|||
import asyncio
|
||||
from typing import Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.common_utils.registry_read_through import RegistryReadThrough
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
|
||||
|
||||
def nothing_loaded(_key: str) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
class ResyncSpy:
|
||||
def __init__(self, found: bool = True, error: Exception | None = None) -> None:
|
||||
|
|
@ -22,7 +29,7 @@ class ResyncSpy:
|
|||
@pytest.mark.asyncio
|
||||
async def test_attempt_returns_true_when_resync_finds_object():
|
||||
spy: Final = ResyncSpy(found=True)
|
||||
read_through: Final = RegistryReadThrough(resync=spy)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, is_loaded=nothing_loaded)
|
||||
|
||||
assert await read_through.attempt("new-model") is True
|
||||
assert spy.calls == ["new-model"]
|
||||
|
|
@ -31,7 +38,7 @@ async def test_attempt_returns_true_when_resync_finds_object():
|
|||
@pytest.mark.asyncio
|
||||
async def test_attempt_found_key_is_not_negative_cached():
|
||||
spy: Final = ResyncSpy(found=True)
|
||||
read_through: Final = RegistryReadThrough(resync=spy)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, is_loaded=nothing_loaded)
|
||||
|
||||
assert await read_through.attempt("new-model") is True
|
||||
assert await read_through.attempt("new-model") is True
|
||||
|
|
@ -41,7 +48,7 @@ async def test_attempt_found_key_is_not_negative_cached():
|
|||
@pytest.mark.asyncio
|
||||
async def test_missing_key_is_negative_cached_within_ttl():
|
||||
spy: Final = ResyncSpy(found=False)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, miss_ttl_seconds=60.0)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, is_loaded=nothing_loaded, miss_ttl_seconds=60.0)
|
||||
|
||||
assert await read_through.attempt("ghost-model") is False
|
||||
assert await read_through.attempt("ghost-model") is False
|
||||
|
|
@ -51,7 +58,7 @@ async def test_missing_key_is_negative_cached_within_ttl():
|
|||
@pytest.mark.asyncio
|
||||
async def test_negative_cache_expires_and_resync_runs_again():
|
||||
spy: Final = ResyncSpy(found=False)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, miss_ttl_seconds=0.05)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, is_loaded=nothing_loaded, miss_ttl_seconds=0.05)
|
||||
|
||||
assert await read_through.attempt("ghost-model") is False
|
||||
await asyncio.sleep(0.1)
|
||||
|
|
@ -62,7 +69,7 @@ async def test_negative_cache_expires_and_resync_runs_again():
|
|||
@pytest.mark.asyncio
|
||||
async def test_resync_exception_returns_false_without_negative_caching():
|
||||
spy: Final = ResyncSpy(error=RuntimeError("db down"))
|
||||
read_through: Final = RegistryReadThrough(resync=spy)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, is_loaded=nothing_loaded)
|
||||
|
||||
assert await read_through.attempt("new-model") is False
|
||||
assert await read_through.attempt("new-model") is False
|
||||
|
|
@ -77,7 +84,7 @@ async def test_concurrent_attempts_for_missing_key_resync_once():
|
|||
return await super().__call__(key)
|
||||
|
||||
spy: Final = SlowResyncSpy(found=False)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, miss_ttl_seconds=60.0)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, is_loaded=nothing_loaded, miss_ttl_seconds=60.0)
|
||||
|
||||
results: Final = await asyncio.gather(*(read_through.attempt("ghost-model") for _ in range(5)))
|
||||
assert results == [False] * 5
|
||||
|
|
@ -87,7 +94,7 @@ async def test_concurrent_attempts_for_missing_key_resync_once():
|
|||
@pytest.mark.asyncio
|
||||
async def test_distinct_keys_do_not_share_negative_cache():
|
||||
spy: Final = ResyncSpy(found=False)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, miss_ttl_seconds=60.0)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, is_loaded=nothing_loaded, miss_ttl_seconds=60.0)
|
||||
|
||||
assert await read_through.attempt("ghost-a") is False
|
||||
assert await read_through.attempt("ghost-b") is False
|
||||
|
|
@ -98,7 +105,11 @@ async def test_distinct_keys_do_not_share_negative_cache():
|
|||
async def test_resync_budget_exhausted_blocks_resync_without_negative_caching():
|
||||
spy: Final = ResyncSpy(found=False)
|
||||
read_through: Final = RegistryReadThrough(
|
||||
resync=spy, miss_ttl_seconds=60.0, max_resyncs_per_window=2, resync_window_seconds=60.0
|
||||
resync=spy,
|
||||
is_loaded=nothing_loaded,
|
||||
miss_ttl_seconds=60.0,
|
||||
max_resyncs_per_window=2,
|
||||
resync_window_seconds=60.0,
|
||||
)
|
||||
|
||||
assert await read_through.attempt("ghost-a") is False
|
||||
|
|
@ -108,10 +119,48 @@ async def test_resync_budget_exhausted_blocks_resync_without_negative_caching():
|
|||
assert read_through._recent_misses.get_cache("ghost-c") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_requests_queued_behind_a_successful_resync_spend_no_budget():
|
||||
from unittest.mock import AsyncMock, call
|
||||
|
||||
entered: Final = asyncio.Event()
|
||||
release: Final = asyncio.Event()
|
||||
new_model_loaded: Final = asyncio.Event()
|
||||
|
||||
async def gated_load(key: str) -> bool:
|
||||
entered.set()
|
||||
await release.wait()
|
||||
if key == "new-model":
|
||||
new_model_loaded.set()
|
||||
return True
|
||||
|
||||
def is_loaded(key: str) -> bool:
|
||||
return key == "new-model" and new_model_loaded.is_set()
|
||||
|
||||
resync: Final = AsyncMock(side_effect=gated_load)
|
||||
read_through: Final = RegistryReadThrough(
|
||||
resync=resync,
|
||||
is_loaded=is_loaded,
|
||||
max_resyncs_per_window=2,
|
||||
resync_window_seconds=60.0,
|
||||
)
|
||||
|
||||
burst: Final = asyncio.gather(*(read_through.attempt("new-model") for _ in range(25)))
|
||||
await entered.wait()
|
||||
release.set()
|
||||
|
||||
assert await burst == [True] * 25
|
||||
assert resync.await_args_list == [call("new-model")]
|
||||
assert await read_through.attempt("other-model") is True
|
||||
assert resync.await_args_list == [call("new-model"), call("other-model")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resync_budget_replenishes_after_window():
|
||||
spy: Final = ResyncSpy(found=True)
|
||||
read_through: Final = RegistryReadThrough(resync=spy, max_resyncs_per_window=1, resync_window_seconds=0.05)
|
||||
read_through: Final = RegistryReadThrough(
|
||||
resync=spy, is_loaded=nothing_loaded, max_resyncs_per_window=1, resync_window_seconds=0.05
|
||||
)
|
||||
|
||||
assert await read_through.attempt("model-a") is True
|
||||
assert await read_through.attempt("model-b") is False
|
||||
|
|
@ -551,3 +600,100 @@ async def test_agent_read_through_hydrates_identity_binding(lookup, clean_agent_
|
|||
assert agent.identity is not None
|
||||
assert agent.identity.model_dump(include=set(binding)) == binding
|
||||
assert clean_agent_registry.get_agent_by_id(agent_id="agent-id").identity == agent.identity
|
||||
|
||||
|
||||
def test_model_is_loaded_matches_router_model_names_and_deployment_ids(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm import Router
|
||||
from litellm.proxy.common_utils.registry_read_through import _model_is_loaded
|
||||
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "loaded-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test"},
|
||||
"model_info": {"id": "loaded-deployment-id"},
|
||||
}
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert _model_is_loaded("loaded-model") is True
|
||||
assert _model_is_loaded("loaded-deployment-id") is True
|
||||
assert _model_is_loaded("model-created-on-a-sibling") is False
|
||||
|
||||
monkeypatch.setattr(proxy_server, "llm_router", None)
|
||||
assert _model_is_loaded("loaded-model") is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_read_through_answers_a_loaded_model_without_reading_the_db(monkeypatch: pytest.MonkeyPatch):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm import Router
|
||||
from litellm.proxy.common_utils.registry_read_through import model_registry_read_through
|
||||
|
||||
prisma_client: Final = MagicMock()
|
||||
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=AssertionError("db read"))
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "wired-loaded-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test"},
|
||||
}
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert await model_registry_read_through.attempt("wired-loaded-model") is True
|
||||
prisma_client.db.litellm_proxymodeltable.find_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_guardrail_read_through_answers_a_loaded_guardrail_without_reading_the_db(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy.common_utils.registry_read_through import guardrail_registry_read_through
|
||||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
from litellm.types.guardrails import Guardrail
|
||||
|
||||
guardrail_id: Final = "wired-loaded-guardrail-id"
|
||||
guardrail_name: Final = "wired-loaded-guardrail"
|
||||
prisma_client: Final = MagicMock()
|
||||
prisma_client.db.litellm_guardrailstable.find_first = AsyncMock(side_effect=AssertionError("db read"))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
|
||||
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
|
||||
guardrail=Guardrail(**dict(FakeGuardrailRow(guardrail_id, guardrail_name)))
|
||||
)
|
||||
try:
|
||||
assert await guardrail_registry_read_through.attempt(guardrail_name) is True
|
||||
prisma_client.db.litellm_guardrailstable.find_first.assert_not_awaited()
|
||||
finally:
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.delete_in_memory_guardrail(guardrail_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_read_through_answers_a_loaded_agent_without_reading_the_db(
|
||||
clean_agent_registry: "AgentRegistry", monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy.common_utils.registry_read_through import agent_registry_read_through
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", False)
|
||||
clean_agent_registry.register_agent(
|
||||
agent_config=AgentResponse.model_validate(
|
||||
FakeAgentRow("wired-loaded-agent-id", "wired-loaded-agent").model_dump()
|
||||
)
|
||||
)
|
||||
|
||||
assert await agent_registry_read_through.attempt("wired-loaded-agent-id") is True
|
||||
assert await agent_registry_read_through.attempt("wired-loaded-agent") is True
|
||||
|
|
|
|||
|
|
@ -411,6 +411,8 @@ def create_proxy_test_client(
|
|||
def fresh_agent_read_through(monkeypatch):
|
||||
from litellm.proxy.common_utils import registry_read_through
|
||||
|
||||
read_through = registry_read_through.RegistryReadThrough(resync=registry_read_through._resync_agents)
|
||||
read_through = registry_read_through.RegistryReadThrough(
|
||||
resync=registry_read_through._resync_agents, is_loaded=registry_read_through._agent_is_loaded
|
||||
)
|
||||
monkeypatch.setattr(registry_read_through, "agent_registry_read_through", read_through)
|
||||
return read_through
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue