diff --git a/solstone/think/services/operations.py b/solstone/think/services/operations.py index 3c0c15959..6e72cca70 100644 --- a/solstone/think/services/operations.py +++ b/solstone/think/services/operations.py @@ -27,11 +27,11 @@ RETRYABLE_CODES = frozenset( ) OPERATION_GRACE_SECONDS = 30.0 # Phases at which an operation is finished — no actionable consent CTA should -# be surfaced (the portal page is already satisfied or moot). Mirrors the JS -# PRIVATE_LINK_TERMINAL_PHASES set in solstone/apps/network/workspace.html — the -# two sit on opposite sides of the Python/JS boundary with no shared source, -# so keep them in lockstep. -TERMINAL_PHASES = frozenset({"enabled", "needs_subscription", "revoked", "error"}) +# be surfaced (the portal page is already satisfied or moot). Individual apps +# may expose narrower product-specific terminal sets in browser code. +TERMINAL_PHASES = frozenset( + {"enabled", "needs_subscription", "revoked", "error", "early_access"} +) class OperationBusyError(RuntimeError): diff --git a/solstone/think/services/portal_client.py b/solstone/think/services/portal_client.py index b2cb2cbef..8c4f9604f 100644 --- a/solstone/think/services/portal_client.py +++ b/solstone/think/services/portal_client.py @@ -245,13 +245,16 @@ def poll_handoff_once( if status == 200: try: - return PollOutcome(kind="success", payload=read_handoff_payload(raw_body)) + payload = read_handoff_payload(raw_body) except ValueError as exc: return PollOutcome( kind="failed", reason="unexpected_payload", detail=str(exc), ) + if payload.get("state") == "early_access": + return PollOutcome(kind="early_access") + return PollOutcome(kind="success", payload=payload) if status == 204: return PollOutcome(kind="continue") return handle_http_status(status) diff --git a/solstone/think/services/spb_handoff.py b/solstone/think/services/spb_handoff.py index 787f524b5..249e99be8 100644 --- a/solstone/think/services/spb_handoff.py +++ b/solstone/think/services/spb_handoff.py @@ -129,6 +129,13 @@ def enable_spb_via_consent( continue if outcome.kind == "failed": return _failed_result(outcome.reason) + if outcome.kind != "success": + return SpbHandoffResult( + outcomes.MALFORMED, + None, + None, + outcomes.MALFORMED, + ) payload = outcome.payload if not isinstance(payload, dict): diff --git a/solstone/think/services/spl_handoff.py b/solstone/think/services/spl_handoff.py index c1bb5dd36..d3635db4f 100644 --- a/solstone/think/services/spl_handoff.py +++ b/solstone/think/services/spl_handoff.py @@ -102,6 +102,8 @@ def enable_spl_via_consent( detail=outcome.detail, ) return outcomes.outcome_for_code(outcomes.MALFORMED, detail=outcome.detail) + if outcome.kind != "success": + return outcomes.outcome_for_code(outcomes.MALFORMED) payload = outcome.payload or {} if not isinstance(payload, dict): diff --git a/solstone/think/services/spp.py b/solstone/think/services/spp.py index 413f85730..13bb69539 100644 --- a/solstone/think/services/spp.py +++ b/solstone/think/services/spp.py @@ -10,7 +10,7 @@ import logging import threading from dataclasses import dataclass from datetime import datetime, timezone -from typing import Any +from typing import Any, Literal from urllib.parse import urlsplit from solstone.think.journal_config import ( @@ -46,10 +46,17 @@ class DisableOutcome: credential_preserved: bool +@dataclass(frozen=True, slots=True) +class AttestationFailure: + kind: Literal["failed", "unreachable"] + reason_code: str + + @dataclass(frozen=True, slots=True) class AttestationState: session: AttestationSession | None = None - failure: str | None = None + failure: AttestationFailure | None = None + last_verified: AttestationSession | None = None _ATTESTATION_LOCK = threading.Lock() @@ -61,19 +68,40 @@ def record_attestation_verified(session: AttestationSession) -> None: global _ATTESTATION_STATE with _ATTESTATION_LOCK: - _ATTESTATION_STATE = AttestationState(session=session, failure=None) + _ATTESTATION_STATE = AttestationState( + session=session, + failure=None, + last_verified=session, + ) -def record_attestation_failed(detail: str) -> None: +def record_attestation_failed( + kind: Literal["failed", "unreachable"], + reason_code: str, +) -> None: """Record an in-process attestation failure.""" global _ATTESTATION_STATE with _ATTESTATION_LOCK: - _ATTESTATION_STATE = AttestationState(session=None, failure=detail) + _ATTESTATION_STATE = AttestationState( + session=None, + failure=AttestationFailure(kind=kind, reason_code=reason_code), + last_verified=_ATTESTATION_STATE.last_verified, + ) def clear_attestation_state() -> None: - """Clear process-local attestation state.""" + """Clear current process-local attestation state while preserving last verified.""" + + global _ATTESTATION_STATE + with _ATTESTATION_LOCK: + _ATTESTATION_STATE = AttestationState( + last_verified=_ATTESTATION_STATE.last_verified + ) + + +def delete_attestation_state() -> None: + """Delete all process-local attestation state.""" global _ATTESTATION_STATE with _ATTESTATION_LOCK: @@ -211,6 +239,7 @@ def disable_confidential() -> DisableOutcome: services = config.setdefault("services", {}) block = services.get("confidential") if not isinstance(block, dict): + delete_attestation_state() return DisableOutcome(was_enabled=False, credential_preserved=False) providers = _providers_block(config) @@ -243,6 +272,7 @@ def disable_confidential() -> DisableOutcome: services.pop("confidential", None) write_journal_config(config) + delete_attestation_state() log.debug("disabled confidential service") return DisableOutcome( was_enabled=True, diff --git a/solstone/think/services/spp_handoff.py b/solstone/think/services/spp_handoff.py index 86263a53d..004a746a3 100644 --- a/solstone/think/services/spp_handoff.py +++ b/solstone/think/services/spp_handoff.py @@ -9,6 +9,7 @@ import logging import time from collections.abc import Callable +from solstone.think.link.paths import LinkState from solstone.think.services import operations, outcomes, portal_client, spp from solstone.think.services.constants import SERVICE_SPP @@ -26,7 +27,8 @@ def build_confidential_handoff_url() -> tuple[str, str, str]: Returns ``(consent_url, nonce, base_url)``. """ - return portal_client.build_consent_url(SERVICE_SPP) + instance_id = LinkState.load_or_create().instance_id + return portal_client.build_consent_url(SERVICE_SPP, instance=instance_id) def _handoff_error_result( @@ -78,6 +80,12 @@ def run_confidential_handoff( "unexpected_payload", detail=outcome.detail, ) + if outcome.kind == "early_access": + return operations.HandoffResult( + phase="early_access", + guidance=None, + retryable=False, + ) if outcome.kind != "success": return _handoff_error_result("unexpected_payload") diff --git a/solstone/think/services/spp_transport.py b/solstone/think/services/spp_transport.py index 45a4a16cc..780a0b498 100644 --- a/solstone/think/services/spp_transport.py +++ b/solstone/think/services/spp_transport.py @@ -6,14 +6,13 @@ from __future__ import annotations import logging -import re import secrets import selectors import socket import threading import time from datetime import datetime, timezone -from typing import Any +from typing import Any, Literal from urllib.parse import urlsplit from OpenSSL import SSL @@ -45,37 +44,40 @@ _LISTENER: socket.socket | None = None _LISTENER_THREAD: threading.Thread | None = None _FORWARDER_BASE_URL: str | None = None _CONFIDENTIAL_BLOCK: dict[str, Any] | None = None -_REASON_CODE_RE = re.compile(r"\(([a-z0-9_]+)\)$") -def _attestation_failed(reason_code: str) -> None: +class ConfidentialEndpointError(RuntimeError): + """Raised when confidential endpoint configuration is unusable.""" + + reason_code = "endpoint_invalid" + + +def _failure_kind(exc: BaseException) -> Literal["failed", "unreachable"]: + if isinstance(exc, RatlsChannelError) and exc.reason_code == "gateway_unreachable": + return "unreachable" + return "failed" + + +def _attestation_failed( + kind: Literal["failed", "unreachable"], + reason_code: str, +) -> None: log.warning("event=confidential_attestation_rejected reason=%s", reason_code) - spp.record_attestation_failed(reason_code) + spp.record_attestation_failed(kind, reason_code) raise AttestationFailedError( f"the confidential attestation transport failed closed ({reason_code})" ) -def _reason_from_attestation_failed(exc: AttestationFailedError) -> str: - match = _REASON_CODE_RE.search(exc.detail.strip()) - if match is not None: - return match.group(1) - return exc.reason_code - - def _endpoint_from_block(block: dict[str, Any]) -> RatlsEndpoint: endpoint_url = str(block.get("endpoint_url") or "") parsed = urlsplit(endpoint_url) if not parsed.hostname: - raise AttestationFailedError( - "the confidential endpoint configuration is invalid (endpoint_invalid)" - ) + raise ConfidentialEndpointError("confidential endpoint hostname is invalid") try: port = parsed.port except ValueError: - raise AttestationFailedError( - "the confidential endpoint configuration is invalid (endpoint_invalid)" - ) + raise ConfidentialEndpointError("confidential endpoint port is invalid") if port is None: port = 443 if parsed.scheme == "https" else 80 return RatlsEndpoint(host=parsed.hostname, port=port) @@ -120,6 +122,7 @@ def _teardown_locked() -> None: def teardown_confidential_transport() -> None: with _LOCK: _teardown_locked() + spp.clear_attestation_state() def _discard_idle_locked(now_monotonic: float) -> None: @@ -177,15 +180,32 @@ def _establish_channel_locked(block: dict[str, Any], now: datetime) -> AttestedC ) except (RatlsChannelError, RatlsVerificationError) as exc: reason_code = exc.reason_code + kind = _failure_kind(exc) _teardown_locked() - _attestation_failed(reason_code) + _attestation_failed(kind, reason_code) + except ConfidentialEndpointError as exc: + _teardown_locked() + _attestation_failed("failed", exc.reason_code) except AttestationFailedError as exc: - reason_code = _reason_from_attestation_failed(exc) _teardown_locked() - _attestation_failed(reason_code) + _attestation_failed("failed", exc.reason_code) except Exception: _teardown_locked() - _attestation_failed("unexpected_error") + _attestation_failed("failed", "unexpected_error") + + +def _establish_and_record_locked(block: dict[str, Any], now: datetime) -> None: + channel = _establish_channel_locked(block, now) + _start_listener_locked() + _POOL.append(channel) + spp.record_attestation_verified( + AttestationSession( + verdict=channel.verdict, + started_at=now, + tpm_heartbeat_at=now, + gpu_reattest_at=now, + ) + ) def verify_confidential_attestation(block: dict[str, Any]) -> None: @@ -211,17 +231,7 @@ def verify_confidential_attestation(block: dict[str, Any]) -> None: "the confidential attestation cadence lapsed (attestation_stale)" ) - channel = _establish_channel_locked(block, now) - _start_listener_locked() - _POOL.append(channel) - spp.record_attestation_verified( - AttestationSession( - verdict=channel.verdict, - started_at=now, - tpm_heartbeat_at=now, - gpu_reattest_at=now, - ) - ) + _establish_and_record_locked(block, now) def confidential_egress_base_url(endpoint_base_url: str) -> str: @@ -245,6 +255,8 @@ def confidential_probe_status( now = now or datetime.now(timezone.utc) state = spp.get_attestation_state() if state.failure is not None: + if state.failure.kind == "unreachable": + return False, "attestation_unreachable" return False, "attestation_failed" if state.session is None: return False, "attestation_not_yet_verified" @@ -253,6 +265,25 @@ def confidential_probe_status( return True, None +def recheck_confidential_attestation() -> None: + """Recheck the configured confidential service without sending inference content.""" + + global _CONFIDENTIAL_BLOCK + + block = spp.confidential_provenance() + if block is None: + return + now = datetime.now(timezone.utc) + with _LOCK: + _CONFIDENTIAL_BLOCK = dict(block) + _teardown_locked() + spp.clear_attestation_state() + try: + _establish_and_record_locked(block, now) + except AttestationFailedError: + return + + def _borrow_channel_locked(now_monotonic: float) -> AttestedChannel | None: _discard_idle_locked(now_monotonic) if not _POOL: diff --git a/tests/services/test_operations.py b/tests/services/test_operations.py index d218f959e..794448e44 100644 --- a/tests/services/test_operations.py +++ b/tests/services/test_operations.py @@ -120,6 +120,21 @@ def test_terminal_enabled_entry_drops_portal_url_in_grace() -> None: assert op["portal_url"] is None +def test_terminal_early_access_entry_drops_portal_url_in_grace() -> None: + operations.start_operation( + "spp", + "enable", + "https://portal.test/spp", + lambda: operations.HandoffResult("early_access", None, False), + ) + _wait_until( + lambda: operations.operation_for_service("spp")["phase"] == "early_access" + ) + op = operations.operation_for_service("spp") + assert op["phase"] == "early_access" + assert op["portal_url"] is None + + def test_non_terminal_phases_keep_portal_url() -> None: started = threading.Event() release = threading.Event() diff --git a/tests/services/test_portal_client_services.py b/tests/services/test_portal_client_services.py index 1939f9b54..9a0a2bcc9 100644 --- a/tests/services/test_portal_client_services.py +++ b/tests/services/test_portal_client_services.py @@ -3,11 +3,32 @@ from __future__ import annotations +import json +from typing import Any + import pytest from solstone.think.services import portal_client +class FakeResponse: + def __init__(self, status: int, payload: dict[str, Any]) -> None: + self.status = status + self._payload = payload + + def __enter__(self) -> FakeResponse: + return self + + def __exit__(self, *_args: object) -> None: + return None + + def getcode(self) -> int: + return self.status + + def read(self) -> bytes: + return json.dumps(self._payload).encode("utf-8") + + def test_scout_is_default_handoff_service() -> None: assert ( portal_client.browser_url("https://services.test", "NONCE") @@ -120,3 +141,19 @@ def test_poll_handoff_unknown_service_never_opens_network(monkeypatch) -> None: "NONCE", service="bogus", ) + + +def test_poll_handoff_early_access_is_terminal_kind(monkeypatch) -> None: + monkeypatch.setattr( + portal_client.urllib.request, + "urlopen", + lambda *_args, **_kwargs: FakeResponse(200, {"state": "early_access"}), + ) + + outcome = portal_client.poll_handoff_once( + "https://services.test", + "NONCE", + service="spp", + ) + + assert outcome == portal_client.PollOutcome(kind="early_access") diff --git a/tests/services/test_scout_handoff.py b/tests/services/test_scout_handoff.py index d632cbd53..c70652ff2 100644 --- a/tests/services/test_scout_handoff.py +++ b/tests/services/test_scout_handoff.py @@ -145,3 +145,20 @@ def test_run_scout_handoff_malformed_apply_does_not_write_journal( assert result.phase == "error" assert result.retryable is False assert _config_bytes(journal_copy) == before + + +def test_run_scout_handoff_unknown_outcome_kind_fails_loudly( + journal_copy: Path, +) -> None: + before = _config_bytes(journal_copy) + + result = scout_handoff.run_scout_handoff( + refresh=True, + nonce="NONCE", + base_url="http://portal.test", + poll_once=lambda *_args, **_kwargs: PollOutcome(kind="early_access"), + ) + + assert result.phase == "error" + assert result.retryable is False + assert _config_bytes(journal_copy) == before diff --git a/tests/services/test_spl_handoff.py b/tests/services/test_spl_handoff.py index 2eae09119..df9d8247f 100644 --- a/tests/services/test_spl_handoff.py +++ b/tests/services/test_spl_handoff.py @@ -693,6 +693,13 @@ def test_run_spl_handoff_maps_terminal_outcomes(journal_copy: Path) -> None: payload={"service": "spl", "state": "bad"}, ), ) + unknown = spl_handoff.run_spl_handoff( + nonce=TEST_NONCE, + base_url=TEST_BASE_URL, + poll_once=lambda *_args, **_kwargs: portal_client.PollOutcome( + kind="early_access", + ), + ) assert revoked.phase == "revoked" assert revoked.retryable is False @@ -700,3 +707,5 @@ def test_run_spl_handoff_maps_terminal_outcomes(journal_copy: Path) -> None: assert expired.retryable is True assert malformed.phase == "error" assert malformed.retryable is False + assert unknown.phase == "error" + assert unknown.retryable is False diff --git a/tests/services/test_spp_attestation_state.py b/tests/services/test_spp_attestation_state.py index 1d5f9b64f..12187651c 100644 --- a/tests/services/test_spp_attestation_state.py +++ b/tests/services/test_spp_attestation_state.py @@ -18,9 +18,9 @@ NOW = datetime(2026, 7, 11, 18, 0, tzinfo=timezone.utc) @pytest.fixture(autouse=True) def _clear_attestation_state(): - spp.clear_attestation_state() + spp.delete_attestation_state() yield - spp.clear_attestation_state() + spp.delete_attestation_state() def _session() -> AttestationSession: @@ -68,10 +68,11 @@ def test_attestation_state_defaults_empty() -> None: assert state.session is None assert state.failure is None + assert state.last_verified is None def test_record_attestation_verified_sets_session_and_clears_failure() -> None: - spp.record_attestation_failed("prior") + spp.record_attestation_failed("failed", "prior") session = _session() spp.record_attestation_verified(session) @@ -79,23 +80,42 @@ def test_record_attestation_verified_sets_session_and_clears_failure() -> None: state = spp.get_attestation_state() assert state.session is session assert state.failure is None + assert state.last_verified is session -def test_record_attestation_failed_sets_failure_and_clears_session() -> None: - spp.record_attestation_verified(_session()) +def test_record_attestation_failed_sets_failure_and_preserves_last_verified() -> None: + session = _session() + spp.record_attestation_verified(session) - spp.record_attestation_failed("gpu_nonce_mismatch") + spp.record_attestation_failed("failed", "gpu_nonce_mismatch") state = spp.get_attestation_state() assert state.session is None - assert state.failure == "gpu_nonce_mismatch" + assert state.failure is not None + assert state.failure.kind == "failed" + assert state.failure.reason_code == "gpu_nonce_mismatch" + assert state.last_verified is session -def test_clear_attestation_state_resets_holder() -> None: - spp.record_attestation_failed("failure") +def test_clear_attestation_state_preserves_last_verified() -> None: + session = _session() + spp.record_attestation_verified(session) + spp.record_attestation_failed("unreachable", "gateway_unreachable") spp.clear_attestation_state() state = spp.get_attestation_state() assert state.session is None assert state.failure is None + assert state.last_verified is session + + +def test_delete_attestation_state_resets_holder() -> None: + spp.record_attestation_verified(_session()) + + spp.delete_attestation_state() + + state = spp.get_attestation_state() + assert state.session is None + assert state.failure is None + assert state.last_verified is None diff --git a/tests/services/test_spp_handoff.py b/tests/services/test_spp_handoff.py index 23ab99b00..3c26d1ce3 100644 --- a/tests/services/test_spp_handoff.py +++ b/tests/services/test_spp_handoff.py @@ -5,6 +5,7 @@ from __future__ import annotations import json from pathlib import Path +from types import SimpleNamespace import pytest @@ -28,6 +29,15 @@ def _stable_portal(monkeypatch: pytest.MonkeyPatch) -> None: spp_handoff.portal_client, "portal_base_url", lambda: "http://portal.test" ) monkeypatch.setattr(spp_handoff.portal_client, "mint_nonce", lambda: "NONCE") + monkeypatch.setattr( + spp_handoff.LinkState, + "load_or_create", + classmethod( + lambda cls: SimpleNamespace( + instance_id="00000000-0000-4000-8000-000000000000" + ) + ), + ) def _config_bytes(journal: Path) -> bytes: @@ -41,7 +51,10 @@ def _config(journal: Path) -> dict: def test_build_confidential_handoff_url_uses_spp_service() -> None: consent_url, nonce, base_url = spp_handoff.build_confidential_handoff_url() - assert consent_url == "http://portal.test/enable/spp?nonce=NONCE" + assert ( + consent_url + == "http://portal.test/enable/spp?nonce=NONCE&instance=00000000-0000-4000-8000-000000000000" + ) assert nonce == "NONCE" assert base_url == "http://portal.test" @@ -121,3 +134,21 @@ def test_run_confidential_handoff_malformed_apply_does_not_write_journal( assert result.phase == "error" assert result.retryable is False assert _config_bytes(journal_copy) == before + + +def test_run_confidential_handoff_early_access_is_terminal_without_write( + journal_copy: Path, +) -> None: + before = _config_bytes(journal_copy) + + result = spp_handoff.run_confidential_handoff( + refresh=True, + nonce="NONCE", + base_url="http://portal.test", + poll_once=lambda *_args, **_kwargs: PollOutcome(kind="early_access"), + ) + + assert result.phase == "early_access" + assert result.guidance is None + assert result.retryable is False + assert _config_bytes(journal_copy) == before diff --git a/tests/services/test_spp_transport.py b/tests/services/test_spp_transport.py index 4913cca4c..57b2d4a94 100644 --- a/tests/services/test_spp_transport.py +++ b/tests/services/test_spp_transport.py @@ -11,9 +11,11 @@ from unittest.mock import Mock import pytest -from solstone.think.models import AttestationStaleError +from solstone.think.models import AttestationFailedError, AttestationStaleError from solstone.think.services import spp, spp_transport from solstone.think.services.spp_attest.cadence import AttestationSession +from solstone.think.services.spp_attest.ratls.channel import RatlsChannelError +from solstone.think.services.spp_attest.ratls.verify import RatlsVerificationError class _FakeChannel: @@ -58,10 +60,10 @@ class _AliveThread: @pytest.fixture(autouse=True) def _clear_transport_state(): - spp.clear_attestation_state() + spp.delete_attestation_state() spp_transport.teardown_confidential_transport() yield - spp.clear_attestation_state() + spp.delete_attestation_state() spp_transport.teardown_confidential_transport() @@ -194,8 +196,11 @@ def test_confidential_probe_status_reads_state_without_attestation( spp.record_attestation_verified(_stale_session(object())) assert spp_transport.confidential_probe_status() == (False, "attestation_stale") - spp.record_attestation_failed("gateway_unreachable") - assert spp_transport.confidential_probe_status() == (False, "attestation_failed") + spp.record_attestation_failed("unreachable", "gateway_unreachable") + assert spp_transport.confidential_probe_status() == ( + False, + "attestation_unreachable", + ) establish.assert_not_called() @@ -252,3 +257,111 @@ def test_wrong_epoch_channel_is_not_checked_out() -> None: assert spp_transport._borrow_channel_locked(time.monotonic()) is None assert stale_epoch_channel.closed is True assert spp_transport._POOL == [] + + +@pytest.mark.parametrize( + ("exc", "kind", "reason_code"), + [ + ( + RatlsChannelError("gateway_unreachable"), + "unreachable", + "gateway_unreachable", + ), + (RatlsChannelError("tls_handshake_failed"), "failed", "tls_handshake_failed"), + (RatlsVerificationError("nonce_mismatch"), "failed", "nonce_mismatch"), + (RuntimeError("boom"), "failed", "unexpected_error"), + ], +) +def test_attestation_failure_buckets_at_transport_catch_site( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + exc: Exception, + kind: str, + reason_code: str, +) -> None: + block = _write_confidential_config(tmp_path, monkeypatch) + monkeypatch.setattr( + spp_transport, "establish_attested_channel", Mock(side_effect=exc) + ) + + with pytest.raises(AttestationFailedError): + spp_transport.verify_confidential_attestation(block) + + failure = spp.get_attestation_state().failure + assert failure is not None + assert failure.kind == kind + assert failure.reason_code == reason_code + + +def test_endpoint_invalid_buckets_as_failed_without_detail_parsing( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + block = _write_confidential_config(tmp_path, monkeypatch) + block["endpoint_url"] = "not-a-url" + + with pytest.raises(AttestationFailedError): + spp_transport.verify_confidential_attestation(block) + + failure = spp.get_attestation_state().failure + assert failure is not None + assert failure.kind == "failed" + assert failure.reason_code == "endpoint_invalid" + + +def test_recheck_confidential_attestation_records_success( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + block = _write_confidential_config(tmp_path, monkeypatch) + _patch_listener(monkeypatch) + verdict = object() + monkeypatch.setattr( + spp_transport, + "establish_attested_channel", + Mock(return_value=_FakeChannel(verdict)), + ) + + spp_transport.recheck_confidential_attestation() + + state = spp.get_attestation_state() + assert state.session is not None + assert state.session.verdict is verdict + assert state.failure is None + assert state.last_verified is state.session + assert spp_transport._FORWARDER_BASE_URL == "http://127.0.0.1:4567" + assert block["endpoint_url"] == "https://spp.example.test:9443" + + +def test_recheck_confidential_attestation_fails_closed_and_preserves_last_verified( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _write_confidential_config(tmp_path, monkeypatch) + prior = _stale_session(object()) + spp.record_attestation_verified(prior) + monkeypatch.setattr( + spp_transport, + "establish_attested_channel", + Mock(side_effect=RatlsChannelError("gateway_unreachable")), + ) + + spp_transport.recheck_confidential_attestation() + + state = spp.get_attestation_state() + assert state.session is None + assert state.last_verified is prior + assert state.failure is not None + assert state.failure.kind == "unreachable" + assert state.failure.reason_code == "gateway_unreachable" + + +def test_recheck_confidential_attestation_off_is_noop( + monkeypatch: pytest.MonkeyPatch, +) -> None: + establish = Mock(side_effect=AssertionError("attestation attempted")) + monkeypatch.setattr(spp_transport, "establish_attested_channel", establish) + + spp_transport.recheck_confidential_attestation() + + establish.assert_not_called() diff --git a/tests/test_backup_spb_handoff.py b/tests/test_backup_spb_handoff.py index aab4e8294..773a4c9f2 100644 --- a/tests/test_backup_spb_handoff.py +++ b/tests/test_backup_spb_handoff.py @@ -227,6 +227,19 @@ def test_malformed_success_payload_returns_malformed() -> None: assert result.reason_code == outcomes.MALFORMED +def test_unknown_outcome_kind_returns_malformed() -> None: + result = _run( + poll_once=lambda *_args, **_kwargs: portal_client.PollOutcome( + kind="early_access", + ), + clock=lambda: 0.0, + wait_seconds=10, + ) + + assert result.state == outcomes.MALFORMED + assert result.reason_code == outcomes.MALFORMED + + def test_deadline_returns_expired() -> None: def fail_poll(*_args, **_kwargs): raise AssertionError("poll_once should not be called after deadline") diff --git a/tests/test_no_implicit_cloud.py b/tests/test_no_implicit_cloud.py index f81d001d5..93cc0fcca 100644 --- a/tests/test_no_implicit_cloud.py +++ b/tests/test_no_implicit_cloud.py @@ -33,10 +33,10 @@ from solstone.think.providers import get_provider_module def _clear_confidential_transport_state(): from solstone.think.services import spp, spp_transport - spp.clear_attestation_state() + spp.delete_attestation_state() spp_transport.teardown_confidential_transport() yield - spp.clear_attestation_state() + spp.delete_attestation_state() spp_transport.teardown_confidential_transport() @@ -111,7 +111,7 @@ def _install_failing_confidential_transport( from solstone.think.services import spp, spp_transport from solstone.think.services.spp_attest.ratls.channel import RatlsChannelError - spp.clear_attestation_state() + spp.delete_attestation_state() spp_transport.teardown_confidential_transport() monkeypatch.setattr(models, "_CONFIDENTIAL_ATTESTATION_VERIFIER", None) establish = Mock(side_effect=RatlsChannelError(reason_code))