From c4e15bf62733b041b95eb467c05e08e16a2921ad Mon Sep 17 00:00:00 2001 From: puffball1567 Date: Mon, 20 Jul 2026 17:00:03 +0900 Subject: [PATCH] feat: add auth and TLS support and rebrand to KoutenDB Implements password and shared-secret (challenge-response) authentication, the encrypted secure transport, and TLS with CA-file verification / server-name override / an explicit insecure opt-out. Shared-secret crypto uses libsodium via the optional PyNaCl `secure` extra. Renames the package to KoutenDB. Co-Authored-By: Claude Opus 4.8 (1M context) --- .github/workflows/ci.yml | 14 +- README.md | 80 ++++++--- koutendb/__init__.py | 3 + {rochedb => koutendb}/client.py | 276 +++++++++++++++++++++++++------- koutendb/secure.py | 87 ++++++++++ pyproject.toml | 26 +-- rochedb/__init__.py | 3 - tests/test_driver.py | 24 +-- tests/test_secure.py | 69 ++++++++ 9 files changed, 474 insertions(+), 108 deletions(-) create mode 100644 koutendb/__init__.py rename {rochedb => koutendb}/client.py (52%) create mode 100644 koutendb/secure.py delete mode 100644 rochedb/__init__.py create mode 100644 tests/test_secure.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 51a811c..9dcf23c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -23,11 +23,11 @@ jobs: - name: Checkout Python driver uses: actions/checkout@v4 - - name: Checkout RocheDB core + - name: Checkout KoutenDB core uses: actions/checkout@v4 with: - repository: puffball1567/rochedb - path: rochedb-core + repository: puffball1567/koutendb + path: koutendb-core - name: Set up Python uses: actions/setup-python@v5 @@ -42,14 +42,14 @@ jobs: - name: Install dependencies run: sudo apt-get update && sudo apt-get install -y libsodium-dev - - name: Build roched + - name: Build koutend run: | - cd rochedb-core + cd koutendb-core nimble install -y - nim c -d:release --nimcache:/tmp/nimcache_roched -o:src/roched src/roched.nim + nim c -d:release --nimcache:/tmp/nimcache_koutend -o:src/koutend src/koutend.nim - name: Install package run: python -m pip install -e . - name: Run tests - run: ROCHEDB_CORE_DIR="${{ github.workspace }}/rochedb-core" python -m unittest discover -s tests + run: KOUTENDB_CORE_DIR="${{ github.workspace }}/koutendb-core" python -m unittest discover -s tests diff --git a/README.md b/README.md index 8c98643..b332c7c 100644 --- a/README.md +++ b/README.md @@ -1,18 +1,18 @@ -# RocheDB Python Driver +# KoutenDB Python Driver -Pure Python TCP driver for [RocheDB](https://github.com/puffball1567/rochedb). +Pure Python TCP driver for [KoutenDB](https://github.com/puffball1567/koutendb). -This driver talks to `roched` over RocheDB's high-level wire protocol. It does -not reimplement RocheDB's ring-key, period, head-angle, or placement rules. -Applications pass a human-readable ring name, and RocheDB returns a typed ID. +This driver talks to `koutend` over KoutenDB's high-level wire protocol. It does +not reimplement KoutenDB's ring-key, period, head-angle, or placement rules. +Applications pass a human-readable ring name, and KoutenDB returns a typed ID. ## Status -- package: PyPI [`rochedb`](https://pypi.org/project/rochedb/) v0.1.3 +- package: PyPI [`koutendb`](https://pypi.org/project/koutendb/) v0.1.3 - current mode: native TCP wire driver - Python: 3.10+ - runtime dependencies: none -- RocheDB core: running `roched` node or cluster +- KoutenDB core: running `koutend` node or cluster Implemented: @@ -23,7 +23,7 @@ Implemented: - `query` / `query_encoded` / `query_text` / `query_json` - codec metadata negotiation with `CODECMETA ON` - `batch_get` -- typed `RocheId` +- typed `KoutenId` - one reconnect retry - context manager support @@ -39,7 +39,7 @@ Planned: Install the published package from PyPI: ```sh -python3 -m pip install rochedb +python3 -m pip install koutendb ``` For local driver development, install from a checkout: @@ -48,21 +48,21 @@ For local driver development, install from a checkout: python3 -m pip install -e . ``` -Build `roched` from the RocheDB core repository: +Build `koutend` from the KoutenDB core repository: ```sh -git clone https://github.com/puffball1567/rochedb.git -cd rochedb +git clone https://github.com/puffball1567/koutendb.git +cd koutendb nimble install -y -nim c -d:release --nimcache:/tmp/nimcache_roched -o:src/roched src/roched.nim +nim c -d:release --nimcache:/tmp/nimcache_koutend -o:src/koutend src/koutend.nim ``` ## Example ```python -from rochedb import RocheClient +from koutendb import KoutenClient -with RocheClient.connect("127.0.0.1:17301") as db: +with KoutenClient.connect("127.0.0.1:17301") as db: doc_id = db.put_json( "docs/japan/support", {"title": "Tokyo support note", "country": "JP"}, @@ -74,21 +74,61 @@ with RocheClient.connect("127.0.0.1:17301") as db: print(db.query_json(doc_id, "{ title }")) ``` +## Authentication and TLS + +`connect` accepts credentials and TLS options. Password auth and TLS use only +the standard library. Shared-secret (`secret_key`) challenge-response and the +encrypted transport it enables additionally need libsodium via PyNaCl — install +the `secure` extra: + +```sh +pip install koutendb[secure] +``` + +Connect over TLS with shared-secret auth, verifying the server against a CA or +self-signed certificate PEM (certificate verification stays on): + +```python +from koutendb import KoutenClient + +db = KoutenClient.connect( + "127.0.0.1:17301", + username="alice", + password="secret", + secret_key="shared-secret", + tls_ca_file="/path/to/server.crt", +) +``` + +Password-only auth over TLS needs no extra dependency: + +```python +db = KoutenClient.connect( + "127.0.0.1:17301", username="alice", password="secret", + tls_ca_file="/path/to/server.crt", +) +``` + +`tls_insecure_skip_verify=True` disables certificate verification. The +connection is then encrypted but unauthenticated and trivially impersonable, so +it is for local smoke tests only — never a production server. Prefer +`tls_ca_file` for self-signed certificates. + ## Test -From this driver repository, point `ROCHEDB_CORE_DIR` at a RocheDB checkout: +From this driver repository, point `KOUTENDB_CORE_DIR` at a KoutenDB checkout: ```sh -ROCHEDB_CORE_DIR=/path/to/rochedb python3 -m unittest discover -s tests +KOUTENDB_CORE_DIR=/path/to/koutendb python3 -m unittest discover -s tests ``` -The test starts a two-node local `roched` cluster and verifies put/get/query, +The test starts a two-node local `koutend` cluster and verifies put/get/query, JSON helpers, codec metadata, BIF opaque payloads, `wire_version`, and `batch_get`. ## Why A Native Wire Driver? The Python driver is intended for API services, scripts, experiments, and -AI/RAG validation where a running RocheDB server or cluster is the natural -boundary. It keeps Python out of RocheDB's placement internals and uses the same +AI/RAG validation where a running KoutenDB server or cluster is the natural +boundary. It keeps Python out of KoutenDB's placement internals and uses the same ring-oriented API that other external drivers should use. diff --git a/koutendb/__init__.py b/koutendb/__init__.py new file mode 100644 index 0000000..7f2093d --- /dev/null +++ b/koutendb/__init__.py @@ -0,0 +1,3 @@ +from .client import EncodedPayload, PayloadCodec, KoutenClient, KoutenError, KoutenId + +__all__ = ["EncodedPayload", "PayloadCodec", "KoutenClient", "KoutenError", "KoutenId"] diff --git a/rochedb/client.py b/koutendb/client.py similarity index 52% rename from rochedb/client.py rename to koutendb/client.py index 0af337f..4cfc213 100644 --- a/rochedb/client.py +++ b/koutendb/client.py @@ -3,12 +3,20 @@ from dataclasses import dataclass import json import socket +import ssl import struct from typing import Any, Iterable, Literal, Optional +from .secure import ( + SecureState, + decrypt_transport_frame, + encrypt_transport_frame, + secret_response_hex, +) -class RocheError(Exception): - """Raised when RocheDB returns an error frame or the TCP connection fails.""" + +class KoutenError(Exception): + """Raised when KoutenDB returns an error frame or the TCP connection fails.""" PayloadCodec = Literal["raw", "json", "nif", "bif"] @@ -16,7 +24,7 @@ class RocheError(Exception): @dataclass(frozen=True) -class RocheId: +class KoutenId: parent: int epoch: int seq: int @@ -25,10 +33,10 @@ class RocheId: head: float @classmethod - def parse(cls, text: str) -> "RocheId": + def parse(cls, text: str) -> "KoutenId": parts = text.split(":") if len(parts) != 6: - raise ValueError("RocheId text must have 6 ':'-separated fields") + raise ValueError("KoutenId text must have 6 ':'-separated fields") return cls( parent=int(parts[0]), epoch=int(parts[1]), @@ -85,15 +93,68 @@ def _codec(codec: str) -> PayloadCodec: return codec # type: ignore[return-value] -class RocheClient: - def __init__(self, peers: str | Iterable[str], timeout: float = 10.0): +class KoutenClient: + def __init__( + self, + peers: str | Iterable[str], + timeout: float = 10.0, + *, + username: str = "", + password: str = "", + auth_token: str = "", + secret_key: str = "", + galaxy: str = "", + tls: bool = False, + tls_ca_file: str = "", + tls_server_name: str = "", + tls_insecure_skip_verify: bool = False, + ): self.peers = _parse_peers(peers) self.timeout = timeout + # A bare token becomes password auth under a reserved "token" user, the + # same normalization the core client applies. + if auth_token and not username: + username, password = "token", auth_token + self.username = username + self.password = password + self.secret_key = secret_key + self.galaxy = galaxy + self.tls = tls or bool(tls_ca_file) or bool(tls_server_name) or tls_insecure_skip_verify + self.tls_ca_file = tls_ca_file + self.tls_server_name = tls_server_name + self.tls_insecure_skip_verify = tls_insecure_skip_verify self._socks: dict[int, socket.socket] = {} + self._secure: dict[int, SecureState] = {} @classmethod - def connect(cls, peers: str | Iterable[str], timeout: float = 10.0) -> "RocheClient": - return cls(peers, timeout=timeout) + def connect( + cls, + peers: str | Iterable[str], + timeout: float = 10.0, + *, + username: str = "", + password: str = "", + auth_token: str = "", + secret_key: str = "", + galaxy: str = "", + tls: bool = False, + tls_ca_file: str = "", + tls_server_name: str = "", + tls_insecure_skip_verify: bool = False, + ) -> "KoutenClient": + return cls( + peers, + timeout=timeout, + username=username, + password=password, + auth_token=auth_token, + secret_key=secret_key, + galaxy=galaxy, + tls=tls, + tls_ca_file=tls_ca_file, + tls_server_name=tls_server_name, + tls_insecure_skip_verify=tls_insecure_skip_verify, + ) def close(self) -> None: for sock in self._socks.values(): @@ -102,8 +163,9 @@ def close(self) -> None: except OSError: pass self._socks.clear() + self._secure.clear() - def __enter__(self) -> "RocheClient": + def __enter__(self) -> "KoutenClient": return self def __exit__(self, exc_type, exc, tb) -> None: @@ -112,13 +174,13 @@ def __exit__(self, exc_type, exc, tb) -> None: def wire_version(self, node: int = 0) -> int: parts = self._rpc(node, "WIREVER") if len(parts) != 2 or parts[0] != "WIREVER": - raise RocheError("WIREVER failed: " + " ".join(parts)) + raise KoutenError("WIREVER failed: " + " ".join(parts)) return int(parts[1]) def health(self, node: int = 0) -> str: parts = self._rpc(node, "HEALTH") if not parts or parts[0] != "OK": - raise RocheError("HEALTH failed: " + " ".join(parts)) + raise KoutenError("HEALTH failed: " + " ".join(parts)) return " ".join(parts[1:]) def put( @@ -128,7 +190,7 @@ def put( vector: Optional[Iterable[float]] = None, codec: PayloadCodec = "raw", node: int = 0, - ) -> RocheId: + ) -> KoutenId: ring_b = ring.encode("utf-8") payload_b = _as_bytes(payload) vec_b = _vec_bytes(vector) @@ -136,8 +198,8 @@ def put( header = f"PUTR {len(ring_b)} {len(payload_b)} {vec_dim} {_codec(codec)}" parts = self._rpc(node, header, ring_b + payload_b + vec_b) if not parts or parts[0] != "ID" or len(parts) != 7: - raise RocheError("PUTR failed: " + " ".join(parts)) - return RocheId( + raise KoutenError("PUTR failed: " + " ".join(parts)) + return KoutenId( parent=int(parts[1]), epoch=int(parts[2]), seq=int(parts[3]), @@ -153,7 +215,7 @@ def put_codec( codec: PayloadCodec, vector: Optional[Iterable[float]] = None, node: int = 0, - ) -> RocheId: + ) -> KoutenId: return self.put(ring, payload, vector=vector, codec=codec, node=node) def put_json( @@ -162,7 +224,7 @@ def put_json( value: Any, vector: Optional[Iterable[float]] = None, node: int = 0, - ) -> RocheId: + ) -> KoutenId: payload = json.dumps(value, separators=(",", ":"), ensure_ascii=False) return self.put(ring, payload, vector=vector, codec="json", node=node) @@ -172,7 +234,7 @@ def put_nif( payload: bytes | bytearray | memoryview | str, vector: Optional[Iterable[float]] = None, node: int = 0, - ) -> RocheId: + ) -> KoutenId: return self.put(ring, payload, vector=vector, codec="nif", node=node) def put_bif( @@ -181,48 +243,48 @@ def put_bif( payload: bytes | bytearray | memoryview, vector: Optional[Iterable[float]] = None, node: int = 0, - ) -> RocheId: + ) -> KoutenId: return self.put(ring, payload, vector=vector, codec="bif", node=node) - def get(self, doc_id: RocheId, node: Optional[int] = None) -> Optional[bytes]: + def get(self, doc_id: KoutenId, node: Optional[int] = None) -> Optional[bytes]: return self._read_with_fallback("GETID", doc_id, b"", node=node) def get_encoded( - self, doc_id: RocheId, node: Optional[int] = None + self, doc_id: KoutenId, node: Optional[int] = None ) -> Optional[EncodedPayload]: return self._read_encoded_with_fallback("GETID", doc_id, b"", node=node) - def get_text(self, doc_id: RocheId, node: Optional[int] = None) -> Optional[str]: + def get_text(self, doc_id: KoutenId, node: Optional[int] = None) -> Optional[str]: value = self.get(doc_id, node=node) return None if value is None else value.decode("utf-8") - def get_json(self, doc_id: RocheId, node: Optional[int] = None) -> Any: + def get_json(self, doc_id: KoutenId, node: Optional[int] = None) -> Any: value = self.get_text(doc_id, node=node) return None if value is None else json.loads(value) def query( - self, doc_id: RocheId, selection: str, node: Optional[int] = None + self, doc_id: KoutenId, selection: str, node: Optional[int] = None ) -> Optional[bytes]: selection_b = selection.encode("utf-8") return self._read_with_fallback("QRYID", doc_id, selection_b, node=node) def query_encoded( - self, doc_id: RocheId, selection: str, node: Optional[int] = None + self, doc_id: KoutenId, selection: str, node: Optional[int] = None ) -> Optional[EncodedPayload]: selection_b = selection.encode("utf-8") return self._read_encoded_with_fallback("QRYID", doc_id, selection_b, node=node) def query_text( - self, doc_id: RocheId, selection: str, node: Optional[int] = None + self, doc_id: KoutenId, selection: str, node: Optional[int] = None ) -> Optional[str]: value = self.query(doc_id, selection, node=node) return None if value is None else value.decode("utf-8") - def query_json(self, doc_id: RocheId, selection: str, node: Optional[int] = None) -> Any: + def query_json(self, doc_id: KoutenId, selection: str, node: Optional[int] = None) -> Any: value = self.query_text(doc_id, selection, node=node) return None if value is None else json.loads(value) - def batch_get(self, ids: Iterable[RocheId], node: int = 0) -> list[Optional[bytes]]: + def batch_get(self, ids: Iterable[KoutenId], node: int = 0) -> list[Optional[bytes]]: id_list = list(ids) body = "".join( f"{doc_id.parent} {doc_id.seq} {doc_id.period} {doc_id.head} {doc_id.t_write}\n" @@ -230,15 +292,15 @@ def batch_get(self, ids: Iterable[RocheId], node: int = 0) -> list[Optional[byte ).encode("utf-8") parts = self._rpc(node, f"BGET {len(id_list)} {len(body)}", body) if len(parts) != 3 or parts[0] != "BVAL": - raise RocheError("BGET failed: " + " ".join(parts)) + raise KoutenError("BGET failed: " + " ".join(parts)) expected = int(parts[1]) - payload = self._read_exact(self._socket_for(node), int(parts[2])) + payload = self._read_exact(node, int(parts[2])) out: list[Optional[bytes]] = [] pos = 0 for _ in range(expected): nl = payload.find(b"\n", pos) if nl < 0: - raise RocheError("BGET payload length header missing") + raise KoutenError("BGET payload length header missing") length = int(payload[pos:nl].decode("utf-8")) pos = nl + 1 if length == 0: @@ -249,7 +311,7 @@ def batch_get(self, ids: Iterable[RocheId], node: int = 0) -> list[Optional[byte return out def _read_id_encoded( - self, op: str, doc_id: RocheId, selection: bytes, node: int + self, op: str, doc_id: KoutenId, selection: bytes, node: int ) -> Optional[EncodedPayload]: header = ( f"{op} {doc_id.parent} {doc_id.epoch} {doc_id.seq} " @@ -259,15 +321,15 @@ def _read_id_encoded( header += f" {len(selection)}" parts = self._rpc(node, header, selection) if not parts: - raise RocheError(f"{op} returned an empty response") + raise KoutenError(f"{op} returned an empty response") if parts[0] == "MISS": return None if parts[0] == "ERR": - raise RocheError(" ".join(parts[1:])) + raise KoutenError(" ".join(parts[1:])) if parts[0] == "FWD": if len(parts) != 7: - raise RocheError("invalid FWD response: " + " ".join(parts)) - fwd = RocheId( + raise KoutenError("invalid FWD response: " + " ".join(parts)) + fwd = KoutenId( parent=int(parts[1]), epoch=int(parts[2]), seq=int(parts[3]), @@ -277,16 +339,16 @@ def _read_id_encoded( ) return self._read_id_encoded(op, fwd, selection, node=node) if parts[0] != "VAL" or len(parts) not in (3, 4): - raise RocheError(f"{op} failed: " + " ".join(parts)) + raise KoutenError(f"{op} failed: " + " ".join(parts)) codec = _codec(parts[3]) if len(parts) == 4 else ("json" if op == "QRYID" else "raw") - return EncodedPayload(self._read_exact(self._socket_for(node), int(parts[2])), codec) + return EncodedPayload(self._read_exact(node, int(parts[2])), codec) - def _read_id(self, op: str, doc_id: RocheId, selection: bytes, node: int) -> Optional[bytes]: + def _read_id(self, op: str, doc_id: KoutenId, selection: bytes, node: int) -> Optional[bytes]: value = self._read_id_encoded(op, doc_id, selection, node=node) return None if value is None else value.payload def _read_with_fallback( - self, op: str, doc_id: RocheId, selection: bytes, node: Optional[int] + self, op: str, doc_id: KoutenId, selection: bytes, node: Optional[int] ) -> Optional[bytes]: if node is not None: return self._read_id(op, doc_id, selection, node=node) @@ -300,7 +362,7 @@ def _read_with_fallback( return None def _read_encoded_with_fallback( - self, op: str, doc_id: RocheId, selection: bytes, node: Optional[int] + self, op: str, doc_id: KoutenId, selection: bytes, node: Optional[int] ) -> Optional[EncodedPayload]: if node is not None: return self._read_id_encoded(op, doc_id, selection, node=node) @@ -317,15 +379,15 @@ def _rpc(self, node: int, header: str, payload: bytes = b"") -> list[str]: last_error: Optional[BaseException] = None for attempt in range(2): try: - sock = self._socket_for(node) - sock.sendall(header.encode("utf-8") + b"\n" + payload) - return self._read_header(sock) + self._socket_for(node) + self._send_frame(node, header, payload) + return self._read_header(node) except OSError as err: last_error = err self._drop_socket(node) if attempt == 1: - raise RocheError(str(err)) from err - raise RocheError(str(last_error)) + raise KoutenError(str(err)) from err + raise KoutenError(str(last_error)) def _socket_for(self, node: int) -> socket.socket: if node < 0 or node >= len(self.peers): @@ -340,18 +402,72 @@ def _socket_for(self, node: int) -> socket.socket: sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) except OSError: pass + if self.tls: + sock = self._wrap_tls(sock, host) + self._socks[node] = sock try: - sock.sendall(b"CODECMETA ON\n") - reply = self._read_header(sock) - if not reply or reply[0] != "OK": - raise RocheError("CODECMETA failed: " + " ".join(reply)) + self._handshake(node) except BaseException: - sock.close() + self._drop_socket(node) raise - self._socks[node] = sock return sock + def _wrap_tls(self, sock: socket.socket, host: str) -> ssl.SSLSocket: + # Verification stays on unless the caller explicitly opts out. A CA file + # anchors trust to a private CA / self-signed certificate while keeping + # certificate and hostname checks; the insecure switch is the only path + # that disables them, and it is for local smoke tests only. + context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + if self.tls_insecure_skip_verify: + context.check_hostname = False + context.verify_mode = ssl.CERT_NONE + else: + if self.tls_ca_file: + context.load_verify_locations(self.tls_ca_file) + else: + context.load_default_certs() + context.check_hostname = True + context.verify_mode = ssl.CERT_REQUIRED + server_hostname = self.tls_server_name or host + try: + return context.wrap_socket(sock, server_hostname=server_hostname) + except ssl.SSLError as err: + sock.close() + raise KoutenError(f"TLS handshake failed: {err}") from err + + def _handshake(self, node: int) -> None: + if self.username: + if self.secret_key: + self._send_frame(node, "AUTHCHAL " + self.username) + chal = self._read_header(node) + if len(chal) < 2 or chal[0] != "CHAL": + raise KoutenError("AUTHCHAL failed: " + " ".join(chal)) + response = secret_response_hex( + self.username, self.password, chal[1], self.secret_key + ) + self._send_frame(node, "AUTHRESP " + response) + reply = self._read_header(node) + if not reply or reply[0] != "OK": + raise KoutenError("AUTHRESP failed: " + " ".join(reply)) + # Every frame after this point is sealed with the transport key. + self._secure[node] = SecureState(self.secret_key, chal[1]) + else: + self._send_frame(node, f"AUTH {self.username} {self.password}") + reply = self._read_header(node) + if not reply or reply[0] != "OK": + raise KoutenError("AUTH failed: " + " ".join(reply)) + if self.galaxy: + self._send_frame(node, "HELLO " + self.galaxy) + reply = self._read_header(node) + if not reply or reply[0] != "OK": + raise KoutenError("HELLO failed: " + " ".join(reply)) + self._send_frame(node, "CODECMETA ON") + reply = self._read_header(node) + if not reply or reply[0] != "OK": + raise KoutenError("CODECMETA failed: " + " ".join(reply)) + def _drop_socket(self, node: int) -> None: + self._secure.pop(node, None) sock = self._socks.pop(node, None) if sock is not None: try: @@ -359,23 +475,69 @@ def _drop_socket(self, node: int) -> None: except OSError: pass - def _read_header(self, sock: socket.socket) -> list[str]: + def _send_frame(self, node: int, header: str, payload: bytes = b"") -> None: + plaintext = header.encode("utf-8") + b"\n" + payload + state = self._secure.get(node) + if state is None: + self._socks[node].sendall(plaintext) + else: + ciphertext = encrypt_transport_frame( + plaintext, state.secret_key, state.challenge_hex + ) + self._socks[node].sendall( + b"SEC " + str(len(ciphertext)).encode("ascii") + b"\n" + ciphertext + ) + + def _read_secure_frame(self, node: int) -> None: + sock = self._socks[node] + header = self._raw_read_line(sock).split(" ") + if len(header) < 2 or header[0] != "SEC": + raise KoutenError("expected secure frame, got: " + " ".join(header)) + ciphertext = self._raw_read_exact(sock, int(header[1])) + state = self._secure[node] + state.buffer += decrypt_transport_frame( + ciphertext, state.secret_key, state.challenge_hex + ) + + def _read_header(self, node: int) -> list[str]: + state = self._secure.get(node) + if state is None: + return self._raw_read_line(self._socks[node]).split(" ") + while True: + nl = state.buffer.find(b"\n") + if nl >= 0: + line = state.buffer[:nl] + state.buffer = state.buffer[nl + 1 :] + return line.decode("utf-8").split(" ") + self._read_secure_frame(node) + + def _read_exact(self, node: int, n: int) -> bytes: + state = self._secure.get(node) + if state is None: + return self._raw_read_exact(self._socks[node], n) + while len(state.buffer) < n: + self._read_secure_frame(node) + out = state.buffer[:n] + state.buffer = state.buffer[n:] + return out + + def _raw_read_line(self, sock: socket.socket) -> str: chunks: list[bytes] = [] while True: b = sock.recv(1) if not b: - raise RocheError("connection closed") + raise KoutenError("connection closed") if b == b"\n": - return b"".join(chunks).decode("utf-8").split(" ") + return b"".join(chunks).decode("utf-8") chunks.append(b) - def _read_exact(self, sock: socket.socket, n: int) -> bytes: + def _raw_read_exact(self, sock: socket.socket, n: int) -> bytes: chunks: list[bytes] = [] remaining = n while remaining > 0: chunk = sock.recv(remaining) if not chunk: - raise RocheError("connection closed") + raise KoutenError("connection closed") chunks.append(chunk) remaining -= len(chunk) return b"".join(chunks) diff --git a/koutendb/secure.py b/koutendb/secure.py new file mode 100644 index 0000000..628e194 --- /dev/null +++ b/koutendb/secure.py @@ -0,0 +1,87 @@ +"""libsodium-backed shared-secret challenge response and secure transport. + +This mirrors the KoutenDB core `kouten/auth` module byte-for-byte so the +Python wire driver interoperates with the Nim server: + +- key derivation uses BLAKE2b (libsodium ``crypto_generichash``), available in + the standard library as ``hashlib.blake2b``; +- the challenge response and transport frames use ``crypto_secretbox`` + (XSalsa20-Poly1305), which requires PyNaCl. PyNaCl is imported lazily so the + driver stays usable for plaintext / password-only connections without it. +""" + +from __future__ import annotations + +import hashlib + +_AUTH_DOMAIN = b"koutendb-auth-v1" +_KEY_BYTES = 32 + + +def _require_secretbox(): + try: + from nacl.secret import SecretBox + except ImportError as exc: # pragma: no cover - exercised via error path + raise RuntimeError( + "secret-key authentication requires PyNaCl. Install it with " + "'pip install koutendb[secure]' or 'pip install pynacl'." + ) from exc + return SecretBox + + +def _blake2b32(data: bytes) -> bytes: + return hashlib.blake2b(data, digest_size=_KEY_BYTES).digest() + + +def _box_key(secret_key: bytes) -> bytes: + return _blake2b32(_AUTH_DOMAIN + b"\0box\0" + secret_key) + + +def _transport_key(secret_key: bytes, challenge_hex: bytes) -> bytes: + return _blake2b32( + _AUTH_DOMAIN + b"\0transport\0" + challenge_hex + b"\0" + secret_key + ) + + +def _auth_message(username: bytes, password: bytes, challenge_hex: bytes) -> bytes: + return _AUTH_DOMAIN + b"\n" + username + b"\n" + password + b"\n" + challenge_hex + + +def secret_response_hex( + username: str, password: str, challenge_hex: str, secret_key: str +) -> str: + """Return the hex-encoded sealed response for an AUTHRESP frame.""" + box = _require_secretbox()(_box_key(secret_key.encode("utf-8"))) + message = _auth_message( + username.encode("utf-8"), password.encode("utf-8"), challenge_hex.encode("utf-8") + ) + return bytes(box.encrypt(message)).hex() + + +def encrypt_transport_frame( + plaintext: bytes, secret_key: str, challenge_hex: str +) -> bytes: + box = _require_secretbox()( + _transport_key(secret_key.encode("utf-8"), challenge_hex.encode("utf-8")) + ) + return bytes(box.encrypt(plaintext)) + + +def decrypt_transport_frame( + ciphertext: bytes, secret_key: str, challenge_hex: str +) -> bytes: + box = _require_secretbox()( + _transport_key(secret_key.encode("utf-8"), challenge_hex.encode("utf-8")) + ) + return bytes(box.decrypt(ciphertext)) + + +class SecureState: + """Decrypted-plaintext buffer for one secured connection.""" + + __slots__ = ("secret_key", "challenge_hex", "buffer") + + def __init__(self, secret_key: str, challenge_hex: str): + self.secret_key = secret_key + self.challenge_hex = challenge_hex + self.buffer = b"" diff --git a/pyproject.toml b/pyproject.toml index 2805a55..92e810c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,16 +3,21 @@ requires = ["setuptools>=68"] build-backend = "setuptools.build_meta" [project] -name = "rochedb" -version = "0.1.3" -description = "Pure Python TCP driver for RocheDB" +name = "koutendb" +version = "0.2.0" +description = "Pure Python TCP driver for KoutenDB" readme = "README.md" requires-python = ">=3.10" license = "Apache-2.0" authors = [ - { name = "RocheDB contributors" } + { name = "KoutenDB contributors" } ] -keywords = ["rochedb", "database", "nosql", "rag", "retrieval"] +keywords = ["koutendb", "database", "nosql", "rag", "retrieval"] + +# TLS uses the standard library `ssl` module (no dependency). Shared-secret +# (secret_key) challenge-response and encrypted transport frames additionally +# need libsodium via PyNaCl; install with the `secure` extra when using it. +dependencies = [] classifiers = [ "Development Status :: 3 - Alpha", "Programming Language :: Python :: 3", @@ -22,11 +27,14 @@ classifiers = [ "Programming Language :: Python :: 3.13", ] +[project.optional-dependencies] +secure = ["pynacl>=1.5"] + [project.urls] -Homepage = "https://github.com/puffball1567/rochedb" -Repository = "https://github.com/puffball1567/rochedb-python" -Issues = "https://github.com/puffball1567/rochedb-python/issues" +Homepage = "https://github.com/puffball1567/koutendb" +Repository = "https://github.com/puffball1567/koutendb-python" +Issues = "https://github.com/puffball1567/koutendb-python/issues" [tool.setuptools.packages.find] where = ["."] -include = ["rochedb*"] +include = ["koutendb*"] diff --git a/rochedb/__init__.py b/rochedb/__init__.py deleted file mode 100644 index de9f993..0000000 --- a/rochedb/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .client import EncodedPayload, PayloadCodec, RocheClient, RocheError, RocheId - -__all__ = ["EncodedPayload", "PayloadCodec", "RocheClient", "RocheError", "RocheId"] diff --git a/tests/test_driver.py b/tests/test_driver.py index 41754b1..fc6013a 100644 --- a/tests/test_driver.py +++ b/tests/test_driver.py @@ -6,28 +6,28 @@ import unittest from pathlib import Path -from rochedb import RocheClient, RocheId +from koutendb import KoutenClient, KoutenId DRIVER_ROOT = Path(__file__).resolve().parents[1] -CORE_ROOT = Path(os.environ.get("ROCHEDB_CORE_DIR", DRIVER_ROOT.parent / "rochedb")) +CORE_ROOT = Path(os.environ.get("KOUTENDB_CORE_DIR", DRIVER_ROOT.parent / "koutendb")) -class RochePythonDriverTest(unittest.TestCase): +class KoutenPythonDriverTest(unittest.TestCase): @classmethod def setUpClass(cls): cls.peers = os.environ.get( - "ROCHE_TEST_PEERS", "127.0.0.1:17831,127.0.0.1:17832" + "KOUTEN_TEST_PEERS", "127.0.0.1:17831,127.0.0.1:17832" ) cls.processes = [] - roched = CORE_ROOT / "src" / "roched" - if not roched.exists(): - raise RuntimeError(f"roched not found: {roched}") + koutend = CORE_ROOT / "src" / "koutend" + if not koutend.exists(): + raise RuntimeError(f"koutend not found: {koutend}") for i in range(2): cls.processes.append( subprocess.Popen( [ - str(roched), + str(koutend), f"--id={i}", f"--peers={cls.peers}", "--slow-tick=1000", @@ -36,7 +36,7 @@ def setUpClass(cls): ) ) - cls.client = RocheClient.connect(cls.peers, timeout=1.0) + cls.client = KoutenClient.connect(cls.peers, timeout=1.0) deadline = time.time() + 5.0 while time.time() < deadline: try: @@ -45,7 +45,7 @@ def setUpClass(cls): return except Exception: time.sleep(0.1) - raise RuntimeError("roched test cluster did not start") + raise RuntimeError("koutend test cluster did not start") @classmethod def tearDownClass(cls): @@ -68,7 +68,7 @@ def test_put_get_query_roundtrip(self): b'{"title":"Shinjuku","country":"JP"}', vector=[1.0, 0.0], ) - self.assertIsInstance(doc_id, RocheId) + self.assertIsInstance(doc_id, KoutenId) self.assertEqual( self.client.get(doc_id), b'{"title":"Shinjuku","country":"JP"}' ) @@ -103,7 +103,7 @@ def test_bif_codec_roundtrip(self): def test_batch_get_and_id_string_roundtrip(self): first = self.client.put("tenant/acme/orders", "order-1") second = self.client.put("tenant/acme/orders", "order-2") - self.assertEqual(RocheId.parse(str(first)), first) + self.assertEqual(KoutenId.parse(str(first)), first) self.assertEqual(self.client.batch_get([first, second]), [b"order-1", b"order-2"]) diff --git a/tests/test_secure.py b/tests/test_secure.py new file mode 100644 index 0000000..1286a51 --- /dev/null +++ b/tests/test_secure.py @@ -0,0 +1,69 @@ +"""Unit tests for the shared-secret transport crypto (no server required). + +These lock down the security-critical properties of secure.py: confidentiality, +authenticated integrity, and key separation. They require the `secure` extra +(PyNaCl); the whole module is skipped if it is not installed. +""" +import os +import sys +import unittest + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from koutendb import secure # noqa: E402 + +try: + import nacl.secret # noqa: F401 + + _HAVE_PYNACL = True +except ImportError: + _HAVE_PYNACL = False + + +@unittest.skipUnless(_HAVE_PYNACL, "PyNaCl (the 'secure' extra) is required") +class TransportCryptoTest(unittest.TestCase): + SECRET = "shared-secret" + CHAL = "0a1b2c3d4e5f" + + def test_ciphertext_hides_plaintext(self): + pt = b"PUTR docs 12 0 json" + ct = secure.encrypt_transport_frame(pt, self.SECRET, self.CHAL) + self.assertNotIn(pt, ct) + # 24-byte nonce + 16-byte Poly1305 tag of overhead. + self.assertEqual(len(ct), len(pt) + 24 + 16) + + def test_roundtrip(self): + pt = b"hello kouten" + ct = secure.encrypt_transport_frame(pt, self.SECRET, self.CHAL) + self.assertEqual(secure.decrypt_transport_frame(ct, self.SECRET, self.CHAL), pt) + + def test_tampered_ciphertext_is_rejected(self): + ct = bytearray(secure.encrypt_transport_frame(b"payload", self.SECRET, self.CHAL)) + ct[-1] ^= 0x01 + with self.assertRaises(Exception): + secure.decrypt_transport_frame(bytes(ct), self.SECRET, self.CHAL) + + def test_wrong_key_cannot_decrypt(self): + ct = secure.encrypt_transport_frame(b"payload", self.SECRET, self.CHAL) + with self.assertRaises(Exception): + secure.decrypt_transport_frame(ct, "other-secret", self.CHAL) + + def test_challenge_separates_transport_keys(self): + ct = secure.encrypt_transport_frame(b"payload", self.SECRET, self.CHAL) + with self.assertRaises(Exception): + secure.decrypt_transport_frame(ct, self.SECRET, "ffffffffffff") + + def test_nonce_is_randomized_per_frame(self): + pt = b"same message" + a = secure.encrypt_transport_frame(pt, self.SECRET, self.CHAL) + b = secure.encrypt_transport_frame(pt, self.SECRET, self.CHAL) + self.assertNotEqual(a, b) + + def test_response_is_lowercase_hex(self): + resp = secure.secret_response_hex("alice", "secret", self.CHAL, self.SECRET) + int(resp, 16) # must be valid hex + self.assertEqual(resp, resp.lower()) + + +if __name__ == "__main__": + unittest.main()