diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 60a40b3ede9..a1c70d78902 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2417,6 +2417,7 @@ class ExperimentalUIJWTToken: user_info: LiteLLM_UserTable, team_id: Optional[str] = None, team_alias: Optional[str] = None, + max_budget: Optional[float] = None, ) -> str: """ Generate a JWT token for CLI authentication with configurable expiration. @@ -2462,6 +2463,7 @@ class ExperimentalUIJWTToken: key_name=session_alias, key_alias=session_alias, expires=expires, + max_budget=max_budget, user_id=user_info.user_id, team_id=_team_id, team_alias=team_alias, diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 199de54ff09..3b626ba6edc 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -2243,8 +2243,8 @@ async def cli_poll_key( key_id: The CLI login session ID team_id: Optional team ID to assign to the JWT. If provided, must be one of user's teams. """ - from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken - from litellm.proxy.proxy_server import user_api_key_cache + from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_team_object, get_user_object + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache try: flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache) @@ -2320,18 +2320,43 @@ async def cli_poll_key( None, ) - # Create user object for JWT generation user_info = LiteLLM_UserTable( user_id=user_id, user_role=session_data["user_role"], models=session_data.get("models", []), - max_budget=litellm.max_ui_session_budget, ) - # Generate CLI JWT on-demand (expiration configurable via LITELLM_CLI_JWT_EXPIRATION_HOURS) - # Pass selected team_id to ensure JWT has correct team + user_db_obj = await get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + ) + user_budget = user_db_obj.max_budget if user_db_obj is not None else None + + team_budget: Optional[float] = None + if team_id is not None: + try: + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + team_budget = team_obj.max_budget + except Exception: + pass + + session_max_budget = ( + litellm.max_ui_session_budget + if user_budget is None and team_budget is None + else None + ) + jwt_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( - user_info=user_info, team_id=team_id, team_alias=team_alias + user_info=user_info, + team_id=team_id, + team_alias=team_alias, + max_budget=session_max_budget, ) # Delete cache entry (single-use) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 6d343af5b15..f56e309a552 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -538,6 +538,26 @@ def test_get_cli_jwt_auth_token_unique_per_session(valid_sso_user_defined_values assert first["key_name"] == second["key_name"] == expected_alias +def test_get_cli_jwt_auth_token_applies_fallback_budget(valid_sso_user_defined_values): + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( + valid_sso_user_defined_values, max_budget=litellm.max_ui_session_budget + ) + decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") + assert decrypted is not None + assert json.loads(decrypted).get("max_budget") == litellm.max_ui_session_budget + + +def test_get_cli_jwt_auth_token_no_fallback_when_budget_provided( + valid_sso_user_defined_values, +): + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( + valid_sso_user_defined_values, max_budget=None + ) + decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") + assert decrypted is not None + assert json.loads(decrypted).get("max_budget") is None + + @pytest.mark.asyncio async def test_default_internal_user_params_with_get_user_object(monkeypatch): """Test that default_internal_user_params is used when creating a new user via get_user_object"""