fix(auth): reread org and access-group policy for delegated grants

Delegated proxy:admin identities set requires_fresh_policy, but the
organization lookup and the team access-group model grants were still
answered from user_api_key_cache and the in-process access-group cache.
Thread check_db_only through get_org_object, get_org_object_for_request,
_organization_max_budget_check, common_checks, can_team_access_model and
_get_models_from_access_groups so a delegated request rereads them, and
fail closed instead of skipping the org budget check when that read fails.

Restore request_dispatched_to_provider_pass_through to its base predicate
and move the router-wide marker into
request_dispatched_to_marked_provider_pass_through, used only by the
delegated route gate. The widened predicate had skipped the
router_settings.model_group_alias rewrite for ordinary virtual keys on
provider routes that bind no endpoint param (/bespoke, /laya,
/comprehendmedical, /transcribe).

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
Tin Chi Lo 2026-10-03 14:15:23 -07:00
parent f115bb4ad0
commit b5ac683a91
7 changed files with 140 additions and 11 deletions

View file

@ -967,6 +967,7 @@ async def common_checks(
)
skip_all_budget_checks: Final = skip_budget_checks or route_skips_budget_checks(route=route)
fresh_policy: Final = valid_token is not None and valid_token.requires_fresh_policy
membership_user_id: Final = (
valid_token.user_id if valid_token is not None and (bool(_model) or not skip_all_budget_checks) else None
@ -979,6 +980,7 @@ async def common_checks(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=fresh_policy,
)
if team_object is not None and membership_user_id is not None
else None
@ -1011,6 +1013,7 @@ async def common_checks(
llm_router=llm_router,
team_model_aliases=(valid_token.team_model_aliases if valid_token else None),
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
check_db_only=fresh_policy,
)
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
@ -4123,6 +4126,8 @@ async def get_org_object(
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
include_budget_table: bool = False,
*,
check_db_only: bool = False,
) -> LiteLLM_OrganizationTable | None:
"""
- Check if org id in proxy Org Table
@ -4147,10 +4152,10 @@ async def get_org_object(
if include_budget_table:
cache_key = f"org_id:{org_id}:with_budget"
# check if in cache
deserialized_org: Final = await user_api_key_cache.async_get_cache(
key=cache_key,
model_type=LiteLLM_OrganizationTable,
deserialized_org: Final = (
None
if check_db_only
else await user_api_key_cache.async_get_cache(key=cache_key, model_type=LiteLLM_OrganizationTable)
)
if deserialized_org is not None:
return deserialized_org
@ -4215,6 +4220,7 @@ async def get_org_object_for_request(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging | None,
check_db_only: bool = False,
) -> LiteLLM_OrganizationTable | None:
try:
org: Final = await get_org_object(
@ -4224,10 +4230,13 @@ async def get_org_object_for_request(
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
include_budget_table=True,
check_db_only=check_db_only,
)
except OrganizationNotFoundError:
return None
except Exception as e:
if check_db_only:
raise
if not PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(e):
verbose_proxy_logger.debug("org lookup failed, continuing without org limits", exc_info=True)
return None
@ -4314,6 +4323,7 @@ async def _get_models_from_access_groups(
prisma_client: DatabaseClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> list[str]:
"""
Collect model names from unified access groups.
@ -4325,6 +4335,7 @@ async def _get_models_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
@ -5138,6 +5149,7 @@ async def can_team_access_model(
team_model_aliases: dict[str, str] | None = None,
key_model_aliases: Mapping[str, str] | None = None,
prisma_client: DatabaseClient | None = None,
check_db_only: bool = False,
) -> Literal[True]:
"""
Returns True if the team can access a specific model.
@ -5162,6 +5174,7 @@ async def can_team_access_model(
models_from_groups: Final = await _get_models_from_access_groups(
access_group_ids=team_access_group_ids,
prisma_client=prisma_client,
check_db_only=check_db_only,
)
if models_from_groups:
return _can_object_call_model(
@ -6354,8 +6367,11 @@ async def _organization_max_budget_check(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
include_budget_table=True,
check_db_only=valid_token.requires_fresh_policy,
)
except Exception:
if valid_token.requires_fresh_policy:
raise
# If organization lookup fails, skip the check
return

View file

@ -1984,10 +1984,12 @@ def request_dispatched_to_pass_through_endpoint(request: Request | None) -> bool
def request_dispatched_to_provider_pass_through(request: Request) -> bool:
return (
getattr(request.scope.get("endpoint"), LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER, False) is True
or "endpoint" in request.path_params
)
"""Built-in provider pass-through handlers (``/anthropic/{endpoint:path}``, ...) bind ``endpoint``."""
return "endpoint" in request.path_params
def request_dispatched_to_marked_provider_pass_through(request: Request) -> bool:
return getattr(request.scope.get("endpoint"), LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER, False) is True
def get_model_from_request(

View file

@ -17,6 +17,7 @@ from litellm.proxy._types import (
from .auth_checks_organization import _user_is_org_admin
from .auth_utils import (
get_request_route_template,
request_dispatched_to_marked_provider_pass_through,
request_dispatched_to_pass_through_endpoint,
request_dispatched_to_provider_pass_through,
)
@ -111,6 +112,7 @@ class RouteChecks:
route.isprintable()
and not request_dispatched_to_pass_through_endpoint(request)
and not request_dispatched_to_provider_pass_through(request)
and not request_dispatched_to_marked_provider_pass_through(request)
and not RouteChecks.check_route_access(template, excluded)
and RouteChecks.check_route_access(template, allowed)
)

View file

@ -2793,6 +2793,7 @@ async def _inherit_org_identity(
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=user_api_key_auth_obj.requires_fresh_policy,
)
if org_object is None:
return

View file

@ -3,12 +3,26 @@ from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from fastapi import HTTPException
from fastapi import HTTPException, Request
import litellm
from litellm.models.access_group import LiteLLM_AccessGroupTable
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import SessionPrincipal
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership, LiteLLM_TeamTable, LiteLLM_UserTable, Member
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_OrganizationTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
LiteLLM_UserTable,
Member,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import common_checks
from litellm.proxy.auth.delegated_oauth import delegated_identity
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import ProxyLogging
@pytest.fixture
@ -47,7 +61,9 @@ async def test_delegated_identity_uses_current_team_limits_and_roster(database:
)
database.writer_db.litellm_teamtable.find_unique = AsyncMock(
return_value=LiteLLM_TeamTable(
team_id="team", models=["allowed-model"], rpm_limit=10,
team_id="team",
models=["allowed-model"],
rpm_limit=10,
members_with_roles=[Member(user_id="admin", role="user")],
)
)
@ -76,3 +92,74 @@ async def test_delegated_identity_fails_closed_when_database_is_unavailable(data
with pytest.raises(HTTPException) as error:
await delegated_identity(SessionPrincipal(user_id="admin", client_id="app"))
assert error.value.status_code == 503
async def _common_checks(token: UserAPIKeyAuth, team: LiteLLM_TeamTable | None = None) -> bool:
logging: Final = MagicMock(spec=ProxyLogging)
logging.budget_alerts = AsyncMock()
return await common_checks(
request_body={"model": "gpt-4"},
team_object=team,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=None,
proxy_logging_obj=logging,
valid_token=token,
request=Request(
{"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": [], "query_string": b""}
),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("fresh", [False, True])
async def test_common_checks_reread_organization_budget_for_delegated_identity(
database: MagicMock, fresh: bool
) -> None:
from litellm.proxy import proxy_server
cached: Final = LiteLLM_OrganizationTable(
organization_id="org",
budget_id="budget",
created_by="admin",
updated_by="admin",
litellm_budget_table=LiteLLM_BudgetTable(max_budget=100.0),
)
await proxy_server.user_api_key_cache.async_set_cache(key="org_id:org:with_budget", value=cached)
database.db.litellm_organizationtable.find_unique = AsyncMock(
return_value=cached.model_copy(update={"litellm_budget_table": LiteLLM_BudgetTable(max_budget=0.0)})
)
token: Final = UserAPIKeyAuth(token="token", user_id="admin", org_id="org")
token.requires_fresh_policy = fresh
if not fresh:
assert await _common_checks(token) is True
return
with pytest.raises(litellm.BudgetExceededError):
await _common_checks(token)
@pytest.mark.asyncio
@pytest.mark.parametrize("fresh", [False, True])
async def test_common_checks_reread_team_access_groups_for_delegated_identity(database: MagicMock, fresh: bool) -> None:
from litellm.proxy import proxy_server
cached: Final = LiteLLM_AccessGroupTable(
access_group_id="group", access_group_name="group", access_model_names=["gpt-4"]
)
await proxy_server.user_api_key_cache.async_set_cache(key="access_group_id:group", value=cached)
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(
return_value=cached.model_copy(update={"access_model_names": []})
)
database.writer_db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
team: Final = LiteLLM_TeamTable(team_id="team", models=["other-model"], access_group_ids=["group"])
token: Final = UserAPIKeyAuth(token="token", user_id="admin", team_id="team")
token.requires_fresh_policy = fresh
if not fresh:
assert await _common_checks(token, team) is True
return
with pytest.raises(ProxyException) as error:
await _common_checks(token, team)
assert error.value.type == ProxyErrorTypes.team_model_access_denied

View file

@ -62,6 +62,26 @@ def test_delegated_admin_scope_limits_api_access(route: str, allowed: bool) -> N
assert RouteChecks.is_delegated_admin_route(route, request) is allowed
def test_marked_provider_pass_through_without_endpoint_param_keeps_model_alias_dispatch() -> None:
from litellm.proxy.auth.auth_utils import (
request_dispatched_to_marked_provider_pass_through,
request_dispatched_to_provider_pass_through,
)
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import laya_proxy_route
request: Final = Request(
{
"type": "http",
"method": "POST",
"path": "/laya/v1/systemone",
"endpoint": laya_proxy_route,
"path_params": {},
}
)
assert request_dispatched_to_marked_provider_pass_through(request)
assert not request_dispatched_to_provider_pass_through(request)
DAILY_ACTIVITY_ROUTE_PAIRS: Final[tuple[tuple[str, str], ...]] = (
("/user/daily/activity", "/user/daily/activity/aggregated"),
("/user/daily/activity", "/user/daily/activity/aggregated/keys"),

View file

@ -373,6 +373,7 @@ async def test_team_member_budget_check_blocks_regenerated_key_after_old_key_exh
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
check_db_only=False,
)
assert "Budget has been exceeded" in str(exc_info.value)
assert "test-user-1" in str(exc_info.value)