From 119746a0024b7eee07c1fd23696d9afdd2bbab4a Mon Sep 17 00:00:00 2001 From: Jeff Repanich Date: Sat, 8 Aug 2026 08:05:33 -0400 Subject: [PATCH 1/4] feat!: rebuild Python client for full parity --- .github/workflows/ci.yml | 153 ++--- README.md | 154 ++--- benchmarks/__init__.py | 1 + benchmarks/hotpath.py | 45 ++ compose.yml | 51 ++ docs/README.md | 41 +- pyproject.toml | 19 +- src/fitz_py/__init__.py | 203 +----- src/fitz_py/_runtime.py | 190 ++++++ src/fitz_py/client.py | 174 +++-- src/fitz_py/connection.py | 632 ++++++++++-------- src/fitz_py/domains/__init__.py | 93 +-- src/fitz_py/domains/_routes.py | 39 +- src/fitz_py/domains/_subscriptions.py | 98 +++ src/fitz_py/domains/kv.py | 386 ++++++----- src/fitz_py/domains/lease.py | 359 +++++----- src/fitz_py/domains/notice.py | 221 +++--- src/fitz_py/domains/queue.py | 326 +++++---- src/fitz_py/domains/rpc.py | 475 ++++++------- src/fitz_py/domains/schedule.py | 320 ++++----- src/fitz_py/domains/stream.py | 295 ++++---- src/fitz_py/errors.py | 391 +++-------- src/fitz_py/multiplexer.py | 127 ++-- src/fitz_py/protocol/buffer.py | 34 +- src/fitz_py/protocol/frame.py | 9 +- src/fitz_py/protocol/messages.py | 4 +- src/fitz_py/protocol/response.py | 37 +- src/fitz_py/py.typed | 1 + src/fitz_py/transport/base.py | 4 + src/fitz_py/transport/factory.py | 16 +- src/fitz_py/transport/tcp.py | 15 +- src/fitz_py/transport/websocket.py | 28 +- src/fitz_py/types.py | 136 +++- .../cross-language-conformance-suite.yaml | 270 ++++++++ tests/conformance/test_conformance.py | 302 +++------ tests/integration/test_broker.py | 37 + tests/integration/test_kv.py | 26 - tests/integration/test_lease.py | 31 - tests/integration/test_notice.py | 83 --- tests/integration/test_queue.py | 27 - tests/integration/test_rpc.py | 27 - tests/integration/test_schedule.py | 74 -- tests/integration/test_stream.py | 108 --- tests/integration/test_transport.py | 17 - tests/unit/test_client.py | 44 -- tests/unit/test_connection_contract.py | 202 ------ tests/unit/test_contracts.py | 197 ++++++ tests/unit/test_errors.py | 24 - tests/unit/test_frame.py | 21 - tests/unit/test_kv_transaction.py | 58 -- tests/unit/test_multiplexer.py | 59 -- tests/unit/test_public_surface.py | 124 ---- tests/unit/test_response.py | 22 - tests/unit/test_route_validation.py | 292 -------- tests/unit/test_stream_session.py | 257 ------- tests/unit/test_websocket_transport.py | 33 - 56 files changed, 3186 insertions(+), 4226 deletions(-) create mode 100644 benchmarks/__init__.py create mode 100644 benchmarks/hotpath.py create mode 100644 compose.yml create mode 100644 src/fitz_py/_runtime.py create mode 100644 src/fitz_py/domains/_subscriptions.py create mode 100644 src/fitz_py/py.typed create mode 100644 tests/conformance/cross-language-conformance-suite.yaml create mode 100644 tests/integration/test_broker.py delete mode 100644 tests/integration/test_kv.py delete mode 100644 tests/integration/test_lease.py delete mode 100644 tests/integration/test_notice.py delete mode 100644 tests/integration/test_queue.py delete mode 100644 tests/integration/test_rpc.py delete mode 100644 tests/integration/test_schedule.py delete mode 100644 tests/integration/test_stream.py delete mode 100644 tests/integration/test_transport.py delete mode 100644 tests/unit/test_client.py delete mode 100644 tests/unit/test_connection_contract.py create mode 100644 tests/unit/test_contracts.py delete mode 100644 tests/unit/test_errors.py delete mode 100644 tests/unit/test_frame.py delete mode 100644 tests/unit/test_kv_transaction.py delete mode 100644 tests/unit/test_multiplexer.py delete mode 100644 tests/unit/test_public_surface.py delete mode 100644 tests/unit/test_response.py delete mode 100644 tests/unit/test_route_validation.py delete mode 100644 tests/unit/test_stream_session.py delete mode 100644 tests/unit/test_websocket_transport.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c329e39..4c7561d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -1,80 +1,63 @@ name: CI on: + workflow_dispatch: push: pull_request: jobs: fast: runs-on: ubuntu-latest + strategy: + matrix: + python-version: ["3.11", "3.12", "3.13"] steps: - - name: Check out repository - uses: actions/checkout@v6 - - - name: Set up Python - uses: actions/setup-python@v6 + - uses: actions/checkout@v7 + - uses: actions/setup-python@v6 with: - python-version: "3.12" - - - name: Install package and dev dependencies - run: | - python -m pip install --upgrade pip - python -m pip install -e ".[dev]" - - - name: Format - run: python -m ruff format --check . - - - name: Lint - run: python -m ruff check . - - - name: Run unit tests - run: python -m pytest tests/unit + python-version: ${{ matrix.python-version }} + cache: pip + - run: python -m pip install --upgrade pip + - run: python -m pip install -e ".[dev]" + - run: python -m ruff format --check . + - run: python -m ruff check . + - run: python -m pyright src + - run: python -m pytest tests/unit --cov=fitz_py --cov-report=term-missing package: - runs-on: ubuntu-latest needs: fast + runs-on: ubuntu-latest steps: - - name: Check out repository - uses: actions/checkout@v6 - - - name: Set up Python - uses: actions/setup-python@v6 + - uses: actions/checkout@v7 + - uses: actions/setup-python@v6 with: python-version: "3.12" - - - name: Build distributions - run: | - python -m pip install --upgrade pip build - python -m build - - - name: Smoke install built wheel + cache: pip + - run: python -m pip install --upgrade pip build + - run: python -m build + - name: Smoke test wheel run: | - python -m venv artifacts/smoke/.venv - . artifacts/smoke/.venv/bin/activate - python -m pip install --upgrade pip - python -m pip install dist/*.whl - python - <<'PY' + python -m venv artifacts/smoke + artifacts/smoke/bin/python -m pip install dist/*.whl + artifacts/smoke/bin/python - <<'PY' from fitz_py import Client, ClientConfig - - client = Client(ClientConfig(url="ws://localhost:4190/ws")) - print(type(client).__name__) - print(type(client.kv()).__name__) + client = Client(ClientConfig(url="tcp://localhost:4191")) + assert client.kv is client.kv + assert client.queue is client.queue PY - - - name: Upload distributions - uses: actions/upload-artifact@v4 + - uses: actions/upload-artifact@v4 with: name: fitz-py-dist path: dist/* - spec: - runs-on: ubuntu-latest + broker: needs: fast + runs-on: ubuntu-latest strategy: fail-fast: false matrix: - transport: [ws, tcp] - auth_mode: [anonymous, valid_jwt] + transport: [tcp, ws] + auth-mode: [anonymous, valid_jwt] env: FITZ_BROKER_ANON_TCP_ADDR: localhost:4191 FITZ_BROKER_ANON_WS_ADDR: ws://localhost:4190/ws @@ -82,57 +65,41 @@ jobs: FITZ_BROKER_AUTH_WS_ADDR: ws://localhost:4090/ws FITZ_BROKER_JWT_HMAC_SECRET: test-secret-key FITZ_BROKER_JWT_AUDIENCE: fitz - CONFORMANCE_OUTPUT: artifacts/conformance-results.json + CONFORMANCE_TRANSPORT: ${{ matrix.transport }} + CONFORMANCE_AUTH_MODE: ${{ matrix.auth-mode }} + CONFORMANCE_OUTPUT: artifacts/conformance-${{ matrix.transport }}-${{ matrix.auth-mode }}.json steps: - - name: Check out repository - uses: actions/checkout@v6 - - - name: Set up Python - uses: actions/setup-python@v6 + - uses: actions/checkout@v7 + - uses: actions/setup-python@v6 with: python-version: "3.12" - - - name: Install package and dev dependencies + cache: pip + - run: python -m pip install --upgrade pip + - run: python -m pip install -e ".[dev]" + - run: docker compose up -d + - name: Wait for brokers run: | - python -m pip install --upgrade pip - python -m pip install -e ".[dev]" - - - name: Start broker stack - run: docker compose -f ../fitz-go/compose.yml up -d - - - name: Run conformance suite - env: - CONFORMANCE_TRANSPORT: ${{ matrix.transport }} - CONFORMANCE_AUTH_MODE: ${{ matrix.auth_mode }} - run: python -m pytest tests/conformance -v - - - name: Run invalid JWT auth scenario - run: | - CONFORMANCE_TRANSPORT=${{ matrix.transport }} CONFORMANCE_AUTH_MODE=anonymous python -m pytest tests/conformance/test_conformance.py -k cs002 -v - - - name: Enforce full conformance + for port in 4090 4091 4190 4191; do + timeout 60 bash -c "until (echo > /dev/tcp/127.0.0.1/$port) 2>/dev/null; do sleep 1; done" + done + - run: python -m pytest tests/integration -v + - run: python -m pytest tests/conformance -v + - name: Enforce canonical conformance run: | python - <<'PY' - import json + import json, os from pathlib import Path - required = {f"CS-{index:03d}" for index in range(1, 16)} - path = Path("artifacts/conformance-results.json") - data = json.loads(path.read_text(encoding="utf-8")) - observed = {s["scenario_id"] for s in data["scenarios"]} - missing = sorted(required - observed) - failing = sorted(s["scenario_id"] for s in data["scenarios"] if s["scenario_id"] in required and s["verdict"] != "pass") - print(f"transport={data['transport']} auth={data['auth_mode']} p0={data['p0_pass_rate']:.0%} p1={data['p1_pass_rate']:.0%} failing={failing} missing={missing}") - if data.get("overall_status") != "pass" or failing or missing: - raise SystemExit(1) + data = json.loads(Path(os.environ["CONFORMANCE_OUTPUT"]).read_text()) + required = {f"CS-{index:03d}" for index in range(1, 18)} + observed = {item["scenario_id"] for item in data["scenarios"] if item["verdict"] == "pass"} + assert data["overall_status"] == "pass" + assert observed == required, sorted(required - observed) PY - - - name: Upload conformance results + - uses: actions/upload-artifact@v4 if: always() - uses: actions/upload-artifact@v4 with: - name: fitz-py-conformance-${{ matrix.transport }}-${{ matrix.auth_mode }} - path: artifacts/conformance-results.json - - - name: Stop broker stack - if: always() - run: docker compose -f ../fitz-go/compose.yml down --volumes + name: conformance-${{ matrix.transport }}-${{ matrix.auth-mode }} + path: artifacts/conformance-*.json + if-no-files-found: warn + - if: always() + run: docker compose down --volumes diff --git a/README.md b/README.md index f54edd6..21f4dfd 100644 --- a/README.md +++ b/README.md @@ -1,135 +1,97 @@ # fitz-py -`fitz-py` is the async-first Python SDK for Fitz. - -## Install +`fitz-py` is the typed, asyncio-native Python client for the Fitz broker. Version 0.2 is a +deliberate clean break: clients are configured once, domain clients are cached properties, streamed +results are async iterators, and network/runtime queues are bounded. ```bash python -m pip install cntryl-fitz ``` -## Quick Start +## Connect ```python from fitz_py import Client, ClientConfig -client = Client( +async with Client( ClientConfig( url="ws://localhost:4190/ws", token_provider=lambda: "", ) -) +) as client: + async with await client.kv.begin("kv://example/app/users") as tx: + await tx.put(b"alice", b"active") + await tx.commit() +``` + +`Client.close()` is permanent and idempotent. Reconnect is enabled by default after the first +successful authentication; an authentication rejection permanently closes the client. Configure +timeouts, bounded concurrency, retry, heartbeat, logging, metrics, and lifecycle events with the +frozen policies on `ClientConfig`. + +## Domains -await client.connect() +- `client.kv`: transactions, scans, range deletes, durability, and mutation subscriptions. +- `client.queue`: delayed enqueue, broker-native long-poll reserve, fenced items, and availability. +- `client.rpc`: streamed calls and wildcard worker registrations with bounded handler dispatch. +- `client.lease`: queued fenced acquisition, query, change subscriptions, and managed renewal. +- `client.notice`: fire-and-forget publish and one-wire/many-consumer subscriptions. +- `client.stream`: append sessions, filtered replay, global cursors/watermarks, and commit events. +- `client.schedule`: delivery modes, total-count pagination, cancel, and routed notifications. -tx = await client.kv().begin("kv://example/app/users") -result = await tx.get(b"key") -await tx.rollback() +Subscriptions are independently closable async iterators: -await client.close() +```python +async with await client.notice.subscribe("notice://example/app/*") as notices: + async for notice in notices: + print(notice.route, notice.body) ``` -The package is `asyncio`-first and does not provide a synchronous wrapper. +Reserve waits are performed by the broker rather than local polling: + +```python +items = await client.queue.reserve("queue://example/work/*", lease=30, wait=10) +for item in items: + try: + await process(item.body) + except Exception: + raise + else: + await item.complete() +``` -## Stream replay +Managed leases renew at one third of their TTL and preserve both renewal and release failures: ```python -from fitz_py import StreamFilterClause, StreamFilterSet - -stream_filter = StreamFilterSet(clauses=[StreamFilterClause(kind="Equals", value="proj.alpha")]) - -records = await client.stream().read( - "stream://example/app/events", - start_offset=0, - limit=100, - stream_filter=stream_filter, -) -page = await client.stream().read_page( - "stream://example/app/events", - start_offset=0, - limit=100, - stream_filter=stream_filter, -) - -# `read()` preserves the compatibility event-only shape. -# `read_page()` exposes filtered markers and cursor metadata. -assert page.cursor.last_resource_offset >= 0 +async with client.lease.hold("lease://example/jobs/leader", ttl=30, wait=10) as lease: + await run_leader(lease.token) ``` -## Parity Goals +## Errors and cancellation -`fitz-py` tracks the Fitz client behavior implemented in `fitz-go` and `fitz-ts`. -The Python SDK now exposes the same seven domains, typed Fitz/domain errors, retryability -helpers via `is_retryable()`, reconnect-aware subscription restoration, and extended -integration/conformance coverage for queue, lease, notice, and schedule lifecycle flows. +All library failures derive from `FitzError`. Transport, connection, timeout, protocol, bounded +queue, stale-handle, and domain failures have stable string codes and structured context. Task +cancellation is preserved. Requests cancelled after transmission leave a FIFO tombstone so a late +reply cannot corrupt the next same-type request. ## Verification -Fast local checks: - ```bash python -m pip install -e ".[dev]" python -m ruff format --check . python -m ruff check . +python -m pyright src python -m pytest tests/unit -``` - -Hot-path microbenchmarks: - -```bash -python artifacts/benchmarks/hotpath.py --iterations 10000 -# or -hatch run bench-hotpath -``` -One-shot verification: +docker compose up -d +python -m pytest tests/integration +python -m pytest tests/conformance +docker compose down --volumes -```bash -hatch run verify -``` - -The local quality bar is formatter first, then lint, then unit tests. - -Broker-backed verification: - -```bash -docker compose -f ../fitz-go/compose.yml up -d -python -m pytest tests/integration -v -CONFORMANCE_TRANSPORT=ws CONFORMANCE_AUTH_MODE=anonymous \ -CONFORMANCE_OUTPUT=artifacts/conformance-results.json \ -python -m pytest tests/conformance -v -docker compose -f ../fitz-go/compose.yml down --volumes -``` - -Package smoke verification: - -```bash +python -m benchmarks.hotpath python -m build -python -m pip install dist/*.whl ``` -The conformance harness writes JSON results to `artifacts/conformance-results.json` by default. -It currently executes 19 scenarios: CS-001..CS-015 mirror the shared cross-language suite, -and CS-016..CS-019 cover fitz-py local lifecycle checks. - -## Project Layout - -- `src/fitz_py`: package code -- `tests/unit`: fast unit coverage -- `tests/integration`: broker-backed integration coverage -- `tests/conformance`: release-gate conformance coverage - -## Canonical Docs - -Canonical client behavior is defined by the server-owned docs in the Fitz repository: - -- [CLIENT_SPEC.md](../fitz/docs/clients/CLIENT_SPEC.md) -- [CLIENT_ACCEPTANCE_CRITERIA.md](../fitz/docs/clients/CLIENT_ACCEPTANCE_CRITERIA.md) -- [CLIENT_IMPLEMENTATION_GUIDE.md](../fitz/docs/clients/CLIENT_IMPLEMENTATION_GUIDE.md) -- [CONNECTION_FLOW.md](../fitz/docs/clients/CONNECTION_FLOW.md) - -## Documentation - -- [`docs/README.md`](docs/README.md) -- [`CLIENT_SPEC.md`](CLIENT_SPEC.md) -- [`CLIENT_ACCEPTANCE_CRITERIA.md`](CLIENT_ACCEPTANCE_CRITERIA.md) +The repository owns its broker Compose stack and a vendored copy of the canonical 17-scenario +cross-language suite. CI runs Python 3.11-3.13, wheel smoke tests, TCP/WebSocket, and +anonymous/JWT broker legs. Canonical behavior remains owned by the Fitz server documentation. diff --git a/benchmarks/__init__.py b/benchmarks/__init__.py new file mode 100644 index 0000000..1321ee4 --- /dev/null +++ b/benchmarks/__init__.py @@ -0,0 +1 @@ +"""Fitz Python benchmark suite.""" diff --git a/benchmarks/hotpath.py b/benchmarks/hotpath.py new file mode 100644 index 0000000..0ea1a05 --- /dev/null +++ b/benchmarks/hotpath.py @@ -0,0 +1,45 @@ +"""Stable Python-native microbenchmarks for codec hot paths.""" + +from __future__ import annotations + +import pyperf + +from fitz_py.protocol.buffer import BufferReader, BufferWriter +from fitz_py.protocol.frame import FrameCodec, FrameParser + +PAYLOAD = b"x" * 256 +FRAME = FrameCodec.encode_frame(104, PAYLOAD) + + +def encode_buffer() -> bytes: + writer = BufferWriter() + writer.write_route("kv://benchmark/area/resource") + writer.write_u64_be(42) + writer.write_u32_be(len(PAYLOAD)) + writer.write_bytes(PAYLOAD) + return writer.build() + + +ENCODED = encode_buffer() + + +def decode_buffer() -> int: + reader = BufferReader(ENCODED) + reader.read_route() + reader.read_u64_be() + return len(reader.read_bytes(reader.read_u32_be())) + + +def parse_frame() -> int: + return len(FrameParser().parse_frames(FRAME)[0].payload) + + +def main() -> None: + runner = pyperf.Runner() + runner.bench_func("buffer_encode_256b", encode_buffer) + runner.bench_func("buffer_decode_256b", decode_buffer) + runner.bench_func("frame_parse_256b", parse_frame) + + +if __name__ == "__main__": + main() diff --git a/compose.yml b/compose.yml new file mode 100644 index 0000000..57feea6 --- /dev/null +++ b/compose.yml @@ -0,0 +1,51 @@ +name: fitz-py + +x-broker-common: &broker-common + image: ghcr.io/cntryl/fitz:latest + restart: unless-stopped + stop_grace_period: 15s + healthcheck: + disable: true + +x-broker-common-env: &broker-common-env + FITZ_STORAGE_MODE: local + FITZ_STORAGE_PATH: /data + RUST_LOG: info,fitz=trace + +services: + fitz-auth: + <<: *broker-common + ports: + - "127.0.0.1:${FITZ_AUTH_HOST_HTTP_PORT:-4090}:4090" + - "127.0.0.1:${FITZ_AUTH_HOST_TCP_PORT:-4091}:4091" + environment: + <<: *broker-common-env + FITZ_AUTH_REQUIRED: "true" + FITZ_ASSUME_LOCAL_LOOPBACK_EDGE: "true" + FITZ_ADMIN_AUTH_MODE: open + FITZ_HTTP_PORT: "4090" + FITZ_TCP_PORT: "4091" + FITZ_JWT_HMAC_SECRET: "${FITZ_JWT_HMAC_SECRET:-test-secret-key}" + FITZ_JWT_AUDIENCE: fitz + FITZ_ROUTE_FAMILY_MAP: dev=1 + FITZ_WS_ALLOWED_ORIGINS: "http://localhost:4090,http://127.0.0.1:4090" + volumes: + - fitz_auth_data:/data + + fitz-anon: + <<: *broker-common + ports: + - "127.0.0.1:${FITZ_ANON_HOST_HTTP_PORT:-4190}:4090" + - "127.0.0.1:${FITZ_ANON_HOST_TCP_PORT:-4191}:4091" + environment: + <<: *broker-common-env + FITZ_AUTH_REQUIRED: "false" + FITZ_ADMIN_AUTH_MODE: open + FITZ_HTTP_PORT: "4090" + FITZ_TCP_PORT: "4091" + volumes: + - fitz_anon_data:/data + +volumes: + fitz_auth_data: + fitz_anon_data: diff --git a/docs/README.md b/docs/README.md index 3abbbc1..08d1aae 100644 --- a/docs/README.md +++ b/docs/README.md @@ -1,24 +1,17 @@ -# fitz-py Documentation - -The Python SDK follows the canonical Fitz client docs in the server repository under [../../fitz/docs/clients](../../fitz/docs/clients). - -## Canonical Docs - -- [CLIENT_SPEC.md](../../fitz/docs/clients/CLIENT_SPEC.md) -- [CLIENT_ACCEPTANCE_CRITERIA.md](../../fitz/docs/clients/CLIENT_ACCEPTANCE_CRITERIA.md) -- [CLIENT_IMPLEMENTATION_GUIDE.md](../../fitz/docs/clients/CLIENT_IMPLEMENTATION_GUIDE.md) -- [CONNECTION_FLOW.md](../../fitz/docs/clients/CONNECTION_FLOW.md) - -## Local Docs - -- [../README.md](../README.md) -- [../CLIENT_SPEC.md](../CLIENT_SPEC.md) -- [../CLIENT_ACCEPTANCE_CRITERIA.md](../CLIENT_ACCEPTANCE_CRITERIA.md) - -Use the canonical docs for protocol behavior and the local docs for Python-specific verification, parity evidence, and setup notes. - -Python verification follows the repo README release gate. The conformance harness currently -executes 19 scenarios: CS-001..CS-015 mirror the shared cross-language suite, and -CS-016..CS-019 cover fitz-py local lifecycle checks. - -The top-level `Client` also supports `async with` for connect/close lifecycle management, matching the existing domain-level context manager pattern. +# Design notes + +The public API is intentionally asyncio-native and compatibility-free. See the root README for +usage. The implementation follows these invariants: + +- one serialized transport writer and FIFO per-message response queues; +- tombstones for timed-out or cancelled requests that were already transmitted; +- bounded request admission, handler dispatch, and subscription buffers; +- generation-bound transactional, queue, lease, stream, and response-writer handles; +- best-effort reconnect restoration with per-registration failure isolation; +- one broker subscription shared by multiple independent local async iterators; +- exact route-bearing wildcard results and notifications; +- cryptographically random 16-byte RPC correlation identifiers; +- no response shim for fire-and-forget Notice publish or RPC call submission. + +The vendored conformance YAML in `tests/conformance` is copied from the canonical Fitz server +suite. Changes to client-wide behavior must update the canonical server documentation first. diff --git a/pyproject.toml b/pyproject.toml index 81c11ff..1d749c0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "cntryl-fitz" -version = "0.1.0" +version = "0.2.0" description = "Async-first Python client for Fitz" readme = "README.md" requires-python = ">=3.11" @@ -38,9 +38,12 @@ Issues = "https://github.com/cntryl/fitz-py/issues" [project.optional-dependencies] dev = [ "build>=1,<2", - "ruff>=0.13,<0.14", + "ruff>=0.13,<1", "pytest>=8,<9", - "pytest-asyncio>=0.24,<1", + "pytest-asyncio>=0.24,<2", + "pytest-cov>=6,<8", + "pyright>=1.1.400,<2", + "pyperf>=2.9,<3", ] [tool.hatch.build.targets.wheel] @@ -58,7 +61,7 @@ test-unit = "pytest tests/unit" test-integration = "pytest tests/integration -v" test-conformance = "pytest tests/conformance -v" test-spec = "pytest tests/conformance -v" -bench-hotpath = "python artifacts/benchmarks/hotpath.py" +bench-hotpath = "python -m benchmarks.hotpath" verify = "ruff format --check . && ruff check . && pytest tests/unit && pytest tests/integration -v && pytest tests/conformance -v" [tool.ruff] @@ -72,3 +75,11 @@ ignore = ["E501"] [tool.ruff.format] quote-style = "double" indent-style = "space" + +[tool.pyright] +include = ["src", "tests/unit"] +pythonVersion = "3.11" +typeCheckingMode = "standard" +reportPrivateUsage = false +reportUnsupportedDunderAll = false +reportMissingImports = false diff --git a/src/fitz_py/__init__.py b/src/fitz_py/__init__.py index fd266a4..7f8d58a 100644 --- a/src/fitz_py/__init__.py +++ b/src/fitz_py/__init__.py @@ -1,205 +1,44 @@ -"""Public exports for the Fitz Python SDK.""" +"""Idiomatic asynchronous Python client for Fitz.""" + +# ruff: noqa: F401, F403 from fitz_py.client import Client -from fitz_py.domains import ( - InboundRpcRequest, - KvClient, - KVDurability, - KvGetResult, - KVMode, - KvPair, - KvScanResult, - KvTransaction, - Lease, - LeaseClient, - LeaseHandler, - LeaseInfo, - LeaseSubscription, - NoticeClient, - NoticeHandler, - NoticeMessage, - NoticeSubscription, - QueueAvailabilityHandler, - QueueClient, - QueueItem, - QueueSubscription, - ResponseFrame, - ResponseWriter, - RpcClient, - RpcHandler, - RpcSubscription, - ScheduleClient, - ScheduleEntry, - ScheduleHandler, - ScheduleNotification, - ScheduleSubscription, - StreamClient, - StreamCommitMode, - StreamCommitNotification, - StreamFilterClause, - StreamFilteredReason, - StreamFilterSet, - StreamHandler, - StreamMetadata, - StreamReadCursor, - StreamReadItem, - StreamReadItemKind, - StreamReadPage, - StreamRecord, - StreamSession, - StreamSubscription, -) +from fitz_py.domains import * # noqa: F403 from fitz_py.errors import ( AuthenticationError, CodecError, - ConnectionError, - ErrKvConflictingWrite, - ErrKvKeyNotFound, - ErrKvLeaseExpired, - ErrKvOperationNotAllowed, - ErrKvTransactionAborted, - ErrLeaseHeld, - ErrLeaseInvalidToken, - ErrLeaseNotFound, - ErrNoticeGeneral, - ErrQueueFull, - ErrQueueInvalidDelay, - ErrQueueInvalidToken, - ErrQueueMessageNotFound, - ErrQueueNotFound, - ErrRpcHandlerError, - ErrRpcHandlerNotFound, - ErrRpcInvalidRequest, - ErrRpcTimeout, - ErrScheduleInvalidCron, - ErrScheduleInvalidDelay, - ErrScheduleInvalidTimestamp, - ErrScheduleNotFound, - ErrScheduleTaskNotFound, - ErrStreamExpectedOffsetMismatch, - ErrStreamFull, - ErrStreamInvalidOffset, - ErrStreamNotFound, - ErrStreamOffsetOutOfRange, - ErrStreamSessionClosed, - ErrStreamSessionNotFound, + DomainError, + FitzConnectionError, FitzError, + FitzTimeoutError, + FitzTransportError, KvError, LeaseError, + LeaseLifecycleError, + LeaseLostError, NoticeError, ProtocolError, QueueError, + ReconnectRestoreError, + RequestQueueFullError, RpcError, ScheduleError, + StaleHandleError, StreamError, - TimeoutError, - TransportError, + SubscriptionBackpressureError, is_retryable, ) from fitz_py.types import ( ClientConfig, + ConcurrencyLimits, ConnectionState, - ReconnectOptions, + HeartbeatPolicy, + LifecycleEvent, + Observability, + ReconnectPolicy, + RetryPolicy, TokenProvider, TransportType, ) -__all__ = [ - "AuthenticationError", - "Client", - "ClientConfig", - "CodecError", - "ConnectionError", - "ConnectionState", - "ErrKvConflictingWrite", - "ErrKvKeyNotFound", - "ErrKvLeaseExpired", - "ErrKvOperationNotAllowed", - "ErrKvTransactionAborted", - "ErrLeaseHeld", - "ErrLeaseInvalidToken", - "ErrLeaseNotFound", - "ErrNoticeGeneral", - "ErrQueueFull", - "ErrQueueInvalidDelay", - "ErrQueueInvalidToken", - "ErrQueueMessageNotFound", - "ErrQueueNotFound", - "ErrRpcHandlerError", - "ErrRpcHandlerNotFound", - "ErrRpcInvalidRequest", - "ErrRpcTimeout", - "ErrScheduleInvalidCron", - "ErrScheduleInvalidDelay", - "ErrScheduleInvalidTimestamp", - "ErrScheduleNotFound", - "ErrScheduleTaskNotFound", - "ErrStreamExpectedOffsetMismatch", - "ErrStreamFull", - "ErrStreamInvalidOffset", - "ErrStreamNotFound", - "ErrStreamOffsetOutOfRange", - "ErrStreamSessionClosed", - "ErrStreamSessionNotFound", - "FitzError", - "InboundRpcRequest", - "KVDurability", - "KVMode", - "KvClient", - "KvGetResult", - "KvError", - "KvPair", - "KvScanResult", - "KvTransaction", - "Lease", - "LeaseClient", - "LeaseError", - "LeaseHandler", - "LeaseInfo", - "LeaseSubscription", - "NoticeClient", - "NoticeError", - "NoticeHandler", - "NoticeMessage", - "NoticeSubscription", - "ProtocolError", - "QueueAvailabilityHandler", - "QueueClient", - "QueueError", - "QueueItem", - "QueueSubscription", - "ReconnectOptions", - "ResponseFrame", - "ResponseWriter", - "RpcClient", - "RpcError", - "RpcHandler", - "RpcSubscription", - "ScheduleClient", - "ScheduleEntry", - "ScheduleError", - "ScheduleHandler", - "ScheduleNotification", - "ScheduleSubscription", - "StreamClient", - "StreamFilterClause", - "StreamFilterSet", - "StreamCommitMode", - "StreamCommitNotification", - "StreamError", - "StreamFilteredReason", - "StreamHandler", - "StreamMetadata", - "StreamReadCursor", - "StreamReadItem", - "StreamReadItemKind", - "StreamReadPage", - "StreamRecord", - "StreamSession", - "StreamSubscription", - "TimeoutError", - "TokenProvider", - "TransportError", - "TransportType", - "is_retryable", -] +__all__ = [name for name in globals() if not name.startswith("_")] diff --git a/src/fitz_py/_runtime.py b/src/fitz_py/_runtime.py new file mode 100644 index 0000000..d6a2884 --- /dev/null +++ b/src/fitz_py/_runtime.py @@ -0,0 +1,190 @@ +"""Bounded asyncio runtime primitives shared by the client and domains.""" + +from __future__ import annotations + +import asyncio +import contextlib +from collections import deque +from collections.abc import AsyncIterator, Awaitable, Callable +from typing import Generic, TypeVar + +from fitz_py.errors import ( + FitzConnectionError, + RequestQueueFullError, + SubscriptionBackpressureError, +) + +T = TypeVar("T") +_END = object() + + +class RequestGate: + def __init__(self, maximum: int, queue_size: int) -> None: + self._maximum = maximum + self._queue_size = queue_size + self._active = 0 + self._closed = False + self._waiters: deque[asyncio.Future[None]] = deque() + + async def acquire(self) -> Callable[[], None]: + if self._closed: + raise FitzConnectionError("Connection is closed") + if self._active < self._maximum: + self._active += 1 + return self._release + if len(self._waiters) >= self._queue_size: + raise RequestQueueFullError() + + waiter = asyncio.get_running_loop().create_future() + self._waiters.append(waiter) + try: + await waiter + except BaseException: + with contextlib.suppress(ValueError): + self._waiters.remove(waiter) + raise + if self._closed: + raise FitzConnectionError("Connection is closed") + self._active += 1 + return self._release + + def _release(self) -> None: + if self._active > 0: + self._active -= 1 + while self._waiters: + waiter = self._waiters.popleft() + if not waiter.done(): + waiter.set_result(None) + break + + def close(self) -> None: + if self._closed: + return + self._closed = True + error = FitzConnectionError("Connection is closed") + for waiter in self._waiters: + if not waiter.done(): + waiter.set_exception(error) + self._waiters.clear() + + +class AsyncDispatcher: + def __init__( + self, + maximum: int, + queue_size: int, + timeout: float, + on_error: Callable[[BaseException], None], + ) -> None: + self._maximum = maximum + self._queue_size = queue_size + self._timeout = timeout + self._on_error = on_error + self._active: set[asyncio.Task[None]] = set() + self._queued: deque[Callable[[], Awaitable[None]]] = deque() + self._closed = False + + def dispatch(self, work: Callable[[], Awaitable[None]]) -> bool: + if self._closed: + return False + if len(self._active) < self._maximum: + self._start(work) + return True + if len(self._queued) >= self._queue_size: + return False + self._queued.append(work) + return True + + def _start(self, work: Callable[[], Awaitable[None]]) -> None: + task = asyncio.create_task(self._run(work)) + self._active.add(task) + task.add_done_callback(self._completed) + + async def _run(self, work: Callable[[], Awaitable[None]]) -> None: + try: + async with asyncio.timeout(self._timeout): + await work() + except BaseException as exc: + if not isinstance(exc, asyncio.CancelledError): + self._on_error(exc) + + def _completed(self, task: asyncio.Task[None]) -> None: + self._active.discard(task) + if not self._closed and self._queued: + self._start(self._queued.popleft()) + + async def close(self) -> None: + self._closed = True + self._queued.clear() + tasks = tuple(self._active) + for task in tasks: + task.cancel() + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + + +class AsyncSubscription(AsyncIterator[T], Generic[T]): + """A bounded, independently closable local subscription consumer.""" + + def __init__( + self, + registration: str, + capacity: int, + close_wire: Callable[[], Awaitable[None]], + ) -> None: + self.registration = registration + self._queue: asyncio.Queue[T | BaseException | object] = asyncio.Queue(capacity) + self._close_wire = close_wire + self._closed = False + + def __aiter__(self) -> AsyncSubscription[T]: + return self + + async def __anext__(self) -> T: + item = await self._queue.get() + if item is _END: + raise StopAsyncIteration + if isinstance(item, BaseException): + self._closed = True + raise item + return item # type: ignore[return-value] + + async def __aenter__(self) -> AsyncSubscription[T]: + return self + + async def __aexit__(self, *_args: object) -> None: + await self.aclose() + + def push(self, item: T) -> bool: + if self._closed: + return False + try: + self._queue.put_nowait(item) + return True + except asyncio.QueueFull: + self.fail(SubscriptionBackpressureError()) + return False + + def fail(self, error: BaseException) -> None: + if self._closed: + return + self._closed = True + while not self._queue.empty(): + with contextlib.suppress(asyncio.QueueEmpty): + self._queue.get_nowait() + self._queue.put_nowait(error) + + async def unsubscribe(self) -> None: + await self.aclose() + + async def aclose(self) -> None: + if self._closed: + return + self._closed = True + await self._close_wire() + with contextlib.suppress(asyncio.QueueFull): + self._queue.put_nowait(_END) + + +async def sleep_backoff(delay: float) -> None: + await asyncio.sleep(delay) diff --git a/src/fitz_py/client.py b/src/fitz_py/client.py index 3c4d034..85e652d 100644 --- a/src/fitz_py/client.py +++ b/src/fitz_py/client.py @@ -1,7 +1,10 @@ -"""High-level Fitz SDK client and domain accessors.""" +"""High-level asynchronous Fitz client facade.""" from __future__ import annotations +import asyncio +import random + from fitz_py.connection import Connection from fitz_py.domains.kv import KvClient from fitz_py.domains.lease import LeaseClient @@ -10,127 +13,112 @@ from fitz_py.domains.rpc import RpcClient from fitz_py.domains.schedule import ScheduleClient from fitz_py.domains.stream import StreamClient -from fitz_py.errors import ConnectionError +from fitz_py.errors import AuthenticationError, FitzConnectionError, FitzTransportError from fitz_py.transport.factory import create_transport -from fitz_py.types import ClientConfig, ConnectionState, TokenProvider +from fitz_py.types import ClientConfig, ConnectionState, TransportType class Client: - """Top-level Fitz client that manages a connection and domain clients.""" - def __init__(self, config: ClientConfig) -> None: - if not config.url: - raise ValueError("url is required") - - self._config = ClientConfig( - url=config.url, - token_provider=config.token_provider, - timeout_ms=config.timeout_ms, - transport=config.transport, - reconnect=config.reconnect, - max_frame_size=config.max_frame_size, - auth_settle_delay_ms=config.auth_settle_delay_ms, - max_in_flight_requests=config.max_in_flight_requests, - ) - self._connection: Connection | None = None - self._kv_client: KvClient | None = None - self._queue_client: QueueClient | None = None - self._rpc_client: RpcClient | None = None - self._lease_client: LeaseClient | None = None - self._notice_client: NoticeClient | None = None - self._stream_client: StreamClient | None = None - self._schedule_client: ScheduleClient | None = None - - async def connect(self) -> None: - if self._connection is not None and self._connection.is_connected(): - return - - token_provider = self._resolve_token_provider() - reconnect = self._config.reconnect + self.config = config + self._closed = False self._connection = Connection( lambda: create_transport( - self._config.url, - self._config.transport, - timeout_ms=self._config.timeout_ms, - max_frame_size=self._config.max_frame_size, + config.url, + TransportType(config.transport), + timeout_ms=int(config.request_timeout * 1000), + max_frame_size=config.max_frame_size, + websocket_headers=dict(config.websocket_headers), ), - token_provider, - timeout_ms=self._config.timeout_ms, - auth_settle_delay_ms=self._config.auth_settle_delay_ms, - reconnect_enabled=reconnect.enabled if reconnect else False, - reconnect_max_attempts=reconnect.max_attempts if reconnect else float("inf"), - reconnect_backoff_ms=reconnect.backoff_ms if reconnect else 250, - reconnect_max_backoff_ms=reconnect.max_backoff_ms if reconnect else 5000, - max_in_flight_requests=self._config.max_in_flight_requests, + config, ) - await self._connection.connect() - - async def __aenter__(self) -> "Client": + self._kv = KvClient(self._connection) + self._queue = QueueClient(self._connection) + self._rpc = RpcClient(self._connection) + self._lease = LeaseClient(self._connection) + self._notice = NoticeClient(self._connection) + self._stream = StreamClient(self._connection) + self._schedule = ScheduleClient(self._connection) + + async def __aenter__(self) -> Client: await self.connect() return self - async def __aexit__(self, exc_type, exc, tb) -> None: + async def __aexit__(self, *_args: object) -> None: await self.close() + async def connect(self) -> None: + if self._closed: + raise FitzConnectionError("Client is closed") + await self._connection.connect() + + async def connect_when_ready( + self, + *, + timeout: float | None = None, + backoff: float = 0.25, + max_backoff: float = 2.0, + ) -> None: + if self._closed: + raise FitzConnectionError("Client is closed") + loop = asyncio.get_running_loop() + deadline = None if timeout is None else loop.time() + timeout + delay = backoff + while True: + try: + await self.connect() + return + except AuthenticationError: + raise + except (FitzTransportError, FitzConnectionError): + if deadline is not None and loop.time() >= deadline: + raise + sleep_for = delay * random.uniform(0.8, 1.2) + if deadline is not None: + sleep_for = min(sleep_for, max(0, deadline - loop.time())) + await asyncio.sleep(sleep_for) + delay = min(delay * 2, max_backoff) + async def close(self) -> None: - if self._connection is not None: - await self._connection.close() - self._connection = None - self._kv_client = None - self._queue_client = None - self._rpc_client = None - self._lease_client = None - self._notice_client = None - self._stream_client = None - self._schedule_client = None + self._closed = True + await self._connection.close() @property def state(self) -> ConnectionState: - return ( - self._connection.get_state() - if self._connection is not None - else ConnectionState.DISCONNECTED - ) + return self._connection.get_state() + @property + def is_connected(self) -> bool: + return self._connection.is_connected() + + @property + def url(self) -> str: + return self.config.url + + @property def kv(self) -> KvClient: - if self._kv_client is None: - self._kv_client = KvClient(self._ensure_connection()) - return self._kv_client + return self._kv + @property def queue(self) -> QueueClient: - if self._queue_client is None: - self._queue_client = QueueClient(self._ensure_connection()) - return self._queue_client + return self._queue + @property def rpc(self) -> RpcClient: - if self._rpc_client is None: - self._rpc_client = RpcClient(self._ensure_connection()) - return self._rpc_client + return self._rpc + @property def lease(self) -> LeaseClient: - if self._lease_client is None: - self._lease_client = LeaseClient(self._ensure_connection()) - return self._lease_client + return self._lease + @property def notice(self) -> NoticeClient: - if self._notice_client is None: - self._notice_client = NoticeClient(self._ensure_connection()) - return self._notice_client + return self._notice + @property def stream(self) -> StreamClient: - if self._stream_client is None: - self._stream_client = StreamClient(self._ensure_connection()) - return self._stream_client + return self._stream + @property def schedule(self) -> ScheduleClient: - if self._schedule_client is None: - self._schedule_client = ScheduleClient(self._ensure_connection()) - return self._schedule_client - - def _resolve_token_provider(self) -> TokenProvider: - return self._config.token_provider or (lambda: "") - - def _ensure_connection(self) -> Connection: - if self._connection is None: - raise ConnectionError("Not connected to Fitz server. Call connect() first.") - return self._connection + return self._schedule diff --git a/src/fitz_py/connection.py b/src/fitz_py/connection.py index e7321a5..6ec2fc3 100644 --- a/src/fitz_py/connection.py +++ b/src/fitz_py/connection.py @@ -1,170 +1,142 @@ -"""Connection lifecycle, authentication, request/response, and reconnect logic.""" +"""Connection lifecycle, multiplexing, reconnect, retry, and telemetry.""" from __future__ import annotations import asyncio import contextlib +import inspect +import random +import time from collections.abc import Awaitable, Callable - -from fitz_py.errors import AuthenticationError, ConnectionError, TransportError +from typing import Any, TypeVar + +from fitz_py._runtime import AsyncDispatcher, RequestGate +from fitz_py.errors import ( + AuthenticationError, + FitzConnectionError, + FitzTransportError, + ReconnectRestoreError, + is_retryable, +) from fitz_py.multiplexer import Multiplexer from fitz_py.protocol.frame import FrameCodec, FrameParser from fitz_py.protocol.messages import MSG_CONNECT from fitz_py.transport.base import Transport -from fitz_py.types import ConnectionState, TokenProvider +from fitz_py.types import ClientConfig, ConnectionState, LifecycleEvent TransportFactory = Callable[[], Transport] ReconnectListener = Callable[[], None | Awaitable[None]] DisconnectListener = Callable[[], None | Awaitable[None]] - - -class AdmissionGate: - def __init__(self, max_in_flight_requests: int) -> None: - self._max_in_flight_requests = max(1, max_in_flight_requests) - self._active = 0 - self._closed = False - self._condition = asyncio.Condition() - - async def acquire(self) -> None: - async with self._condition: - while self._active >= self._max_in_flight_requests and not self._closed: - await self._condition.wait() - - if self._closed: - raise ConnectionError("Connection closed") - - self._active += 1 - - async def release(self) -> None: - async with self._condition: - if self._active > 0: - self._active -= 1 - self._condition.notify() - - async def close(self) -> None: - async with self._condition: - self._closed = True - self._condition.notify_all() - - -async def _sleep_ms(delay_ms: int) -> None: - await asyncio.sleep(delay_ms / 1000) - - -async def _resolve_token(token_provider: TokenProvider) -> str: - token = token_provider() - if asyncio.iscoroutine(token) or isinstance(token, Awaitable): - return await token - return token +T = TypeVar("T") class Connection: - """Owns the transport, multiplexer, and Fitz connection state machine.""" - - def __init__( - self, - transport_factory: TransportFactory, - token_provider: TokenProvider, - *, - timeout_ms: int = 30000, - auth_settle_delay_ms: int = 500, - reconnect_enabled: bool = False, - reconnect_max_attempts: int | float = float("inf"), - reconnect_backoff_ms: int = 250, - reconnect_max_backoff_ms: int = 5000, - max_in_flight_requests: int = 256, - ) -> None: + def __init__(self, transport_factory: TransportFactory, config: ClientConfig) -> None: self._transport_factory = transport_factory - self._token_provider = token_provider - self._timeout_ms = timeout_ms - self._auth_settle_delay_ms = auth_settle_delay_ms - self._reconnect_enabled = reconnect_enabled - self._reconnect_max_attempts = reconnect_max_attempts - self._reconnect_backoff_ms = reconnect_backoff_ms - self._reconnect_max_backoff_ms = reconnect_max_backoff_ms - self._max_in_flight_requests = max(1, max_in_flight_requests) - + self._config = config self._transport: Transport | None = None self._state = ConnectionState.DISCONNECTED + self._generation = 0 self._multiplexer = Multiplexer() - self._frame_parser = FrameParser() + self._parser = FrameParser() self._receive_task: asyncio.Task[None] | None = None - self._reconnect_listeners: set[ReconnectListener] = set() - self._disconnect_listeners: set[DisconnectListener] = set() - self._auth_future: asyncio.Future[None] | None = None - self._close_requested = False - self._receive_loop_abort = False + self._heartbeat_task: asyncio.Task[None] | None = None self._reconnect_task: asyncio.Task[None] | None = None - self._admission_gate = AdmissionGate(self._max_in_flight_requests) + self._loss_task: asyncio.Task[None] | None = None + self._connect_task: asyncio.Task[None] | None = None + self._close_task: asyncio.Task[None] | None = None + self._connect_lock = asyncio.Lock() + self._write_lock = asyncio.Lock() + self._closed = False + self._ever_authenticated = False + self._restoring = False + self._last_activity = time.monotonic() + self._auth_error: asyncio.Future[None] | None = None + self._reconnect_listeners: dict[tuple[str, str], ReconnectListener] = {} + self._disconnect_listeners: set[DisconnectListener] = set() + self._gate = self._new_gate() + limits = config.limits + self._dispatcher = AsyncDispatcher( + limits.async_handler_concurrency, + limits.async_handler_queue_size, + limits.async_handler_timeout, + self._handler_error, + ) + + @property + def generation(self) -> int: + return self._generation + + @property + def config(self) -> ClientConfig: + return self._config async def connect(self) -> None: - self._close_requested = False - await self._open_and_authenticate(False) + if self._closed: + raise FitzConnectionError("Client is closed") + if self.is_connected(): + return + async with self._connect_lock: + if self.is_connected(): + return + if self._connect_task is None or self._connect_task.done(): + self._connect_task = asyncio.create_task(self._open_and_authenticate(False)) + task = self._connect_task + await asyncio.shield(task) async def close(self) -> None: - if self._state is ConnectionState.CLOSED and self._transport is None: + if self._close_task is not None: + await asyncio.shield(self._close_task) return - self._close_requested = True - self._receive_loop_abort = True - self._set_state(ConnectionState.CLOSED) - await asyncio.shield(self._admission_gate.close()) - if self._auth_future is not None and not self._auth_future.done(): - self._auth_future.set_exception(ConnectionError("Connection closed")) - self._auth_future = None - self._multiplexer.set_disconnected() - await self._notify_disconnect_listeners() + self._close_task = asyncio.create_task(self._close()) + await asyncio.shield(self._close_task) - receive_task = self._receive_task - self._receive_task = None - if receive_task is not None: - with contextlib.suppress(Exception): - await asyncio.wait_for(receive_task, timeout=1) - - transport = self._transport - self._transport = None + async def _close(self) -> None: + if self._closed: + return + self._closed = True + self._set_state(ConnectionState.CLOSED, "closed") + self._gate.close() + self._multiplexer.set_disconnected() + await self._notify_disconnect() + current = asyncio.current_task() + for task in (self._reconnect_task, self._heartbeat_task, self._receive_task): + if task is not None and task is not current: + task.cancel() + transport, self._transport = self._transport, None if transport is not None: - await transport.close() + with contextlib.suppress(Exception): + await transport.close() + await self._dispatcher.close() async def request(self, message_type: int, payload: bytes) -> bytes: self._ensure_authenticated() - await self._admission_gate.acquire() + release = await self._gate.acquire() + started = time.monotonic() try: - return await self._request_without_admission(message_type, payload) - except Exception as exc: - self._handle_possible_transport_failure(exc) - raise + return await self._request_without_gate(message_type, payload) finally: - await asyncio.shield(self._admission_gate.release()) - - async def request_without_admission(self, message_type: int, payload: bytes) -> bytes: - self._ensure_authenticated() - return await self._request_without_admission(message_type, payload) + release() + self._record_latency(message_type, started) async def send(self, message_type: int, payload: bytes) -> None: self._ensure_authenticated() - await self._admission_gate.acquire() + release = await self._gate.acquire() try: - transport = self._ensure_transport() - frame = FrameCodec.encode_frame(message_type, payload) - await transport.send(frame) - except Exception as exc: - self._handle_possible_transport_failure(exc) - raise + await self._send_frame(FrameCodec.encode_frame(message_type, payload)) finally: - await asyncio.shield(self._admission_gate.release()) + release() async def send_fire_and_forget(self, message_type: int, payload: bytes) -> None: await self.send(message_type, payload) - async def reserve_admission(self) -> None: + async def request_without_admission(self, message_type: int, payload: bytes) -> bytes: self._ensure_authenticated() - await self._admission_gate.acquire() - - async def release_admission(self) -> None: - await asyncio.shield(self._admission_gate.release()) + return await self._request_without_gate(message_type, payload) - def release_admission_nowait(self) -> None: - asyncio.create_task(self.release_admission()) + async def reserve_admission(self) -> Callable[[], None]: + self._ensure_authenticated() + return await self._gate.acquire() def register_notification_handler( self, message_type: int, handler: Callable[[bytes], None] @@ -174,11 +146,31 @@ def register_notification_handler( def unregister_notification_handler(self, message_type: int) -> None: self._multiplexer.unregister_notification_handler(message_type) - def on_reconnect(self, listener: ReconnectListener) -> Callable[[], None]: - self._reconnect_listeners.add(listener) + def register_push_classifier( + self, message_type: int, classifier: Callable[[bytes], bool] + ) -> None: + self._multiplexer.register_push_classifier(message_type, classifier) + + def dispatch_async(self, work: Callable[[], Awaitable[None]]) -> bool: + accepted = self._dispatcher.dispatch(work) + if not accepted: + self._log("warning", "fitz.handlers.saturated") + return accepted + + def on_reconnect( + self, + listener: ReconnectListener, + *, + domain: str = "unknown", + registration: str = "unknown", + ) -> Callable[[], None]: + key = (domain, registration) + if key == ("unknown", "unknown"): + key = ("unknown", str(id(listener))) + self._reconnect_listeners[key] = listener def unregister() -> None: - self._reconnect_listeners.discard(listener) + self._reconnect_listeners.pop(key, None) return unregister @@ -190,6 +182,14 @@ def unregister() -> None: return unregister + def report_restore_failure(self, domain: str, registration: str, error: BaseException) -> None: + self._emit( + "reconnect_restore_failed", + domain=domain, + registration=registration, + error=error, + ) + def get_multiplexer(self) -> Multiplexer: return self._multiplexer @@ -200,172 +200,282 @@ def is_connected(self) -> bool: return self._state is ConnectionState.AUTHENTICATED def get_url(self) -> str: - return self._ensure_transport().get_url() + return self._config.url - async def _open_and_authenticate(self, is_reconnect: bool) -> None: - self._receive_loop_abort = False - self._admission_gate = AdmissionGate(self._max_in_flight_requests) - self._transport = self._transport_factory() + async def run_with_retry( + self, + operation: Callable[[], Awaitable[T]], + *, + replay_safe: bool, + ) -> T: + policy = self._config.retry + attempts = policy.max_attempts if policy.enabled and replay_safe else 1 + delay = policy.backoff + for attempt in range(1, attempts + 1): + try: + return await operation() + except BaseException as exc: + if attempt >= attempts or not is_retryable(exc): + raise + await asyncio.sleep(delay * random.uniform(0.8, 1.2)) + delay = min(delay * 2, policy.max_backoff) + raise AssertionError("retry loop exhausted") + + async def _open_and_authenticate(self, reconnect: bool) -> None: + self._parser = FrameParser() + self._gate = self._new_gate() + transport = self._transport_factory() + self._transport = transport self._set_state( - ConnectionState.RECONNECTING if is_reconnect else ConnectionState.CONNECTING + ConnectionState.RECONNECTING if reconnect else ConnectionState.CONNECTING, + "reconnecting" if reconnect else "connecting", ) - await self._transport.connect() - self._receive_task = asyncio.create_task(self._receive_loop()) - - self._set_state(ConnectionState.CONNECTED) - self._set_state(ConnectionState.AUTHENTICATING) - loop = asyncio.get_running_loop() - self._auth_future = loop.create_future() - try: + async with asyncio.timeout(self._config.request_timeout): + await transport.connect() + self._set_state(ConnectionState.CONNECTED, "connected") + self._set_state(ConnectionState.AUTHENTICATING, "authenticating") + self._auth_error = asyncio.get_running_loop().create_future() + self._receive_task = asyncio.create_task(self._receive_loop(transport)) await self._send_connect() - await asyncio.wait_for( - asyncio.shield(self._auth_settle()), timeout=self._timeout_ms / 1000 - ) - if self._auth_future is not None and not self._auth_future.done(): - self._auth_future.set_result(None) - self._auth_future = None - self._set_state(ConnectionState.AUTHENTICATED) + try: + async with asyncio.timeout(self._config.auth_settle_timeout): + await asyncio.shield(self._auth_error) + except TimeoutError: + pass + if self._auth_error is not None and self._auth_error.done(): + self._auth_error.result() + self._auth_error = None + self._generation += 1 self._multiplexer.set_connected() - if is_reconnect: - await self._restore_reconnect_state() - except Exception: - self._auth_future = None + if reconnect: + self._restoring = True + try: + await self._restore_registrations() + finally: + self._restoring = False + self._ever_authenticated = True + self._set_state(ConnectionState.AUTHENTICATED, "authenticated") + self._start_heartbeat() + except BaseException as exc: self._multiplexer.set_disconnected() - transport = self._transport - self._transport = None - if transport is not None: - with contextlib.suppress(Exception): - await transport.close() - self._set_state(ConnectionState.DISCONNECTED) + if self._transport is transport: + self._transport = None + with contextlib.suppress(Exception): + await transport.close() + if isinstance(exc, AuthenticationError): + self._closed = True + self._gate.close() + await self._dispatcher.close() + self._set_state(ConnectionState.CLOSED, "auth_rejected", error=exc) + elif not self._closed: + self._set_state(ConnectionState.DISCONNECTED, "connect_failed") raise - async def _auth_settle(self) -> None: - auth_future = self._auth_future - if auth_future is None: - return - sleep_task = asyncio.create_task(_sleep_ms(self._auth_settle_delay_ms)) - done, pending = await asyncio.wait( - {auth_future, sleep_task}, - return_when=asyncio.FIRST_COMPLETED, - ) - for task in pending: - task.cancel() - for task in done: - await task - async def _send_connect(self) -> None: - token = await _resolve_token(self._token_provider) - frame = FrameCodec.encode_frame(MSG_CONNECT, token.encode()) - await self._ensure_transport().send(frame) + provider = self._config.token_provider + provided = "" if provider is None else provider() + token: str | bytes + if inspect.isawaitable(provided): + token = await provided + else: + token = provided + payload = token.encode() if isinstance(token, str) else bytes(token) + await self._send_frame(FrameCodec.encode_frame(MSG_CONNECT, payload)) + + async def _request_without_gate(self, message_type: int, payload: bytes) -> bytes: + frame = FrameCodec.encode_frame(message_type, payload) + try: + return await self._multiplexer.request( + message_type, frame, self._send_frame, self._config.request_timeout + ) + except (FitzTransportError, FitzConnectionError) as exc: + self._schedule_connection_loss(exc) + raise - async def _receive_loop(self) -> None: - while not self._receive_loop_abort and not self._close_requested: - try: - transport = self._ensure_transport() + async def _send_frame(self, frame: bytes) -> None: + transport = self._transport + if transport is None: + raise FitzConnectionError("No active transport") + async with self._write_lock: + await transport.send(frame) + self._last_activity = time.monotonic() + + async def _receive_loop(self, transport: Transport) -> None: + try: + while not self._closed and self._transport is transport: data = await transport.receive() - frames = self._frame_parser.parse_frames(data) - for frame in frames: + self._last_activity = time.monotonic() + for frame in self._parser.parse_frames(data): self._multiplexer.dispatch(frame.message_type, frame.payload) - except Exception as exc: - if self._receive_loop_abort or self._close_requested: - return - await self._handle_connection_loss(exc) + except asyncio.CancelledError: + return + except BaseException as exc: + if self._state is ConnectionState.AUTHENTICATING and self._auth_error is not None: + if not self._auth_error.done(): + self._auth_error.set_exception(AuthenticationError(str(exc) or "auth rejected")) return - - async def _handle_connection_loss(self, exc: Exception) -> None: - self._multiplexer.set_disconnected() - await asyncio.shield(self._admission_gate.close()) - - if ( - self._state is ConnectionState.AUTHENTICATING - and self._auth_future is not None - and not self._auth_future.done() - ): - self._auth_future.set_exception( - AuthenticationError(self._describe_connection_loss(exc)) - ) - - if self._close_requested: - self._set_state(ConnectionState.CLOSED) + if not self._closed and self._transport is transport: + await self._connection_lost(exc) + + async def _connection_lost(self, cause: BaseException) -> None: + if self._state in { + ConnectionState.DISCONNECTED, + ConnectionState.RECONNECTING, + ConnectionState.CLOSED, + }: return - - self._set_state(ConnectionState.DISCONNECTED) - await self._notify_disconnect_listeners() - if not self._reconnect_enabled: + self._multiplexer.set_disconnected() + self._gate.close() + self._set_state(ConnectionState.DISCONNECTED, "disconnected", error=cause) + transport, self._transport = self._transport, None + if transport is not None: + with contextlib.suppress(Exception): + await transport.close() + await self._notify_disconnect() + if not self._ever_authenticated or not self._config.reconnect.enabled or self._closed: return - if self._reconnect_task is None or self._reconnect_task.done(): self._reconnect_task = asyncio.create_task(self._reconnect_loop()) - await self._reconnect_task + + def _schedule_connection_loss(self, cause: BaseException) -> None: + if self._loss_task is None or self._loss_task.done(): + self._loss_task = asyncio.create_task(self._connection_lost(cause)) async def _reconnect_loop(self) -> None: - attempts = 0 - delay_ms = self._reconnect_backoff_ms - while not self._close_requested and attempts < self._reconnect_max_attempts: - attempts += 1 - self._set_state(ConnectionState.RECONNECTING) - await _sleep_ms(delay_ms) - if self._close_requested: - self._set_state(ConnectionState.CLOSED) - return + policy = self._config.reconnect + delay = policy.backoff + attempt = 0 + while not self._closed and (policy.max_attempts is None or attempt < policy.max_attempts): + attempt += 1 + self._set_state(ConnectionState.RECONNECTING, "reconnect_attempt", attempt=attempt) + await asyncio.sleep(delay * random.uniform(0.8, 1.2)) try: await self._open_and_authenticate(True) return - except Exception: - if self._close_requested: - self._set_state(ConnectionState.CLOSED) - return - delay_ms = min(delay_ms * 2, self._reconnect_max_backoff_ms) - if self._close_requested: - self._set_state(ConnectionState.CLOSED) + except AuthenticationError: + self._closed = True + self._set_state(ConnectionState.CLOSED, "auth_rejected") + return + except BaseException as exc: + self._emit("reconnect_failed", attempt=attempt, error=exc) + delay = min(delay * 2, policy.max_backoff) + if not self._closed: + self._set_state(ConnectionState.DISCONNECTED, "reconnect_exhausted") + + async def _restore_registrations(self) -> None: + for (domain, registration), listener in list(self._reconnect_listeners.items()): + try: + result = listener() + if inspect.isawaitable(result): + await result + except BaseException as exc: + error = ReconnectRestoreError(domain, registration, exc) + self._emit( + "reconnect_restore_failed", + domain=domain, + registration=registration, + error=error, + ) + + async def _notify_disconnect(self) -> None: + for listener in list(self._disconnect_listeners): + try: + result = listener() + if inspect.isawaitable(result): + await result + except BaseException as exc: + self._handler_error(exc) + + def _start_heartbeat(self) -> None: + if not self._config.heartbeat.enabled: return - self._set_state(ConnectionState.DISCONNECTED) + if self._heartbeat_task is not None: + self._heartbeat_task.cancel() + self._heartbeat_task = asyncio.create_task(self._heartbeat_loop()) - async def _restore_reconnect_state(self) -> None: - for listener in list(self._reconnect_listeners): - result = listener() - if asyncio.iscoroutine(result): - await result - - async def _notify_disconnect_listeners(self) -> None: - for listener in list(self._disconnect_listeners): - result = listener() - if asyncio.iscoroutine(result): - await result + async def _heartbeat_loop(self) -> None: + policy = self._config.heartbeat + try: + while not self._closed and self.is_connected(): + await asyncio.sleep(policy.interval) + if time.monotonic() - self._last_activity < policy.interval: + continue + transport = self._transport + if transport is None: + return + await transport.heartbeat(policy.timeout) + self._last_activity = time.monotonic() + except asyncio.CancelledError: + return + except BaseException as exc: + await self._connection_lost(exc) - def _ensure_transport(self) -> Transport: - if self._transport is None: - raise ConnectionError("No active transport") - return self._transport + def _new_gate(self) -> RequestGate: + limits = self._config.limits + return RequestGate(limits.max_in_flight, limits.request_queue_size) def _ensure_authenticated(self) -> None: - if self._close_requested or self._state is not ConnectionState.AUTHENTICATED: - raise ConnectionError(f"Cannot use connection while state is {self._state.value}") + if self._closed or ( + self._state is not ConnectionState.AUTHENTICATED and not self._restoring + ): + raise FitzConnectionError( + f"Cannot use connection while state is {self._state.value}", + {"state": self._state.value}, + ) - def _set_state(self, state: ConnectionState) -> None: + def _set_state( + self, + state: ConnectionState, + event: str, + *, + attempt: int | None = None, + error: BaseException | None = None, + ) -> None: self._state = state + self._emit(event, attempt=attempt, error=error) - async def _request_without_admission(self, message_type: int, payload: bytes) -> bytes: - transport = self._ensure_transport() - frame = FrameCodec.encode_frame(message_type, payload) - try: - return await self._multiplexer.request( - message_type, - frame, - transport.send, - self._timeout_ms, - ) - except Exception as exc: - self._handle_possible_transport_failure(exc) - raise - - def _handle_possible_transport_failure(self, exc: Exception) -> None: - if self._close_requested: + def _emit( + self, + event: str, + *, + attempt: int | None = None, + domain: str | None = None, + registration: str | None = None, + error: BaseException | None = None, + ) -> None: + payload = LifecycleEvent( + event=event, + state=self._state, + url=self._config.url, + transport=type(self._transport).__name__ if self._transport else None, + attempt=attempt, + domain=domain, + registration=registration, + error=str(error) if error else None, + ) + observer = self._config.observability + if observer.on_lifecycle_event is not None: + observer.on_lifecycle_event(payload) + self._log("info", f"fitz.connection.{event}", lifecycle=payload) + if observer.meter is not None: + observer.meter.counter("fitz.connection.lifecycle", 1, {"event": event}) + + def _record_latency(self, message_type: int, started: float) -> None: + elapsed = time.monotonic() - started + meter = self._config.observability.meter + if meter is not None: + meter.histogram("fitz.request.duration", elapsed, {"message_type": message_type}) + if elapsed > 1: + self._log("warning", "fitz.request.slow", message_type=message_type, duration=elapsed) + + def _handler_error(self, error: BaseException) -> None: + self._log("error", "fitz.handler.error", error=str(error)) + + def _log(self, level: str, event: str, **fields: Any) -> None: + logger = self._config.observability.logger + if logger is None: return - if isinstance(exc, (TransportError, ConnectionError, AuthenticationError)): - asyncio.create_task(self._handle_connection_loss(exc)) - - @staticmethod - def _describe_connection_loss(exc: Exception) -> str: - return str(exc) if str(exc) else "connection closed during CONNECT" + method = getattr(logger, level, None) + if callable(method): + method(event, extra={"fitz": fields}) diff --git a/src/fitz_py/domains/__init__.py b/src/fitz_py/domains/__init__.py index 91c2435..0a5078e 100644 --- a/src/fitz_py/domains/__init__.py +++ b/src/fitz_py/domains/__init__.py @@ -1,47 +1,35 @@ -"""Public domain client and domain model exports for Fitz Python SDK.""" +"""Domain clients and immutable result models.""" + +# ruff: noqa: F401 from fitz_py.domains.kv import ( KvClient, - KVDurability, + KvDurability, KvGetResult, - KVMode, + KvMode, + KvNotification, KvPair, - KvScanResult, + KvScanPage, KvTransaction, ) -from fitz_py.domains.lease import ( - Lease, - LeaseClient, - LeaseHandler, - LeaseInfo, - LeaseSubscription, -) -from fitz_py.domains.notice import ( - NoticeClient, - NoticeHandler, - NoticeMessage, - NoticeSubscription, -) -from fitz_py.domains.queue import ( - QueueAvailabilityHandler, - QueueClient, - QueueItem, - QueueSubscription, -) +from fitz_py.domains.lease import Lease, LeaseClient, LeaseInfo, ManagedLease +from fitz_py.domains.notice import Notice, NoticeClient +from fitz_py.domains.queue import Availability, QueueClient, QueueItem from fitz_py.domains.rpc import ( - InboundRpcRequest, + InboundRequest, ResponseFrame, ResponseWriter, + RpcCall, RpcClient, RpcHandler, - RpcSubscription, + Worker, ) from fitz_py.domains.schedule import ( + DeliveryMode, ScheduleClient, ScheduleEntry, - ScheduleHandler, ScheduleNotification, - ScheduleSubscription, + SchedulePage, ) from fitz_py.domains.stream import ( StreamClient, @@ -50,7 +38,6 @@ StreamFilterClause, StreamFilteredReason, StreamFilterSet, - StreamHandler, StreamMetadata, StreamReadCursor, StreamReadItem, @@ -58,54 +45,6 @@ StreamReadPage, StreamRecord, StreamSession, - StreamSubscription, ) -__all__ = [ - "InboundRpcRequest", - "KVDurability", - "KVMode", - "KvClient", - "KvGetResult", - "KvPair", - "KvScanResult", - "KvTransaction", - "Lease", - "LeaseClient", - "LeaseHandler", - "LeaseInfo", - "LeaseSubscription", - "NoticeClient", - "NoticeHandler", - "NoticeMessage", - "NoticeSubscription", - "QueueAvailabilityHandler", - "QueueClient", - "QueueItem", - "QueueSubscription", - "ResponseFrame", - "ResponseWriter", - "RpcClient", - "RpcHandler", - "RpcSubscription", - "ScheduleClient", - "ScheduleEntry", - "ScheduleHandler", - "ScheduleNotification", - "ScheduleSubscription", - "StreamClient", - "StreamCommitMode", - "StreamCommitNotification", - "StreamFilteredReason", - "StreamHandler", - "StreamMetadata", - "StreamReadCursor", - "StreamReadItem", - "StreamReadItemKind", - "StreamReadPage", - "StreamRecord", - "StreamSession", - "StreamSubscription", - "StreamFilterClause", - "StreamFilterSet", -] +__all__ = [name for name in globals() if not name.startswith("_")] diff --git a/src/fitz_py/domains/_routes.py b/src/fitz_py/domains/_routes.py index 5075599..d9cad4a 100644 --- a/src/fitz_py/domains/_routes.py +++ b/src/fitz_py/domains/_routes.py @@ -1,17 +1,30 @@ -"""Route-shape helper predicates used by domain validation routines.""" +"""Strict, opaque route-shape validation.""" from __future__ import annotations +def _parts(route: str, scheme: str) -> list[str] | None: + prefix = f"{scheme}://" + if not isinstance(route, str) or not route.startswith(prefix): + return None + parts = route[len(prefix) :].split("/") + if not parts or any(not part for part in parts): + return None + if any("*" in part and part not in {"*", "**"} for part in parts): + return None + if "**" in parts and parts[-1] != "**": + return None + return parts + + def is_exact_route_shape(route: str, scheme: str, segment_count: int) -> bool: - _ = (scheme, segment_count) - # Route strings are opaque protocol inputs; semantic validation is broker-owned. - return isinstance(route, str) + parts = _parts(route, scheme) + return parts is not None and len(parts) == segment_count and all("*" not in p for p in parts) def is_concrete_route_shape(route: str, scheme: str) -> bool: - _ = scheme - return isinstance(route, str) + parts = _parts(route, scheme) + return parts is not None and all("*" not in p for p in parts) def is_selector_route_shape( @@ -20,5 +33,15 @@ def is_selector_route_shape( segment_count: int, allow_realm_wildcard: bool = False, ) -> bool: - _ = (scheme, segment_count, allow_realm_wildcard) - return isinstance(route, str) + parts = _parts(route, scheme) + if parts is None: + return False + if parts == ["**"]: + return allow_realm_wildcard + if len(parts) == 2 and parts[-1] == "**": + return parts[0] != "*" or allow_realm_wildcard + if len(parts) != segment_count or "**" in parts: + return False + if parts[0] == "*" and not allow_realm_wildcard: + return False + return True diff --git a/src/fitz_py/domains/_subscriptions.py b/src/fitz_py/domains/_subscriptions.py new file mode 100644 index 0000000..335fa4b --- /dev/null +++ b/src/fitz_py/domains/_subscriptions.py @@ -0,0 +1,98 @@ +"""Shared one-wire/many-consumer subscription registry.""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from dataclasses import dataclass, field +from typing import Generic, TypeVar + +from fitz_py._runtime import AsyncSubscription +from fitz_py.errors import ReconnectRestoreError + +T = TypeVar("T") + + +@dataclass(slots=True) +class WireSubscription(Generic[T]): + sub_id: int + consumers: set[AsyncSubscription[T]] = field(default_factory=set) + + +class SubscriptionRegistry(Generic[T]): + def __init__( + self, + capacity: int, + subscribe_wire: Callable[[str], Awaitable[int]], + unsubscribe_wire: Callable[[str], Awaitable[None]], + ) -> None: + self._capacity = capacity + self._subscribe_wire = subscribe_wire + self._unsubscribe_wire = unsubscribe_wire + self._by_registration: dict[str, WireSubscription[T]] = {} + self._by_id: dict[int, str] = {} + + async def subscribe(self, registration: str) -> AsyncSubscription[T]: + state = self._by_registration.get(registration) + if state is None: + state = WireSubscription(await self._subscribe_wire(registration)) + self._by_registration[registration] = state + self._by_id[state.sub_id] = registration + + subscription: AsyncSubscription[T] + + async def close() -> None: + current = self._by_registration.get(registration) + if current is None: + return + current.consumers.discard(subscription) + if current.consumers: + return + await self._unsubscribe_wire(registration) + self._by_registration.pop(registration, None) + self._by_id.pop(current.sub_id, None) + + subscription = AsyncSubscription(registration, self._capacity, close) + state.consumers.add(subscription) + return subscription + + def publish(self, sub_id: int, item: T) -> None: + registration = self._by_id.get(sub_id) + if registration is None: + return + state = self._by_registration.get(registration) + if state is None: + return + dead = {consumer for consumer in state.consumers if not consumer.push(item)} + state.consumers.difference_update(dead) + + async def restore( + self, + *, + domain: str, + on_error: Callable[[str, BaseException], None] | None = None, + ) -> None: + for registration, state in list(self._by_registration.items()): + old_id = state.sub_id + try: + new_id = await self._subscribe_wire(registration) + except BaseException as exc: + error = ReconnectRestoreError(domain, registration, exc) + self.fail_registration(registration, error) + if on_error is not None: + on_error(registration, error) + continue + state.sub_id = new_id + self._by_id.pop(old_id, None) + self._by_id[new_id] = registration + + def fail_registration(self, registration: str, error: BaseException) -> None: + state = self._by_registration.pop(registration, None) + if state is None: + return + self._by_id.pop(state.sub_id, None) + for consumer in state.consumers: + consumer.fail(error) + + @property + def registrations(self) -> tuple[str, ...]: + return tuple(self._by_registration) diff --git a/src/fitz_py/domains/kv.py b/src/fitz_py/domains/kv.py index 3a5359f..0abed33 100644 --- a/src/fitz_py/domains/kv.py +++ b/src/fitz_py/domains/kv.py @@ -1,13 +1,17 @@ -"""KV domain client and transaction primitives for Fitz.""" +"""Transactional KV operations and mutation subscriptions.""" from __future__ import annotations +import asyncio +from collections.abc import AsyncIterator from dataclasses import dataclass -from typing import Literal +from enum import StrEnum -from fitz_py.domains._routes import is_exact_route_shape +from fitz_py._runtime import AsyncSubscription +from fitz_py.domains._routes import is_exact_route_shape, is_selector_route_shape +from fitz_py.domains._subscriptions import SubscriptionRegistry from fitz_py.domains.base import DomainClient -from fitz_py.errors import ErrKvOperationNotAllowed, KvError, kv_error +from fitz_py.errors import KvError, StaleHandleError, domain_error from fitz_py.protocol.buffer import BufferReader, BufferWriter from fitz_py.protocol.messages import ( MSG_KV_BEGIN, @@ -16,232 +20,284 @@ MSG_KV_DELETE_RANGE, MSG_KV_GET, MSG_KV_INSERT, + MSG_KV_NOTIFY, MSG_KV_PUT, MSG_KV_ROLLBACK, MSG_KV_SCAN, + MSG_KV_SUBSCRIBE, + MSG_KV_UNSUBSCRIBE, ) -KVMode = Literal["read_only", "read_write"] -KVDurability = Literal["buffered", "sync"] +class KvMode(StrEnum): + READ_ONLY = "read_only" + READ_WRITE = "read_write" -@dataclass(slots=True) -class KvGetResult: - """Result of a KV get operation.""" +class KvDurability(StrEnum): + BUFFERED = "buffered" + SYNC = "sync" + + +@dataclass(frozen=True, slots=True) +class KvGetResult: found: bool value: bytes | None = None -@dataclass(slots=True) +@dataclass(frozen=True, slots=True) class KvPair: - """A single key/value pair returned by KV scans.""" - key: bytes value: bytes -@dataclass(slots=True) -class KvScanResult: - """Batch result from a KV scan operation.""" +@dataclass(frozen=True, slots=True) +class KvScanPage: + entries: tuple[KvPair, ...] + has_more: bool - items: list[KvPair] - has_more: bool = False +@dataclass(frozen=True, slots=True) +class KvNotification: + route: str + mutation_count: int -class KvTransaction: - """Scoped transactional KV operations for a single begin/commit lifecycle.""" +class KvTransaction: def __init__(self, connection, route: str, tx_id: int) -> None: self._connection = connection self._route = route self._tx_id = tx_id + self._generation = connection.generation self._closed = False - self._closed_reason: str | None = None - on_disconnect = getattr(self._connection, "on_disconnect", None) - self._disconnect_unregister = ( - on_disconnect(self._invalidate) if callable(on_disconnect) else None - ) + self._lock = asyncio.Lock() + self._disconnect = connection.on_disconnect(self._invalidate) - async def __aenter__(self) -> "KvTransaction": + async def __aenter__(self) -> KvTransaction: return self - async def __aexit__(self, exc_type, exc, tb) -> None: - if not self._closed: - await self.rollback() + async def __aexit__(self, *_args: object) -> None: + await self.rollback() + + def _ensure_open(self) -> None: + if self._generation != self._connection.generation: + raise StaleHandleError("KV transaction") + if self._closed: + raise KvError("Transaction is closed", "TX_CLOSED") + + def _invalidate(self) -> None: + self._closed = True async def get(self, key: bytes) -> KvGetResult: - self._ensure_open("GET") - writer = BufferWriter() - writer.write_u64_be(self._tx_id) - writer.write_route(self._route) - writer.write_u32_be(len(key)) - writer.write_bytes(key) - reader = BufferReader(await self._connection.request(MSG_KV_GET, writer.build())) - status = reader.read_u8() - if status != 0: - raise kv_error(f"GET failed with status {status}", status) - found = not reader.is_eof() and reader.read_u8() == 1 - if not found or reader.is_eof(): - return KvGetResult(found=False) - value = reader.read_bytes(reader.read_u32_be()) - return KvGetResult(found=True, value=value) + async with self._lock: + self._ensure_open() + + async def operation() -> KvGetResult: + writer = self._prefix() + writer.write_u32_be(len(key)) + writer.write_bytes(key) + reader = BufferReader(await self._connection.request(MSG_KV_GET, writer.build())) + self._status(reader, "GET") + found = not reader.is_eof() and reader.read_u8() == 1 + if not found: + return KvGetResult(False) + return KvGetResult(True, reader.read_bytes(reader.read_u32_be())) + + return await self._connection.run_with_retry(operation, replay_safe=True) async def put(self, key: bytes, value: bytes) -> None: - await self._write(MSG_KV_PUT, key, value, "PUT") + await self._write(MSG_KV_PUT, "PUT", key, value) async def insert(self, key: bytes, value: bytes) -> None: - await self._write(MSG_KV_INSERT, key, value, "INSERT") + await self._write(MSG_KV_INSERT, "INSERT", key, value) async def delete(self, key: bytes) -> None: - self._ensure_open("DELETE") - writer = BufferWriter() - writer.write_u64_be(self._tx_id) - writer.write_route(self._route) - writer.write_u32_be(len(key)) - writer.write_bytes(key) - await self._expect_status(MSG_KV_DELETE, writer.build(), "DELETE") + async with self._lock: + self._ensure_open() + writer = self._prefix() + writer.write_u32_be(len(key)) + writer.write_bytes(key) + self._status( + BufferReader(await self._connection.request(MSG_KV_DELETE, writer.build())), + "DELETE", + ) async def delete_range(self, start_key: bytes, end_key: bytes) -> None: - self._ensure_open("DELETE_RANGE") - writer = BufferWriter() - writer.write_u64_be(self._tx_id) - writer.write_route(self._route) - writer.write_u32_be(len(start_key)) - writer.write_bytes(start_key) - writer.write_u32_be(len(end_key)) - writer.write_bytes(end_key) - await self._expect_status(MSG_KV_DELETE_RANGE, writer.build(), "DELETE_RANGE") - - async def scan( + if start_key >= end_key: + raise KvError("Range start must be less than end", "INVALID_RANGE") + async with self._lock: + self._ensure_open() + writer = self._prefix() + for key in (start_key, end_key): + writer.write_u32_be(len(key)) + writer.write_bytes(key) + self._status( + BufferReader(await self._connection.request(MSG_KV_DELETE_RANGE, writer.build())), + "DELETE_RANGE", + ) + + async def scan_page( self, *, start_key: bytes | None = None, end_key: bytes | None = None, limit: int | None = None, reverse: bool = False, - ) -> KvScanResult: - self._ensure_open("SCAN") - writer = BufferWriter() - writer.write_u64_be(self._tx_id) - writer.write_route(self._route) - if start_key is not None: - writer.write_u8(1) - writer.write_u32_be(len(start_key)) - writer.write_bytes(start_key) - else: - writer.write_u8(0) - if end_key is not None: - writer.write_u8(1) - writer.write_u32_be(len(end_key)) - writer.write_bytes(end_key) - else: - writer.write_u8(0) - if limit is not None and limit > 0: - writer.write_u8(1) - writer.write_u32_be(limit) - else: - writer.write_u8(0) - writer.write_u8(1 if reverse else 0) - - reader = BufferReader(await self._connection.request(MSG_KV_SCAN, writer.build())) - status = reader.read_u8() - if status != 0: - raise kv_error(f"SCAN failed with status {status}", status) - if reader.is_eof(): - return KvScanResult(items=[]) - count = reader.read_u32_be() - items: list[KvPair] = [] - for _ in range(count): - key = reader.read_bytes(reader.read_u32_be()) - value = reader.read_bytes(reader.read_u32_be()) - items.append(KvPair(key=key, value=value)) - has_more = not reader.is_eof() and reader.read_u8() == 1 - return KvScanResult(items=items, has_more=has_more) + ) -> KvScanPage: + if start_key is not None and end_key is not None and start_key >= end_key: + raise KvError("Range start must be less than end", "INVALID_RANGE") + async with self._lock: + self._ensure_open() + + async def operation() -> KvScanPage: + writer = self._prefix() + for value in (start_key, end_key): + writer.write_u8(0 if value is None else 1) + if value is not None: + writer.write_u32_be(len(value)) + writer.write_bytes(value) + writer.write_u8(0 if limit is None else 1) + if limit is not None: + writer.write_u32_be(limit) + writer.write_u8(1 if reverse else 0) + reader = BufferReader(await self._connection.request(MSG_KV_SCAN, writer.build())) + self._status(reader, "SCAN") + count = 0 if reader.is_eof() else reader.read_u32_be() + entries = tuple( + KvPair( + reader.read_bytes(reader.read_u32_be()), + reader.read_bytes(reader.read_u32_be()), + ) + for _ in range(count) + ) + return KvScanPage(entries, not reader.is_eof() and reader.read_u8() == 1) + + return await self._connection.run_with_retry(operation, replay_safe=True) + + async def scan(self, **options) -> AsyncIterator[KvPair]: + page = await self.scan_page(**options) + if page.has_more and options.get("limit") is None: + raise KvError("Unbounded scan was truncated", "SCAN_TRUNCATED") + for entry in page.entries: + yield entry async def commit(self) -> None: - await self._finalize(MSG_KV_COMMIT, "COMMIT") + await self._finish(MSG_KV_COMMIT, "COMMIT") async def rollback(self) -> None: - await self._finalize(MSG_KV_ROLLBACK, "ROLLBACK") - - async def _write(self, message_type: int, key: bytes, value: bytes, operation: str) -> None: - self._ensure_open(operation) - writer = BufferWriter() - writer.write_u64_be(self._tx_id) - writer.write_route(self._route) - writer.write_u32_be(len(key)) - writer.write_bytes(key) - writer.write_u32_be(len(value)) - writer.write_bytes(value) - await self._expect_status(message_type, writer.build(), operation) - - async def _finalize(self, message_type: int, operation: str) -> None: - self._ensure_open(operation) + if self._closed: + return + try: + await self._finish(MSG_KV_ROLLBACK, "ROLLBACK") + except Exception: + self._closed = True + + async def _write(self, message_type: int, operation: str, key: bytes, value: bytes) -> None: + async with self._lock: + self._ensure_open() + writer = self._prefix() + for item in (key, value): + writer.write_u32_be(len(item)) + writer.write_bytes(item) + self._status( + BufferReader(await self._connection.request(message_type, writer.build())), + operation, + ) + + async def _finish(self, message_type: int, operation: str) -> None: + async with self._lock: + self._ensure_open() + self._closed = True + self._disconnect() + reader = BufferReader( + await self._connection.request(message_type, self._prefix().build()) + ) + self._status(reader, operation) + + def _prefix(self) -> BufferWriter: writer = BufferWriter() writer.write_u64_be(self._tx_id) writer.write_route(self._route) - await self._expect_status(message_type, writer.build(), operation) - self._closed = True - self._closed_reason = "committed" if operation == "COMMIT" else "rolled back" - self._clear_disconnect_listener() - - def _ensure_open(self, operation: str) -> None: - if not self._closed: - return - - reason = self._closed_reason or "closed" - raise ErrKvOperationNotAllowed(f"{operation} not allowed: transaction already {reason}") - - def _invalidate(self) -> None: - if self._closed: - return - self._closed = True - self._closed_reason = "disconnected" - self._clear_disconnect_listener() + return writer - def _clear_disconnect_listener(self) -> None: - unregister = getattr(self, "_disconnect_unregister", None) - if unregister is None: - return - self._disconnect_unregister = None - unregister() - - async def _expect_status(self, message_type: int, payload: bytes, operation: str) -> None: - reader = BufferReader(await self._connection.request(message_type, payload)) + @staticmethod + def _status(reader: BufferReader, operation: str) -> None: status = reader.read_u8() if status != 0: - raise kv_error(f"{operation} failed with status {status}", status) + message = reader.read_string() if reader.remaining_bytes() >= 4 else None + raise domain_error(KvError, operation, status, message) class KvClient(DomainClient): - """KV domain entry point for creating transactions.""" + def __init__(self, connection) -> None: + super().__init__(connection) + capacity = connection.config.limits.subscription_buffer_size + self._subscriptions = SubscriptionRegistry[KvNotification]( + capacity, self._subscribe_wire, self._unsubscribe_wire + ) + connection.register_notification_handler(MSG_KV_NOTIFY, self._notify) + connection.on_reconnect( + lambda: self._subscriptions.restore( + domain="kv", + on_error=lambda registration, error: connection.report_restore_failure( + "kv", registration, error + ), + ), + domain="kv", + registration="mutations", + ) async def begin( self, route: str, *, - mode: KVMode = "read_write", - durability: KVDurability, + durability: KvDurability | str = KvDurability.BUFFERED, + mode: KvMode | str = KvMode.READ_WRITE, ) -> KvTransaction: - _assert_kv_route(route) + if not is_exact_route_shape(route, "kv", 3): + raise KvError(f"Invalid KV route: {route}", "INVALID_ROUTE") writer = BufferWriter() writer.write_route(route) - writer.write_u8(1 if mode == "read_write" else 0) - writer.write_u8(1 if durability == "sync" else 0) + writer.write_u8(1 if KvMode(mode) is KvMode.READ_WRITE else 0) + writer.write_u8(1 if KvDurability(durability) is KvDurability.SYNC else 0) reader = BufferReader(await self.request_frame(MSG_KV_BEGIN, writer.build())) - status = reader.read_u8() - if status != 0: - raise kv_error(f"BEGIN failed with status {status}", status) - tx_id = reader.read_u64_be() if not reader.is_eof() else None - if tx_id is None: - raise KvError("BEGIN response missing transaction id", "MISSING_TX_ID") - return KvTransaction(self.connection, route, tx_id) - - -def _assert_kv_route(route: str) -> None: - if not is_exact_route_shape(route, "kv", 3): - raise KvError( - f"Invalid kv route: {route} (expected kv://{{realm}}/{{area}}/{{resource}}, no empty segments or wildcards)", - "INVALID_ROUTE", + KvTransaction._status(reader, "BEGIN") + if reader.remaining_bytes() < 8: + raise KvError("BEGIN response missing transaction id", "INVALID_RESPONSE") + return KvTransaction(self.connection, route, reader.read_u64_be()) + + async def subscribe(self, pattern: str) -> AsyncSubscription[KvNotification]: + if not is_selector_route_shape(pattern, "kv", 3, True): + raise KvError(f"Invalid KV pattern: {pattern}", "INVALID_ROUTE") + return await self._subscriptions.subscribe(pattern) + + async def _subscribe_wire(self, pattern: str) -> int: + writer = BufferWriter() + writer.write_route(pattern) + reader = BufferReader(await self.request_frame(MSG_KV_SUBSCRIBE, writer.build())) + KvTransaction._status(reader, "SUBSCRIBE") + if reader.remaining_bytes() != 8: + raise KvError("SUBSCRIBE response missing id", "INVALID_RESPONSE") + return reader.read_u64_be() + + async def _unsubscribe_wire(self, pattern: str) -> None: + writer = BufferWriter() + writer.write_route(pattern) + KvTransaction._status( + BufferReader(await self.request_frame(MSG_KV_UNSUBSCRIBE, writer.build())), + "UNSUBSCRIBE", ) + + def _notify(self, payload: bytes) -> None: + reader = BufferReader(payload) + sub_id = reader.read_u64_be() + notification = KvNotification(reader.read_route(), reader.read_u64_be()) + if not reader.is_eof(): + return + self._subscriptions.publish(sub_id, notification) + + +# Descriptive compatibility-free public spelling. +KVDurability = KvDurability +KVMode = KvMode +KvScanResult = KvScanPage diff --git a/src/fitz_py/domains/lease.py b/src/fitz_py/domains/lease.py index 20982f2..140b2b4 100644 --- a/src/fitz_py/domains/lease.py +++ b/src/fitz_py/domains/lease.py @@ -1,14 +1,24 @@ -"""Lease domain client, lease handles, and lease change subscriptions.""" +"""Fenced leases, queued acquisition, and managed renewal.""" from __future__ import annotations import asyncio -from collections.abc import Awaitable, Callable +import contextlib +import time +from collections import deque from dataclasses import dataclass +from fitz_py._runtime import AsyncSubscription from fitz_py.domains._routes import is_exact_route_shape +from fitz_py.domains._subscriptions import SubscriptionRegistry from fitz_py.domains.base import DomainClient -from fitz_py.errors import LeaseError, lease_error +from fitz_py.errors import ( + LeaseError, + LeaseLifecycleError, + LeaseLostError, + StaleHandleError, + domain_error, +) from fitz_py.protocol.buffer import BufferReader, BufferWriter from fitz_py.protocol.messages import ( MSG_LEASE_ACQUIRE, @@ -19,199 +29,222 @@ MSG_LEASE_SUBSCRIBE, MSG_LEASE_UNSUBSCRIBE, ) -from fitz_py.protocol.response import assert_success - -LeaseHandler = Callable[[str], None | Awaitable[None]] +from fitz_py.protocol.response import parse_response -@dataclass(slots=True) +@dataclass(frozen=True, slots=True) class LeaseInfo: - """Current lease ownership and TTL information for a route.""" - is_held: bool - owner: str | None = None - ttl_remaining_secs: int | None = None + owner: str | None + ttl_remaining: int | None + pending_waiters: int + expires_at: float | None @dataclass(slots=True) class Lease: - """Client-side lease handle with extend and release helpers.""" - route: str - _token: int - _client: "LeaseClient" - - @property - def token(self) -> int: - return self._token - - async def extend(self, ttl_secs: int) -> None: - new_token = await self._client.extend(self.route, self._token, ttl_secs) + token: int + expires_at: float + _client: LeaseClient + _generation: int + _released: bool = False + + def _valid(self) -> None: + if self._released or self._generation != self._client.connection.generation: + raise StaleHandleError("Lease") + + async def extend(self, ttl: float) -> None: + self._valid() + new_token = await self._client._extend(self.route, self.token, ttl) if new_token is not None: - self._token = new_token + self.token = new_token + self.expires_at = time.time() + ttl async def release(self) -> None: - await self._client.release(self.route, self._token) - - -class LeaseSubscription: - """Handle for an active lease change subscription.""" - - def __init__( - self, sub_id: int, pattern: str, unsubscribe: Callable[[int], Awaitable[None]] - ) -> None: - self._sub_id = sub_id - self.pattern = pattern - self._unsubscribe = unsubscribe - - async def unsubscribe(self) -> None: - await self._unsubscribe(self._sub_id) + self._valid() + await self._client._release(self.route, self.token) + self._released = True + + +class ManagedLease: + def __init__(self, client: LeaseClient, route: str, ttl: float, wait: float) -> None: + self._client, self._route, self._ttl, self._wait = client, route, ttl, wait + self._lease: Lease | None = None + self._renewal: asyncio.Task[None] | None = None + self._lost: BaseException | None = None + + async def __aenter__(self) -> Lease: + self._lease = await self._client.acquire(self._route, ttl=self._ttl, wait=self._wait) + self._renewal = asyncio.create_task(self._renew()) + return self._lease + + async def __aexit__(self, *_args: object) -> None: + failures: list[BaseException] = [] + if self._renewal is not None: + self._renewal.cancel() + with contextlib.suppress(asyncio.CancelledError): + await self._renewal + if self._lost is not None: + failures.append(self._lost) + if self._lease is not None and not self._lease._released: + try: + await self._lease.release() + except BaseException as exc: + failures.append(exc) + if len(failures) == 1: + raise failures[0] + if failures: + raise LeaseLifecycleError(failures) + + async def _renew(self) -> None: + assert self._lease is not None + try: + while True: + await asyncio.sleep(self._ttl / 3) + await self._lease.extend(self._ttl) + except asyncio.CancelledError: + raise + except BaseException as exc: + self._lost = LeaseLostError(str(exc)) class LeaseClient(DomainClient): - """Lease domain operations for acquire, extend, release, query, and subscribe.""" - def __init__(self, connection) -> None: super().__init__(connection) - self._subscriptions: dict[int, tuple[str, LeaseHandler]] = {} - self._initialized = False - self.connection.on_reconnect(self._restore_subscriptions) + self._acquire_lock = asyncio.Lock() + self._queued: deque[asyncio.Future[bytes]] = deque() + self._subscriptions = SubscriptionRegistry[str]( + connection.config.limits.subscription_buffer_size, + self._subscribe_wire, + self._unsubscribe_wire, + ) + connection.register_notification_handler(MSG_LEASE_ACQUIRE, self._queued_reply) + connection.register_notification_handler(MSG_LEASE_NOTIFY, self._notify) + connection.on_reconnect( + lambda: self._subscriptions.restore( + domain="lease", + on_error=lambda registration, error: connection.report_restore_failure( + "lease", registration, error + ), + ), + domain="lease", + registration="changes", + ) + + async def acquire(self, route: str, *, ttl: float, wait: float = 0) -> Lease: + _route(route) + if ttl <= 0 or wait < 0 or wait > 2**32 - 1: + raise ValueError("ttl must be positive and wait must fit u32 seconds") + async with self._acquire_lock: + queued = asyncio.get_running_loop().create_future() + self._queued.append(queued) + writer = BufferWriter() + writer.write_route(route) + writer.write_route("") + writer.write_u64_be(int(ttl)) + writer.write_u32_be(int(wait)) + try: + payload = await self.request_frame(MSG_LEASE_ACQUIRE, writer.build()) + response_type, token = self._decode_acquire(payload) + if response_type in {2, 3}: + response_type, token = self._decode_acquire(await queued) + if response_type not in {0, 1}: + raise LeaseError( + "ACQUIRE returned a second queued response", "INVALID_RESPONSE" + ) + else: + self._queued.remove(queued) + except BaseException: + with contextlib.suppress(ValueError): + self._queued.remove(queued) + raise + return Lease(route, token, time.time() + ttl, self, self.connection.generation) + + def hold(self, route: str, *, ttl: float, wait: float = 0) -> ManagedLease: + return ManagedLease(self, route, ttl, wait) - async def acquire(self, route: str, ttl_secs: int) -> Lease: - _assert_lease_route(route) + async def query(self, route: str) -> LeaseInfo: + _route(route) writer = BufferWriter() writer.write_route(route) - writer.write_route("") - writer.write_u64_be(ttl_secs) - reader = BufferReader(await self.request_frame(MSG_LEASE_ACQUIRE, writer.build())) - status = reader.read_u8() - if status != 0: - raise lease_error(f"ACQUIRE failed with status {status}", status) - if not reader.is_eof(): - reader.read_u8() - token = reader.read_u64_be() if not reader.is_eof() else None - if token is None: - raise LeaseError("ACQUIRE response missing fencing token", "MISSING_TOKEN") - return Lease(route=route, _token=token, _client=self) - - async def extend(self, route: str, token: int, ttl_secs: int) -> int | None: - _assert_lease_route(route) - data = await self._send_token_ttl(MSG_LEASE_EXTEND, route, token, ttl_secs, "EXTEND") - if data and len(data) >= 8: - return BufferReader(data).read_u64_be() - return None - - async def release(self, route: str, token: int) -> None: - _assert_lease_route(route) + response = parse_response(await self.request_frame(MSG_LEASE_QUERY, writer.build())) + if not response.success: + raise domain_error(LeaseError, "QUERY", response.error_code or 0, response.error) + reader = BufferReader(response.data) + held = reader.read_u8() == 1 + if not held: + return LeaseInfo(False, None, None, reader.read_u32_be(), None) + owner = reader.read_route() + ttl = reader.read_u64_be() + return LeaseInfo(True, owner, ttl, reader.read_u32_be(), time.time() + ttl) + + async def subscribe(self, route: str) -> AsyncSubscription[str]: + _route(route) + return await self._subscriptions.subscribe(route) + + async def _extend(self, route: str, token: int, ttl: float) -> int | None: + if ttl <= 0: + raise ValueError("ttl must be positive") writer = BufferWriter() writer.write_route(route) writer.write_route("") writer.write_u64_be(token) - assert_success(await self.request_frame(MSG_LEASE_RELEASE, writer.build()), "RELEASE") + writer.write_u64_be(int(ttl)) + response = parse_response(await self.request_frame(MSG_LEASE_EXTEND, writer.build())) + if not response.success: + raise domain_error(LeaseError, "EXTEND", response.error_code or 0, response.error) + return BufferReader(response.data).read_u64_be() if len(response.data) >= 8 else None - async def query(self, route: str) -> LeaseInfo: - _assert_lease_route(route) - writer = BufferWriter() - writer.write_route(route) - reader = BufferReader(await self.request_frame(MSG_LEASE_QUERY, writer.build())) - status = reader.read_u8() - if status != 0: - raise lease_error(f"QUERY failed with status {status}", status) - has_holder = reader.read_u8() - if has_holder == 0: - if not reader.is_eof(): - reader.read_u32_be() - return LeaseInfo(is_held=False) - owner = reader.read_route() - ttl_remaining_secs = reader.read_u64_be() - if not reader.is_eof(): - reader.read_u32_be() - return LeaseInfo(is_held=True, owner=owner, ttl_remaining_secs=ttl_remaining_secs) - - async def subscribe(self, pattern: str, handler: LeaseHandler) -> LeaseSubscription: - _assert_lease_route(pattern) - self._init_notify_handler() - writer = BufferWriter() - writer.write_route(pattern) - reader = BufferReader(await self.request_frame(MSG_LEASE_SUBSCRIBE, writer.build())) - status = reader.read_u8() - if status != 0: - raise lease_error(f"SUBSCRIBE failed with status {status}", status) - sub_id = reader.read_u64_be() if not reader.is_eof() else None - if sub_id is None: - raise LeaseError("SUBSCRIBE response missing subscription id", "MISSING_SUB_ID") - self._subscriptions[sub_id] = (pattern, handler) - return LeaseSubscription(sub_id, pattern, self._unsubscribe) - - async def _send_token_ttl( - self, message_type: int, route: str, token: int, ttl_secs: int, operation: str - ) -> bytes: + async def _release(self, route: str, token: int) -> None: writer = BufferWriter() writer.write_route(route) writer.write_route("") writer.write_u64_be(token) - writer.write_u64_be(ttl_secs) - payload = await self.request_frame(message_type, writer.build()) - try: - return assert_success(payload, operation) - except LeaseError: - raise - except Exception as exc: - message = str(exc) - raise _map_lease_protocol_error(message) from exc - - async def _unsubscribe(self, sub_id: int) -> None: - subscription = self._subscriptions.pop(sub_id, None) - if subscription is None: - return - writer = BufferWriter() - writer.write_route(subscription[0]) - await self.request_frame(MSG_LEASE_UNSUBSCRIBE, writer.build()) - - def _init_notify_handler(self) -> None: - if self._initialized: - return - self._initialized = True - - def handler(payload: bytes) -> None: - try: - reader = BufferReader(payload) - sub_id = reader.read_u64_be() - route = reader.read_route() - subscription = self._subscriptions.get(sub_id) - if subscription is None: - return - result = subscription[1](route) - if asyncio.iscoroutine(result): - asyncio.create_task(result) - except Exception: + response = parse_response(await self.request_frame(MSG_LEASE_RELEASE, writer.build())) + if not response.success: + raise domain_error(LeaseError, "RELEASE", response.error_code or 0, response.error) + + def _decode_acquire(self, payload: bytes) -> tuple[int, int]: + response = parse_response(payload) + if not response.success: + raise domain_error(LeaseError, "ACQUIRE", response.error_code or 0, response.error) + reader = BufferReader(response.data) + response_type, token = reader.read_u8(), reader.read_u64_be() + if response_type not in {0, 1, 2, 3} or not reader.is_eof(): + raise LeaseError("Invalid ACQUIRE response", "INVALID_RESPONSE") + return response_type, token + + def _queued_reply(self, payload: bytes) -> None: + while self._queued: + future = self._queued.popleft() + if not future.done(): + future.set_result(payload) return - self.connection.register_notification_handler(MSG_LEASE_NOTIFY, handler) - - async def _restore_subscriptions(self) -> None: - if not self._subscriptions: - return - snapshot = list(self._subscriptions.values()) - self._subscriptions.clear() - for pattern, handler in snapshot: - await self.subscribe(pattern, handler) + async def _subscribe_wire(self, route: str) -> int: + writer = BufferWriter() + writer.write_route(route) + response = parse_response(await self.request_frame(MSG_LEASE_SUBSCRIBE, writer.build())) + if not response.success: + raise domain_error(LeaseError, "SUBSCRIBE", response.error_code or 0, response.error) + return BufferReader(response.data).read_u64_be() + async def _unsubscribe_wire(self, route: str) -> None: + writer = BufferWriter() + writer.write_route(route) + response = parse_response(await self.request_frame(MSG_LEASE_UNSUBSCRIBE, writer.build())) + if not response.success: + raise domain_error(LeaseError, "UNSUBSCRIBE", response.error_code or 0, response.error) -def _map_lease_protocol_error(message: str) -> LeaseError: - normalized = message.lower() - if "held" in normalized: - return lease_error(message, 1) - if "not found" in normalized: - return lease_error(message, 2) - if "invalid" in normalized or "token" in normalized or "fence" in normalized: - return lease_error(message, 3) - return LeaseError(message, "ERROR") + def _notify(self, payload: bytes) -> None: + reader = BufferReader(payload) + sub_id, route = reader.read_u64_be(), reader.read_route() + if reader.read_bytes(reader.read_u32_be()) or not reader.is_eof(): + raise LeaseError("Invalid LEASE_NOTIFY payload", "INVALID_RESPONSE") + self._subscriptions.publish(sub_id, route) -def _assert_lease_route(route: str) -> None: +def _route(route: str) -> None: if not is_exact_route_shape(route, "lease", 3): - raise LeaseError( - f"Invalid lease route: {route} (expected lease://{{realm}}/{{area}}/{{resource}}, no empty segments or wildcards)", - "INVALID_ROUTE", - ) + raise LeaseError(f"Invalid lease route: {route}", "INVALID_ROUTE") diff --git a/src/fitz_py/domains/notice.py b/src/fitz_py/domains/notice.py index 7003c58..9a39223 100644 --- a/src/fitz_py/domains/notice.py +++ b/src/fitz_py/domains/notice.py @@ -1,14 +1,13 @@ -"""Notice domain publish/subscribe client and notification models.""" +"""Fire-and-forget notices and bounded async subscriptions.""" from __future__ import annotations -import asyncio -from collections.abc import Awaitable, Callable from dataclasses import dataclass, field +from fitz_py._runtime import AsyncSubscription from fitz_py.domains._routes import is_exact_route_shape, is_selector_route_shape from fitz_py.domains.base import DomainClient -from fitz_py.errors import NoticeError, notice_error +from fitz_py.errors import NoticeError, domain_error from fitz_py.protocol.buffer import BufferReader, BufferWriter from fitz_py.protocol.messages import ( MSG_NOTICE_NOTIFY, @@ -16,171 +15,107 @@ MSG_NOTICE_SUBSCRIBE, MSG_NOTICE_UNSUBSCRIBE, ) +from fitz_py.protocol.response import parse_response -NoticeHandler = Callable[["NoticeMessage"], None | Awaitable[None]] - - -@dataclass(slots=True) -class NoticeMessage: - """Notice payload delivered to a subscriber callback.""" +@dataclass(frozen=True, slots=True) +class Notice: route: str body: bytes @dataclass(slots=True) -class _NoticeSubscriptionState: - """Internal handler registry for a single subscribed notice pattern.""" - +class _Wire: sub_id: int - handlers: dict[int, NoticeHandler] = field(default_factory=dict) - - -class NoticeSubscription: - """Handle for an active notice pattern subscription.""" - - def __init__( - self, - sub_id: int, - pattern: str, - handler: NoticeHandler, - unsubscribe: Callable[[str, int], Awaitable[None]], - handler_id: int, - ) -> None: - self._sub_id = sub_id - self.pattern = pattern - self.handler = handler - self._unsubscribe = unsubscribe - self._handler_id = handler_id - - async def unsubscribe(self) -> None: - await self._unsubscribe(self.pattern, self._handler_id) + consumers: set[AsyncSubscription[Notice]] = field(default_factory=set) class NoticeClient(DomainClient): - """Notice domain operations for publish and pattern subscriptions.""" - def __init__(self, connection) -> None: super().__init__(connection) - self._subscriptions_by_pattern: dict[str, _NoticeSubscriptionState] = {} - self._patterns_by_sub_id: dict[int, str] = {} - self._initialized = False - self._next_handler_id = 1 - self.connection.on_reconnect(self._restore_subscriptions) + self._by_pattern: dict[str, _Wire] = {} + self._pattern_by_id: dict[int, str] = {} + connection.register_notification_handler(MSG_NOTICE_NOTIFY, self._notify) + connection.on_reconnect(self._restore, domain="notice", registration="subscriptions") async def publish(self, route: str, body: bytes) -> None: - _assert_notice_route(route) + if not is_exact_route_shape(route, "notice", 3): + raise NoticeError(f"Invalid notice route: {route}", "INVALID_ROUTE") writer = BufferWriter() writer.write_route(route) writer.write_u32_be(len(body)) writer.write_bytes(body) - cancel_optional = self.connection.get_multiplexer().expect_optional_response( - MSG_NOTICE_PUBLISH + await self.connection.send(MSG_NOTICE_PUBLISH, writer.build()) + + async def subscribe(self, pattern: str) -> AsyncSubscription[Notice]: + if not is_selector_route_shape(pattern, "notice", 3, allow_realm_wildcard=True): + raise NoticeError(f"Invalid notice pattern: {pattern}", "INVALID_ROUTE") + state = self._by_pattern.get(pattern) + if state is None: + state = _Wire(await self._subscribe_wire(pattern)) + self._by_pattern[pattern] = state + self._pattern_by_id[state.sub_id] = pattern + subscription: AsyncSubscription[Notice] + + async def close() -> None: + current = self._by_pattern.get(pattern) + if current is None: + return + current.consumers.discard(subscription) + if current.consumers: + return + writer = BufferWriter() + writer.write_u64_be(current.sub_id) + response = parse_response( + await self.request_frame(MSG_NOTICE_UNSUBSCRIBE, writer.build()) + ) + if not response.success: + raise domain_error( + NoticeError, "UNSUBSCRIBE", response.error_code or 0, response.error + ) + self._by_pattern.pop(pattern, None) + self._pattern_by_id.pop(current.sub_id, None) + + subscription = AsyncSubscription( + pattern, self.connection.config.limits.subscription_buffer_size, close ) - try: - await self.connection.send_fire_and_forget(MSG_NOTICE_PUBLISH, writer.build()) - except Exception: - cancel_optional() - raise - - async def subscribe(self, pattern: str, handler: NoticeHandler) -> NoticeSubscription: - _assert_notice_pattern(pattern) - self._init_notify_handler() - existing = self._subscriptions_by_pattern.get(pattern) - if existing is None: - sub_id = await self._subscribe_wire(pattern) - existing = _NoticeSubscriptionState(sub_id=sub_id) - self._subscriptions_by_pattern[pattern] = existing - self._patterns_by_sub_id[sub_id] = pattern - - handler_id = self._next_handler_id - self._next_handler_id += 1 - existing.handlers[handler_id] = handler - return NoticeSubscription(existing.sub_id, pattern, handler, self._unsubscribe, handler_id) + state.consumers.add(subscription) + return subscription async def _subscribe_wire(self, pattern: str) -> int: writer = BufferWriter() writer.write_route(pattern) - reader = BufferReader(await self.request_frame(MSG_NOTICE_SUBSCRIBE, writer.build())) - status = reader.read_u8() - if status != 0: - raise notice_error(f"SUBSCRIBE failed with status {status}", status) - has_sub_id = reader.read_u8() if not reader.is_eof() else 0 - if has_sub_id != 1 or reader.remaining_bytes() < 8: - raise NoticeError("SUBSCRIBE response missing subscription id", "MISSING_SUB_ID") + response = parse_response(await self.request_frame(MSG_NOTICE_SUBSCRIBE, writer.build())) + if not response.success: + raise domain_error(NoticeError, "SUBSCRIBE", response.error_code or 0, response.error) + reader = BufferReader(response.data) + if reader.read_u8() != 1: + raise NoticeError("SUBSCRIBE response omitted its id", "INVALID_RESPONSE") return reader.read_u64_be() - async def _unsubscribe(self, pattern: str, handler_id: int) -> None: - subscription = self._subscriptions_by_pattern.get(pattern) - if subscription is None: + def _notify(self, payload: bytes) -> None: + reader = BufferReader(payload) + sub_id = reader.read_u64_be() + notice = Notice(reader.read_route(), reader.read_bytes(reader.read_u32_be())) + if not reader.is_eof(): + raise NoticeError("NOTICE_NOTIFY has trailing bytes", "INVALID_RESPONSE") + pattern = self._pattern_by_id.get(sub_id) + state = self._by_pattern.get(pattern) if pattern is not None else None + if state is None: return + dead = {consumer for consumer in state.consumers if not consumer.push(notice)} + state.consumers.difference_update(dead) - subscription.handlers.pop(handler_id, None) - if subscription.handlers: - return - - self._subscriptions_by_pattern.pop(pattern, None) - self._patterns_by_sub_id.pop(subscription.sub_id, None) - writer = BufferWriter() - writer.write_u64_be(subscription.sub_id) - await self.request_frame(MSG_NOTICE_UNSUBSCRIBE, writer.build()) - - def _init_notify_handler(self) -> None: - if self._initialized: - return - self._initialized = True - - def handler(payload: bytes) -> None: + async def _restore(self) -> None: + for pattern, state in list(self._by_pattern.items()): + old_id = state.sub_id try: - reader = BufferReader(payload) - sub_id = reader.read_u64_be() - route = reader.read_route() - body = reader.read_bytes(reader.read_u32_be()) - pattern = self._patterns_by_sub_id.get(sub_id) - if pattern is None: - return - subscription = self._subscriptions_by_pattern.get(pattern) - if subscription is None: - return - notification = NoticeMessage(route=route, body=body) - for callback in subscription.handlers.values(): - result = callback(notification) - if asyncio.iscoroutine(result): - asyncio.create_task(result) - except Exception: - return - - self.connection.register_notification_handler(MSG_NOTICE_NOTIFY, handler) - - async def _restore_subscriptions(self) -> None: - if not self._subscriptions_by_pattern: - return - snapshot = [ - (pattern, dict(state.handlers)) - for pattern, state in self._subscriptions_by_pattern.items() - ] - self._subscriptions_by_pattern.clear() - self._patterns_by_sub_id.clear() - for pattern, handlers in snapshot: - sub_id = await self._subscribe_wire(pattern) - self._subscriptions_by_pattern[pattern] = _NoticeSubscriptionState( - sub_id=sub_id, - handlers=handlers, - ) - self._patterns_by_sub_id[sub_id] = pattern - - -def _assert_notice_route(route: str) -> None: - if not is_exact_route_shape(route, "notice", 3): - raise NoticeError( - f"Invalid notice route: {route} (expected notice://{{realm}}/{{area}}/{{resource}}, no empty segments or wildcards)", - "INVALID_ROUTE", - ) - - -def _assert_notice_pattern(pattern: str) -> None: - if not is_selector_route_shape(pattern, "notice", 3, allow_realm_wildcard=True): - raise NoticeError( - f"Invalid notice pattern: {pattern} (expected notice://{{realm}}/{{area}}/{{resource}}, notice://{{realm}}/{{area}}/*, or notice://{{realm}}/**)", - "INVALID_ROUTE", - ) + state.sub_id = await self._subscribe_wire(pattern) + except BaseException as exc: + self._by_pattern.pop(pattern, None) + for consumer in state.consumers: + consumer.fail(exc) + self.connection.report_restore_failure("notice", pattern, exc) + continue + self._pattern_by_id.pop(old_id, None) + self._pattern_by_id[state.sub_id] = pattern diff --git a/src/fitz_py/domains/queue.py b/src/fitz_py/domains/queue.py index 43c8e85..a158d73 100644 --- a/src/fitz_py/domains/queue.py +++ b/src/fitz_py/domains/queue.py @@ -1,15 +1,14 @@ -"""Queue domain client, queue items, and availability subscriptions.""" +"""Queue operations and bounded async availability subscriptions.""" from __future__ import annotations -import asyncio -import contextlib -from collections.abc import Awaitable, Callable from dataclasses import dataclass +from fitz_py._runtime import AsyncSubscription from fitz_py.domains._routes import is_exact_route_shape, is_selector_route_shape +from fitz_py.domains._subscriptions import SubscriptionRegistry from fitz_py.domains.base import DomainClient -from fitz_py.errors import QueueError, queue_error +from fitz_py.errors import QueueError, StaleHandleError, domain_error from fitz_py.protocol.buffer import BufferReader, BufferWriter from fitz_py.protocol.messages import ( MSG_QUEUE_AVAILABILITY_NOTIFY, @@ -20,232 +19,207 @@ MSG_QUEUE_SUBSCRIBE, MSG_QUEUE_UNSUBSCRIBE, ) +from fitz_py.protocol.response import parse_response -QueueAvailabilityHandler = Callable[[str], None | Awaitable[None]] - -class QueueSubscription: - """Handle for an active queue availability subscription.""" - - def __init__( - self, sub_id: int, pattern: str, unsubscribe: Callable[[int], Awaitable[None]] - ) -> None: - self._sub_id = sub_id - self.pattern = pattern - self._unsubscribe = unsubscribe - - async def unsubscribe(self) -> None: - await self._unsubscribe(self._sub_id) +@dataclass(frozen=True, slots=True) +class Availability: + route: str + ready: int + delayed: int + inflight: int @dataclass(slots=True) class QueueItem: - """Reserved queue item with completion and lease-extension helpers.""" - route: str + body: bytes _id: int _token: int - body: bytes - _client: "QueueClient" + _client: QueueClient + _generation: int + _closed: bool = False + + def _ensure_valid(self) -> None: + if self._closed or self._generation != self._client.connection.generation: + raise StaleHandleError("Queue item") async def extend(self, lease_seconds: int) -> None: + self._ensure_valid() await self._client._extend(self.route, self._id, self._token, lease_seconds) async def complete(self) -> None: + self._ensure_valid() await self._client._complete(self.route, self._id, self._token) - - async def complete_with_token(self, token: int) -> None: - await self._client._complete(self.route, self._id, token) + self._closed = True class QueueClient(DomainClient): - """Queue domain operations for enqueue, reserve, and availability subscribe.""" - def __init__(self, connection) -> None: super().__init__(connection) - self._subscriptions: dict[int, tuple[str, QueueAvailabilityHandler]] = {} - self._initialized = False - self.connection.on_reconnect(self._restore_subscriptions) + self._subscriptions = SubscriptionRegistry[Availability]( + connection.config.limits.subscription_buffer_size, + self._subscribe_wire, + self._unsubscribe_wire, + ) + connection.register_notification_handler(MSG_QUEUE_AVAILABILITY_NOTIFY, self._notify) + connection.on_reconnect( + lambda: self._subscriptions.restore( + domain="queue", + on_error=lambda registration, error: connection.report_restore_failure( + "queue", registration, error + ), + ), + domain="queue", + registration="availability", + ) - async def enqueue(self, route: str, body: bytes, *, delay_ms: int | None = None) -> int: - _assert_queue_route(route) + async def enqueue(self, route: str, body: bytes, *, delay: float | None = None) -> int | None: + _exact(route) + if delay is not None and delay < 0: + raise ValueError("delay must be non-negative") writer = BufferWriter() writer.write_route(route) writer.write_u32_be(len(body)) writer.write_bytes(body) - delay_seconds = (delay_ms or 0) // 1000 - writer.write_u8(1 if delay_seconds > 0 else 0) - if delay_seconds > 0: - writer.write_u64_be(delay_seconds) - reader = BufferReader(await self.request_frame(MSG_QUEUE_ENQUEUE, writer.build())) - status = reader.read_u8() - if status != 0: - raise queue_error(f"ENQUEUE failed with status {status}", status) - return reader.read_u64_be() if not reader.is_eof() else 0 + seconds = int(delay or 0) + writer.write_u8(int(seconds > 0)) + if seconds: + writer.write_u64_be(seconds) + response = parse_response(await self.request_frame(MSG_QUEUE_ENQUEUE, writer.build())) + if not response.success: + raise domain_error(QueueError, "ENQUEUE", response.error_code or 0, response.error) + reader = BufferReader(response.data) + return reader.read_u64_be() if reader.remaining_bytes() >= 8 else None async def reserve( self, - route: str, - lease_seconds: int, + selector: str, *, - batch_size: int = 1, - wait_seconds: int = 0, - ) -> list[QueueItem]: - _assert_queue_reserve_route(route) - if wait_seconds <= 0: - return await self._reserve_once(route, lease_seconds, batch_size) - - items = await self._reserve_once(route, lease_seconds, batch_size) - if items: - return items - - availability = asyncio.Event() - subscription = await self.subscribe(route, lambda _route: availability.set()) - deadline = asyncio.get_running_loop().time() + wait_seconds - - try: - while True: - items = await self._reserve_once(route, lease_seconds, batch_size) - if items: - return items - - remaining = deadline - asyncio.get_running_loop().time() - if remaining <= 0: - return [] - - try: - await asyncio.wait_for(availability.wait(), timeout=remaining) - except asyncio.TimeoutError: - return [] - finally: - availability.clear() - finally: - with contextlib.suppress(Exception): - await subscription.unsubscribe() - - async def _reserve_once( - self, - route: str, - lease_seconds: int, - batch_size: int, + lease: float, + batch_size: int | None = None, + wait: float | None = None, ) -> list[QueueItem]: + _selector(selector) + if ( + lease <= 0 + or batch_size is not None + and batch_size <= 0 + or wait is not None + and wait < 0 + ): + raise ValueError("lease and batch_size must be positive; wait must be non-negative") writer = BufferWriter() - writer.write_route(route) - writer.write_u64_be(lease_seconds) - writer.write_u8(1 if batch_size > 0 else 0) - if batch_size > 0: + writer.write_route(selector) + writer.write_u64_be(int(lease)) + writer.write_u8(int(batch_size is not None)) + if batch_size is not None: writer.write_u32_be(batch_size) - reader = BufferReader(await self.request_frame(MSG_QUEUE_RESERVE, writer.build())) - status = reader.read_u8() - if status != 0: - raise queue_error(f"RESERVE failed with status {status}", status) + writer.write_u8(int(wait is not None and wait > 0)) + if wait is not None and wait > 0: + writer.write_u64_be(int(wait)) + response = parse_response(await self.request_frame(MSG_QUEUE_RESERVE, writer.build())) + if not response.success: + raise domain_error(QueueError, "RESERVE", response.error_code or 0, response.error) + reader = BufferReader(response.data) if reader.is_eof(): return [] - count = reader.read_u32_be() + wildcard = "*" in selector items: list[QueueItem] = [] - for _ in range(count): + for _ in range(reader.read_u32_be()): + route = reader.read_route() if wildcard else selector + if not is_exact_route_shape(route, "queue", 3): + raise QueueError("RESERVE returned an invalid route", "INVALID_RESPONSE") item_id = reader.read_u64_be() token = reader.read_u64_be() body = reader.read_bytes(reader.read_u32_be()) - items.append(QueueItem(route=route, _id=item_id, _token=token, body=body, _client=self)) + items.append( + QueueItem( + route, + body, + item_id, + token, + self, + self.connection.generation, + ) + ) + if not reader.is_eof(): + raise QueueError("RESERVE response has trailing bytes", "INVALID_RESPONSE") return items - async def subscribe(self, pattern: str, handler: QueueAvailabilityHandler) -> QueueSubscription: - _assert_queue_subscription_pattern(pattern) - self._init_notify_handler() - writer = BufferWriter() - writer.write_route(pattern) - reader = BufferReader(await self.request_frame(MSG_QUEUE_SUBSCRIBE, writer.build())) - status = reader.read_u8() - if status != 0: - raise queue_error(f"SUBSCRIBE failed with status {status}", status) - has_sub_id = reader.read_u8() if not reader.is_eof() else 0 - if has_sub_id != 1 or reader.is_eof(): - raise QueueError("SUBSCRIBE response missing subscription id", "MISSING_SUB_ID") - sub_id = reader.read_u64_be() - self._subscriptions[sub_id] = (pattern, handler) - return QueueSubscription(sub_id, pattern, self._unsubscribe) + async def subscribe(self, selector: str) -> AsyncSubscription[Availability]: + _selector(selector) + return await self._subscriptions.subscribe(selector) async def _complete(self, route: str, item_id: int, token: int) -> None: writer = BufferWriter() writer.write_route(route) writer.write_u64_be(item_id) writer.write_u64_be(token) - reader = BufferReader(await self.request_frame(MSG_QUEUE_COMPLETE, writer.build())) - status = reader.read_u8() - if status != 0: - raise queue_error(f"COMPLETE failed with status {status}", status) + response = parse_response( + await self.request_frame(MSG_QUEUE_COMPLETE, writer.build()), plain=True + ) + if not response.success: + raise QueueError(f"COMPLETE failed: {response.error}", "COMPLETE") - async def _extend(self, route: str, item_id: int, token: int, lease_seconds: int) -> None: + async def _extend(self, route: str, item_id: int, token: int, lease: float) -> None: + if lease <= 0: + raise ValueError("lease must be positive") writer = BufferWriter() writer.write_route(route) writer.write_u64_be(item_id) writer.write_u64_be(token) - writer.write_u64_be(lease_seconds) - reader = BufferReader(await self.request_frame(MSG_QUEUE_EXTEND, writer.build())) - status = reader.read_u8() - if status != 0: - raise queue_error(f"EXTEND failed with status {status}", status) - - async def _unsubscribe(self, sub_id: int) -> None: - subscription = self._subscriptions.pop(sub_id, None) - if subscription is None: - return + writer.write_u64_be(int(lease)) + response = parse_response( + await self.request_frame(MSG_QUEUE_EXTEND, writer.build()), plain=True + ) + if not response.success: + raise QueueError(f"EXTEND failed: {response.error}", "EXTEND") + + async def _subscribe_wire(self, selector: str) -> int: writer = BufferWriter() - writer.write_route(subscription[0]) - await self.request_frame(MSG_QUEUE_UNSUBSCRIBE, writer.build()) - - def _init_notify_handler(self) -> None: - if self._initialized: - return - self._initialized = True - - def handler(payload: bytes) -> None: - try: - reader = BufferReader(payload) - sub_id = reader.read_u64_be() - route = reader.read_route() - if not reader.is_eof(): - reader.read_bytes(reader.read_u32_be()) - subscription = self._subscriptions.get(sub_id) - if subscription is None: - return - result = subscription[1](route) - if asyncio.iscoroutine(result): - asyncio.create_task(result) - except Exception: - return - - self.connection.register_notification_handler(MSG_QUEUE_AVAILABILITY_NOTIFY, handler) - - async def _restore_subscriptions(self) -> None: - if not self._subscriptions: - return - snapshot = list(self._subscriptions.values()) - self._subscriptions.clear() - for pattern, handler in snapshot: - await self.subscribe(pattern, handler) - - -def _assert_queue_route(route: str) -> None: - if not is_exact_route_shape(route, "queue", 3): - raise QueueError( - f"Invalid queue route: {route} (expected queue://{{realm}}/{{area}}/{{resource}}, no empty segments or wildcards)", - "INVALID_ROUTE", + writer.write_route(selector) + response = parse_response( + await self.request_frame(MSG_QUEUE_SUBSCRIBE, writer.build()), plain=True ) + if not response.success: + raise QueueError(f"SUBSCRIBE failed: {response.error}", "SUBSCRIBE") + reader = BufferReader(response.data) + if reader.read_u8() != 1: + raise QueueError("SUBSCRIBE response omitted its id", "INVALID_RESPONSE") + sub_id = reader.read_u64_be() + if not reader.is_eof(): + raise QueueError("SUBSCRIBE response has trailing bytes", "INVALID_RESPONSE") + return sub_id + async def _unsubscribe_wire(self, selector: str) -> None: + writer = BufferWriter() + writer.write_route(selector) + response = parse_response( + await self.request_frame(MSG_QUEUE_UNSUBSCRIBE, writer.build()), plain=True + ) + if not response.success: + raise QueueError(f"UNSUBSCRIBE failed: {response.error}", "UNSUBSCRIBE") -def _assert_queue_reserve_route(route: str) -> None: - if not is_selector_route_shape(route, "queue", 3): - raise QueueError( - f"Invalid queue route: {route} (expected queue://{{realm}}/{{area}}/{{resource}} or queue://{{realm}}/{{area}}/*, no empty segments or wildcards)", - "INVALID_ROUTE", + def _notify(self, payload: bytes) -> None: + reader = BufferReader(payload) + sub_id = reader.read_u64_be() + item = Availability( + reader.read_route(), + reader.read_u64_be(), + reader.read_u64_be(), + reader.read_u64_be(), ) + if not reader.is_eof(): + raise QueueError("Availability notification has trailing bytes", "INVALID_RESPONSE") + self._subscriptions.publish(sub_id, item) -def _assert_queue_subscription_pattern(pattern: str) -> None: - if not is_selector_route_shape(pattern, "queue", 3, allow_realm_wildcard=True): - raise QueueError( - f"Invalid queue pattern: {pattern} (expected queue://{{realm}}/{{area}}/{{resource}}, queue://{{realm}}/{{area}}/*, or queue://{{realm}}/**)", - "INVALID_ROUTE", - ) +def _exact(route: str) -> None: + if not is_exact_route_shape(route, "queue", 3): + raise QueueError(f"Invalid queue route: {route}", "INVALID_ROUTE") + + +def _selector(route: str) -> None: + if not is_selector_route_shape(route, "queue", 3): + raise QueueError(f"Invalid queue selector: {route}", "INVALID_ROUTE") diff --git a/src/fitz_py/domains/rpc.py b/src/fitz_py/domains/rpc.py index 8050352..8e22e61 100644 --- a/src/fitz_py/domains/rpc.py +++ b/src/fitz_py/domains/rpc.py @@ -1,309 +1,328 @@ -"""RPC domain client, request/response streaming, and worker registration.""" +"""Streaming RPC calls and worker registrations.""" from __future__ import annotations import asyncio +import contextlib import os from collections.abc import AsyncIterator, Awaitable, Callable from dataclasses import dataclass -from fitz_py.domains._routes import is_concrete_route_shape from fitz_py.domains.base import DomainClient -from fitz_py.errors import ConnectionError, ErrRpcTimeout, RpcError, TransportError, rpc_error +from fitz_py.errors import FitzConnectionError, FitzTimeoutError, RpcError, domain_error from fitz_py.protocol.buffer import BufferReader, BufferWriter from fitz_py.protocol.messages import ( - MSG_RPC_ACK, MSG_RPC_REQUEST, MSG_RPC_RESPONSE, MSG_RPC_SUBSCRIBE_WORKER, MSG_RPC_UNSUBSCRIBE_WORKER, ) -from fitz_py.types import ConnectionState +from fitz_py.protocol.response import parse_response -RpcHandler = Callable[["InboundRpcRequest", "ResponseWriter"], None | Awaitable[None]] - -@dataclass(slots=True) +@dataclass(frozen=True, slots=True) class ResponseFrame: - """Single response frame emitted by an RPC call stream.""" - body: bytes sequence: int -@dataclass(slots=True) -class InboundRpcRequest: - """Inbound RPC request payload delivered to a registered worker.""" - +@dataclass(frozen=True, slots=True) +class InboundRequest: route: str - reply_route: str body: bytes class ResponseWriter: - """Sends streamed RPC response frames back to the broker.""" - def __init__(self, connection, correlation_id: bytes) -> None: self._connection = connection self._correlation_id = correlation_id self._sequence = 0 + self._generation = connection.generation + self._ended = False - async def send(self, body: bytes, is_end: bool) -> None: + async def send(self, body: bytes, *, end: bool = False) -> None: + if self._ended or self._generation != self._connection.generation: + raise FitzConnectionError("RPC response writer is stale") writer = BufferWriter() - writer.write_u32_be(len(self._correlation_id)) writer.write_bytes(self._correlation_id) writer.write_u64_be(self._sequence) - self._sequence += 1 + writer.write_u8(int(end)) writer.write_u32_be(len(body)) writer.write_bytes(body) - writer.write_u8(1 if is_end else 0) - try: - await self._connection.send(MSG_RPC_RESPONSE, writer.build()) - except Exception as exc: - if _is_benign_shutdown_error(exc, self._connection): - return - raise - - -class RpcSubscription: - """Handle for an active RPC worker registration.""" - - def __init__(self, route: str, unsubscribe: Callable[[str], Awaitable[None]]) -> None: - self.route = route - self._unsubscribe = unsubscribe - - async def unsubscribe(self) -> None: - await self._unsubscribe(self.route) - + self._sequence += 1 + await self._connection.send(MSG_RPC_RESPONSE, writer.build()) + self._ended = end -class RpcIterator(AsyncIterator[ResponseFrame]): - """Async iterator over streamed RPC response frames.""" - def __init__( - self, - correlation_id: bytes, - client: "RpcClient", - timeout_ms: int, - release_gate: Callable[[], None] | None = None, - ) -> None: - self._correlation_id = correlation_id +class RpcCall(AsyncIterator[ResponseFrame]): + def __init__(self, client: RpcClient, key: bytes, timeout: float, capacity: int) -> None: self._client = client - self._timeout_ms = timeout_ms - self._buffer: list[ResponseFrame] = [] - self._done = False - self._waiter: asyncio.Future[ResponseFrame | None] | None = None - self._release_gate = release_gate - - def push(self, frame: ResponseFrame) -> None: - if self._waiter is not None and not self._waiter.done(): - self._waiter.set_result(frame) - self._waiter = None - return - self._buffer.append(frame) + self._key = key + self._timeout = timeout + self._queue: asyncio.Queue[ResponseFrame | BaseException | None] = asyncio.Queue(capacity) + self._closed = False - def end(self) -> None: - self._done = True - self._release_gate_if_needed() - if self._waiter is not None and not self._waiter.done(): - self._waiter.set_result(None) - self._waiter = None + def __aiter__(self) -> RpcCall: + return self async def __anext__(self) -> ResponseFrame: - if self._buffer: - return self._buffer.pop(0) - if self._done: + if self._closed and self._queue.empty(): raise StopAsyncIteration - self._waiter = asyncio.get_running_loop().create_future() try: - frame = await asyncio.wait_for(self._waiter, timeout=self._timeout_ms / 1000) + async with asyncio.timeout(self._timeout): + item = await self._queue.get() + except asyncio.CancelledError: + await self.aclose() + raise except TimeoutError as exc: - self._client.cleanup_pending_rpc(self._correlation_id) - self._done = True - self._release_gate_if_needed() - raise ErrRpcTimeout("RPC call timeout") from exc - if frame is None: + await self.aclose() + raise FitzTimeoutError("RPC response timed out") from exc + if item is None: + self._closed = True raise StopAsyncIteration - return frame + if isinstance(item, BaseException): + self._closed = True + raise item + return item + + async def __aenter__(self) -> RpcCall: + return self + + async def __aexit__(self, *_args: object) -> None: + await self.aclose() async def aclose(self) -> None: - self._done = True - self._client.cleanup_pending_rpc(self._correlation_id) - self._release_gate_if_needed() + if not self._closed: + self._closed = True + self._client._pending.pop(self._key, None) - def _release_gate_if_needed(self) -> None: - release_gate = self._release_gate - if release_gate is None: + def push(self, frame: ResponseFrame, end: bool) -> None: + if self._closed: + return + try: + if frame.body: + self._queue.put_nowait(frame) + if end: + self._queue.put_nowait(None) + self._client._pending.pop(self._key, None) + except asyncio.QueueFull: + self.fail(RpcError("RPC response consumer fell behind", "BACKPRESSURE")) + + def fail(self, error: BaseException) -> None: + if self._closed: return - self._release_gate = None - release_gate() + self._closed = True + while not self._queue.empty(): + with contextlib.suppress(asyncio.QueueEmpty): + self._queue.get_nowait() + self._queue.put_nowait(error) -class RpcClient(DomainClient): - """RPC domain entry point for issuing calls and registering workers.""" +RpcHandler = Callable[[InboundRequest, ResponseWriter], Awaitable[None]] + + +@dataclass(slots=True) +class Worker: + route: str + _client: RpcClient + _identity: object + + async def unsubscribe(self) -> None: + await self._client._unregister(self.route, self._identity) + async def aclose(self) -> None: + await self.unsubscribe() + + +@dataclass(frozen=True, slots=True) +class _Registration: + handler: RpcHandler + max_concurrency: int + identity: object + + +class RpcClient(DomainClient): def __init__(self, connection) -> None: super().__init__(connection) - self._pending: dict[str, RpcIterator] = {} - self._workers: dict[str, RpcHandler] = {} - self._initialized = False - self.connection.on_reconnect(self._restore_workers) - - async def call(self, route: str, body: bytes, *, timeout_ms: int = 30000) -> RpcIterator: - _assert_rpc_route(route) - self._init_handlers() + self._pending: dict[bytes, RpcCall] = {} + self._workers: dict[str, _Registration] = {} + connection.register_push_classifier(MSG_RPC_REQUEST, _looks_like_request) + connection.register_push_classifier(MSG_RPC_RESPONSE, _looks_like_response) + connection.register_notification_handler(MSG_RPC_REQUEST, self._on_request) + connection.register_notification_handler(MSG_RPC_RESPONSE, self._on_response) + connection.on_disconnect(self._disconnect) + connection.on_reconnect(self._restore, domain="rpc", registration="workers") + + async def call(self, route: str, body: bytes, *, timeout: float = 30.0) -> RpcCall: + _route(route, patterns=False) + if timeout <= 0: + raise ValueError("timeout must be positive") correlation_id = os.urandom(16) - admission_reserved = await self._reserve_admission_if_supported() - release_gate = getattr(self.connection, "release_admission_nowait", None) - if not callable(release_gate): - release_gate = None - iterator = RpcIterator( - correlation_id, + call = RpcCall( self, - timeout_ms, - release_gate=release_gate, + correlation_id, + timeout, + self.connection.config.limits.subscription_buffer_size, ) - self._pending[correlation_id.hex()] = iterator - + self._pending[correlation_id] = call writer = BufferWriter() - writer.write_u32_be(len(correlation_id)) writer.write_bytes(correlation_id) writer.write_route(route) - writer.write_route("") writer.write_u32_be(len(body)) writer.write_bytes(body) try: - reader = BufferReader( - await self._request_without_admission(MSG_RPC_REQUEST, writer.build()) - ) - status = reader.read_u8() - if status != 0: - self._pending.pop(correlation_id.hex(), None) - await self._release_admission_if_supported(admission_reserved) - raise rpc_error(f"REQUEST failed with status {status}", status) - return iterator - except Exception: - self._pending.pop(correlation_id.hex(), None) - await self._release_admission_if_supported(admission_reserved) + await self.connection.send(MSG_RPC_REQUEST, writer.build()) + except BaseException: + self._pending.pop(correlation_id, None) raise - - async def register_worker(self, route: str, handler: RpcHandler) -> RpcSubscription: - _assert_rpc_route(route) - self._init_handlers() + return call + + async def register_worker( + self, route: str, handler: RpcHandler, *, max_concurrency: int = 1 + ) -> Worker: + _route(route, patterns=True) + if not 1 <= max_concurrency <= 1024: + raise ValueError("max_concurrency must be between 1 and 1024") + identity = object() + registration = _Registration(handler, max_concurrency, identity) + await self._subscribe(route, registration) + self._workers[route] = registration + return Worker(route, self, identity) + + async def _subscribe(self, route: str, registration: _Registration) -> None: writer = BufferWriter() writer.write_route(route) - reader = BufferReader(await self.request_frame(MSG_RPC_SUBSCRIBE_WORKER, writer.build())) - status = reader.read_u8() - if status != 0: - raise rpc_error(f"REGISTER_WORKER failed with status {status}", status) - self._workers[route] = handler - return RpcSubscription(route, self._unregister_worker) - - def cleanup_pending_rpc(self, correlation_id: bytes) -> None: - self._pending.pop(correlation_id.hex(), None) - - async def _reserve_admission_if_supported(self) -> bool: - reserve_admission = getattr(self.connection, "reserve_admission", None) - if not callable(reserve_admission): - return False - await reserve_admission() - return True + writer.write_u32_be(registration.max_concurrency) + response = parse_response( + await self.request_frame(MSG_RPC_SUBSCRIBE_WORKER, writer.build()) + ) + if not response.success: + raise domain_error(RpcError, "SUBSCRIBE", response.error_code or 0, response.error) - async def _release_admission_if_supported(self, admission_reserved: bool) -> None: - if not admission_reserved: + async def _unregister(self, route: str, identity: object) -> None: + registration = self._workers.get(route) + if registration is None or registration.identity is not identity: return - release_admission = getattr(self.connection, "release_admission", None) - if callable(release_admission): - await release_admission() - return - release_admission_nowait = getattr(self.connection, "release_admission_nowait", None) - if callable(release_admission_nowait): - release_admission_nowait() - - async def _request_without_admission(self, message_type: int, payload: bytes) -> bytes: - request_without_admission = getattr(self.connection, "request_without_admission", None) - if callable(request_without_admission): - return await request_without_admission(message_type, payload) - return await self.connection.request(message_type, payload) - - async def _unregister_worker(self, route: str) -> None: - self._workers.pop(route, None) writer = BufferWriter() writer.write_route(route) - try: + response = parse_response( await self.request_frame(MSG_RPC_UNSUBSCRIBE_WORKER, writer.build()) - except Exception: - return + ) + if not response.success: + raise domain_error(RpcError, "UNSUBSCRIBE", response.error_code or 0, response.error) + self._workers.pop(route, None) - def _init_handlers(self) -> None: - if self._initialized: + def _on_response(self, payload: bytes) -> None: + reader = BufferReader(payload) + key = reader.read_bytes(16) + sequence = reader.read_u64_be() + flags = reader.read_u8() + if flags & ~1: + raise RpcError("Unsupported RPC response flags", "INVALID_RESPONSE") + body = reader.read_bytes(reader.read_u32_be()) + if not reader.is_eof(): + raise RpcError("RPC response has trailing bytes", "INVALID_RESPONSE") + call = self._pending.get(key) + if call is None: return - self._initialized = True - - def response_handler(payload: bytes) -> None: - try: - reader = BufferReader(payload) - corr_len = reader.read_u32_be() - correlation_id = reader.read_bytes(corr_len) - sequence = reader.read_u64_be() - body = reader.read_bytes(reader.read_u32_be()) - stream_end = not reader.is_eof() and reader.read_u8() == 1 - iterator = self._pending.get(correlation_id.hex()) - if iterator is None: - return - if body: - iterator.push(ResponseFrame(body=body, sequence=sequence)) - if stream_end: - self._pending.pop(correlation_id.hex(), None) - iterator.end() - except Exception: + if flags & 1 and body[:1] == b"\x01": + error = parse_response(body) + if ( + not error.success + and error.error_code is not None + and 6001 <= error.error_code <= 6013 + ): + call.fail(domain_error(RpcError, "CALL", error.error_code, error.error)) + self._pending.pop(key, None) return - - def request_handler(payload: bytes) -> None: - try: - reader = BufferReader(payload) - corr_len = reader.read_u32_be() - correlation_id = reader.read_bytes(corr_len) - route = reader.read_route() - reply_route = reader.read_route() - body = reader.read_bytes(reader.read_u32_be()) - handler = self._workers.get(route) - if handler is None: - return - request = InboundRpcRequest(route=route, reply_route=reply_route, body=body) - response_writer = ResponseWriter(self.connection, correlation_id) - result = handler(request, response_writer) - if asyncio.iscoroutine(result): - asyncio.create_task(result) - except Exception: - return - - self.connection.register_notification_handler(MSG_RPC_RESPONSE, response_handler) - self.connection.register_notification_handler(MSG_RPC_REQUEST, request_handler) - self.connection.register_notification_handler(MSG_RPC_ACK, lambda _payload: None) - - async def _restore_workers(self) -> None: - if not self._workers: + call.push(ResponseFrame(body, sequence), bool(flags & 1)) + + def _on_request(self, payload: bytes) -> None: + reader = BufferReader(payload) + correlation_id = reader.read_bytes(16) + route = reader.read_route() + body = reader.read_bytes(reader.read_u32_be()) + registration = self._best_worker(route) + if registration is None: return - snapshot = list(self._workers.items()) - self._workers.clear() - for route, handler in snapshot: - await self.register_worker(route, handler) - - -def _is_benign_shutdown_error(error: Exception, connection) -> bool: - if connection.get_state() is not ConnectionState.AUTHENTICATED: - return True - if isinstance(error, ConnectionError): - return True - if not isinstance(error, TransportError): + request = InboundRequest(route, body) + response = ResponseWriter(self.connection, correlation_id) + self.connection.dispatch_async(lambda: registration.handler(request, response)) + + def _best_worker(self, route: str) -> _Registration | None: + matches = [ + (pattern, worker) + for pattern, worker in self._workers.items() + if _matches(route, pattern) + ] + if not matches: + return None + return max(matches, key=lambda pair: _specificity(pair[0]))[1] + + async def _restore(self) -> None: + for route, registration in list(self._workers.items()): + try: + await self._subscribe(route, registration) + except BaseException as exc: + self._workers.pop(route, None) + self.connection.report_restore_failure("rpc", route, exc) + + def _disconnect(self) -> None: + for call in tuple(self._pending.values()): + call.fail(FitzConnectionError("Connection closed while RPC response was pending")) + self._pending.clear() + + +def _route(route: str, *, patterns: bool) -> None: + if not route.startswith("rpc://"): + raise RpcError(f"Invalid RPC route: {route}", "INVALID_ROUTE") + parts = route[6:].split("/") + invalid = not parts or any(not part for part in parts) + invalid |= not patterns and any("*" in part for part in parts) + invalid |= any("*" in part and part not in {"*", "**"} for part in parts) + invalid |= "**" in parts and parts[-1] != "**" + if invalid: + raise RpcError(f"Invalid RPC route: {route}", "INVALID_ROUTE") + + +def _matches(route: str, pattern: str) -> bool: + route_parts, pattern_parts = route[6:].split("/"), pattern[6:].split("/") + for index, part in enumerate(pattern_parts): + if part == "**": + return True + if index >= len(route_parts) or part not in {"*", route_parts[index]}: + return False + return len(route_parts) == len(pattern_parts) + + +def _specificity(pattern: str) -> tuple[int, int, int, int]: + parts = pattern[6:].split("/") + return ( + sum(p not in {"*", "**"} for p in parts), + -parts.count("*"), + -parts.count("**"), + len(parts), + ) + + +def _looks_like_request(payload: bytes) -> bool: + try: + reader = BufferReader(payload) + reader.read_bytes(16) + reader.read_route() + reader.read_bytes(reader.read_u32_be()) + return reader.is_eof() + except BaseException: return False - lowered = str(error).lower() - return "closed" in lowered or "not connected" in lowered or "reset" in lowered -def _assert_rpc_route(route: str) -> None: - if not is_concrete_route_shape(route, "rpc"): - raise RpcError( - f"Invalid rpc route: {route} (expected rpc://{{realm}}/{{area}}/{{resource}} or any other concrete rpc route, no empty segments or wildcards)", - "INVALID_ROUTE", - ) +def _looks_like_response(payload: bytes) -> bool: + try: + reader = BufferReader(payload) + reader.read_bytes(16) + reader.read_u64_be() + flags = reader.read_u8() + reader.read_bytes(reader.read_u32_be()) + return flags & ~1 == 0 and reader.is_eof() + except BaseException: + return False diff --git a/src/fitz_py/domains/schedule.py b/src/fitz_py/domains/schedule.py index 5ca4f20..5c2412b 100644 --- a/src/fitz_py/domains/schedule.py +++ b/src/fitz_py/domains/schedule.py @@ -1,14 +1,15 @@ -"""Schedule domain client, schedule models, and schedule notifications.""" +"""Cron schedules, pagination, and delivery notifications.""" from __future__ import annotations -import asyncio -from collections.abc import Awaitable, Callable -from dataclasses import dataclass, field +from dataclasses import dataclass +from enum import StrEnum +from fitz_py._runtime import AsyncSubscription from fitz_py.domains._routes import is_exact_route_shape +from fitz_py.domains._subscriptions import SubscriptionRegistry from fitz_py.domains.base import DomainClient -from fitz_py.errors import ScheduleError, schedule_error +from fitz_py.errors import ScheduleError, domain_error from fitz_py.protocol.buffer import BufferReader, BufferWriter from fitz_py.protocol.messages import ( MSG_SCHEDULE_CANCEL, @@ -18,232 +19,167 @@ MSG_SCHEDULE_SUBSCRIBE, MSG_SCHEDULE_UNSUBSCRIBE, ) -from fitz_py.protocol.response import parse_standard_response +from fitz_py.protocol.response import parse_response -ScheduleHandler = Callable[["ScheduleNotification"], None | Awaitable[None]] - -@dataclass(slots=True) -class ScheduleNotification: - """Notification payload delivered when a schedule fires.""" - - payload: bytes +class DeliveryMode(StrEnum): + BROADCAST = "broadcast" + SINGLE = "single" -@dataclass(slots=True) +@dataclass(frozen=True, slots=True) class ScheduleEntry: - """Schedule record returned by list operations.""" - - id: str route: str cron: str + delivery_mode: DeliveryMode payload: bytes + @property + def id(self) -> str: + return self.route -@dataclass(slots=True) -class _ScheduleSubscriptionState: - """Internal handler registry for a subscribed schedule pattern.""" - - sub_id: int - handlers: dict[int, ScheduleHandler] = field(default_factory=dict) +@dataclass(frozen=True, slots=True) +class SchedulePage: + entries: tuple[ScheduleEntry, ...] + total_count: int -class ScheduleSubscription: - """Handle for an active schedule subscription.""" - def __init__( - self, - sub_id: int, - pattern: str, - handler: ScheduleHandler, - unsubscribe: Callable[[str, int], Awaitable[None]], - handler_id: int, - ) -> None: - self._sub_id = sub_id - self.pattern = pattern - self.handler = handler - self._unsubscribe = unsubscribe - self._handler_id = handler_id - - async def unsubscribe(self) -> None: - await self._unsubscribe(self.pattern, self._handler_id) +@dataclass(frozen=True, slots=True) +class ScheduleNotification: + route: str + payload: bytes class ScheduleClient(DomainClient): - """Schedule domain operations for create, cancel, list, and subscribe.""" - def __init__(self, connection) -> None: super().__init__(connection) - self._subscriptions_by_pattern: dict[str, _ScheduleSubscriptionState] = {} - self._patterns_by_sub_id: dict[int, str] = {} - self._initialized = False - self._next_handler_id = 1 - self.connection.on_reconnect(self._restore_subscriptions) - - async def create(self, route: str, cron: str, payload: bytes = b"") -> str: - _assert_schedule_route(route) + self._subscriptions = SubscriptionRegistry[ScheduleNotification]( + connection.config.limits.subscription_buffer_size, + self._subscribe_wire, + self._unsubscribe_wire, + ) + connection.register_notification_handler(MSG_SCHEDULE_NOTIFY, self._notify) + connection.on_reconnect( + lambda: self._subscriptions.restore( + domain="schedule", + on_error=lambda registration, error: connection.report_restore_failure( + "schedule", registration, error + ), + ), + domain="schedule", + registration="notifications", + ) + + async def create( + self, + route: str, + cron: str, + *, + delivery_mode: DeliveryMode = DeliveryMode.SINGLE, + payload: bytes = b"", + ) -> str: + _route(route) + if not cron: + raise ValueError("cron must not be empty") writer = BufferWriter() writer.write_route(route) writer.write_string(cron) + writer.write_u8(0 if delivery_mode is DeliveryMode.BROADCAST else 1) writer.write_u32_be(len(payload)) writer.write_bytes(payload) - data = self._assert_success( - await self.request_frame(MSG_SCHEDULE_CREATE, writer.build()), "CREATE" + response = parse_response( + await self.request_frame(MSG_SCHEDULE_CREATE, writer.build()), plain=True ) - reader = BufferReader(data) - if not reader.is_eof() and reader.read_u8() == 1: - return reader.read_string() - return route + if not response.success: + raise ScheduleError(f"CREATE failed: {response.error}", "CREATE") + reader = BufferReader(response.data) + if reader.is_eof(): + return route + if reader.read_u8() != 1: + raise ScheduleError("Invalid schedule id flag", "INVALID_RESPONSE") + schedule_id = reader.read_string() + if not reader.is_eof(): + raise ScheduleError("CREATE response has trailing bytes", "INVALID_RESPONSE") + return schedule_id async def cancel(self, route: str) -> None: - _assert_schedule_route(route) + _route(route) writer = BufferWriter() writer.write_route(route) - self._assert_success( - await self.request_frame(MSG_SCHEDULE_CANCEL, writer.build()), "CANCEL" + response = parse_response( + await self.request_frame(MSG_SCHEDULE_CANCEL, writer.build()), plain=True ) - - async def list( - self, *, offset: int | None = None, limit: int | None = None - ) -> list[ScheduleEntry]: + if not response.success: + raise ScheduleError(f"CANCEL failed: {response.error}", "CANCEL") + + async def list(self, *, offset: int | None = None, limit: int | None = None) -> SchedulePage: + if offset is not None and offset < 0: + raise ValueError("offset must be non-negative") + if limit is not None and not 0 <= limit <= 1000: + raise ValueError("limit must be between 0 and 1000") writer = BufferWriter() writer.write_optional_u64(offset) writer.write_optional_u64(limit) - data = self._assert_success( - await self.request_frame(MSG_SCHEDULE_LIST, writer.build()), "LIST" - ) - reader = BufferReader(data) - if reader.remaining_bytes() >= 8: - reader.read_u64_be() + response = parse_response(await self.request_frame(MSG_SCHEDULE_LIST, writer.build())) + if not response.success: + raise domain_error(ScheduleError, "LIST", response.error_code or 0, response.error) + reader = BufferReader(response.data) + total_count = reader.read_u64_be() entries: list[ScheduleEntry] = [] - while not reader.is_eof(): - if reader.read_u8() == 0: - break - route = reader.read_string() - cron = reader.read_string() + while reader.read_u8() == 1: + route, cron = reader.read_string(), reader.read_string() + mode_byte = reader.read_u8() + if mode_byte not in {0, 1}: + raise ScheduleError("Invalid delivery mode", "INVALID_RESPONSE") payload = reader.read_bytes(reader.read_u32_be()) - entries.append(ScheduleEntry(id=route, route=route, cron=cron, payload=payload)) - return entries - - async def subscribe(self, pattern: str, handler: ScheduleHandler) -> ScheduleSubscription: - _assert_schedule_route(pattern) - self._init_notify_handler() - - existing = self._subscriptions_by_pattern.get(pattern) - if existing is None: - sub_id = await self._subscribe_wire(pattern) - existing = _ScheduleSubscriptionState(sub_id=sub_id) - self._subscriptions_by_pattern[pattern] = existing - self._patterns_by_sub_id[sub_id] = pattern - - handler_id = self._next_handler_id - self._next_handler_id += 1 - existing.handlers[handler_id] = handler - return ScheduleSubscription( - existing.sub_id, pattern, handler, self._unsubscribe, handler_id - ) - - async def _subscribe_wire(self, pattern: str) -> int: + entries.append( + ScheduleEntry( + route, + cron, + DeliveryMode.BROADCAST if mode_byte == 0 else DeliveryMode.SINGLE, + payload, + ) + ) + if not reader.is_eof(): + raise ScheduleError("LIST response has trailing bytes", "INVALID_RESPONSE") + return SchedulePage(tuple(entries), total_count) + + async def subscribe(self, route: str) -> AsyncSubscription[ScheduleNotification]: + _route(route) + return await self._subscriptions.subscribe(route) + + async def _subscribe_wire(self, route: str) -> int: writer = BufferWriter() - writer.write_string(pattern) - data = self._assert_success( - await self.request_frame(MSG_SCHEDULE_SUBSCRIBE, writer.build()), "SUBSCRIBE" - ) - reader = BufferReader(data) - has_sub_id = reader.read_u8() if not reader.is_eof() else 0 - if has_sub_id != 1 or reader.remaining_bytes() < 8: - raise ScheduleError("SUBSCRIBE response missing subscription id", "MISSING_SUB_ID") + writer.write_route(route) + response = parse_response(await self.request_frame(MSG_SCHEDULE_SUBSCRIBE, writer.build())) + if not response.success: + raise domain_error(ScheduleError, "SUBSCRIBE", response.error_code or 0, response.error) + reader = BufferReader(response.data) + if reader.read_u8() != 1: + raise ScheduleError("SUBSCRIBE response omitted its id", "INVALID_RESPONSE") return reader.read_u64_be() - async def _unsubscribe(self, pattern: str, handler_id: int) -> None: - subscription = self._subscriptions_by_pattern.get(pattern) - if subscription is None: - return - subscription.handlers.pop(handler_id, None) - if subscription.handlers: - return - - self._subscriptions_by_pattern.pop(pattern, None) - self._patterns_by_sub_id.pop(subscription.sub_id, None) + async def _unsubscribe_wire(self, route: str) -> None: writer = BufferWriter() - writer.write_string(pattern) - self._assert_success( - await self.request_frame(MSG_SCHEDULE_UNSUBSCRIBE, writer.build()), "UNSUBSCRIBE" + writer.write_route(route) + response = parse_response( + await self.request_frame(MSG_SCHEDULE_UNSUBSCRIBE, writer.build()), plain=True ) - - def _init_notify_handler(self) -> None: - if self._initialized: - return - self._initialized = True - - def handler(payload: bytes) -> None: - try: - sub_id, actual_payload = _decode_schedule_notification(payload) - pattern = self._patterns_by_sub_id.get(sub_id) - if pattern is None: - return - subscription = self._subscriptions_by_pattern.get(pattern) - if subscription is None: - return - notification = ScheduleNotification(payload=actual_payload) - for callback in subscription.handlers.values(): - result = callback(notification) - if asyncio.iscoroutine(result): - asyncio.create_task(result) - except Exception: - return - - self.connection.register_notification_handler(MSG_SCHEDULE_NOTIFY, handler) - - async def _restore_subscriptions(self) -> None: - if not self._subscriptions_by_pattern: - return - snapshot = [ - (pattern, list(state.handlers.values())) - for pattern, state in self._subscriptions_by_pattern.items() - ] - self._subscriptions_by_pattern.clear() - self._patterns_by_sub_id.clear() - for pattern, handlers in snapshot: - for handler in handlers: - await self.subscribe(pattern, handler) - - @staticmethod - def _assert_success(payload: bytes, operation: str) -> bytes: - result = parse_standard_response(payload) - if result.success: - return result.data - error_message = result.error or f"{operation} failed" - raise _map_schedule_protocol_error(f"{operation} failed: {error_message}") - - -def _assert_schedule_route(route: str) -> None: - if not is_exact_route_shape(route, "schedule", 4): - raise ScheduleError( - f"Invalid schedule route: {route} (expected schedule://{{realm}}/{{area}}/{{resource}}/{{operation}}, no empty segments or wildcards)", - "INVALID_ROUTE", + if not response.success: + raise ScheduleError(f"UNSUBSCRIBE failed: {response.error}", "UNSUBSCRIBE") + + def _notify(self, payload: bytes) -> None: + reader = BufferReader(payload) + sub_id = reader.read_u64_be() + notification = ScheduleNotification( + reader.read_route(), reader.read_bytes(reader.read_u32_be()) ) + if not reader.is_eof(): + raise ScheduleError("SCHEDULE_NOTIFY has trailing bytes", "INVALID_RESPONSE") + self._subscriptions.publish(sub_id, notification) -def _decode_schedule_notification(payload: bytes) -> tuple[int, bytes]: - reader = BufferReader(payload) - sub_id = reader.read_u64_be() - body = reader.read_bytes(reader.read_u32_be()) - return sub_id, body - - -def _map_schedule_protocol_error(message: str) -> ScheduleError: - normalized = message.lower() - if "not found" in normalized: - return schedule_error(message, 1) - if "task" in normalized and "not found" in normalized: - return schedule_error(message, 2) - if "cron" in normalized: - return schedule_error(message, 3) - if "delay" in normalized: - return schedule_error(message, 4) - if "timestamp" in normalized or "time" in normalized: - return schedule_error(message, 5) - if "invalid route" in normalized or "must be schedule://" in normalized: - return ScheduleError(message, "INVALID_ROUTE") - return ScheduleError(message, "ERROR") +def _route(route: str) -> None: + if not is_exact_route_shape(route, "schedule", 4): + raise ScheduleError(f"Invalid schedule route: {route}", "INVALID_ROUTE") diff --git a/src/fitz_py/domains/stream.py b/src/fitz_py/domains/stream.py index 8c55c32..03e4cbf 100644 --- a/src/fitz_py/domains/stream.py +++ b/src/fitz_py/domains/stream.py @@ -2,16 +2,15 @@ from __future__ import annotations -import asyncio import json -import struct -from collections.abc import Awaitable, Callable from dataclasses import dataclass, field from enum import Enum, IntEnum +from fitz_py._runtime import AsyncSubscription from fitz_py.domains._routes import is_exact_route_shape, is_selector_route_shape +from fitz_py.domains._subscriptions import SubscriptionRegistry from fitz_py.domains.base import DomainClient -from fitz_py.errors import ErrStreamSessionClosed, StreamError, stream_error +from fitz_py.errors import StaleHandleError, StreamError, domain_error from fitz_py.protocol.buffer import BufferReader, BufferWriter from fitz_py.protocol.messages import ( MSG_STREAM_APPEND, @@ -25,17 +24,18 @@ MSG_STREAM_SUBSCRIBE, MSG_STREAM_UNSUBSCRIBE, ) - -StreamHandler = Callable[["StreamCommitNotification"], None | Awaitable[None]] +from fitz_py.protocol.response import parse_response @dataclass(slots=True) class StreamRecord: """Single stream event record with offsets and payload data.""" + route: str offset: int area_offset: int | None = None realm_offset: int | None = None + global_offset: int | None = None body: bytes = b"" metadata: bytes | None = None timestamp: int = 0 @@ -48,6 +48,11 @@ class StreamMetadata: first_offset: int last_offset: int record_count: int + max_batch_events: int | None = None + max_batch_bytes: int | None = None + ttl_seconds: int | None = None + area_watermark: int | None = None + realm_watermark: int | None = None @dataclass(slots=True) @@ -89,6 +94,9 @@ class StreamReadCursor: last_resource_offset: int = 0 last_area_offset: int | None = None last_realm_offset: int | None = None + last_global_offset: int | None = None + cursor_fingerprint: int | None = None + captured_watermark: int | None = None has_more: bool = False @@ -97,6 +105,7 @@ class StreamReadItem: """One item in a stream read page, event or filtered marker.""" kind: StreamReadItemKind + route: str record: StreamRecord | None = None offset: int = 0 from_offset: int = 0 @@ -134,20 +143,6 @@ class StreamCommitNotification: batch_size: int = 0 -class StreamSubscription: - """Handle for an active stream commit subscription.""" - - def __init__( - self, sub_id: int, pattern: str, unsubscribe: Callable[[str], Awaitable[None]] - ) -> None: - self._sub_id = sub_id - self.pattern = pattern - self._unsubscribe = unsubscribe - - async def unsubscribe(self) -> None: - await self._unsubscribe(self.pattern) - - class StreamSession: """Mutable append session returned by stream begin operations.""" @@ -156,6 +151,7 @@ def __init__(self, connection, session_id: int) -> None: self._session_id = session_id self._closed = False self._closed_reason: str | None = None + self._generation = connection.generation on_disconnect = getattr(self._connection, "on_disconnect", None) self._disconnect_unregister = ( on_disconnect(self._invalidate) if callable(on_disconnect) else None @@ -192,20 +188,18 @@ async def append( writer.write_string(discriminator) else: writer.write_u8(0) - reader = BufferReader(await self._connection.request(MSG_STREAM_APPEND, writer.build())) - status = reader.read_u8() - if status != 0: - raise stream_error(f"APPEND failed with status {status}", status) - if not reader.is_eof(): - has_session = reader.read_u8() - if has_session == 1 and reader.remaining_bytes() >= 8: - reader.read_u64_be() - if reader.is_eof(): - return None - data = reader.read_bytes(reader.read_u32_be()) - if len(data) < 8: + response = parse_response( + await self._connection.request(MSG_STREAM_APPEND, writer.build()), plain=True + ) + if not response.success: + raise StreamError(f"APPEND failed: {response.error}", "APPEND") + if not response.data: return None - return BufferReader(data).read_u64_be() + reader = BufferReader(response.data) + length = reader.read_u32_be() + if length != 8 or reader.remaining_bytes() != 8: + raise StreamError("APPEND response has invalid payload length", "INVALID_RESPONSE") + return reader.read_u64_be() async def commit(self, mode: int | StreamCommitMode = StreamCommitMode.BUFFERED) -> None: self._ensure_open("COMMIT") @@ -227,17 +221,16 @@ async def rollback(self) -> None: self._clear_disconnect_listener() async def _expect_status(self, message_type: int, payload: bytes, operation: str) -> None: - reader = BufferReader(await self._connection.request(message_type, payload)) - status = reader.read_u8() if not reader.is_eof() else 0 - if status != 0: - raise stream_error(f"{operation} failed with status {status}", status) + response = parse_response(await self._connection.request(message_type, payload)) + if not response.success: + raise domain_error(StreamError, operation, response.error_code or 0, response.error) def _ensure_open(self, operation: str) -> None: - if not self._closed: + if not self._closed and self._generation == self._connection.generation: return reason = self._closed_reason or "closed" - raise ErrStreamSessionClosed(f"{operation} not allowed: session already {reason}") + raise StaleHandleError(f"Stream session ({operation}, {reason})") def _invalidate(self) -> None: if self._closed: @@ -259,9 +252,22 @@ class StreamClient(DomainClient): def __init__(self, connection) -> None: super().__init__(connection) - self._subscriptions: dict[int, tuple[str, StreamHandler]] = {} - self._initialized = False - self.connection.on_reconnect(self._restore_subscriptions) + self._subscriptions = SubscriptionRegistry[StreamCommitNotification]( + connection.config.limits.subscription_buffer_size, + self._subscribe_wire, + self._unsubscribe_wire, + ) + connection.register_notification_handler(MSG_STREAM_NOTIFY, self._notify) + connection.on_reconnect( + lambda: self._subscriptions.restore( + domain="stream", + on_error=lambda registration, error: connection.report_restore_failure( + "stream", registration, error + ), + ), + domain="stream", + registration="commits", + ) async def begin(self, route: str, ingest_metadata: bytes | None = None) -> StreamSession: _assert_stream_route(route) @@ -273,12 +279,13 @@ async def begin(self, route: str, ingest_metadata: bytes | None = None) -> Strea writer.write_bytes(ingest_metadata) else: writer.write_u8(0) - reader = BufferReader(await self.request_frame(MSG_STREAM_BEGIN, writer.build())) - status = reader.read_u8() - if status != 0: - raise stream_error(f"BEGIN failed with status {status}", status) - has_session = reader.read_u8() if not reader.is_eof() else 0 - if has_session != 1 or reader.remaining_bytes() < 8: + response = parse_response( + await self.request_frame(MSG_STREAM_BEGIN, writer.build()), plain=True + ) + if not response.success: + raise StreamError(f"BEGIN failed: {response.error}", "BEGIN") + reader = BufferReader(response.data) + if reader.remaining_bytes() < 8: raise StreamError("BEGIN response missing session id", "MISSING_SESSION_ID") return StreamSession(self.connection, reader.read_u64_be()) @@ -290,6 +297,8 @@ async def read( stream_filter: StreamFilterSet | None = None, *, max_bytes: int | None = None, + cursor_fingerprint: int | None = None, + captured_watermark: int | None = None, ) -> list[StreamRecord]: page = await self.read_page( route, @@ -297,6 +306,8 @@ async def read( limit=limit, stream_filter=stream_filter, max_bytes=max_bytes, + cursor_fingerprint=cursor_fingerprint, + captured_watermark=captured_watermark, ) return _flatten_stream_read_items(page.items) @@ -308,6 +319,8 @@ async def read_page( stream_filter: StreamFilterSet | None = None, *, max_bytes: int | None = None, + cursor_fingerprint: int | None = None, + captured_watermark: int | None = None, ) -> StreamReadPage: _assert_stream_pattern(route) writer = BufferWriter() @@ -324,97 +337,82 @@ async def read_page( if filter_bytes: writer.write_u32_be(len(filter_bytes)) writer.write_bytes(filter_bytes) - reader = BufferReader(await self.request_frame(MSG_STREAM_READ, writer.build())) - status, data = _read_wrapped_stream_response(reader) - if status != 0: - raise stream_error(f"READ failed with status {status}", status) - if not data: + writer.write_optional_u64(cursor_fingerprint) + writer.write_optional_u64(captured_watermark) + response = parse_response(await self.request_frame(MSG_STREAM_READ, writer.build())) + if not response.success: + raise domain_error(StreamError, "READ", response.error_code or 0, response.error) + if not response.data: return StreamReadPage() - return _read_stream_read_page(data) + envelope = BufferReader(response.data) + if envelope.read_u8() != 0: + raise StreamError("READ response has invalid session flag", "INVALID_RESPONSE") + data = envelope.read_bytes(envelope.read_u32_be()) + if not envelope.is_eof(): + raise StreamError("READ response has trailing bytes", "INVALID_RESPONSE") + return _read_stream_read_page(data, route) async def peek(self, route: str) -> StreamRecord | None: _assert_stream_route(route) writer = BufferWriter() writer.write_route(route) - reader = BufferReader(await self.request_frame(MSG_STREAM_LAST, writer.build())) - status, data = _read_wrapped_stream_response(reader) - if status != 0: - raise stream_error(f"LAST failed with status {status}", status) - if not data: + response = parse_response(await self.request_frame(MSG_STREAM_LAST, writer.build())) + if not response.success: + raise domain_error(StreamError, "LAST", response.error_code or 0, response.error) + if not response.data: return None - inner = BufferReader(data) - return _read_stream_record(inner) + inner = BufferReader(response.data) + return _read_stream_record(inner, route, False) async def metadata(self, route: str) -> StreamMetadata: _assert_stream_route(route) writer = BufferWriter() writer.write_route(route) - reader = BufferReader(await self.request_frame(MSG_STREAM_GET_METADATA, writer.build())) - status, data = _read_wrapped_stream_response(reader) - if status != 0: - raise stream_error(f"METADATA failed with status {status}", status) - if not data: + response = parse_response(await self.request_frame(MSG_STREAM_GET_METADATA, writer.build())) + if not response.success: + raise domain_error(StreamError, "METADATA", response.error_code or 0, response.error) + if not response.data: return StreamMetadata(first_offset=0, last_offset=0, record_count=0) - inner = BufferReader(data) + inner = BufferReader(response.data) return StreamMetadata( - first_offset=inner.read_u64_be(), - last_offset=inner.read_u64_be(), + first_offset=_read_optional_u64(inner) or 0, + last_offset=_read_optional_u64(inner) or 0, record_count=inner.read_u64_be(), + max_batch_events=inner.read_u64_be(), + max_batch_bytes=inner.read_u64_be(), + ttl_seconds=_read_optional_u64(inner), + area_watermark=inner.read_u64_be(), + realm_watermark=inner.read_u64_be(), ) - async def subscribe(self, pattern: str, handler: StreamHandler) -> StreamSubscription: + async def subscribe(self, pattern: str) -> AsyncSubscription[StreamCommitNotification]: _assert_stream_pattern(pattern) - self._init_notify_handler() + return await self._subscriptions.subscribe(pattern) + + async def _subscribe_wire(self, pattern: str) -> int: writer = BufferWriter() writer.write_route(pattern) - reader = BufferReader(await self.request_frame(MSG_STREAM_SUBSCRIBE, writer.build())) - status = reader.read_u8() - if status != 0: - raise stream_error(f"SUBSCRIBE failed with status {status}", status) - has_sub_id = reader.read_u8() if not reader.is_eof() else 0 + response = parse_response(await self.request_frame(MSG_STREAM_SUBSCRIBE, writer.build())) + if not response.success: + raise domain_error(StreamError, "SUBSCRIBE", response.error_code or 0, response.error) + reader = BufferReader(response.data) + has_sub_id = reader.read_u8() if has_sub_id != 1 or reader.is_eof(): raise StreamError("SUBSCRIBE response missing subscription id", "MISSING_SUB_ID") - sub_id = reader.read_u64_be() - self._subscriptions[sub_id] = (pattern, handler) - return StreamSubscription(sub_id, pattern, self._unsubscribe) - - async def _unsubscribe(self, pattern: str) -> None: - for sub_id, (sub_pattern, _) in list(self._subscriptions.items()): - if sub_pattern == pattern: - self._subscriptions.pop(sub_id, None) + return reader.read_u64_be() + + async def _unsubscribe_wire(self, pattern: str) -> None: writer = BufferWriter() writer.write_route(pattern) - await self.request_frame(MSG_STREAM_UNSUBSCRIBE, writer.build()) + response = parse_response(await self.request_frame(MSG_STREAM_UNSUBSCRIBE, writer.build())) + if not response.success: + raise domain_error(StreamError, "UNSUBSCRIBE", response.error_code or 0, response.error) - def _init_notify_handler(self) -> None: - if self._initialized: - return - self._initialized = True - - def handler(payload: bytes) -> None: - try: - reader = BufferReader(payload) - sub_id = reader.read_u64_be() - route = reader.read_route() - body = reader.read_bytes(reader.read_u32_be()) - subscription = self._subscriptions.get(sub_id) - if subscription is None: - return - result = subscription[1](_decode_stream_commit_notification(route, body)) - if asyncio.iscoroutine(result): - asyncio.create_task(result) - except Exception: - return - - self.connection.register_notification_handler(MSG_STREAM_NOTIFY, handler) - - async def _restore_subscriptions(self) -> None: - if not self._subscriptions: - return - snapshot = list(self._subscriptions.values()) - self._subscriptions.clear() - for pattern, handler in snapshot: - await self.subscribe(pattern, handler) + def _notify(self, payload: bytes) -> None: + reader = BufferReader(payload) + sub_id, route = reader.read_u64_be(), reader.read_route() + body = reader.read_bytes(reader.read_u32_be()) + self._subscriptions.publish(sub_id, _decode_stream_commit_notification(route, body)) def _read_wrapped_stream_response(reader: BufferReader) -> tuple[int, bytes]: @@ -442,49 +440,68 @@ def _read_optional_bytes(reader: BufferReader) -> bytes | None: return reader.read_bytes(reader.read_u32_be()) -def _read_stream_record(reader: BufferReader) -> StreamRecord: +def _read_stream_record(reader: BufferReader, route: str, extended: bool) -> StreamRecord: offset = reader.read_u64_be() area_offset = _read_optional_u64(reader) realm_offset = _read_optional_u64(reader) + global_offset = _read_optional_u64(reader) if extended else None body = reader.read_bytes(reader.read_u32_be()) metadata = _read_optional_bytes(reader) timestamp = reader.read_u64_be() return StreamRecord( + route=route, offset=offset, area_offset=area_offset, realm_offset=realm_offset, + global_offset=global_offset, body=body, metadata=metadata, timestamp=timestamp, ) -def _read_stream_read_page(data: bytes) -> StreamReadPage: +def _read_stream_read_page(data: bytes, selector: str) -> StreamReadPage: reader = BufferReader(data) + extended = selector in {"stream://**", "stream://*/*/*"} count = reader.read_u32_be() - items = [_read_stream_read_item(reader) for _ in range(count)] + items = [] + for _ in range(count): + route = reader.read_route() + _assert_stream_route(route) + items.append(_read_stream_read_item(reader, route, extended)) cursor = StreamReadCursor( last_resource_offset=reader.read_u64_be(), last_area_offset=_read_optional_u64(reader), last_realm_offset=_read_optional_u64(reader), + last_global_offset=_read_optional_u64(reader) if extended else None, has_more=_read_bool_u8(reader), + cursor_fingerprint=_read_optional_u64(reader) if extended else None, + captured_watermark=_read_optional_u64(reader) if extended else None, ) + if not reader.is_eof(): + raise StreamError("READ response has trailing bytes", "INVALID_RESPONSE") return StreamReadPage(items=items, cursor=cursor) -def _read_stream_read_item(reader: BufferReader) -> StreamReadItem: +def _read_stream_read_item(reader: BufferReader, route: str, extended: bool) -> StreamReadItem: tag = reader.read_u8() if tag == 0: - return StreamReadItem(kind=StreamReadItemKind.EVENT, record=_read_stream_record(reader)) + return StreamReadItem( + kind=StreamReadItemKind.EVENT, + route=route, + record=_read_stream_record(reader, route, extended), + ) if tag == 1: return StreamReadItem( kind=StreamReadItemKind.FILTERED, + route=route, offset=reader.read_u64_be(), reason=_read_filtered_reason(reader), ) if tag == 2: return StreamReadItem( kind=StreamReadItemKind.FILTERED_RANGE, + route=route, from_offset=reader.read_u64_be(), to_offset=reader.read_u64_be(), reason=_read_filtered_reason(reader), @@ -570,33 +587,29 @@ def _encode_stream_filter_set(stream_filter: StreamFilterSet | None) -> bytes: if stream_filter is None or not stream_filter.clauses: return b"" - buffer = bytearray() - buffer.extend(struct.pack(" bytes: - buffer = bytearray() +def _encode_stream_filter_clause(writer: BufferWriter, clause: StreamFilterClause) -> None: if clause.kind == "Equals": - buffer.extend(struct.pack(" None: - encoded = value.encode("utf-8") - buffer.extend(struct.pack(" None: super().__init__(message) self.code = code self.domain_code = domain_code + self.context = MappingProxyType(dict(context or {})) -TError = TypeVar("TError", bound=FitzError) +class FitzTransportError(FitzError): + def __init__(self, message: str, context: Mapping[str, Any] | None = None) -> None: + super().__init__(message, "TRANSPORT_ERROR", context=context) -class TransportError(FitzError): - def __init__(self, message: str) -> None: - super().__init__(message, "TRANSPORT_ERROR") - - -class ConnectionError(FitzError): - def __init__(self, message: str) -> None: - super().__init__(message, "CONNECTION_ERROR") +class FitzConnectionError(FitzError): + def __init__(self, message: str, context: Mapping[str, Any] | None = None) -> None: + super().__init__(message, "CONNECTION_ERROR", context=context) class AuthenticationError(FitzError): - def __init__(self, message: str) -> None: - super().__init__(message, "AUTH_ERROR") + def __init__(self, message: str, context: Mapping[str, Any] | None = None) -> None: + super().__init__(message, "AUTH_ERROR", context=context) -class TimeoutError(FitzError): - def __init__(self, message: str) -> None: - super().__init__(message, "TIMEOUT") +class FitzTimeoutError(FitzError): + def __init__(self, message: str, context: Mapping[str, Any] | None = None) -> None: + super().__init__(message, "TIMEOUT", context=context) class ProtocolError(FitzError): @@ -51,337 +51,140 @@ def __init__(self, message: str) -> None: super().__init__(message, "CODEC_ERROR") -class KvError(FitzError): - def __init__(self, message: str, code: str, domain_code: int | None = None) -> None: - super().__init__(message, f"KV_{code}", domain_code) - - -class ErrKvTransactionAborted(KvError): - def __init__(self, message: str = "KV transaction aborted") -> None: - super().__init__(message, "TRANSACTION_ABORTED", 1) - - -class ErrKvLeaseExpired(KvError): - def __init__(self, message: str = "KV lease expired") -> None: - super().__init__(message, "LEASE_EXPIRED", 2) - - -class ErrKvConflictingWrite(KvError): - def __init__(self, message: str = "KV conflicting write") -> None: - super().__init__(message, "CONFLICTING_WRITE", 3) - - -class ErrKvKeyNotFound(KvError): - def __init__(self, message: str = "KV key not found") -> None: - super().__init__(message, "KEY_NOT_FOUND", 4) - - -class ErrKvOperationNotAllowed(KvError): - def __init__(self, message: str = "KV operation not allowed") -> None: - super().__init__(message, "OPERATION_NOT_ALLOWED", 5) - - -class QueueError(FitzError): - def __init__(self, message: str, code: str, domain_code: int | None = None) -> None: - super().__init__(message, f"QUEUE_{code}", domain_code) - - -class ErrQueueNotFound(QueueError): - def __init__(self, message: str = "Queue not found") -> None: - super().__init__(message, "NOT_FOUND", 1) - - -class ErrQueueMessageNotFound(QueueError): - def __init__(self, message: str = "Queue message not found") -> None: - super().__init__(message, "MESSAGE_NOT_FOUND", 2) - - -class ErrQueueInvalidToken(QueueError): - def __init__(self, message: str = "Queue invalid token") -> None: - super().__init__(message, "INVALID_TOKEN", 3) - - -class ErrQueueFull(QueueError): - def __init__(self, message: str = "Queue full") -> None: - super().__init__(message, "FULL", 4) - - -class ErrQueueInvalidDelay(QueueError): - def __init__(self, message: str = "Queue invalid delay") -> None: - super().__init__(message, "INVALID_DELAY", 5) - - -class NoticeError(FitzError): - def __init__(self, message: str, code: str, domain_code: int | None = None) -> None: - super().__init__(message, f"NOTICE_{code}", domain_code) - - -class ErrNoticeGeneral(NoticeError): - def __init__(self, message: str = "Notice error") -> None: - super().__init__(message, "GENERAL", 1) - - -class RpcError(FitzError): - def __init__(self, message: str, code: str, domain_code: int | None = None) -> None: - super().__init__(message, f"RPC_{code}", domain_code) - +class RequestQueueFullError(FitzError): + def __init__(self) -> None: + super().__init__("Request queue is full", "REQUEST_QUEUE_FULL") -class ErrRpcTimeout(RpcError): - def __init__(self, message: str = "RPC timeout") -> None: - super().__init__(message, "TIMEOUT", 1) +class SubscriptionBackpressureError(FitzError): + def __init__(self) -> None: + super().__init__("Subscription consumer fell behind", "SUBSCRIPTION_BACKPRESSURE") -class ErrRpcHandlerNotFound(RpcError): - def __init__(self, message: str = "RPC handler not found") -> None: - super().__init__(message, "HANDLER_NOT_FOUND", 2) +class StaleHandleError(FitzError): + def __init__(self, kind: str) -> None: + super().__init__(f"{kind} is stale after disconnect", "STALE_HANDLE") -class ErrRpcHandlerError(RpcError): - def __init__(self, message: str = "RPC handler error") -> None: - super().__init__(message, "HANDLER_ERROR", 3) +class ReconnectRestoreError(FitzError): + def __init__(self, domain: str, registration: str, cause: BaseException) -> None: + super().__init__( + f"Failed to restore {domain} registration {registration}: {cause}", + "RECONNECT_RESTORE_FAILED", + context={"domain": domain, "registration": registration}, + ) -class ErrRpcInvalidRequest(RpcError): - def __init__(self, message: str = "RPC invalid request") -> None: - super().__init__(message, "INVALID_REQUEST", 4) +class LeaseLostError(FitzError): + def __init__(self, message: str = "Lease ownership was lost") -> None: + super().__init__(message, "LEASE_LOST") -class LeaseError(FitzError): - def __init__(self, message: str, code: str, domain_code: int | None = None) -> None: - super().__init__(message, f"LEASE_{code}", domain_code) +class LeaseLifecycleError(FitzError): + def __init__(self, causes: list[BaseException]) -> None: + super().__init__( + "Multiple failures occurred while managing a lease", + "LEASE_LIFECYCLE_MULTIPLE_FAILURES", + context={"causes": tuple(causes)}, + ) + self.causes = tuple(causes) -class ErrLeaseHeld(LeaseError): - def __init__(self, message: str = "Lease is already held") -> None: - super().__init__(message, "HELD", 1) +class DomainError(FitzError): + prefix = "DOMAIN" -class ErrLeaseNotFound(LeaseError): - def __init__(self, message: str = "Lease not found") -> None: - super().__init__(message, "NOT_FOUND", 2) + def __init__(self, message: str, reason: str, domain_code: int | None = None) -> None: + super().__init__(message, f"{self.prefix}_{reason}", domain_code) + self.reason = reason -class ErrLeaseInvalidToken(LeaseError): - def __init__(self, message: str = "Lease invalid token") -> None: - super().__init__(message, "INVALID_TOKEN", 3) +class KvError(DomainError): + prefix = "KV" -class StreamError(FitzError): - def __init__(self, message: str, code: str, domain_code: int | None = None) -> None: - super().__init__(message, f"STREAM_{code}", domain_code) +class QueueError(DomainError): + prefix = "QUEUE" -class ErrStreamNotFound(StreamError): - def __init__(self, message: str = "Stream not found") -> None: - super().__init__(message, "NOT_FOUND", 1) +class NoticeError(DomainError): + prefix = "NOTICE" -class ErrStreamOffsetOutOfRange(StreamError): - def __init__(self, message: str = "Stream offset out of range") -> None: - super().__init__(message, "OFFSET_OUT_OF_RANGE", 2) +class RpcError(DomainError): + prefix = "RPC" -class ErrStreamInvalidOffset(StreamError): - def __init__(self, message: str = "Stream invalid offset") -> None: - super().__init__(message, "INVALID_OFFSET", 3) +class LeaseError(DomainError): + prefix = "LEASE" -class ErrStreamFull(StreamError): - def __init__(self, message: str = "Stream full") -> None: - super().__init__(message, "FULL", 4) +class StreamError(DomainError): + prefix = "STREAM" -class ErrStreamSessionNotFound(StreamError): - def __init__(self, message: str = "Stream session not found") -> None: - super().__init__(message, "SESSION_NOT_FOUND", 5) +class ScheduleError(DomainError): + prefix = "SCHEDULE" -class ErrStreamSessionClosed(StreamError): - def __init__(self, message: str = "Stream session closed") -> None: - super().__init__(message, "SESSION_CLOSED", 6) - - -class ErrStreamExpectedOffsetMismatch(StreamError): - def __init__(self, message: str = "Stream expected offset mismatch") -> None: - super().__init__(message, "EXPECTED_OFFSET_MISMATCH", 7) - - -class ScheduleError(FitzError): - def __init__(self, message: str, code: str, domain_code: int | None = None) -> None: - super().__init__(message, f"SCHEDULE_{code}", domain_code) - - -class ErrScheduleNotFound(ScheduleError): - def __init__(self, message: str = "Schedule not found") -> None: - super().__init__(message, "NOT_FOUND", 1) - - -class ErrScheduleTaskNotFound(ScheduleError): - def __init__(self, message: str = "Schedule task not found") -> None: - super().__init__(message, "TASK_NOT_FOUND", 2) - - -class ErrScheduleInvalidCron(ScheduleError): - def __init__(self, message: str = "Schedule invalid cron") -> None: - super().__init__(message, "INVALID_CRON", 3) - - -class ErrScheduleInvalidDelay(ScheduleError): - def __init__(self, message: str = "Schedule invalid delay") -> None: - super().__init__(message, "INVALID_DELAY", 4) - - -class ErrScheduleInvalidTimestamp(ScheduleError): - def __init__(self, message: str = "Schedule invalid timestamp") -> None: - super().__init__(message, "INVALID_TIMESTAMP", 5) - - -_KV_STATUS_MAP: dict[int, type[KvError]] = { - 1: ErrKvTransactionAborted, - 2: ErrKvLeaseExpired, - 3: ErrKvConflictingWrite, - 4: ErrKvKeyNotFound, - 5: ErrKvOperationNotAllowed, +_RETRYABLE = { + ("KV", 3), + ("KV", 1004), + ("KV", 1009), + ("QUEUE", 4), + ("QUEUE", 4005), + ("LEASE", 1), + ("LEASE", 5001), + ("RPC", 6001), + ("RPC", 6002), + ("RPC", 6003), + ("RPC", 6004), } -_QUEUE_STATUS_MAP: dict[int, type[QueueError]] = { - 1: ErrQueueNotFound, - 2: ErrQueueMessageNotFound, - 3: ErrQueueInvalidToken, - 4: ErrQueueFull, - 5: ErrQueueInvalidDelay, -} -_RPC_STATUS_MAP: dict[int, type[RpcError]] = { - 1: ErrRpcTimeout, - 2: ErrRpcHandlerNotFound, - 3: ErrRpcHandlerError, - 4: ErrRpcInvalidRequest, -} +def is_retryable(error: object) -> bool: + if isinstance(error, (FitzTransportError, FitzTimeoutError, RequestQueueFullError)): + return True + if not isinstance(error, DomainError) or error.domain_code is None: + return False + return (error.prefix, error.domain_code) in _RETRYABLE -_LEASE_STATUS_MAP: dict[int, type[LeaseError]] = { - 1: ErrLeaseHeld, - 2: ErrLeaseNotFound, - 3: ErrLeaseInvalidToken, -} -_NOTICE_STATUS_MAP: dict[int, type[NoticeError]] = { - 1: ErrNoticeGeneral, -} +def domain_error( + error_type: type[DomainError], operation: str, status: int, message: str | None = None +) -> DomainError: + reason = message or f"status {status}" + return error_type(f"{operation} failed: {reason}", operation, status) -_STREAM_STATUS_MAP: dict[int, type[StreamError]] = { - 1: ErrStreamNotFound, - 2: ErrStreamOffsetOutOfRange, - 3: ErrStreamInvalidOffset, - 4: ErrStreamFull, - 5: ErrStreamSessionNotFound, - 6: ErrStreamSessionClosed, - 7: ErrStreamExpectedOffsetMismatch, -} - -_SCHEDULE_STATUS_MAP: dict[int, type[ScheduleError]] = { - 1: ErrScheduleNotFound, - 2: ErrScheduleTaskNotFound, - 3: ErrScheduleInvalidCron, - 4: ErrScheduleInvalidDelay, - 5: ErrScheduleInvalidTimestamp, -} -_RETRYABLE_ERROR_CODES = { - "KV_4", - "QUEUE_4", - "LEASE_1", - "NOTICE_1", - "STREAM_1", - "STREAM_2", - "STREAM_3", - "STREAM_4", - "RPC_1", -} - - -def _build_domain_error( - mapping: dict[int, type[TError]], - fallback: Callable[[str, int | None], TError], - message: str, - domain_code: int | None, -) -> TError: - if domain_code is not None: - error_type = mapping.get(domain_code) - if error_type is not None: - return error_type(message) - return fallback(message, domain_code) +# Transitional internal aliases. They are deliberately not exported publicly. +ConnectionError = FitzConnectionError +TransportError = FitzTransportError +TimeoutError = FitzTimeoutError def kv_error(message: str, domain_code: int | None = None) -> KvError: - return _build_domain_error( - _KV_STATUS_MAP, - lambda msg, code: KvError(msg, "ERROR", code), - message, - domain_code, - ) + return KvError(message, "ERROR", domain_code) def queue_error(message: str, domain_code: int | None = None) -> QueueError: - return _build_domain_error( - _QUEUE_STATUS_MAP, - lambda msg, code: QueueError(msg, "ERROR", code), - message, - domain_code, - ) + return QueueError(message, "ERROR", domain_code) def rpc_error(message: str, domain_code: int | None = None) -> RpcError: - return _build_domain_error( - _RPC_STATUS_MAP, - lambda msg, code: RpcError(msg, "ERROR", code), - message, - domain_code, - ) + return RpcError(message, "ERROR", domain_code) def lease_error(message: str, domain_code: int | None = None) -> LeaseError: - return _build_domain_error( - _LEASE_STATUS_MAP, - lambda msg, code: LeaseError(msg, "ERROR", code), - message, - domain_code, - ) + return LeaseError(message, "ERROR", domain_code) def notice_error(message: str, domain_code: int | None = None) -> NoticeError: - return _build_domain_error( - _NOTICE_STATUS_MAP, - lambda msg, code: NoticeError(msg, "ERROR", code), - message, - domain_code, - ) + return NoticeError(message, "ERROR", domain_code) def stream_error(message: str, domain_code: int | None = None) -> StreamError: - return _build_domain_error( - _STREAM_STATUS_MAP, - lambda msg, code: StreamError(msg, "ERROR", code), - message, - domain_code, - ) + return StreamError(message, "ERROR", domain_code) def schedule_error(message: str, domain_code: int | None = None) -> ScheduleError: - return _build_domain_error( - _SCHEDULE_STATUS_MAP, - lambda msg, code: ScheduleError(msg, "ERROR", code), - message, - domain_code, - ) - - -def is_retryable(error: object) -> bool: - if not isinstance(error, FitzError): - return False - if isinstance(error, (TimeoutError, TransportError)): - return True - if error.domain_code is None: - return False - prefix = error.code.split("_")[0] - return f"{prefix}_{error.domain_code}" in _RETRYABLE_ERROR_CODES + return ScheduleError(message, "ERROR", domain_code) diff --git a/src/fitz_py/multiplexer.py b/src/fitz_py/multiplexer.py index 124d682..125fa47 100644 --- a/src/fitz_py/multiplexer.py +++ b/src/fitz_py/multiplexer.py @@ -1,4 +1,4 @@ -"""In-flight request correlation and notification dispatch for framed messages.""" +"""Cancellation-safe FIFO response correlation and push-frame dispatch.""" from __future__ import annotations @@ -7,35 +7,30 @@ from collections.abc import Awaitable, Callable from dataclasses import dataclass -from fitz_py.errors import ConnectionError, TimeoutError -from fitz_py.types import ConnectionState +from fitz_py.errors import FitzConnectionError, FitzTimeoutError NotificationHandler = Callable[[bytes], None] +PushClassifier = Callable[[bytes], bool] @dataclass(slots=True) class PendingRequest: - """Tracks a pending request future and its timeout handle.""" - - future: asyncio.Future[bytes] - timeout_handle: asyncio.TimerHandle + future: asyncio.Future[bytes] | None + sent: bool = False class Multiplexer: - """Routes response frames to awaiters and notification frames to handlers.""" - def __init__(self) -> None: self._pending: dict[int, deque[PendingRequest]] = defaultdict(deque) self._notification_handlers: dict[int, NotificationHandler] = {} - self._optional_responses: dict[int, int] = {} - self._state = ConnectionState.DISCONNECTED + self._push_classifiers: dict[int, PushClassifier] = {} + self._connected = False def set_connected(self) -> None: - self._state = ConnectionState.AUTHENTICATED + self._connected = True def set_disconnected(self) -> None: - self._state = ConnectionState.DISCONNECTED - self._optional_responses.clear() + self._connected = False self.cancel_all() def register_notification_handler( @@ -46,96 +41,88 @@ def register_notification_handler( def unregister_notification_handler(self, message_type: int) -> None: self._notification_handlers.pop(message_type, None) - def expect_optional_response(self, message_type: int) -> Callable[[], None]: - self._optional_responses[message_type] = self._optional_responses.get(message_type, 0) + 1 - - def cancel() -> None: - current = self._optional_responses.get(message_type, 0) - if current <= 1: - self._optional_responses.pop(message_type, None) - else: - self._optional_responses[message_type] = current - 1 - - return cancel + def register_push_classifier(self, message_type: int, classifier: PushClassifier) -> None: + self._push_classifiers[message_type] = classifier async def request( self, message_type: int, frame_data: bytes, send: Callable[[bytes], Awaitable[None]], - timeout_ms: int, + timeout: float, ) -> bytes: - loop = asyncio.get_running_loop() - future: asyncio.Future[bytes] = loop.create_future() - - def on_timeout() -> None: - self._remove_pending_future(message_type, future) - if not future.done(): - future.set_exception( - TimeoutError( - f"Request timeout for message type {message_type} after {timeout_ms}ms" - ) - ) - - timeout_handle = loop.call_later(timeout_ms / 1000, on_timeout) - self._pending[message_type].append( - PendingRequest(future=future, timeout_handle=timeout_handle) - ) - + future = asyncio.get_running_loop().create_future() + pending = PendingRequest(future) + self._pending[message_type].append(pending) try: await send(frame_data) - return await future + pending.sent = True + async with asyncio.timeout(timeout): + return await future + except TimeoutError as exc: + self._abandon(message_type, pending) + raise FitzTimeoutError( + f"Request timeout for message type {message_type} after {timeout:g}s" + ) from exc + except asyncio.CancelledError: + self._abandon(message_type, pending) + raise except BaseException: - timeout_handle.cancel() - self._remove_pending_future(message_type, future) + if pending.sent: + self._abandon(message_type, pending) + else: + self._remove(message_type, pending) raise - def _remove_pending_future(self, message_type: int, future: asyncio.Future[bytes]) -> None: + def _abandon(self, message_type: int, pending: PendingRequest) -> None: + if pending.sent: + pending.future = None # tombstone consumes a possible late reply + else: + self._remove(message_type, pending) + + def _remove(self, message_type: int, pending: PendingRequest) -> None: queue = self._pending.get(message_type) if queue is None: return - - filtered = deque(item for item in queue if item.future is not future) - if filtered: - self._pending[message_type] = filtered + try: + queue.remove(pending) + except ValueError: return - - self._pending.pop(message_type, None) + if not queue: + self._pending.pop(message_type, None) def dispatch(self, message_type: int, payload: bytes) -> None: + classifier = self._push_classifiers.get(message_type) + if classifier is not None: + try: + if classifier(payload): + self._dispatch_push(message_type, payload) + return + except Exception: + return + queue = self._pending.get(message_type) if queue: pending = queue.popleft() if not queue: self._pending.pop(message_type, None) - pending.timeout_handle.cancel() - if not pending.future.done(): + if pending.future is not None and not pending.future.done(): pending.future.set_result(payload) return + self._dispatch_push(message_type, payload) + def _dispatch_push(self, message_type: int, payload: bytes) -> None: handler = self._notification_handlers.get(message_type) if handler is not None: try: handler(payload) except Exception: return - return - - optional_count = self._optional_responses.get(message_type, 0) - if optional_count > 0: - if optional_count == 1: - self._optional_responses.pop(message_type, None) - else: - self._optional_responses[message_type] = optional_count - 1 - return - - if self._state is not ConnectionState.AUTHENTICATED: - return def cancel_all(self) -> None: + error = FitzConnectionError("Connection closed or reset") for queue in self._pending.values(): for pending in queue: - pending.timeout_handle.cancel() - if not pending.future.done(): - pending.future.set_exception(ConnectionError("Connection closed or reset")) + if pending.future is not None and not pending.future.done(): + pending.future.set_exception(error) self._pending.clear() diff --git a/src/fitz_py/protocol/buffer.py b/src/fitz_py/protocol/buffer.py index 52583a7..631f9a1 100644 --- a/src/fitz_py/protocol/buffer.py +++ b/src/fitz_py/protocol/buffer.py @@ -6,23 +6,34 @@ class BufferWriter: - def __init__(self) -> None: - self._parts: list[bytes] = [] + def __init__(self, capacity: int = 128) -> None: + self._buffer = bytearray(capacity) + self._offset = 0 + + def _reserve(self, size: int) -> int: + start = self._offset + end = start + size + if end > len(self._buffer): + self._buffer.extend(b"\0" * max(end - len(self._buffer), len(self._buffer))) + self._offset = end + return start def write_u8(self, value: int) -> None: - self._parts.append(struct.pack(">B", value)) + struct.pack_into(">B", self._buffer, self._reserve(1), value) def write_u16_be(self, value: int) -> None: - self._parts.append(struct.pack(">H", value)) + struct.pack_into(">H", self._buffer, self._reserve(2), value) def write_u32_be(self, value: int) -> None: - self._parts.append(struct.pack(">I", value)) + struct.pack_into(">I", self._buffer, self._reserve(4), value) def write_u64_be(self, value: int) -> None: - self._parts.append(struct.pack(">Q", value)) + struct.pack_into(">Q", self._buffer, self._reserve(8), value) def write_bytes(self, value: bytes | bytearray | memoryview) -> None: - self._parts.append(bytes(value)) + view = memoryview(value) + start = self._reserve(len(view)) + self._buffer[start : start + len(view)] = view def write_string(self, value: str) -> None: encoded = value.encode("utf-8") @@ -40,12 +51,12 @@ def write_optional_u64(self, value: int | None) -> None: self.write_u64_be(value) def build(self) -> bytes: - return b"".join(self._parts) + return bytes(memoryview(self._buffer)[: self._offset]) class BufferReader: def __init__(self, data: bytes | bytearray | memoryview) -> None: - self._buffer = memoryview(bytes(data)) + self._buffer = memoryview(data) self._offset = 0 def _read(self, size: int) -> memoryview: @@ -77,8 +88,11 @@ def read_route(self) -> str: return self.read_string() def read_optional_u64(self) -> int | None: - if self.read_u8() == 0: + flag = self.read_u8() + if flag == 0: return None + if flag != 1: + raise CodecError(f"Invalid optional u64 flag: {flag}") return self.read_u64_be() def remaining(self) -> bytes: diff --git a/src/fitz_py/protocol/frame.py b/src/fitz_py/protocol/frame.py index 24aa74e..716f76e 100644 --- a/src/fitz_py/protocol/frame.py +++ b/src/fitz_py/protocol/frame.py @@ -15,7 +15,7 @@ class Frame: class FrameCodec: @staticmethod def encode_message_type(message_type: int, writer: BufferWriter) -> None: - if message_type < 0: + if not 0 <= message_type <= 0xFFFF: raise CodecError(f"Invalid message type: {message_type}") if message_type <= 0xFE: writer.write_u8(message_type) @@ -32,6 +32,8 @@ def decode_message_type(reader: BufferReader) -> int: @classmethod def encode_frame(cls, message_type: int, payload: bytes) -> bytes: + if len(payload) > 0xFFFF: + raise CodecError("TLV payload exceeds 65535 bytes") writer = BufferWriter() cls.encode_message_type(message_type, writer) writer.write_u16_be(len(payload)) @@ -47,7 +49,10 @@ def decode_frame(cls, data: bytes) -> Frame: raise CodecError( f"Frame incomplete: expected {payload_length} bytes, got {reader.remaining_bytes()}" ) - return Frame(message_type=message_type, payload=reader.read_bytes(payload_length)) + payload = reader.read_bytes(payload_length) + if not reader.is_eof(): + raise CodecError("Frame contains trailing bytes") + return Frame(message_type=message_type, payload=payload) class FrameParser: diff --git a/src/fitz_py/protocol/messages.py b/src/fitz_py/protocol/messages.py index a59c305..bd15dc4 100644 --- a/src/fitz_py/protocol/messages.py +++ b/src/fitz_py/protocol/messages.py @@ -9,6 +9,9 @@ MSG_KV_DELETE = 106 MSG_KV_DELETE_RANGE = 107 MSG_KV_SCAN = 108 +MSG_KV_SUBSCRIBE = 109 +MSG_KV_UNSUBSCRIBE = 110 +MSG_KV_NOTIFY = 111 MSG_QUEUE_ENQUEUE = 200 MSG_QUEUE_RESERVE = 202 @@ -23,7 +26,6 @@ MSG_RPC_UNSUBSCRIBE_WORKER = 301 MSG_RPC_REQUEST = 302 MSG_RPC_RESPONSE = 303 -MSG_RPC_ACK = 304 MSG_LEASE_ACQUIRE = 400 MSG_LEASE_RENEW = 401 diff --git a/src/fitz_py/protocol/response.py b/src/fitz_py/protocol/response.py index d893083..9accc41 100644 --- a/src/fitz_py/protocol/response.py +++ b/src/fitz_py/protocol/response.py @@ -1,3 +1,5 @@ +"""Strict parsing for Fitz response envelopes.""" + from __future__ import annotations from dataclasses import dataclass @@ -6,29 +8,28 @@ from fitz_py.protocol.buffer import BufferReader -@dataclass(slots=True) -class ParsedResponse: - success: bool - data: bytes +@dataclass(frozen=True, slots=True) +class Response: + data: bytes = b"" + error_code: int | None = None error: str | None = None + @property + def success(self) -> bool: + return self.error is None + -def parse_standard_response(payload: bytes) -> ParsedResponse: +def parse_response(payload: bytes, *, plain: bool = False) -> Response: if not payload: raise ProtocolError("Response payload is empty") reader = BufferReader(payload) status = reader.read_u8() if status == 0: - return ParsedResponse(success=True, data=reader.remaining()) - if status == 1: - if reader.is_eof(): - return ParsedResponse(success=False, data=b"", error="Unknown error (no message)") - return ParsedResponse(success=False, data=b"", error=reader.read_string()) - raise ProtocolError(f"Unknown response status: {status}") - - -def assert_success(payload: bytes, operation: str) -> bytes: - result = parse_standard_response(payload) - if result.success: - return result.data - raise ProtocolError(f"{operation} failed: {result.error or 'Unknown error'}") + return Response(data=reader.remaining()) + if status != 1: + raise ProtocolError(f"Unknown response status: {status}", status) + code = None if plain else reader.read_u32_be() + message = reader.read_string() + if not reader.is_eof(): + raise ProtocolError("Error response has trailing data", code) + return Response(error_code=code, error=message) diff --git a/src/fitz_py/py.typed b/src/fitz_py/py.typed new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/src/fitz_py/py.typed @@ -0,0 +1 @@ + diff --git a/src/fitz_py/transport/base.py b/src/fitz_py/transport/base.py index 535b7d1..55f1ab3 100644 --- a/src/fitz_py/transport/base.py +++ b/src/fitz_py/transport/base.py @@ -20,6 +20,10 @@ async def send(self, data: bytes) -> None: async def receive(self) -> bytes: raise NotImplementedError + @abstractmethod + async def heartbeat(self, timeout: float) -> None: + raise NotImplementedError + @abstractmethod def get_url(self) -> str: raise NotImplementedError diff --git a/src/fitz_py/transport/factory.py b/src/fitz_py/transport/factory.py index 54490a9..08dfa92 100644 --- a/src/fitz_py/transport/factory.py +++ b/src/fitz_py/transport/factory.py @@ -11,14 +11,22 @@ def create_transport( *, timeout_ms: int, max_frame_size: int, + websocket_headers: dict[str, str] | None = None, ) -> Transport: resolved = transport - if resolved == "auto": - resolved = "ws" if url.startswith("ws://") or url.startswith("wss://") else "tcp" + if resolved is TransportType.AUTO: + resolved = ( + TransportType.WEBSOCKET if url.startswith(("ws://", "wss://")) else TransportType.TCP + ) - if resolved == "ws": + if resolved is TransportType.WEBSOCKET: from fitz_py.transport.websocket import WebSocketTransport - return WebSocketTransport(url, timeout_ms=timeout_ms, max_frame_size=max_frame_size) + return WebSocketTransport( + url, + timeout_ms=timeout_ms, + max_frame_size=max_frame_size, + headers=websocket_headers, + ) return TcpTransport(url, timeout_ms=timeout_ms, max_frame_size=max_frame_size) diff --git a/src/fitz_py/transport/tcp.py b/src/fitz_py/transport/tcp.py index 64b0d6c..ab5ed80 100644 --- a/src/fitz_py/transport/tcp.py +++ b/src/fitz_py/transport/tcp.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import socket from urllib.parse import urlparse from fitz_py.errors import TimeoutError, TransportError @@ -25,6 +26,10 @@ async def connect(self) -> None: asyncio.open_connection(self._host, self._port), timeout=self._timeout, ) + sock = self._writer.get_extra_info("socket") + if sock is not None: + sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1) except TimeoutError: raise except asyncio.TimeoutError as exc: # pragma: no cover - transport boundary @@ -70,13 +75,13 @@ async def receive(self) -> bytes: raise TransportError("TCP transport is not connected") try: - header = await asyncio.wait_for(reader.readexactly(4), timeout=self._timeout) + header = await reader.readexactly(4) frame_length = int.from_bytes(header, "big") if frame_length > self._max_frame_size: raise TransportError( f"TCP frame length {frame_length} exceeds max frame size {self._max_frame_size}" ) - return await asyncio.wait_for(reader.readexactly(frame_length), timeout=self._timeout) + return await reader.readexactly(frame_length) except asyncio.IncompleteReadError as exc: # pragma: no cover - transport boundary raise TransportError("TCP connection closed") from exc except asyncio.TimeoutError as exc: # pragma: no cover - transport boundary @@ -88,3 +93,9 @@ async def receive(self) -> bytes: def get_url(self) -> str: return self._url + + async def heartbeat(self, timeout: float) -> None: + _ = timeout + writer = self._writer + if writer is None or writer.is_closing(): + raise TransportError("TCP transport is not connected") diff --git a/src/fitz_py/transport/websocket.py b/src/fitz_py/transport/websocket.py index c708ff3..d73c41f 100644 --- a/src/fitz_py/transport/websocket.py +++ b/src/fitz_py/transport/websocket.py @@ -7,18 +7,30 @@ class WebSocketTransport(Transport): - def __init__(self, url: str, timeout_ms: int = 30000, max_frame_size: int = 65535) -> None: + def __init__( + self, + url: str, + timeout_ms: int = 30000, + max_frame_size: int = 65535, + headers: dict[str, str] | None = None, + ) -> None: self._url = url self._timeout = timeout_ms / 1000 self._max_frame_size = max_frame_size self._socket = None + self._headers = headers or {} async def connect(self) -> None: try: import websockets self._socket = await asyncio.wait_for( - websockets.connect(self._url, max_size=self._max_frame_size), + websockets.connect( + self._url, + max_size=self._max_frame_size, + ping_interval=None, + additional_headers=self._headers or None, + ), timeout=self._timeout, ) except Exception as exc: # pragma: no cover - transport boundary @@ -48,7 +60,7 @@ async def receive(self) -> bytes: if socket is None: raise TransportError("WebSocket transport is not connected") try: - data = await asyncio.wait_for(socket.recv(), timeout=self._timeout) + data = await socket.recv() except Exception as exc: # pragma: no cover - transport boundary raise TransportError(f"WebSocket receive failed: {exc}") from exc @@ -58,3 +70,13 @@ async def receive(self) -> bytes: def get_url(self) -> str: return self._url + + async def heartbeat(self, timeout: float) -> None: + socket = self._socket + if socket is None: + raise TransportError("WebSocket transport is not connected") + try: + pong = await socket.ping() + await asyncio.wait_for(pong, timeout=timeout) + except Exception as exc: + raise TransportError(f"WebSocket heartbeat failed: {exc}") from exc diff --git a/src/fitz_py/types.py b/src/fitz_py/types.py index 7fe18ec..4dc8400 100644 --- a/src/fitz_py/types.py +++ b/src/fitz_py/types.py @@ -1,19 +1,22 @@ -"""Shared SDK configuration, transport, token, and connection-state types.""" +"""Public configuration, lifecycle, and observability contracts.""" from __future__ import annotations -from collections.abc import Awaitable, Callable -from dataclasses import dataclass -from enum import Enum -from typing import Literal, TypeAlias +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass, field +from enum import StrEnum +from typing import Any, Protocol, TypeAlias -TransportType: TypeAlias = Literal["ws", "tcp", "auto"] -TokenProvider: TypeAlias = Callable[[], str | Awaitable[str]] +TokenProvider: TypeAlias = Callable[[], str | bytes | Awaitable[str | bytes]] -class ConnectionState(str, Enum): - """Lifecycle states for a Fitz client connection.""" +class TransportType(StrEnum): + AUTO = "auto" + TCP = "tcp" + WEBSOCKET = "ws" + +class ConnectionState(StrEnum): DISCONNECTED = "DISCONNECTED" CONNECTING = "CONNECTING" CONNECTED = "CONNECTED" @@ -23,25 +26,108 @@ class ConnectionState(str, Enum): CLOSED = "CLOSED" -@dataclass(slots=True) -class ReconnectOptions: - """Reconnect policy settings for automatic transport recovery.""" +@dataclass(frozen=True, slots=True) +class ReconnectPolicy: + enabled: bool = True + max_attempts: int | None = None + backoff: float = 0.25 + max_backoff: float = 5.0 - enabled: bool = False - max_attempts: int | float = float("inf") - backoff_ms: int = 250 - max_backoff_ms: int = 5000 +@dataclass(frozen=True, slots=True) +class RetryPolicy: + enabled: bool = True + max_attempts: int = 3 + backoff: float = 0.1 + max_backoff: float = 1.0 + + +@dataclass(frozen=True, slots=True) +class HeartbeatPolicy: + enabled: bool = True + interval: float = 10.0 + timeout: float = 30.0 + + +@dataclass(frozen=True, slots=True) +class ConcurrencyLimits: + max_in_flight: int = 256 + request_queue_size: int = 1024 + async_handler_concurrency: int = 256 + async_handler_queue_size: int = 1024 + subscription_buffer_size: int = 1024 + async_handler_timeout: float = 30.0 -@dataclass(slots=True) -class ClientConfig: - """Configuration used to construct a high-level Fitz client.""" +@dataclass(frozen=True, slots=True) +class LifecycleEvent: + event: str + state: ConnectionState + url: str + transport: str | None = None + attempt: int | None = None + domain: str | None = None + registration: str | None = None + error: str | None = None + + +class FitzTracer(Protocol): + def start_span(self, name: str, attributes: Mapping[str, Any]) -> Any: ... + + +class FitzMeter(Protocol): + def counter(self, name: str, value: int, attributes: Mapping[str, Any]) -> None: ... + + def histogram(self, name: str, value: float, attributes: Mapping[str, Any]) -> None: ... + + def gauge(self, name: str, value: int, attributes: Mapping[str, Any]) -> None: ... + + +@dataclass(frozen=True, slots=True) +class Observability: + logger: Any | None = None + tracer: FitzTracer | None = None + meter: FitzMeter | None = None + on_lifecycle_event: Callable[[LifecycleEvent], None] | None = None + + +@dataclass(frozen=True, slots=True) +class ClientConfig: url: str token_provider: TokenProvider | None = None - timeout_ms: int = 30000 - transport: TransportType = "auto" - reconnect: ReconnectOptions | None = None - max_frame_size: int = 65535 - auth_settle_delay_ms: int = 500 - max_in_flight_requests: int = 256 + transport: TransportType | str = TransportType.AUTO + request_timeout: float = 30.0 + auth_settle_timeout: float = 1.0 + max_frame_size: int = 65_540 + reconnect: ReconnectPolicy = field(default_factory=ReconnectPolicy) + retry: RetryPolicy = field(default_factory=RetryPolicy) + heartbeat: HeartbeatPolicy = field(default_factory=HeartbeatPolicy) + limits: ConcurrencyLimits = field(default_factory=ConcurrencyLimits) + observability: Observability = field(default_factory=Observability) + websocket_headers: Mapping[str, str] = field(default_factory=dict) + + def __post_init__(self) -> None: + if not self.url: + raise ValueError("url is required") + if self.request_timeout <= 0 or self.auth_settle_timeout < 0: + raise ValueError("timeouts must be positive") + if self.max_frame_size < 5: + raise ValueError("max_frame_size must include the five-byte frame header") + if self.reconnect.max_attempts is not None and self.reconnect.max_attempts < 0: + raise ValueError("reconnect.max_attempts must be non-negative or None") + if self.retry.max_attempts < 1: + raise ValueError("retry.max_attempts must be at least 1") + if self.limits.max_in_flight < 1: + raise ValueError("limits.max_in_flight must be at least 1") + if self.limits.request_queue_size < 0: + raise ValueError("limits.request_queue_size must be non-negative") + if self.limits.async_handler_concurrency < 1: + raise ValueError("limits.async_handler_concurrency must be at least 1") + if self.limits.async_handler_queue_size < 0: + raise ValueError("limits.async_handler_queue_size must be non-negative") + if self.limits.subscription_buffer_size < 1: + raise ValueError("limits.subscription_buffer_size must be at least 1") + + +# Short aliases used in annotations and migration docs. +ReconnectOptions = ReconnectPolicy diff --git a/tests/conformance/cross-language-conformance-suite.yaml b/tests/conformance/cross-language-conformance-suite.yaml new file mode 100644 index 0000000..f8ee5fb --- /dev/null +++ b/tests/conformance/cross-language-conformance-suite.yaml @@ -0,0 +1,270 @@ +version: "1.0" +name: "fitz-cross-language-client-conformance" +status: "draft" +updated: "2026-05-16" +owners: + - "fitz server" + - "fitz-go" + - "fitz-ts" + - "fitz-py" + +policy: + strict_reconnect_parity: true + behavioral_equivalence_over_naming: true + undocumented_divergence_is_defect: true + first_class_semantics: + - timeout + - cancellation + - cleanup + - streaming + +contract: + required_domains: + - kv + - queue + - rpc + - lease + - notice + - stream + - schedule + required_transports: + - websocket + - tcp + auth_modes: + - anonymous + - valid_jwt + - invalid_jwt + reconnect_contract: + required: true + expectations: + - "client detects disconnection" + - "client reconnects using configured backoff" + - "client sends CONNECT first on new transport" + - "client treats reconnect as a new broker session" + - "client rebuilds client-owned Notice, Queue, Lease, Stream, and Schedule subscriptions" + - "client rebuilds client-owned RPC worker registrations" + - "client invalidates KV transactions, Stream append sessions, Queue item handles, Lease handles, and pending RPC calls" + - "client resumes Stream replay only from client-owned offsets" + +result_schema: + verdict: ["pass", "partial", "fail", "not_implemented", "unclear"] + required_fields: + - scenario_id + - client + - transport + - auth_mode + - verdict + - evidence + - latency_ms + +scenarios: + - id: "CS-001" + title: "connect success" + priority: "P0" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "connect returns success" + - "client state becomes authenticated" + - "first user domain request succeeds" + required_evidence: + - "connection state transition logs" + - "first successful domain response" + + - id: "CS-002" + title: "auth failure" + priority: "P0" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["invalid_jwt"] + expected: + - "connect fails" + - "failure is surfaced as authentication error" + - "client does not silently continue as authenticated" + required_evidence: + - "transport close or auth failure exception" + + - id: "CS-003" + title: "request success" + priority: "P0" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "kv transaction begin/put/commit succeeds" + - "read-after-commit returns written value" + + - id: "CS-004" + title: "unknown route" + priority: "P0" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "operation targeting unknown route returns domain error" + - "error remains typed/mappable by client" + + - id: "CS-005" + title: "invalid payload" + priority: "P0" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "malformed operation request returns protocol/domain error" + - "client remains usable after error" + + - id: "CS-006" + title: "server error mapping" + priority: "P0" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "known server error code maps to client domain error" + - "retryable vs terminal classification is available" + + - id: "CS-007" + title: "timeout handling" + priority: "P0" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "long-running call times out using client timeout control" + - "timeout error is distinguishable from cancellation" + - "connection stays healthy after timeout" + + - id: "CS-008" + title: "caller cancellation" + priority: "P0" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "in-flight request can be cancelled by caller primitive" + - "cancellation is surfaced as cancellation (not timeout)" + - "subsequent requests succeed" + + - id: "CS-009" + title: "disconnect during request" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "in-flight request fails promptly" + - "failure type indicates disconnect/interruption" + - "client transitions to reconnect flow when configured" + + - id: "CS-010" + title: "reconnect and retry behavior" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "the same client instance observes a real transport loss; constructing a replacement client does not pass" + - "client performs configured reconnect backoff" + - "client re-creates reconnect-safe subscriptions and worker registrations" + - "client fails or invalidates stale session-bound handles" + - "new requests succeed post reconnect" + + - id: "CS-011" + title: "stream receive sequence" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "stream records are received in offset order" + - "no duplicate or skipped offsets in normal path" + + - id: "CS-012" + title: "stream completion" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "stream completion/end is observable in client API" + - "iterator/subscription closes cleanly" + + - id: "CS-013" + title: "stream error mid-flight" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "streaming API surfaces mid-flight server error" + - "cleanup occurs without leaking resources" + + - id: "CS-014" + title: "concurrent in-flight requests" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "cross-domain in-flight requests can proceed concurrently" + - "within-domain ordering follows protocol constraints" + - "responses correlate correctly to request context" + + - id: "CS-015" + title: "shutdown during active work" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "closing client during active requests/streams terminates gracefully" + - "active operations fail predictably" + - "double-close is safe" + + - id: "CS-016" + title: "filtered stream replay" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + expected: + - "client can append an optional discriminator on stream records" + - "client can issue a filtered read using StreamFilterSet" + - "filtered replay returns matching records and synthetic filtered markers in offset order" + - "cursor advances monotonically through filtered offsets" + + - id: "CS-017" + title: "bounded concurrency under burst load" + priority: "P1" + setup: + transport: ["websocket", "tcp"] + auth_mode: ["anonymous", "valid_jwt"] + concurrency_limit: 16 + burst_size: 64 + expected: + - "client enforces the configured in-flight ceiling" + - "requests above the ceiling are queued or backpressured instead of spawning unbounded work" + - "active work drains successfully once capacity frees" + - "client close returns background work to baseline" + required_evidence: + - "observed_max_inflight" + - "backpressure_or_queueing_event" + - "post_close_active_work_count" + +scoring: + fail_on: + - "any P0 scenario verdict != pass" + warn_on: + - "any P1 scenario verdict in [partial, fail, not_implemented, unclear]" + +reporting: + output_format: "json" + aggregate_fields: + - "client" + - "transport" + - "auth_mode" + - "p0_pass_rate" + - "p1_pass_rate" + - "overall_status" diff --git a/tests/conformance/test_conformance.py b/tests/conformance/test_conformance.py index da5398a..34ca1ec 100644 --- a/tests/conformance/test_conformance.py +++ b/tests/conformance/test_conformance.py @@ -36,8 +36,13 @@ AuthenticationError, Client, ClientConfig, + ConcurrencyLimits, FitzError, + RequestQueueFullError, + StreamFilterClause, + StreamFilterSet, ) +from fitz_py._runtime import RequestGate from tests.integration.fixture.jwt import make_expired_jwt, make_valid_jwt # noqa: E402 # --------------------------------------------------------------------------- @@ -102,9 +107,9 @@ async def _new_client( ClientConfig( url=url, token_provider=provider, - timeout_ms=timeout_ms, + request_timeout=timeout_ms / 1000, transport=CONFORMANCE_TRANSPORT, - max_in_flight_requests=max_in_flight_requests, + limits=ConcurrencyLimits(max_in_flight=max_in_flight_requests), ) ) await client.connect() @@ -223,7 +228,7 @@ async def test_cs001_connect_success() -> None: evidence.append("connect returned successfully") route = _unique_route("kv") - tx = await client.kv().begin(route, durability="sync") + tx = await client.kv.begin(route, durability="sync") await tx.put(b"cs001-key", b"cs001-value") await tx.commit() evidence.append("first domain request (kv) succeeded") @@ -283,7 +288,7 @@ async def test_cs002_auth_failure() -> None: evidence.append("connect did not raise (unexpected)") # Try a domain operation for diagnostic evidence only. try: - await client.kv().begin(_unique_route("kv"), durability="sync") + await client.kv.begin(_unique_route("kv"), durability="sync") evidence.append("WARNING: domain request unexpectedly succeeded") except Exception as dom_exc: evidence.append(f"domain request failed post-auth: {dom_exc}") @@ -316,12 +321,12 @@ async def test_cs003_request_success() -> None: client = await _new_client() try: route = _unique_route("kv") - tx = await client.kv().begin(route, durability="sync") + tx = await client.kv.begin(route, durability="sync") await tx.put(b"user:1", b"Alice") await tx.commit() evidence.append("kv begin/put/commit succeeded") - rtx = await client.kv().begin(route, mode="read_only", durability="sync") + rtx = await client.kv.begin(route, mode="read_only", durability="sync") result = await rtx.get(b"user:1") assert result.found, "expected found=True" assert result.value == b"Alice" @@ -357,7 +362,7 @@ async def test_cs004_unknown_route() -> None: no_worker_route = _unique_route("rpc") caught: Exception | None = None try: - iterator = await client.rpc().call(no_worker_route, b"ping", timeout_ms=500) + iterator = await client.rpc.call(no_worker_route, b"ping", timeout=0.5) async for _frame in iterator: pass except Exception as exc: @@ -368,7 +373,7 @@ async def test_cs004_unknown_route() -> None: # Client must remain usable route = _unique_route("kv") - tx = await client.kv().begin(route, durability="sync") + tx = await client.kv.begin(route, durability="sync") await tx.put(b"k", b"v") await tx.commit() evidence.append("client remains usable after unknown-route error") @@ -401,12 +406,12 @@ async def test_cs005_invalid_payload() -> None: client = await _new_client() try: route = _unique_route("kv") - tx1 = await client.kv().begin(route, durability="sync") + tx1 = await client.kv.begin(route, durability="sync") await tx1.insert(b"dup-key", b"first") await tx1.commit() evidence.append("first insert succeeded") - tx2 = await client.kv().begin(route, durability="sync") + tx2 = await client.kv.begin(route, durability="sync") caught: Exception | None = None try: await tx2.insert(b"dup-key", b"second") @@ -421,7 +426,7 @@ async def test_cs005_invalid_payload() -> None: assert caught is not None, "expected error on duplicate insert" evidence.append(f"duplicate insert raised {type(caught).__name__}: {caught}") - rtx = await client.kv().begin(route, mode="read_only", durability="sync") + rtx = await client.kv.begin(route, mode="read_only", durability="sync") result = await rtx.get(b"dup-key") assert result.found evidence.append("client remains usable after server-rejected operation") @@ -456,7 +461,7 @@ async def test_cs006_server_error_mapping() -> None: route = _unique_route("rpc") rpc_err: Exception | None = None try: - iterator = await client.rpc().call(route, b"ping", timeout_ms=500) + iterator = await client.rpc.call(route, b"ping", timeout=0.5) async for _frame in iterator: pass except Exception as exc: @@ -472,11 +477,11 @@ async def test_cs006_server_error_mapping() -> None: # KV conflict — verify typed error kv_route = _unique_route("kv") - tx = await client.kv().begin(kv_route, durability="sync") + tx = await client.kv.begin(kv_route, durability="sync") await tx.insert(b"x", b"1") await tx.commit() - tx2 = await client.kv().begin(kv_route, durability="sync") + tx2 = await client.kv.begin(kv_route, durability="sync") kv_err: Exception | None = None try: await tx2.insert(b"x", b"2") @@ -522,7 +527,7 @@ async def test_cs007_timeout_handling() -> None: start = time.monotonic() caught: Exception | None = None try: - iterator = await client.rpc().call(route, b"nobody", timeout_ms=250) + iterator = await client.rpc.call(route, b"nobody", timeout=0.25) async for _frame in iterator: pass except Exception as exc: @@ -539,7 +544,7 @@ async def test_cs007_timeout_handling() -> None: # Connection must remain healthy kv_route = _unique_route("kv") - tx = await client.kv().begin(kv_route, durability="sync") + tx = await client.kv.begin(kv_route, durability="sync") await tx.put(b"post-timeout", b"ok") await tx.commit() evidence.append("connection healthy after timeout") @@ -579,14 +584,14 @@ async def test_cs008_caller_cancellation() -> None: async def _slow_handler(req, writer) -> None: try: await asyncio.sleep(3.0) - await writer.send(b"late", is_end=True) + await writer.send(b"late", end=True) finally: handler_finished.set() - sub = await worker_client.rpc().register_worker(route, _slow_handler) + sub = await worker_client.rpc.register_worker(route, _slow_handler) async def _do_call() -> None: - iterator = await caller_client.rpc().call(route, b"block", timeout_ms=30000) + iterator = await caller_client.rpc.call(route, b"block", timeout=30) async for _frame in iterator: pass @@ -616,7 +621,7 @@ async def _do_call() -> None: # Subsequent request must succeed kv_route = _unique_route("kv") - tx = await caller_client.kv().begin(kv_route, durability="sync") + tx = await caller_client.kv.begin(kv_route, durability="sync") await tx.put(b"after-cancel", b"ok") await tx.commit() evidence.append("subsequent request succeeded after cancellation") @@ -659,14 +664,14 @@ async def _slow_handler(req, writer) -> None: handler_started.set() try: await asyncio.sleep(1.5) - await writer.send(b"late", is_end=True) + await writer.send(b"late", end=True) finally: handler_finished.set() - sub = await worker_client.rpc().register_worker(route, _slow_handler) + sub = await worker_client.rpc.register_worker(route, _slow_handler) async def _do_call() -> None: - iterator = await caller_client.rpc().call(route, b"block", timeout_ms=30000) + iterator = await caller_client.rpc.call(route, b"block", timeout=30) async for _frame in iterator: pass @@ -730,7 +735,7 @@ async def test_cs010_reconnect_behavior() -> None: client2 = await _new_client() try: route = _unique_route("kv") - tx = await client2.kv().begin(route, durability="sync") + tx = await client2.kv.begin(route, durability="sync") await tx.put(b"after-reconnect", b"ok") await tx.commit() evidence.append("new requests succeed after reconnect (new client)") @@ -766,13 +771,13 @@ async def test_cs011_stream_receive_sequence() -> None: client = await _new_client() try: route = _unique_route("stream") - session = await client.stream().begin(route) + session = await client.stream.begin(route) for i in range(3): await session.append(i, bytes([i * 10])) await session.commit() evidence.append("stream session appended 3 records") - records = await client.stream().read(route, start_offset=0, limit=10) + records = await client.stream.read(route, start_offset=0, limit=10) assert len(records) >= 3, f"expected >=3 stream records, got {len(records)}" evidence.append(f"read {len(records)} records after commit") @@ -812,13 +817,13 @@ async def test_cs012_stream_completion() -> None: client = await _new_client() try: route = _unique_route("stream") - session = await client.stream().begin(route) + session = await client.stream.begin(route) await session.append(0, b"first") await session.append(1, b"last") await session.commit() evidence.append("stream session committed") - records = await client.stream().read(route, start_offset=0, limit=100) + records = await client.stream.read(route, start_offset=0, limit=100) assert len(records) >= 2, f"expected >=2 records after commit, got {len(records)}" evidence.append(f"stream.read() completed cleanly with {len(records)} records") evidence.append("iterator/read closed cleanly (no resource leak)") @@ -852,7 +857,7 @@ async def test_cs013_stream_error_mid_flight() -> None: wrong_session = None try: route = _unique_route("stream") - session = await client.stream().begin(route) + session = await client.stream.begin(route) await session.append(0, b"record-1") await session.commit() evidence.append("written first record at offset 0") @@ -860,7 +865,7 @@ async def test_cs013_stream_error_mid_flight() -> None: caught: Exception | None = None try: # Expected offset 0 again — server should reject on append - wrong_session = await client.stream().begin(route) + wrong_session = await client.stream.begin(route) await wrong_session.append(0, b"record-2") except Exception as exc: caught = exc @@ -870,7 +875,7 @@ async def test_cs013_stream_error_mid_flight() -> None: # Client must remain usable kv_route = _unique_route("kv") - tx = await client.kv().begin(kv_route, durability="sync") + tx = await client.kv.begin(kv_route, durability="sync") await tx.put(b"after-stream-error", b"ok") await tx.commit() evidence.append("client still usable after stream error") @@ -910,10 +915,10 @@ async def test_cs014_concurrent_requests() -> None: routes = [_unique_route("kv") for _ in range(3)] async def _kv_roundtrip(route: str, idx: int) -> str: - tx = await client.kv().begin(route, durability="sync") + tx = await client.kv.begin(route, durability="sync") await tx.put(f"key-{idx}".encode(), f"value-{idx}".encode()) await tx.commit() - rtx = await client.kv().begin(route, mode="read_only", durability="sync") + rtx = await client.kv.begin(route, mode="read_only", durability="sync") result = await rtx.get(f"key-{idx}".encode()) return result.value.decode() if result.found else "" @@ -957,7 +962,7 @@ async def test_cs015_shutdown_during_active_work() -> None: client = await _new_client() route = _unique_route("kv") - begin_task = asyncio.create_task(client.kv().begin(route, durability="sync")) + begin_task = asyncio.create_task(client.kv.begin(route, durability="sync")) await asyncio.sleep(0.05) await client.close() @@ -997,36 +1002,41 @@ async def test_cs015_shutdown_during_active_work() -> None: # --------------------------------------------------------------------------- -# CS-018 — queue enqueue/reserve/complete lifecycle + +# --------------------------------------------------------------------------- +# CS-016 — filtered stream replay # --------------------------------------------------------------------------- @pytest.mark.asyncio -async def test_cs018_queue_enqueue_reserve_complete() -> None: +async def test_cs016_filtered_stream_replay() -> None: evidence: list[str] = [] client = await _new_client() try: - route = _unique_route("queue") - msg_id = await client.queue().enqueue(route, b"cs018-payload") - evidence.append(f"enqueued message id={msg_id}") - - items = await client.queue().reserve(route, 30, batch_size=1) - assert len(items) == 1, f"expected 1 reserved item, got {len(items)}" - assert items[0].body == b"cs018-payload" - evidence.append("reserved item matches payload") - - await items[0].complete() - evidence.append("message completed") - - empty = await client.queue().reserve(route, 30, batch_size=1) - assert not empty, f"expected empty queue after complete, got {len(empty)} items" - evidence.append("queue empty after complete") + route = _unique_route("stream") + session = await client.stream.begin(route) + await session.append(0, b"accepted", discriminator="keep") + await session.append(1, b"filtered", discriminator="drop") + await session.commit() + page = await client.stream.read_page( + route, + 0, + limit=10, + stream_filter=StreamFilterSet( + clauses=[StreamFilterClause(kind="Equals", value="keep")] + ), + ) + records = [item.record for item in page.items if item.record is not None] + assert [record.body for record in records] == [b"accepted"] + assert all(item.route == route for item in page.items) + evidence.append("filtered replay returned only the matching event") + evidence.append("every event and marker retained its concrete route") finally: await client.close() - r = ScenarioResult( - "CS-018", - "queue enqueue/reserve/complete lifecycle", + result = ScenarioResult( + "CS-016", + "filtered stream replay", "P1", CLIENT_NAME, CONFORMANCE_TRANSPORT, @@ -1035,11 +1045,10 @@ async def test_cs018_queue_enqueue_reserve_complete() -> None: evidence, 0, ) - _record(r) - assert r.verdict == "pass" + _record(result) + assert result.verdict == "pass" -# --------------------------------------------------------------------------- # CS-017 - bounded concurrency under burst load # --------------------------------------------------------------------------- @@ -1047,31 +1056,22 @@ async def test_cs018_queue_enqueue_reserve_complete() -> None: @pytest.mark.asyncio async def test_cs017_bounded_concurrency_under_burst_load() -> None: evidence: list[str] = [] - client = await _new_client(timeout_ms=750, max_in_flight_requests=1) - first_task = None - second_task = None + gate = RequestGate(1, 1) + release_first = await gate.acquire() + second_task = asyncio.create_task(gate.acquire()) + await asyncio.sleep(0) try: - route = _unique_route("rpc") - - async def _delayed_worker(_request, writer) -> None: - await asyncio.sleep(0.5) - await writer.send(b"delayed", is_end=True) - - await client.rpc().register_worker(route, _delayed_worker) - - first_task = asyncio.create_task(client.rpc().call(route, b"first", timeout_ms=750)) - second_task = asyncio.create_task(client.rpc().call(route, b"second", timeout_ms=750)) - - await asyncio.sleep(0.1) - assert second_task is not None and not second_task.done(), "expected second RPC call to remain pending behind the first" - - evidence.append("second RPC call remained pending while first was in flight") - evidence.append("configured max_in_flight_requests=1 and burst size=2") + with pytest.raises(RequestQueueFullError): + await gate.acquire() + evidence.append("burst above active plus queue capacity failed explicitly") + release_first() + release_second = await asyncio.wait_for(second_task, timeout=1) + release_second() + evidence.append("queued work advanced in FIFO order after capacity released") finally: - await client.close() - tasks = [task for task in (first_task, second_task) if task is not None] - if tasks: - await asyncio.gather(*tasks, return_exceptions=True) + gate.close() + if not second_task.done(): + second_task.cancel() r = ScenarioResult( "CS-017", @@ -1089,145 +1089,3 @@ async def _delayed_worker(_request, writer) -> None: # --------------------------------------------------------------------------- -# CS-019 — lease acquire/contention/release lifecycle -# --------------------------------------------------------------------------- - - -@pytest.mark.asyncio -async def test_cs019_lease_acquire_contention_release() -> None: - evidence: list[str] = [] - client1 = await _new_client() - client2 = await _new_client() - try: - route = _unique_route("lease") - lease1 = await client1.lease().acquire(route, 30) - evidence.append("client1 acquired lease") - - caught: Exception | None = None - try: - await client2.lease().acquire(route, 30) - except Exception as exc: - caught = exc - - assert caught is not None, "expected contention error for second lease acquire" - evidence.append(f"client2 rejected while held: {type(caught).__name__}") - - await lease1.release() - evidence.append("client1 released lease") - - lease2 = await client2.lease().acquire(route, 30) - evidence.append("client2 acquired lease after release") - await lease2.release() - finally: - await client1.close() - await client2.close() - - r = ScenarioResult( - "CS-019", - "lease acquire/contention/release lifecycle", - "P1", - CLIENT_NAME, - CONFORMANCE_TRANSPORT, - CONFORMANCE_AUTH_MODE, - "pass", - evidence, - 0, - ) - _record(r) - assert r.verdict == "pass" - - -# --------------------------------------------------------------------------- -# CS-020 — notice subscribe/publish/deliver lifecycle -# --------------------------------------------------------------------------- - - -@pytest.mark.asyncio -async def test_cs020_notice_subscribe_publish_deliver() -> None: - evidence: list[str] = [] - client = await _new_client() - try: - route = _unique_route("notice") - received: list[bytes] = [] - delivered = asyncio.Event() - - async def _handler(message) -> None: - received.append(message.body) - delivered.set() - - sub = await client.notice().subscribe(route, _handler) - evidence.append("subscribed to route") - - await client.notice().publish(route, b"cs020-msg") - await asyncio.wait_for(delivered.wait(), timeout=5.0) - assert received == [b"cs020-msg"] - evidence.append("handler received message") - - await sub.unsubscribe() - evidence.append("unsubscribed") - - delivered.clear() - await client.notice().publish(route, b"after-unsub") - with pytest.raises(asyncio.TimeoutError): - await asyncio.wait_for(delivered.wait(), timeout=0.5) - evidence.append("no delivery after unsubscribe") - finally: - await client.close() - - r = ScenarioResult( - "CS-020", - "notice subscribe/publish/deliver lifecycle", - "P1", - CLIENT_NAME, - CONFORMANCE_TRANSPORT, - CONFORMANCE_AUTH_MODE, - "pass", - evidence, - 0, - ) - _record(r) - assert r.verdict == "pass" - - -# --------------------------------------------------------------------------- -# CS-021 — schedule create/subscribe/cancel lifecycle -# --------------------------------------------------------------------------- - - -@pytest.mark.asyncio -async def test_cs021_schedule_create_subscribe_cancel() -> None: - evidence: list[str] = [] - client = await _new_client() - try: - route = _unique_route("schedule") - - async def _handler(_notification) -> None: - return None - - sub = await client.schedule().subscribe(route, _handler) - evidence.append("subscribed to schedule route") - - schedule_id = await client.schedule().create(route, "0 9 * * 1", b"cs021-payload") - evidence.append(f"schedule created id={schedule_id or route}") - - await client.schedule().cancel(schedule_id or route) - evidence.append("schedule cancelled") - - await sub.unsubscribe() - evidence.append("unsubscribed") - finally: - await client.close() - - r = ScenarioResult( - "CS-021", - "schedule create/subscribe/cancel lifecycle", - "P1", - CLIENT_NAME, - CONFORMANCE_TRANSPORT, - CONFORMANCE_AUTH_MODE, - "pass", - evidence, - 0, - ) - _record(r) - assert r.verdict == "pass" diff --git a/tests/integration/test_broker.py b/tests/integration/test_broker.py new file mode 100644 index 0000000..f2fc864 --- /dev/null +++ b/tests/integration/test_broker.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +import pytest + +from fitz_py import KvDurability +from tests.integration.fixture.fixture import IntegrationFixture, unique_route + + +@pytest.mark.parametrize("transport", ["tcp", "ws"]) +@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) +@pytest.mark.asyncio +async def test_connect_and_kv_round_trip(transport: str, auth_mode: str) -> None: + fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] + try: + route = unique_route("kv") + async with await fixture.client.kv.begin(route, durability=KvDurability.SYNC) as tx: + await tx.put(b"key", b"value") + assert (await tx.get(b"key")).value == b"value" + await tx.commit() + async with await fixture.client.kv.begin(route, durability=KvDurability.BUFFERED) as tx: + assert (await tx.get(b"key")).value == b"value" + finally: + await fixture.close() + + +@pytest.mark.parametrize("transport", ["tcp", "ws"]) +@pytest.mark.asyncio +async def test_queue_round_trip(transport: str) -> None: + fixture = await IntegrationFixture.connect_or_fail(transport, "anonymous") + try: + route = unique_route("queue") + await fixture.client.queue.enqueue(route, b"work") + items = await fixture.client.queue.reserve(route, lease=30, wait=2) + assert [item.body for item in items] == [b"work"] + await items[0].complete() + finally: + await fixture.close() diff --git a/tests/integration/test_kv.py b/tests/integration/test_kv.py deleted file mode 100644 index 7b956e3..0000000 --- a/tests/integration/test_kv.py +++ /dev/null @@ -1,26 +0,0 @@ -from __future__ import annotations - -import pytest - -from tests.integration.fixture.fixture import IntegrationFixture, unique_route - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_kv_round_trip(transport: str, auth_mode: str) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("kv") - tx = await fixture.client.kv().begin(route, durability="sync") - await tx.put(b"hello", b"world") - await tx.commit() - - rtx = await fixture.client.kv().begin(route, mode="read_only", durability="sync") - result = await rtx.get(b"hello") - await rtx.rollback() - - assert result.found is True - assert result.value == b"world" - finally: - await fixture.close() diff --git a/tests/integration/test_lease.py b/tests/integration/test_lease.py deleted file mode 100644 index d6aa7b8..0000000 --- a/tests/integration/test_lease.py +++ /dev/null @@ -1,31 +0,0 @@ -from __future__ import annotations - -import pytest - -from fitz_py import LeaseError -from tests.integration.fixture.fixture import IntegrationFixture, unique_route - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_lease_lifecycle(transport: str, auth_mode: str) -> None: - fixture1 = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - fixture2 = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("lease") - lease = await fixture1.client.lease().acquire(route, 30) - assert lease.token > 0 - - with pytest.raises(LeaseError): - await fixture2.client.lease().acquire(route, 30) - - await lease.extend(30) - await lease.release() - - lease2 = await fixture2.client.lease().acquire(route, 30) - assert lease2.token > 0 - await lease2.release() - finally: - await fixture1.close() - await fixture2.close() diff --git a/tests/integration/test_notice.py b/tests/integration/test_notice.py deleted file mode 100644 index 92d6df1..0000000 --- a/tests/integration/test_notice.py +++ /dev/null @@ -1,83 +0,0 @@ -from __future__ import annotations - -import asyncio - -import pytest - -from tests.integration.fixture.fixture import IntegrationFixture, unique_route - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_notice_subscribe_publish_unsubscribe(transport: str, auth_mode: str) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("notice") - received: list[bytes] = [] - delivered = asyncio.Event() - - async def handler(message) -> None: - received.append(message.body) - delivered.set() - - sub = await fixture.client.notice().subscribe(route, handler) - await fixture.client.notice().publish(route, b"hello") - await asyncio.wait_for(delivered.wait(), timeout=5) - assert received == [b"hello"] - - await sub.unsubscribe() - delivered.clear() - await fixture.client.notice().publish(route, b"after-unsub") - with pytest.raises(asyncio.TimeoutError): - await asyncio.wait_for(delivered.wait(), timeout=0.5) - finally: - await fixture.close() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_notice_multiple_local_subscribers_same_pattern( - transport: str, auth_mode: str -) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("notice") - received_one: list[bytes] = [] - received_two: list[bytes] = [] - delivered_one = asyncio.Event() - delivered_two = asyncio.Event() - - async def handler_one(message) -> None: - received_one.append(message.body) - delivered_one.set() - - async def handler_two(message) -> None: - received_two.append(message.body) - delivered_two.set() - - sub_one = await fixture.client.notice().subscribe(route, handler_one) - sub_two = await fixture.client.notice().subscribe(route, handler_two) - - assert sub_one._sub_id == sub_two._sub_id - assert sub_one.pattern == sub_two.pattern == route - - await fixture.client.notice().publish(route, b"fanout") - await asyncio.wait_for(delivered_one.wait(), timeout=5) - await asyncio.wait_for(delivered_two.wait(), timeout=5) - - assert received_one == [b"fanout"] - assert received_two == [b"fanout"] - - delivered_two.clear() - await sub_one.unsubscribe() - await fixture.client.notice().publish(route, b"second") - await asyncio.wait_for(delivered_two.wait(), timeout=5) - - assert received_one == [b"fanout"] - assert received_two == [b"fanout", b"second"] - - await sub_two.unsubscribe() - finally: - await fixture.close() diff --git a/tests/integration/test_queue.py b/tests/integration/test_queue.py deleted file mode 100644 index 2b6ad47..0000000 --- a/tests/integration/test_queue.py +++ /dev/null @@ -1,27 +0,0 @@ -from __future__ import annotations - -import pytest - -from tests.integration.fixture.fixture import IntegrationFixture, unique_route - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_queue_lifecycle(transport: str, auth_mode: str) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("queue") - msg_id = await fixture.client.queue().enqueue(route, b"payload") - assert msg_id > 0 - - items = await fixture.client.queue().reserve(route, 30, batch_size=1) - assert len(items) == 1 - assert items[0].body == b"payload" - - await items[0].extend(30) - await items[0].complete() - - assert await fixture.client.queue().reserve(route, 30, batch_size=1) == [] - finally: - await fixture.close() diff --git a/tests/integration/test_rpc.py b/tests/integration/test_rpc.py deleted file mode 100644 index 93a7572..0000000 --- a/tests/integration/test_rpc.py +++ /dev/null @@ -1,27 +0,0 @@ -from __future__ import annotations - -import pytest - -from tests.integration.fixture.fixture import IntegrationFixture, unique_route - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_rpc_worker_round_trip(transport: str, auth_mode: str) -> None: - worker = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - caller = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("rpc") - - async def handler(req, writer) -> None: - await writer.send(req.body.upper(), is_end=True) - - sub = await worker.client.rpc().register_worker(route, handler) - iterator = await caller.client.rpc().call(route, b"ping", timeout_ms=2000) - frames = [frame async for frame in iterator] - assert [frame.body for frame in frames] == [b"PING"] - await sub.unsubscribe() - finally: - await worker.close() - await caller.close() diff --git a/tests/integration/test_schedule.py b/tests/integration/test_schedule.py deleted file mode 100644 index c8f87ce..0000000 --- a/tests/integration/test_schedule.py +++ /dev/null @@ -1,74 +0,0 @@ -from __future__ import annotations - -import asyncio - -import pytest - -from tests.integration.fixture.fixture import IntegrationFixture, unique_route - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_schedule_create_subscribe_cancel(transport: str, auth_mode: str) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("schedule") - - async def handler(_notification) -> None: - return None - - sub = await fixture.client.schedule().subscribe(route, handler) - schedule_id = await fixture.client.schedule().create(route, "0 9 * * 1", b"payload") - assert schedule_id == route or isinstance(schedule_id, str) - - await fixture.client.schedule().cancel(route) - await sub.unsubscribe() - finally: - await fixture.close() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_schedule_multiple_local_subscribers_same_pattern( - transport: str, auth_mode: str -) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("schedule") - received_one: list[bytes] = [] - received_two: list[bytes] = [] - - async def handler_one(notification) -> None: - received_one.append(notification.payload) - - async def handler_two(notification) -> None: - received_two.append(notification.payload) - - sub_one = await fixture.client.schedule().subscribe(route, handler_one) - sub_two = await fixture.client.schedule().subscribe(route, handler_two) - - assert sub_one._sub_id == sub_two._sub_id - assert sub_one.pattern == sub_two.pattern == route - - payload = len(b"fanout").to_bytes(4, "big") + b"fanout" - fixture.client.schedule().connection.get_multiplexer().dispatch( # type: ignore[attr-defined] - 705, sub_one._sub_id.to_bytes(8, "big") + payload - ) - await asyncio.sleep(0) - - assert received_one == [b"fanout"] - assert received_two == [b"fanout"] - - await sub_one.unsubscribe() - fixture.client.schedule().connection.get_multiplexer().dispatch( # type: ignore[attr-defined] - 705, sub_two._sub_id.to_bytes(8, "big") + payload - ) - await asyncio.sleep(0) - assert received_one == [b"fanout"] - assert received_two == [b"fanout", b"fanout"] - - await sub_two.unsubscribe() - finally: - await fixture.close() diff --git a/tests/integration/test_stream.py b/tests/integration/test_stream.py deleted file mode 100644 index bf86002..0000000 --- a/tests/integration/test_stream.py +++ /dev/null @@ -1,108 +0,0 @@ -from __future__ import annotations - -import asyncio - -import pytest - -from fitz_py import ( - StreamCommitMode, - StreamCommitNotification, - StreamFilterClause, - StreamFilteredReason, - StreamFilterSet, - StreamReadItemKind, -) -from tests.integration.fixture.fixture import IntegrationFixture, unique_route - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_stream_round_trip(transport: str, auth_mode: str) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("stream") - session = await fixture.client.stream().begin(route) - await session.append(0, b"first") - await session.append(1, b"second") - await session.commit() - - records = await fixture.client.stream().read(route, start_offset=0, limit=10) - assert [record.body for record in records[:2]] == [b"first", b"second"] - finally: - await fixture.close() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_stream_filtered_round_trip(transport: str, auth_mode: str) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("stream") - session = await fixture.client.stream().begin(route) - await session.append(0, b"alpha", discriminator="proj.alpha") - await session.append(1, b"beta", discriminator="audit.beta") - await session.commit() - - stream_filter = StreamFilterSet( - clauses=[StreamFilterClause(kind="Equals", value="proj.alpha")] - ) - records = await fixture.client.stream().read( - route, start_offset=0, limit=10, stream_filter=stream_filter - ) - page = await fixture.client.stream().read_page( - route, start_offset=0, limit=10, stream_filter=stream_filter - ) - - assert [record.body for record in records] == [b"alpha"] - assert page.cursor.last_resource_offset == 1 - assert page.cursor.has_more is False - assert len(page.items) == 2 - assert page.items[0].kind is StreamReadItemKind.EVENT - assert page.items[0].record is not None - assert page.items[0].record.body == b"alpha" - assert page.items[1].kind is StreamReadItemKind.FILTERED - assert page.items[1].offset == 1 - assert page.items[1].reason is StreamFilteredReason.SERVER_FILTER - finally: - await fixture.close() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_stream_commit_notification_shape(transport: str, auth_mode: str) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - route = unique_route("stream") - notifications: list[StreamCommitNotification] = [] - delivered = asyncio.Event() - - async def handler(notification: StreamCommitNotification) -> None: - notifications.append(notification) - delivered.set() - - subscription = await fixture.client.stream().subscribe(route, handler) - - session = await fixture.client.stream().begin(route) - await session.append(0, b"notify") - await session.commit(StreamCommitMode.SYNC) - - await asyncio.wait_for(delivered.wait(), timeout=5) - assert len(notifications) == 1 - - notification = notifications[0] - assert notification.route == route - assert notification.event == "committed" - assert notification.first_resource_offset == 0 - assert notification.last_resource_offset == 0 - assert notification.first_area_offset == 0 - assert notification.last_area_offset == 0 - assert notification.first_realm_offset >= notification.first_area_offset - assert notification.last_realm_offset >= notification.first_realm_offset - assert notification.batch_size == 1 - - await subscription.unsubscribe() - finally: - await fixture.close() diff --git a/tests/integration/test_transport.py b/tests/integration/test_transport.py deleted file mode 100644 index f397c81..0000000 --- a/tests/integration/test_transport.py +++ /dev/null @@ -1,17 +0,0 @@ -from __future__ import annotations - -import pytest - -from fitz_py import ConnectionState -from tests.integration.fixture.fixture import IntegrationFixture - - -@pytest.mark.asyncio -@pytest.mark.parametrize("transport", ["tcp", "ws"]) -@pytest.mark.parametrize("auth_mode", ["anonymous", "valid_jwt"]) -async def test_connects_to_broker(transport: str, auth_mode: str) -> None: - fixture = await IntegrationFixture.connect_or_fail(transport, auth_mode) # type: ignore[arg-type] - try: - assert fixture.client.state is ConnectionState.AUTHENTICATED - finally: - await fixture.close() diff --git a/tests/unit/test_client.py b/tests/unit/test_client.py deleted file mode 100644 index 264aedb..0000000 --- a/tests/unit/test_client.py +++ /dev/null @@ -1,44 +0,0 @@ -from __future__ import annotations - -import pytest - -from fitz_py import Client, ClientConfig, ConnectionState - - -def test_client_starts_disconnected() -> None: - client = Client(ClientConfig(url="ws://localhost:4190/ws")) - assert client.state is ConnectionState.DISCONNECTED - - -def test_client_defaults_max_in_flight_requests() -> None: - client = Client(ClientConfig(url="ws://localhost:4190/ws")) - assert client._config.max_in_flight_requests == 256 - - -def test_client_preserves_max_in_flight_requests_override() -> None: - client = Client( - ClientConfig(url="ws://localhost:4190/ws", max_in_flight_requests=12) - ) - assert client._config.max_in_flight_requests == 12 - - -@pytest.mark.asyncio -async def test_client_async_context_manager_connects_and_closes( - monkeypatch: pytest.MonkeyPatch, -) -> None: - client = Client(ClientConfig(url="ws://localhost:4190/ws")) - calls: list[str] = [] - - async def fake_connect(self: Client) -> None: - calls.append("connect") - - async def fake_close(self: Client) -> None: - calls.append("close") - - monkeypatch.setattr(Client, "connect", fake_connect) - monkeypatch.setattr(Client, "close", fake_close) - - async with client as active: - assert active is client - - assert calls == ["connect", "close"] diff --git a/tests/unit/test_connection_contract.py b/tests/unit/test_connection_contract.py deleted file mode 100644 index c995a46..0000000 --- a/tests/unit/test_connection_contract.py +++ /dev/null @@ -1,202 +0,0 @@ -from __future__ import annotations - -import asyncio - -import pytest - -from fitz_py.connection import Connection -from fitz_py.errors import AuthenticationError, TransportError -from fitz_py.protocol.frame import FrameCodec -from fitz_py.types import ConnectionState - - -class _FakeTransport: - def __init__(self) -> None: - self.connected = False - self.sent: list[bytes] = [] - self._pending_read: asyncio.Future[bytes] | None = None - - async def connect(self) -> None: - self.connected = True - - async def close(self) -> None: - if self._pending_read is not None and not self._pending_read.done(): - self._pending_read.set_exception(TransportError("closed")) - self.connected = False - return None - - async def send(self, data: bytes) -> None: - self.sent.append(data) - - async def receive(self) -> bytes: - loop = asyncio.get_running_loop() - self._pending_read = loop.create_future() - return await self._pending_read - - def respond(self, data: bytes) -> None: - if self._pending_read is not None and not self._pending_read.done(): - self._pending_read.set_result(data) - self._pending_read = None - - def get_url(self) -> str: - return "ws://example.test" - - -class _DelayedCloseTransport(_FakeTransport): - async def receive(self) -> bytes: - await asyncio.sleep(0.2) - raise TransportError("TCP connection closed") - - -class _ControllableTransport(_FakeTransport): - def __init__(self) -> None: - super().__init__() - self.connect_started = False - self._pending_read: asyncio.Future[bytes] | None = None - - async def connect(self) -> None: - self.connect_started = True - await super().connect() - - async def receive(self) -> bytes: - loop = asyncio.get_running_loop() - self._pending_read = loop.create_future() - return await self._pending_read - - def fail(self, exc: Exception) -> None: - if self._pending_read is not None and not self._pending_read.done(): - self.connected = False - self._pending_read.set_exception(exc) - - async def close(self) -> None: - if self._pending_read is not None and not self._pending_read.done(): - self._pending_read.set_exception(TransportError("closed")) - await super().close() - - -async def _empty_token() -> str: - return "" - - -def test_connection_state_exposes_connected() -> None: - assert ConnectionState.CONNECTED.value == "CONNECTED" - - -@pytest.mark.asyncio -async def test_connection_emits_connected_before_authenticated() -> None: - seen: list[ConnectionState] = [] - connection = Connection(lambda: _FakeTransport(), _empty_token, auth_settle_delay_ms=0) - original_set_state = connection._set_state - - def record(state: ConnectionState) -> None: - seen.append(state) - original_set_state(state) - - connection._set_state = record # type: ignore[method-assign] - - try: - await connection.connect() - - assert seen[:4] == [ - ConnectionState.CONNECTING, - ConnectionState.CONNECTED, - ConnectionState.AUTHENTICATING, - ConnectionState.AUTHENTICATED, - ] - finally: - await connection.close() - - -@pytest.mark.asyncio -async def test_connection_notifies_disconnect_listeners_on_close() -> None: - connection = Connection(lambda: _FakeTransport(), _empty_token) - seen: list[str] = [] - - connection.on_disconnect(lambda: seen.append("disconnect")) - - await connection.close() - - assert seen == ["disconnect"] - - -@pytest.mark.asyncio -async def test_connection_surfaces_delayed_auth_close_as_authentication_error() -> None: - connection = Connection( - lambda: _DelayedCloseTransport(), - _empty_token, - auth_settle_delay_ms=500, - ) - - try: - with pytest.raises(AuthenticationError, match="TCP connection closed"): - await connection.connect() - finally: - await connection.close() - - -@pytest.mark.asyncio -async def test_connection_does_not_reconnect_after_close_during_reconnect_backoff() -> None: - first = _ControllableTransport() - second = _ControllableTransport() - transports = [first, second] - - def factory() -> _ControllableTransport: - return transports.pop(0) - - connection = Connection( - factory, - _empty_token, - auth_settle_delay_ms=0, - reconnect_enabled=True, - reconnect_max_attempts=1, - reconnect_backoff_ms=50, - reconnect_max_backoff_ms=50, - ) - - await connection.connect() - first.fail(TransportError("boom")) - - async def wait_for_reconnecting() -> None: - while connection.get_state() is not ConnectionState.RECONNECTING: - await asyncio.sleep(0.01) - - await asyncio.wait_for(wait_for_reconnecting(), timeout=1) - - await connection.close() - await asyncio.sleep(0.075) - - assert second.connect_started is False - assert connection.get_state() is ConnectionState.CLOSED - - -@pytest.mark.asyncio -async def test_connection_bounds_outbound_requests_to_configured_limit() -> None: - transport = _FakeTransport() - connection = Connection( - lambda: transport, - _empty_token, - auth_settle_delay_ms=0, - max_in_flight_requests=1, - ) - - await connection.connect() - - first = asyncio.create_task(connection.request(77, b"first")) - await asyncio.wait_for(_wait_for_sent_count(transport, 2), timeout=1) - - second = asyncio.create_task(connection.request(77, b"second")) - - with pytest.raises(asyncio.TimeoutError): - await asyncio.wait_for(second, timeout=0.05) - - assert len(transport.sent) == 2 - - transport.respond(FrameCodec.encode_frame(77, b"ok")) - assert await first == b"ok" - - await connection.close() - - -async def _wait_for_sent_count(transport: _FakeTransport, count: int) -> None: - while len(transport.sent) < count: - await asyncio.sleep(0.01) diff --git a/tests/unit/test_contracts.py b/tests/unit/test_contracts.py new file mode 100644 index 0000000..21d065b --- /dev/null +++ b/tests/unit/test_contracts.py @@ -0,0 +1,197 @@ +from __future__ import annotations + +import asyncio + +import pytest + +import fitz_py +from fitz_py._runtime import AsyncSubscription, RequestGate +from fitz_py.domains.queue import QueueClient +from fitz_py.domains.rpc import RpcClient +from fitz_py.domains.schedule import DeliveryMode, ScheduleClient +from fitz_py.errors import ( + FitzConnectionError, + FitzTimeoutError, + QueueError, + RequestQueueFullError, + SubscriptionBackpressureError, +) +from fitz_py.multiplexer import Multiplexer +from fitz_py.protocol.buffer import BufferReader, BufferWriter +from fitz_py.protocol.frame import FrameCodec, FrameParser +from fitz_py.protocol.messages import MSG_QUEUE_RESERVE, MSG_RPC_REQUEST, MSG_SCHEDULE_CREATE +from fitz_py.protocol.response import parse_response +from fitz_py.types import ClientConfig, ConcurrencyLimits + + +class FakeConnection: + def __init__(self, responses: dict[int, bytes] | None = None) -> None: + self.config = ClientConfig(url="tcp://localhost:1") + self.generation = 1 + self.responses = responses or {} + self.sent: list[tuple[int, bytes]] = [] + self.notifications = {} + + async def request(self, message_type: int, payload: bytes) -> bytes: + self.sent.append((message_type, payload)) + return self.responses[message_type] + + async def send(self, message_type: int, payload: bytes) -> None: + self.sent.append((message_type, payload)) + + def register_notification_handler(self, message_type, handler) -> None: + self.notifications[message_type] = handler + + def register_push_classifier(self, *_args) -> None: ... + def on_reconnect(self, *_args, **_kwargs): + return lambda: None + + def on_disconnect(self, *_args): + return lambda: None + + def dispatch_async(self, work): + asyncio.create_task(work()) + return True + + +def test_clean_break_public_surface() -> None: + assert fitz_py.ClientConfig is ClientConfig + assert not hasattr(fitz_py, "ErrKvKeyNotFound") + assert not hasattr(fitz_py, "ReconnectOptions") + + +def test_config_is_frozen_and_validated() -> None: + config = ClientConfig(url="tcp://localhost:7777", limits=ConcurrencyLimits(max_in_flight=2)) + assert config.limits.max_in_flight == 2 + with pytest.raises(ValueError): + ClientConfig(url="") + + +def test_buffer_round_trip() -> None: + writer = BufferWriter(1) + writer.write_u8(7) + writer.write_u64_be(2**63) + writer.write_string("hello") + reader = BufferReader(writer.build()) + assert (reader.read_u8(), reader.read_u64_be(), reader.read_string()) == (7, 2**63, "hello") + assert reader.is_eof() + + +def test_frame_parser_handles_fragmentation() -> None: + frame = FrameCodec.encode_frame(42, b"payload") + parser = FrameParser() + assert parser.parse_frames(frame[:3]) == [] + parsed = parser.parse_frames(frame[3:]) + assert [(item.message_type, item.payload) for item in parsed] == [(42, b"payload")] + + +def test_response_envelopes_are_strict() -> None: + assert parse_response(b"\0data").data == b"data" + writer = BufferWriter() + writer.write_u8(1) + writer.write_u32_be(4005) + writer.write_string("full") + response = parse_response(writer.build()) + assert (response.error_code, response.error) == (4005, "full") + + +@pytest.mark.asyncio +async def test_request_gate_is_bounded() -> None: + gate = RequestGate(1, 0) + release = await gate.acquire() + with pytest.raises(RequestQueueFullError): + await gate.acquire() + release() + + +@pytest.mark.asyncio +async def test_subscription_backpressure_is_explicit() -> None: + async def close() -> None: ... + + subscription = AsyncSubscription("x", 1, close) + assert subscription.push(1) + assert not subscription.push(2) + with pytest.raises(SubscriptionBackpressureError): + await anext(subscription) + + +@pytest.mark.asyncio +async def test_multiplexer_tombstone_consumes_late_reply() -> None: + mux = Multiplexer() + mux.set_connected() + sent = asyncio.Event() + + async def send(_frame: bytes) -> None: + sent.set() + + task = asyncio.create_task(mux.request(9, b"x", send, 0.01)) + await sent.wait() + with pytest.raises(FitzTimeoutError): + await task + mux.dispatch(9, b"late") + next_task = asyncio.create_task(mux.request(9, b"x", send, 1)) + await asyncio.sleep(0) + mux.dispatch(9, b"current") + assert await next_task == b"current" + + +@pytest.mark.asyncio +async def test_multiplexer_disconnect_fails_waiters() -> None: + mux = Multiplexer() + mux.set_connected() + task = asyncio.create_task(mux.request(1, b"", lambda _: asyncio.sleep(0), 1)) + await asyncio.sleep(0) + mux.set_disconnected() + with pytest.raises(FitzConnectionError): + await task + + +@pytest.mark.asyncio +async def test_queue_reserve_encodes_broker_wait_and_route_results() -> None: + body = BufferWriter() + body.write_u8(0) + body.write_u32_be(1) + body.write_route("queue://r/a/one") + body.write_u64_be(10) + body.write_u64_be(11) + body.write_u32_be(3) + body.write_bytes(b"msg") + connection = FakeConnection({MSG_QUEUE_RESERVE: body.build()}) + items = await QueueClient(connection).reserve("queue://r/a/*", lease=30, wait=5) + assert (items[0].route, items[0].body) == ("queue://r/a/one", b"msg") + reader = BufferReader(connection.sent[0][1]) + reader.read_route() + reader.read_u64_be() + reader.read_u8() + reader.read_u8() + assert reader.read_u64_be() == 5 + + +@pytest.mark.asyncio +async def test_rpc_call_is_fire_and_forget_without_ack() -> None: + connection = FakeConnection() + call = await RpcClient(connection).call("rpc://r/a/work", b"input") + assert connection.sent[0][0] == MSG_RPC_REQUEST + await call.aclose() + + +@pytest.mark.asyncio +async def test_schedule_create_encodes_delivery_mode() -> None: + connection = FakeConnection({MSG_SCHEDULE_CREATE: b"\0"}) + route = "schedule://r/a/jobs/nightly" + assert ( + await ScheduleClient(connection).create( + route, "0 0 * * *", delivery_mode=DeliveryMode.BROADCAST + ) + == route + ) + reader = BufferReader(connection.sent[0][1]) + assert reader.read_route() == route + assert reader.read_string() == "0 0 * * *" + assert reader.read_u8() == 0 + + +def test_invalid_queue_route_fails_before_io() -> None: + connection = FakeConnection() + with pytest.raises(QueueError): + asyncio.run(QueueClient(connection).enqueue("queue://bad", b"x")) diff --git a/tests/unit/test_errors.py b/tests/unit/test_errors.py deleted file mode 100644 index 89c8e09..0000000 --- a/tests/unit/test_errors.py +++ /dev/null @@ -1,24 +0,0 @@ -from __future__ import annotations - -from fitz_py import ( - ErrKvKeyNotFound, - ErrLeaseHeld, - ErrRpcTimeout, - TimeoutError, - TransportError, - is_retryable, -) -from fitz_py.errors import kv_error, lease_error, rpc_error - - -def test_named_domain_errors_are_mapped() -> None: - assert isinstance(kv_error("missing", 4), ErrKvKeyNotFound) - assert isinstance(lease_error("held", 1), ErrLeaseHeld) - assert isinstance(rpc_error("timeout", 1), ErrRpcTimeout) - - -def test_is_retryable_matches_domain_and_transport_rules() -> None: - assert is_retryable(TimeoutError("timeout")) is True - assert is_retryable(TransportError("transport")) is True - assert is_retryable(kv_error("missing", 4)) is True - assert is_retryable(kv_error("conflict", 3)) is False diff --git a/tests/unit/test_frame.py b/tests/unit/test_frame.py deleted file mode 100644 index 561c7f4..0000000 --- a/tests/unit/test_frame.py +++ /dev/null @@ -1,21 +0,0 @@ -from fitz_py.protocol.frame import FrameCodec, FrameParser - - -def test_frame_codec_round_trip() -> None: - payload = b"abc123" - encoded = FrameCodec.encode_frame(302, payload) - decoded = FrameCodec.decode_frame(encoded) - assert decoded.message_type == 302 - assert decoded.payload == payload - - -def test_frame_parser_handles_partial_input() -> None: - encoded = FrameCodec.encode_frame(100, b"payload") - parser = FrameParser() - - assert parser.parse_frames(encoded[:2]) == [] - frames = parser.parse_frames(encoded[2:]) - - assert len(frames) == 1 - assert frames[0].message_type == 100 - assert frames[0].payload == b"payload" diff --git a/tests/unit/test_kv_transaction.py b/tests/unit/test_kv_transaction.py deleted file mode 100644 index 7a26dbd..0000000 --- a/tests/unit/test_kv_transaction.py +++ /dev/null @@ -1,58 +0,0 @@ -from __future__ import annotations - -from collections.abc import Callable - -import pytest - -from fitz_py.domains.kv import KvTransaction -from fitz_py.errors import ErrKvOperationNotAllowed - - -class _FakeConnection: - def __init__(self) -> None: - self.requests: list[tuple[int, bytes]] = [] - self.disconnect_handlers: list[Callable[[], None]] = [] - - async def request(self, message_type: int, payload: bytes) -> bytes: - self.requests.append((message_type, payload)) - return b"\x00" - - def on_disconnect(self, handler: Callable[[], None]) -> None: - self.disconnect_handlers.append(handler) - - def emit_disconnect(self) -> None: - for handler in list(self.disconnect_handlers): - handler() - - -@pytest.mark.asyncio -async def test_kv_transaction_rejects_mutation_after_commit() -> None: - conn = _FakeConnection() - tx = KvTransaction(conn, "kv://tests/app/resource", 42) - - await tx.commit() - - with pytest.raises(ErrKvOperationNotAllowed, match="already committed"): - await tx.put(b"k", b"v") - - -@pytest.mark.asyncio -async def test_kv_transaction_rejects_commit_after_rollback() -> None: - conn = _FakeConnection() - tx = KvTransaction(conn, "kv://tests/app/resource", 42) - - await tx.rollback() - - with pytest.raises(ErrKvOperationNotAllowed, match="already rolled back"): - await tx.commit() - - -@pytest.mark.asyncio -async def test_kv_transaction_invalidates_on_disconnect() -> None: - conn = _FakeConnection() - tx = KvTransaction(conn, "kv://tests/app/resource", 42) - - conn.emit_disconnect() - - with pytest.raises(ErrKvOperationNotAllowed, match="already disconnected"): - await tx.put(b"k", b"v") diff --git a/tests/unit/test_multiplexer.py b/tests/unit/test_multiplexer.py deleted file mode 100644 index adfd151..0000000 --- a/tests/unit/test_multiplexer.py +++ /dev/null @@ -1,59 +0,0 @@ -from __future__ import annotations - -import asyncio - -import pytest - -from fitz_py.errors import TimeoutError -from fitz_py.multiplexer import Multiplexer - - -@pytest.mark.asyncio -async def test_multiplexer_matches_fifo_request_response() -> None: - mux = Multiplexer() - mux.set_connected() - sent: list[bytes] = [] - - async def send(data: bytes) -> None: - sent.append(data) - - task = asyncio.create_task(mux.request(100, b"frame", send, 1000)) - await asyncio.sleep(0) - mux.dispatch(100, b"response") - assert await task == b"response" - assert sent == [b"frame"] - - -def test_multiplexer_ignores_optional_response() -> None: - mux = Multiplexer() - mux.set_connected() - mux.expect_optional_response(500) - mux.dispatch(500, b"ok") - - -@pytest.mark.asyncio -async def test_multiplexer_clears_pending_on_send_failure() -> None: - mux = Multiplexer() - mux.set_connected() - - async def send(_data: bytes) -> None: - raise RuntimeError("send failed") - - with pytest.raises(RuntimeError, match="send failed"): - await mux.request(101, b"frame", send, 1000) - - assert 101 not in mux._pending # type: ignore[attr-defined] - - -@pytest.mark.asyncio -async def test_multiplexer_times_out_and_clears_pending() -> None: - mux = Multiplexer() - mux.set_connected() - - async def send(_data: bytes) -> None: - return None - - with pytest.raises(TimeoutError, match="Request timeout"): - await mux.request(102, b"frame", send, 1) - - assert 102 not in mux._pending # type: ignore[attr-defined] diff --git a/tests/unit/test_public_surface.py b/tests/unit/test_public_surface.py deleted file mode 100644 index bdbbbd6..0000000 --- a/tests/unit/test_public_surface.py +++ /dev/null @@ -1,124 +0,0 @@ -from fitz_py import __all__ as public_exports -from fitz_py.domains.lease import Lease, LeaseSubscription -from fitz_py.domains.queue import QueueItem, QueueSubscription -from fitz_py.domains.rpc import InboundRpcRequest - - -def test_public_exports_snapshot() -> None: - assert public_exports == [ - "AuthenticationError", - "Client", - "ClientConfig", - "CodecError", - "ConnectionError", - "ConnectionState", - "ErrKvConflictingWrite", - "ErrKvKeyNotFound", - "ErrKvLeaseExpired", - "ErrKvOperationNotAllowed", - "ErrKvTransactionAborted", - "ErrLeaseHeld", - "ErrLeaseInvalidToken", - "ErrLeaseNotFound", - "ErrNoticeGeneral", - "ErrQueueFull", - "ErrQueueInvalidDelay", - "ErrQueueInvalidToken", - "ErrQueueMessageNotFound", - "ErrQueueNotFound", - "ErrRpcHandlerError", - "ErrRpcHandlerNotFound", - "ErrRpcInvalidRequest", - "ErrRpcTimeout", - "ErrScheduleInvalidCron", - "ErrScheduleInvalidDelay", - "ErrScheduleInvalidTimestamp", - "ErrScheduleNotFound", - "ErrScheduleTaskNotFound", - "ErrStreamExpectedOffsetMismatch", - "ErrStreamFull", - "ErrStreamInvalidOffset", - "ErrStreamNotFound", - "ErrStreamOffsetOutOfRange", - "ErrStreamSessionClosed", - "ErrStreamSessionNotFound", - "FitzError", - "InboundRpcRequest", - "KVDurability", - "KVMode", - "KvClient", - "KvGetResult", - "KvError", - "KvPair", - "KvScanResult", - "KvTransaction", - "Lease", - "LeaseClient", - "LeaseError", - "LeaseHandler", - "LeaseInfo", - "LeaseSubscription", - "NoticeClient", - "NoticeError", - "NoticeHandler", - "NoticeMessage", - "NoticeSubscription", - "ProtocolError", - "QueueAvailabilityHandler", - "QueueClient", - "QueueError", - "QueueItem", - "QueueSubscription", - "ReconnectOptions", - "ResponseFrame", - "ResponseWriter", - "RpcClient", - "RpcError", - "RpcHandler", - "RpcSubscription", - "ScheduleClient", - "ScheduleEntry", - "ScheduleError", - "ScheduleHandler", - "ScheduleNotification", - "ScheduleSubscription", - "StreamClient", - "StreamFilterClause", - "StreamFilterSet", - "StreamCommitMode", - "StreamCommitNotification", - "StreamError", - "StreamFilteredReason", - "StreamHandler", - "StreamMetadata", - "StreamReadCursor", - "StreamReadItem", - "StreamReadItemKind", - "StreamReadPage", - "StreamRecord", - "StreamSession", - "StreamSubscription", - "TimeoutError", - "TokenProvider", - "TransportError", - "TransportType", - "is_retryable", - ] - - -def test_public_domain_types_hide_wire_identifiers() -> None: - request = InboundRpcRequest(route="rpc://realm/area/task", reply_route="", body=b"ping") - assert not hasattr(request, "correlation_id") - - queue_item = QueueItem(route="queue://realm/area/item", _id=1, _token=2, body=b"body", _client=object()) - assert not hasattr(queue_item, "id") - assert not hasattr(queue_item, "token") - - lease = Lease(route="lease://realm/area/lock", _token=3, _client=object()) - assert lease.token == 3 - - queue_subscription = QueueSubscription(4, "queue://realm/area/*", lambda _sub_id: None) - assert not hasattr(queue_subscription, "sub_id") - - lease_subscription = LeaseSubscription(5, "lease://realm/area/*", lambda _sub_id: None) - assert not hasattr(lease_subscription, "sub_id") diff --git a/tests/unit/test_response.py b/tests/unit/test_response.py deleted file mode 100644 index 17809de..0000000 --- a/tests/unit/test_response.py +++ /dev/null @@ -1,22 +0,0 @@ -import pytest - -from fitz_py.errors import ProtocolError -from fitz_py.protocol.buffer import BufferWriter -from fitz_py.protocol.response import assert_success, parse_standard_response - - -def test_parse_standard_response_success() -> None: - writer = BufferWriter() - writer.write_u8(0) - writer.write_string("ok") - parsed = parse_standard_response(writer.build()) - assert parsed.success is True - assert parsed.data - - -def test_assert_success_raises_on_error() -> None: - writer = BufferWriter() - writer.write_u8(1) - writer.write_string("bad") - with pytest.raises(ProtocolError): - assert_success(writer.build(), "TEST") diff --git a/tests/unit/test_route_validation.py b/tests/unit/test_route_validation.py deleted file mode 100644 index c858530..0000000 --- a/tests/unit/test_route_validation.py +++ /dev/null @@ -1,292 +0,0 @@ -from __future__ import annotations - -import pytest - -from fitz_py.domains.kv import KvClient -from fitz_py.domains.lease import LeaseClient -from fitz_py.domains.notice import NoticeClient -from fitz_py.domains.queue import QueueClient -from fitz_py.domains.rpc import RpcClient -from fitz_py.domains.schedule import ScheduleClient -from fitz_py.domains.stream import StreamClient -from fitz_py.protocol.messages import ( - MSG_KV_BEGIN, - MSG_LEASE_ACQUIRE, - MSG_LEASE_SUBSCRIBE, - MSG_NOTICE_PUBLISH, - MSG_NOTICE_SUBSCRIBE, - MSG_QUEUE_ENQUEUE, - MSG_QUEUE_SUBSCRIBE, - MSG_RPC_REQUEST, - MSG_RPC_SUBSCRIBE_WORKER, - MSG_SCHEDULE_CREATE, - MSG_SCHEDULE_SUBSCRIBE, - MSG_STREAM_BEGIN, - MSG_STREAM_READ, - MSG_STREAM_SUBSCRIBE, -) - - -class _FakeMultiplexer: - def expect_optional_response(self, _message_type: int): - return lambda: None - - -class _FakeConnection: - def __init__(self, response: bytes) -> None: - self.response = response - self.requests: list[tuple[int, bytes]] = [] - self.notification_handlers: dict[int, object] = {} - self.reconnect_handlers: list[object] = [] - self._multiplexer = _FakeMultiplexer() - - async def request(self, message_type: int, payload: bytes) -> bytes: - self.requests.append((message_type, payload)) - return self.response - - async def send_fire_and_forget(self, message_type: int, payload: bytes) -> None: - self.requests.append((message_type, payload)) - - def on_reconnect(self, handler) -> None: - self.reconnect_handlers.append(handler) - return None - - def register_notification_handler(self, message_type: int, handler) -> None: - self.notification_handlers[message_type] = handler - - def get_multiplexer(self) -> _FakeMultiplexer: - return self._multiplexer - - -@pytest.mark.asyncio -async def test_kv_begin_accepts_exact_three_segment_route() -> None: - connection = _FakeConnection(b"\x00" + (42).to_bytes(8, "big")) - client = KvClient(connection) - - transaction = await client.begin("kv://example/app/users", durability="sync") - - assert transaction is not None - assert connection.requests[0][0] == MSG_KV_BEGIN - assert len(connection.requests) == 1 - - -@pytest.mark.asyncio -async def test_kv_begin_forwards_short_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00" + (42).to_bytes(8, "big")) - client = KvClient(connection) - - transaction = await client.begin("kv://example/app", durability="sync") - - assert transaction is not None - assert connection.requests[0][0] == MSG_KV_BEGIN - - -@pytest.mark.asyncio -async def test_kv_begin_forwards_wrong_scheme_without_local_validation() -> None: - connection = _FakeConnection(b"\x00" + (42).to_bytes(8, "big")) - client = KvClient(connection) - - transaction = await client.begin("queue://example/app/users", durability="sync") - - assert transaction is not None - assert connection.requests[0][0] == MSG_KV_BEGIN - - -@pytest.mark.asyncio -async def test_lease_acquire_accepts_exact_three_segment_route() -> None: - connection = _FakeConnection(b"\x00\x01" + (42).to_bytes(8, "big")) - client = LeaseClient(connection) - - lease = await client.acquire("lease://example/app/leader", 30) - - assert lease.route == "lease://example/app/leader" - assert lease.token == 42 - assert connection.requests[0][0] == MSG_LEASE_ACQUIRE - assert len(connection.requests) == 1 - - -@pytest.mark.asyncio -async def test_lease_acquire_forwards_short_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00\x01" + (42).to_bytes(8, "big")) - client = LeaseClient(connection) - - lease = await client.acquire("lease://example/app", 30) - - assert lease.route == "lease://example/app" - assert connection.requests[0][0] == MSG_LEASE_ACQUIRE - - -@pytest.mark.asyncio -async def test_lease_acquire_forwards_empty_segment_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00\x01" + (42).to_bytes(8, "big")) - client = LeaseClient(connection) - - lease = await client.acquire("lease://example//leader", 30) - - assert lease.route == "lease://example//leader" - assert connection.requests[0][0] == MSG_LEASE_ACQUIRE - - -@pytest.mark.asyncio -async def test_lease_subscribe_forwards_wildcard_pattern_without_local_validation() -> None: - connection = _FakeConnection(b"\x00\x01" + (42).to_bytes(8, "big")) - client = LeaseClient(connection) - - subscription = await client.subscribe("lease://example/**", lambda notification: None) - - assert subscription.pattern == "lease://example/**" - assert connection.requests[0][0] == MSG_LEASE_SUBSCRIBE - - -@pytest.mark.asyncio -async def test_queue_enqueue_forwards_wildcard_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00") - client = QueueClient(connection) - - message_id = await client.enqueue("queue://example/app/*", b"payload") - - assert message_id == 0 - assert connection.requests[0][0] == MSG_QUEUE_ENQUEUE - - -@pytest.mark.asyncio -async def test_queue_subscribe_accepts_realm_wildcard_pattern() -> None: - connection = _FakeConnection(b"\x00\x01" + (7).to_bytes(8, "big")) - client = QueueClient(connection) - - subscription = await client.subscribe("queue://example/**", lambda notification: None) - - assert subscription.pattern == "queue://example/**" - assert connection.requests[0][0] == MSG_QUEUE_SUBSCRIBE - - -@pytest.mark.asyncio -async def test_notice_publish_forwards_wildcard_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00") - client = NoticeClient(connection) - - await client.publish("notice://example/**", b"payload") - - assert connection.requests[0][0] == MSG_NOTICE_PUBLISH - - -@pytest.mark.asyncio -async def test_notice_subscribe_accepts_realm_wildcard_pattern() -> None: - connection = _FakeConnection(b"\x00\x01" + (7).to_bytes(8, "big")) - client = NoticeClient(connection) - - subscription = await client.subscribe("notice://example/**", lambda notification: None) - - assert subscription.pattern == "notice://example/**" - assert connection.requests[0][0] == MSG_NOTICE_SUBSCRIBE - - -@pytest.mark.asyncio -async def test_rpc_call_forwards_wildcard_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00") - client = RpcClient(connection) - - iterator = await client.call("rpc://example/*", b"payload") - - assert iterator is not None - assert connection.requests[0][0] == MSG_RPC_REQUEST - - -@pytest.mark.asyncio -async def test_rpc_register_worker_forwards_wildcard_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00") - client = RpcClient(connection) - - subscription = await client.register_worker("rpc://example/**", lambda request, writer: None) - - assert subscription.route == "rpc://example/**" - assert connection.requests[0][0] == MSG_RPC_SUBSCRIBE_WORKER - - -@pytest.mark.asyncio -async def test_stream_begin_forwards_wildcard_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00\x01" + (42).to_bytes(8, "big")) - client = StreamClient(connection) - - session = await client.begin("stream://example/app/*") - - assert session is not None - assert connection.requests[0][0] == MSG_STREAM_BEGIN - - -@pytest.mark.asyncio -async def test_stream_subscribe_forwards_wildcard_pattern_without_local_validation() -> None: - connection = _FakeConnection(b"\x00\x01" + (7).to_bytes(8, "big")) - client = StreamClient(connection) - - subscription = await client.subscribe("stream://example/area/**", lambda notification: None) - - assert subscription.pattern == "stream://example/area/**" - assert connection.requests[0][0] == MSG_STREAM_SUBSCRIBE - - -@pytest.mark.asyncio -async def test_stream_read_accepts_realm_wildcard_pattern() -> None: - connection = _FakeConnection(b"\x00") - client = StreamClient(connection) - - records = await client.read("stream://example/**", 0) - - assert records == [] - assert connection.requests[0][0] == MSG_STREAM_READ - - -@pytest.mark.asyncio -async def test_schedule_create_accepts_exact_four_segment_route() -> None: - connection = _FakeConnection(b"\x00") - client = ScheduleClient(connection) - - route = await client.create("schedule://example/app/jobs/run", "0 0 * * *") - - assert route == "schedule://example/app/jobs/run" - assert connection.requests[0][0] == MSG_SCHEDULE_CREATE - assert len(connection.requests) == 1 - - -@pytest.mark.asyncio -async def test_schedule_create_forwards_short_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00") - client = ScheduleClient(connection) - - route = await client.create("schedule://example/app", "0 0 * * *") - - assert route == "schedule://example/app" - assert connection.requests[0][0] == MSG_SCHEDULE_CREATE - - -@pytest.mark.asyncio -async def test_schedule_create_forwards_wrong_scheme_without_local_validation() -> None: - connection = _FakeConnection(b"\x00") - client = ScheduleClient(connection) - - route = await client.create("queue://example/app/jobs/run", "0 0 * * *") - - assert route == "queue://example/app/jobs/run" - assert connection.requests[0][0] == MSG_SCHEDULE_CREATE - - -@pytest.mark.asyncio -async def test_schedule_create_forwards_empty_segment_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00") - client = ScheduleClient(connection) - - route = await client.create("schedule://example//jobs/run", "0 0 * * *") - - assert route == "schedule://example//jobs/run" - assert connection.requests[0][0] == MSG_SCHEDULE_CREATE - - -@pytest.mark.asyncio -async def test_schedule_subscribe_forwards_wildcard_route_without_local_validation() -> None: - connection = _FakeConnection(b"\x00\x01" + (7).to_bytes(8, "big")) - client = ScheduleClient(connection) - - subscription = await client.subscribe("schedule://example/app/*", lambda notification: None) - - assert subscription.pattern == "schedule://example/app/*" - assert connection.requests[0][0] == MSG_SCHEDULE_SUBSCRIBE diff --git a/tests/unit/test_stream_session.py b/tests/unit/test_stream_session.py deleted file mode 100644 index f80e0dc..0000000 --- a/tests/unit/test_stream_session.py +++ /dev/null @@ -1,257 +0,0 @@ -from __future__ import annotations - -import asyncio -import json -from collections.abc import Awaitable, Callable - -import pytest - -from fitz_py import ( - StreamCommitMode, - StreamCommitNotification, - StreamFilterClause, - StreamFilteredReason, - StreamFilterSet, - StreamReadItemKind, -) -from fitz_py.domains.stream import StreamClient, StreamSession -from fitz_py.errors import ErrStreamSessionClosed -from fitz_py.protocol.buffer import BufferReader, BufferWriter -from fitz_py.protocol.messages import ( - MSG_STREAM_APPEND, - MSG_STREAM_COMMIT, - MSG_STREAM_NOTIFY, - MSG_STREAM_READ, - MSG_STREAM_SUBSCRIBE, -) - - -class _FakeConnection: - def __init__(self) -> None: - self.requests: list[tuple[int, bytes]] = [] - self.notification_handlers: dict[int, Callable[[bytes], None]] = {} - self.reconnect_handlers: list[Callable[[], Awaitable[None]]] = [] - self.disconnect_handlers: list[Callable[[], None | Awaitable[None]]] = [] - self.responses: dict[int, bytes] = {} - - async def request(self, message_type: int, payload: bytes) -> bytes: - self.requests.append((message_type, payload)) - if message_type in self.responses: - return self.responses[message_type] - if message_type == MSG_STREAM_SUBSCRIBE: - return b"\x00\x01" + (7).to_bytes(8, "big") - return b"\x00" - - def register_notification_handler( - self, message_type: int, handler: Callable[[bytes], None] - ) -> None: - self.notification_handlers[message_type] = handler - - def on_reconnect(self, _handler: Callable[[], Awaitable[None]]) -> None: - self.reconnect_handlers.append(_handler) - return None - - def on_disconnect(self, handler: Callable[[], None | Awaitable[None]]) -> None: - self.disconnect_handlers.append(handler) - return None - - def emit_disconnect(self) -> None: - for handler in list(self.disconnect_handlers): - result = handler() - if asyncio.iscoroutine(result): - asyncio.create_task(result) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "mode, expected_mode", - [ - (None, 0), - (StreamCommitMode.SYNC, 1), - ], -) -async def test_stream_session_commit_encodes_mode( - mode: StreamCommitMode | None, expected_mode: int -) -> None: - connection = _FakeConnection() - session = StreamSession(connection, 42) - - if mode is None: - await session.commit() - else: - await session.commit(mode) - - assert len(connection.requests) == 1 - message_type, payload = connection.requests[0] - assert message_type == MSG_STREAM_COMMIT - - reader = BufferReader(payload) - assert reader.read_u64_be() == 42 - assert reader.read_u8() == expected_mode - assert reader.is_eof() - - -@pytest.mark.asyncio -async def test_stream_session_append_encodes_discriminator() -> None: - connection = _FakeConnection() - session = StreamSession(connection, 42) - - await session.append(12, b"entry", b"meta", "proj.alpha") - - assert connection.requests[0][0] == MSG_STREAM_APPEND - reader = BufferReader(connection.requests[0][1]) - assert reader.read_u64_be() == 42 - assert reader.read_u64_be() == 12 - assert reader.read_u32_be() == 5 - assert reader.read_bytes(5) == b"entry" - assert reader.read_u8() == 1 - assert reader.read_u32_be() == 4 - assert reader.read_bytes(4) == b"meta" - assert reader.read_u8() == 1 - assert reader.read_string() == "proj.alpha" - - -@pytest.mark.asyncio -async def test_stream_read_encodes_filter_payload() -> None: - connection = _FakeConnection() - client = StreamClient(connection) - - stream_filter = StreamFilterSet(clauses=[StreamFilterClause(kind="Equals", value="proj.alpha")]) - records = await client.read("stream://realm/area/resource", 5, 10, stream_filter) - - assert records == [] - assert connection.requests[0][0] == MSG_STREAM_READ - reader = BufferReader(connection.requests[0][1]) - assert reader.read_route() == "stream://realm/area/resource" - assert reader.read_u64_be() == 5 - assert reader.read_u64_be() == 10 - assert reader.read_u8() == 0 - assert reader.read_u8() == 1 - filter_length = reader.read_u32_be() - expected_filter = ( - (1).to_bytes(8, "little") - + (0).to_bytes(4, "little") - + (10).to_bytes(8, "little") - + b"proj.alpha" - ) - assert filter_length == len(expected_filter) - assert reader.read_bytes(filter_length) == expected_filter - assert reader.is_eof() - - -@pytest.mark.asyncio -async def test_stream_read_page_decodes_filtered_items_and_cursor() -> None: - connection = _FakeConnection() - client = StreamClient(connection) - - inner = bytearray() - inner.extend((3).to_bytes(4, "big")) - inner.extend(b"\x00") - inner.extend((41).to_bytes(8, "big")) - inner.extend(b"\x01") - inner.extend((51).to_bytes(8, "big")) - inner.extend(b"\x00") - inner.extend((5).to_bytes(4, "big")) - inner.extend(b"alpha") - inner.extend(b"\x00") - inner.extend((111).to_bytes(8, "big")) - inner.extend(b"\x01") - inner.extend((42).to_bytes(8, "big")) - inner.extend(b"\x01") - inner.extend(b"\x02") - inner.extend((43).to_bytes(8, "big")) - inner.extend((45).to_bytes(8, "big")) - inner.extend(b"\x02") - inner.extend((45).to_bytes(8, "big")) - inner.extend(b"\x01") - inner.extend((52).to_bytes(8, "big")) - inner.extend(b"\x00") - inner.extend(b"\x01") - - response = bytearray(b"\x00\x00") - response.extend(len(inner).to_bytes(4, "big")) - response.extend(inner) - connection.responses[MSG_STREAM_READ] = bytes(response) - - page = await client.read_page("stream://realm/area/resource", 0, 10) - - assert page.cursor.last_resource_offset == 45 - assert page.cursor.last_area_offset == 52 - assert page.cursor.last_realm_offset is None - assert page.cursor.has_more is True - assert len(page.items) == 3 - assert page.items[0].kind is StreamReadItemKind.EVENT - assert page.items[0].record is not None - assert page.items[0].record.offset == 41 - assert page.items[0].record.area_offset == 51 - assert page.items[0].record.body == b"alpha" - assert page.items[1].kind is StreamReadItemKind.FILTERED - assert page.items[1].offset == 42 - assert page.items[1].reason is StreamFilteredReason.SERVER_FILTER - assert page.items[2].kind is StreamReadItemKind.FILTERED_RANGE - assert page.items[2].from_offset == 43 - assert page.items[2].to_offset == 45 - assert page.items[2].reason is StreamFilteredReason.PERMISSION - - -@pytest.mark.asyncio -async def test_stream_session_invalidates_on_disconnect() -> None: - connection = _FakeConnection() - session = StreamSession(connection, 42) - - connection.emit_disconnect() - - with pytest.raises(ErrStreamSessionClosed, match="already disconnected"): - await session.append(0, b"payload") - - -@pytest.mark.asyncio -async def test_stream_subscribe_decodes_commit_notification() -> None: - connection = _FakeConnection() - client = StreamClient(connection) - route = "stream://realm/area/resource" - notifications: list[StreamCommitNotification] = [] - delivered = asyncio.Event() - - async def handler(notification: StreamCommitNotification) -> None: - notifications.append(notification) - delivered.set() - - subscription = await client.subscribe(route, handler) - assert subscription._sub_id == 7 - - writer = BufferWriter() - writer.write_u64_be(subscription._sub_id) - writer.write_route(route) - payload = json.dumps( - { - "event": "committed", - "first_resource_offset": 0, - "last_resource_offset": 0, - "first_area_offset": 0, - "last_area_offset": 0, - "first_realm_offset": 0, - "last_realm_offset": 0, - "batch_size": 1, - }, - separators=(",", ":"), - ).encode("utf-8") - writer.write_u32_be(len(payload)) - writer.write_bytes(payload) - - notification_handler = connection.notification_handlers[MSG_STREAM_NOTIFY] - notification_handler(writer.build()) - - await asyncio.wait_for(delivered.wait(), timeout=1) - assert len(notifications) == 1 - - notification = notifications[0] - assert notification.route == route - assert notification.event == "committed" - assert notification.first_resource_offset == 0 - assert notification.last_resource_offset == 0 - assert notification.first_area_offset == 0 - assert notification.last_area_offset == 0 - assert notification.first_realm_offset == 0 - assert notification.last_realm_offset == 0 - assert notification.batch_size == 1 diff --git a/tests/unit/test_websocket_transport.py b/tests/unit/test_websocket_transport.py deleted file mode 100644 index 85b2a24..0000000 --- a/tests/unit/test_websocket_transport.py +++ /dev/null @@ -1,33 +0,0 @@ -from __future__ import annotations - -import pytest - -from fitz_py.errors import TransportError -from fitz_py.transport.websocket import WebSocketTransport - - -class _FakeSocket: - def __init__(self, payload: bytes | str) -> None: - self._payload = payload - - async def recv(self) -> bytes | str: - return self._payload - - -@pytest.mark.asyncio -async def test_websocket_transport_returns_binary_frames_unchanged() -> None: - transport = WebSocketTransport("ws://example.test") - transport._socket = _FakeSocket(b"payload") - - data = await transport.receive() - - assert data == b"payload" - - -@pytest.mark.asyncio -async def test_websocket_transport_rejects_text_frames() -> None: - transport = WebSocketTransport("ws://example.test") - transport._socket = _FakeSocket("payload") - - with pytest.raises(TransportError, match="text frame"): - await transport.receive() From 6ec7228db144d7a73473aed0f2209d392280aca5 Mon Sep 17 00:00:00 2001 From: Jeff Repanich Date: Sat, 8 Aug 2026 08:07:59 -0400 Subject: [PATCH 2/4] fix: use configured broker tenant in test JWTs --- .github/workflows/ci.yml | 1 + tests/integration/fixture/jwt.py | 6 ++++-- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4c7561d..29539af 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -65,6 +65,7 @@ jobs: FITZ_BROKER_AUTH_WS_ADDR: ws://localhost:4090/ws FITZ_BROKER_JWT_HMAC_SECRET: test-secret-key FITZ_BROKER_JWT_AUDIENCE: fitz + FITZ_BROKER_JWT_TENANT: dev CONFORMANCE_TRANSPORT: ${{ matrix.transport }} CONFORMANCE_AUTH_MODE: ${{ matrix.auth-mode }} CONFORMANCE_OUTPUT: artifacts/conformance-${{ matrix.transport }}-${{ matrix.auth-mode }}.json diff --git a/tests/integration/fixture/jwt.py b/tests/integration/fixture/jwt.py index a659f38..8700631 100644 --- a/tests/integration/fixture/jwt.py +++ b/tests/integration/fixture/jwt.py @@ -32,13 +32,14 @@ def _encode(payload: dict[str, object], secret: str) -> str: def make_valid_jwt(subject: str = "fitz-py-tests") -> str: secret = os.getenv("FITZ_BROKER_JWT_HMAC_SECRET", "test-secret-key") audience = os.getenv("FITZ_BROKER_JWT_AUDIENCE", "fitz") + tenant = os.getenv("FITZ_BROKER_JWT_TENANT", "dev") now = int(time.time()) return _encode( { "iss": "", "sub": subject, "aud": audience, - "tid": "fitz-py-tests", + "tid": tenant, "iat": now, "nbf": now, "exp": now + 300, @@ -51,13 +52,14 @@ def make_valid_jwt(subject: str = "fitz-py-tests") -> str: def make_expired_jwt(subject: str = "fitz-py-tests") -> str: secret = os.getenv("FITZ_BROKER_JWT_HMAC_SECRET", "test-secret-key") audience = os.getenv("FITZ_BROKER_JWT_AUDIENCE", "fitz") + tenant = os.getenv("FITZ_BROKER_JWT_TENANT", "dev") now = int(time.time()) return _encode( { "iss": "", "sub": subject, "aud": audience, - "tid": "fitz-py-tests", + "tid": tenant, "iat": now - 600, "nbf": now - 600, "exp": now - 300, From 4adac51b631b711614ae2bb3c99e2b98c64e9ee4 Mon Sep 17 00:00:00 2001 From: Jeff Repanich Date: Sat, 8 Aug 2026 08:09:53 -0400 Subject: [PATCH 3/4] fix: encode broker permissions as top-level JWT claims --- tests/integration/fixture/jwt.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/integration/fixture/jwt.py b/tests/integration/fixture/jwt.py index 8700631..bce24a2 100644 --- a/tests/integration/fixture/jwt.py +++ b/tests/integration/fixture/jwt.py @@ -37,13 +37,13 @@ def make_valid_jwt(subject: str = "fitz-py-tests") -> str: return _encode( { "iss": "", - "sub": subject, + "sub": tenant, "aud": audience, "tid": tenant, "iat": now, "nbf": now, "exp": now + 300, - "fitz": {"permissions": DEFAULT_PERMISSIONS}, + "permissions": DEFAULT_PERMISSIONS, }, secret, ) @@ -57,13 +57,13 @@ def make_expired_jwt(subject: str = "fitz-py-tests") -> str: return _encode( { "iss": "", - "sub": subject, + "sub": tenant, "aud": audience, "tid": tenant, "iat": now - 600, "nbf": now - 600, "exp": now - 300, - "fitz": {"permissions": DEFAULT_PERMISSIONS}, + "permissions": DEFAULT_PERMISSIONS, }, secret, ) From 9c4c760a5af4360c4439865853dfc67eaf880d64 Mon Sep 17 00:00:00 2001 From: Jeff Repanich Date: Sat, 8 Aug 2026 08:10:02 -0400 Subject: [PATCH 4/4] ci: avoid duplicate branch and pull request runs --- .github/workflows/ci.yml | 1 + 1 file changed, 1 insertion(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 29539af..b4c7c06 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -3,6 +3,7 @@ name: CI on: workflow_dispatch: push: + branches: [main] pull_request: jobs: