fix(sso): preserve omitted settings on partial update

This commit is contained in:
Ganni Galea Curmi 2026-05-26 21:14:26 -04:00
parent 73e9071311
commit b964b1b543
2 changed files with 108 additions and 8 deletions

View file

@ -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]

View file

@ -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()