From 4d4a8ebf036cd5d088da160e649779e1019bf5fa Mon Sep 17 00:00:00 2001 From: Jer Miller Date: Sun, 24 May 2026 00:27:37 -0600 Subject: [PATCH] refactor(link): extract bundle + TunnelClient dialer; migrate ObserverClient Move observer-specific PL bundle parsing into link.bundle as load_client_identity(bundle_dir), parameterized by bundle directory instead of an observer XDG label. The shared bundle module also owns PL_BUNDLE_FILES and endpoint labeling. Consolidate the PL dial race, direct/relay dialing, event-loop thread, cached TunnelSession, sync request, stream_request, and close lifecycle into link.dialer.TunnelClient. ObserverClient now constructs TunnelClient(identity, relay_url) lazily on first PL request and no longer owns duplicate loop, dial, or session-cache code. Move TlsError out of client.py into link.tls so dialer.py can import the shared error type without pulling in the full client module. Do not keep a backwards-compat re-export from client.py. Co-Authored-By: Claude Opus 4.7 (1M context) --- .../observer/tests/test_observer_client.py | 3 +- solstone/observe/observer_client.py | 315 ++++-------------- solstone/think/link/bundle.py | 64 ++++ solstone/think/link/client.py | 14 +- solstone/think/link/dialer.py | 302 +++++++++++++++++ solstone/think/link/tls.py | 8 + tests/link/test_bundle.py | 70 ++++ tests/link/test_dialer_unit.py | 70 ++-- tests/link/test_lan_direct.py | 3 +- 9 files changed, 549 insertions(+), 300 deletions(-) create mode 100644 solstone/think/link/bundle.py create mode 100644 solstone/think/link/dialer.py create mode 100644 solstone/think/link/tls.py create mode 100644 tests/link/test_bundle.py diff --git a/solstone/apps/observer/tests/test_observer_client.py b/solstone/apps/observer/tests/test_observer_client.py index 708d76251..d97a957bc 100644 --- a/solstone/apps/observer/tests/test_observer_client.py +++ b/solstone/apps/observer/tests/test_observer_client.py @@ -151,8 +151,9 @@ def test_observer_client_pl_requires_bundle( } } + client = ObserverClient("main-stream") with pytest.raises(ValueError, match="bundle not found"): - ObserverClient("main-stream") + client.relay_event("tract", "event") def test_auto_registration(mock_session, mock_config, mock_journal, tmp_path): diff --git a/solstone/observe/observer_client.py b/solstone/observe/observer_client.py index c7be7df76..0223b723d 100644 --- a/solstone/observe/observer_client.py +++ b/solstone/observe/observer_client.py @@ -3,7 +3,6 @@ from __future__ import annotations -import asyncio import json import logging import os @@ -20,15 +19,10 @@ import requests from urllib3.filepost import encode_multipart_formdata from solstone.apps.observer.routes import OBSERVER_CALLOSUM_SSE_ROUTE -from solstone.think.link.ca import cert_fingerprint -from solstone.think.link.client import ( - Client, - ClientIdentity, - EnrolledDevice, - StreamResetError, - TlsError, - TunnelSession, -) +from solstone.think.link.bundle import load_client_identity +from solstone.think.link.client import StreamResetError +from solstone.think.link.dialer import TunnelClient, TunnelRequestError +from solstone.think.link.tls import TlsError from solstone.think.utils import get_config, get_journal, read_service_port logger = logging.getLogger(__name__) @@ -39,13 +33,6 @@ MAX_RETRIES = 3 UPLOAD_TIMEOUT = 300 EVENT_TIMEOUT = 30 CALLOSUM_RECONNECT_BACKOFF = [1, 2, 4, 8, 16, 30] -PL_BUNDLE_FILES = { - "private.pem", - "cert.pem", - "chain.pem", - "home_attestation.jwt", - "peer.json", -} class UploadResult(NamedTuple): @@ -65,51 +52,6 @@ def _spl_bundle_dir(label: str) -> Path: return root / "solstone-observer" / "spl" / label -def _load_pl_identity(label: str) -> ClientIdentity: - bundle_dir = _spl_bundle_dir(label) - if not bundle_dir.is_dir(): - raise ValueError(f"observe.observer.spl_label bundle not found: {bundle_dir}") - - missing = sorted( - name for name in PL_BUNDLE_FILES if not (bundle_dir / name).exists() - ) - if missing: - raise ValueError( - "observe.observer.spl_label bundle missing required file(s): " - + ", ".join(missing) - ) - - private_key_pem = (bundle_dir / "private.pem").read_text(encoding="utf-8") - client_cert_pem = (bundle_dir / "cert.pem").read_text(encoding="utf-8") - ca_chain_pem = (bundle_dir / "chain.pem").read_text(encoding="utf-8") - home_attestation = (bundle_dir / "home_attestation.jwt").read_text(encoding="utf-8") - try: - peer = json.loads((bundle_dir / "peer.json").read_text(encoding="utf-8")) - except json.JSONDecodeError as exc: - raise ValueError(f"invalid peer.json in {bundle_dir}: {exc}") from exc - - local_endpoints = peer.get("local_endpoints") or () - if not isinstance(local_endpoints, list): - raise ValueError("peer.json local_endpoints must be a list") - - return ClientIdentity( - private_key_pem=private_key_pem, - client_cert_pem=client_cert_pem, - ca_chain_pem=ca_chain_pem, - fingerprint=cert_fingerprint(client_cert_pem), - home_instance_id=str(peer.get("instance_id") or ""), - home_label=str(peer.get("home_label") or ""), - home_attestation=home_attestation, - local_endpoints=tuple(local_endpoints), - ) - - -def _endpoint_label(endpoint: dict[str, object]) -> str: - host = str(endpoint.get("ip") or endpoint.get("host") or "?") - port = endpoint.get("port") or 7657 - return f"lan-direct {host}:{port}" - - def cleanup_draft(draft_dir: str) -> None: """Remove all files in a draft directory and delete the directory.""" try: @@ -181,13 +123,9 @@ class ObserverClient: self._callosum_stop = threading.Event() self._callosum_response: requests.Response | None = None self._callosum_error: Exception | None = None - self._pl_loop: asyncio.AbstractEventLoop | None = None - self._pl_loop_thread: threading.Thread | None = None - self._pl_session: TunnelSession | None = None - self._pl_session_lock: asyncio.Lock | None = None - self._pl_identity: ClientIdentity | None = None - self._pl_relay_url: str | None = None - self._pl_enrolled: EnrolledDevice | None = None + self._tunnel: TunnelClient | None = None + self._spl_label: str | None = None + self._spl_relay_url: str | None = None self._pl_fingerprint_prefix: str | None = None if self._pair_mode == "pl": @@ -206,11 +144,8 @@ class ObserverClient: raise ValueError( "observe.observer.spl_relay_url is required when pair_mode=pl" ) - self._pl_identity = _load_pl_identity(spl_label) - self._pl_relay_url = spl_relay_url.rstrip("/") - self._pl_fingerprint_prefix = self._pl_identity.fingerprint.replace( - "sha256:", "" - )[:16] + self._spl_label = spl_label + self._spl_relay_url = spl_relay_url.rstrip("/") self._auto_register = False def _persist_key(self, key: str) -> None: @@ -282,53 +217,17 @@ class ObserverClient: time.sleep(delay) logger.error(f"Registration failed after {MAX_RETRIES} attempts") - def _ensure_pl_loop(self) -> asyncio.AbstractEventLoop: - if self._pl_loop is not None and self._pl_loop.is_running(): - return self._pl_loop - - loop = asyncio.new_event_loop() - ready = threading.Event() - - def run_loop() -> None: - asyncio.set_event_loop(loop) - self._pl_session_lock = asyncio.Lock() - ready.set() - loop.run_forever() - - thread = threading.Thread( - target=run_loop, - name=f"observer-pl-{self._name}", - daemon=True, - ) - thread.start() - ready.wait() - self._pl_loop = loop - self._pl_loop_thread = thread - return loop - - def _run_pl(self, coro): - loop = self._ensure_pl_loop() - future = asyncio.run_coroutine_threadsafe(coro, loop) - return future.result() - - async def _get_pl_session(self) -> TunnelSession: - if self._pl_identity is None: - raise TlsError("PL identity not loaded") - if self._pl_session_lock is None: - self._pl_session_lock = asyncio.Lock() - async with self._pl_session_lock: - if self._pl_session is not None: - return self._pl_session - self._pl_session = await self._open_tunnel() - return self._pl_session - - async def _close_pl_session(self) -> None: - session = self._pl_session - self._pl_session = None - if session is not None: - await session.close() - - async def _pl_request( + def _pl_tunnel(self) -> TunnelClient: + if self._tunnel is not None: + return self._tunnel + if self._spl_label is None or self._spl_relay_url is None: + raise TlsError("PL identity not configured") + identity = load_client_identity(_spl_bundle_dir(self._spl_label)) + self._pl_fingerprint_prefix = identity.fingerprint.replace("sha256:", "")[:16] + self._tunnel = TunnelClient(identity, self._spl_relay_url) + return self._tunnel + + def _pl_request( self, method: str, path: str, @@ -336,83 +235,17 @@ class ObserverClient: headers: dict[str, str] | None = None, body: bytes = b"", ) -> PlRequestResult: - session = await self._get_pl_session() try: - status, response_headers, response_body = await session.request( + status, response_headers, response_body = self._pl_tunnel().request( method, path, headers=headers, body=body, ) return PlRequestResult(status, response_headers, response_body) - except (ConnectionError, OSError, StreamResetError): - await self._close_pl_session() + except TunnelRequestError: raise - async def _open_tunnel(self) -> TunnelSession: - if self._pl_identity is None: - raise TlsError("PL identity not loaded") - - attempts: list[tuple[str, Any]] = [] - for endpoint in self._pl_identity.local_endpoints: - label = _endpoint_label(endpoint) - attempts.append((label, self._dial_direct_endpoint(endpoint))) - if self._pl_relay_url: - attempts.append(("spl-relay", self._dial_relay())) - if not attempts: - raise TlsError("no PL dial attempts configured") - - tasks = {asyncio.create_task(coro): label for label, coro in attempts} - pending = set(tasks) - failures: dict[str, BaseException] = {} - - while pending: - done, pending = await asyncio.wait( - pending, - return_when=asyncio.FIRST_COMPLETED, - ) - for task in done: - label = tasks[task] - try: - session = task.result() - except BaseException as exc: - failures[label] = exc - continue - for loser in pending: - loser.cancel() - if pending: - await asyncio.gather(*pending, return_exceptions=True) - return session - - detail = "; ".join( - f"{label}: {type(exc).__name__}: {exc}" for label, exc in failures.items() - ) - raise TlsError(f"all PL dial attempts failed: {detail}") - - async def _dial_direct_endpoint(self, endpoint: dict[str, object]) -> TunnelSession: - if self._pl_identity is None: - raise TlsError("PL identity not loaded") - host = str(endpoint.get("ip") or endpoint.get("host") or "").strip() - if not host: - raise TlsError("LAN endpoint missing ip") - port_value = endpoint.get("port") or 7657 - try: - port = int(port_value) - except (TypeError, ValueError) as exc: - raise TlsError(f"LAN endpoint has invalid port: {port_value!r}") from exc - enrolled = EnrolledDevice(device_token="", identity=self._pl_identity) - return await Client.dial_direct(host, enrolled, port=port) - - async def _dial_relay(self) -> TunnelSession: - if self._pl_identity is None or self._pl_relay_url is None: - raise TlsError("PL relay is not configured") - if self._pl_enrolled is None: - self._pl_enrolled = Client.enroll_device( - self._pl_relay_url, - self._pl_identity, - ) - return await Client.dial(self._pl_relay_url, self._pl_enrolled) - def upload_segment( self, day: str, @@ -544,13 +377,11 @@ class ObserverClient: return UploadResult(False) body, content_type = encode_multipart_formdata(fields) - result = self._run_pl( - self._pl_request( - "POST", - "/app/observer/ingest", - headers={"Content-Type": content_type}, - body=body, - ) + result = self._pl_request( + "POST", + "/app/observer/ingest", + headers={"Content-Type": content_type}, + body=body, ) if result.status == 200: @@ -568,7 +399,13 @@ class ObserverClient: result.status, result.body.decode("utf-8", errors="replace"), ) - except (ConnectionError, OSError, StreamResetError, TlsError) as exc: + except ( + ConnectionError, + OSError, + StreamResetError, + TlsError, + TunnelRequestError, + ) as exc: logger.warning("PL upload attempt %s failed: %s", attempt + 1, exc) if attempt < len(RETRY_BACKOFF) - 1: time.sleep(delay) @@ -607,13 +444,11 @@ class ObserverClient: payload = {"tract": tract, "event": event, **fields} body = json.dumps(payload).encode("utf-8") try: - result = self._run_pl( - self._pl_request( - "POST", - "/app/observer/ingest/event", - headers={"Content-Type": "application/json"}, - body=body, - ) + result = self._pl_request( + "POST", + "/app/observer/ingest/event", + headers={"Content-Type": "application/json"}, + body=body, ) if result.status == 200: return True @@ -627,7 +462,13 @@ class ObserverClient: result.body.decode("utf-8", errors="replace"), ) return False - except (ConnectionError, OSError, StreamResetError, TlsError) as exc: + except ( + ConnectionError, + OSError, + StreamResetError, + TlsError, + TunnelRequestError, + ) as exc: logger.debug("PL event relay failed: %s", exc) return False @@ -720,7 +561,16 @@ class ObserverClient: backoff_index += 1 def _callosum_loop_pl(self, callback: Callable[[dict], None]) -> None: - if self._revoked or not self._pl_fingerprint_prefix: + if self._revoked: + return + try: + tunnel = self._pl_tunnel() + except Exception as exc: + self._callosum_error = exc + logger.debug("PL callosum tunnel setup failed: %s", exc) + return + if not self._pl_fingerprint_prefix: + self._callosum_error = RuntimeError("PL identity fingerprint not loaded") return path = OBSERVER_CALLOSUM_SSE_ROUTE.replace( @@ -731,10 +581,11 @@ class ObserverClient: while not self._callosum_stop.is_set(): chunks: queue.Queue[bytes | Exception | None] = queue.Queue() - loop = self._ensure_pl_loop() - future = asyncio.run_coroutine_threadsafe( - self._pl_callosum_reader(path, chunks), - loop, + future = tunnel.stream_request( + "GET", + path, + headers={"Accept": "text/event-stream"}, + chunks=chunks, ) data_lines: list[str] = [] text_buffer = "" @@ -750,6 +601,8 @@ class ObserverClient: break if isinstance(item, Exception): self._callosum_error = item + if isinstance(item, PermissionError): + self._revoked = True if self._revoked: return break @@ -775,37 +628,6 @@ class ObserverClient: if backoff_index < len(CALLOSUM_RECONNECT_BACKOFF) - 1: backoff_index += 1 - async def _pl_callosum_reader( - self, - path: str, - chunks: queue.Queue[bytes | Exception | None], - ) -> None: - try: - session = await self._get_pl_session() - status, _headers, initial_body, stream = await session.stream_request( - "GET", - path, - headers={"Accept": "text/event-stream"}, - ) - if status == 200: - if initial_body: - chunks.put(initial_body) - async for chunk in stream.read(): - chunks.put(chunk) - return - if status in {401, 403}: - self._revoked = True - chunks.put(RuntimeError(f"Callosum subscription rejected ({status})")) - return - chunks.put(RuntimeError(f"Callosum subscription failed ({status})")) - except (ConnectionError, OSError, StreamResetError, TlsError) as exc: - await self._close_pl_session() - chunks.put(exc) - except Exception as exc: - chunks.put(exc) - finally: - chunks.put(None) - def _consume_callosum_response( self, response: requests.Response, @@ -900,12 +722,7 @@ class ObserverClient: and self._callosum_thread is not threading.current_thread() ): self._callosum_thread.join(timeout=5.0) - if self._pl_loop is not None: - try: - self._run_pl(self._close_pl_session()) - except Exception as exc: - logger.debug("PL session close failed: %s", exc) - self._pl_loop.call_soon_threadsafe(self._pl_loop.stop) - if self._pl_loop_thread is not None and self._pl_loop_thread.is_alive(): - self._pl_loop_thread.join(timeout=5.0) + if self._tunnel is not None: + self._tunnel.close() + self._tunnel = None self._session.close() diff --git a/solstone/think/link/bundle.py b/solstone/think/link/bundle.py new file mode 100644 index 000000000..7ee7c12a0 --- /dev/null +++ b/solstone/think/link/bundle.py @@ -0,0 +1,64 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import json +from pathlib import Path + +from solstone.think.link.ca import cert_fingerprint +from solstone.think.link.client import ClientIdentity + +PL_BUNDLE_FILES = { + "private.pem", + "cert.pem", + "chain.pem", + "home_attestation.jwt", + "peer.json", +} + + +def endpoint_label(endpoint: dict[str, object]) -> str: + host = str(endpoint.get("ip") or endpoint.get("host") or "?") + port = endpoint.get("port") or 7657 + return f"lan-direct {host}:{port}" + + +def load_client_identity(bundle_dir: Path) -> ClientIdentity: + if not bundle_dir.is_dir(): + raise ValueError(f"PL bundle not found: {bundle_dir}") + + missing = sorted( + name for name in PL_BUNDLE_FILES if not (bundle_dir / name).exists() + ) + if missing: + raise ValueError( + "missing PL bundle file: " + + ", ".join(str(bundle_dir / name) for name in missing) + ) + + private_key_pem = (bundle_dir / "private.pem").read_text(encoding="utf-8") + client_cert_pem = (bundle_dir / "cert.pem").read_text(encoding="utf-8") + ca_chain_pem = (bundle_dir / "chain.pem").read_text(encoding="utf-8") + home_attestation = (bundle_dir / "home_attestation.jwt").read_text(encoding="utf-8") + try: + peer = json.loads((bundle_dir / "peer.json").read_text(encoding="utf-8")) + except json.JSONDecodeError as exc: + raise ValueError(f"invalid peer.json in {bundle_dir}: {exc}") from exc + + local_endpoints = peer.get("local_endpoints", []) + if local_endpoints is None: + local_endpoints = [] + if not isinstance(local_endpoints, list): + raise ValueError("peer.json local_endpoints must be a list") + + return ClientIdentity( + private_key_pem=private_key_pem, + client_cert_pem=client_cert_pem, + ca_chain_pem=ca_chain_pem, + fingerprint=cert_fingerprint(client_cert_pem), + home_instance_id=str(peer.get("instance_id") or ""), + home_label=str(peer.get("home_label") or ""), + home_attestation=home_attestation, + local_endpoints=tuple(local_endpoints), + ) diff --git a/solstone/think/link/client.py b/solstone/think/link/client.py index bb6baad81..99eb40c6e 100644 --- a/solstone/think/link/client.py +++ b/solstone/think/link/client.py @@ -54,16 +54,13 @@ from solstone.convey.secure_listener.framing import ( parse_window_credit, ) from solstone.think.link.ca import cert_fingerprint +from solstone.think.link.tls import TlsError as _TlsError LOG = logging.getLogger(__name__) _CONNECT_TIMEOUT_SECONDS = 15 _HTTP_TIMEOUT_SECONDS = 30 -class TlsError(RuntimeError): - """Raised when the client-side TLS handshake or tunnel aborts.""" - - class StreamResetError(ConnectionError): """Raised when the peer sends a RESET frame for an active stream.""" @@ -598,7 +595,7 @@ async def _open_tunnel_session( timeout=_CONNECT_TIMEOUT_SECONDS, ) if inbound is None: - raise TlsError("transport closed during TLS handshake") + raise _TlsError("transport closed during TLS handshake") outbound, plaintext = _drive_tls_client(tls, inbound=inbound) if outbound: await transport.send(outbound) @@ -692,7 +689,7 @@ def _drive_tls_client( except SSL.WantReadError: pass except SSL.Error as exc: - raise TlsError(f"send failed: {exc}") from exc + raise _TlsError(f"send failed: {exc}") from exc if not state.handshake_done: try: @@ -701,7 +698,7 @@ def _drive_tls_client( except SSL.WantReadError: pass except SSL.Error as exc: - raise TlsError(f"handshake failed: {exc}") from exc + raise _TlsError(f"handshake failed: {exc}") from exc plaintext_in = bytearray() if state.handshake_done: @@ -713,7 +710,7 @@ def _drive_tls_client( except SSL.ZeroReturnError: break except SSL.Error as exc: - raise TlsError(f"recv failed: {exc}") from exc + raise _TlsError(f"recv failed: {exc}") from exc if not chunk: break plaintext_in.extend(chunk) @@ -871,7 +868,6 @@ __all__ = [ "EncryptedTransport", "EnrolledDevice", "StreamResetError", - "TlsError", "TunnelSession", "_http_request_bytes", ] diff --git a/solstone/think/link/dialer.py b/solstone/think/link/dialer.py new file mode 100644 index 000000000..81a21dc85 --- /dev/null +++ b/solstone/think/link/dialer.py @@ -0,0 +1,302 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import asyncio +import queue +import threading +import time +from collections.abc import Awaitable +from concurrent.futures import Future +from typing import Any, Self + +from solstone.think.link.bundle import endpoint_label +from solstone.think.link.client import ( + Client, + ClientIdentity, + EnrolledDevice, + StreamResetError, + TunnelSession, +) +from solstone.think.link.tls import TlsError + + +class TunnelRequestError(ConnectionError): + def __init__(self, reason: str, detail: str) -> None: + super().__init__(f"{reason}: {detail}" if detail else reason) + self.reason = reason + self.detail = detail + + +async def _dial_direct_endpoint( + client: Client, + endpoint: dict[str, object], + identity: ClientIdentity, + deadline: float | None = None, +) -> TunnelSession: + host = str(endpoint.get("ip") or endpoint.get("host") or "").strip() + if not host: + raise TlsError("LAN endpoint missing ip") + port_value = endpoint.get("port") or 7657 + try: + port = int(port_value) + except (TypeError, ValueError) as exc: + raise TlsError(f"LAN endpoint has invalid port: {port_value!r}") from exc + enrolled = EnrolledDevice(device_token="", identity=identity) + return await _with_deadline( + client.dial_direct(host, enrolled, port=port), + deadline, + ) + + +async def _dial_relay( + client: Client, + relay_url: str, + identity: ClientIdentity, + deadline: float | None = None, +) -> TunnelSession: + enrolled = client.enroll_device(relay_url, identity) + return await _with_deadline(client.dial(relay_url, enrolled), deadline) + + +async def _with_deadline(coro: Awaitable[Any], deadline: float | None) -> Any: + if deadline is None: + return await coro + timeout = max(0.0, deadline - time.monotonic()) + return await asyncio.wait_for(coro, timeout=timeout) + + +async def open_tunnel( + identity: ClientIdentity, + relay_url: str | None, + *, + deadline: float | None = None, +) -> TunnelSession: + client = Client() + attempts: list[tuple[str, Any]] = [] + for endpoint in identity.local_endpoints: + label = endpoint_label(endpoint) + attempts.append( + (label, _dial_direct_endpoint(client, endpoint, identity, deadline)) + ) + if relay_url: + attempts.append( + ( + "spl-relay", + _dial_relay(client, relay_url.rstrip("/"), identity, deadline), + ) + ) + if not attempts: + raise TlsError("no PL dial attempts configured") + + tasks = {asyncio.create_task(coro): label for label, coro in attempts} + pending = set(tasks) + failures: dict[str, BaseException] = {} + + while pending: + done, pending = await asyncio.wait( + pending, + return_when=asyncio.FIRST_COMPLETED, + ) + for task in done: + label = tasks[task] + try: + session = task.result() + except BaseException as exc: + failures[label] = exc + continue + for loser in pending: + loser.cancel() + if pending: + await asyncio.gather(*pending, return_exceptions=True) + return session + + detail = "; ".join( + f"{label}: {type(exc).__name__}: {exc}" for label, exc in failures.items() + ) + raise TlsError(f"all PL dial attempts failed: {detail}") + + +class TunnelClient: + def __init__(self, identity: ClientIdentity, relay_url: str | None) -> None: + self._identity = identity + self._relay_url = relay_url.rstrip("/") if relay_url else None + self._loop: asyncio.AbstractEventLoop | None = None + self._loop_thread: threading.Thread | None = None + self._session: TunnelSession | None = None + self._session_lock: asyncio.Lock | None = None + self._closed = False + + def _ensure_loop(self) -> asyncio.AbstractEventLoop: + if self._closed: + raise TunnelRequestError("closed", "tunnel client is closed") + if self._loop is not None and self._loop.is_running(): + return self._loop + + loop = asyncio.new_event_loop() + ready = threading.Event() + + def run_loop() -> None: + asyncio.set_event_loop(loop) + self._session_lock = asyncio.Lock() + ready.set() + loop.run_forever() + + thread = threading.Thread( + target=run_loop, + name=f"link-tunnel-{self._identity.home_instance_id}", + daemon=True, + ) + thread.start() + ready.wait() + self._loop = loop + self._loop_thread = thread + return loop + + def _run(self, coro: Awaitable[Any]) -> Any: + loop = self._ensure_loop() + future = asyncio.run_coroutine_threadsafe(coro, loop) + return future.result() + + async def _get_session_async(self) -> TunnelSession: + if self._session_lock is None: + self._session_lock = asyncio.Lock() + async with self._session_lock: + if self._session is not None: + return self._session + self._session = await open_tunnel(self._identity, self._relay_url) + return self._session + + def _get_session(self) -> TunnelSession: + return self._run(self._get_session_async()) + + async def _close_session_async(self) -> None: + session = self._session + self._session = None + if session is not None: + await session.close() + + def _close_session(self) -> None: + if self._loop is None or not self._loop.is_running(): + self._session = None + return + self._run(self._close_session_async()) + + def request( + self, + method: str, + path: str, + *, + headers: dict[str, str] | None = None, + body: bytes = b"", + ) -> tuple[int, dict[str, str], bytes]: + try: + session = self._get_session() + return self._run( + session.request( + method, + path, + headers=headers or {}, + body=body, + ) + ) + except (ConnectionError, OSError, StreamResetError, TlsError) as exc: + self._close_session() + raise TunnelRequestError(type(exc).__name__, str(exc)) from exc + + def stream_request( + self, + method: str, + path: str, + *, + headers: dict[str, str] | None = None, + body: bytes = b"", + chunks: queue.Queue[bytes | Exception | None] | None = None, + ) -> Future[None] | tuple[int, dict[str, str], bytes, Any]: + if chunks is None: + return self._run( + self._stream_request_async( + method, + path, + headers=headers or {}, + body=body, + ) + ) + loop = self._ensure_loop() + return asyncio.run_coroutine_threadsafe( + self._stream_to_queue( + method, + path, + headers=headers or {}, + body=body, + chunks=chunks, + ), + loop, + ) + + async def _stream_request_async( + self, + method: str, + path: str, + *, + headers: dict[str, str], + body: bytes, + ) -> tuple[int, dict[str, str], bytes, Any]: + session = await self._get_session_async() + return await session.stream_request(method, path, headers=headers, body=body) + + async def _stream_to_queue( + self, + method: str, + path: str, + *, + headers: dict[str, str], + body: bytes, + chunks: queue.Queue[bytes | Exception | None], + ) -> None: + try: + status, _headers, initial_body, stream = await self._stream_request_async( + method, + path, + headers=headers, + body=body, + ) + if status == 200: + if initial_body: + chunks.put(initial_body) + async for chunk in stream.read(): + chunks.put(chunk) + return + if status in {401, 403}: + chunks.put(PermissionError(f"stream request rejected ({status})")) + return + chunks.put(RuntimeError(f"stream request failed ({status})")) + except (ConnectionError, OSError, StreamResetError, TlsError) as exc: + await self._close_session_async() + chunks.put(TunnelRequestError(type(exc).__name__, str(exc))) + except Exception as exc: + chunks.put(exc) + finally: + chunks.put(None) + + def close(self) -> None: + if self._closed: + return + if self._loop is not None and self._loop.is_running(): + try: + self._run(self._close_session_async()) + except Exception: + pass + self._loop.call_soon_threadsafe(self._loop.stop) + if self._loop_thread is not None and self._loop_thread.is_alive(): + self._loop_thread.join(timeout=5.0) + self._loop = None + self._loop_thread = None + self._closed = True + + def __enter__(self) -> Self: + return self + + def __exit__(self, *_exc: object) -> None: + self.close() diff --git a/solstone/think/link/tls.py b/solstone/think/link/tls.py new file mode 100644 index 000000000..c9f617383 --- /dev/null +++ b/solstone/think/link/tls.py @@ -0,0 +1,8 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + + +class TlsError(RuntimeError): + """Raised when the client-side TLS handshake or tunnel aborts.""" diff --git a/tests/link/test_bundle.py b/tests/link/test_bundle.py new file mode 100644 index 000000000..18d7f9605 --- /dev/null +++ b/tests/link/test_bundle.py @@ -0,0 +1,70 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest +from cryptography.hazmat.primitives import serialization + +from solstone.think.link.bundle import load_client_identity +from solstone.think.link.ca import cert_fingerprint, generate_ca + + +def _write_bundle(tmp_path: Path) -> Path: + bundle = tmp_path / "peer" + bundle.mkdir() + ca = generate_ca(tmp_path / "ca") + cert_pem = ca.cert.public_bytes(serialization.Encoding.PEM).decode("ascii") + (bundle / "private.pem").write_text("private", encoding="utf-8") + (bundle / "cert.pem").write_text(cert_pem, encoding="utf-8") + (bundle / "chain.pem").write_text(cert_pem, encoding="utf-8") + (bundle / "home_attestation.jwt").write_text("jwt", encoding="utf-8") + (bundle / "peer.json").write_text( + json.dumps( + { + "label": "host-a", + "instance_id": "12345678-1234-1234-1234-123456789abc", + "home_label": "solstone", + "local_endpoints": [{"ip": "127.0.0.1", "port": 7657}], + } + ), + encoding="utf-8", + ) + return bundle + + +def test_load_client_identity_happy_path(tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + cert_pem = (bundle / "cert.pem").read_text(encoding="utf-8") + + identity = load_client_identity(bundle) + + assert identity.private_key_pem == "private" + assert identity.fingerprint == cert_fingerprint(cert_pem) + assert identity.home_instance_id == "12345678-1234-1234-1234-123456789abc" + assert identity.home_label == "solstone" + assert identity.home_attestation == "jwt" + assert identity.local_endpoints == ({"ip": "127.0.0.1", "port": 7657},) + + +def test_load_client_identity_missing_file(tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + (bundle / "chain.pem").unlink() + + with pytest.raises(ValueError, match="missing PL bundle file") as exc_info: + load_client_identity(bundle) + + assert str(bundle / "chain.pem") in str(exc_info.value) + + +def test_load_client_identity_bad_peer_json(tmp_path: Path) -> None: + bundle = _write_bundle(tmp_path) + (bundle / "peer.json").write_text("{", encoding="utf-8") + + with pytest.raises(ValueError, match="invalid peer.json") as exc_info: + load_client_identity(bundle) + + assert str(bundle) in str(exc_info.value) diff --git a/tests/link/test_dialer_unit.py b/tests/link/test_dialer_unit.py index b5c8f055b..89aa86737 100644 --- a/tests/link/test_dialer_unit.py +++ b/tests/link/test_dialer_unit.py @@ -1,7 +1,7 @@ # SPDX-License-Identifier: AGPL-3.0-only # Copyright (c) 2026 sol pbc -"""Unit tests for observer PL dial orchestration.""" +"""Unit tests for paired-link dial orchestration.""" from __future__ import annotations @@ -9,16 +9,17 @@ import asyncio import pytest -from solstone.observe.observer_client import ObserverClient +from solstone.think.link import dialer from solstone.think.link.client import ( Client, ClientIdentity, EnrolledDevice, StreamResetError, - TlsError, TunnelSession, _http_request_bytes, ) +from solstone.think.link.dialer import TunnelClient, TunnelRequestError +from solstone.think.link.tls import TlsError def test_link_client_public_imports() -> None: @@ -49,30 +50,18 @@ def _identity(*, endpoints: tuple[dict[str, object], ...]) -> ClientIdentity: ) -def _client_for_identity(identity: ClientIdentity) -> ObserverClient: - client = object.__new__(ObserverClient) - client._pl_identity = identity - client._pl_relay_url = None - client._pl_enrolled = None - client._pl_session = None - client._pl_session_lock = None - return client - - @pytest.mark.asyncio async def test_lan_direct_race_picks_first_and_cancels_loser(monkeypatch) -> None: - client = _client_for_identity( - _identity( - endpoints=( - {"ip": "10.0.0.1", "port": 7657}, - {"ip": "10.0.0.2", "port": 7657}, - ) + identity = _identity( + endpoints=( + {"ip": "10.0.0.1", "port": 7657}, + {"ip": "10.0.0.2", "port": 7657}, ) ) cancelled: list[str] = [] winner = object() - async def dial_direct(endpoint: dict[str, object]): + async def dial_direct(_client, endpoint, _identity, _deadline=None): if endpoint["ip"] == "10.0.0.2": await asyncio.sleep(0) return winner @@ -82,30 +71,27 @@ async def test_lan_direct_race_picks_first_and_cancels_loser(monkeypatch) -> Non cancelled.append(str(endpoint["ip"])) raise - monkeypatch.setattr(client, "_dial_direct_endpoint", dial_direct) + monkeypatch.setattr(dialer, "_dial_direct_endpoint", dial_direct) - assert await client._open_tunnel() is winner + assert await dialer.open_tunnel(identity, None) is winner assert cancelled == ["10.0.0.1"] @pytest.mark.asyncio async def test_all_fail_error_names_every_attempt(monkeypatch) -> None: - client = _client_for_identity( - _identity(endpoints=({"ip": "10.0.0.1", "port": 7657},)) - ) - client._pl_relay_url = "https://relay.test" + identity = _identity(endpoints=({"ip": "10.0.0.1", "port": 7657},)) - async def dial_direct(_endpoint: dict[str, object]): + async def dial_direct(_client, _endpoint, _identity, _deadline=None): raise TlsError("lan failed") - async def dial_relay(): + async def dial_relay(_client, _relay_url, _identity, _deadline=None): raise OSError("relay failed") - monkeypatch.setattr(client, "_dial_direct_endpoint", dial_direct) - monkeypatch.setattr(client, "_dial_relay", dial_relay) + monkeypatch.setattr(dialer, "_dial_direct_endpoint", dial_direct) + monkeypatch.setattr(dialer, "_dial_relay", dial_relay) with pytest.raises(TlsError) as exc_info: - await client._open_tunnel() + await dialer.open_tunnel(identity, "https://relay.test") message = str(exc_info.value) assert "lan-direct 10.0.0.1:7657" in message @@ -114,11 +100,7 @@ async def test_all_fail_error_names_every_attempt(monkeypatch) -> None: assert "relay failed" in message -@pytest.mark.asyncio -async def test_cached_session_drops_on_stream_reset() -> None: - client = _client_for_identity(_identity(endpoints=())) - client._pl_session_lock = asyncio.Lock() - +def test_cached_session_drops_on_stream_reset(monkeypatch) -> None: class ResetSession: def __init__(self) -> None: self.closed = False @@ -130,10 +112,18 @@ async def test_cached_session_drops_on_stream_reset() -> None: self.closed = True session = ResetSession() - client._pl_session = session - with pytest.raises(StreamResetError): - await client._pl_request("GET", "/") + async def open_tunnel(_identity, _relay_url): + return session + + monkeypatch.setattr(dialer, "open_tunnel", open_tunnel) + client = TunnelClient(_identity(endpoints=()), None) + try: + with pytest.raises(TunnelRequestError) as exc_info: + client.request("GET", "/") + finally: + client.close() + assert exc_info.value.reason == "StreamResetError" assert session.closed is True - assert client._pl_session is None + assert client._session is None diff --git a/tests/link/test_lan_direct.py b/tests/link/test_lan_direct.py index 47bb0533b..e502e60bf 100644 --- a/tests/link/test_lan_direct.py +++ b/tests/link/test_lan_direct.py @@ -13,7 +13,8 @@ from pathlib import Path import pytest -from tests.link.client import Client, EnrolledDevice, TlsError +from solstone.think.link.tls import TlsError +from tests.link.client import Client, EnrolledDevice from tests.link.live_helpers import ( CONVEY_PASSWORD, list_devices, -- 2.51.2