mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
f115bb4ad0
commit
b5ac683a91
7 changed files with 140 additions and 11 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue