diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index f1c1d870d1..9c99b625e9 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -561,6 +561,41 @@ async def update_sso_settings(sso_config: SSOConfig): }, ) + # Remove SSO-related env vars from config.environment_variables + try: + env_var_entry = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "environment_variables"} + ) + + # If no environment_variables entry exists, nothing to clean up + if env_var_entry is not None: + if env_var_entry.param_value is not None: + if isinstance(env_var_entry.param_value, str): + environment_variables = json.loads(env_var_entry.param_value) + else: + environment_variables = dict(env_var_entry.param_value) + else: + environment_variables = {} + + env_vars_to_remove = set(env_var_mapping.values()) + filtered_env_vars = { + key: value + for key, value in environment_variables.items() + if key not in env_vars_to_remove + } + + await prisma_client.db.litellm_config.update( + where={"param_name": "environment_variables"}, + data={ + "param_value": json.dumps(filtered_env_vars, default=str), + }, + ) + except Exception as e: + raise HTTPException( + status_code=500, + detail={"error": f"Error updating environment_variables: {str(e)}"}, + ) + return { "message": "SSO settings updated successfully", "status": "success", diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index c15bdba800..d3c9915119 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -306,6 +306,9 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_config = MagicMock() + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_config.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) # Mock encryption to return values as-is @@ -380,6 +383,18 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_config = MagicMock() + + env_var_entry = MagicMock() + env_var_entry.param_value = json.dumps( + { + "GOOGLE_CLIENT_ID": "old_google_id", + "MICROSOFT_CLIENT_SECRET": "old_secret", + "PROXY_BASE_URL": "old_proxy_url", + } + ) + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=env_var_entry) + mock_prisma.db.litellm_config.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) # Mock encryption to return values as-is @@ -440,6 +455,17 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_config = MagicMock() + env_var_entry = MagicMock() + env_var_entry.param_value = json.dumps( + { + "GOOGLE_CLIENT_ID": "old_google_id", + "MICROSOFT_CLIENT_SECRET": "old_secret", + "PROXY_BASE_URL": "old_proxy_url", + } + ) + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=env_var_entry) + mock_prisma.db.litellm_config.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) # Mock encryption to return values as-is @@ -492,6 +518,17 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_config = MagicMock() + + env_var_entry = MagicMock() + env_var_entry.param_value = json.dumps( + { + "GOOGLE_CLIENT_ID": "test_existing_google_id", + "MICROSOFT_CLIENT_SECRET": "test_existing_microsoft_secret", + } + ) + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=env_var_entry) + mock_prisma.db.litellm_config.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) # Mock encryption to return values as-is @@ -551,6 +588,9 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_config = MagicMock() + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_config.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) # Mock encryption to return values as-is @@ -835,6 +875,9 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() upsert_mock = AsyncMock() mock_prisma.db.litellm_ssoconfig.upsert = upsert_mock + mock_prisma.db.litellm_config = MagicMock() + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_config.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) @@ -892,6 +935,99 @@ class TestProxySettingEndpoints: assert create_sso_settings["google_client_secret"] == "encrypted_new_google_secret" assert create_sso_settings["proxy_base_url"] == "encrypted_https://new.example.com" + def test_update_sso_settings_removes_sso_env_vars_from_config( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """Ensure SSO-related env vars are deleted from stored config""" + import json + from unittest.mock import AsyncMock, MagicMock + + monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key") + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_ssoconfig = MagicMock() + mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + + env_var_entry = MagicMock() + env_var_entry.param_value = json.dumps( + { + "GOOGLE_CLIENT_ID": "old_google_id", + "GENERIC_TOKEN_ENDPOINT": "old_endpoint", + "UNCHANGED_ENV": "keep_me", + } + ) + mock_prisma.db.litellm_config = MagicMock() + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=env_var_entry) + mock_prisma.db.litellm_config.update = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + from litellm.proxy.proxy_server import proxy_config + + monkeypatch.setattr( + proxy_config, + "_encrypt_env_variables", + lambda environment_variables: environment_variables, + ) + + response = client.patch( + "/update/sso_settings", json={"google_client_id": "new_google_id"} + ) + + assert response.status_code == 200 + mock_prisma.db.litellm_config.find_unique.assert_called_once() + mock_prisma.db.litellm_config.update.assert_called_once() + update_call = mock_prisma.db.litellm_config.update.call_args + updated_env_vars = json.loads(update_call.kwargs["data"]["param_value"]) + assert "GOOGLE_CLIENT_ID" not in updated_env_vars + assert "GENERIC_TOKEN_ENDPOINT" not in updated_env_vars + assert updated_env_vars["UNCHANGED_ENV"] == "keep_me" + + def test_update_sso_settings_preserves_non_sso_env_vars( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """Ensure env vars outside SSO mapping remain unchanged""" + import json + from unittest.mock import AsyncMock, MagicMock + + monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key") + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_ssoconfig = MagicMock() + mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + + env_var_entry = MagicMock() + env_var_entry.param_value = { + "UNRELATED_ENV": "keep_this", + "ANOTHER_ENV": "also_keep", + } + mock_prisma.db.litellm_config = MagicMock() + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=env_var_entry) + mock_prisma.db.litellm_config.update = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + from litellm.proxy.proxy_server import proxy_config + + monkeypatch.setattr( + proxy_config, + "_encrypt_env_variables", + lambda environment_variables: environment_variables, + ) + + response = client.patch( + "/update/sso_settings", json={"microsoft_client_id": "new_microsoft_id"} + ) + + assert response.status_code == 200 + mock_prisma.db.litellm_config.find_unique.assert_called_once() + mock_prisma.db.litellm_config.update.assert_called_once() + update_call = mock_prisma.db.litellm_config.update.call_args + updated_env_vars = json.loads(update_call.kwargs["data"]["param_value"]) + assert updated_env_vars == env_var_entry.param_value + def test_get_sso_settings_empty_database(self, mock_proxy_config, mock_auth, monkeypatch): """Test getting SSO settings when database table is empty""" from unittest.mock import AsyncMock, MagicMock