Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 9 additions & 5 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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 }"))
```

Expand All @@ -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?

Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
4 changes: 2 additions & 2 deletions rochedb/__init__.py
Original file line number Diff line number Diff line change
@@ -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"]
99 changes: 92 additions & 7 deletions rochedb/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]] = []
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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))
Expand All @@ -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,
Expand All @@ -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")
Expand All @@ -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]:
Expand Down Expand Up @@ -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}"
Expand All @@ -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]
Expand All @@ -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):
Expand Down Expand Up @@ -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

Expand Down
20 changes: 20 additions & 0 deletions tests/test_driver.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import os
import json
import subprocess
import sys
import time
Expand Down Expand Up @@ -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")
Expand Down
Loading