From d18a7c02359cd827d0ff15058861de5c2600a96f Mon Sep 17 00:00:00 2001 From: Jer Miller Date: Sun, 19 Apr 2026 20:01:29 -0600 Subject: [PATCH] link: add live integration test for pair+dial roundtrip --- Makefile | 20 ++- tests/link/client.py | 5 +- tests/link/live_helpers.py | 226 +++++++++++++++++++++++++++++++++ tests/link/test_integration.py | 101 +++++++++++++++ 4 files changed, 345 insertions(+), 7 deletions(-) create mode 100644 tests/link/live_helpers.py create mode 100644 tests/link/test_integration.py diff --git a/Makefile b/Makefile index fc6c3131f..162325d64 100644 --- a/Makefile +++ b/Makefile @@ -275,6 +275,7 @@ review: .installed # Test environment - use fixtures journal for all tests TEST_ENV = _SOLSTONE_JOURNAL_OVERRIDE=tests/fixtures/journal +LINK_LIVE_TESTS = --ignore=tests/link/test_integration.py --ignore=tests/link/test_privacy_scan.py # Venv tool shortcuts PYTEST := $(VENV_BIN)/pytest @@ -288,7 +289,7 @@ format-check: .installed # Run core tests (excluding integration and app tests) test: .installed format-check @echo "Running core tests..." - $(TEST_ENV) $(PYTEST) tests/ -q --cov=. --ignore=tests/integration + $(TEST_ENV) $(PYTEST) tests/ -q --cov=. --ignore=tests/integration $(LINK_LIVE_TESTS) # Run app tests test-apps: .installed @@ -317,7 +318,9 @@ test-only: .installed # Run integration tests test-integration: .installed @echo "Running integration tests..." - $(TEST_ENV) $(PYTEST) tests/integration/ -v --tb=short --timeout=20 + @STATUS=0; \ + $(TEST_ENV) $(PYTEST) tests/integration/ tests/link/test_integration.py tests/link/test_privacy_scan.py -v --tb=short --timeout=20 || STATUS=$$?; \ + if [ "$$STATUS" -ne 0 ] && [ "$$STATUS" -ne 5 ]; then exit $$STATUS; fi # Run specific integration test test-integration-only: .installed @@ -326,12 +329,19 @@ test-integration-only: .installed echo "Example: make test-integration-only TEST=test_api.py"; \ exit 1; \ fi - $(TEST_ENV) $(PYTEST) tests/integration/$(TEST) --timeout=20 + @TARGET="$(TEST)"; \ + case "$$TARGET" in \ + tests/*|-*) ;; \ + *) TARGET="tests/integration/$$TARGET" ;; \ + esac; \ + STATUS=0; \ + $(TEST_ENV) $(PYTEST) "$$TARGET" --timeout=20 || STATUS=$$?; \ + if [ "$$STATUS" -ne 0 ] && [ "$$STATUS" -ne 5 ]; then exit $$STATUS; fi # Run all tests (core + apps + integration) test-all: .installed @echo "Running all tests (core + apps + integration)..." - $(TEST_ENV) $(PYTEST) tests/ -v --cov=. && $(TEST_ENV) $(PYTEST) apps/ -v --cov=. --cov-append + $(TEST_ENV) $(PYTEST) tests/ -v --cov=. --ignore=tests/integration $(LINK_LIVE_TESTS) && $(TEST_ENV) $(PYTEST) apps/ -v --cov=. --cov-append # Auto-format and fix code, then report any remaining issues format: .installed @@ -491,7 +501,7 @@ watch: .installed # Generate coverage report (core + apps, excluding core integration tests) coverage: .installed - $(TEST_ENV) $(PYTEST) tests/ --cov=. --cov-report=html --cov-report=term --ignore=tests/integration + $(TEST_ENV) $(PYTEST) tests/ --cov=. --cov-report=html --cov-report=term --ignore=tests/integration $(LINK_LIVE_TESTS) $(TEST_ENV) $(PYTEST) apps/ --cov=. --cov-report=html --cov-report=term --cov-append @echo "Coverage report generated in htmlcov/index.html" diff --git a/tests/link/client.py b/tests/link/client.py index 15bf5de43..f78a62212 100644 --- a/tests/link/client.py +++ b/tests/link/client.py @@ -231,6 +231,8 @@ class _DialerMultiplexer: return self._closed = True for state in list(self._streams.values()): + if state.reset_reason is None: + state.reset_reason = RESET_INTERNAL_ERROR self._close_stream(state, forget=True) async def _dispatch(self, frame: Frame) -> None: @@ -361,8 +363,7 @@ class TunnelSession: if self._closed.is_set(): return self._mux.close() - if not self._ws.closed: - await self._ws.close() + await self._ws.close() await self._reader_task self._closed.set() diff --git a/tests/link/live_helpers.py b/tests/link/live_helpers.py new file mode 100644 index 000000000..5acc51eb5 --- /dev/null +++ b/tests/link/live_helpers.py @@ -0,0 +1,226 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import contextlib +import json +import os +import queue +import subprocess +import sys +import threading +import time +from collections.abc import Iterator +from pathlib import Path +from typing import Any + +import pytest +import requests +from werkzeug.security import generate_password_hash +from werkzeug.serving import make_server + +RELAY_URL = "https://spl-relay-staging.jer-3f2.workers.dev" +CONVEY_PASSWORD = "pytest-link-pass" +_READY_LINE = "listen WS open" + + +def skip_unless_live_relay() -> None: + if os.environ.get("SPL_RELAY_LIVE_TESTS", "1") == "0": + pytest.skip( + "SPL_RELAY_LIVE_TESTS=0; skipping live relay tests", + allow_module_level=True, + ) + try: + response = requests.get(f"{RELAY_URL}/", timeout=5) + if response.status_code != 200: + pytest.skip( + f"relay unreachable: status {response.status_code}", + allow_module_level=True, + ) + except Exception as exc: # noqa: BLE001 + pytest.skip(f"relay unreachable: {exc}", allow_module_level=True) + + +class LinkProcessCapture: + def __init__(self, proc: subprocess.Popen[str]) -> None: + self.proc = proc + self.stdout_lines: list[str] = [] + self.stderr_lines: list[str] = [] + self._queue: queue.Queue[tuple[str, str]] = queue.Queue() + self._threads = [ + threading.Thread( + target=self._drain, + args=(proc.stdout, self.stdout_lines, "stdout", self._queue), + daemon=True, + ), + threading.Thread( + target=self._drain, + args=(proc.stderr, self.stderr_lines, "stderr", self._queue), + daemon=True, + ), + ] + for thread in self._threads: + thread.start() + + def wait_for_line(self, needle: str, timeout: float) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if any(needle in line for line in self.stderr_lines): + return + if self.proc.poll() is not None: + break + remaining = max(0.0, deadline - time.monotonic()) + try: + stream, line = self._queue.get(timeout=min(0.25, remaining)) + except queue.Empty: + continue + if stream == "stderr" and needle in line: + return + raise RuntimeError( + f"link service never emitted {needle!r}; stderr tail:\n" + + "".join(self.stderr_lines[-50:]) + ) + + @property + def stdout_text(self) -> str: + return "".join(self.stdout_lines) + + @property + def stderr_text(self) -> str: + return "".join(self.stderr_lines) + + def stop(self) -> None: + if self.proc.poll() is None: + self.proc.terminate() + try: + self.proc.wait(timeout=5) + except subprocess.TimeoutExpired: + self.proc.kill() + self.proc.wait(timeout=5) + for pipe in (self.proc.stdout, self.proc.stderr): + if pipe is not None: + pipe.close() + for thread in self._threads: + thread.join(timeout=1) + + @staticmethod + def _drain( + pipe: Any, + target: list[str], + stream: str, + line_queue: queue.Queue[tuple[str, str]], + ) -> None: + if pipe is None: + return + try: + for line in pipe: + target.append(line) + line_queue.put((stream, line)) + finally: + return + + +@contextlib.contextmanager +def running_convey_server(journal_path: Path) -> Iterator[str]: + from convey import create_app + + _prepare_journal(journal_path) + app = create_app(str(journal_path)) + server = make_server("127.0.0.1", 0, app) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield f"http://127.0.0.1:{server.server_port}" + finally: + server.shutdown() + thread.join(timeout=5) + + +@contextlib.contextmanager +def running_link_service( + journal_path: Path, + *, + relay_url: str = RELAY_URL, +) -> Iterator[LinkProcessCapture]: + _prepare_journal(journal_path) + repo_root = Path(__file__).resolve().parents[2] + sol_bin = Path(sys.executable).with_name("sol") + env = os.environ.copy() + env["_SOLSTONE_JOURNAL_OVERRIDE"] = str(journal_path) + env["SOL_LINK_RELAY_URL"] = relay_url + env["SOL_SKIP_SUPERVISOR_CHECK"] = "1" + env["PYTHONUNBUFFERED"] = "1" + env["PATH"] = f"{repo_root / '.venv' / 'bin'}:{env.get('PATH', '')}" + proc = subprocess.Popen( + [str(sol_bin), "link", "-v"], + cwd=repo_root, + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + bufsize=1, + ) + capture = LinkProcessCapture(proc) + try: + capture.wait_for_line(_READY_LINE, timeout=15) + yield capture + finally: + capture.stop() + + +def list_devices(base_url: str) -> list[dict[str, Any]]: + response = requests.get(f"{base_url}/app/link/api/devices", timeout=10) + response.raise_for_status() + payload = response.json() + assert isinstance(payload, dict) + devices = payload.get("devices") + assert isinstance(devices, list) + return devices + + +def unpair_device(base_url: str, fingerprint: str) -> dict[str, Any]: + response = requests.post( + f"{base_url}/app/link/unpair", + json={"fingerprint": fingerprint}, + timeout=10, + ) + response.raise_for_status() + payload = response.json() + assert isinstance(payload, dict) + return payload + + +def runtime_texts( + journal_path: Path, + capture: LinkProcessCapture, +) -> dict[str, str]: + out = { + "stdout": capture.stdout_text, + "stderr": capture.stderr_text, + } + extra_paths = [journal_path / "health" / "supervisor.log"] + extra_paths.extend(sorted((journal_path / "link").glob("*.log"))) + for path in extra_paths: + if path.exists(): + out[str(path.relative_to(journal_path))] = path.read_text("utf-8") + return out + + +def _prepare_journal(journal_path: Path) -> None: + config_path = journal_path / "config" / "journal.json" + config_path.parent.mkdir(parents=True, exist_ok=True) + config_path.write_text( + json.dumps( + { + "convey": { + "password_hash": generate_password_hash(CONVEY_PASSWORD), + "trust_localhost": True, + }, + "setup": {"completed_at": 1}, + }, + indent=2, + ) + + "\n", + encoding="utf-8", + ) diff --git a/tests/link/test_integration.py b/tests/link/test_integration.py new file mode 100644 index 000000000..e2a24390b --- /dev/null +++ b/tests/link/test_integration.py @@ -0,0 +1,101 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import asyncio +import base64 +import json +from pathlib import Path + +import pytest + +from tests.link.client import Client, StreamResetError +from tests.link.live_helpers import ( + CONVEY_PASSWORD, + RELAY_URL, + list_devices, + running_convey_server, + running_link_service, + skip_unless_live_relay, + unpair_device, +) + +pytestmark = pytest.mark.integration +skip_unless_live_relay() + + +@pytest.mark.asyncio +@pytest.mark.timeout(60) +async def test_pair_enroll_dial_roundtrip( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + tmp_journal = tmp_path / "journal" + tmp_journal.mkdir() + monkeypatch.setenv("_SOLSTONE_JOURNAL_OVERRIDE", str(tmp_journal)) + + with ( + running_convey_server(tmp_journal) as base_url, + running_link_service(tmp_journal), + ): + identity = Client.pair(base_url, device_label="pytest-device") + assert identity.home_instance_id + assert identity.client_cert_pem.startswith("-----BEGIN CERTIFICATE-----") + assert len(identity.home_attestation.split(".")) == 3 + + enrolled = Client.enroll_device(RELAY_URL, identity) + assert enrolled.device_token + + before = next( + device + for device in list_devices(base_url) + if device["fingerprint"] == identity.fingerprint + ) + assert before["last_seen_at"] is None + + session = await Client.dial(RELAY_URL, enrolled) + async with session: + auth = base64.b64encode(f":{CONVEY_PASSWORD}".encode("utf-8")).decode( + "ascii" + ) + status, headers, body = await session.request( + "GET", + "/", + headers={"authorization": f"Basic {auth}"}, + ) + assert status == 302 + assert headers["location"] == "/app/home/" + + status, headers, body = await session.request( + "GET", + "/app/link/api/status", + headers={"authorization": f"Basic {auth}"}, + ) + assert status == 200 + assert headers["content-type"] == "application/json" + status_payload = json.loads(body) + assert status_payload["instance_id"] == identity.home_instance_id + + after = next( + device + for device in list_devices(base_url) + if device["fingerprint"] == identity.fingerprint + ) + assert after["last_seen_at"] is not None + + unpaired = unpair_device(base_url, identity.fingerprint) + assert unpaired["unpaired"] == identity.fingerprint + + await asyncio.sleep(1) + failed = await Client.dial(RELAY_URL, enrolled) + async with failed: + with pytest.raises(StreamResetError): + auth = base64.b64encode(f":{CONVEY_PASSWORD}".encode("utf-8")).decode( + "ascii" + ) + await failed.request( + "GET", + "/app/link/api/status", + headers={"authorization": f"Basic {auth}"}, + ) -- 2.51.2