summaryrefslogtreecommitdiffstats
path: root/packages/meshbay-node/src/meshbay_node/hub_client.py
blob: d91b945f0e151105f6369846bcf902a3f73075d1 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
"""
MeshBay Node — Hub client.

Handles all communication from the node to a Mesh Hub:
  - User registration (first run)
  - Login → JWT (access token + refresh token)
  - JWT offline verification and auto-refresh
  - Node announcement (endpoint_hint)
  - GEK bundle retrieval for a group
  - User public key lookup (for GEK wrapping)

JWT verification is done locally using the hub's cached Ed25519 public key.
The hub is only contacted for login and refresh — not for every request.
"""

import base64
import json
import logging
import time
from dataclasses import dataclass, field
from pathlib import Path

import httpx
import jwt
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey
from cryptography.hazmat.primitives import serialization

from meshbay_common.crypto import pk_to_b64, unwrap_gek
from meshbay_node.keystore import NodeKeys

log = logging.getLogger(__name__)

TOKEN_REFRESH_MARGIN = 300   # refresh access token 5 min before expiry


@dataclass
class HubSession:
    hub_url:       str
    username:      str
    user_id:       str
    access_token:  str
    refresh_token: str
    hub_pk_pem:    bytes         # cached hub Ed25519 public key
    node_id:       str = ""
    _token_exp:    int = 0

    @property
    def auth_headers(self) -> dict:
        return {"Authorization": f"Bearer {self.access_token}"}

    @property
    def token_expires_in(self) -> int:
        return max(0, self._token_exp - int(time.time()))

    @property
    def token_needs_refresh(self) -> bool:
        return self.token_expires_in < TOKEN_REFRESH_MARGIN


@dataclass
class HubConfig:
    hub_url:  str
    username: str
    password: str
    cache_dir: Path = field(default_factory=lambda: Path.home() / ".config" / "meshbay")

    @property
    def hub_pk_cache_path(self) -> Path:
        safe = self.hub_url.replace("://", "_").replace("/", "_").replace(":", "_")
        return self.cache_dir / f"hub_pk_{safe}.pem"


# ── Hub client ────────────────────────────────────────────────────────────────

class HubClient:
    """Async hub client. Use as async context manager or call close() explicitly."""

    def __init__(self, config: HubConfig, keys: NodeKeys):
        self._config  = config
        self._keys    = keys
        self._http    = httpx.AsyncClient(timeout=15, base_url=config.hub_url)
        self._session: HubSession | None = None

    async def __aenter__(self):
        return self

    async def __aexit__(self, *_):
        await self.close()

    async def close(self):
        await self._http.aclose()

    # ── Hub public key ────────────────────────────────────────────────────────

    async def _fetch_hub_pk(self) -> bytes:
        """Fetch and cache hub Ed25519 public key PEM."""
        cache = self._config.hub_pk_cache_path
        if cache.exists():
            log.debug("Hub PK loaded from cache: %s", cache)
            return cache.read_bytes()

        r = await self._http.get("/v1/hub/pubkey")
        r.raise_for_status()
        pem = r.json()["pk_hub_pem"].encode()

        cache.parent.mkdir(parents=True, exist_ok=True)
        cache.write_bytes(pem)
        log.info("Hub PK fetched and cached: %s", cache)
        return pem

    # ── Registration ──────────────────────────────────────────────────────────

    async def register(self) -> str:
        """Register this node's user on the hub. Returns user_id. Idempotent (409 ok)."""
        r = await self._http.post("/v1/users/register", json={
            "username":        self._config.username,
            "password":        self._config.password,
            "pk_user_ed25519": self._keys.pk_ed25519_b64,
            "pk_user_x25519":  self._keys.pk_x25519_b64,
        })
        if r.status_code == 201:
            log.info("Registered user '%s' on hub", self._config.username)
            return r.json()["user_id"]
        if r.status_code == 409:
            log.debug("User '%s' already registered", self._config.username)
            return ""
        r.raise_for_status()
        return ""

    # ── Login ─────────────────────────────────────────────────────────────────

    async def login(self) -> HubSession:
        """Login, verify JWT offline, return HubSession."""
        hub_pk_pem = await self._fetch_hub_pk()

        r = await self._http.post("/v1/users/login", json={
            "username": self._config.username,
            "password": self._config.password,
        })
        r.raise_for_status()
        data = r.json()

        access_token  = data["access_token"]
        refresh_token = data["refresh_token"]

        # Verify offline — if this passes, the hub's identity is confirmed
        decoded = jwt.decode(access_token, hub_pk_pem, algorithms=["EdDSA"])
        assert decoded["pk_user"] == self._keys.pk_ed25519_b64, \
            "Hub returned token for wrong public key"
        assert "jti" in decoded, "Hub token missing jti — hub is outdated"

        self._session = HubSession(
            hub_url=self._config.hub_url,
            username=self._config.username,
            user_id=decoded["sub"],
            access_token=access_token,
            refresh_token=refresh_token,
            hub_pk_pem=hub_pk_pem,
            _token_exp=decoded["exp"],
        )
        log.info("Logged in as '%s' (exp in %ds)", self._config.username,
                 self._session.token_expires_in)
        return self._session

    async def refresh_token(self) -> None:
        """Refresh the access token using the refresh token."""
        if self._session is None:
            raise RuntimeError("Not logged in")

        r = await self._http.post("/v1/users/token/refresh", json={
            "refresh_token": self._session.refresh_token,
        })
        r.raise_for_status()
        new_token = r.json()["access_token"]

        decoded = jwt.decode(new_token, self._session.hub_pk_pem, algorithms=["EdDSA"])
        self._session.access_token = new_token
        self._session._token_exp   = decoded["exp"]
        log.debug("Access token refreshed (exp in %ds)", self._session.token_expires_in)

    async def ensure_fresh_token(self) -> None:
        """Auto-refresh token if close to expiry."""
        if self._session and self._session.token_needs_refresh:
            await self.refresh_token()

    # ── Node announcement ─────────────────────────────────────────────────────

    async def announce_node(self, endpoint_hint: str | None = None) -> str:
        """Announce this node to the hub. Returns node_id."""
        if self._session is None:
            raise RuntimeError("Not logged in")
        await self.ensure_fresh_token()

        r = await self._http.post("/v1/nodes/announce", json={
            "pk_node":       self._keys.pk_ed25519_b64,
            "endpoint_hint": endpoint_hint,
        }, headers=self._session.auth_headers)
        r.raise_for_status()
        node_id = r.json()["node_id"]
        self._session.node_id = node_id
        log.info("Node announced: %s (hint=%s)", node_id[:8], endpoint_hint)
        return node_id

    # ── GEK retrieval ─────────────────────────────────────────────────────────

    async def fetch_gek(self, group_id: str) -> bytes:
        """
        Fetch and unwrap the GEK bundle for a group.
        Returns the raw GEK bytes.
        """
        if self._session is None:
            raise RuntimeError("Not logged in")
        await self.ensure_fresh_token()

        r = await self._http.get(f"/v1/groups/{group_id}/gek",
                                 headers=self._session.auth_headers)
        if r.status_code == 404:
            raise LookupError(f"No GEK bundle found for group {group_id!r}")
        r.raise_for_status()

        bundle     = r.json()
        sk_x_raw   = self._keys.sk_x25519.private_bytes(
            serialization.Encoding.Raw, serialization.PrivateFormat.Raw,
            serialization.NoEncryption())
        pk_x_raw   = base64.b64decode(self._keys.pk_x25519_b64)

        gek = unwrap_gek(bundle, sk_x_raw, pk_x_raw)
        log.info("GEK unwrapped for group %s", group_id[:8])
        return gek

    # ── User pubkey lookup ────────────────────────────────────────────────────

    async def get_user_pubkeys(self, username: str) -> dict:
        """Return {'pk_ed25519': str, 'pk_x25519': str} for a user."""
        if self._session is None:
            raise RuntimeError("Not logged in")
        await self.ensure_fresh_token()

        r = await self._http.get(f"/v1/users/{username}/pubkeys",
                                 headers=self._session.auth_headers)
        if r.status_code == 404:
            raise LookupError(f"User not found: {username!r}")
        r.raise_for_status()
        return r.json()

    # ── Convenience: full startup sequence ───────────────────────────────────

    async def startup(self, endpoint_hint: str | None = None) -> HubSession:
        """
        Full startup sequence: register (idempotent) → login → announce node.
        Returns an active HubSession.
        """
        await self.register()
        session = await self.login()
        await self.announce_node(endpoint_hint)
        return session