diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e6de6d0142f..4117d4adc99 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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 diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index fffda2e98df..08538edcd2b 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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( diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 74d26aa1689..f372e97c831 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -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) ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index eb9a49b965e..020d777f12e 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/tests/unit/proxy/auth/test_delegated_oauth.py b/tests/unit/proxy/auth/test_delegated_oauth.py index 0f7aa4edc9a..d0f587035ac 100644 --- a/tests/unit/proxy/auth/test_delegated_oauth.py +++ b/tests/unit/proxy/auth/test_delegated_oauth.py @@ -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 diff --git a/tests/unit/proxy/auth/test_route_checks.py b/tests/unit/proxy/auth/test_route_checks.py index 154cab0151a..834981faeb9 100644 --- a/tests/unit/proxy/auth/test_route_checks.py +++ b/tests/unit/proxy/auth/test_route_checks.py @@ -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"), diff --git a/tests/unit/proxy/auth/test_team_member_budget.py b/tests/unit/proxy/auth/test_team_member_budget.py index b38a953d189..f1209beec90 100644 --- a/tests/unit/proxy/auth/test_team_member_budget.py +++ b/tests/unit/proxy/auth/test_team_member_budget.py @@ -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)