diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8ca0869..6159256 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -15,9 +15,10 @@ concurrency: jobs: test: timeout-minutes: 25 - runs-on: ubuntu-latest + runs-on: ${{ matrix.os }} strategy: matrix: + os: [ubuntu-latest, macos-latest] python-version: ["3.10", "3.12"] steps: - name: Checkout Python driver @@ -27,7 +28,7 @@ jobs: uses: actions/checkout@v4 with: repository: puffball1567/koutendb - ref: v0.12.0 + ref: e36b424bcfd9cd0dfa24ae121f4b4dd028b0eaac path: koutendb-core - name: Set up Python @@ -36,21 +37,37 @@ jobs: python-version: ${{ matrix.python-version }} - name: Install Nim + if: runner.os == 'Linux' uses: jiro4989/setup-nim-action@v2 with: nim-version: "2.2.10" - - name: Install dependencies + - name: Install dependencies on Linux + if: runner.os == 'Linux' run: sudo apt-get update && sudo apt-get install -y libsodium-dev + - name: Install dependencies on macOS + if: runner.os == 'macOS' + run: | + brew install nim libsodium openssl@3 + sodium="$(brew --prefix libsodium)" + ssl="$(brew --prefix openssl@3)" + echo "$ssl/bin" >> "$GITHUB_PATH" + echo "LIBRARY_PATH=$sodium/lib" >> "$GITHUB_ENV" + echo "CPATH=$sodium/include" >> "$GITHUB_ENV" + echo "DYLD_LIBRARY_PATH=$sodium/lib:$ssl/lib" >> "$GITHUB_ENV" + - name: Build koutend run: | cd koutendb-core nimble install -y - nim c -d:release --nimcache:/tmp/nimcache_koutend -o:src/koutend src/koutend.nim + nim c -d:release -d:ssl --nimcache:/tmp/nimcache_koutend -o:src/koutend src/koutend.nim - name: Install package - run: python -m pip install -e . + run: python -m pip install -e '.[secure]' - name: Run tests run: KOUTENDB_CORE_DIR="${{ github.workspace }}/koutendb-core" python -m unittest discover -s tests + + - name: Shared native TCP conformance + run: python koutendb-core/scripts/native_driver_conformance.py --server koutendb-core/src/koutend -- python "$PWD/tests/tcp_adapter.py" diff --git a/MANIFEST.in b/MANIFEST.in new file mode 100644 index 0000000..dad544f --- /dev/null +++ b/MANIFEST.in @@ -0,0 +1,3 @@ +include docs/native-tcp.md +include docs/native-tcp-validation.md +include tests/tcp_adapter.py diff --git a/README.md b/README.md index c1f168f..1d2aa39 100644 --- a/README.md +++ b/README.md @@ -6,9 +6,12 @@ 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. +Version 0.3.0 adds stricter framing, version negotiation, typed errors and safe +retry behavior. See [the TCP safety and migration guide](docs/native-tcp.md). + ## Status -- package: PyPI [`koutendb`](https://pypi.org/project/koutendb/) v0.2.1 +- package: PyPI [`koutendb`](https://pypi.org/project/koutendb/) v0.3.0 - current mode: native TCP wire driver - Python: 3.10+ - runtime dependencies: none @@ -24,9 +27,9 @@ Implemented: - codec metadata negotiation with `CODECMETA ON` - `batch_get` - direct owner redirects from extended `FWD ... owner` responses -- routed multi-node `batch_get` fallback with stable input ordering +- ordered `batch_get` using epoch-aware GETID requests - typed `KoutenId` -- one reconnect retry +- at most one reconnect retry for reads; no automatic write replay - context manager support - username/password, shared-secret transport, and TLS authentication diff --git a/docs/native-tcp-validation.md b/docs/native-tcp-validation.md new file mode 100644 index 0000000..db08ead --- /dev/null +++ b/docs/native-tcp-validation.md @@ -0,0 +1,22 @@ +# Native TCP Validation + +Date: 2026-09-19 + +Local Linux verification against the shared KoutenDB conformance harness at +core commit `e36b424bcfd9cd0dfa24ae121f4b4dd028b0eaac`: + +- All 27 scripted protocol/failure cases passed. +- All six real-server modes passed: plain, password, token, secret, TLS, TLS+secret. +- Verified Unicode, empty/binary data and 1 MiB payload round trips. +- Verified invalid credentials, certificate rejection, hostname mismatch, + bounded redirects, partial frames, timeout, disconnection and unsafe-write replay prevention. + +- Unit, two-node integration and crypto tests: 17 passed. +- Empty payloads and duplicate IDs retain their values and ordering in batch_get. +- Wheel and source distribution build: passed; transport/error modules included. + +The GitHub workflow runs the shared conformance matrix on Linux and macOS. +See the release commit's workflow checks for CI results. The local results above +are correctness/integration checks, not load or long-duration operational tests. + +See [native TCP usage and reproduction commands](native-tcp.md). diff --git a/docs/native-tcp.md b/docs/native-tcp.md new file mode 100644 index 0000000..eb03d0c --- /dev/null +++ b/docs/native-tcp.md @@ -0,0 +1,114 @@ +# Native TCP Safety Update + +Python already uses native TCP and still needs no libkoutendb or FFI. This +release hardens that transport rather than adding another one. + +```python +from koutendb import KoutenClient, IndeterminateWriteException + +with KoutenClient.connect( + ["127.0.0.1:17301"], timeout=3, + read_timeout=5, write_timeout=5, +) as db: + document_id = db.put_json("articles", {"title": "Hello"}) + print(db.get_json(document_id)) +``` + +For TLS set `tls=True`, `tls_ca_file="ca.pem"` and +`tls_server_name="db.example.com"`. Omit the CA file to use system roots. +Credentials are `username`, `password`, `auth_token`, `secret_key` and `galaxy`. +Shared-secret mode requires `pip install 'koutendb[secure]'`; ordinary TLS uses +Python's standard library. + +New exceptions inherit KoutenError, so existing broad error handlers still work: +ConnectionException, ConnectionTimeoutException, AuthenticationException, +ProtocolException, VersionMismatchException, ServerException and +IndeterminateWriteException. + +Construction remains lazy. The first operation negotiates the protocol before +use. A client serializes operations with a lock; close is now terminal. +For local development install this checkout with `pip install -e '.[secure]'`. +DNS resolution is OS-controlled and may outlast the connection timeout. + +## Behavior Changes + +Unsafe automatic write replay and all-peer miss probing have been removed. +Reads follow explicit server redirects only. Keep the peer ordering consistent +with the server cluster. + +`batch_get` now uses ordered GETID requests. Wire-v1 BGET omits the epoch and +conflates an empty value with a miss. This prioritizes correct identity and +empty-value handling, but means one request per ID rather than one BGET frame. +Do not expect the previous batch throughput. The return shape and input ordering +are unchanged. + +```sh +KOUTENDB_CORE_DIR=../koutendb python3 -m unittest discover -s tests +bash ../koutendb/scripts/native_driver_conformance.sh python3 "$PWD/tests/tcp_adapter.py" +``` + +## Server Setup + +Run a TLS-enabled `koutend` build. For a local-only first test: + +```sh +koutend --id=0 --peers=127.0.0.1:17301 --data=./kouten-data +``` + +Keep plaintext connections on localhost or an isolated, trusted private network. +A Docker network is not a substitute for access control. Use verified TLS when +traffic crosses a trust boundary. For password authentication, start the server +with `--user=app --password=...`; prefer the server's configuration/secret +management facilities for production rather than putting secrets in shell history. + +Native TCP implements wire version 1: WIREVER, CODECMETA, PUTR, GETID, QRYID, +HEALTH, authentication and bounded FWD handling. It is not a replacement for +every embedded/admin API. It uses server-provided IDs and does not calculate +ring placement or orbit ownership. Peer ordering must match the server cluster +configuration, because explicit redirect owners are node indexes. + +## Safety Contract + +- Every new connection authenticates, checks WIREVER and enables codec metadata + before sending application requests. Unsupported versions fail closed. +- Headers are bounded to 8 KiB; payload frames default to at most 64 MiB. + The configurable payload cap cannot exceed that hard limit. +- Partial reads/writes are handled. A read deadline covers the complete response, + not a fresh timeout for every fragment. +- A read may reconnect and retry once. An unknown write outcome is never retried. +- After a broken or malformed response the connection is discarded. +- Redirects default to eight hops (configurable up to 32), and an out-of-range + owner is rejected. Missing values do not trigger a scan of every server. +- CA and hostname verification are enabled by default. TLS 1.2 is the minimum. + Insecure verification bypass is explicitly development-only. +- Password/token and shared-secret challenge authentication are supported. + Library transport errors do not include raw server error text or credentials. + +A successful send is not proof that a write committed. If the connection breaks +or the reply is malformed after a PUT may have been sent, handle an +**indeterminate write** separately from a definite server rejection. Do not +blindly repeat the insert or assume a fallback database is now authoritative. +Reconcile at the application level until a server-side idempotency contract is +available. + +The pre-v1 protocol is version-checked, not promised compatible with future +versions. Authentication errors, protocol errors, connection failures, timeouts, +server rejections and indeterminate writes are distinguishable. + +## Verification + +The adapter in this repository runs against KoutenDB's language-independent +`scripts/native_driver_conformance.py` suite, pinned in CI to core commit +`e36b424bcfd9cd0dfa24ae121f4b4dd028b0eaac`. + +The shared matrix covers 27 scripted cases: fragmented/empty/Unicode/binary +responses, missing values, projections, invalid lengths/codecs/headers, redacted +server errors, version mismatches, connection loss, partial-response retry, +timeouts, backpressure, redirects and poisoned-connection disposal. +Six real-server configurations cover plaintext, password, token, shared-secret, +TLS and TLS plus shared-secret; these include 1 MiB round trips, invalid +credentials, untrusted certificates and hostname mismatch. + +These are bounded correctness/integration checks, not endurance or throughput +benchmarks. Linux results are checked locally; Linux/macOS CI must pass before +release. Existing embedded regressions remain separate from native TCP checks. diff --git a/koutendb/__init__.py b/koutendb/__init__.py index 7f2093d..ee59373 100644 --- a/koutendb/__init__.py +++ b/koutendb/__init__.py @@ -1,3 +1,10 @@ from .client import EncodedPayload, PayloadCodec, KoutenClient, KoutenError, KoutenId +from .errors import (ConnectionException, ConnectionTimeoutException, + AuthenticationException, ProtocolException, + VersionMismatchException, ServerException, + IndeterminateWriteException) -__all__ = ["EncodedPayload", "PayloadCodec", "KoutenClient", "KoutenError", "KoutenId"] +__all__ = ["EncodedPayload", "PayloadCodec", "KoutenClient", "KoutenError", "KoutenId", + "ConnectionException", "ConnectionTimeoutException", "AuthenticationException", + "ProtocolException", "VersionMismatchException", "ServerException", + "IndeterminateWriteException"] diff --git a/koutendb/client.py b/koutendb/client.py index cc9611a..f3a91de 100644 --- a/koutendb/client.py +++ b/koutendb/client.py @@ -2,25 +2,23 @@ from dataclasses import dataclass import json -import socket -import ssl +import math +import re import struct +from threading import RLock from typing import Any, Iterable, Literal, Optional -from .secure import ( - SecureState, - decrypt_transport_frame, - encrypt_transport_frame, - secret_response_hex, -) +from .errors import (KoutenError, ConnectionException, ProtocolException, + ServerException, IndeterminateWriteException) +from .transport import Connection, MAX_FRAME, expect, number - -class KoutenError(Exception): - """Raised when KoutenDB returns an error frame or the TCP connection fails.""" +PayloadCodec = Literal["raw", "json", "nif", "bif"] -PayloadCodec = Literal["raw", "json", "nif", "bif"] -_PAYLOAD_CODECS = {"raw", "json", "nif", "bif"} +def _codec(value: str) -> PayloadCodec: + if value not in ("raw", "json", "nif", "bif"): + raise ValueError("Unsupported payload codec") + return value # type: ignore[return-value] @dataclass(frozen=True) @@ -32,25 +30,26 @@ class KoutenId: period: float head: float + def __post_init__(self): + for value, bits in [(self.parent, 64), (self.epoch, 32), (self.seq, 32)]: + if type(value) is not int or not 0 <= value < (1 << bits): + raise ValueError("ID integer out of range") + if not all(math.isfinite(v) for v in (self.t_write, self.period, self.head)) or self.period <= 0: + raise ValueError("Invalid ID coordinates") + @classmethod def parse(cls, text: str) -> "KoutenId": parts = text.split(":") if len(parts) != 6: - raise ValueError("KoutenId text must have 6 ':'-separated fields") - return cls( - parent=int(parts[0]), - epoch=int(parts[1]), - seq=int(parts[2]), - t_write=float(parts[3]), - period=float(parts[4]), - head=float(parts[5]), - ) + raise ValueError("KoutenId requires six fields") + if not all(re.fullmatch(r"[0-9]+", v) for v in parts[:3]): + raise ValueError("Invalid ID integer") + if not all(re.fullmatch(r"[+-]?(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][+-]?\d+)?", v, re.ASCII) for v in parts[3:]): + raise ValueError("Invalid ID coordinate") + return cls(*(int(v) for v in parts[:3]), *(float(v) for v in parts[3:])) def __str__(self) -> str: - return ( - f"{self.parent}:{self.epoch}:{self.seq}:" - f"{self.t_write}:{self.period}:{self.head}" - ) + return f"{self.parent}:{self.epoch}:{self.seq}:{self.t_write}:{self.period}:{self.head}" @dataclass(frozen=True) @@ -61,524 +60,248 @@ class EncodedPayload: 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]] = [] + if not 1 <= len(values) <= 64: + raise ValueError("Provide 1..64 ordered peers") + result = [] for value in values: - host, sep, port = value.rpartition(":") - if not sep or not host or not port: - raise ValueError(f"invalid peer '{value}', expected host:port") - parsed.append((host, int(port))) - if not parsed: - raise ValueError("peers must not be empty") - return parsed - - -def _vec_bytes(vector: Optional[Iterable[float]]) -> bytes: - if vector is None: - return b"" - values = [float(v) for v in vector] - if not values: - return b"" - return struct.pack("<" + "f" * len(values), *values) - - -def _as_bytes(payload: bytes | bytearray | memoryview | str) -> bytes: - if isinstance(payload, str): - return payload.encode("utf-8") - 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] + match = re.fullmatch(r"(?:([a-zA-Z0-9._-]+)|\[([a-fA-F0-9:]+)\]):([0-9]{1,5})", value) + if not match or not 1 <= int(match[3]) <= 65535: + raise ValueError("Invalid TCP peer") + result.append((match[1] or match[2], int(match[3]))) + return result 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, - ): + """Synchronous native TCP client. Operations on one client are serialized. + + Construction stays lazy for compatibility; every new connection negotiates + wire version and codec metadata before sending any application request. + """ + + 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, + read_timeout: float | None = None, write_timeout: float | None = None, + max_frame_bytes: int = MAX_FRAME, max_redirects: int = 8, + retry_reads: bool = True): 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. + self.read_timeout = timeout if read_timeout is None else read_timeout + self.write_timeout = timeout if write_timeout is None else write_timeout + for value in (self.timeout, self.read_timeout, self.write_timeout): + if isinstance(value, bool) or not math.isfinite(value) or not 0 < value <= 3600: + raise ValueError("Timeout must be in (0, 3600] seconds") + if type(max_frame_bytes) is not int or not 1 <= max_frame_bytes <= MAX_FRAME: + raise ValueError("Invalid frame limit") + if type(max_redirects) is not int or not 0 <= max_redirects <= 32: + raise ValueError("Invalid redirect limit") + self.max_frame_bytes, self.max_redirects = max_frame_bytes, max_redirects + for value in (retry_reads, tls, tls_insecure_skip_verify): + if type(value) is not bool: + raise ValueError("Boolean options must be bool values") + self.retry_reads = retry_reads if auth_token and not username: username, password = "token", auth_token - self.username = username - self.password = password - self.secret_key = secret_key + for value in (username, password, galaxy): + if not isinstance(value, str) or len(value.encode()) > 1024 or any(ord(c) <= 32 or ord(c) == 127 for c in value): + raise ValueError("Invalid authentication or galaxy field") + if not username and (password or secret_key): + raise ValueError("Authentication requires username") + self.username, self.password, self.secret_key = username, password, 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_ca_file, self.tls_server_name = tls_ca_file, tls_server_name self.tls_insecure_skip_verify = tls_insecure_skip_verify - self._socks: dict[int, socket.socket] = {} - self._secure: dict[int, SecureState] = {} + self._connections: dict[int, Connection] = {} + self._closed = False + self._lock = RLock() @classmethod - 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(): - try: - sock.close() - except OSError: - pass - self._socks.clear() - self._secure.clear() - - def __enter__(self) -> "KoutenClient": + 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, + read_timeout: float | None = None, write_timeout: float | None = None, + max_frame_bytes: int = MAX_FRAME, max_redirects: int = 8, + retry_reads: bool = True) -> "KoutenClient": + return cls(peers, 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, + read_timeout=read_timeout, write_timeout=write_timeout, + max_frame_bytes=max_frame_bytes, max_redirects=max_redirects, + retry_reads=retry_reads) + + def __repr__(self): + return f"KoutenClient(closed={self._closed})" + + def _drop(self): + for connection in self._connections.values(): + connection.close() + self._connections.clear() + + def close(self): + with self._lock: + self._closed = True + self._drop() + + def __enter__(self): return self - def __exit__(self, exc_type, exc, tb) -> None: + def __exit__(self, exc_type, exc, tb): self.close() + def _connection(self, node: int) -> Connection: + if self._closed: + raise ConnectionException("TCP client is closed") + if type(node) is not int or not 0 <= node < len(self.peers): + raise ProtocolException("Node out of range") + if node not in self._connections: + self._connections[node] = Connection(self.peers[node], self) + return self._connections[node] + + def _read_operation(self, operation): + with self._lock: + for attempt in range(2): + try: + return operation() + except KoutenError as error: + self._drop() + if self._closed or not self.retry_reads or attempt or not isinstance(error, ConnectionException): + raise + def wire_version(self, node: int = 0) -> int: - parts = self._rpc(node, "WIREVER") - if len(parts) != 2 or parts[0] != "WIREVER": - raise KoutenError("WIREVER failed: " + " ".join(parts)) - return int(parts[1]) + def operation(): + parts = self._connection(node).exchange("WIREVER") + expect(parts, "WIREVER", 2) + return number(parts[1], 1) + return self._read_operation(operation) def health(self, node: int = 0) -> str: - parts = self._rpc(node, "HEALTH") - if not parts or parts[0] != "OK": - raise KoutenError("HEALTH failed: " + " ".join(parts)) - return " ".join(parts[1:]) - - def put( - self, - ring: str, - payload: bytes | bytearray | memoryview | str, - vector: Optional[Iterable[float]] = None, - codec: PayloadCodec = "raw", - node: int = 0, - ) -> KoutenId: + def operation(): + parts = self._connection(node).exchange("HEALTH") + if parts[0] == "ERR": + raise ServerException("Health request rejected") + if len(parts) < 2 or parts[0] != "OK" or not re.fullmatch(r"node=[0-9]+", parts[1]): + raise ProtocolException("Invalid health response") + return " ".join(parts[1:]) + return self._read_operation(operation) + + def put(self, ring: str, payload: bytes | bytearray | memoryview | str, + vector: Optional[Iterable[float]] = None, codec: PayloadCodec = "raw", + node: int = 0) -> KoutenId: + _codec(codec) 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} {_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 KoutenError("PUTR failed: " + " ".join(parts)) - return KoutenId( - parent=int(parts[1]), - epoch=int(parts[2]), - seq=int(parts[3]), - t_write=float(parts[4]), - period=float(parts[5]), - 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, - ) -> KoutenId: - return self.put(ring, payload, vector=vector, codec=codec, node=node) - - def put_json( - self, - ring: str, - value: Any, - vector: Optional[Iterable[float]] = None, - node: int = 0, - ) -> KoutenId: - payload = json.dumps(value, separators=(",", ":"), ensure_ascii=False) - 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, - ) -> KoutenId: - 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, - ) -> KoutenId: - return self.put(ring, payload, vector=vector, codec="bif", node=node) + payload_b = payload.encode("utf-8") if isinstance(payload, str) else bytes(payload) + values = [] if vector is None else [float(v) for v in vector] + if not all(math.isfinite(v) for v in values): + raise ValueError("Vector must contain finite numbers") + if not ring_b or len(ring_b) + len(payload_b) + len(values) * 4 > self.max_frame_bytes: + raise ValueError("Invalid ring or request exceeds limit") + body = ring_b + payload_b + struct.pack("<" + "f" * len(values), *values) + with self._lock: + attempted = [False] + try: + connection = self._connection(node) + connection.send(f"PUTR {len(ring_b)} {len(payload_b)} {len(values)} {codec}", body, attempted) + reply = connection.header() + expect(reply, "ID", 7) + try: + return KoutenId.parse(":".join(reply[1:])) + except ValueError: + raise ProtocolException("Invalid returned ID") from None + except KoutenError as error: + self._drop() + if attempted[0] and isinstance(error, (ConnectionException, ProtocolException)): + raise IndeterminateWriteException("Write outcome unknown; do not automatically retry") from None + raise + + def put_codec(self, ring: str, payload: bytes | bytearray | memoryview | str, + codec: PayloadCodec, vector: Optional[Iterable[float]] = None, + node: int = 0) -> KoutenId: + return self.put(ring, payload, vector, codec, node) + + def put_json(self, ring: str, value: Any, vector: Optional[Iterable[float]] = None, + node: int = 0) -> KoutenId: + return self.put(ring, json.dumps(value, separators=(",", ":"), ensure_ascii=False, allow_nan=False), vector, "json", node) + + def put_nif(self, ring: str, payload: bytes | bytearray | memoryview | str, + vector: Optional[Iterable[float]] = None, node: int = 0) -> KoutenId: + return self.put(ring, payload, vector, "nif", node) + + def put_bif(self, ring: str, payload: bytes | bytearray | memoryview, + vector: Optional[Iterable[float]] = None, node: int = 0) -> KoutenId: + return self.put(ring, payload, vector, "bif", node) + + def _read(self, doc_id: KoutenId, selection: str | None, node: int | None): + body = b"" if selection is None else selection.encode("utf-8") + if len(body) > self.max_frame_bytes: + raise ValueError("Selection exceeds limit") + def operation(): + current, target = doc_id, 0 if node is None else node + for redirects in range(self.max_redirects + 1): + connection = self._connection(target) + fields = str(current).replace(":", " ") + connection.send(f"GETID {fields}" if selection is None else f"QRYID {fields} {len(body)}", body) + reply = connection.header() + if reply[0] in ("MISS", "GONE"): + expect(reply, reply[0], 1) + return None + if reply[0] == "FWD": + if len(reply) not in (7, 8) or redirects == self.max_redirects: + raise ProtocolException("Invalid or excessive redirect") + try: + current = KoutenId.parse(":".join(reply[1:7])) + except ValueError: + raise ProtocolException("Invalid redirect ID") from None + if len(reply) == 8: + target = number(reply[7], len(self.peers) - 1) + continue + expect(reply, "VAL", 4) + number(reply[1], len(self.peers) - 1) + size = number(reply[2], self.max_frame_bytes) + try: + codec = _codec(reply[3]) + except ValueError: + raise ProtocolException("Unknown response codec") from None + return EncodedPayload(connection.read(size), codec) + raise ProtocolException("Excessive redirect") + return self._read_operation(operation) + + def get_encoded(self, doc_id: KoutenId, node: Optional[int] = None) -> Optional[EncodedPayload]: + return self._read(doc_id, None, node) 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: KoutenId, node: Optional[int] = None - ) -> Optional[EncodedPayload]: - return self._read_encoded_with_fallback("GETID", doc_id, b"", node=node) + result = self.get_encoded(doc_id, node) + return None if result is None else result.payload 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") + result = self.get(doc_id, node) + return None if result is None else result.decode("utf-8") 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: 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: 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: 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") + result = self.get_text(doc_id, node) + return None if result is None else json.loads(result) + + def query_encoded(self, doc_id: KoutenId, selection: str, node: Optional[int] = None) -> Optional[EncodedPayload]: + return self._read(doc_id, selection, node) + + def query(self, doc_id: KoutenId, selection: str, node: Optional[int] = None) -> Optional[bytes]: + result = self.query_encoded(doc_id, selection, node) + return None if result is None else result.payload + + def query_text(self, doc_id: KoutenId, selection: str, node: Optional[int] = None) -> Optional[str]: + result = self.query(doc_id, selection, node) + return None if result is None else result.decode("utf-8") 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[KoutenId], node: Optional[int] = None - ) -> list[Optional[bytes]]: - id_list = list(ids) - if not id_list: - return [] - if node is not None: - return self._batch_get_node(id_list, node) - - result: list[Optional[bytes]] = [None] * len(id_list) - missing = list(range(len(id_list))) - for peer_node in range(len(self.peers)): - if not missing: - break - values = self._batch_get_node( - [id_list[index] for index in missing], peer_node - ) - still_missing: list[int] = [] - for index, value in zip(missing, values): - if value is None: - still_missing.append(index) - else: - result[index] = value - missing = still_missing - return result - - def _batch_get_node( - self, id_list: list[KoutenId], node: int - ) -> list[Optional[bytes]]: - body = "".join( - f"{doc_id.parent} {doc_id.seq} {doc_id.period} {doc_id.head} {doc_id.t_write}\n" - for doc_id in id_list - ).encode("utf-8") - parts = self._rpc(node, f"BGET {len(id_list)} {len(body)}", body) - if len(parts) != 3 or parts[0] != "BVAL": - raise KoutenError("BGET failed: " + " ".join(parts)) - expected = int(parts[1]) - 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 KoutenError("BGET payload length header missing") - length = int(payload[pos:nl].decode("utf-8")) - pos = nl + 1 - if length == 0: - out.append(None) - else: - out.append(payload[pos : pos + length]) - pos += length - return out - - def _read_id_encoded( - self, - op: str, - doc_id: KoutenId, - selection: bytes, - node: int, - redirects_left: int = 2, - ) -> 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}" - ) - if op == "QRYID": - header += f" {len(selection)}" - parts = self._rpc(node, header, selection) - if not parts: - raise KoutenError(f"{op} returned an empty response") - if parts[0] in ("MISS", "GONE"): - return None - if parts[0] == "ERR": - raise KoutenError(" ".join(parts[1:])) - if parts[0] == "FWD": - if len(parts) not in (7, 8): - raise KoutenError("invalid FWD response: " + " ".join(parts)) - if redirects_left <= 0: - raise KoutenError("too many FWD redirects") - fwd = KoutenId( - parent=int(parts[1]), - epoch=int(parts[2]), - seq=int(parts[3]), - t_write=float(parts[4]), - period=float(parts[5]), - head=float(parts[6]), - ) - target_node = int(parts[7]) if len(parts) == 8 else node - return self._read_id_encoded( - op, - fwd, - selection, - node=target_node, - redirects_left=redirects_left - 1, - ) - if parts[0] != "VAL" or len(parts) not in (3, 4): - 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(node, int(parts[2])), codec) - - 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: KoutenId, selection: bytes, node: Optional[int] - ) -> Optional[bytes]: - if node is not None: - return self._read_id(op, doc_id, selection, node=node) - first = self._read_id(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(op, doc_id, selection, node=peer_node) - if value is not None: - return value - return None - - def _read_encoded_with_fallback( - 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) - 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): - try: - 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 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): - raise IndexError(f"node out of range: {node}") - sock = self._socks.get(node) - if sock is not None: - return sock - host, port = self.peers[node] - sock = socket.create_connection((host, port), timeout=self.timeout) - sock.settimeout(self.timeout) - try: - 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: - self._handshake(node) - except BaseException: - self._drop_socket(node) - raise - 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: - sock.close() - except OSError: - pass - - 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 KoutenError("connection closed") - if b == b"\n": - return b"".join(chunks).decode("utf-8") - chunks.append(b) - - 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 KoutenError("connection closed") - chunks.append(chunk) - remaining -= len(chunk) - return b"".join(chunks) + result = self.query_text(doc_id, selection, node) + return None if result is None else json.loads(result) + + def batch_get(self, ids: Iterable[KoutenId], node=None) -> list[Optional[bytes]]: + # BGET v1 omits epochs and conflates empty payloads with misses. GETID + # preserves identity, empty values, and server-directed ownership. + return [self.get(doc_id, node) for doc_id in ids] diff --git a/koutendb/errors.py b/koutendb/errors.py new file mode 100644 index 0000000..7aa0256 --- /dev/null +++ b/koutendb/errors.py @@ -0,0 +1,30 @@ +class KoutenError(Exception): + """Base class for sanitized KoutenDB transport errors.""" + + +class ConnectionException(KoutenError): + pass + + +class ConnectionTimeoutException(ConnectionException): + pass + + +class AuthenticationException(KoutenError): + pass + + +class ProtocolException(KoutenError): + pass + + +class VersionMismatchException(ProtocolException): + pass + + +class ServerException(KoutenError): + pass + + +class IndeterminateWriteException(KoutenError): + """The request may have committed. Do not automatically repeat it.""" diff --git a/koutendb/transport.py b/koutendb/transport.py new file mode 100644 index 0000000..eb38ca8 --- /dev/null +++ b/koutendb/transport.py @@ -0,0 +1,187 @@ +"""Bounded synchronous framing; independent of database placement logic.""" +from __future__ import annotations + +import socket +import ssl +import time + +from .errors import (ConnectionException, ConnectionTimeoutException, + ProtocolException, AuthenticationException, ServerException, + VersionMismatchException) +from .secure import (SecureState, encrypt_transport_frame, + decrypt_transport_frame, secret_response_hex) + +HEADER_LIMIT = 8192 +MAX_FRAME = 64 * 1024 * 1024 + + +def number(text: str, maximum: int) -> int: + if not text or len(text) > 20 or not text.isascii() or not text.isdecimal(): + raise ProtocolException("Invalid unsigned wire integer") + value = int(text) + if value > maximum: + raise ProtocolException("Wire integer exceeds limit") + return value + + +def expect(parts: list[str], tag: str, count: int, auth: bool = False) -> None: + if parts[0] == "ERR": + if auth: + raise AuthenticationException("Authentication or galaxy rejected") + raise ServerException("Server rejected request") + if len(parts) != count or parts[0] != tag: + raise ProtocolException("Invalid wire response") + + +class Connection: + def __init__(self, peer: tuple[str, int], config): + self.config = config + self.secure: SecureState | None = None + self.deadline = 0.0 + self.socket: socket.socket | None = None + try: + sock = socket.create_connection(peer, timeout=config.timeout) + self.socket = sock + sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + if config.tls: + context = ssl.create_default_context(cafile=config.tls_ca_file or None) + context.minimum_version = ssl.TLSVersion.TLSv1_2 + if config.tls_insecure_skip_verify: + context.check_hostname = False + context.verify_mode = ssl.CERT_NONE + self.socket = context.wrap_socket( + sock, server_hostname=config.tls_server_name or peer[0]) + self._handshake() + except TimeoutError: + self.close() + raise ConnectionTimeoutException("TCP connection timed out") from None + except (OSError, ValueError): + self.close() + raise ConnectionException("Unable to connect or configure TLS") from None + except BaseException: + self.close() + raise + + def close(self): + if self.socket is not None: + self.socket.close() + self.socket = None + self.secure = None + + def send(self, header: str, body: bytes = b"", attempt: list[bool] | None = None): + if len(header) > HEADER_LIMIT or any(c in header for c in "\r\n\0"): + raise ValueError("Invalid request header") + if len(body) > self.config.max_frame_bytes: + raise ValueError("Request exceeds frame limit") + if self.socket is None: + raise ConnectionException("TCP connection closed") + frame = header.encode("utf-8") + b"\n" + body + if self.secure: + frame = encrypt_transport_frame(frame, self.secure.secret_key, + self.secure.challenge_hex) + frame = f"SEC {len(frame)}\n".encode("ascii") + frame + try: + self.socket.settimeout(self.config.write_timeout) + if attempt is not None: + attempt[0] = True + self.socket.sendall(frame) + except TimeoutError: + raise ConnectionTimeoutException("TCP write timed out") from None + except OSError: + raise ConnectionException("TCP write failed") from None + self.deadline = time.monotonic() + self.config.read_timeout + + def _raw(self, n: int) -> bytes: + result = bytearray() + while len(result) < n: + remaining = self.deadline - time.monotonic() + if remaining <= 0: + raise ConnectionTimeoutException("TCP read timed out") + if self.socket is None: + raise ConnectionException("TCP connection closed") + try: + self.socket.settimeout(remaining) + chunk = self.socket.recv(min(n - len(result), 65536)) + except TimeoutError: + raise ConnectionTimeoutException("TCP read timed out") from None + except OSError: + raise ConnectionException("TCP read failed") from None + if not chunk: + raise ConnectionException("TCP connection closed") + result.extend(chunk) + return bytes(result) + + @staticmethod + def _parse_line(line: bytes) -> list[str]: + line = line.removesuffix(b"\r") + if not line or any(b < 32 or b > 126 for b in line): + raise ProtocolException("Invalid response header") + return line.decode("ascii").split(" ") + + def _line(self, reader) -> list[str]: + line = bytearray() + while True: + byte = reader(1) + if byte == b"\n": + return self._parse_line(bytes(line)) + if len(line) >= HEADER_LIMIT: + raise ProtocolException("Response header exceeds limit") + line.extend(byte) + + def _fill(self): + parts = self._line(self._raw) + expect(parts, "SEC", 2) + n = number(parts[1], self.config.max_frame_bytes + HEADER_LIMIT + 41) + if n < 40: + raise ProtocolException("Invalid encrypted frame length") + ciphertext = self._raw(n) + state = self.secure + assert state is not None + try: + plaintext = decrypt_transport_frame(ciphertext, state.secret_key, + state.challenge_hex) + except Exception: + raise ProtocolException("Encrypted frame authentication failed") from None + if len(state.buffer) + len(plaintext) > self.config.max_frame_bytes + HEADER_LIMIT + 1: + raise ProtocolException("Encrypted response exceeds limit") + state.buffer += plaintext + + def read(self, n: int) -> bytes: + if self.secure is None: + return self._raw(n) + while len(self.secure.buffer) < n: + self._fill() + result = self.secure.buffer[:n] + self.secure.buffer = self.secure.buffer[n:] + return result + + def header(self) -> list[str]: + return self._line(self.read) + + def exchange(self, header: str) -> list[str]: + self.send(header) + return self.header() + + def _handshake(self): + c = self.config + if c.username: + if c.secret_key: + chal = self.exchange("AUTHCHAL " + c.username) + expect(chal, "CHAL", 2, auth=True) + if len(chal[1]) != 64 or any(x not in "0123456789abcdefABCDEF" for x in chal[1]): + raise ProtocolException("Invalid authentication challenge") + response = secret_response_hex(c.username, c.password, chal[1], c.secret_key) + expect(self.exchange("AUTHRESP " + response), "OK", 2, auth=True) + self.secure = SecureState(c.secret_key, chal[1]) + else: + expect(self.exchange(f"AUTH {c.username} {c.password}"), "OK", 2, auth=True) + if c.galaxy: + expect(self.exchange("HELLO " + c.galaxy), "OK", 2, auth=True) + version = self.exchange("WIREVER") + expect(version, "WIREVER", 2, auth=True) + if version[1] != "1": + raise VersionMismatchException("Unsupported KoutenDB wire version") + reply = self.exchange("CODECMETA ON") + expect(reply, "OK", 2) + if reply[1] != "codec-metadata": + raise ProtocolException("Codec metadata negotiation failed") diff --git a/pyproject.toml b/pyproject.toml index cebc296..758a44e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "koutendb" -version = "0.2.1" +version = "0.3.0" description = "Pure Python TCP driver for KoutenDB" readme = "README.md" requires-python = ">=3.10" diff --git a/tests/tcp_adapter.py b/tests/tcp_adapter.py new file mode 100644 index 0000000..7e02c29 --- /dev/null +++ b/tests/tcp_adapter.py @@ -0,0 +1,55 @@ +"""JSONL adapter for the shared core native-driver conformance suite.""" +import base64 +import json +from pathlib import Path +import sys + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from koutendb import KoutenClient, KoutenId + +client = None +names = { + "authToken": "auth_token", "secretKey": "secret_key", "tlsCaFile": "tls_ca_file", + "tlsServerName": "tls_server_name", "tlsInsecureSkipVerify": "tls_insecure_skip_verify", + "readTimeout": "read_timeout", "writeTimeout": "write_timeout", + "maxFrameBytes": "max_frame_bytes", "maxRedirects": "max_redirects", "retryReads": "retry_reads", +} + +for line in sys.stdin: + try: + request = json.loads(line) + op = request["op"] + if op == "connect": + if client: + client.close() + options = {"timeout": 1, "read_timeout": 1, "write_timeout": 1} + options.update({names.get(k, k): v for k, v in request.get("options", {}).items()}) + for key in ("timeout", "readTimeout", "writeTimeout"): + if key in request: + options[names.get(key, key)] = request[key] + client = KoutenClient.connect(request["peers"], **options) + client._connection(0) + result = "connected" + elif op == "close": + client.close() + result = "closed" + elif op == "debug": + result = repr(client) + elif op == "health": + result = client.health() + elif op == "put": + result = str(client.put_codec(request["ring"], base64.b64decode(request["payload"]), request.get("codec", "raw"))) + elif op == "putJson": + result = str(client.put_json(request["ring"], request["value"])) + elif op == "get": + value = client.get_encoded(KoutenId.parse(request["id"])) + result = None if value is None else {"payload": base64.b64encode(value.payload).decode(), "codec": value.codec} + elif op == "getJson": + result = client.get_json(KoutenId.parse(request["id"])) + elif op == "query": + result = client.query_json(KoutenId.parse(request["id"]), request["selection"]) + else: + raise ValueError("Unsupported adapter operation") + print(json.dumps({"ok": True, "result": result}), flush=True) + except Exception as error: + print(json.dumps({"ok": False, "error": type(error).__name__, "message": str(error)}), flush=True) diff --git a/tests/test_driver.py b/tests/test_driver.py index fc6013a..a599355 100644 --- a/tests/test_driver.py +++ b/tests/test_driver.py @@ -106,6 +106,11 @@ def test_batch_get_and_id_string_roundtrip(self): self.assertEqual(KoutenId.parse(str(first)), first) self.assertEqual(self.client.batch_get([first, second]), [b"order-1", b"order-2"]) + def test_batch_preserves_empty_payload_and_duplicates(self): + empty = self.client.put("docs/empty", b"") + self.assertEqual(self.client.batch_get([empty, empty]), [b"", b""]) + self.assertEqual(self.client.batch_get([]), []) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_safety.py b/tests/test_safety.py new file mode 100644 index 0000000..38ddbe2 --- /dev/null +++ b/tests/test_safety.py @@ -0,0 +1,38 @@ +import unittest +from koutendb import KoutenClient, KoutenId, ConnectionException +from koutendb.transport import number +from koutendb.errors import ProtocolException + + +class SafetyTests(unittest.TestCase): + def test_id_bounds(self): + value = "18446744073709551615:4294967295:4294967295:1:60:0" + self.assertEqual(KoutenId.parse(value).parent, 2**64 - 1) + for value in ["-1:0:1:1:60:0", "1:0:1:nan:60:0", "1:0:1:1:0:0", "1:4294967296:1:1:60:0"]: + with self.subTest(value=value), self.assertRaises(ValueError): + KoutenId.parse(value) + + def test_invalid_configuration(self): + for options in [{"timeout": 0}, {"read_timeout": float("nan")}, + {"write_timeout": -1}, {"max_frame_bytes": 2**40}, + {"max_redirects": 100}, {"username": "u\nHEALTH"}]: + with self.subTest(options=options), self.assertRaises(ValueError): + KoutenClient("localhost:17301", **options) + + def test_closed_and_redacted(self): + client = KoutenClient("localhost:17301", username="u", password="sensitive") + self.assertNotIn("sensitive", repr(client)) + client.close() + client.close() + with self.assertRaises(ConnectionException): + client.health() + + def test_wire_integer(self): + self.assertEqual(number("0", 100), 0) + for value in ["-1", "1e3", "101", "\uff11"]: + with self.subTest(value=value), self.assertRaises(ProtocolException): + number(value, 100) + + +if __name__ == "__main__": + unittest.main()