From b964b1b54354ed0dff0cbd1eb6ad55e8e1d7b786 Mon Sep 17 00:00:00 2001 From: Ganni Galea Curmi Date: Tue, 26 May 2026 21:14:26 -0400 Subject: [PATCH] fix(sso): preserve omitted settings on partial update --- .../proxy_setting_endpoints.py | 44 +++++++++--- .../test_proxy_setting_endpoints.py | 72 +++++++++++++++++++ 2 files changed, 108 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 07e2ca71950..430998de695 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -31,6 +31,16 @@ _SSO_SENSITIVE_FIELDS: Set[str] = { } +def _get_sso_settings_dict_from_db_record( + sso_db_record: Optional[Any], +) -> Dict[str, Any]: + if sso_db_record is None or not sso_db_record.sso_settings: + return {} + if isinstance(sso_db_record.sso_settings, str): + return json.loads(sso_db_record.sso_settings) + return dict(sso_db_record.sso_settings) + + class IPAddress(BaseModel): ip: str @@ -669,12 +679,8 @@ async def get_sso_settings(): where={"id": "sso_config"} ) - # Initialize with defaults - sso_settings_dict = {} - - if sso_db_record and sso_db_record.sso_settings: - # Load settings from database - sso_settings_dict = dict(sso_db_record.sso_settings) + # Load settings from database + sso_settings_dict = _get_sso_settings_dict_from_db_record(sso_db_record) role_mappings_data = sso_settings_dict.pop("role_mappings", None) role_mappings = None @@ -820,8 +826,30 @@ async def update_sso_settings(sso_config: SSOConfig): if "general_settings" not in config: config["general_settings"] = {} - # Update environment variables in config and in memory - sso_data = sso_config.model_dump() + existing_sso_db_record = await prisma_client.db.litellm_ssoconfig.find_unique( + where={"id": "sso_config"} + ) + existing_sso_settings = _get_sso_settings_dict_from_db_record( + existing_sso_db_record + ) + existing_structured_settings = {} + for structured_field in ("role_mappings", "team_mappings"): + if structured_field in existing_sso_settings: + existing_structured_settings[structured_field] = existing_sso_settings.pop( + structured_field + ) + + existing_sso_data = proxy_config._decrypt_and_set_db_env_variables( + environment_variables=existing_sso_settings + ) + existing_sso_data.update(existing_structured_settings) + + # Update environment variables in config and in memory. PATCH semantics: + # omitted fields keep their current stored value, explicit null/empty clears. + sso_data = { + **existing_sso_data, + **sso_config.model_dump(exclude_unset=True), + } for field_name, value in sso_data.items(): if field_name in env_var_mapping: env_var_name = env_var_mapping[field_name] 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 ae217aca16e..c9e69424f71 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 @@ -397,6 +397,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config.update = AsyncMock() @@ -466,6 +467,69 @@ class TestProxySettingEndpoints: create_sso_settings = json.loads(create_data["sso_settings"]) assert create_sso_settings["google_client_id"] == "new_google_client_id" + def test_update_sso_settings_partial_update_preserves_existing_secrets( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """Test that omitted SSO fields keep their stored values during PATCH updates.""" + 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) + monkeypatch.setenv("GOOGLE_CLIENT_SECRET", "existing_google_secret") + + mock_prisma = MagicMock() + existing_sso_record = MagicMock() + existing_sso_record.sso_settings = { + "google_client_id": "existing_google_client_id", + "google_client_secret": "existing_google_secret", + "microsoft_client_secret": "existing_microsoft_secret", + } + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock( + return_value=existing_sso_record + ) + 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) + + from litellm.proxy.proxy_server import proxy_config + + monkeypatch.setattr( + proxy_config, + "_decrypt_and_set_db_env_variables", + lambda environment_variables: dict(environment_variables), + ) + monkeypatch.setattr( + proxy_config, + "_encrypt_env_variables", + lambda environment_variables: environment_variables, + ) + + response = client.patch( + "/update/sso_settings", json={"ui_access_mode": "admin_only"} + ) + + assert response.status_code == 200 + settings = response.json()["settings"] + assert settings["ui_access_mode"] == "admin_only" + assert settings["google_client_secret"] == "existing_google_secret" + assert settings["microsoft_client_secret"] == "existing_microsoft_secret" + assert os.environ["GOOGLE_CLIENT_SECRET"] == "existing_google_secret" + + call_args = mock_prisma.db.litellm_ssoconfig.upsert.call_args + stored_sso_settings = json.loads( + call_args.kwargs["data"]["update"]["sso_settings"] + ) + assert stored_sso_settings["ui_access_mode"] == "admin_only" + assert stored_sso_settings["google_client_id"] == "existing_google_client_id" + assert stored_sso_settings["google_client_secret"] == "existing_google_secret" + assert ( + stored_sso_settings["microsoft_client_secret"] + == "existing_microsoft_secret" + ) + def test_update_sso_settings_with_null_values_clears_env_vars( self, mock_proxy_config, mock_auth, monkeypatch ): @@ -479,6 +543,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config = MagicMock() env_var_entry = MagicMock() @@ -558,6 +623,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config = MagicMock() env_var_entry = MagicMock() env_var_entry.param_value = json.dumps( @@ -628,6 +694,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config = MagicMock() env_var_entry = MagicMock() @@ -705,6 +772,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config.update = AsyncMock() @@ -1351,6 +1419,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() upsert_mock = AsyncMock() mock_prisma.db.litellm_ssoconfig.upsert = upsert_mock + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config.update = AsyncMock() @@ -1430,6 +1499,7 @@ class TestProxySettingEndpoints: mock_prisma.db = MagicMock() mock_prisma.db.litellm_ssoconfig = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) env_var_entry = MagicMock() env_var_entry.param_value = json.dumps( @@ -1481,6 +1551,7 @@ class TestProxySettingEndpoints: mock_prisma.db = MagicMock() mock_prisma.db.litellm_ssoconfig = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) env_var_entry = MagicMock() env_var_entry.param_value = { @@ -1652,6 +1723,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config.update = AsyncMock()