From 054ac7b3de60436587b7f3bf451ef38bca6f8797 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 27 Aug 2026 02:17:24 +0000 Subject: [PATCH] fix(proxy): restore deployment CyberArk env config on delete and roll back on persist failure Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../config_override_endpoints.py | 92 +++++++++++++------ litellm/proxy/proxy_server.py | 7 +- .../test_config_override_endpoints.py | 75 +++++++++++++++ 3 files changed, 140 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index d0151747493..072c733f0c2 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -258,15 +258,61 @@ def _clear_hashicorp_vault_state(proxy_config: "ProxyConfig") -> None: proxy_config._last_hashicorp_vault_config = None # pyright: ignore[reportPrivateUsage] # proxy-internal change-detection cache -def _clear_cyberark_state(proxy_config: "ProxyConfig") -> None: - """Clear all CyberArk state: env vars, secret manager, and change-detection cache.""" - _set_env_vars({}, CYBERARK_ENV_VAR_MAPPING) # mutable-ok: empty payload unsets all mapped env vars +def _snapshot_cyberark_boot_env(proxy_config: "ProxyConfig") -> None: + """Capture deployment-provided CYBERARK_* env vars once, before the first DB-driven overwrite.""" + if proxy_config._cyberark_boot_env is None: # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot + proxy_config._cyberark_boot_env = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot + + +def _restore_cyberark_runtime(proxy_config: "ProxyConfig", env_values: Mapping[str, str | None]) -> None: + """Restore CYBERARK_* env vars and reinitialize (or drop) the secret manager to match them.""" + _set_env_vars(env_values, CYBERARK_ENV_VAR_MAPPING) + if env_values.get("cyberark_api_base"): + try: + proxy_config.initialize_secret_manager(key_management_system="cyberark") + except Exception: # noqa: BLE001 # restore is best-effort; fall through to dropping the manager + verbose_proxy_logger.exception("Failed to restore previous CyberArk configuration") + else: + return if litellm._key_management_system == KeyManagementSystem.CYBERARK: # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage litellm.secret_manager_client = None litellm._key_management_system = None # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage + # Force the vault reload to re-init from its own row so no manager is stranded inactive + proxy_config._last_hashicorp_vault_config = None # pyright: ignore[reportPrivateUsage] # proxy-internal change-detection cache + + +def _clear_cyberark_state(proxy_config: "ProxyConfig") -> None: + """Drop DB-driven CyberArk state, restoring deployment-provided env vars if any.""" + boot_env: Final[Mapping[str, str | None]] = ( + proxy_config._cyberark_boot_env or {} # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot + ) + _restore_cyberark_runtime(proxy_config, boot_env) proxy_config._last_cyberark_config = None # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage +async def _persist_cyberark_config( + prisma_client: "PrismaClient", + proxy_config: "ProxyConfig", + config_data: Mapping[str, object], +) -> dict[str, object]: + """Encrypt and upsert the CyberArk config row; returns the stored (encrypted) payload.""" + encrypted_data: Final = proxy_config._encrypt_env_variables(dict(config_data)) # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage + config_value: Final = safe_dumps(encrypted_data) + await _config_overrides_table(prisma_client).upsert( + where={"config_type": "cyberark"}, # mutable-ok: prisma upsert payload + data={ # mutable-ok: prisma upsert payload + "create": { # mutable-ok: prisma upsert payload + "config_type": "cyberark", + "config_value": config_value, + }, + "update": { # mutable-ok: prisma upsert payload + "config_value": config_value, + }, + }, + ) + return safe_json_loads(config_value) + + # --- Hashicorp Vault endpoints --- @@ -647,8 +693,7 @@ async def update_cyberark_config( if field not in config_data and env_values.get(field): config_data[field] = env_values[field] - # Strip empty strings — they signal "clear this field" - config_data = {k: v for k, v in config_data.items() if v != ""} # mutable-ok: dict # rebind-ok: strip clears + config_data = {k: v for k, v in config_data.items() if v != ""} # mutable-ok: dict # rebind-ok: "" means clear has_api_base: Final = bool(config_data.get("cyberark_api_base")) has_api_key_auth: Final = bool(config_data.get("cyberark_api_key")) @@ -667,10 +712,8 @@ async def update_cyberark_config( "provide an API Key, or both Client Certificate and Client Key", ) - # Snapshot current env vars so we can restore on failure + _snapshot_cyberark_boot_env(proxy_config) previous_env: Final = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) - - # Set env vars and verify the secret manager can initialize before persisting _set_env_vars(config_data, CYBERARK_ENV_VAR_MAPPING) try: @@ -683,24 +726,17 @@ async def update_cyberark_config( detail=f"Failed to initialize secret manager: {e}", ) - # Only persist to DB after successful init - encrypted_data: Final = proxy_config._encrypt_env_variables(config_data) # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage - config_value: Final = safe_dumps(encrypted_data) - await _config_overrides_table(prisma_client).upsert( - where={"config_type": "cyberark"}, # mutable-ok: prisma upsert payload - data={ # mutable-ok: prisma upsert payload - "create": { # mutable-ok: prisma upsert payload - "config_type": "cyberark", - "config_value": config_value, - }, - "update": { # mutable-ok: prisma upsert payload - "config_value": config_value, - }, - }, - ) - - # Update change-detection cache so the background reload doesn't redundantly re-init - proxy_config._last_cyberark_config = safe_json_loads(config_value) # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage + try: + proxy_config._last_cyberark_config = await _persist_cyberark_config( # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage + prisma_client, proxy_config, config_data + ) + except Exception as e: # noqa: BLE001 # persistence failure must roll back the runtime state set above + _restore_cyberark_runtime(proxy_config, previous_env) + verbose_proxy_logger.exception("Error persisting CyberArk configuration: %s", str(e)) + raise HTTPException( + status_code=500, + detail=f"Failed to persist CyberArk configuration: {e}", + ) before_config: Final = existing_decrypted if existing_decrypted is not None else env_values action: Final[AUDIT_ACTIONS] = "updated" if existing_record is not None else "created" @@ -758,8 +794,6 @@ async def get_cyberark_config( if db_record is not None and db_record.config_value is not None: config_data: Final = _parse_config_value(db_record.config_value) - - # Decrypt then mask sensitive fields so plaintext secrets are never sent to the UI decrypted_data: Final[Mapping[str, object]] = proxy_config._decrypt_db_variables(config_data) # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage masked_data: Final = _mask_sensitive_fields(decrypted_data, CYBERARK_SENSITIVE_FIELDS) @@ -870,7 +904,6 @@ async def test_cyberark_connection( detail="CyberArk is not configured. Save a configuration first.", ) - # Step 1: Authenticate (exercises API-key or TLS cert login) try: headers: Final[Mapping[str, str]] = await asyncio.to_thread(client._get_request_headers) # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage except Exception as e: # noqa: BLE001 # surface any auth failure as a 502 with detail @@ -879,7 +912,6 @@ async def test_cyberark_connection( detail=f"CyberArk authentication failed: {e}", ) - # Step 2: Verify the token is valid via /whoami try: async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8a0439fb879..604fc874638 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4239,6 +4239,7 @@ class ProxyConfig: self._last_semantic_filter_config: dict[str, object] | None = None self._last_hashicorp_vault_config: dict[str, object] | None = None self._last_cyberark_config: dict[str, object] | None = None # mutable-ok: change-detection cache + self._cyberark_boot_env: dict[str, str | None] | None = None # mutable-ok: deployment env snapshot, set once self.worker_registry: list[WorkerRegistryEntry] = [] self.config_sync_subscriber: ConfigSyncSubscriber | None = None self.auth_cache_invalidation_subscriber: AuthCacheInvalidationSubscriber | None = None @@ -7096,6 +7097,7 @@ class ProxyConfig: _get_current_env_values, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module _parse_config_value, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module _set_env_vars, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module + _snapshot_cyberark_boot_env, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module ) try: @@ -7121,18 +7123,15 @@ class ProxyConfig: if self._last_cyberark_config == config_data: return - # Decrypt all fields and set env vars decrypted_data: Final = self._decrypt_db_variables(config_data) - # Snapshot current env vars so we can restore on failure + _snapshot_cyberark_boot_env(self) previous_env: Final = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) _set_env_vars(decrypted_data, CYBERARK_ENV_VAR_MAPPING) - # Reinitialize the secret manager try: self.initialize_secret_manager(key_management_system="cyberark") except Exception: - # Restore previous working env vars instead of wiping all _set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) raise diff --git a/tests/test_litellm/proxy/management_endpoints/test_config_override_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_config_override_endpoints.py index a4dfa686aaf..66b0a6555e0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_config_override_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_config_override_endpoints.py @@ -45,6 +45,7 @@ def _make_mock_proxy_config(): cfg = MagicMock() cfg.initialize_secret_manager = MagicMock() cfg._last_hashicorp_vault_config = None + cfg._cyberark_boot_env = None cfg._encrypt_env_variables = MagicMock( side_effect=lambda d: {k: f"enc_{v}" for k, v in d.items()} ) @@ -396,6 +397,7 @@ async def test_cyberark_validation_errors_and_access_control(client, monkeypatch mock_prisma, mock_db = _make_mock_db() mock_cfg = MagicMock() mock_cfg._last_cyberark_config = {"cyberark_api_base": "old"} + mock_cfg._cyberark_boot_env = None monkeypatch.setattr(ps, "prisma_client", mock_prisma) monkeypatch.setattr(ps, "proxy_config", mock_cfg) old_client, old_kms = litellm.secret_manager_client, litellm._key_management_system @@ -454,6 +456,79 @@ async def test_cyberark_validation_errors_and_access_control(client, monkeypatch _cleanup() +@pytest.mark.asyncio +async def test_cyberark_delete_restores_deployment_env_config(client, monkeypatch): + """Deleting the DB override must restore env vars the deployment started with, + and reinitialize the manager from them, instead of wiping CyberArk entirely.""" + mock_prisma, mock_db = _make_mock_db() + mock_cfg = _make_mock_proxy_config() + mock_cfg._last_cyberark_config = None + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "proxy_config", mock_cfg) + old_client, old_kms = litellm.secret_manager_client, litellm._key_management_system + _set_admin() + + try: + monkeypatch.setenv("CYBERARK_API_BASE", "https://conjur.boot.com") + monkeypatch.setenv("CYBERARK_API_KEY", "boot-key") + + r = client.post( + CYBERARK_URL, + json={"cyberark_api_base": "https://conjur.db.com", "cyberark_api_key": "db-key"}, + ) + assert r.status_code == 200 + assert os.environ["CYBERARK_API_BASE"] == "https://conjur.db.com" + + mock_cfg.initialize_secret_manager.reset_mock() + r = client.delete(CYBERARK_URL) + assert r.status_code == 200 + assert os.environ["CYBERARK_API_BASE"] == "https://conjur.boot.com" + assert os.environ["CYBERARK_API_KEY"] == "boot-key" + mock_cfg.initialize_secret_manager.assert_called_with(key_management_system="cyberark") + assert mock_cfg._last_cyberark_config is None + finally: + litellm.secret_manager_client = old_client # test-quality-ok: endpoint hot-reloads litellm globals; test must set and restore them + litellm._key_management_system = old_kms # test-quality-ok: endpoint hot-reloads litellm globals; test must set and restore them + _cleanup() + + +@pytest.mark.asyncio +async def test_cyberark_persist_failure_rolls_back_runtime_state(client, monkeypatch): + """If the DB upsert fails after the manager was reinitialized, the endpoint + must restore the previous env vars and reinitialize from them, so this pod + does not keep serving credentials that were never committed to the DB.""" + mock_prisma, mock_db = _make_mock_db() + mock_cfg = _make_mock_proxy_config() + mock_cfg._last_cyberark_config = None + mock_db.upsert = AsyncMock(side_effect=Exception("db write failed")) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "proxy_config", mock_cfg) + old_client, old_kms = litellm.secret_manager_client, litellm._key_management_system + _set_admin() + + try: + monkeypatch.setenv("CYBERARK_API_BASE", "https://conjur.prev.com") + monkeypatch.setenv("CYBERARK_API_KEY", "prev-key") + + r = client.post( + CYBERARK_URL, + json={"cyberark_api_base": "https://conjur.new.com", "cyberark_api_key": "new-key"}, + ) + assert r.status_code == 500 + assert "persist" in r.json()["detail"].lower() + assert os.environ["CYBERARK_API_BASE"] == "https://conjur.prev.com" + assert os.environ["CYBERARK_API_KEY"] == "prev-key" + # last call must be the rollback reinit against the restored env + assert ( + mock_cfg.initialize_secret_manager.call_args_list[-1].kwargs["key_management_system"] == "cyberark" + ) + assert os.environ.get("CYBERARK_API_BASE") != "https://conjur.new.com" + finally: + litellm.secret_manager_client = old_client # test-quality-ok: endpoint hot-reloads litellm globals; test must set and restore them + litellm._key_management_system = old_kms # test-quality-ok: endpoint hot-reloads litellm globals; test must set and restore them + _cleanup() + + @pytest.mark.asyncio async def test_cyberark_audit_log_redacts_values(client, monkeypatch): monkeypatch.setattr(litellm, "store_audit_logs", True)