fix(proxy): greptile review fixes

four things from the bot review:

1. validate_models_exist bailed early with (False, all_names) when
   llm_router was None - never even checked known_access_groups.
   db-only setups couldn't create a group of groups at all. now
   we check known_groups even without a router.

2. get_group_memberships_from_db only caught AttributeError +
   TypeError, so a transient DB error would 500 every /v1/models
   call. broadened to Exception - the fallback now actually
   catches what the docstring says it does.

3. delete_group_membership_edges does WHERE parent OR child but
   only parent had an index. added @@index([child_group]) + a
   tiny migration.

4. get_group_memberships_from_db was being called on every
   model-listing request. added a 60s in-process TTL cache via
   get_cached_group_memberships, with explicit invalidation on
   every membership write. same shape as the existing
   llm_router.get_model_access_groups() cache.

9 tests for the above.

refs #28032
This commit is contained in:
Ashwin Upadhyay 2026-05-17 01:13:47 +05:30
parent 917b28588c
commit c5e2c6cc63
10 changed files with 259 additions and 23 deletions

View file

@ -0,0 +1,2 @@
-- CreateIndex
CREATE INDEX IF NOT EXISTS "LiteLLM_AccessGroupMembership_child_group_idx" ON "LiteLLM_AccessGroupMembership"("child_group");

View file

@ -1388,4 +1388,5 @@ model LiteLLM_AccessGroupMembership {
@@unique([parent_group, child_group])
@@index([parent_group])
@@index([child_group])
}

View file

@ -6,6 +6,7 @@ Endpoints here:
"""
import json
import time
from typing import Any, Dict, List, Optional, Set, Tuple
from fastapi import APIRouter, Depends, HTTPException
@ -32,6 +33,29 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import
router = APIRouter()
# ---------------------------------------------------------------------------
# Per-process membership-map cache
# ---------------------------------------------------------------------------
# get_group_memberships_from_db is called on every /v1/models and /model/info
# request via get_available_models_for_user. Without a cache that's one extra
# Prisma roundtrip per request - bad under burst traffic. We cache the map
# in process memory for a short TTL and invalidate explicitly on writes
# (upsert/delete) so the consistency window inside the writing process is
# zero. Across processes, eventual consistency is bounded by the TTL -
# matches today's behavior for llm_router.get_model_access_groups() which is
# also per-process.
_MEMBERSHIPS_CACHE_TTL_SECONDS = 60.0
_MEMBERSHIPS_CACHE: Optional[Tuple[float, Dict[str, List[str]]]] = None
def invalidate_group_memberships_cache() -> None:
"""Drop the in-process membership cache. Call after any write that
mutates the LiteLLM_AccessGroupMembership table."""
global _MEMBERSHIPS_CACHE
_MEMBERSHIPS_CACHE = None
def validate_models_exist(
model_names: List[str],
llm_router,
@ -44,11 +68,17 @@ def validate_models_exist(
Returns:
Tuple[bool, List[str]]: (all_valid, missing_names)
"""
known_groups = known_access_groups or set()
if llm_router is None:
return False, model_names
# DB-only deployment: no in-memory router means we cannot validate
# real model names, but known_access_groups is still authoritative
# for nested-group composition. Anything not in known_groups is
# reported as missing (fail-closed).
missing = [m for m in model_names if m not in known_groups]
return (len(missing) == 0, missing)
router_model_names = set(llm_router.get_model_names())
known_groups = known_access_groups or set()
missing = [
m for m in model_names if m not in router_model_names and m not in known_groups
]
@ -87,17 +117,18 @@ async def get_group_memberships_from_db(
Build parent_group -> [child_groups] map from the membership table.
Single query, in-memory bucketing - no N+1.
Resilient by design: if the table isn't available (Prisma client predates
this migration, the proxy started before `prisma migrate deploy` finished,
or the membership Prisma model was stripped from a downstream build) we
return an empty map. The auth path then falls back to today's flat-group
semantics instead of 500-ing the whole request.
Resilient by design: any failure to read the membership table (missing
Prisma model, migration race, transient DB/network error, query timeout)
degrades to an empty map. The auth path then falls back to today's
flat-group semantics instead of 500-ing the whole request. We log at
debug so ops can correlate fallback periods with incidents without
drowning normal traffic in warnings.
"""
try:
rows = await prisma_client.db.litellm_accessgroupmembership.find_many()
except (AttributeError, TypeError) as e:
except Exception as e: # noqa: BLE001 - intentional broad catch on auth path
verbose_proxy_logger.debug(
"litellm_accessgroupmembership unavailable - "
"litellm_accessgroupmembership read failed - "
"skipping nested group resolution: %s",
e,
)
@ -109,6 +140,25 @@ async def get_group_memberships_from_db(
return memberships
async def get_cached_group_memberships(
prisma_client: PrismaClient,
) -> Dict[str, List[str]]:
"""
TTL-cached wrapper around get_group_memberships_from_db. Hot-path
callers (model-listing endpoints) should use this; tests and write
paths that need fresh data can call the underlying helper directly.
"""
global _MEMBERSHIPS_CACHE
now = time.monotonic()
if _MEMBERSHIPS_CACHE is not None:
cached_at, value = _MEMBERSHIPS_CACHE
if now - cached_at < _MEMBERSHIPS_CACHE_TTL_SECONDS:
return value
fresh = await get_group_memberships_from_db(prisma_client=prisma_client)
_MEMBERSHIPS_CACHE = (now, fresh)
return fresh
async def upsert_group_memberships(
parent_group: str,
child_groups: List[str],
@ -142,6 +192,7 @@ async def upsert_group_memberships(
data=rows,
skip_duplicates=True,
)
invalidate_group_memberships_cache()
return result
@ -164,6 +215,7 @@ async def delete_group_membership_edges(
]
}
)
invalidate_group_memberships_cache()
return result
@ -850,6 +902,7 @@ async def update_access_group(
await prisma_client.db.litellm_accessgroupmembership.delete_many(
where={"parent_group": access_group}
)
invalidate_group_memberships_cache()
# Step 2: re-add membership using the appropriate write path
if use_model_ids:

View file

@ -367,7 +367,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
router as key_management_router,
)
from litellm.proxy.management_endpoints.model_access_group_management_endpoints import (
get_group_memberships_from_db,
get_cached_group_memberships,
router as model_access_group_management_router,
)
from litellm.proxy.management_endpoints.model_management_endpoints import (
@ -11978,11 +11978,11 @@ async def model_info_v1( # noqa: PLR0915
proxy_model_list = llm_router.get_model_names()
model_access_groups = llm_router.get_model_access_groups()
# Parent->child edges for nested access groups. Empty when no DB is
# configured, preserving today's flat behavior.
# Parent->child edges for nested access groups (TTL-cached per process).
# Empty when no DB is configured, preserving today's flat behavior.
group_memberships: Dict[str, List[str]] = {}
if prisma_client is not None:
group_memberships = await get_group_memberships_from_db(
group_memberships = await get_cached_group_memberships(
prisma_client=prisma_client
)

View file

@ -1388,4 +1388,5 @@ model LiteLLM_AccessGroupMembership {
@@unique([parent_group, child_group])
@@index([parent_group])
@@index([child_group])
}

View file

@ -5848,7 +5848,7 @@ async def get_available_models_for_user(
get_team_models,
)
from litellm.proxy.management_endpoints.model_access_group_management_endpoints import (
get_group_memberships_from_db,
get_cached_group_memberships,
)
from litellm.proxy.management_endpoints.team_endpoints import validate_membership
@ -5860,11 +5860,12 @@ async def get_available_models_for_user(
proxy_model_list = llm_router.get_model_names()
model_access_groups = llm_router.get_model_access_groups()
# Parent->child edges for nested access groups. Empty when no DB is
# configured (e.g. SDK-only mode), preserving today's flat behavior.
# Parent->child edges for nested access groups (TTL-cached per process).
# Empty when no DB is configured (e.g. SDK-only mode), preserving
# today's flat behavior.
group_memberships: Dict[str, List[str]] = {}
if prisma_client is not None:
group_memberships = await get_group_memberships_from_db(
group_memberships = await get_cached_group_memberships(
prisma_client=prisma_client
)

View file

@ -1388,4 +1388,5 @@ model LiteLLM_AccessGroupMembership {
@@unique([parent_group, child_group])
@@index([parent_group])
@@index([child_group])
}

View file

@ -541,16 +541,29 @@ def test_validate_models_exist_reports_missing_in_input_order():
assert missing == ["z-missing", "y-missing"]
def test_validate_models_exist_with_null_router_returns_false():
"""No router - everything reports as missing (matches today's defensive behavior)."""
def test_validate_models_exist_with_null_router_still_accepts_known_groups():
"""DB-only deployment: llm_router is None but known_access_groups is still authoritative
for nested-group composition - only names not in known_groups are reported missing.
"""
all_valid, missing = validate_models_exist(
model_names=["any"],
model_names=["image", "reasoning"],
llm_router=None,
known_access_groups={"any"},
known_access_groups={"image", "reasoning"},
)
assert all_valid is True
assert missing == []
def test_validate_models_exist_with_null_router_rejects_unknown_real_models():
"""Without a router we can't validate real model names, so anything not in
known_access_groups is fail-closed reported as missing."""
all_valid, missing = validate_models_exist(
model_names=["gpt-4", "image"],
llm_router=None,
known_access_groups={"image"},
)
# Without a router we can't say what's a model, so we fall back to fail-closed
assert all_valid is False
assert missing == ["any"]
assert missing == ["gpt-4"]
def test_resolve_with_empty_models_and_empty_memberships_returns_empty():

View file

@ -0,0 +1,150 @@
"""
Cache-behavior tests for the nested-access-group membership map (#28032).
Hot-path callers go through get_cached_group_memberships() which TTL-caches
get_group_memberships_from_db() and is invalidated by every membership
write. These tests pin the cache hit/miss/invalidation semantics so the
optimization can't silently break later.
"""
import os
import sys
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
sys.path.insert(0, os.path.abspath("../../.."))
import pytest
import litellm.proxy.management_endpoints.model_access_group_management_endpoints as mgmt
from litellm.proxy.management_endpoints.model_access_group_management_endpoints import (
delete_group_membership_edges,
get_cached_group_memberships,
invalidate_group_memberships_cache,
upsert_group_memberships,
)
def _row(parent: str, child: str) -> SimpleNamespace:
return SimpleNamespace(parent_group=parent, child_group=child)
def _make_prisma(membership_rows=None):
membership_rows = membership_rows or []
db = MagicMock()
db.litellm_accessgroupmembership = MagicMock()
db.litellm_accessgroupmembership.find_many = AsyncMock(return_value=membership_rows)
db.litellm_accessgroupmembership.create_many = AsyncMock(return_value=0)
db.litellm_accessgroupmembership.delete_many = AsyncMock(return_value=0)
client = MagicMock()
client.db = db
return client
@pytest.fixture(autouse=True)
def _reset_cache_between_tests():
"""Module-level cache state must not leak between tests."""
invalidate_group_memberships_cache()
yield
invalidate_group_memberships_cache()
@pytest.mark.asyncio
async def test_cache_miss_then_hit_avoids_second_db_query():
prisma = _make_prisma(membership_rows=[_row("project-x", "image")])
first = await get_cached_group_memberships(prisma_client=prisma)
second = await get_cached_group_memberships(prisma_client=prisma)
assert first == second == {"project-x": ["image"]}
# Only the first call should hit the DB
prisma.db.litellm_accessgroupmembership.find_many.assert_awaited_once()
@pytest.mark.asyncio
async def test_cache_invalidation_forces_db_refetch():
prisma = _make_prisma(membership_rows=[_row("project-x", "image")])
await get_cached_group_memberships(prisma_client=prisma)
invalidate_group_memberships_cache()
await get_cached_group_memberships(prisma_client=prisma)
assert prisma.db.litellm_accessgroupmembership.find_many.await_count == 2
@pytest.mark.asyncio
async def test_upsert_invalidates_cache():
"""Writing edges must drop the cache so the next read sees the change."""
prisma = _make_prisma(membership_rows=[_row("project-x", "image")])
prisma.db.litellm_accessgroupmembership.create_many = AsyncMock(return_value=1)
await get_cached_group_memberships(prisma_client=prisma) # populates cache
await upsert_group_memberships(
parent_group="project-x",
child_groups=["reasoning"],
prisma_client=prisma,
)
await get_cached_group_memberships(prisma_client=prisma) # must re-fetch
assert prisma.db.litellm_accessgroupmembership.find_many.await_count == 2
@pytest.mark.asyncio
async def test_delete_edges_invalidates_cache():
"""Deleting edges must drop the cache too."""
prisma = _make_prisma(membership_rows=[_row("project-x", "image")])
prisma.db.litellm_accessgroupmembership.delete_many = AsyncMock(return_value=1)
await get_cached_group_memberships(prisma_client=prisma)
await delete_group_membership_edges(access_group="project-x", prisma_client=prisma)
await get_cached_group_memberships(prisma_client=prisma)
assert prisma.db.litellm_accessgroupmembership.find_many.await_count == 2
@pytest.mark.asyncio
async def test_cache_expires_after_ttl(monkeypatch):
"""When monotonic time advances past the TTL, the next read re-fetches."""
prisma = _make_prisma(membership_rows=[_row("project-x", "image")])
# Freeze time; advance past TTL between calls
now = [1000.0]
monkeypatch.setattr(mgmt.time, "monotonic", lambda: now[0])
await get_cached_group_memberships(prisma_client=prisma)
now[0] += mgmt._MEMBERSHIPS_CACHE_TTL_SECONDS + 1
await get_cached_group_memberships(prisma_client=prisma)
assert prisma.db.litellm_accessgroupmembership.find_many.await_count == 2
@pytest.mark.asyncio
async def test_cache_within_ttl_does_not_refetch(monkeypatch):
"""Reads inside the TTL window stay served from cache."""
prisma = _make_prisma(membership_rows=[_row("project-x", "image")])
now = [1000.0]
monkeypatch.setattr(mgmt.time, "monotonic", lambda: now[0])
await get_cached_group_memberships(prisma_client=prisma)
now[0] += mgmt._MEMBERSHIPS_CACHE_TTL_SECONDS - 1
await get_cached_group_memberships(prisma_client=prisma)
prisma.db.litellm_accessgroupmembership.find_many.assert_awaited_once()
@pytest.mark.asyncio
async def test_cache_falls_through_empty_dict_on_error_path():
"""When the underlying helper returns {} due to a DB error, the cache
still stores it - we don't want to retry on every single request."""
prisma = _make_prisma()
prisma.db.litellm_accessgroupmembership.find_many = AsyncMock(
side_effect=ConnectionError("postgres unreachable")
)
first = await get_cached_group_memberships(prisma_client=prisma)
second = await get_cached_group_memberships(prisma_client=prisma)
assert first == second == {}
# Only one DB attempt; subsequent calls served from the cached {}
prisma.db.litellm_accessgroupmembership.find_many.assert_awaited_once()

View file

@ -113,6 +113,20 @@ async def test_get_group_memberships_returns_empty_when_table_attribute_missing(
assert await get_group_memberships_from_db(prisma_client=NoMembershipTable()) == {}
@pytest.mark.asyncio
async def test_get_group_memberships_returns_empty_on_transient_db_error():
"""
Generic DB error (connection timeout, Prisma query failure, network blip)
must NOT propagate as a 500 on the auth path - we fall back to empty so
model-listing requests keep working until the DB recovers.
"""
prisma = _make_prisma()
prisma.db.litellm_accessgroupmembership.find_many = AsyncMock(
side_effect=ConnectionError("postgres unreachable")
)
assert await get_group_memberships_from_db(prisma_client=prisma) == {}
# ---------------------------------------------------------------------------
# upsert_group_memberships
# ---------------------------------------------------------------------------