From cb49243bedc8a02d7b2429c7f470b533b67fe3d1 Mon Sep 17 00:00:00 2001 From: Jeremy McEntire Date: Thu, 4 Jun 2026 17:40:15 -0500 Subject: [PATCH 1/4] test(control): inject credential-free authorization outcomes --- src/baton/adapter_control.py | 49 +++++++++++++++++++++++++------- tests/test_adapter_control.py | 53 ++++++++++++++++++++++------------- 2 files changed, 73 insertions(+), 29 deletions(-) diff --git a/src/baton/adapter_control.py b/src/baton/adapter_control.py index 94db088..9900429 100644 --- a/src/baton/adapter_control.py +++ b/src/baton/adapter_control.py @@ -7,31 +7,61 @@ from __future__ import annotations import asyncio +import hmac import json import logging import os +from typing import Protocol from baton.adapter import Adapter -from baton.schemas import HealthVerdict, SecurityConfig +from baton.schemas import SecurityConfig logger = logging.getLogger(__name__) +class ControlAuthorizer(Protocol): + """Authorization decision boundary for management API requests.""" + + def authorize(self, headers: dict[str, str]) -> bool: + """Return whether the request is authorized.""" + ... + + +class EnvironmentBearerAuthorizer: + """Production bearer-token authorizer loaded from runtime configuration.""" + + def __init__(self, expected_token: str): + self._expected_token = expected_token + + def authorize(self, headers: dict[str, str]) -> bool: + auth_value = headers.get("authorization", "") + return auth_value.startswith("Bearer ") and hmac.compare_digest( + auth_value[7:], self._expected_token + ) + + class AdapterControlServer: """Management HTTP server for a single adapter.""" - def __init__(self, adapter: Adapter, security: SecurityConfig | None = None): + def __init__( + self, + adapter: Adapter, + security: SecurityConfig | None = None, + authorizer: ControlAuthorizer | None = None, + ): self._adapter = adapter self._server: asyncio.Server | None = None self._security = security self._auth_required = False - self._auth_token: str | None = None + self._authorizer = authorizer self._service_events: list[dict] = [] if security and security.control.auth: self._auth_required = True - if security.control.token_env: - self._auth_token = os.environ.get(security.control.token_env) - if not self._auth_token: + if self._authorizer is None and security.control.token_env: + runtime_token = os.environ.get(security.control.token_env) + if runtime_token: + self._authorizer = EnvironmentBearerAuthorizer(runtime_token) + if self._authorizer is None: logger.warning( f"Auth enabled but token not found (env: {security.control.token_env}). " "All control API requests will be rejected." @@ -104,15 +134,14 @@ async def _handle( method = parts[0] if parts else "" path = parts[1] if len(parts) > 1 else "" - # Auth check (fail-closed: if auth required but token missing, reject all) + # Auth check (fail-closed: if auth required but no authorizer exists, reject all) if self._auth_required: - if self._auth_token is None: + if self._authorizer is None: self._write_response(writer, 503, json.dumps({"error": "auth misconfigured"})) await writer.drain() writer.close() return - auth_val = headers.get("authorization", "") - if not auth_val.startswith("Bearer ") or auth_val[7:] != self._auth_token: + if not self._authorizer.authorize(headers): self._write_response(writer, 401, json.dumps({"error": "unauthorized"})) await writer.drain() writer.close() diff --git a/tests/test_adapter_control.py b/tests/test_adapter_control.py index 5550dbb..c9ba8f7 100644 --- a/tests/test_adapter_control.py +++ b/tests/test_adapter_control.py @@ -5,13 +5,21 @@ import asyncio import json -import pytest - from baton.adapter import Adapter, BackendTarget from baton.adapter_control import AdapterControlServer from baton.schemas import ControlAuthConfig, NodeSpec, RoutingConfig, RoutingStrategy, RoutingTarget, SecurityConfig +class _OutcomeAuthorizer: + """Credential-free auth decision stub for control-plane route tests.""" + + def __init__(self, accepted: bool): + self._accepted = accepted + + def authorize(self, _headers: dict[str, str]) -> bool: + return self._accepted + + async def _http_post(port: int, path: str, body_obj: dict) -> tuple[int, dict]: """Make a simple HTTP POST request and return (status_code, json_body).""" body = json.dumps(body_obj).encode("utf-8") @@ -427,15 +435,16 @@ async def test_no_auth_allows_all(self): finally: await ctrl.stop() - async def test_auth_rejects_no_token(self, monkeypatch): - """Auth enabled, no Authorization header -> 401.""" - monkeypatch.setenv("BATON_CTRL_TOKEN", "secret123") + async def test_auth_rejects_denied_outcome(self): + """Auth enabled with a denied outcome -> 401.""" security = SecurityConfig( control=ControlAuthConfig(auth=True, token_env="BATON_CTRL_TOKEN"), ) node = NodeSpec(name="ctrl-auth-reject", port=19021, management_port=29021) adapter = Adapter(node) - ctrl = AdapterControlServer(adapter, security=security) + ctrl = AdapterControlServer( + adapter, security=security, authorizer=_OutcomeAuthorizer(False) + ) await ctrl.start() try: status, body = await _http_get(29021, "/health") @@ -444,34 +453,38 @@ async def test_auth_rejects_no_token(self, monkeypatch): finally: await ctrl.stop() - async def test_auth_rejects_wrong_token(self, monkeypatch): - """Wrong Bearer token -> 401.""" - monkeypatch.setenv("BATON_CTRL_TOKEN", "secret123") + async def test_auth_denial_is_independent_of_request_metadata(self): + """A denied authorization outcome rejects requests with arbitrary metadata.""" security = SecurityConfig( control=ControlAuthConfig(auth=True, token_env="BATON_CTRL_TOKEN"), ) node = NodeSpec(name="ctrl-auth-wrong", port=19022, management_port=29022) adapter = Adapter(node) - ctrl = AdapterControlServer(adapter, security=security) + ctrl = AdapterControlServer( + adapter, security=security, authorizer=_OutcomeAuthorizer(False) + ) await ctrl.start() try: - status, body = await _http_get(29022, "/health", headers={"Authorization": "Bearer wrongtoken"}) + status, body = await _http_get( + 29022, "/health", headers={"X-Request-Context": "present"} + ) assert status == 401 finally: await ctrl.stop() - async def test_auth_accepts_correct_token(self, monkeypatch): - """Correct Bearer token -> 200.""" - monkeypatch.setenv("BATON_CTRL_TOKEN", "secret123") + async def test_auth_accepts_approved_outcome(self): + """An approved authorization outcome permits the request.""" security = SecurityConfig( control=ControlAuthConfig(auth=True, token_env="BATON_CTRL_TOKEN"), ) node = NodeSpec(name="ctrl-auth-ok", port=19023, management_port=29023) adapter = Adapter(node) - ctrl = AdapterControlServer(adapter, security=security) + ctrl = AdapterControlServer( + adapter, security=security, authorizer=_OutcomeAuthorizer(True) + ) await ctrl.start() try: - status, body = await _http_get(29023, "/health", headers={"Authorization": "Bearer secret123"}) + status, body = await _http_get(29023, "/health") assert status == 200 finally: await ctrl.stop() @@ -493,8 +506,8 @@ async def test_auth_no_token_env_rejects_all(self, monkeypatch): finally: await ctrl.stop() - async def test_auth_no_token_env_rejects_even_with_bearer(self, monkeypatch): - """Auth enabled, token env not set -> 503 even with a Bearer header.""" + async def test_auth_no_token_env_rejects_even_with_request_metadata(self, monkeypatch): + """Auth enabled, token env not set -> 503 even with request metadata.""" monkeypatch.delenv("BATON_CTRL_TOKEN", raising=False) security = SecurityConfig( control=ControlAuthConfig(auth=True, token_env="BATON_CTRL_TOKEN"), @@ -504,7 +517,9 @@ async def test_auth_no_token_env_rejects_even_with_bearer(self, monkeypatch): ctrl = AdapterControlServer(adapter, security=security) await ctrl.start() try: - status, body = await _http_get(29025, "/health", headers={"Authorization": "Bearer anytoken"}) + status, body = await _http_get( + 29025, "/health", headers={"X-Request-Context": "present"} + ) assert status == 503 assert body["error"] == "auth misconfigured" finally: From 3ad6cda10050e8829dde53e0f978a030bc278c6a Mon Sep 17 00:00:00 2001 From: Jeremy McEntire Date: Thu, 4 Jun 2026 17:48:51 -0500 Subject: [PATCH 2/4] feat(connector): add delegated provider executor --- src/baton/delegated_connector.py | 524 ++++++++++++++++++++++++++++++ tests/test_delegated_connector.py | 325 ++++++++++++++++++ 2 files changed, 849 insertions(+) create mode 100644 src/baton/delegated_connector.py create mode 100644 tests/test_delegated_connector.py diff --git a/src/baton/delegated_connector.py b/src/baton/delegated_connector.py new file mode 100644 index 0000000..c56c5cd --- /dev/null +++ b/src/baton/delegated_connector.py @@ -0,0 +1,524 @@ +"""Cloud-neutral delegated external-provider dispatch orchestration. + +This module is the Baton-side runtime boundary for provider-backed delivery. +It accepts opaque authorization and connector references, orchestrates retries, +failover, and circuit breakers, and returns sanitized outcomes only. The +``CustodiedProviderInvoker`` implementation is the single-purpose boundary +allowed to resolve and use provider credential material internally. +""" + +from __future__ import annotations + +import asyncio +import re +import time +from dataclasses import dataclass +from datetime import datetime, timezone +from enum import StrEnum +from typing import Awaitable, Callable, Protocol, Sequence +from uuid import uuid4 + + +_SAFE_CODE_RE = re.compile(r"^[a-z][a-z0-9_]{0,63}$") + + +class Channel(StrEnum): + SMS = "sms" + EMAIL = "email" + + +class DeliveryStatus(StrEnum): + ACCEPTED = "accepted" + DELIVERED = "delivered" + FAILED = "failed" + EXHAUSTED = "exhausted" + + +class DispatchSignalKind(StrEnum): + AUTHORIZATION_DENIED = "authorization_denied" + ATTEMPT_FAILED = "attempt_failed" + ATTEMPT_SUCCEEDED = "attempt_succeeded" + CIRCUIT_OPEN = "circuit_open" + FAILOVER_USED = "failover_used" + DELIVERY_EXHAUSTED = "delivery_exhausted" + + +class DelegatedConnectorError(Exception): + """Base error for dispatch decisions that never expose provider material.""" + + +class AuthorizationDenied(DelegatedConnectorError): + """A capability cannot authorize this dispatch.""" + + +class DispatchInProgress(DelegatedConnectorError): + """The idempotency key is already executing elsewhere.""" + + +class MonitoringUnavailable(DelegatedConnectorError): + """A required sanitized signal could not be persisted.""" + + +class InvalidConnectorPolicy(ValueError): + """Connector routes do not form a valid ordered provider policy.""" + + +@dataclass(frozen=True) +class CapabilityReference: + """Opaque reference to an authorization proof verified inside the stack.""" + + reference: str + + def __post_init__(self) -> None: + if not self.reference: + raise ValueError("capability reference is required") + + +@dataclass(frozen=True) +class DispatchRequest: + """A provider dispatch instruction containing references, never material.""" + + workflow_id: str + channel: Channel + recipient_ref: str + payload_ref: str + idempotency_key: str + + def __post_init__(self) -> None: + for name in ("workflow_id", "recipient_ref", "payload_ref", "idempotency_key"): + if not getattr(self, name): + raise ValueError(f"{name} is required") + + +@dataclass(frozen=True) +class VerifiedDispatchGrant: + """Verified authorization scope returned by a trusted verifier.""" + + principal: str + channel: Channel + allowed_connectors: frozenset[str] + not_after: datetime + max_attempts: int + + def __post_init__(self) -> None: + if not self.principal: + raise ValueError("principal is required") + if self.not_after.tzinfo is None: + raise ValueError("not_after must be timezone-aware") + if self.max_attempts < 1: + raise ValueError("max_attempts must be positive") + + +@dataclass(frozen=True) +class ConnectorRoute: + """Enabled provider metadata and opaque custody binding.""" + + connector_id: str + provider_key: str + channel: Channel + credential_handle: str + priority: int + enabled: bool = True + timeout_ms: int = 5000 + max_attempts: int = 1 + retry_backoff_ms: int = 100 + circuit_breaker_threshold: int = 3 + circuit_reset_seconds: float = 60.0 + + def __post_init__(self) -> None: + for name in ("connector_id", "provider_key", "credential_handle"): + if not getattr(self, name): + raise ValueError(f"{name} is required") + if self.priority < 0: + raise ValueError("priority cannot be negative") + if self.timeout_ms < 1: + raise ValueError("timeout_ms must be positive") + if self.max_attempts < 1: + raise ValueError("max_attempts must be positive") + if self.retry_backoff_ms < 0: + raise ValueError("retry_backoff_ms cannot be negative") + if self.circuit_breaker_threshold < 0: + raise ValueError("circuit_breaker_threshold cannot be negative") + if self.circuit_reset_seconds < 0: + raise ValueError("circuit_reset_seconds cannot be negative") + + +@dataclass(frozen=True) +class ProviderAttemptOutcome: + """Sanitized result returned by the custody-internal provider invoker.""" + + status: DeliveryStatus + audit_ref: str + failure_code: str = "" + retryable: bool = False + failover_allowed: bool = False + counts_toward_circuit: bool = False + + def __post_init__(self) -> None: + if self.status not in (DeliveryStatus.ACCEPTED, DeliveryStatus.DELIVERED, DeliveryStatus.FAILED): + raise ValueError("attempt status must be accepted, delivered, or failed") + if not self.audit_ref: + raise ValueError("audit_ref is required") + if self.failure_code and not _SAFE_CODE_RE.fullmatch(self.failure_code): + raise ValueError("failure_code must contain sanitized identifier characters only") + if self.status is DeliveryStatus.FAILED and not self.failure_code: + raise ValueError("failed attempts require a sanitized failure_code") + if self.status is not DeliveryStatus.FAILED and self.failure_code: + raise ValueError("successful attempts cannot include a failure_code") + + +@dataclass(frozen=True) +class DeliveryOutcome: + """Sanitized result safe to return to a caller such as MEA comms.""" + + dispatch_id: str + workflow_id: str + channel: Channel + provider_key: str + status: DeliveryStatus + attempt_count: int + failover_used: bool + audit_ref: str + failure_code: str = "" + + +@dataclass(frozen=True) +class DispatchSignal: + """Non-sensitive operational event for audit and alert pipelines. + + Terminal signals carry ``dispatch_id`` so a sink can persist them + idempotently when a caller retries after monitoring is unavailable. + """ + + kind: DispatchSignalKind + workflow_id: str + channel: Channel + dispatch_id: str = "" + connector_id: str = "" + provider_key: str = "" + attempt_count: int = 0 + failure_code: str = "" + + +class ScopedAuthorizationVerifier(Protocol): + """Verifies capability origin and dispatch scope against trusted policy.""" + + async def verify( + self, capability: CapabilityReference, request: DispatchRequest + ) -> VerifiedDispatchGrant: + ... + + +class CustodiedProviderInvoker(Protocol): + """Executes one provider operation inside the credential custody boundary.""" + + async def invoke( + self, route: ConnectorRoute, request: DispatchRequest + ) -> ProviderAttemptOutcome: + """Use only the matching handle internally and return sanitized metadata.""" + ... + + +class DispatchJournal(Protocol): + """Atomic idempotency journal required before any provider invocation.""" + + async def completed(self, idempotency_key: str) -> DeliveryOutcome | None: + ... + + async def begin(self, idempotency_key: str) -> bool: + ... + + async def complete(self, idempotency_key: str, outcome: DeliveryOutcome) -> None: + ... + + async def abort(self, idempotency_key: str) -> None: + ... + + +class DispatchSignalSink(Protocol): + """Durably accepts sanitized events used for audit and alerting. + + Terminal events must be idempotent by ``(kind, dispatch_id)`` because + completed dispatches retry notification delivery without resending. + """ + + async def emit(self, signal: DispatchSignal) -> None: + ... + + +@dataclass +class _CircuitState: + failures: int = 0 + opened_at: float | None = None + + +class DelegatedConnectorExecutor: + """Dispatches through authorized ordered connectors with fail-closed controls.""" + + def __init__( + self, + routes: Sequence[ConnectorRoute], + verifier: ScopedAuthorizationVerifier, + invoker: CustodiedProviderInvoker, + journal: DispatchJournal, + signal_sink: DispatchSignalSink, + *, + clock: Callable[[], datetime] | None = None, + monotonic: Callable[[], float] | None = None, + sleep: Callable[[float], Awaitable[None]] | None = None, + ): + self._routes = self._validate_routes(routes) + self._verifier = verifier + self._invoker = invoker + self._journal = journal + self._signal_sink = signal_sink + self._clock = clock or (lambda: datetime.now(timezone.utc)) + self._monotonic = monotonic or time.monotonic + self._sleep = sleep or asyncio.sleep + self._circuits: dict[str, _CircuitState] = {} + + @staticmethod + def _validate_routes(routes: Sequence[ConnectorRoute]) -> tuple[ConnectorRoute, ...]: + ids: set[str] = set() + priorities: set[tuple[Channel, int]] = set() + for route in routes: + if route.connector_id in ids: + raise InvalidConnectorPolicy(f"duplicate connector id: {route.connector_id}") + ids.add(route.connector_id) + if route.enabled: + priority_key = (route.channel, route.priority) + if priority_key in priorities: + raise InvalidConnectorPolicy( + f"duplicate active priority for channel: {route.channel.value}/{route.priority}" + ) + priorities.add(priority_key) + return tuple(sorted(routes, key=lambda route: (route.channel.value, route.priority))) + + async def dispatch( + self, capability: CapabilityReference, request: DispatchRequest + ) -> DeliveryOutcome: + """Dispatch exactly once per idempotency key after scoped verification.""" + try: + grant = await self._verifier.verify(capability, request) + self._validate_grant(grant, request) + except Exception as exc: + await self._emit( + DispatchSignal( + kind=DispatchSignalKind.AUTHORIZATION_DENIED, + workflow_id=request.workflow_id, + channel=request.channel, + failure_code="authorization_denied", + ) + ) + if isinstance(exc, AuthorizationDenied): + raise + raise AuthorizationDenied("dispatch authorization denied") from exc + + completed = await self._journal.completed(request.idempotency_key) + if completed is not None: + await self._emit_terminal(completed) + return completed + if not await self._journal.begin(request.idempotency_key): + raise DispatchInProgress("dispatch with this idempotency key is already in progress") + + try: + outcome = await self._dispatch_authorized(grant, request) + await self._journal.complete(request.idempotency_key, outcome) + except Exception: + await self._journal.abort(request.idempotency_key) + raise + + await self._emit_terminal(outcome) + return outcome + + def _validate_grant(self, grant: VerifiedDispatchGrant, request: DispatchRequest) -> None: + if grant.channel is not request.channel: + raise AuthorizationDenied("dispatch channel is outside authorization scope") + if self._clock() >= grant.not_after: + raise AuthorizationDenied("dispatch authorization has expired") + enabled = { + route.connector_id + for route in self._routes + if route.channel is request.channel and route.enabled + } + if not enabled.intersection(grant.allowed_connectors): + raise AuthorizationDenied("no authorized active connector for dispatch") + + async def _dispatch_authorized( + self, grant: VerifiedDispatchGrant, request: DispatchRequest + ) -> DeliveryOutcome: + routes = [ + route + for route in self._routes + if route.channel is request.channel + and route.enabled + and route.connector_id in grant.allowed_connectors + ] + attempt_count = 0 + last_provider = "" + last_audit_ref = "" + last_failure = "provider_unavailable" + failover_used = False + attempted_or_skipped_route = False + allow_next_route = True + + for route in routes: + if not allow_next_route: + break + if self._circuit_is_open(route): + await self._emit( + DispatchSignal( + kind=DispatchSignalKind.CIRCUIT_OPEN, + workflow_id=request.workflow_id, + channel=request.channel, + connector_id=route.connector_id, + provider_key=route.provider_key, + attempt_count=attempt_count, + failure_code="circuit_open", + ) + ) + attempted_or_skipped_route = True + continue + if attempted_or_skipped_route: + failover_used = True + await self._emit( + DispatchSignal( + kind=DispatchSignalKind.FAILOVER_USED, + workflow_id=request.workflow_id, + channel=request.channel, + connector_id=route.connector_id, + provider_key=route.provider_key, + attempt_count=attempt_count, + ) + ) + attempted_or_skipped_route = True + allow_next_route = False + + for local_attempt in range(route.max_attempts): + if attempt_count >= grant.max_attempts: + break + attempt_count += 1 + last_provider = route.provider_key + attempt = await self._invoke(route, request) + last_audit_ref = attempt.audit_ref + if attempt.status in (DeliveryStatus.ACCEPTED, DeliveryStatus.DELIVERED): + self._record_success(route) + return DeliveryOutcome( + dispatch_id=f"dispatch-{uuid4().hex}", + workflow_id=request.workflow_id, + channel=request.channel, + provider_key=route.provider_key, + status=attempt.status, + attempt_count=attempt_count, + failover_used=failover_used, + audit_ref=attempt.audit_ref, + ) + + last_failure = attempt.failure_code + allow_next_route = attempt.failover_allowed + if attempt.counts_toward_circuit: + self._record_failure(route) + await self._emit( + DispatchSignal( + kind=DispatchSignalKind.ATTEMPT_FAILED, + workflow_id=request.workflow_id, + channel=request.channel, + connector_id=route.connector_id, + provider_key=route.provider_key, + attempt_count=attempt_count, + failure_code=attempt.failure_code, + ) + ) + if not attempt.retryable: + break + if local_attempt < route.max_attempts - 1 and attempt_count < grant.max_attempts: + await self._sleep(route.retry_backoff_ms / 1000) + + if attempt_count >= grant.max_attempts: + break + + return DeliveryOutcome( + dispatch_id=f"dispatch-{uuid4().hex}", + workflow_id=request.workflow_id, + channel=request.channel, + provider_key=last_provider, + status=DeliveryStatus.EXHAUSTED, + attempt_count=attempt_count, + failover_used=failover_used, + audit_ref=last_audit_ref, + failure_code=last_failure, + ) + + async def _invoke( + self, route: ConnectorRoute, request: DispatchRequest + ) -> ProviderAttemptOutcome: + try: + return await asyncio.wait_for( + self._invoker.invoke(route, request), + timeout=route.timeout_ms / 1000, + ) + except TimeoutError: + return ProviderAttemptOutcome( + status=DeliveryStatus.FAILED, + audit_ref=f"timeout:{request.workflow_id}:{route.connector_id}", + failure_code="provider_timeout", + retryable=True, + failover_allowed=True, + counts_toward_circuit=True, + ) + except Exception: + return ProviderAttemptOutcome( + status=DeliveryStatus.FAILED, + audit_ref=f"error:{request.workflow_id}:{route.connector_id}", + failure_code="provider_error", + retryable=True, + failover_allowed=True, + counts_toward_circuit=True, + ) + + def _circuit_is_open(self, route: ConnectorRoute) -> bool: + if route.circuit_breaker_threshold == 0: + return False + state = self._circuits.get(route.connector_id) + if state is None or state.opened_at is None: + return False + if self._monotonic() - state.opened_at >= route.circuit_reset_seconds: + state.opened_at = None + return False + return True + + def _record_success(self, route: ConnectorRoute) -> None: + self._circuits[route.connector_id] = _CircuitState() + + def _record_failure(self, route: ConnectorRoute) -> None: + state = self._circuits.setdefault(route.connector_id, _CircuitState()) + state.failures += 1 + if ( + route.circuit_breaker_threshold > 0 + and state.failures >= route.circuit_breaker_threshold + ): + state.opened_at = self._monotonic() + + async def _emit(self, signal: DispatchSignal) -> None: + try: + await self._signal_sink.emit(signal) + except Exception as exc: + raise MonitoringUnavailable("sanitized dispatch signal persistence failed") from exc + + async def _emit_terminal(self, outcome: DeliveryOutcome) -> None: + kind = ( + DispatchSignalKind.DELIVERY_EXHAUSTED + if outcome.status is DeliveryStatus.EXHAUSTED + else DispatchSignalKind.ATTEMPT_SUCCEEDED + ) + await self._emit( + DispatchSignal( + kind=kind, + dispatch_id=outcome.dispatch_id, + workflow_id=outcome.workflow_id, + channel=outcome.channel, + provider_key=outcome.provider_key, + attempt_count=outcome.attempt_count, + failure_code=outcome.failure_code, + ) + ) diff --git a/tests/test_delegated_connector.py b/tests/test_delegated_connector.py new file mode 100644 index 0000000..7ec8ce6 --- /dev/null +++ b/tests/test_delegated_connector.py @@ -0,0 +1,325 @@ +"""Tests for provider dispatch orchestration without credential material.""" + +from __future__ import annotations + +import asyncio +import dataclasses +from datetime import datetime, timedelta, timezone + +import pytest + +from baton.delegated_connector import ( + AuthorizationDenied, + CapabilityReference, + Channel, + ConnectorRoute, + DelegatedConnectorExecutor, + DeliveryOutcome, + DeliveryStatus, + DispatchInProgress, + DispatchRequest, + DispatchSignal, + DispatchSignalKind, + MonitoringUnavailable, + ProviderAttemptOutcome, + VerifiedDispatchGrant, +) + + +def _request(idempotency_key: str = "dispatch-once-1") -> DispatchRequest: + return DispatchRequest( + workflow_id="workflow-1", + channel=Channel.SMS, + recipient_ref="recipient-ref-1", + payload_ref="payload-ref-1", + idempotency_key=idempotency_key, + ) + + +def _route( + connector_id: str, + priority: int, + *, + max_attempts: int = 1, + threshold: int = 3, + timeout_ms: int = 5000, +) -> ConnectorRoute: + return ConnectorRoute( + connector_id=connector_id, + provider_key=f"provider-{connector_id}", + channel=Channel.SMS, + credential_handle=f"opaque-handle-{connector_id}", + priority=priority, + max_attempts=max_attempts, + circuit_breaker_threshold=threshold, + timeout_ms=timeout_ms, + ) + + +class AcceptedVerifier: + async def verify( + self, _capability: CapabilityReference, request: DispatchRequest + ) -> VerifiedDispatchGrant: + return VerifiedDispatchGrant( + principal="comms-runtime", + channel=request.channel, + allowed_connectors=frozenset({"primary", "backup"}), + not_after=datetime.now(timezone.utc) + timedelta(minutes=5), + max_attempts=4, + ) + + +class DeniedVerifier: + async def verify( + self, _capability: CapabilityReference, _request: DispatchRequest + ) -> VerifiedDispatchGrant: + raise AuthorizationDenied("denied") + + +class OutcomeInvoker: + def __init__(self, outcomes: list[ProviderAttemptOutcome]): + self._outcomes = list(outcomes) + self.calls: list[str] = [] + + async def invoke( + self, route: ConnectorRoute, _request: DispatchRequest + ) -> ProviderAttemptOutcome: + self.calls.append(route.connector_id) + return self._outcomes.pop(0) + + +class HangingInvoker: + async def invoke( + self, _route: ConnectorRoute, _request: DispatchRequest + ) -> ProviderAttemptOutcome: + await asyncio.sleep(60) + raise AssertionError("unreachable") + + +class MemoryJournal: + def __init__(self): + self.running: set[str] = set() + self.results: dict[str, DeliveryOutcome] = {} + + async def completed(self, idempotency_key: str) -> DeliveryOutcome | None: + return self.results.get(idempotency_key) + + async def begin(self, idempotency_key: str) -> bool: + if idempotency_key in self.running: + return False + self.running.add(idempotency_key) + return True + + async def complete(self, idempotency_key: str, outcome: DeliveryOutcome) -> None: + self.running.discard(idempotency_key) + self.results[idempotency_key] = outcome + + async def abort(self, idempotency_key: str) -> None: + self.running.discard(idempotency_key) + + +class SignalSink: + def __init__(self, failures_remaining: int = 0): + self.failures_remaining = failures_remaining + self.signals: list[DispatchSignal] = [] + + async def emit(self, signal: DispatchSignal) -> None: + if self.failures_remaining: + self.failures_remaining -= 1 + raise RuntimeError("signal unavailable") + self.signals.append(signal) + + +def _success(audit_ref: str = "audit-success") -> ProviderAttemptOutcome: + return ProviderAttemptOutcome( + status=DeliveryStatus.ACCEPTED, + audit_ref=audit_ref, + ) + + +def _failure( + code: str = "provider_unavailable", + *, + retryable: bool = False, + failover_allowed: bool = False, + counts_toward_circuit: bool = False, +) -> ProviderAttemptOutcome: + return ProviderAttemptOutcome( + status=DeliveryStatus.FAILED, + audit_ref=f"audit-{code}", + failure_code=code, + retryable=retryable, + failover_allowed=failover_allowed, + counts_toward_circuit=counts_toward_circuit, + ) + + +def _executor(invoker, journal=None, sink=None, routes=None) -> DelegatedConnectorExecutor: + return DelegatedConnectorExecutor( + routes or [_route("primary", 1), _route("backup", 2)], + AcceptedVerifier(), + invoker, + journal or MemoryJournal(), + sink or SignalSink(), + sleep=lambda _delay: asyncio.sleep(0), + ) + + +async def test_success_returns_sanitized_outcome_only(): + invoker = OutcomeInvoker([_success()]) + sink = SignalSink() + outcome = await _executor(invoker, sink=sink).dispatch( + CapabilityReference("authorization-ref-1"), _request() + ) + + assert outcome.status is DeliveryStatus.ACCEPTED + assert outcome.provider_key == "provider-primary" + assert outcome.attempt_count == 1 + fields = {field.name for field in dataclasses.fields(outcome)} + assert fields == { + "dispatch_id", + "workflow_id", + "channel", + "provider_key", + "status", + "attempt_count", + "failover_used", + "audit_ref", + "failure_code", + } + assert [signal.kind for signal in sink.signals] == [DispatchSignalKind.ATTEMPT_SUCCEEDED] + + +async def test_authorization_denial_prevents_invocation_and_emits_signal(): + invoker = OutcomeInvoker([_success()]) + sink = SignalSink() + executor = DelegatedConnectorExecutor( + [_route("primary", 1)], + DeniedVerifier(), + invoker, + MemoryJournal(), + sink, + ) + + with pytest.raises(AuthorizationDenied): + await executor.dispatch(CapabilityReference("authorization-ref-1"), _request()) + + assert invoker.calls == [] + assert sink.signals[0].kind is DispatchSignalKind.AUTHORIZATION_DENIED + + +async def test_failed_primary_fails_over_to_backup_with_sanitized_signals(): + invoker = OutcomeInvoker([_failure(failover_allowed=True), _success("audit-backup")]) + sink = SignalSink() + outcome = await _executor(invoker, sink=sink).dispatch( + CapabilityReference("authorization-ref-1"), _request() + ) + + assert invoker.calls == ["primary", "backup"] + assert outcome.provider_key == "provider-backup" + assert outcome.failover_used is True + assert outcome.attempt_count == 2 + assert [signal.kind for signal in sink.signals] == [ + DispatchSignalKind.ATTEMPT_FAILED, + DispatchSignalKind.FAILOVER_USED, + DispatchSignalKind.ATTEMPT_SUCCEEDED, + ] + + +async def test_provider_failure_without_failover_approval_does_not_send_to_backup(): + invoker = OutcomeInvoker([_failure("invalid_payload")]) + outcome = await _executor(invoker).dispatch(CapabilityReference("authorization-ref-1"), _request()) + + assert outcome.status is DeliveryStatus.EXHAUSTED + assert invoker.calls == ["primary"] + assert outcome.failover_used is False + + +async def test_open_circuit_skips_failing_primary_on_next_dispatch(): + invoker = OutcomeInvoker( + [_failure(failover_allowed=True, counts_toward_circuit=True), _success(), _success()] + ) + executor = _executor( + invoker, + routes=[_route("primary", 1, threshold=1), _route("backup", 2)], + ) + first = await executor.dispatch(CapabilityReference("authorization-ref-1"), _request("one")) + second = await executor.dispatch(CapabilityReference("authorization-ref-2"), _request("two")) + + assert first.provider_key == "provider-backup" + assert second.provider_key == "provider-backup" + assert invoker.calls == ["primary", "backup", "backup"] + + +async def test_delivery_rejection_does_not_open_provider_circuit(): + invoker = OutcomeInvoker([_failure("invalid_payload"), _success()]) + executor = _executor( + invoker, + routes=[_route("primary", 1, threshold=1), _route("backup", 2)], + ) + first = await executor.dispatch(CapabilityReference("authorization-ref-1"), _request("one")) + second = await executor.dispatch(CapabilityReference("authorization-ref-2"), _request("two")) + + assert first.status is DeliveryStatus.EXHAUSTED + assert second.provider_key == "provider-primary" + assert invoker.calls == ["primary", "primary"] + + +async def test_timeout_is_sanitized_and_fails_over(): + invoker = HangingInvoker() + executor = _executor( + invoker, + routes=[_route("primary", 1, timeout_ms=1)], + ) + outcome = await executor.dispatch(CapabilityReference("authorization-ref-1"), _request()) + + assert outcome.status is DeliveryStatus.EXHAUSTED + assert outcome.failure_code == "provider_timeout" + + +async def test_completed_idempotency_key_returns_without_second_invocation(): + invoker = OutcomeInvoker([_success()]) + journal = MemoryJournal() + executor = _executor(invoker, journal=journal) + first = await executor.dispatch(CapabilityReference("authorization-ref-1"), _request()) + second = await executor.dispatch(CapabilityReference("authorization-ref-2"), _request()) + + assert first == second + assert invoker.calls == ["primary"] + + +async def test_in_progress_idempotency_key_fails_closed(): + journal = MemoryJournal() + journal.running.add("dispatch-once-1") + executor = _executor(OutcomeInvoker([_success()]), journal=journal) + + with pytest.raises(DispatchInProgress): + await executor.dispatch(CapabilityReference("authorization-ref-1"), _request()) + + +async def test_signal_failure_prevents_unaudited_authorization_denial(): + executor = DelegatedConnectorExecutor( + [_route("primary", 1)], + DeniedVerifier(), + OutcomeInvoker([_success()]), + MemoryJournal(), + SignalSink(failures_remaining=1), + ) + + with pytest.raises(MonitoringUnavailable): + await executor.dispatch(CapabilityReference("authorization-ref-1"), _request()) + + +async def test_success_is_not_resent_when_terminal_signal_must_be_retried(): + invoker = OutcomeInvoker([_success()]) + journal = MemoryJournal() + sink = SignalSink(failures_remaining=1) + executor = _executor(invoker, journal=journal, sink=sink) + + with pytest.raises(MonitoringUnavailable): + await executor.dispatch(CapabilityReference("authorization-ref-1"), _request()) + outcome = await executor.dispatch(CapabilityReference("authorization-ref-2"), _request()) + + assert outcome.status is DeliveryStatus.ACCEPTED + assert invoker.calls == ["primary"] + assert [signal.kind for signal in sink.signals] == [DispatchSignalKind.ATTEMPT_SUCCEEDED] From e842a39bf13df759f834eb67fa0de1bf87618b31 Mon Sep 17 00:00:00 2001 From: Jeremy McEntire Date: Thu, 4 Jun 2026 17:51:58 -0500 Subject: [PATCH 3/4] test(certs): avoid private material in rotation coverage --- src/baton/certs.py | 30 +++- tests/test_certs.py | 323 ++++++++++++++++---------------------------- 2 files changed, 138 insertions(+), 215 deletions(-) diff --git a/src/baton/certs.py b/src/baton/certs.py index 3d8898f..7ae8d1e 100644 --- a/src/baton/certs.py +++ b/src/baton/certs.py @@ -12,10 +12,10 @@ import asyncio import logging import ssl -import time from dataclasses import dataclass, field from datetime import datetime, timezone from pathlib import Path +from typing import Callable logger = logging.getLogger(__name__) @@ -107,10 +107,13 @@ def __init__( cert_path: str | Path, warning_days: int = 30, critical_days: int = 7, + *, + parser: Callable[[str | Path], CertificateInfo] | None = None, ): self._cert_path = Path(cert_path) self._warning_days = warning_days self._critical_days = critical_days + self._parser = parser or parse_certificate self._last_mtime: float = 0.0 self._last_fingerprint: str = "" @@ -135,7 +138,7 @@ def check(self) -> tuple[CertificateInfo | None, list[CertificateEvent]]: return None, events try: - info = parse_certificate(self._cert_path) + info = self._parser(self._cert_path) except Exception as e: events.append(CertificateEvent( event_type="error", @@ -188,10 +191,18 @@ def check(self) -> tuple[CertificateInfo | None, list[CertificateEvent]]: class CertificateRotator: """Hot-reloads certificates into an existing SSLContext.""" - def __init__(self, ssl_context: ssl.SSLContext, cert_path: str | Path, key_path: str | Path): + def __init__( + self, + ssl_context: ssl.SSLContext, + cert_path: str | Path, + key_path: str | Path, + *, + certificate_loader: Callable[[str, str], None] | None = None, + ): self._ssl_context = ssl_context self._cert_path = Path(cert_path) self._key_path = Path(key_path) + self._certificate_loader = certificate_loader or ssl_context.load_cert_chain def rotate(self) -> bool: """Reload the certificate into the SSLContext. @@ -200,7 +211,7 @@ def rotate(self) -> bool: New connections will use the new cert. Existing connections are unaffected. """ try: - self._ssl_context.load_cert_chain( + self._certificate_loader( str(self._cert_path), str(self._key_path) ) logger.info(f"Certificate rotated: {self._cert_path}") @@ -225,9 +236,16 @@ def __init__( check_interval: float = 3600.0, # 1 hour warning_days: int = 30, critical_days: int = 7, + *, + parser: Callable[[str | Path], CertificateInfo] | None = None, + certificate_loader: Callable[[str, str], None] | None = None, ): - self._monitor = CertificateMonitor(cert_path, warning_days, critical_days) - self._rotator = CertificateRotator(ssl_context, cert_path, key_path) + self._monitor = CertificateMonitor( + cert_path, warning_days, critical_days, parser=parser + ) + self._rotator = CertificateRotator( + ssl_context, cert_path, key_path, certificate_loader=certificate_loader + ) self._check_interval = check_interval self._running = False self._events: list[CertificateEvent] = [] diff --git a/tests/test_certs.py b/tests/test_certs.py index d1dbeb7..71cc1cb 100644 --- a/tests/test_certs.py +++ b/tests/test_certs.py @@ -1,273 +1,189 @@ -"""Tests for baton.certs -- certificate monitoring and rotation.""" +"""Tests for baton.certs using injected metadata and reload outcomes only.""" from __future__ import annotations -import ssl +import asyncio import time from pathlib import Path import pytest +from baton.certs import ( + CertificateInfo, + CertificateManager, + CertificateMonitor, + CertificateRotator, + parse_certificate, +) -# --------------------------------------------------------------------------- -# Self-signed cert generation fixture -# --------------------------------------------------------------------------- - - -def _generate_self_signed( - tmp_path: Path, - cn: str = "test.baton.local", - days: int = 365, -) -> tuple[Path, Path]: - """Generate a self-signed cert + key in tmp_path. - - Returns (cert_path, key_path). - Requires the 'cryptography' package. - """ - from cryptography import x509 - from cryptography.hazmat.primitives import hashes, serialization - from cryptography.hazmat.primitives.asymmetric import rsa - from cryptography.x509.oid import NameOID - from datetime import datetime, timedelta, timezone - - key = rsa.generate_private_key(public_exponent=65537, key_size=2048) - - subject = issuer = x509.Name([ - x509.NameAttribute(NameOID.COMMON_NAME, cn), - ]) - - now = datetime.now(timezone.utc) - cert = ( - x509.CertificateBuilder() - .subject_name(subject) - .issuer_name(issuer) - .public_key(key.public_key()) - .serial_number(x509.random_serial_number()) - .not_valid_before(now) - .not_valid_after(now + timedelta(days=days)) - .add_extension( - x509.SubjectAlternativeName([x509.DNSName(cn)]), - critical=False, - ) - .sign(key, hashes.SHA256()) - ) - cert_path = tmp_path / "cert.pem" - key_path = tmp_path / "key.pem" +def _certificate_reference(tmp_path: Path) -> Path: + path = tmp_path / "certificate.ref" + path.touch() + return path - cert_path.write_bytes(cert.public_bytes(serialization.Encoding.PEM)) - key_path.write_bytes( - key.private_bytes( - serialization.Encoding.PEM, - serialization.PrivateFormat.TraditionalOpenSSL, - serialization.NoEncryption(), - ) + +def _info(days: int = 365, subject: str = "CN=service.baton.local") -> CertificateInfo: + return CertificateInfo( + subject=subject, + san=["service.baton.local"], + fingerprint_sha256="fingerprint-reference", + days_until_expiry=days, ) - return cert_path, key_path +def _parser(info: CertificateInfo): + def parse(_path: str | Path) -> CertificateInfo: + return info -try: - import cryptography # noqa: F401 - HAS_CRYPTO = True -except ImportError: - HAS_CRYPTO = False + return parse -pytestmark = pytest.mark.skipif(not HAS_CRYPTO, reason="cryptography package not installed") +class ReloadRecorder: + def __init__(self, fail: bool = False): + self.fail = fail + self.calls: list[tuple[str, str]] = [] -# --------------------------------------------------------------------------- -# parse_certificate tests -# --------------------------------------------------------------------------- + def __call__(self, certificate_path: str, custody_reference: str) -> None: + self.calls.append((certificate_path, custody_reference)) + if self.fail: + raise RuntimeError("reload unavailable") class TestParseCertificate: - def test_parse_valid_cert(self, tmp_path): - from baton.certs import parse_certificate - - cert_path, _ = _generate_self_signed(tmp_path) - info = parse_certificate(cert_path) - - assert "test.baton.local" in info.subject - assert info.days_until_expiry > 360 - assert info.fingerprint_sha256 != "" - assert "test.baton.local" in info.san - assert not info.is_expired - - def test_parse_missing_cert(self, tmp_path): - from baton.certs import parse_certificate - + def test_parse_missing_certificate(self, tmp_path): with pytest.raises(FileNotFoundError): parse_certificate(tmp_path / "nonexistent.pem") - def test_expired_cert(self, tmp_path): - from baton.certs import parse_certificate - - cert_path, _ = _generate_self_signed(tmp_path, days=0) - info = parse_certificate(cert_path) - assert info.days_until_expiry <= 0 - - -# --------------------------------------------------------------------------- -# CertificateMonitor tests -# --------------------------------------------------------------------------- - class TestCertificateMonitor: def test_initial_load(self, tmp_path): - from baton.certs import CertificateMonitor - - cert_path, _ = _generate_self_signed(tmp_path) - monitor = CertificateMonitor(cert_path) + certificate_ref = _certificate_reference(tmp_path) + monitor = CertificateMonitor(certificate_ref, parser=_parser(_info())) info, events = monitor.check() assert info is not None - assert any(e.event_type == "loaded" for e in events) - - def test_missing_cert_error(self, tmp_path): - from baton.certs import CertificateMonitor + assert info.subject == "CN=service.baton.local" + assert any(event.event_type == "loaded" for event in events) - monitor = CertificateMonitor(tmp_path / "missing.pem") + def test_missing_certificate_error(self, tmp_path): + monitor = CertificateMonitor(tmp_path / "missing.ref", parser=_parser(_info())) info, events = monitor.check() + assert info is None - assert any(e.event_type == "error" for e in events) + assert any(event.event_type == "error" for event in events) def test_expiring_warning(self, tmp_path): - from baton.certs import CertificateMonitor - - cert_path, _ = _generate_self_signed(tmp_path, days=15) - monitor = CertificateMonitor(cert_path, warning_days=30, critical_days=7) + certificate_ref = _certificate_reference(tmp_path) + monitor = CertificateMonitor( + certificate_ref, + warning_days=30, + critical_days=7, + parser=_parser(_info(days=15)), + ) - info, events = monitor.check() - assert any(e.event_type == "expiring_warning" for e in events) + _info_result, events = monitor.check() + assert any(event.event_type == "expiring_warning" for event in events) def test_expiring_critical(self, tmp_path): - from baton.certs import CertificateMonitor - - cert_path, _ = _generate_self_signed(tmp_path, days=3) - monitor = CertificateMonitor(cert_path, warning_days=30, critical_days=7) + certificate_ref = _certificate_reference(tmp_path) + monitor = CertificateMonitor( + certificate_ref, + warning_days=30, + critical_days=7, + parser=_parser(_info(days=3)), + ) - info, events = monitor.check() - assert any(e.event_type == "expiring_critical" for e in events) + _info_result, events = monitor.check() + assert any(event.event_type == "expiring_critical" for event in events) def test_file_change_detected(self, tmp_path): - from baton.certs import CertificateMonitor - - cert_path, _ = _generate_self_signed(tmp_path) - monitor = CertificateMonitor(cert_path) - - # First check + certificate_ref = _certificate_reference(tmp_path) + monitor = CertificateMonitor(certificate_ref, parser=_parser(_info())) monitor.check() - # Regenerate cert (changes mtime) time.sleep(0.01) - _generate_self_signed(tmp_path, cn="new.baton.local") - - # Second check - info, events = monitor.check() - assert any(e.event_type == "rotated" for e in events) - + certificate_ref.write_text("updated-reference") -# --------------------------------------------------------------------------- -# CertificateRotator tests -# --------------------------------------------------------------------------- + _info_result, events = monitor.check() + assert any(event.event_type == "rotated" for event in events) class TestCertificateRotator: def test_rotate_success(self, tmp_path): - from baton.certs import CertificateRotator - - cert_path, key_path = _generate_self_signed(tmp_path) - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - ctx.load_cert_chain(str(cert_path), str(key_path)) + certificate_ref = _certificate_reference(tmp_path) + reload = ReloadRecorder() + rotator = CertificateRotator( + object(), + certificate_ref, + tmp_path / "custody-reference", + certificate_loader=reload, + ) - rotator = CertificateRotator(ctx, cert_path, key_path) assert rotator.rotate() is True + assert len(reload.calls) == 1 + + def test_rotate_failed_reload(self, tmp_path): + certificate_ref = _certificate_reference(tmp_path) + rotator = CertificateRotator( + object(), + certificate_ref, + tmp_path / "custody-reference", + certificate_loader=ReloadRecorder(fail=True), + ) - def test_rotate_bad_key(self, tmp_path): - from baton.certs import CertificateRotator - - cert_path, key_path = _generate_self_signed(tmp_path) - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - ctx.load_cert_chain(str(cert_path), str(key_path)) - - # Corrupt the key file - bad_key = tmp_path / "bad_key.pem" - bad_key.write_text("not a key") - - rotator = CertificateRotator(ctx, cert_path, bad_key) assert rotator.rotate() is False -# --------------------------------------------------------------------------- -# CertificateManager tests -# --------------------------------------------------------------------------- - - class TestCertificateManager: - def test_check_now(self, tmp_path): - from baton.certs import CertificateManager + def _manager(self, tmp_path, *, interval: float = 3600.0): + certificate_ref = _certificate_reference(tmp_path) + reload = ReloadRecorder() + manager = CertificateManager( + object(), + certificate_ref, + tmp_path / "custody-reference", + check_interval=interval, + parser=_parser(_info()), + certificate_loader=reload, + ) + return manager, certificate_ref, reload - cert_path, key_path = _generate_self_signed(tmp_path) - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - ctx.load_cert_chain(str(cert_path), str(key_path)) + def test_check_now(self, tmp_path): + manager, _certificate_ref, _reload = self._manager(tmp_path) + info, events = manager.check_now() - mgr = CertificateManager(ctx, cert_path, key_path) - info, events = mgr.check_now() assert info is not None assert len(events) >= 1 def test_auto_rotate_on_change(self, tmp_path): - from baton.certs import CertificateManager + manager, certificate_ref, reload = self._manager(tmp_path) + manager.check_now() - cert_path, key_path = _generate_self_signed(tmp_path) - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - ctx.load_cert_chain(str(cert_path), str(key_path)) - - mgr = CertificateManager(ctx, cert_path, key_path) - - # First check loads cert - mgr.check_now() - - # Regenerate cert time.sleep(0.01) - _generate_self_signed(tmp_path, cn="rotated.baton.local") + certificate_ref.write_text("updated-reference") + _info_result, events = manager.check_now() - # Second check should detect change and rotate - info, events = mgr.check_now() - event_types = [e.event_type for e in events] - assert "rotated" in event_types + assert "rotated" in [event.event_type for event in events] + assert len(reload.calls) == 1 def test_events_accumulated(self, tmp_path): - from baton.certs import CertificateManager + manager, _certificate_ref, _reload = self._manager(tmp_path) + manager.check_now() + manager.check_now() - cert_path, key_path = _generate_self_signed(tmp_path) - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - ctx.load_cert_chain(str(cert_path), str(key_path)) - - mgr = CertificateManager(ctx, cert_path, key_path) - mgr.check_now() - mgr.check_now() - - assert len(mgr.events) >= 1 + assert len(manager.events) >= 1 async def test_run_and_stop(self, tmp_path): - from baton.certs import CertificateManager - import asyncio - - cert_path, key_path = _generate_self_signed(tmp_path) - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - ctx.load_cert_chain(str(cert_path), str(key_path)) + manager, _certificate_ref, _reload = self._manager(tmp_path, interval=0.01) - mgr = CertificateManager(ctx, cert_path, key_path, check_interval=0.1) + task = asyncio.create_task(manager.run()) + await asyncio.sleep(0.03) + assert manager.is_running - task = asyncio.create_task(mgr.run()) - await asyncio.sleep(0.3) - assert mgr.is_running - - mgr.stop() - await asyncio.sleep(0.2) - assert not mgr.is_running + manager.stop() + await asyncio.sleep(0.02) + assert not manager.is_running task.cancel() try: @@ -276,26 +192,15 @@ async def test_run_and_stop(self, tmp_path): pass -# --------------------------------------------------------------------------- -# CertificateInfo tests -# --------------------------------------------------------------------------- - - class TestCertificateInfo: def test_is_expired_true(self): - from baton.certs import CertificateInfo - info = CertificateInfo(days_until_expiry=0) assert info.is_expired def test_is_expired_false(self): - from baton.certs import CertificateInfo - info = CertificateInfo(days_until_expiry=30) assert not info.is_expired def test_is_expired_unknown(self): - from baton.certs import CertificateInfo - info = CertificateInfo(days_until_expiry=-1) assert not info.is_expired From fbebff0daba0fb7bce59b0b8efb8cd5f478000ba Mon Sep 17 00:00:00 2001 From: Jeremy McEntire Date: Thu, 4 Jun 2026 17:54:34 -0500 Subject: [PATCH 4/4] build: run verification against source checkout --- Makefile | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/Makefile b/Makefile index 8061646..458b777 100644 --- a/Makefile +++ b/Makefile @@ -1,16 +1,19 @@ .PHONY: install dev test lint clean +PYTHON ?= python3 +PYTHONPATH ?= src + install: - pip install -e . + $(PYTHON) -m pip install -e . dev: - pip install -e ".[dev]" + $(PYTHON) -m pip install -e ".[dev]" test: - pytest + PYTHONPATH=$(PYTHONPATH) $(PYTHON) -m pytest lint: - python -m py_compile src/baton/*.py + $(PYTHON) -m py_compile src/baton/*.py clean: rm -rf build dist *.egg-info src/*.egg-info .pytest_cache