mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-14 16:25:29 +00:00
Merge pull request #17668 from BerriAI/litellm_sso_config_2
[Fix] Remove SSO Config Values from Config Table on SSO Update
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user