mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(team): refresh team cache on team_model_add/delete (LIT-3244)
team_model_add and team_model_delete wrote to the DB but did not invalidate the in-memory LiteLLM_TeamTableCachedObj used by common_checks. After the v1.83.14 common_checks centralization made team.models authoritative on /v1/files and /v1/vector_stores/*, adding a Team-BYOK model silently failed to grant the new public model name to team members until the cache TTL expired (and a removed model kept working until then on the symmetric path). Extract the cache-refresh snippet from update_team into a small helper and apply it consistently at all three team-write sites.
This commit is contained in:
parent
1b141bc588
commit
3d3ae40467
2 changed files with 151 additions and 8 deletions
|
|
@ -64,6 +64,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_cache_team_object,
|
||||
allowed_route_check_inside_route,
|
||||
can_org_access_model,
|
||||
get_org_object,
|
||||
|
|
@ -129,6 +130,33 @@ def _sanitize_for_log(value: Any) -> str:
|
|||
return text.replace("\r", "").replace("\n", "")
|
||||
|
||||
|
||||
async def _refresh_cached_team(
|
||||
team_row: Any,
|
||||
user_api_key_cache: Any,
|
||||
proxy_logging_obj: Any,
|
||||
) -> None:
|
||||
"""
|
||||
Refresh the in-memory cached team object after a DB write.
|
||||
|
||||
Every endpoint that mutates `litellm_teamtable` must call this so the
|
||||
cached `LiteLLM_TeamTableCachedObj` used by `common_checks` stays in
|
||||
sync. Without this, subsequent auth checks read a stale team and can
|
||||
403 on permissions the DB has already granted (or, symmetrically,
|
||||
keep granting permissions the DB has already revoked).
|
||||
|
||||
`team_row` is the Prisma row returned by `update`/`find_unique` on
|
||||
`litellm_teamtable`. It is converted to `LiteLLM_TeamTableCachedObj`
|
||||
via `model_dump()` to match the cache shape `_cache_team_object`
|
||||
expects.
|
||||
"""
|
||||
await _cache_team_object(
|
||||
team_id=team_row.team_id,
|
||||
team_table=LiteLLM_TeamTableCachedObj(**team_row.model_dump()),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
async def _verify_team_access(
|
||||
team_obj: LiteLLM_TeamTable,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -1585,7 +1613,6 @@ async def update_team( # noqa: PLR0915
|
|||
```
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.auth.auth_checks import _cache_team_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
litellm_proxy_admin_name,
|
||||
llm_router,
|
||||
|
|
@ -1865,9 +1892,8 @@ async def update_team( # noqa: PLR0915
|
|||
verbose_proxy_logger.info(
|
||||
"Successfully updated team - %s, info", team_row.team_id
|
||||
)
|
||||
await _cache_team_object(
|
||||
team_id=team_row.team_id,
|
||||
team_table=LiteLLM_TeamTableCachedObj(**team_row.model_dump()),
|
||||
await _refresh_cached_team(
|
||||
team_row=team_row,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
@ -4560,7 +4586,11 @@ async def team_model_add(
|
|||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail={"error": "No db connected"})
|
||||
|
|
@ -4599,6 +4629,12 @@ async def team_model_add(
|
|||
where={"team_id": data.team_id}, data={"models": updated_models}
|
||||
)
|
||||
|
||||
await _refresh_cached_team(
|
||||
team_row=updated_team,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return updated_team
|
||||
|
||||
|
||||
|
|
@ -4631,7 +4667,11 @@ async def team_model_delete(
|
|||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail={"error": "No db connected"})
|
||||
|
|
@ -4675,6 +4715,12 @@ async def team_model_delete(
|
|||
where={"team_id": data.team_id}, data={"models": updated_models}
|
||||
)
|
||||
|
||||
await _refresh_cached_team(
|
||||
team_row=updated_team,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return updated_team
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1540,6 +1540,99 @@ def test_add_new_models_to_team_with_existing_models():
|
|||
assert updated_models.sort() == ["model1", "model2", "model3", "model4"].sort()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint_name",
|
||||
["team_model_add", "team_model_delete"],
|
||||
)
|
||||
async def test_team_model_add_delete_refresh_team_cache(endpoint_name):
|
||||
"""
|
||||
Regression pin for LIT-3244 vector-store BYOK 403.
|
||||
|
||||
`team_model_add` and `team_model_delete` mutate `team.models` in the
|
||||
DB. Without a cache refresh, the in-memory `LiteLLM_TeamTableCachedObj`
|
||||
used by `common_checks` stays stale and team members 403 on a model
|
||||
the DB has just granted (or, symmetrically, keep using a model the DB
|
||||
has just revoked).
|
||||
|
||||
Pin: after the DB update, the endpoint must call `_cache_team_object`
|
||||
with the updated team row so the cached team stays in sync.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LitellmUserRoles,
|
||||
TeamModelAddRequest,
|
||||
TeamModelDeleteRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
team_model_add,
|
||||
team_model_delete,
|
||||
)
|
||||
|
||||
mock_request = Mock(spec=Request)
|
||||
mock_user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id"
|
||||
)
|
||||
|
||||
existing_team = MagicMock()
|
||||
existing_team.model_dump.return_value = {
|
||||
"team_id": "team-1234",
|
||||
"models": ["bedrock-claude-sonnet-4", "openai/*"],
|
||||
}
|
||||
|
||||
updated_team = MagicMock()
|
||||
updated_team.team_id = "team-1234"
|
||||
updated_team.model_dump.return_value = {
|
||||
"team_id": "team-1234",
|
||||
"models": ["bedrock-claude-sonnet-4", "openai/*", "team-byok-1"],
|
||||
}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client,
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._cache_team_object"
|
||||
) as mock_cache_team,
|
||||
):
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=existing_team
|
||||
)
|
||||
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(
|
||||
return_value=updated_team
|
||||
)
|
||||
mock_cache_team.return_value = None
|
||||
|
||||
if endpoint_name == "team_model_add":
|
||||
await team_model_add(
|
||||
data=TeamModelAddRequest(team_id="team-1234", models=["team-byok-1"]),
|
||||
http_request=mock_request,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
else:
|
||||
await team_model_delete(
|
||||
data=TeamModelDeleteRequest(team_id="team-1234", models=["openai/*"]),
|
||||
http_request=mock_request,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# The pin: cache refresh must run with the updated team row.
|
||||
assert mock_cache_team.await_count == 1, (
|
||||
f"{endpoint_name} must call _cache_team_object exactly once "
|
||||
f"after the DB update (LIT-3244 regression pin); "
|
||||
f"got await_count={mock_cache_team.await_count}"
|
||||
)
|
||||
call_kwargs = mock_cache_team.await_args.kwargs
|
||||
assert call_kwargs["team_id"] == "team-1234"
|
||||
# The cached object must be built from the *updated* row, not the
|
||||
# pre-mutation `existing_team` — that's the whole point.
|
||||
assert call_kwargs["team_table"].team_id == "team-1234"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_team_team_member_budget_not_passed_to_db():
|
||||
"""
|
||||
|
|
@ -1568,7 +1661,9 @@ async def test_update_team_team_member_budget_not_passed_to_db():
|
|||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
patch("litellm.proxy.auth.auth_checks._cache_team_object") as mock_cache_team,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._cache_team_object"
|
||||
) as mock_cache_team,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table"
|
||||
) as mock_upsert_budget,
|
||||
|
|
@ -1999,7 +2094,9 @@ async def test_update_team_with_team_member_budget_duration():
|
|||
patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"),
|
||||
patch("litellm.proxy.auth.auth_checks._cache_team_object") as mock_cache_team,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._cache_team_object"
|
||||
) as mock_cache_team,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table"
|
||||
) as mock_upsert_budget,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue