This commit is contained in:
Harshit1259 2026-08-26 14:30:20 -04:00 • committed by GitHub
commit 558096ccc2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 173 additions and 1 deletions

View file

@ -3,7 +3,7 @@ import asyncio
import json
import os
from collections import Counter
from collections.abc import Mapping
from collections.abc import Mapping, MutableMapping
from typing import (
Any,
Final,
@ -958,6 +958,34 @@ async def get_sso_settings():
return result
def _restore_masked_sso_secrets(
sso_data: MutableMapping[str, object], # mutable-ok: stored secrets are restored into the caller's config in place
before_sso_data: Mapping[str, object] | None,
) -> None:
"""Keep a stored SSO secret when the client round-trips its masked value.
Secret fields are masked before they are sent to the UI (see
get_sso_settings), so a client that edits only unrelated fields sends the
masked placeholder back unchanged. A plaintext secret never contains the
mask character, so an incoming secret that still carries the mask is treated
as "unchanged" and the previously effective secret is restored. This stops a
partial edit from overwriting e.g. the OAuth client_secret with
`abcd****wxyz` and breaking SSO login. An empty value is an intentional
clear and a genuinely new secret has no mask, so both pass through unchanged.
The effective secret is the stored row value, falling back to the process
environment -- the same precedence get_sso_settings uses when it masks the
field -- so a secret configured via an environment variable is preserved
too, not only database-stored ones. See #38177.
"""
for secret_field in SSO_SECRET_FIELDS:
db_secret = before_sso_data.get(secret_field) if before_sso_data else None
stored_secret = db_secret or os.environ.get(SSO_FIELD_ENV_VARS.get(secret_field, ""))
incoming_secret = sso_data.get(secret_field)
if stored_secret and isinstance(incoming_secret, str) and "*" in incoming_secret:
sso_data[secret_field] = stored_secret # rebind-ok: intentional in-place restore of a masked secret
@router.patch(
"/update/sso_settings",
tags=["SSO Settings"],
@ -1021,6 +1049,8 @@ async def update_sso_settings(
# Update environment variables in config and in memory
sso_data: Final = sso_config.model_dump()
_restore_masked_sso_secrets(sso_data, before_sso_data)
for field_name, value in sso_data.items():
if field_name in SSO_FIELD_ENV_VARS:
env_var_name = SSO_FIELD_ENV_VARS[field_name]

View file

@ -0,0 +1,142 @@
"""Regression tests for #38177: a partial SSO settings update must not overwrite
a stored client secret with the masked placeholder the UI sends back, while
still allowing an intentional clear and a genuinely new secret."""
import json
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi.testclient import TestClient
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
from litellm.proxy.config_resolvers.sso import SSO_SECRET_FIELDS
from litellm.proxy.proxy_server import app
client = TestClient(app)
REAL_SECRET = "real_generic_secret_ABCD1234"
@pytest.fixture
def mock_auth():
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
async def _override():
return {"user_id": "test_user"}
app.dependency_overrides[user_api_key_auth] = _override
yield
app.dependency_overrides.pop(user_api_key_auth, None)
def _masked(secret):
return mask_sensitive_keys({"generic_client_secret": secret}, set(SSO_SECRET_FIELDS))["generic_client_secret"]
def _mock_prisma(monkeypatch, existing_settings):
"""Wire a prisma client whose SSO row returns existing_settings (or None)."""
record = None
if existing_settings is not None:
record = MagicMock()
record.sso_settings = existing_settings
mock_prisma = MagicMock()
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=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, "get_config", AsyncMock(return_value={}))
monkeypatch.setattr(proxy_config, "_encrypt_env_variables", lambda environment_variables: environment_variables)
monkeypatch.setattr(proxy_config, "_decrypt_db_variables", lambda stored: stored)
return mock_prisma
def _stored_secret(mock_prisma):
return json.loads(mock_prisma.db.litellm_ssoconfig.upsert.call_args.kwargs["data"]["update"]["sso_settings"])
def test_database_stored_secret_is_preserved(mock_auth, monkeypatch):
"""A masked round-trip must keep a database-stored secret."""
monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key")
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.delenv("GENERIC_CLIENT_SECRET", raising=False)
mock_prisma = _mock_prisma(
monkeypatch,
{"generic_client_id": "cid", "generic_client_secret": REAL_SECRET, "proxy_base_url": "https://old.example.com"},
)
edited = {
"generic_client_id": "cid",
"generic_client_secret": _masked(REAL_SECRET),
"proxy_base_url": "https://new.example.com",
}
resp = client.patch("/update/sso_settings", json=edited)
assert resp.status_code == 200
stored = _stored_secret(mock_prisma)
assert stored["generic_client_secret"] == REAL_SECRET
assert stored["proxy_base_url"] == "https://new.example.com"
def test_environment_sourced_secret_is_preserved(mock_auth, monkeypatch):
"""A masked round-trip must keep a secret configured via env, even with no DB row."""
monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key")
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.setenv("GENERIC_CLIENT_SECRET", REAL_SECRET)
mock_prisma = _mock_prisma(monkeypatch, None) # nothing in the database
edited = {
"generic_client_id": "cid",
"generic_client_secret": _masked(REAL_SECRET),
"proxy_base_url": "https://new.example.com",
}
resp = client.patch("/update/sso_settings", json=edited)
assert resp.status_code == 200
stored = _stored_secret(mock_prisma)
assert stored["generic_client_secret"] == REAL_SECRET
def test_empty_secret_clears_intentionally(mock_auth, monkeypatch):
"""An empty incoming value is an intentional clear, not a masked round-trip."""
monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key")
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.delenv("GENERIC_CLIENT_SECRET", raising=False)
mock_prisma = _mock_prisma(
monkeypatch,
{"generic_client_id": "cid", "generic_client_secret": REAL_SECRET},
)
edited = {"generic_client_id": "cid", "generic_client_secret": "", "proxy_base_url": "https://x.example.com"}
resp = client.patch("/update/sso_settings", json=edited)
assert resp.status_code == 200
stored = _stored_secret(mock_prisma)
assert stored["generic_client_secret"] == ""
def test_new_secret_is_saved(mock_auth, monkeypatch):
"""A genuinely new secret (no mask character) replaces the stored one."""
monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key")
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
monkeypatch.delenv("GENERIC_CLIENT_SECRET", raising=False)
mock_prisma = _mock_prisma(
monkeypatch,
{"generic_client_id": "cid", "generic_client_secret": REAL_SECRET},
)
edited = {
"generic_client_id": "cid",
"generic_client_secret": "brand_new_secret_9999",
"proxy_base_url": "https://x",
}
resp = client.patch("/update/sso_settings", json=edited)
assert resp.status_code == 200
stored = _stored_secret(mock_prisma)
assert stored["generic_client_secret"] == "brand_new_secret_9999"