mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(sso): preserve omitted settings on partial update
This commit is contained in:
parent
73e9071311
commit
b964b1b543
2 changed files with 108 additions and 8 deletions
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue