summaryrefslogtreecommitdiffstats
path: root/packages/meshbay-node/src/meshbay_node/transport/client.py
blob: 34303653d4b449e3dbc54b4498d0ddf331bf43de (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
"""
MeshBay — TCP+TLS chunk client (MNP v1).

Used by the web client (or other nodes) to fetch files from a Mesh Node.
Verifies Ed25519 chunk signatures using the node's public key from the hub.
"""

import asyncio
import base64
import logging
import struct
from pathlib import Path

import blake3
import msgpack
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey

from meshbay_common import MNP_VERSION
from meshbay_common.crypto import (
    chunk_key as derive_chunk_key,
    decrypt_chunk,
    verify_chunk_signature,
)
from meshbay_common.protocol import MNP
from meshbay_node.transport.tls_cert import client_ssl_context

log = logging.getLogger(__name__)

MAX_MSG = 64 * 1024 * 1024


async def _send(writer, obj):
    data = msgpack.packb(obj, use_bin_type=True)
    writer.write(struct.pack(">I", len(data)) + data)
    await writer.drain()

async def _recv(reader):
    header = await reader.readexactly(4)
    length = struct.unpack(">I", header)[0]
    if length > MAX_MSG:
        raise ValueError(f"Message too large: {length}")
    return msgpack.unpackb(await reader.readexactly(length), raw=False)


class ChunkClient:
    """
    Async client for fetching encrypted chunks from a ChunkServer.

    Usage:
        async with ChunkClient(host, port, jwt_token, gek, pk_node_b64) as client:
            data = await client.fetch_chunk(file_id, chunk_index=0)
    """

    def __init__(
        self,
        host: str,
        port: int,
        jwt_token: str,
        gek: bytes,
        pk_node_b64: str,     # node's Ed25519 PK from hub — used for sig verification
    ):
        self._host       = host
        self._port       = port
        self._jwt_token  = jwt_token
        self._gek        = gek
        self._pk_node    = Ed25519PublicKey.from_public_bytes(
            base64.b64decode(pk_node_b64))
        self._reader: asyncio.StreamReader | None = None
        self._writer: asyncio.StreamWriter | None = None

    async def __aenter__(self):
        await self.connect()
        return self

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

    async def connect(self) -> None:
        ssl_ctx = client_ssl_context()
        self._reader, self._writer = await asyncio.open_connection(
            self._host, self._port, ssl=ssl_ctx)

        # MNP handshake
        await _send(self._writer, {
            "type":  MNP.HANDSHAKE,
            "v":     MNP_VERSION,
            "token": self._jwt_token,
        })
        ack = await _recv(self._reader)
        if ack.get("type") != MNP.HANDSHAKE_ACK:
            raise ConnectionError(f"Handshake rejected: {ack}")
        log.debug("Connected to node %s:%d", self._host, self._port)

    async def close(self) -> None:
        if self._writer:
            self._writer.close()
            await self._writer.wait_closed()

    async def fetch_index(self) -> bytes:
        """Request the Mesh Group Index. Returns raw wire bytes (encrypted)."""
        await _send(self._writer, {"type": MNP.INDEX_SYNC, "v": MNP_VERSION})
        msg = await _recv(self._reader)
        return base64.b64decode(msg["index_b64"])

    async def fetch_chunk(self, file_id: str, chunk_index: int) -> bytes:
        """
        Fetch, verify, and decrypt one chunk.
        Returns plaintext bytes.
        """
        await _send(self._writer, {
            "type":        MNP.FILE_REQUEST,
            "v":           MNP_VERSION,
            "file_id":     file_id,
            "chunk_index": chunk_index,
        })
        msg = await _recv(self._reader)

        if msg.get("type") == "error":
            raise LookupError(msg.get("detail", "Unknown error"))

        ct        = base64.b64decode(msg["ct_b64"])
        nonce     = base64.b64decode(msg["nonce_b64"])
        ct_hash   = base64.b64decode(msg["ct_hash_b64"])
        pt_hash   = base64.b64decode(msg["pt_hash_b64"])
        sig       = base64.b64decode(msg["sig_b64"])
        file_hash = base64.b64decode(msg["file_hash_b64"])
        ci        = msg["chunk_index"]

        # 1. Verify Ed25519 signature
        verify_chunk_signature(self._pk_node, ci, nonce, ct_hash, sig)

        # 2. Verify ciphertext hash
        if blake3.blake3(ct).digest() != ct_hash:
            raise ValueError("Ciphertext hash mismatch")

        # 3. Decrypt
        ckey = derive_chunk_key(self._gek, file_hash, ci)
        plaintext = decrypt_chunk(ckey, nonce, ct)

        # 4. Verify plaintext hash
        if blake3.blake3(plaintext).digest() != pt_hash:
            raise ValueError("Plaintext hash mismatch after decryption")

        return plaintext