diff --git a/tests/test_pairing_devices.py b/tests/test_pairing_devices.py new file mode 100644 index 000000000..6a9d70c29 --- /dev/null +++ b/tests/test_pairing_devices.py @@ -0,0 +1,161 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import json + +from think.pairing import devices + + +def _devices_path(journal_copy): + return journal_copy / "config" / "paired_devices.json" + + +def test_load_devices_returns_empty_for_missing_store(journal_copy): + assert devices.load_devices() == [] + + +def test_register_load_and_remove_round_trip(journal_copy): + device = devices.register_device( + name="Phone", + platform="ios", + public_key="ssh-ed25519 AAAAphone", + session_key_hash="sha256:abc", + bundle_id="org.solpbc.solstone-swift", + app_version="0.1.0", + paired_at="2026-04-20T15:31:02Z", + ) + + loaded = devices.load_devices() + + assert loaded == [ + { + "id": device["id"], + "name": "Phone", + "platform": "ios", + "public_key": "ssh-ed25519 AAAAphone", + "session_key_hash": "sha256:abc", + "bundle_id": "org.solpbc.solstone-swift", + "app_version": "0.1.0", + "paired_at": "2026-04-20T15:31:02Z", + "last_seen_at": None, + } + ] + assert devices.find_device_by_id(device["id"]) == loaded[0] + assert devices.find_device_by_session_key_hash("sha256:abc") == loaded[0] + assert devices.remove_device(device["id"]) is True + assert devices.load_devices() == [] + + +def test_register_device_upserts_by_public_key(journal_copy): + first = devices.register_device( + name="Phone", + platform="ios", + public_key="ssh-ed25519 AAAAsame", + session_key_hash="sha256:first", + bundle_id="org.solpbc.solstone-swift", + app_version="0.1.0", + paired_at="2026-04-20T15:31:02Z", + ) + second = devices.register_device( + name="Phone 2", + platform="ios", + public_key="ssh-ed25519 AAAAsame", + session_key_hash="sha256:second", + bundle_id="org.solpbc.solstone-swift", + app_version="0.2.0", + paired_at="2026-04-20T16:00:00Z", + ) + + assert first["id"] == second["id"] + assert devices.load_devices() == [ + { + "id": first["id"], + "name": "Phone 2", + "platform": "ios", + "public_key": "ssh-ed25519 AAAAsame", + "session_key_hash": "sha256:second", + "bundle_id": "org.solpbc.solstone-swift", + "app_version": "0.2.0", + "paired_at": "2026-04-20T16:00:00Z", + "last_seen_at": None, + } + ] + + +def test_touch_last_seen_updates_existing_device(journal_copy): + device = devices.register_device( + name="Phone", + platform="ios", + public_key="ssh-ed25519 AAAAseen", + session_key_hash="sha256:seen", + bundle_id="org.solpbc.solstone-swift", + app_version="0.1.0", + paired_at="2026-04-20T15:31:02Z", + ) + + touched = devices.touch_last_seen(device["id"], last_seen_at="2026-04-20T16:01:00Z") + + assert touched is True + assert devices.find_device_by_id(device["id"]) is not None + assert ( + devices.find_device_by_id(device["id"])["last_seen_at"] + == "2026-04-20T16:01:00Z" + ) + + +def test_load_devices_recovers_from_malformed_store(journal_copy, caplog): + path = _devices_path(journal_copy) + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text('{"devices": "bad"}', encoding="utf-8") + + assert devices.load_devices() == [] + assert "paired device store unreadable" in caplog.text + + healed = devices.register_device( + name="Phone", + platform="ios", + public_key="ssh-ed25519 AAAAheal", + session_key_hash="sha256:heal", + bundle_id="org.solpbc.solstone-swift", + app_version="0.1.0", + paired_at="2026-04-20T15:31:02Z", + ) + payload = json.loads(path.read_text(encoding="utf-8")) + + assert payload == { + "devices": [ + { + "id": healed["id"], + "name": "Phone", + "platform": "ios", + "public_key": "ssh-ed25519 AAAAheal", + "session_key_hash": "sha256:heal", + "bundle_id": "org.solpbc.solstone-swift", + "app_version": "0.1.0", + "paired_at": "2026-04-20T15:31:02Z", + "last_seen_at": None, + } + ] + } + + +def test_status_view_redacts_secret_fields(journal_copy): + device = devices.register_device( + name="Phone", + platform="ios", + public_key="ssh-ed25519 AAAAredact", + session_key_hash="sha256:redact", + bundle_id="org.solpbc.solstone-swift", + app_version="0.1.0", + paired_at="2026-04-20T15:31:02Z", + ) + + assert devices.status_view(device) == { + "id": device["id"], + "name": "Phone", + "platform": "ios", + "paired_at": "2026-04-20T15:31:02Z", + "last_seen_at": None, + } diff --git a/tests/test_pairing_keys.py b/tests/test_pairing_keys.py new file mode 100644 index 000000000..e54d7fa63 --- /dev/null +++ b/tests/test_pairing_keys.py @@ -0,0 +1,63 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec, ed25519, rsa + +from think.pairing import keys + + +def _openssh_public_key(public_key) -> str: + return public_key.public_bytes( + encoding=serialization.Encoding.OpenSSH, + format=serialization.PublicFormat.OpenSSH, + ).decode("utf-8") + + +def test_validate_public_key_accepts_ssh_ed25519(): + key = ed25519.Ed25519PrivateKey.generate().public_key() + encoded = _openssh_public_key(key) + + assert keys.validate_public_key(encoded) == encoded + + +def test_validate_public_key_rejects_non_ed25519_algorithms(): + rsa_key = rsa.generate_private_key( + public_exponent=65537, key_size=2048 + ).public_key() + ecdsa_key = ec.generate_private_key(ec.SECP256R1()).public_key() + + for candidate in (_openssh_public_key(rsa_key), _openssh_public_key(ecdsa_key)): + try: + keys.validate_public_key(candidate) + except ValueError as exc: + assert str(exc) == "public key must be ssh-ed25519" + else: + raise AssertionError("expected non-ed25519 key to be rejected") + + +def test_validate_public_key_rejects_malformed_and_oversized_values(): + for candidate, expected in ( + ("ssh-ed25519 AAAA-not-valid", "public key is invalid"), + ("ssh-ed25519 " + ("A" * 2049), "public key is too long"), + ): + try: + keys.validate_public_key(candidate) + except ValueError as exc: + assert str(exc) == expected + else: + raise AssertionError("expected invalid public key to be rejected") + + +def test_session_key_helpers(): + session_key = keys.generate_session_key() + session_hash = keys.hash_session_key(session_key) + + assert session_key.startswith("dsk_") + assert session_hash.startswith("sha256:") + assert len(session_hash) == len("sha256:") + 64 + assert keys.mask_session_key(session_key) == ( + f"...{session_key[-4:]} (len={len(session_key)})" + ) diff --git a/tests/test_pairing_tokens.py b/tests/test_pairing_tokens.py new file mode 100644 index 000000000..dd1bdf04e --- /dev/null +++ b/tests/test_pairing_tokens.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +from think.pairing import tokens + + +def setup_function() -> None: + tokens._TOKENS.clear() + + +def test_create_token_uses_expected_shape_and_metadata(): + token = tokens.create_token(ttl_seconds=600, now=1000) + + assert token.token.startswith("ptk_") + assert token.issued_at == 1000 + assert token.expires_at == 1600 + assert token.ttl_seconds == 600 + assert token.consumed_at is None + + +def test_create_token_clamps_ttl(): + low = tokens.create_token(ttl_seconds=1, now=1000) + high = tokens.create_token(ttl_seconds=9999, now=1000) + + assert low.ttl_seconds == 60 + assert low.expires_at == 1060 + assert high.ttl_seconds == 3600 + assert high.expires_at == 4600 + + +def test_consume_token_marks_token_used_once(): + created = tokens.create_token(ttl_seconds=600, now=1000) + + first = tokens.consume_token(created.token, now=1100) + second = tokens.consume_token(created.token, now=1101) + peeked = tokens.peek_token(created.token, now=1101) + + assert first is not None + assert first.consumed_at == 1100 + assert second is None + assert peeked is not None + assert peeked.consumed_at == 1100 + + +def test_consume_token_rejects_expired_token(): + created = tokens.create_token(ttl_seconds=60, now=1000) + + assert tokens.consume_token(created.token, now=1060) is None + expired = tokens.peek_token(created.token, now=1060) + assert expired is not None + assert expired.expires_at == 1060 + assert expired.consumed_at is None + + +def test_purge_expired_tokens_removes_only_expired_entries(): + keep = tokens.create_token(ttl_seconds=600, now=1000) + drop = tokens.create_token(ttl_seconds=60, now=1000) + + purged = tokens.purge_expired_tokens(now=1060) + + assert purged == 1 + assert tokens.peek_token(drop.token, now=1060) is None + assert tokens.peek_token(keep.token, now=1060) is not None diff --git a/think/pairing/devices.py b/think/pairing/devices.py new file mode 100644 index 000000000..fc347c393 --- /dev/null +++ b/think/pairing/devices.py @@ -0,0 +1,213 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +"""Paired-device storage.""" + +from __future__ import annotations + +import json +import logging +import secrets +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, TypedDict + +from think.entities.core import atomic_write +from think.utils import get_journal + +logger = logging.getLogger(__name__) + + +class Device(TypedDict): + id: str + name: str + platform: str + public_key: str + session_key_hash: str + bundle_id: str + app_version: str + paired_at: str + last_seen_at: str | None + + +def _devices_path() -> Path: + return Path(get_journal()) / "config" / "paired_devices.json" + + +def _utc_now_iso() -> str: + return ( + datetime.now(timezone.utc).isoformat(timespec="seconds").replace("+00:00", "Z") + ) + + +def _empty_store() -> dict[str, list[Device]]: + return {"devices": []} + + +def _clean_str(value: Any) -> str: + return str(value or "").strip() + + +def _validate_timestamp(value: Any) -> str | None: + if value is None: + return None + if not isinstance(value, str) or not value.strip(): + raise ValueError("device timestamp fields must be strings or null") + return value + + +def _validate_store(payload: Any) -> list[Device]: + if not isinstance(payload, dict): + raise ValueError("paired device store must be a JSON object") + devices = payload.get("devices") + if not isinstance(devices, list): + raise ValueError("paired device store must contain a devices list") + normalized: list[Device] = [] + for device in devices: + if not isinstance(device, dict): + raise ValueError("paired device rows must be JSON objects") + normalized.append( + Device( + id=_require_field(device, "id"), + name=_require_field(device, "name"), + platform=_require_field(device, "platform"), + public_key=_require_field(device, "public_key"), + session_key_hash=_require_field(device, "session_key_hash"), + bundle_id=_require_field(device, "bundle_id"), + app_version=_require_field(device, "app_version"), + paired_at=_require_field(device, "paired_at"), + last_seen_at=_validate_timestamp(device.get("last_seen_at")), + ) + ) + return normalized + + +def _require_field(device: dict[str, Any], field: str) -> str: + value = _clean_str(device.get(field)) + if not value: + raise ValueError(f"paired device row missing required field: {field}") + return value + + +def _read_store() -> list[Device]: + path = _devices_path() + if not path.exists(): + return [] + try: + payload = json.loads(path.read_text(encoding="utf-8")) + return _validate_store(payload) + except Exception as exc: + logger.warning("paired device store unreadable path=%s error=%s", path, exc) + return [] + + +def _write_store(devices: list[Device]) -> None: + payload = json.dumps({"devices": devices}, indent=2, ensure_ascii=False) + "\n" + atomic_write(_devices_path(), payload, prefix=".paired_devices_") + + +def load_devices() -> list[Device]: + return _read_store() + + +def find_device_by_id(device_id: str) -> Device | None: + target = _clean_str(device_id) + if not target: + return None + for device in load_devices(): + if device["id"] == target: + return device + return None + + +def find_device_by_session_key_hash(session_key_hash: str) -> Device | None: + target = _clean_str(session_key_hash) + if not target: + return None + for device in load_devices(): + if device["session_key_hash"] == target: + return device + return None + + +def register_device( + *, + name: str, + platform: str, + public_key: str, + session_key_hash: str, + bundle_id: str, + app_version: str, + paired_at: str | None = None, +) -> Device: + devices = load_devices() + row: Device = Device( + id=f"dev_{secrets.token_urlsafe(16)}", + name=_clean_str(name), + platform=_clean_str(platform), + public_key=_clean_str(public_key), + session_key_hash=_clean_str(session_key_hash), + bundle_id=_clean_str(bundle_id), + app_version=_clean_str(app_version), + paired_at=_clean_str(paired_at) or _utc_now_iso(), + last_seen_at=None, + ) + for index, device in enumerate(devices): + if device["public_key"] != row["public_key"]: + continue + row["id"] = device["id"] + devices[index] = row + _write_store(devices) + return row + devices.append(row) + _write_store(devices) + return row + + +def touch_last_seen(device_id: str, *, last_seen_at: str | None = None) -> bool: + target = _clean_str(device_id) + if not target: + return False + devices = load_devices() + timestamp = _clean_str(last_seen_at) or _utc_now_iso() + for device in devices: + if device["id"] != target: + continue + device["last_seen_at"] = timestamp + _write_store(devices) + return True + return False + + +def remove_device(device_id: str) -> bool: + target = _clean_str(device_id) + if not target: + return False + devices = load_devices() + remaining = [device for device in devices if device["id"] != target] + if len(remaining) == len(devices): + return False + _write_store(remaining) + return True + + +def status_view(device: Device) -> dict[str, Any]: + return { + "id": device["id"], + "name": device["name"], + "platform": device["platform"], + "paired_at": device["paired_at"], + "last_seen_at": device["last_seen_at"], + } + + +__all__ = [ + "Device", + "find_device_by_id", + "find_device_by_session_key_hash", + "load_devices", + "register_device", + "remove_device", + "status_view", + "touch_last_seen", +] diff --git a/think/pairing/keys.py b/think/pairing/keys.py new file mode 100644 index 000000000..24aed3c59 --- /dev/null +++ b/think/pairing/keys.py @@ -0,0 +1,55 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +"""Pairing key validation and session-key helpers.""" + +from __future__ import annotations + +import hashlib +import secrets + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey + +MAX_PUBLIC_KEY_LENGTH = 2048 + + +def validate_public_key(public_key: str) -> str: + candidate = str(public_key or "").strip() + if not candidate: + raise ValueError("public key is required") + if len(candidate) > MAX_PUBLIC_KEY_LENGTH: + raise ValueError("public key is too long") + try: + parsed = serialization.load_ssh_public_key(candidate.encode("utf-8")) + except Exception as exc: + raise ValueError("public key is invalid") from exc + if not isinstance(parsed, Ed25519PublicKey): + raise ValueError("public key must be ssh-ed25519") + return parsed.public_bytes( + encoding=serialization.Encoding.OpenSSH, + format=serialization.PublicFormat.OpenSSH, + ).decode("utf-8") + + +def generate_session_key() -> str: + return f"dsk_{secrets.token_urlsafe(32)}" + + +def hash_session_key(session_key: str) -> str: + digest = hashlib.sha256(session_key.encode("utf-8")).hexdigest() + return f"sha256:{digest}" + + +def mask_session_key(session_key: str) -> str: + value = str(session_key or "") + return f"...{value[-4:]} (len={len(value)})" + + +__all__ = [ + "MAX_PUBLIC_KEY_LENGTH", + "generate_session_key", + "hash_session_key", + "mask_session_key", + "validate_public_key", +] diff --git a/think/pairing/tokens.py b/think/pairing/tokens.py new file mode 100644 index 000000000..884931487 --- /dev/null +++ b/think/pairing/tokens.py @@ -0,0 +1,105 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +"""In-memory pairing token store.""" + +from __future__ import annotations + +import secrets +import threading +import time +from dataclasses import dataclass, replace + +from think.pairing.config import ( + MAX_TOKEN_TTL_SECONDS, + MIN_TOKEN_TTL_SECONDS, + get_token_ttl_seconds, +) + + +@dataclass(frozen=True) +class PairingToken: + token: str + issued_at: int + expires_at: int + ttl_seconds: int + consumed_at: int | None + + +_TOKENS: dict[str, PairingToken] = {} +_TOKENS_LOCK = threading.Lock() + + +def _now(now: int | None) -> int: + return int(now if now is not None else time.time()) + + +def _clamp_ttl(ttl_seconds: int) -> int: + return max(MIN_TOKEN_TTL_SECONDS, min(MAX_TOKEN_TTL_SECONDS, ttl_seconds)) + + +def _purge_expired_locked(now: int, *, exclude: str | None = None) -> int: + expired = [ + token + for token, entry in _TOKENS.items() + if token != exclude and entry.expires_at <= now + ] + for token in expired: + del _TOKENS[token] + return len(expired) + + +def create_token( + *, ttl_seconds: int | None = None, now: int | None = None +) -> PairingToken: + ts = _now(now) + effective_ttl = _clamp_ttl( + get_token_ttl_seconds() if ttl_seconds is None else int(ttl_seconds) + ) + entry = PairingToken( + token=f"ptk_{secrets.token_urlsafe(32)}", + issued_at=ts, + expires_at=ts + effective_ttl, + ttl_seconds=effective_ttl, + consumed_at=None, + ) + with _TOKENS_LOCK: + _purge_expired_locked(ts) + _TOKENS[entry.token] = entry + return entry + + +def consume_token(token: str, *, now: int | None = None) -> PairingToken | None: + ts = _now(now) + with _TOKENS_LOCK: + _purge_expired_locked(ts, exclude=token) + entry = _TOKENS.get(token) + if entry is None: + return None + if entry.expires_at <= ts or entry.consumed_at is not None: + return None + consumed = replace(entry, consumed_at=ts) + _TOKENS[token] = consumed + return consumed + + +def peek_token(token: str, *, now: int | None = None) -> PairingToken | None: + ts = _now(now) + with _TOKENS_LOCK: + _purge_expired_locked(ts, exclude=token) + return _TOKENS.get(token) + + +def purge_expired_tokens(*, now: int | None = None) -> int: + ts = _now(now) + with _TOKENS_LOCK: + return _purge_expired_locked(ts) + + +__all__ = [ + "PairingToken", + "consume_token", + "create_token", + "peek_token", + "purge_expired_tokens", +]