diff --git a/litellm-rust/crates/python-bridge/src/secrets/mod.rs b/litellm-rust/crates/python-bridge/src/secrets/mod.rs index edf6aead53a..cddd47d08db 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/mod.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/mod.rs @@ -21,10 +21,10 @@ use crate::{coercion::FieldSpec, python_settings::PythonSettings}; const NATIVE: FieldSpec = FieldSpec::new("native", |field| field.schema_bool()); -/// Where a Rust route reads provider secrets from. Python's `get_secret_str` until a -/// `SecretManagerRule` in `catalog.py` moves the configured system off `PYTHON_ONLY`, then the -/// native secret manager. A bare extension module without the litellm package reads the process -/// environment. +/// Where a Rust route reads provider secrets from. Python's `get_secret_str` when the settings +/// projection reports no configured secret manager, otherwise the resolved manager: the native +/// backend when the configured client captures one, the Python callback when it does not. A bare +/// extension module without the litellm package reads the process environment. pub(crate) fn source(py: Python<'_>) -> PyResult> { let Some(snapshot) = PythonSettings::SecretManager.read_or_unset(py)? else { return Ok(Arc::new(EnvironmentSecrets::python_compatible( diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 40c6456431d..73e95c24de4 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -1,4 +1,4 @@ -"""Ordered rollout policy for routes, loggers, and secret managers. +"""Ordered rollout policy for routes and loggers. The first matching rule wins; unmatched contexts stay on Python. Native admission separately decides whether the selected implementation can execute. @@ -12,7 +12,6 @@ from typing import Final, TypeAlias from litellm.rust_bridge.configuration import Decision, Rollout from litellm.rust_bridge.configuration import decision as _decision -from litellm.types.secret_managers.main import KeyManagementSystem class Route(str, Enum): @@ -49,20 +48,6 @@ class RouteRule: ) -@dataclass(frozen=True, slots=True) -class SecretManagerContext: - system: str - - -@dataclass(frozen=True, slots=True) -class SecretManagerRule: - rollout: Rollout - systems: frozenset[str] | None = None - - def matches(self, context: Context) -> bool: - return isinstance(context, SecretManagerContext) and (self.systems is None or context.system in self.systems) - - @dataclass(frozen=True, slots=True) class LoggerContext: pass @@ -76,8 +61,8 @@ class LoggerRule: return isinstance(context, LoggerContext) -Context: TypeAlias = RouteContext | SecretManagerContext | LoggerContext -Rule: TypeAlias = RouteRule | SecretManagerRule | LoggerRule +Context: TypeAlias = RouteContext | LoggerContext +Rule: TypeAlias = RouteRule | LoggerRule Rules: TypeAlias = tuple[Rule, ...] RULES: Final[Rules] = ( @@ -91,15 +76,6 @@ RULES: Final[Rules] = ( RouteRule(Route.TOKEN_COUNTER, Rollout.PYTHON_ONLY), RouteRule(Route.TOKENIZER, Rollout.PYTHON_ONLY), RouteRule(Route.TRANSCRIPTION, Rollout.RUST_REQUIRED, providers=frozenset({"bedrock"})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.GOOGLE_KMS.value})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.AZURE_KEY_VAULT.value})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.AWS_SECRET_MANAGER.value})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.GOOGLE_SECRET_MANAGER.value})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.HASHICORP_VAULT.value})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.CYBERARK.value})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.LOCAL.value})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.AWS_KMS.value})), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.CUSTOM.value})), ) diff --git a/litellm/rust_bridge/secret_manager.py b/litellm/rust_bridge/secret_manager.py index ca7dc2e6434..548f65cfc26 100644 --- a/litellm/rust_bridge/secret_manager.py +++ b/litellm/rust_bridge/secret_manager.py @@ -1,17 +1,11 @@ from __future__ import annotations import os -from collections.abc import Awaitable, Mapping +from collections.abc import Mapping from dataclasses import dataclass, field from importlib import import_module -from typing import Final, Protocol, runtime_checkable +from typing import Final -import httpx -from pydantic import JsonValue - -from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import Rules, SecretManagerContext, decision -from litellm.rust_bridge.configuration import Decision from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem @@ -163,123 +157,3 @@ def native_secret_manager_config(client: object) -> NativeSecretManagerConfig | ): return _capture(client, adapter) return None - - -class NativeSecretManagerRuntime(Protocol): - @property - def system(self) -> str: ... - - def read_secret(self, name: str, settings: Mapping[str, object] | None = None) -> JsonValue: ... - - -@runtime_checkable -class NativeSecretManagerFactory(Protocol): - @staticmethod - def from_client(client: object) -> NativeSecretManagerRuntime | None: ... - - -def _factory(value: object) -> NativeSecretManagerFactory | None: - return value if isinstance(value, NativeSecretManagerFactory) and callable(value.from_client) else None - - -NATIVE_SECRET_MANAGER: Final = NativeBinding("_SecretManagerRuntime", validate=_factory) - - -def resolve_native_secret_manager( - client: object, - system: str, - rules: Rules | None = None, - *, - binding: NativeBinding[NativeSecretManagerFactory] = NATIVE_SECRET_MANAGER, -) -> NativeSecretManagerRuntime | None: - if system in ("custom", "local"): - return None - selected: Final = decision(SecretManagerContext(system=system), rules) - if selected is Decision.PYTHON: - return None - factory: Final = binding.load() - if factory is None: - if selected is Decision.RUST_REQUIRED: - raise RuntimeError("Rust secret manager runtime is unavailable") - return None - runtime: Final = factory.from_client(client) - if runtime is not None and runtime.system != system: - raise ValueError("Native secret manager system does not match configuration") - return runtime - - -@runtime_checkable -class NativeProviderReader(Protocol): - def sync_read_secret( - self, - secret_name: str, - optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, - ) -> str | None: ... - - def async_read_secret( - self, - secret_name: str, - optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, - ) -> Awaitable[str | None]: ... - - -def resolve_native_provider_reader( - client: object, - system: str, - rules: Rules | None = None, - *, - binding: NativeBinding[NativeSecretManagerFactory] = NATIVE_SECRET_MANAGER, -) -> NativeProviderReader | None: - runtime: Final = resolve_native_secret_manager(client, system, rules, binding=binding) - if runtime is None: - return None - if not isinstance(runtime, NativeProviderReader): - raise TypeError("Rust secret manager provider reads are unavailable") - return runtime - - -@runtime_checkable -class NativeProviderWriter(Protocol): - def async_write_secret( - self, - secret_name: str, - secret_value: str, - description: str | None = None, - optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, - tags: object = None, - ) -> Awaitable[dict[str, JsonValue]]: ... - - def async_delete_secret( - self, - secret_name: str, - recovery_window_in_days: int | None = 7, - optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, - ) -> Awaitable[dict[str, JsonValue]]: ... - - def async_rotate_secret( - self, - current_secret_name: str, - new_secret_name: str, - new_secret_value: str, - optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, - ) -> Awaitable[dict[str, JsonValue]]: ... - - -def resolve_native_provider_writer( - client: object, - system: str, - rules: Rules | None = None, - *, - binding: NativeBinding[NativeSecretManagerFactory] = NATIVE_SECRET_MANAGER, -) -> NativeProviderWriter | None: - runtime: Final = resolve_native_secret_manager(client, system, rules, binding=binding) - if runtime is None: - return None - if not isinstance(runtime, NativeProviderWriter): - raise TypeError("Rust secret manager provider writes are unavailable") - return runtime diff --git a/litellm/rust_bridge/settings.py b/litellm/rust_bridge/settings.py index 7861f50574f..60a58f1aaba 100644 --- a/litellm/rust_bridge/settings.py +++ b/litellm/rust_bridge/settings.py @@ -1,10 +1,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Final - -if TYPE_CHECKING: - from litellm.rust_bridge.catalog import Rules +from typing import Final @dataclass(frozen=True, slots=True) @@ -62,10 +59,8 @@ class SecretManagerBinding: settings_object: object -def secret_manager(rules: Rules | None = None) -> SecretManager: +def secret_manager() -> SecretManager: import litellm - from litellm.rust_bridge.catalog import SecretManagerContext, decision - from litellm.rust_bridge.configuration import Decision from litellm.secret_managers.main import ( _should_read_secret_from_secret_manager, # pyright: ignore[reportPrivateUsage] # canonical resolver is private ) @@ -74,12 +69,7 @@ def secret_manager(rules: Rules | None = None) -> SecretManager: system: Final = ( litellm._key_management_system # pyright: ignore[reportPrivateUsage] # canonical key management globals are private ) - native: Final = ( - readable - and system is not None - and decision(SecretManagerContext(system=system.value), rules) is not Decision.PYTHON - ) - return SecretManager(readable=readable, native=native) + return SecretManager(readable=readable, native=readable and system is not None) def secret_manager_binding() -> SecretManagerBinding: diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index 8b78980eb73..1bd4a4fa566 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -30,7 +30,6 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, ) -from litellm.rust_bridge.secret_manager import resolve_native_provider_reader from litellm.secret_managers.main import get_secret_str from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem @@ -140,10 +139,6 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): secret_name=secret_name, primary_secret_name=primary_secret_name ) - native: Final = resolve_native_provider_reader(self, "aws_secret_manager") - if native is not None: - return await native.async_read_secret(secret_name, optional_params, timeout) - endpoint_url, headers, body = self._prepare_request( action="GetSecretValue", secret_name=secret_name, @@ -196,10 +191,6 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): secret_name=secret_name, primary_secret_name=primary_secret_name ) - native: Final = resolve_native_provider_reader(self, "aws_secret_manager") - if native is not None: - return native.sync_read_secret(secret_name, optional_params, timeout) - endpoint_url, headers, body = self._prepare_request( action="GetSecretValue", secret_name=secret_name, diff --git a/litellm/secret_managers/cyberark_secret_manager.py b/litellm/secret_managers/cyberark_secret_manager.py index b869db51c85..2c827e36979 100644 --- a/litellm/secret_managers/cyberark_secret_manager.py +++ b/litellm/secret_managers/cyberark_secret_manager.py @@ -16,7 +16,6 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.rust_bridge.secret_manager import resolve_native_provider_reader, resolve_native_provider_writer from litellm.types.secret_managers.main import KeyManagementSystem from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name @@ -199,10 +198,6 @@ class CyberArkSecretManager(BaseSecretManager): Returns: Optional[str]: The secret value if found, None otherwise """ - native: Final = resolve_native_provider_reader(self, "cyberark") - if native is not None: - return await native.async_read_secret(secret_name, optional_params, timeout) - # Check cache first if self.cache.get_cache(secret_name) is not None: return self.cache.get_cache(secret_name) @@ -249,10 +244,6 @@ class CyberArkSecretManager(BaseSecretManager): Returns: Optional[str]: The secret value if found, None otherwise """ - native: Final = resolve_native_provider_reader(self, "cyberark") - if native is not None: - return native.sync_read_secret(secret_name, optional_params, timeout) - # Check cache first if self.cache.get_cache(secret_name) is not None: return self.cache.get_cache(secret_name) @@ -302,12 +293,6 @@ class CyberArkSecretManager(BaseSecretManager): Returns: dict: Response containing status and details of the operation """ - native: Final = resolve_native_provider_writer(self, "cyberark") - if native is not None: - return await native.async_write_secret( - secret_name, secret_value, description, optional_params, timeout, tags - ) - async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, params={"ssl_verify": self.ssl_verify}, @@ -353,10 +338,6 @@ class CyberArkSecretManager(BaseSecretManager): Returns: dict: Response indicating operation not supported """ - native: Final = resolve_native_provider_writer(self, "cyberark") - if native is not None: - return await native.async_delete_secret(secret_name, recovery_window_in_days, optional_params, timeout) - verbose_logger.warning( "CyberArk Conjur does not support direct secret deletion. Secrets must be removed through policy updates." ) @@ -377,15 +358,6 @@ class CyberArkSecretManager(BaseSecretManager): optional_params: dict | None = None, timeout: float | httpx.Timeout | None = None, ) -> dict: - native: Final = resolve_native_provider_writer(self, "cyberark") - if native is not None: - return await native.async_rotate_secret( - current_secret_name, - new_secret_name, - new_secret_value, - optional_params, - timeout, - ) return await super().async_rotate_secret( current_secret_name, new_secret_name, diff --git a/litellm/secret_managers/dispatch.py b/litellm/secret_managers/dispatch.py deleted file mode 100644 index 4912771eb3c..00000000000 --- a/litellm/secret_managers/dispatch.py +++ /dev/null @@ -1,30 +0,0 @@ -from typing import Final - -from pydantic import JsonValue - -from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import Rules -from litellm.rust_bridge.secret_manager import ( - NATIVE_SECRET_MANAGER, - NativeSecretManagerFactory, - resolve_native_secret_manager, -) -from litellm.secret_managers.secret_manager_handler import get_secret_from_manager as python_get_secret_from_manager -from litellm.types.secret_managers.main import KeyManagementSettings - - -def get_secret_from_manager( - client: object, - key_manager: str, - secret_name: str, - key_management_settings: KeyManagementSettings | None = None, - *, - rules: Rules | None = None, - binding: NativeBinding[NativeSecretManagerFactory] = NATIVE_SECRET_MANAGER, -) -> JsonValue: - native: Final = resolve_native_secret_manager(client, key_manager, rules, binding=binding) - if native is None: - return python_get_secret_from_manager(client, key_manager, secret_name, key_management_settings) - return native.read_secret( - secret_name, key_management_settings.model_dump(mode="json") if key_management_settings is not None else None - ) diff --git a/litellm/secret_managers/google_secret_manager.py b/litellm/secret_managers/google_secret_manager.py index f8de9b25e96..a29e5ab96e1 100644 --- a/litellm/secret_managers/google_secret_manager.py +++ b/litellm/secret_managers/google_secret_manager.py @@ -9,7 +9,6 @@ from litellm.constants import SECRET_MANAGER_REFRESH_INTERVAL from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase from litellm.llms.custom_httpx.http_handler import _get_httpx_client from litellm.proxy._types import CommonProxyErrors -from litellm.rust_bridge.secret_manager import resolve_native_provider_reader from litellm.types.secret_managers.main import KeyManagementSystem @@ -62,10 +61,6 @@ class GoogleSecretManager(GCSBucketBase): Returns: str: The secret value if successful, None otherwise. """ - native: Final = resolve_native_provider_reader(self, "google_secret_manager") - if native is not None: - return native.sync_read_secret(secret_name) - if self.always_read_secret_manager is not True: cached_secret: Final = self.cache.get_cache(secret_name) if cached_secret is not None: diff --git a/litellm/secret_managers/hashicorp_secret_manager.py b/litellm/secret_managers/hashicorp_secret_manager.py index d130481f8fc..5a05ba8f965 100644 --- a/litellm/secret_managers/hashicorp_secret_manager.py +++ b/litellm/secret_managers/hashicorp_secret_manager.py @@ -15,7 +15,6 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.rust_bridge.secret_manager import resolve_native_provider_reader, resolve_native_provider_writer from litellm.types.secret_managers.main import KeyManagementSystem from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name @@ -406,10 +405,6 @@ class HashicorpSecretManager(BaseSecretManager): secret_name is just the path inside the KV mount (e.g., 'myapp/config'). Returns the entire data dict from data.data, or None on failure. """ - native: Final = resolve_native_provider_reader(self, "hashicorp_vault") - if native is not None: - return await native.async_read_secret(secret_name, optional_params, timeout) - async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, ) @@ -441,10 +436,6 @@ class HashicorpSecretManager(BaseSecretManager): secret_name is just the path inside the KV mount (e.g., 'myapp/config'). Returns the entire data dict from data.data, or None on failure. """ - native: Final = resolve_native_provider_reader(self, "hashicorp_vault") - if native is not None: - return native.sync_read_secret(secret_name, optional_params, timeout) - sync_client: Final = _get_httpx_client() try: target: Final = self._build_secret_target(secret_name, optional_params) @@ -485,12 +476,6 @@ class HashicorpSecretManager(BaseSecretManager): Returns: dict: Response containing status and details of the operation """ - native: Final = resolve_native_provider_writer(self, "hashicorp_vault") - if native is not None: - return await native.async_write_secret( - secret_name, secret_value, description, optional_params, timeout, tags - ) - async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, params={"timeout": timeout}, @@ -540,12 +525,6 @@ class HashicorpSecretManager(BaseSecretManager): On success, returns the response from async_write_secret. On error, returns {"status": "error", "message": "error message"} """ - native: Final = resolve_native_provider_writer(self, "hashicorp_vault") - if native is not None: - return await native.async_rotate_secret( - current_secret_name, new_secret_name, new_secret_value, optional_params, timeout - ) - async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, params={"timeout": timeout}, @@ -692,10 +671,6 @@ class HashicorpSecretManager(BaseSecretManager): Returns: dict: Response containing status and details of the operation """ - native: Final = resolve_native_provider_writer(self, "hashicorp_vault") - if native is not None: - return await native.async_delete_secret(secret_name, recovery_window_in_days, optional_params, timeout) - async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, params={"timeout": timeout}, diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index a8420d3cbc6..716407d6def 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -12,10 +12,10 @@ import litellm from litellm._logging import verbose_logger from litellm.caching.caching import DualCache from litellm.llms.custom_httpx.http_handler import HTTPHandler -from litellm.secret_managers.dispatch import get_secret_from_manager from litellm.secret_managers.get_azure_ad_token_provider import ( get_azure_ad_token_provider, ) +from litellm.secret_managers.secret_manager_handler import get_secret_from_manager from litellm.types.llms.base import LiteLLMBaseModel oidc_cache: Final = DualCache() diff --git a/tests/unit/rust_bridge/test_catalog.py b/tests/unit/rust_bridge/test_catalog.py index 95e98fe98da..4242294a33a 100644 --- a/tests/unit/rust_bridge/test_catalog.py +++ b/tests/unit/rust_bridge/test_catalog.py @@ -9,15 +9,13 @@ from litellm.rust_bridge import catalog, configuration from litellm.rust_bridge.catalog import ( Context, LoggerContext, + LoggerRule, Route, RouteContext, RouteRule, Rules, - SecretManagerContext, - SecretManagerRule, ) from litellm.rust_bridge.configuration import Decision, Rollout -from litellm.types.secret_managers.main import KeyManagementSystem @pytest.fixture(autouse=True) @@ -65,23 +63,6 @@ def test_missing_rule_stays_on_python_even_when_rust_is_enabled(monkeypatch: pyt assert catalog.decision(RouteContext(route), rules=()) is Decision.PYTHON -@pytest.mark.parametrize( - "context", - ( - *(SecretManagerContext(system.value) for system in KeyManagementSystem), - SecretManagerContext("unknown"), - ), -) -def test_backend_rollouts_stay_on_python_when_global_rust_is_enabled( - monkeypatch: pytest.MonkeyPatch, context: Context -) -> None: - configuration.rust(True) - monkeypatch.setenv("LITELLM_RUST", "1") - - assert catalog.rollout(context) is Rollout.PYTHON_ONLY - assert catalog.decision(context) is Decision.PYTHON - - def test_logger_rollout_obeys_the_global_switch() -> None: assert catalog.rollout(LoggerContext()) is Rollout.RUST_OPT_IN assert catalog.decision(LoggerContext()) is Decision.PYTHON @@ -134,14 +115,12 @@ def test_ocr_has_no_python_path_to_opt_out_to( (RouteContext(Route.OCR, provider="local"), Decision.RUST_REQUIRED), (RouteContext(Route.OCR, provider="other"), Decision.PYTHON), (RouteContext(Route.MESSAGES, provider="local"), Decision.PYTHON), - (SecretManagerContext("local"), Decision.PYTHON), - (SecretManagerContext("other"), Decision.RUST_REQUIRED), + (LoggerContext(), Decision.RUST_REQUIRED), ), ) def test_mixed_rules_select_only_the_matching_domain(context: Context, expected: Decision) -> None: rules: Final[Rules] = ( - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), - SecretManagerRule(Rollout.RUST_REQUIRED), + LoggerRule(Rollout.RUST_REQUIRED), RouteRule(Route.OCR, Rollout.RUST_REQUIRED, providers=frozenset({"local"})), RouteRule(Route.OCR, Rollout.PYTHON_ONLY), ) @@ -149,7 +128,7 @@ def test_mixed_rules_select_only_the_matching_domain(context: Context, expected: assert catalog.decision(context, rules) is expected -@pytest.mark.parametrize("context", (RouteContext(Route.OCR), SecretManagerContext("local"))) +@pytest.mark.parametrize("context", (RouteContext(Route.OCR), LoggerContext())) @pytest.mark.parametrize( ("rollout", "process", "environment", "expected"), ( @@ -176,20 +155,19 @@ def test_all_domains_share_rollout_switches_and_first_match( monkeypatch.setenv("LITELLM_RUST", environment) rules: Final[Rules] = ( RouteRule(Route.OCR, rollout), - SecretManagerRule(rollout), + LoggerRule(rollout), RouteRule(Route.OCR, Rollout.RUST_REQUIRED), - SecretManagerRule(Rollout.RUST_REQUIRED), ) assert catalog.decision(context, rules) is expected assert catalog.decision(context, ()) is Decision.PYTHON -@pytest.mark.parametrize("context", (RouteContext(Route.OCR), SecretManagerContext("local"))) +@pytest.mark.parametrize("context", (RouteContext(Route.OCR),)) def test_empty_constraints_match_nothing(context: Context) -> None: rules: Final[Rules] = ( RouteRule(Route.OCR, Rollout.RUST_REQUIRED, providers=frozenset()), - SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset()), + RouteRule(Route.OCR, Rollout.RUST_REQUIRED, models=frozenset()), ) assert catalog.decision(context, rules) is Decision.PYTHON diff --git a/tests/unit/rust_bridge/test_dispatch.py b/tests/unit/rust_bridge/test_dispatch.py index 974e943ea9f..9d53db26ba8 100644 --- a/tests/unit/rust_bridge/test_dispatch.py +++ b/tests/unit/rust_bridge/test_dispatch.py @@ -6,7 +6,7 @@ import pytest from litellm.rust_bridge import configuration from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule, Rules, SecretManagerRule +from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule, Rules from litellm.rust_bridge.configuration import Rollout from litellm.rust_bridge.dispatch import PublicDispatch from litellm.rust_bridge.runtime import NO_PYTHON, NoPythonImplementationError @@ -23,7 +23,7 @@ def binding() -> NativeBinding[object]: return bound -@pytest.mark.parametrize("rules", ((), (SecretManagerRule(Rollout.RUST_REQUIRED),))) +@pytest.mark.parametrize("rules", ((), (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),))) def test_route_without_rules_forwards_before_request_projection(rules: Rules) -> None: stream: Final[Iterator[int]] = iter((1, 2)) @@ -97,7 +97,7 @@ def test_native_stream_result_is_not_consumed_or_wrapped() -> None: request: Final = Request(model="streaming-model") stream: Final[Iterator[int]] = iter((1, 2)) rules: Final[Rules] = ( - SecretManagerRule(Rollout.PYTHON_ONLY), + RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY), RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED), ) dispatch: Final = PublicDispatch( @@ -125,7 +125,7 @@ def test_native_stream_result_is_not_consumed_or_wrapped() -> None: @pytest.mark.asyncio -@pytest.mark.parametrize("rules", ((), (SecretManagerRule(Rollout.RUST_REQUIRED),))) +@pytest.mark.parametrize("rules", ((), (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),))) async def test_async_route_without_rules_preserves_async_iterator_result(rules: Rules) -> None: async def chunks() -> AsyncGenerator[int, None]: yield 1 diff --git a/tests/unit/rust_bridge/test_secret_manager.py b/tests/unit/rust_bridge/test_secret_manager.py index e3e65a4e2bb..6392e1d1e98 100644 --- a/tests/unit/rust_bridge/test_secret_manager.py +++ b/tests/unit/rust_bridge/test_secret_manager.py @@ -4,12 +4,9 @@ import asyncio import inspect import json import re -from collections.abc import Mapping -from dataclasses import dataclass from functools import partial -from importlib import import_module from types import SimpleNamespace -from typing import Final, Never +from typing import TYPE_CHECKING, Final from urllib.parse import urlsplit import httpx @@ -20,26 +17,17 @@ from botocore.credentials import Credentials from pydantic import JsonValue import litellm -from litellm.rust_bridge import bindings -from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import Rules, SecretManagerRule -from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.secret_manager import ( - NativeSecretManagerFactory, - NativeSecretManagerRuntime, - capture_secret_manager, - native_secret_manager_config, - resolve_native_provider_reader, - resolve_native_provider_writer, - resolve_native_secret_manager, -) +from litellm.rust_bridge.secret_manager import capture_secret_manager, native_secret_manager_config from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 from litellm.secret_managers.cyberark_secret_manager import CyberArkSecretManager -from litellm.secret_managers.dispatch import get_secret_from_manager from litellm.secret_managers.hashicorp_secret_manager import HashicorpSecretManager +from litellm.secret_managers.secret_manager_handler import get_secret_from_manager from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service +if TYPE_CHECKING: + from litellm.rust_bridge._native import _SecretManagerRuntime + @pytest.fixture(autouse=True) def preserve_manager_globals(monkeypatch: pytest.MonkeyPatch) -> None: @@ -78,6 +66,13 @@ def _vault_body(value: str) -> dict[str, object]: } +def _native_handle(client: object) -> "_SecretManagerRuntime": + native: Final = pytest.importorskip("litellm.rust_bridge._native") + handle: Final = native._SecretManagerRuntime.from_client(client) + assert handle is not None + return handle + + @pytest.mark.parametrize("system", ("aws_secret_manager", "hashicorp_vault", "cyberark")) async def test_python_native_handle_reuses_backend_across_sync_and_async_reads(system: str) -> None: native: Final = pytest.importorskip("litellm.rust_bridge._native") @@ -162,39 +157,16 @@ def test_configuration_replacement_rebuilds_without_invalidating_existing_handle assert first.read_secret("KEY") == "first" -def test_custom_subclass_keeps_its_python_reader_under_native_selection() -> None: +def test_custom_subclass_is_not_captured_for_a_native_handle() -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + class CustomManager(AWSSecretsManagerV2): def sync_read_secret(self, secret_name: str, primary_secret_name: str | None = None) -> str: return f"custom:{secret_name}" - class DecliningFactory: - @staticmethod - def from_client(client: object) -> NativeSecretManagerRuntime | None: - assert native_secret_manager_config(client) is None - return None - - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(DecliningFactory) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"aws_secret_manager"})),) - assert ( - get_secret_from_manager( - CustomManager(aws_region_name="us-east-1"), "aws_secret_manager", "KEY", rules=rules, binding=binding - ) - == "custom:KEY" - ) - - -def test_read_dispatch_uses_native_backend_with_explicit_rules(monkeypatch: pytest.MonkeyPatch) -> None: - native: Final = pytest.importorskip("litellm.rust_bridge._native") - with recording_service() as server: - server.default_response = ResponseSpec(body=_vault_body("native-value")) - manager: Final = _vault(monkeypatch, server.base_url) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"hashicorp_vault"})),) - assert get_secret_from_manager(manager, "hashicorp_vault", "KEY", rules=rules) == "native-value" - runtime: Final = native._SecretManagerRuntime.from_client(manager) - assert runtime is not None - assert runtime.read_secret("KEY") == "native-value" - assert len(server.requests) == 1 + manager: Final = CustomManager(aws_region_name="us-east-1") + assert native._SecretManagerRuntime.from_client(manager) is None + assert manager.sync_read_secret("KEY") == "custom:KEY" @pytest.mark.parametrize("system", ("google_secret_manager", "hashicorp_vault", "cyberark")) @@ -220,55 +192,20 @@ def test_direct_builtin_constructor_can_use_native_without_registration(monkeypa assert handle.read_secret("KEY") == "native-value" -def test_binding_rejects_a_different_backend_before_reading() -> None: +def test_uncaptured_clients_have_no_native_handle() -> None: native: Final = pytest.importorskip("litellm.rust_bridge._native") - with recording_service() as server: - server.expected_requests = 0 - handle: Final = native._SecretManagerRuntime.from_config( - "hashicorp_vault", - {"HCP_VAULT_ADDR": server.base_url, "HCP_VAULT_TOKEN": "token"}, - enterprise_enabled=True, - ) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"aws_secret_manager"})),) - with pytest.raises(ValueError, match="system does not match"): - resolve_native_secret_manager(handle, "aws_secret_manager", rules) - - -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_OPT_OUT)) -def test_python_selection_and_missing_extension_preserve_python_reader( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout -) -> None: - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(None) - with recording_service() as server: - server.default_response = ResponseSpec(body=_vault_body("python-value")) - manager: Final = _vault(monkeypatch, server.base_url) - rules: Final[Rules] = (SecretManagerRule(rollout, systems=frozenset({"hashicorp_vault"})),) - assert ( - get_secret_from_manager(manager, "hashicorp_vault", "KEY", rules=rules, binding=binding) == "python-value" - ) - assert server.requests[0].headers["x-vault-token"] == "token" - - -def test_required_native_missing_extension_does_not_read_python(monkeypatch: pytest.MonkeyPatch) -> None: - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(None) - with recording_service() as server: - server.expected_requests = 0 - manager: Final = _vault(monkeypatch, server.base_url) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"hashicorp_vault"})),) - with pytest.raises(RuntimeError, match="unavailable"): - get_secret_from_manager(manager, "hashicorp_vault", "KEY", rules=rules, binding=binding) + assert native._SecretManagerRuntime.from_client(SimpleNamespace()) is None def test_native_failure_is_not_replayed_in_python(monkeypatch: pytest.MonkeyPatch) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: server.default_response = ResponseSpec(status=403, body={"errors": ["denied"]}) + server.expected_requests = 1 manager: Final = _vault(monkeypatch, server.base_url) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_OPT_OUT, systems=frozenset({"hashicorp_vault"})),) + handle: Final = _native_handle(manager) with pytest.raises(ValueError, match="HashiCorp Vault"): - get_secret_from_manager(manager, "hashicorp_vault", "KEY", rules=rules) + handle.read_secret("KEY") assert len(server.requests) == 1 @@ -328,96 +265,6 @@ def test_same_named_custom_client_is_not_captured() -> None: assert native_secret_manager_config(client) is None -@dataclass(slots=True) -class _RecordingRuntime: - system: str - result: str | None - calls: tuple[tuple[str, Mapping[str, object] | None], ...] = () - - def read_secret(self, name: str, settings: Mapping[str, object] | None = None) -> str | None: - self.calls = (*self.calls, (name, settings)) - return self.result - - -@pytest.mark.parametrize("value", (None, " value\n")) -@pytest.mark.parametrize("settings", (None, KeyManagementSettings(primary_secret_name="primary"))) -def test_native_dispatch_forwards_settings_and_preserves_missing_or_unmodified_values( - value: str | None, settings: KeyManagementSettings | None -) -> None: - client: Final = object() - runtime: Final = _RecordingRuntime("aws_secret_manager", value) - - class Factory: - @staticmethod - def from_client(candidate: object) -> NativeSecretManagerRuntime: - assert candidate is client - return runtime - - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(Factory) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({runtime.system})),) - - result: Final = get_secret_from_manager(client, runtime.system, "KEY", settings, rules=rules, binding=binding) - - assert result == value - assert runtime.calls == (("KEY", settings.model_dump(mode="json") if settings is not None else None),) - - -def test_native_system_mismatch_is_rejected_before_reading() -> None: - runtime: Final = _RecordingRuntime("hashicorp_vault", "wrong-provider") - - class Factory: - @staticmethod - def from_client(client: object) -> NativeSecretManagerRuntime: - return runtime - - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(Factory) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"aws_secret_manager"})),) - - with pytest.raises(ValueError, match="system does not match"): - get_secret_from_manager(object(), "aws_secret_manager", "KEY", rules=rules, binding=binding) - - assert runtime.calls == () - - -@pytest.mark.parametrize("system", ("custom", "local")) -def test_python_only_manager_types_never_construct_a_native_backend(system: str) -> None: - class ForbiddenFactory: - @staticmethod - def from_client(client: object) -> NativeSecretManagerRuntime: - raise AssertionError("custom and local clients cannot use native backends") - - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(ForbiddenFactory) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({system})),) - - assert resolve_native_secret_manager(object(), system, rules, binding=binding) is None - - -@pytest.mark.parametrize("factory", (None, SimpleNamespace(from_client="not callable"))) -def test_invalid_native_factories_are_reported_as_unavailable(monkeypatch: pytest.MonkeyPatch, factory: object) -> None: - monkeypatch.setattr(bindings, "get_native_bridge", lambda: SimpleNamespace(_SecretManagerRuntime=factory)) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"aws_secret_manager"})),) - - with pytest.raises(RuntimeError, match="runtime is unavailable"): - resolve_native_secret_manager(object(), "aws_secret_manager", rules) - - -def test_native_binding_accepts_a_callable_factory(monkeypatch: pytest.MonkeyPatch) -> None: - runtime: Final = _RecordingRuntime("aws_secret_manager", "value") - - class Factory: - @staticmethod - def from_client(client: object) -> NativeSecretManagerRuntime: - return runtime - - monkeypatch.setattr(bindings, "get_native_bridge", lambda: SimpleNamespace(_SecretManagerRuntime=Factory)) - rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({runtime.system})),) - - assert resolve_native_secret_manager(object(), runtime.system, rules) is runtime - - @pytest.mark.parametrize("value", ("text", "", True, False, 42, 2**100, [1, "two"], {"nested": True}, None)) async def test_aws_primary_values_match_python_handler(monkeypatch: pytest.MonkeyPatch, value: JsonValue) -> None: native: Final = pytest.importorskip("litellm.rust_bridge._native") @@ -429,14 +276,10 @@ async def test_aws_primary_values_match_python_handler(monkeypatch: pytest.Monke monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") settings: Final = KeyManagementSettings(primary_secret_name="primary") - reference: Final = get_secret_from_manager( - manager, "aws_secret_manager", "KEY", settings, rules=(SecretManagerRule(Rollout.PYTHON_ONLY),) - ) - actual: Final = get_secret_from_manager( - manager, "aws_secret_manager", "KEY", settings, rules=(SecretManagerRule(Rollout.RUST_REQUIRED),) - ) + reference: Final = get_secret_from_manager(manager, "aws_secret_manager", "KEY", settings) handle: Final = native._SecretManagerRuntime.from_client(manager) assert handle is not None + actual: Final = handle.read_secret("KEY", settings.model_dump(mode="json")) asynchronous: Final = await handle.read_secret_async("KEY", settings.model_dump(mode="json")) assert type(actual) is type(reference) is type(value) assert actual == reference == value @@ -445,6 +288,27 @@ async def test_aws_primary_values_match_python_handler(monkeypatch: pytest.Monke assert tuple(json.loads(request.raw_body) for request in server.requests) == ({"SecretId": "primary"},) * 3 +@pytest.mark.parametrize("primary", (None, "primary")) +async def test_aws_primary_values_match_python_reads(monkeypatch: pytest.MonkeyPatch, primary: str | None) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body={"SecretString": '{"KEY":"value"}'}) + server.expected_requests = 3 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + settings: Final = KeyManagementSettings(primary_secret_name=primary) + reference: Final = get_secret_from_manager(manager, "aws_secret_manager", "KEY", settings) + handle: Final = _native_handle(manager) + settings_dump: Final = settings.model_dump(mode="json") + synchronous: Final = handle.read_secret("KEY", settings_dump) + asynchronous: Final = await handle.read_secret_async("KEY", settings_dump) + assert synchronous == reference + assert asynchronous == reference + assert tuple(json.loads(request.raw_body) for request in server.requests) == ({"SecretId": primary or "KEY"},) * 3 + + @pytest.mark.parametrize("primary", (None, "primary")) @pytest.mark.parametrize( ("status", "body"), @@ -459,7 +323,6 @@ async def test_aws_primary_values_match_python_handler(monkeypatch: pytest.Monke def test_aws_absence_and_failed_reads_match_python_without_environment_fallback( monkeypatch: pytest.MonkeyPatch, primary: str | None, status: int, body: dict[str, str] ) -> None: - pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: server.default_response = ResponseSpec(status=status, body=body) server.expected_requests = 2 @@ -469,22 +332,9 @@ def test_aws_absence_and_failed_reads_match_python_without_environment_fallback( monkeypatch.setenv("KEY", "must-not-fall-back") manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") settings: Final = KeyManagementSettings(primary_secret_name=primary) - monkeypatch.setattr(litellm, "secret_manager_client", manager) - monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.AWS_SECRET_MANAGER) - monkeypatch.setattr(litellm, "_key_management_settings", settings) - main: Final = import_module("litellm.secret_managers.main") - monkeypatch.setattr( - main, - "get_secret_from_manager", - partial(get_secret_from_manager, rules=(SecretManagerRule(Rollout.PYTHON_ONLY),)), - ) - reference: Final = litellm.get_secret("KEY", "must-not-default") - monkeypatch.setattr( - main, - "get_secret_from_manager", - partial(get_secret_from_manager, rules=(SecretManagerRule(Rollout.RUST_REQUIRED),)), - ) - actual: Final = litellm.get_secret("KEY", "must-not-default") + reference: Final = get_secret_from_manager(manager, "aws_secret_manager", "KEY", settings) + handle: Final = _native_handle(manager) + actual: Final = handle.read_secret("KEY", settings.model_dump(mode="json")) assert actual == reference assert actual == ("" if primary is None and body.get("SecretString") == "" else None) @@ -503,43 +353,22 @@ async def test_aws_primary_json_errors_preserve_python_exception_details( manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") settings: Final = KeyManagementSettings(primary_secret_name="primary") with pytest.raises((json.JSONDecodeError, AttributeError)) as reference: - get_secret_from_manager( - manager, "aws_secret_manager", "KEY", settings, rules=(SecretManagerRule(Rollout.PYTHON_ONLY),) - ) + get_secret_from_manager(manager, "aws_secret_manager", "KEY", settings) + handle: Final = _native_handle(manager) + settings_dump: Final = settings.model_dump(mode="json") with pytest.raises(type(reference.value)) as actual: - get_secret_from_manager( - manager, "aws_secret_manager", "KEY", settings, rules=(SecretManagerRule(Rollout.RUST_REQUIRED),) - ) + handle.read_secret("KEY", settings_dump) assert actual.value.args == reference.value.args if isinstance(reference.value, json.JSONDecodeError): assert isinstance(actual.value, json.JSONDecodeError) assert (actual.value.doc, actual.value.pos) == (reference.value.doc, reference.value.pos) - handle: Final = native._SecretManagerRuntime.from_client(manager) - assert handle is not None with pytest.raises(type(reference.value)) as asynchronous: - await handle.read_secret_async("KEY", settings.model_dump(mode="json")) + await handle.read_secret_async("KEY", settings_dump) assert asynchronous.value.args == reference.value.args -def _select_provider_reads(monkeypatch: pytest.MonkeyPatch, module_name: str, rollout: Rollout) -> None: - module: Final = import_module(module_name) - monkeypatch.setattr( - module, - "resolve_native_provider_reader", - partial(resolve_native_provider_reader, rules=(SecretManagerRule(rollout),)), - ) - if rollout is Rollout.RUST_REQUIRED: - monkeypatch.setattr(module, "_get_httpx_client", _forbid_python_http) - monkeypatch.setattr(module, "get_async_httpx_client", _forbid_python_http) - - -def _forbid_python_http(*args: object, **kwargs: object) -> Never: - raise AssertionError("native reads must not construct a Python HTTP client") - - -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) async def test_public_aws_reads_preserve_coroutines_and_per_call_credentials( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout + monkeypatch: pytest.MonkeyPatch, ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as initial, recording_service() as selected: @@ -550,7 +379,6 @@ async def test_public_aws_reads_preserve_coroutines_and_per_call_credentials( monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "environment-secret") monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", initial.base_url) manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") - _select_provider_reads(monkeypatch, "litellm.secret_managers.aws_secret_manager_v2", rollout) options: Final = { "aws_region_name": "us-west-2", "aws_bedrock_runtime_endpoint": selected.base_url, @@ -592,10 +420,7 @@ async def test_public_aws_reads_preserve_coroutines_and_per_call_credentials( ) -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) -async def test_public_aws_primary_reads_ignore_operation_overrides_like_python( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout -) -> None: +async def test_public_aws_primary_reads_ignore_operation_overrides(monkeypatch: pytest.MonkeyPatch) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server, recording_service() as unused: server.default_response = ResponseSpec(body={"SecretString": '{"KEY":true}'}) @@ -605,7 +430,6 @@ async def test_public_aws_primary_reads_ignore_operation_overrides_like_python( monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") - _select_provider_reads(monkeypatch, "litellm.secret_managers.aws_secret_manager_v2", rollout) options: Final = {"aws_bedrock_runtime_endpoint": unused.base_url} assert manager.sync_read_secret("KEY", options, 0, "primary") is True assert ( @@ -615,10 +439,7 @@ async def test_public_aws_primary_reads_ignore_operation_overrides_like_python( assert tuple(json.loads(request.raw_body) for request in server.requests) == ({"SecretId": "primary"},) * 2 -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) -async def test_public_aws_bootstrap_names_only_bypass_sync_reads( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout -) -> None: +async def test_public_aws_bootstrap_names_only_bypass_sync_reads(monkeypatch: pytest.MonkeyPatch) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: server.default_response = ResponseSpec(body={"SecretString": "remote-access"}) @@ -627,16 +448,14 @@ async def test_public_aws_bootstrap_names_only_bypass_sync_reads( monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "environment-secret") monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") - _select_provider_reads(monkeypatch, "litellm.secret_managers.aws_secret_manager_v2", rollout) assert manager.sync_read_secret("AWS_ACCESS_KEY_ID") == "environment-access" assert server.requests == [] assert await manager.async_read_secret("AWS_ACCESS_KEY_ID") == "remote-access" -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) @pytest.mark.parametrize("timeout", (0.05, httpx.Timeout(1, read=0.05))) async def test_public_aws_read_timeouts_follow_the_python_http_handler( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout, timeout: float | httpx.Timeout + monkeypatch: pytest.MonkeyPatch, timeout: float | httpx.Timeout ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: @@ -646,21 +465,16 @@ async def test_public_aws_read_timeouts_follow_the_python_http_handler( monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") - _select_provider_reads(monkeypatch, "litellm.secret_managers.aws_secret_manager_v2", rollout) assert manager.sync_read_secret("KEY", timeout=timeout) is None assert await manager.async_read_secret("KEY", timeout=timeout) is None -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) -async def test_public_vault_reads_keep_overrides_cache_and_coroutines( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout -) -> None: +async def test_public_vault_reads_keep_overrides_cache_and_coroutines(monkeypatch: pytest.MonkeyPatch) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: server.default_response = ResponseSpec(body=_vault_body("value")) server.expected_requests = 2 manager: Final = _vault(monkeypatch, server.base_url) - _select_provider_reads(monkeypatch, "litellm.secret_managers.hashicorp_secret_manager", rollout) options: Final = {"secret_manager_settings": {"mount": "team", "path_prefix": "keys", "data": "key"}} pending: Final = manager.async_read_secret("KEY", options) assert inspect.iscoroutine(pending) @@ -672,24 +486,21 @@ async def test_public_vault_reads_keep_overrides_cache_and_coroutines( assert urlsplit(server.requests[1].path).path == "/v1/secret/data/KEY" -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) @pytest.mark.parametrize("status", (404, 403)) async def test_public_vault_failed_reads_return_none_without_replay( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout, status: int + monkeypatch: pytest.MonkeyPatch, status: int ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: server.default_response = ResponseSpec(status=status, body={"errors": ["unavailable"]}) server.expected_requests = 2 manager: Final = _vault(monkeypatch, server.base_url) - _select_provider_reads(monkeypatch, "litellm.secret_managers.hashicorp_secret_manager", rollout) assert manager.sync_read_secret("KEY") is None assert await manager.async_read_secret("KEY") is None -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) async def test_public_cyberark_reads_reuse_authentication_and_cached_values( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout + monkeypatch: pytest.MonkeyPatch, ) -> None: pytest.importorskip("litellm.rust_bridge._native") from litellm.proxy import proxy_server @@ -705,58 +516,25 @@ async def test_public_cyberark_reads_reuse_authentication_and_cached_values( monkeypatch.setenv("CYBERARK_ACCOUNT", "account") monkeypatch.setenv("CYBERARK_USERNAME", "reader") manager: Final = CyberArkSecretManager() - _select_provider_reads(monkeypatch, "litellm.secret_managers.cyberark_secret_manager", rollout) pending: Final = manager.async_read_secret(secret_name="KEY", timeout=0) assert inspect.iscoroutine(pending) assert server.requests == [] assert await asyncio.create_task(pending) == '"secret-value"' - assert manager.sync_read_secret("KEY", timeout=0) == ( - '"secret-value"' if rollout is Rollout.RUST_REQUIRED else "secret-value" - ) + assert manager.sync_read_secret("KEY", timeout=0) == "secret-value" assert tuple(request.path for request in server.requests) == ( "/authn/account/reader/authenticate", "/secrets/account/variable/KEY", ) -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) -async def test_public_native_selection_and_missing_extension_keep_the_python_method( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout -) -> None: - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(None) - module: Final = import_module("litellm.secret_managers.aws_secret_manager_v2") - monkeypatch.setattr( - module, - "resolve_native_provider_reader", - partial(resolve_native_provider_reader, rules=(SecretManagerRule(rollout),), binding=binding), - ) - with recording_service() as server: - server.default_response = ResponseSpec(body={"SecretString": "python-value"}) - server.expected_requests = 0 if rollout is Rollout.RUST_REQUIRED else 1 - monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test-access") - monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") - monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) - manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") - pending: Final = manager.async_read_secret("KEY") - if rollout is Rollout.RUST_REQUIRED: - with pytest.raises(RuntimeError, match="runtime is unavailable"): - await pending - else: - assert await pending == "python-value" - - -def test_public_aws_bootstrap_read_does_not_initialize_a_backend(monkeypatch: pytest.MonkeyPatch) -> None: - module: Final = import_module("litellm.secret_managers.aws_secret_manager_v2") - monkeypatch.setattr(module, "resolve_native_provider_reader", _forbid_python_http) +def test_public_aws_bootstrap_read_does_not_hit_the_provider(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("AWS_ACCESS_KEY_ID", "bootstrap-value") assert AWSSecretsManagerV2().sync_read_secret("AWS_ACCESS_KEY_ID") == "bootstrap-value" -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) @pytest.mark.parametrize("override", ("", 42)) def test_public_vault_prefix_overrides_match_python_string_conversion( - monkeypatch: pytest.MonkeyPatch, rollout: Rollout, override: str | int + monkeypatch: pytest.MonkeyPatch, override: str | int ) -> None: pytest.importorskip("litellm.rust_bridge._native") monkeypatch.setenv("HCP_VAULT_PATH_PREFIX", "default-prefix") @@ -764,64 +542,12 @@ def test_public_vault_prefix_overrides_match_python_string_conversion( server.default_response = ResponseSpec(body=_vault_body("value")) server.expected_requests = 1 manager: Final = _vault(monkeypatch, server.base_url) - _select_provider_reads(monkeypatch, "litellm.secret_managers.hashicorp_secret_manager", rollout) assert manager.sync_read_secret("KEY", {"path_prefix": override}) == "value" assert urlsplit(server.requests[0].path).path == ( f"/v1/secret/data/{override}/KEY" if override else "/v1/secret/data/KEY" ) -@pytest.mark.parametrize("value", ("native-value", None)) -def test_public_google_reader_uses_the_selected_binding_without_replaying_python( - monkeypatch: pytest.MonkeyPatch, value: str | None -) -> None: - from litellm.proxy import proxy_server - from litellm.secret_managers.google_secret_manager import GoogleSecretManager - - class Manager(GoogleSecretManager): - def sync_construct_request_headers(self) -> dict[str, str]: - raise AssertionError("selected native reads must not construct Python auth headers") - - class Reader(_RecordingRuntime): - def sync_read_secret( - self, - secret_name: str, - optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, - ) -> str | None: - return self.read_secret(secret_name) - - async def async_read_secret( - self, - secret_name: str, - optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, - ) -> str | None: - return self.sync_read_secret(secret_name) - - monkeypatch.setattr(proxy_server, "premium_user", True) - monkeypatch.setenv("GOOGLE_SECRET_MANAGER_PROJECT_ID", "project") - manager: Final = Manager() - runtime: Final = Reader("google_secret_manager", value) - - class Factory: - @staticmethod - def from_client(candidate: object) -> NativeSecretManagerRuntime: - assert candidate is manager - return runtime - - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(Factory) - module: Final = import_module("litellm.secret_managers.google_secret_manager") - monkeypatch.setattr( - module, - "resolve_native_provider_reader", - partial(resolve_native_provider_reader, rules=(SecretManagerRule(Rollout.RUST_REQUIRED),), binding=binding), - ) - assert manager.get_secret_from_google_secret_manager(secret_name="KEY") == value - assert runtime.calls == (("KEY", None),) - - def _cyberark(monkeypatch: pytest.MonkeyPatch, address: str) -> CyberArkSecretManager: from litellm.proxy import proxy_server @@ -833,20 +559,8 @@ def _cyberark(monkeypatch: pytest.MonkeyPatch, address: str) -> CyberArkSecretMa return CyberArkSecretManager() -def _select_cyberark_mutations(monkeypatch: pytest.MonkeyPatch, rollout: Rollout) -> None: - module: Final = import_module("litellm.secret_managers.cyberark_secret_manager") - _select_provider_reads(monkeypatch, module.__name__, rollout) - monkeypatch.setattr( - module, - "resolve_native_provider_writer", - partial(resolve_native_provider_writer, rules=(SecretManagerRule(rollout),)), - ) - - -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) async def test_public_cyberark_writes_and_deletes_share_the_read_cache( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: @@ -854,7 +568,6 @@ async def test_public_cyberark_writes_and_deletes_share_the_read_cache( server.enqueue(ResponseSpec(body=body)) server.expected_requests = 5 manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, rollout) assert manager.sync_read_secret("KEY") == "old" pending: Final = manager.async_write_secret( "KEY", @@ -910,22 +623,17 @@ async def test_public_cyberark_write_errors_match_python_without_http_retries( for response in responses * 2: server.enqueue(response) server.expected_requests = len(responses) * 2 - reference_manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, Rollout.PYTHON_ONLY) - reference: Final = await reference_manager.async_write_secret("KEY", "value") - native_manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) - actual: Final = await native_manager.async_write_secret("KEY", "value") + reference: Final = await _cyberark(monkeypatch, server.base_url).async_write_secret("KEY", "value") + native_handle: Final = _native_handle(_cyberark(monkeypatch, server.base_url)) + actual: Final = await native_handle.async_write_secret("KEY", "value") assert actual == reference assert tuple(actual) == tuple(reference) assert actual["status"] == "error" assert str(status) in actual["message"] -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) async def test_public_cyberark_write_recovers_from_initial_policy_authentication_failure( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: @@ -934,7 +642,6 @@ async def test_public_cyberark_write_recovers_from_initial_policy_authentication server.enqueue(ResponseSpec(body={})) server.expected_requests = 3 manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, rollout) assert await manager.async_write_secret("KEY", "value") == { "status": "success", "message": "Secret KEY written successfully", @@ -946,40 +653,32 @@ async def test_public_cyberark_write_recovers_from_initial_policy_authentication ) -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) -@pytest.mark.parametrize("name", ("../KEY", "line\nKEY", "a\u2028b")) +@pytest.mark.parametrize("name", ("../KEY", "line\nKEY", "a
b")) async def test_public_cyberark_write_rejects_unsafe_names_before_authentication( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, name: str, ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: server.expected_requests = 0 manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, rollout) assert await manager.async_write_secret(name, "value") == { "status": "error", "message": f"Invalid secret_name {name!r}", } -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) @pytest.mark.parametrize("same_name", (False, True)) async def test_public_cyberark_rotation_returns_the_write_response_and_retains_old_alias( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, same_name: bool, ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: for body in (b"token", b"old-value", {}, {}): server.enqueue(ResponseSpec(body=body)) - if rollout is Rollout.RUST_REQUIRED: - server.enqueue(ResponseSpec(body=b"new-value")) - server.expected_requests = 5 if rollout is Rollout.RUST_REQUIRED else 4 + server.expected_requests = 4 manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, rollout) new_name: Final = "OLD" if same_name else "NEW" pending: Final = manager.async_rotate_secret("OLD", new_name, "new-value", {"ignored": object()}, 0) assert inspect.iscoroutine(pending) @@ -988,16 +687,12 @@ async def test_public_cyberark_rotation_returns_the_write_response_and_retains_o "status": "success", "message": f"Secret {new_name} written successfully", } - assert tuple(request.method for request in server.requests) == ( - ("POST", "GET", "POST", "POST", "GET") - if rollout is Rollout.RUST_REQUIRED - else ("POST", "GET", "POST", "POST") - ) + assert tuple(request.method for request in server.requests) == ("POST", "GET", "POST", "POST") assert server.requests[3].raw_body == b"new-value" @pytest.mark.parametrize("replacement", (None, b"wrong-value")) -async def test_public_cyberark_rotation_requires_a_fresh_matching_replacement( +async def test_public_native_cyberark_rotation_requires_a_fresh_matching_replacement( monkeypatch: pytest.MonkeyPatch, replacement: bytes | None, ) -> None: @@ -1008,11 +703,11 @@ async def test_public_cyberark_rotation_requires_a_fresh_matching_replacement( server.enqueue(ResponseSpec(status=404 if replacement is None else 200, body=replacement)) server.expected_requests = 5 manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) + handle: Final = _native_handle(manager) message: Final = "Failed to verify new secret NEW" if replacement is None else "New secret value mismatch" with pytest.raises(ValueError, match=message): - await manager.async_rotate_secret("OLD", "NEW", "new-value") - assert manager.sync_read_secret("OLD") == "old-value" + await handle.async_rotate_secret("OLD", "NEW", "new-value") + assert handle.sync_read_secret("OLD") == "old-value" assert tuple(request.path for request in server.requests) == ( "/authn/account/reader/authenticate", "/secrets/account/variable/OLD", @@ -1022,10 +717,8 @@ async def test_public_cyberark_rotation_requires_a_fresh_matching_replacement( ) -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) async def test_public_cyberark_cached_authentication_does_not_retry_denied_reads( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: @@ -1034,7 +727,6 @@ async def test_public_cyberark_cached_authentication_does_not_retry_denied_reads server.enqueue(ResponseSpec(status=401, body={})) server.expected_requests = 3 manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, rollout) assert manager.sync_read_secret("OLD") == "value" assert await manager.async_read_secret("NEW") is None @@ -1052,16 +744,13 @@ async def test_cyberark_handler_errors_match_python_after_cached_authentication_ server.enqueue(response) server.expected_requests = 6 reference_manager: Final = _cyberark(monkeypatch, server.base_url) - python_rules: Final = (SecretManagerRule(Rollout.PYTHON_ONLY),) - native_rules: Final = (SecretManagerRule(Rollout.RUST_REQUIRED),) - assert get_secret_from_manager(reference_manager, "cyberark", "OLD", rules=python_rules) == "value" + assert get_secret_from_manager(reference_manager, "cyberark", "OLD") == "value" with pytest.raises(ValueError, match="No secret found in CyberArk Secret Manager for NEW") as reference: - get_secret_from_manager(reference_manager, "cyberark", "NEW", rules=python_rules) - native_manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) - assert get_secret_from_manager(native_manager, "cyberark", "OLD", rules=native_rules) == "value" + get_secret_from_manager(reference_manager, "cyberark", "NEW") + native_handle: Final = _native_handle(_cyberark(monkeypatch, server.base_url)) + assert native_handle.read_secret("OLD") == "value" with pytest.raises(ValueError, match="No secret found in CyberArk Secret Manager for NEW") as actual: - get_secret_from_manager(native_manager, "cyberark", "NEW", rules=native_rules) + native_handle.read_secret("NEW") assert actual.value.args == reference.value.args @@ -1072,89 +761,31 @@ async def test_public_cyberark_connection_errors_match_python( with recording_service() as server: server.expected_requests = 0 address: Final = server.base_url - reference_manager: Final = _cyberark(monkeypatch, address) - _select_cyberark_mutations(monkeypatch, Rollout.PYTHON_ONLY) - reference: Final = await reference_manager.async_write_secret("KEY", "value") - native_manager: Final = _cyberark(monkeypatch, address) - _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) - actual: Final = await native_manager.async_write_secret("KEY", "value") + reference: Final = await _cyberark(monkeypatch, address).async_write_secret("KEY", "value") + native_handle: Final = _native_handle(_cyberark(monkeypatch, address)) + actual: Final = await native_handle.async_write_secret("KEY", "value") assert actual == reference assert actual["status"] == "error" -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) -@pytest.mark.parametrize("operation", ("write", "delete", "rotate")) -async def test_public_cyberark_mutations_preserve_missing_extension_selection( - monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, - operation: str, -) -> None: - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(None) - module: Final = import_module("litellm.secret_managers.cyberark_secret_manager") - _select_provider_reads(monkeypatch, module.__name__, Rollout.PYTHON_ONLY) - monkeypatch.setattr( - module, - "resolve_native_provider_writer", - partial(resolve_native_provider_writer, rules=(SecretManagerRule(rollout),), binding=binding), - ) - with recording_service() as server: - bodies: Final = ( - () - if rollout is Rollout.RUST_REQUIRED or operation == "delete" - else ((b"token", b"old", {}, {}) if operation == "rotate" else (b"token", {}, {})) - ) - for body in bodies: - server.enqueue(ResponseSpec(body=body)) - server.expected_requests = len(bodies) - manager: Final = _cyberark(monkeypatch, server.base_url) - call: Final = { - "write": partial(manager.async_write_secret, "KEY", "value"), - "delete": partial(manager.async_delete_secret, "KEY"), - "rotate": partial(manager.async_rotate_secret, "OLD", "NEW", "value"), - }[operation] - pending: Final = call() - assert inspect.iscoroutine(pending) - assert server.requests == [] - if rollout is Rollout.RUST_REQUIRED: - with pytest.raises(RuntimeError, match="runtime is unavailable"): - await pending - else: - result: Final = await pending - assert result["status"] == ("not_supported" if operation == "delete" else "success") - - -async def test_public_cyberark_rotation_stops_after_a_failed_write(monkeypatch: pytest.MonkeyPatch) -> None: +async def test_public_native_cyberark_rotation_stops_after_a_failed_write(monkeypatch: pytest.MonkeyPatch) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: for body in (b"token", b"old-value", {}): server.enqueue(ResponseSpec(body=body)) server.enqueue(ResponseSpec(status=401, body={})) server.expected_requests = 4 - manager: Final = _cyberark(monkeypatch, server.base_url) - _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) - response: Final = await manager.async_rotate_secret("OLD", "NEW", "new-value") + handle: Final = _native_handle(_cyberark(monkeypatch, server.base_url)) + response: Final = await handle.async_rotate_secret("OLD", "NEW", "new-value") assert response["status"] == "error" assert "401" in response["message"] - assert manager.sync_read_secret("OLD") == "old-value" + assert handle.sync_read_secret("OLD") == "old-value" assert tuple(request.method for request in server.requests) == ("POST", "GET", "POST", "POST") -def _select_vault_mutations(monkeypatch: pytest.MonkeyPatch, rollout: Rollout) -> None: - module: Final = import_module("litellm.secret_managers.hashicorp_secret_manager") - _select_provider_reads(monkeypatch, module.__name__, rollout) - monkeypatch.setattr( - module, - "resolve_native_provider_writer", - partial(resolve_native_provider_writer, rules=(SecretManagerRule(rollout),)), - ) - - -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) @pytest.mark.parametrize("description", (None, "", "purpose")) async def test_public_vault_writes_preserve_complete_responses_and_request_fields( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, description: str | None, ) -> None: pytest.importorskip("litellm.rust_bridge._native") @@ -1168,7 +799,6 @@ async def test_public_vault_writes_preserve_complete_responses_and_request_field server.enqueue(ResponseSpec(body=response)) server.expected_requests = 1 manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, rollout) options: Final = { "secret_manager_settings": {"namespace": "team", "mount": "kv", "path_prefix": "app", "data": "token"} } @@ -1204,19 +834,16 @@ async def test_public_vault_mutation_http_errors_match_python_without_retry( server.default_response = ResponseSpec(status=status, body={"errors": ["denied"]}) server.expected_requests = 2 options: Final = {"namespace": "team", "mount": "kv", "path_prefix": "prefix"} - reference_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) reference: Final = ( - await reference_manager.async_write_secret("KEY", "value", optional_params=options) + await _vault(monkeypatch, server.base_url).async_write_secret("KEY", "value", optional_params=options) if operation == "write" - else await reference_manager.async_delete_secret("KEY", optional_params=options) + else await _vault(monkeypatch, server.base_url).async_delete_secret("KEY", optional_params=options) ) - native_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + native_handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) actual: Final = ( - await native_manager.async_write_secret("KEY", "value", optional_params=options) + await native_handle.async_write_secret("KEY", "value", optional_params=options) if operation == "write" - else await native_manager.async_delete_secret("KEY", optional_params=options) + else await native_handle.async_delete_secret("KEY", optional_params=options) ) assert actual == reference assert tuple(actual) == tuple(reference) @@ -1233,20 +860,15 @@ async def test_public_vault_write_response_conversion_matches_python( with recording_service() as server: server.default_response = ResponseSpec(body=body) server.expected_requests = 2 - reference_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) - reference: Final = await reference_manager.async_write_secret("KEY", "value") - native_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) - actual: Final = await native_manager.async_write_secret("KEY", "value") + reference: Final = await _vault(monkeypatch, server.base_url).async_write_secret("KEY", "value") + native_handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) + actual: Final = await native_handle.async_write_secret("KEY", "value") assert type(actual) is type(reference) assert actual == reference -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) async def test_public_vault_deletion_invalidates_cached_fields( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: @@ -1255,7 +877,6 @@ async def test_public_vault_deletion_invalidates_cached_fields( server.enqueue(ResponseSpec(body=_vault_body("after-delete"))) server.expected_requests = 3 manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, rollout) assert manager.sync_read_secret("KEY") == "old" pending: Final = manager.async_delete_secret("KEY", None, {"ignored": object()}, 2) assert inspect.iscoroutine(pending) @@ -1265,12 +886,10 @@ async def test_public_vault_deletion_invalidates_cached_fields( assert tuple(request.method for request in server.requests) == ("GET", "DELETE", "GET") -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) @pytest.mark.parametrize("same_name", (False, True)) @pytest.mark.parametrize("delete_status", (204, 403)) async def test_public_vault_rotation_preserves_response_and_best_effort_deletion( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, same_name: bool, delete_status: int, ) -> None: @@ -1284,7 +903,6 @@ async def test_public_vault_rotation_preserves_response_and_best_effort_deletion server.enqueue(ResponseSpec(status=delete_status, body=b"")) server.expected_requests = 3 if same_name else 4 manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, rollout) new_name: Final = "OLD" if same_name else "NEW" pending: Final = manager.async_rotate_secret("OLD", new_name, "replacement", timeout=2) assert inspect.iscoroutine(pending) @@ -1324,12 +942,9 @@ async def test_public_vault_rotation_failure_messages_and_request_counts_match_p for response in responses * 2: server.enqueue(response) server.expected_requests = len(responses) * 2 - reference_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) - reference: Final = await reference_manager.async_rotate_secret("OLD", "NEW", "value") - native_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) - actual: Final = await native_manager.async_rotate_secret("OLD", "NEW", "value") + reference: Final = await _vault(monkeypatch, server.base_url).async_rotate_secret("OLD", "NEW", "value") + native_handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) + actual: Final = await native_handle.async_rotate_secret("OLD", "NEW", "value") assert actual == reference assert actual["status"] == "error" expected_paths: Final = ( @@ -1355,12 +970,9 @@ async def test_public_vault_rotation_mismatches_do_not_delete_the_old_alias( for response in responses * 2: server.enqueue(response) server.expected_requests = 6 - reference_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) - reference: Final = await reference_manager.async_rotate_secret("OLD", "NEW", "value") - native_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) - actual: Final = await native_manager.async_rotate_secret("OLD", "NEW", "value") + reference: Final = await _vault(monkeypatch, server.base_url).async_rotate_secret("OLD", "NEW", "value") + native_handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) + actual: Final = await native_handle.async_rotate_secret("OLD", "NEW", "value") assert actual == reference assert actual["status"] == "error" assert all(request.method != "DELETE" for request in server.requests) @@ -1376,18 +988,16 @@ async def test_public_vault_mutation_timeouts_match_python( server.default_response = ResponseSpec(body=_vault_body("value"), delay=0.25) server.expected_requests = 2 reference_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) reference: Final = await { "write": partial(reference_manager.async_write_secret, "KEY", "value"), "delete": partial(reference_manager.async_delete_secret, "KEY"), "rotate": partial(reference_manager.async_rotate_secret, "OLD", "NEW", "value"), }[operation](timeout=0.05) - native_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + native_handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) actual: Final = await { - "write": partial(native_manager.async_write_secret, "KEY", "value"), - "delete": partial(native_manager.async_delete_secret, "KEY"), - "rotate": partial(native_manager.async_rotate_secret, "OLD", "NEW", "value"), + "write": partial(native_handle.async_write_secret, "KEY", "value"), + "delete": partial(native_handle.async_delete_secret, "KEY"), + "rotate": partial(native_handle.async_rotate_secret, "OLD", "NEW", "value"), }[operation](timeout=0.05) if operation == "write": assert isinstance(actual["message"], str) @@ -1406,18 +1016,15 @@ async def test_public_vault_mutation_timeouts_match_python( assert actual["status"] == "error" -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) @pytest.mark.parametrize("operation", ("write", "delete", "rotate")) async def test_public_vault_unsafe_names_fail_before_authentication( monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, operation: str, ) -> None: pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: server.expected_requests = 0 manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, rollout) result: Final = await { "write": partial(manager.async_write_secret, "../KEY", "value"), "delete": partial(manager.async_delete_secret, "../KEY"), @@ -1438,18 +1045,16 @@ async def test_public_vault_authentication_errors_match_python( server.default_response = ResponseSpec(status=403, body={"errors": ["denied"]}) server.expected_requests = 2 reference_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) reference: Final = await { "write": partial(reference_manager.async_write_secret, "KEY", "value"), "delete": partial(reference_manager.async_delete_secret, "KEY"), "rotate": partial(reference_manager.async_rotate_secret, "OLD", "NEW", "value"), }[operation]() - native_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + native_handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) actual: Final = await { - "write": partial(native_manager.async_write_secret, "KEY", "value"), - "delete": partial(native_manager.async_delete_secret, "KEY"), - "rotate": partial(native_manager.async_rotate_secret, "OLD", "NEW", "value"), + "write": partial(native_handle.async_write_secret, "KEY", "value"), + "delete": partial(native_handle.async_delete_secret, "KEY"), + "rotate": partial(native_handle.async_rotate_secret, "OLD", "NEW", "value"), }[operation]() assert actual == reference assert actual["status"] == "error" @@ -1465,12 +1070,9 @@ async def test_public_vault_rotation_stops_on_a_success_response_containing_an_e for response in (ResponseSpec(body=_vault_body("old")), ResponseSpec(body=result)) * 2: server.enqueue(response) server.expected_requests = 4 - reference_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) - reference: Final = await reference_manager.async_rotate_secret("OLD", "NEW", "value") - native_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) - actual: Final = await native_manager.async_rotate_secret("OLD", "NEW", "value") + reference: Final = await _vault(monkeypatch, server.base_url).async_rotate_secret("OLD", "NEW", "value") + native_handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) + actual: Final = await native_handle.async_rotate_secret("OLD", "NEW", "value") assert actual == reference == result assert tuple(request.method for request in server.requests) == ("GET", "POST", "GET", "POST") @@ -1482,11 +1084,10 @@ async def test_public_native_vault_write_invalidates_stale_cached_values(monkeyp server.enqueue(ResponseSpec(body={"data": {"version": 2}})) server.enqueue(ResponseSpec(body=_vault_body("new"))) server.expected_requests = 3 - manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) - assert manager.sync_read_secret("KEY") == "old" - assert await manager.async_write_secret("KEY", "new") == {"data": {"version": 2}} - assert manager.sync_read_secret("KEY") == "new" + handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) + assert handle.sync_read_secret("KEY") == "old" + assert await handle.async_write_secret("KEY", "new") == {"data": {"version": 2}} + assert handle.sync_read_secret("KEY") == "new" assert tuple(request.method for request in server.requests) == ("GET", "POST", "GET") @@ -1496,53 +1097,13 @@ async def test_public_native_vault_write_rejects_description_overwriting_the_sec pytest.importorskip("litellm.rust_bridge._native") with recording_service() as server: server.expected_requests = 0 - manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) - assert await manager.async_write_secret("KEY", "value", "description", {"data": "description"}) == { + handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) + assert await handle.async_write_secret("KEY", "value", "description", {"data": "description"}) == { "status": "error", "message": "HashiCorp Vault data key conflicts with description", } -@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) -@pytest.mark.parametrize("operation", ("write", "delete", "rotate")) -async def test_public_vault_mutations_preserve_missing_extension_selection( - monkeypatch: pytest.MonkeyPatch, - rollout: Rollout, - operation: str, -) -> None: - binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) - binding.override(None) - module: Final = import_module("litellm.secret_managers.hashicorp_secret_manager") - _select_provider_reads(monkeypatch, module.__name__, Rollout.PYTHON_ONLY) - monkeypatch.setattr( - module, - "resolve_native_provider_writer", - partial(resolve_native_provider_writer, rules=(SecretManagerRule(rollout),), binding=binding), - ) - with recording_service() as server: - server.default_response = ResponseSpec(body=_vault_body("value")) - server.expected_requests = 0 if rollout is Rollout.RUST_REQUIRED else (4 if operation == "rotate" else 1) - manager: Final = _vault(monkeypatch, server.base_url) - pending: Final = { - "write": partial(manager.async_write_secret, "KEY", "value"), - "delete": partial(manager.async_delete_secret, "KEY"), - "rotate": partial(manager.async_rotate_secret, "OLD", "NEW", "value"), - }[operation]() - assert inspect.iscoroutine(pending) - assert server.requests == [] - if rollout is Rollout.RUST_REQUIRED: - with pytest.raises(RuntimeError, match="runtime is unavailable"): - await pending - else: - result: Final = await pending - assert result == ( - {"status": "success", "message": "Secret KEY deleted successfully"} - if operation == "delete" - else _vault_body("value") - ) - - @pytest.mark.parametrize( "body", (None, [], 42, 2**100, {"data": 2**100}, {"data": None}, {"data": {"data": []}}, {}, {"data": {}}) ) @@ -1560,12 +1121,9 @@ async def test_public_vault_rotation_preserves_malformed_verification_errors( for response in responses * 2: server.enqueue(response) server.expected_requests = 6 - reference_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) - reference: Final = await reference_manager.async_rotate_secret("OLD", "NEW", "value") - native_manager: Final = _vault(monkeypatch, server.base_url) - _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) - actual: Final = await native_manager.async_rotate_secret("OLD", "NEW", "value") + reference: Final = await _vault(monkeypatch, server.base_url).async_rotate_secret("OLD", "NEW", "value") + native_handle: Final = _native_handle(_vault(monkeypatch, server.base_url)) + actual: Final = await native_handle.async_rotate_secret("OLD", "NEW", "value") assert actual == reference assert actual["status"] == "error" assert all(request.method != "DELETE" for request in server.requests) diff --git a/tests/unit/rust_bridge/test_settings.py b/tests/unit/rust_bridge/test_settings.py index 7650195b3c5..ea943afe970 100644 --- a/tests/unit/rust_bridge/test_settings.py +++ b/tests/unit/rust_bridge/test_settings.py @@ -6,9 +6,7 @@ import pytest import litellm from litellm.integrations.custom_secret_manager import CustomSecretManager from litellm.llms.custom_httpx.http_handler import default_user_agent -from litellm.rust_bridge import catalog, settings -from litellm.rust_bridge.catalog import SecretManagerRule -from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge import settings from litellm.secret_managers.main import get_secret_str from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem @@ -81,11 +79,6 @@ class _VaultSecrets(CustomSecretManager): return self.secrets.get(secret_name) -_RUST_FOR_CUSTOM: Final = ( - SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({KeyManagementSystem.CUSTOM.value})), -) - - @pytest.mark.parametrize( ("access_mode", "readable"), [("read_only", True), ("read_and_write", True), ("write_only", False)], @@ -98,38 +91,30 @@ def test_secret_manager_is_readable_only_when_litellm_would_read_secrets_from_it monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.CUSTOM) monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(access_mode=access_mode)) - assert settings.secret_manager(rules=()) == settings.SecretManager(readable=readable, native=False) + assert settings.secret_manager() == settings.SecretManager(readable=readable, native=readable) assert (get_secret_str("MISTRAL_API_KEY") == "vault-key") is readable @pytest.mark.parametrize( - ("system", "access_mode", "rules", "native"), + ("system", "access_mode", "native"), [ - (KeyManagementSystem.CUSTOM, "read_only", _RUST_FOR_CUSTOM, True), - (KeyManagementSystem.CUSTOM, "read_only", (), False), - ( - KeyManagementSystem.CUSTOM, - "read_only", - (SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.CUSTOM.value})),), - False, - ), - (KeyManagementSystem.CUSTOM, "write_only", _RUST_FOR_CUSTOM, False), - (None, "read_only", _RUST_FOR_CUSTOM, False), - (KeyManagementSystem.AWS_SECRET_MANAGER, "read_only", _RUST_FOR_CUSTOM, False), + (KeyManagementSystem.CUSTOM, "read_only", True), + (KeyManagementSystem.AWS_SECRET_MANAGER, "read_only", True), + (KeyManagementSystem.CUSTOM, "write_only", False), + (None, "read_only", False), ], ) -def test_secret_manager_is_native_only_when_the_rules_select_rust_for_its_system( +def test_secret_manager_is_native_only_when_a_readable_system_is_configured( monkeypatch: pytest.MonkeyPatch, system: KeyManagementSystem | None, access_mode: str, - rules: catalog.Rules, native: bool, ) -> None: monkeypatch.setattr(litellm, "secret_manager_client", _VaultSecrets({})) monkeypatch.setattr(litellm, "_key_management_system", system) monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(access_mode=access_mode)) - assert settings.secret_manager(rules=rules) == settings.SecretManager( + assert settings.secret_manager() == settings.SecretManager( readable=access_mode != "write_only", native=native ) @@ -137,7 +122,7 @@ def test_secret_manager_is_native_only_when_the_rules_select_rust_for_its_system def test_secret_manager_is_not_readable_without_a_client(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(litellm, "secret_manager_client", None) - assert settings.secret_manager(rules=_RUST_FOR_CUSTOM) == settings.SecretManager(readable=False, native=False) + assert settings.secret_manager() == settings.SecretManager(readable=False, native=False) def test_secret_manager_projects_custom_settings(monkeypatch: pytest.MonkeyPatch) -> None: