from __future__ import annotations import asyncio import logging from collections import defaultdict from typing import Any from yuxi.channel.secrets.models import ( DEFAULT_SECRET_PROVIDER_ALIAS, SecretIssue, SecretRef, SecretResolveResult, SecretSource, is_secret_ref, is_valid_exec_secret_ref_id, parse_secret_ref, secret_ref_key, validate_ref_id, ) from yuxi.channel.secrets.providers import ( EnvProvider, ExecProvider, FileProvider, SecretsProvider, ) logger = logging.getLogger(__name__) DEFAULT_PROVIDER_CONCURRENCY = 4 DEFAULT_MAX_REFS_PER_PROVIDER = 512 def _provider_key(source: SecretSource, alias: str | None = None) -> str: return f"{source.value}:{alias or DEFAULT_SECRET_PROVIDER_ALIAS}" def _parse_provider_key(key: str) -> tuple[SecretSource, str] | None: parts = key.split(":", 1) if len(parts) != 2: return None try: return SecretSource(parts[0]), parts[1] except ValueError: return None class SecretProviderResolutionError(Exception): def __init__(self, source: SecretSource, provider: str, message: str): self.source = source self.provider = provider super().__init__(message) class SecretRefResolutionError(Exception): def __init__(self, source: SecretSource, provider: str, ref_id: str, message: str): self.source = source self.provider = provider self.ref_id = ref_id super().__init__(message) class SecretResolver: def __init__( self, providers: dict[SecretSource, SecretsProvider] | dict[str, SecretsProvider] | None = None, max_provider_concurrency: int = DEFAULT_PROVIDER_CONCURRENCY, max_refs_per_provider: int = DEFAULT_MAX_REFS_PER_PROVIDER, ): self._providers: dict[str, SecretsProvider] = {} if providers: for key, provider in providers.items(): if isinstance(key, SecretSource): self._providers[_provider_key(key)] = provider else: self._providers[key] = provider else: self._providers = { _provider_key(SecretSource.ENV): EnvProvider(), _provider_key(SecretSource.FILE): FileProvider(), _provider_key(SecretSource.EXEC): ExecProvider(), } self._max_provider_concurrency = max_provider_concurrency self._max_refs_per_provider = max_refs_per_provider def register_provider( self, source: SecretSource, provider: SecretsProvider, alias: str | None = None, ) -> None: self._providers[_provider_key(source, alias)] = provider async def resolve(self, config: dict[str, Any], path: str = "") -> SecretResolveResult: result = SecretResolveResult(resolved={}, issues=[], warnings=[]) await self._walk(config, result.resolved, path, result) return result async def resolve_secret_ref_values(self, refs: list[SecretRef]) -> dict[str, str | None]: if not refs: return {} unique_refs: dict[str, SecretRef] = {} for ref in refs: ref_id = ref.ref.strip() if not ref_id: raise ValueError("Secret reference id is empty.") if ref.source == SecretSource.EXEC and not is_valid_exec_secret_ref_id(ref_id): raise ValueError( f"Invalid exec secret ref id '{ref_id}': must match /^[A-Za-z0-9][A-Za-z0-9._:/-]{{0,255}}$/" ) normalized = SecretRef(source=ref.source, ref=ref_id, provider=ref.provider) unique_refs[secret_ref_key(normalized)] = normalized grouped: dict[str, list[SecretRef]] = defaultdict(list) for ref in unique_refs.values(): provider_alias = ref.provider or DEFAULT_SECRET_PROVIDER_ALIAS group_key = _provider_key(ref.source, provider_alias) grouped[group_key].append(ref) for group_key, group_refs in grouped.items(): if len(group_refs) > self._max_refs_per_provider: raise SecretProviderResolutionError( source=group_refs[0].source, provider=group_refs[0].provider or DEFAULT_SECRET_PROVIDER_ALIAS, message=f"Provider exceeded maxRefsPerProvider ({self._max_refs_per_provider}).", ) semaphore = asyncio.Semaphore(self._max_provider_concurrency) async def resolve_group(group_refs: list[SecretRef]) -> dict[str, str | None]: async with semaphore: first = group_refs[0] provider = self._providers.get(_provider_key(first.source, first.provider)) if provider is None: raise SecretProviderResolutionError( source=first.source, provider=first.provider or DEFAULT_SECRET_PROVIDER_ALIAS, message=( f"No provider registered for source '{first.source.value}'" f" with alias '{first.provider or DEFAULT_SECRET_PROVIDER_ALIAS}'" ), ) batch_result = await provider.resolve_batch(group_refs) return {secret_ref_key(ref): batch_result.get(ref.ref) for ref in group_refs} tasks = [resolve_group(g) for g in grouped.values()] results = await asyncio.gather(*tasks, return_exceptions=True) resolved: dict[str, str | None] = {} errors: list[Exception] = [] for r in results: if isinstance(r, Exception): errors.append(r) else: resolved.update(r) if errors: if len(errors) == 1: raise errors[0] raise ExceptionGroup("Secret resolution failed for one or more providers", errors) return resolved async def resolve_secret_ref_value(self, ref: SecretRef) -> str | None: ref_id = ref.ref.strip() if not ref_id: raise ValueError("Secret reference id is empty.") normalized = SecretRef(source=ref.source, ref=ref_id, provider=ref.provider) resolved = await self.resolve_secret_ref_values([normalized]) return resolved.get(secret_ref_key(normalized)) async def _walk( self, src: Any, dst: dict, path: str, result: SecretResolveResult, ) -> None: if not isinstance(src, dict): return for key, value in src.items(): current_path = f"{path}.{key}" if path else key if is_secret_ref(value): resolved = await self._resolve_ref(value, current_path, result) dst[key] = resolved if resolved is not None else "" elif isinstance(value, dict): dst[key] = {} await self._walk(value, dst[key], current_path, result) elif isinstance(value, list): dst[key] = await self._walk_list(value, current_path, result) else: dst[key] = value async def _walk_list( self, src: list, path: str, result: SecretResolveResult, ) -> list: resolved_list = [] for i, item in enumerate(src): item_path = f"{path}[{i}]" if is_secret_ref(item): resolved = await self._resolve_ref(item, item_path, result) resolved_list.append(resolved if resolved is not None else "") elif isinstance(item, dict): item_resolved: dict = {} await self._walk(item, item_resolved, item_path, result) resolved_list.append(item_resolved) elif isinstance(item, list): resolved_list.append(await self._walk_list(item, item_path, result)) else: resolved_list.append(item) return resolved_list async def _resolve_ref( self, value: dict, path: str, result: SecretResolveResult, ) -> str | None: ref = parse_secret_ref(value) if ref is None: issue = SecretIssue(path=path, message="Invalid $secret format", severity="error") result.issues.append(issue) return None validation_error = validate_ref_id(ref.source, ref.ref) if validation_error: result.issues.append(SecretIssue(path=path, message=validation_error, ref=ref)) return None provider = self._providers.get(_provider_key(ref.source, ref.provider)) if provider is None: result.issues.append( SecretIssue( path=path, message=( f"No provider registered for source '{ref.source.value}'" f" with alias '{ref.provider or DEFAULT_SECRET_PROVIDER_ALIAS}'" ), ref=ref, ) ) return None try: resolved = await provider.resolve(ref) except asyncio.CancelledError: raise except Exception as e: result.issues.append(SecretIssue(path=path, message=f"Resolution failed: {e}", ref=ref)) return None if resolved is None: result.warnings.append( SecretIssue( path=path, message=f"Secret not found: {ref.source.value}:{ref.ref}", ref=ref, severity="warning", ) ) return resolved async def resolve_secrets( config: dict[str, Any], providers: dict[SecretSource, SecretsProvider] | None = None, ) -> SecretResolveResult: resolver = SecretResolver(providers) return await resolver.resolve(config) def create_exec_provider( command: str, args: list[str] | None = None, timeout_ms: int = 5000, max_output_bytes: int = 1024 * 1024, pass_env: list[str] | None = None, extra_env: dict[str, str] | None = None, trusted_dirs: set[str] | None = None, ) -> ExecProvider: return ExecProvider( command=command, args=args or [], trusted_dirs=trusted_dirs or set(), timeout_ms=timeout_ms, max_output_bytes=max_output_bytes, pass_env=pass_env or [], extra_env=extra_env or {}, ) def create_resolver_with_exec( command: str, args: list[str] | None = None, alias: str = "default", timeout_ms: int = 5000, max_output_bytes: int = 1024 * 1024, pass_env: list[str] | None = None, extra_env: dict[str, str] | None = None, env_allowlist: set[str] | None = None, file_base_dir: str | None = None, max_provider_concurrency: int = DEFAULT_PROVIDER_CONCURRENCY, max_refs_per_provider: int = DEFAULT_MAX_REFS_PER_PROVIDER, ) -> SecretResolver: provider = create_exec_provider( command=command, args=args, timeout_ms=timeout_ms, max_output_bytes=max_output_bytes, pass_env=pass_env, extra_env=extra_env, ) env_provider = EnvProvider(allowlist=env_allowlist) if env_allowlist else EnvProvider() file_provider = FileProvider(base_dir=file_base_dir) if file_base_dir else FileProvider() resolver = SecretResolver( max_provider_concurrency=max_provider_concurrency, max_refs_per_provider=max_refs_per_provider, ) resolver.register_provider(SecretSource.ENV, env_provider, alias=None) resolver.register_provider(SecretSource.FILE, file_provider, alias=None) resolver.register_provider(SecretSource.EXEC, provider, alias=alias) return resolver