From 7f743882e13b03e180b6f84960d238df44390263 Mon Sep 17 00:00:00 2001 From: puffball1567 Date: Mon, 13 Jul 2026 23:10:08 +0900 Subject: [PATCH] feat: add codec-aware wire API --- README.md | 14 ++++--- pyproject.toml | 2 +- rochedb/__init__.py | 4 +- rochedb/client.py | 99 ++++++++++++++++++++++++++++++++++++++++---- tests/test_driver.py | 20 +++++++++ 5 files changed, 124 insertions(+), 15 deletions(-) diff --git a/README.md b/README.md index 31c1f45..8c98643 100644 --- a/README.md +++ b/README.md @@ -8,7 +8,7 @@ Applications pass a human-readable ring name, and RocheDB returns a typed ID. ## Status -- package: PyPI [`rochedb`](https://pypi.org/project/rochedb/) v0.1.2 +- package: PyPI [`rochedb`](https://pypi.org/project/rochedb/) v0.1.3 - current mode: native TCP wire driver - Python: 3.10+ - runtime dependencies: none @@ -18,9 +18,10 @@ Implemented: - persistent TCP connections - `wire_version` / `health` -- `put` / `put_json` -- `get` / `get_text` / `get_json` -- `query` / `query_text` / `query_json` +- `put` / `put_codec` / `put_json` / `put_nif` / `put_bif` +- `get` / `get_encoded` / `get_text` / `get_json` +- `query` / `query_encoded` / `query_text` / `query_json` +- codec metadata negotiation with `CODECMETA ON` - `batch_get` - typed `RocheId` - one reconnect retry @@ -30,6 +31,7 @@ Planned: - authentication / secret-key handshake support - retrieve / atlas wire APIs once the public wire contract is finalized for drivers +- ring-read filters/projection once the public wire contract is finalized for drivers - connection pooling ## Install @@ -68,6 +70,7 @@ with RocheClient.connect("127.0.0.1:17301") as db: ) print(db.get_json(doc_id)) + print(db.get_encoded(doc_id).codec) print(db.query_json(doc_id, "{ title }")) ``` @@ -80,7 +83,8 @@ ROCHEDB_CORE_DIR=/path/to/rochedb python3 -m unittest discover -s tests ``` The test starts a two-node local `roched` cluster and verifies put/get/query, -JSON helpers, `wire_version`, and `batch_get`. +JSON helpers, codec metadata, BIF opaque payloads, `wire_version`, and +`batch_get`. ## Why A Native Wire Driver? diff --git a/pyproject.toml b/pyproject.toml index 84ace2f..2805a55 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "rochedb" -version = "0.1.2" +version = "0.1.3" description = "Pure Python TCP driver for RocheDB" readme = "README.md" requires-python = ">=3.10" diff --git a/rochedb/__init__.py b/rochedb/__init__.py index 50a1d0d..de9f993 100644 --- a/rochedb/__init__.py +++ b/rochedb/__init__.py @@ -1,3 +1,3 @@ -from .client import RocheClient, RocheError, RocheId +from .client import EncodedPayload, PayloadCodec, RocheClient, RocheError, RocheId -__all__ = ["RocheClient", "RocheError", "RocheId"] +__all__ = ["EncodedPayload", "PayloadCodec", "RocheClient", "RocheError", "RocheId"] diff --git a/rochedb/client.py b/rochedb/client.py index b51772c..0af337f 100644 --- a/rochedb/client.py +++ b/rochedb/client.py @@ -4,13 +4,17 @@ import json import socket import struct -from typing import Any, Iterable, Optional +from typing import Any, Iterable, Literal, Optional class RocheError(Exception): """Raised when RocheDB returns an error frame or the TCP connection fails.""" +PayloadCodec = Literal["raw", "json", "nif", "bif"] +_PAYLOAD_CODECS = {"raw", "json", "nif", "bif"} + + @dataclass(frozen=True) class RocheId: parent: int @@ -41,6 +45,12 @@ def __str__(self) -> str: ) +@dataclass(frozen=True) +class EncodedPayload: + payload: bytes + codec: PayloadCodec + + def _parse_peers(peers: str | Iterable[str]) -> list[tuple[str, int]]: values = peers.split(",") if isinstance(peers, str) else list(peers) parsed: list[tuple[str, int]] = [] @@ -69,6 +79,12 @@ def _as_bytes(payload: bytes | bytearray | memoryview | str) -> bytes: return bytes(payload) +def _codec(codec: str) -> PayloadCodec: + if codec not in _PAYLOAD_CODECS: + raise ValueError(f"unsupported payload codec: {codec}") + return codec # type: ignore[return-value] + + class RocheClient: def __init__(self, peers: str | Iterable[str], timeout: float = 10.0): self.peers = _parse_peers(peers) @@ -110,13 +126,14 @@ def put( ring: str, payload: bytes | bytearray | memoryview | str, vector: Optional[Iterable[float]] = None, + codec: PayloadCodec = "raw", node: int = 0, ) -> RocheId: ring_b = ring.encode("utf-8") payload_b = _as_bytes(payload) vec_b = _vec_bytes(vector) vec_dim = len(vec_b) // 4 - header = f"PUTR {len(ring_b)} {len(payload_b)} {vec_dim}" + 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)) @@ -129,6 +146,16 @@ def put( head=float(parts[6]), ) + def put_codec( + self, + ring: str, + payload: bytes | bytearray | memoryview | str, + codec: PayloadCodec, + vector: Optional[Iterable[float]] = None, + node: int = 0, + ) -> RocheId: + return self.put(ring, payload, vector=vector, codec=codec, node=node) + def put_json( self, ring: str, @@ -137,11 +164,34 @@ def put_json( node: int = 0, ) -> RocheId: payload = json.dumps(value, separators=(",", ":"), ensure_ascii=False) - return self.put(ring, payload, vector=vector, node=node) + return self.put(ring, payload, vector=vector, codec="json", node=node) + + def put_nif( + self, + ring: str, + payload: bytes | bytearray | memoryview | str, + vector: Optional[Iterable[float]] = None, + node: int = 0, + ) -> RocheId: + return self.put(ring, payload, vector=vector, codec="nif", node=node) + + def put_bif( + self, + ring: str, + payload: bytes | bytearray | memoryview, + vector: Optional[Iterable[float]] = None, + node: int = 0, + ) -> RocheId: + return self.put(ring, payload, vector=vector, codec="bif", node=node) def get(self, doc_id: RocheId, 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 + ) -> 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]: value = self.get(doc_id, node=node) return None if value is None else value.decode("utf-8") @@ -156,6 +206,12 @@ def query( 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 + ) -> 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 ) -> Optional[str]: @@ -192,7 +248,9 @@ def batch_get(self, ids: Iterable[RocheId], node: int = 0) -> list[Optional[byte pos += length return out - def _read_id(self, op: str, doc_id: RocheId, selection: bytes, node: int) -> Optional[bytes]: + def _read_id_encoded( + self, op: str, doc_id: RocheId, selection: bytes, node: int + ) -> Optional[EncodedPayload]: header = ( f"{op} {doc_id.parent} {doc_id.epoch} {doc_id.seq} " f"{doc_id.t_write} {doc_id.period} {doc_id.head}" @@ -217,10 +275,15 @@ def _read_id(self, op: str, doc_id: RocheId, selection: bytes, node: int) -> Opt period=float(parts[5]), head=float(parts[6]), ) - return self._read_id(op, fwd, selection, node=node) - if parts[0] != "VAL" or len(parts) != 3: + 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)) - return self._read_exact(self._socket_for(node), int(parts[2])) + 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) + + def _read_id(self, op: str, doc_id: RocheId, 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] @@ -236,6 +299,20 @@ def _read_with_fallback( return value return None + def _read_encoded_with_fallback( + self, op: str, doc_id: RocheId, selection: bytes, node: Optional[int] + ) -> Optional[EncodedPayload]: + if node is not None: + return self._read_id_encoded(op, doc_id, selection, node=node) + first = self._read_id_encoded(op, doc_id, selection, node=0) + if first is not None or len(self.peers) == 1: + return first + for peer_node in range(1, len(self.peers)): + value = self._read_id_encoded(op, doc_id, selection, node=peer_node) + if value is not None: + return value + return None + def _rpc(self, node: int, header: str, payload: bytes = b"") -> list[str]: last_error: Optional[BaseException] = None for attempt in range(2): @@ -263,6 +340,14 @@ def _socket_for(self, node: int) -> socket.socket: sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) except OSError: pass + 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)) + except BaseException: + sock.close() + raise self._socks[node] = sock return sock diff --git a/tests/test_driver.py b/tests/test_driver.py index e668fee..41754b1 100644 --- a/tests/test_driver.py +++ b/tests/test_driver.py @@ -1,4 +1,5 @@ import os +import json import subprocess import sys import time @@ -79,6 +80,25 @@ def test_json_helpers(self): ) self.assertEqual(self.client.get_json(doc_id)["title"], "Python driver") self.assertEqual(self.client.query_json(doc_id, "{ kind }"), {"kind": "example"}) + encoded = self.client.get_encoded(doc_id) + self.assertIsNotNone(encoded) + assert encoded is not None + self.assertEqual(encoded.codec, "json") + self.assertEqual(json.loads(encoded.payload.decode("utf-8"))["title"], "Python driver") + + projected = self.client.query_encoded(doc_id, "{ title }") + self.assertIsNotNone(projected) + assert projected is not None + self.assertEqual(projected.codec, "json") + self.assertEqual(json.loads(projected.payload.decode("utf-8")), {"title": "Python driver"}) + + def test_bif_codec_roundtrip(self): + doc_id = self.client.put_bif("artifacts/bif", b"\x01\x02\x03\x04") + encoded = self.client.get_encoded(doc_id) + self.assertIsNotNone(encoded) + assert encoded is not None + self.assertEqual(encoded.codec, "bif") + self.assertEqual(encoded.payload, b"\x01\x02\x03\x04") def test_batch_get_and_id_string_roundtrip(self): first = self.client.put("tenant/acme/orders", "order-1")