From 6ac26dbb13f8330e5baa0565b000fb5c9f188741 Mon Sep 17 00:00:00 2001 From: Prathamesh Gawas Date: Tue, 23 Jun 2026 18:24:37 +0530 Subject: [PATCH] fix(sso): enforce user limit for non-premium SSO users and add corresponding tests (#29677) --- litellm/proxy/management_endpoints/ui_sso.py | 37 ++-- .../proxy/management_endpoints/test_ui_sso.py | 202 ++++++++++++++++++ 2 files changed, 227 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 73ec56a82e6..cabb5c398ee 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -856,17 +856,7 @@ async def google_login( ####### Check if user is a Enterprise / Premium User ####### if microsoft_client_id is not None or google_client_id is not None or generic_client_id is not None: if premium_user is not True: - # Check if under 'free SSO user' limit - if prisma_client is not None: - total_users = await UserRepository(prisma_client).table.count() - if total_users and total_users > 5: - raise ProxyException( - message="You must be a LiteLLM Enterprise user to use SSO for more than 5 users. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this", - type=ProxyErrorTypes.auth_error, - param="premium_user", - code=status.HTTP_403_FORBIDDEN, - ) - else: + if prisma_client is None: raise ProxyException( message=CommonProxyErrors.db_not_connected_error.value, type=ProxyErrorTypes.auth_error, @@ -1588,6 +1578,8 @@ async def get_user_info_from_db( ) return user_info + except ProxyException: + raise except Exception as e: verbose_proxy_logger.exception(f"[Non-Blocking] Error trying to add sso user to db: {e}") @@ -2200,6 +2192,7 @@ async def cli_poll_key( async def insert_sso_user( result_openid: Optional[Union[OpenID, dict]], user_defined_values: Optional[SSOUserDefinedValues] = None, + prisma_client: Optional[PrismaClient] = None, ) -> NewUserResponse: """ Helper function to create a New User in LiteLLM DB after a successful SSO login @@ -2207,11 +2200,18 @@ async def insert_sso_user( Args: result_openid (OpenID): User information in OpenID format if the login was successful. user_defined_values (Optional[SSOUserDefinedValues], optional): LiteLLM SSOValues / fields that were read + prisma_client (Optional[PrismaClient], optional): Prisma client instance. When provided, + uses this instance directly instead of re-importing from proxy_server. Returns: Tuple[str, str]: User ID and User Role """ - verbose_proxy_logger.debug(f"Inserting SSO user into DB. User values: {user_defined_values}") + verbose_proxy_logger.debug( + f"Inserting SSO user into DB. User values: {user_defined_values}" + ) + from litellm.proxy.proxy_server import premium_user + + if result_openid is None: raise ValueError("result_openid is None") if isinstance(result_openid, dict): @@ -2220,6 +2220,16 @@ async def insert_sso_user( if user_defined_values is None: raise ValueError("user_defined_values is None") + if not premium_user and prisma_client is not None: + # Check if under 'free SSO user' limit + total_users = await prisma_client.db.litellm_usertable.count() + if total_users is not None and total_users >= 5: + raise ProxyException( + message="You must be a LiteLLM Enterprise user to use SSO for more than 5 users. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this", + type=ProxyErrorTypes.auth_error, + param="premium_user", + code=status.HTTP_403_FORBIDDEN, + ) # Apply default_internal_user_params if litellm.default_internal_user_params: # Preserve the SSO-extracted role if it's a valid LiteLLM role, @@ -2777,8 +2787,11 @@ class SSOAuthenticationHandler: user_info = await insert_sso_user( result_openid=result, user_defined_values=user_defined_values, + prisma_client=prisma_client, ) return user_info + except ProxyException: + raise except Exception as e: verbose_proxy_logger.exception(f"Error upserting SSO user into LiteLLM DB: {e}") return user_info diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 976048b9521..9b1523e1d73 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -769,6 +769,131 @@ def test_generic_response_convertor_normalizes_email(): assert result.display_name == "Test User" +@pytest.mark.asyncio +async def test_insert_sso_user_blocks_when_at_user_limit(): + """ + Non-premium: insert_sso_user raises ProxyException when the DB already has 5 users. + """ + from litellm.proxy._types import ProxyException, SSOUserDefinedValues + from litellm.proxy.management_endpoints.ui_sso import insert_sso_user + from litellm.proxy.management_endpoints.types import CustomOpenID + + mock_prisma = MagicMock() + mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=5) + + mock_openid = CustomOpenID( + id="new-user-1", + email="new@example.com", + display_name="New User", + provider="google", + team_ids=[], + ) + user_defined_values: SSOUserDefinedValues = { + "user_id": "new-user-1", + "user_email": "new@example.com", + "user_role": None, + "max_budget": None, + "budget_duration": None, + "models": [], + } + + with patch("litellm.proxy.proxy_server.premium_user", False): + with pytest.raises(ProxyException) as exc_info: + await insert_sso_user( + result_openid=mock_openid, + user_defined_values=user_defined_values, + prisma_client=mock_prisma, + ) + + assert str(exc_info.value.code) == "403" + + +@pytest.mark.asyncio +async def test_insert_sso_user_allows_when_under_user_limit(): + """ + Non-premium: insert_sso_user does NOT raise when DB has fewer than 5 users. + """ + from litellm.proxy._types import NewUserResponse, SSOUserDefinedValues + from litellm.proxy.management_endpoints.ui_sso import insert_sso_user + from litellm.proxy.management_endpoints.types import CustomOpenID + + mock_prisma = MagicMock() + mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=3) + + mock_openid = CustomOpenID( + id="new-user-2", + email="new2@example.com", + display_name="New User 2", + provider="google", + team_ids=[], + ) + user_defined_values: SSOUserDefinedValues = { + "user_id": "new-user-2", + "user_email": "new2@example.com", + "user_role": None, + "max_budget": None, + "budget_duration": None, + "models": [], + } + + mock_response = NewUserResponse(user_id="new-user-2", key="sk-test", teams=None) + + with patch("litellm.proxy.proxy_server.premium_user", False): + with patch( + "litellm.proxy.management_endpoints.ui_sso.new_user", + return_value=mock_response, + ): + result = await insert_sso_user( + result_openid=mock_openid, + user_defined_values=user_defined_values, + prisma_client=mock_prisma, + ) + + assert result.user_id == "new-user-2" + + +@pytest.mark.asyncio +async def test_upsert_sso_user_propagates_proxy_exception_for_new_user(): + """ + upsert_sso_user must not swallow ProxyException raised by insert_sso_user + (e.g. the non-premium user-limit enforcement). + """ + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + from litellm.proxy.management_endpoints.types import CustomOpenID + + mock_prisma = MagicMock() + mock_prisma.db.litellm_usertable.count = AsyncMock(return_value=5) + + sso_result = CustomOpenID( + id="over-limit-user", + email="overlimit@example.com", + display_name="Over Limit", + provider="google", + team_ids=[], + ) + user_defined_values = { + "user_id": "over-limit-user", + "user_email": "overlimit@example.com", + "user_role": None, + "max_budget": None, + "budget_duration": None, + "models": [], + } + + with patch("litellm.proxy.proxy_server.premium_user", False): + with pytest.raises(ProxyException) as exc_info: + await SSOAuthenticationHandler.upsert_sso_user( + result=sso_result, + user_info=None, + user_email="overlimit@example.com", + user_defined_values=user_defined_values, + prisma_client=mock_prisma, + ) + + assert str(exc_info.value.code) == "403" + + @pytest.mark.asyncio async def test_upsert_sso_user_updates_role_for_existing_user(): """ @@ -7131,3 +7256,80 @@ async def test_legacy_login_page_hides_credentials_hint_via_general_settings(): assert response.status_code == 200 assert "Default Credentials" not in body assert "MASTER_KEY" not in body + + +@pytest.mark.asyncio +async def test_get_user_info_from_db_propagates_proxy_exception(): + """ + get_user_info_from_db must re-raise ProxyException instead of swallowing it + (e.g. the non-premium user-limit ProxyException raised by upsert_sso_user). + """ + from litellm.proxy._types import ProxyErrorTypes, ProxyException + from litellm.proxy.management_endpoints.ui_sso import get_user_info_from_db + from litellm.proxy.management_endpoints.types import CustomOpenID + from starlette import status + + sso_result = CustomOpenID( + id="test-user", + email="test@example.com", + display_name="Test", + provider="google", + team_ids=[], + ) + expected_exc = ProxyException( + message="limit reached", + type=ProxyErrorTypes.auth_error, + param="premium_user", + code=status.HTTP_403_FORBIDDEN, + ) + + with patch( + "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.upsert_sso_user", + side_effect=expected_exc, + ), patch( + "litellm.proxy.management_endpoints.ui_sso.get_existing_user_info_from_db", + return_value=None, + ): + with pytest.raises(ProxyException) as exc_info: + await get_user_info_from_db( + result=sso_result, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + user_email="test@example.com", + user_defined_values={ + "user_id": "test-user", + "user_email": "test@example.com", + "user_role": None, + "max_budget": None, + "budget_duration": None, + "models": [], + }, + ) + + assert exc_info.value is expected_exc + + +@pytest.mark.asyncio +async def test_google_login_raises_when_sso_configured_and_no_db(): + """ + Non-premium: google_login raises ProxyException (db_not_connected) when an SSO + provider env var is set but prisma_client is None. + """ + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.ui_sso import google_login + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://proxy.example.com/" + + with ( + patch.dict(os.environ, {"GOOGLE_CLIENT_ID": "fake-client-id"}, clear=False), + patch("litellm.proxy.proxy_server.premium_user", False), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", None), + ): + with pytest.raises(ProxyException) as exc_info: + await google_login(request=mock_request) + + assert str(exc_info.value.code) == "403"