diff --git a/docs/realtime.md b/docs/realtime.md index 739d0849..e1bc3525 100644 --- a/docs/realtime.md +++ b/docs/realtime.md @@ -101,10 +101,11 @@ await client.realtime.remove_channel("public:messages", channel_type="postgres") ``` Insert a row from another client during the listening period. +Automatic row lookup requires a primary key named `id`. It reads the current row after the notification; rapid updates may have already changed that row. Matching insert and update notifications in the `public` schema can fetch full rows using the subscription's user token. Compatible row lookups are batched while publication order is preserved. Defaults are a 20 millisecond window and 50 rows; set `fetch_batch_window_ms` and `fetch_max_batch_size` on the channel to change them. -Use `auto_fetch=False` or `set_database_name(None)` to retain lightweight notifications without row lookups. +Set `auto_fetch=False` when first creating the channel or `set_database_name(None)` to retain lightweight notifications without row lookups. Missing rows, failed lookups, and non-public schemas retain the lightweight notification. Deletes use `old_record` or the row ID and do not query the database. diff --git a/features/postgres_changes.py b/features/postgres_changes.py new file mode 100644 index 00000000..441f95a7 --- /dev/null +++ b/features/postgres_changes.py @@ -0,0 +1,128 @@ +from __future__ import annotations + +import asyncio +from datetime import datetime +from typing import TYPE_CHECKING +from uuid import uuid4 + +if TYPE_CHECKING: + from collections.abc import Mapping + + from contract_support import ContractWorld + + from volcano_sdk.models import JSONValue + from volcano_sdk.realtime import Channel, PostgresChange + + +class ChangeObserver: + def __init__(self, channel: Channel, table: str, row_id: JSONValue) -> None: + self.row_id = row_id + self.events: list[PostgresChange] = [] + self.queue: asyncio.Queue[PostgresChange] = asyncio.Queue() + self.inserts: list[PostgresChange] = [] + self.wrong_table: list[PostgresChange] = [] + self.stops = [ + channel.on_postgres_changes( + "*", schema="public", table=table, callback=self.record + ), + channel.on_postgres_changes( + "INSERT", schema="public", table=table, callback=self.record_insert + ), + channel.on_postgres_changes( + "*", + schema="public", + table=table + "_other", + callback=self.record_wrong_table, + ), + ] + + def owns(self, change: PostgresChange) -> bool: + identity = change.record.get("id") if change.record is not None else change.id + return identity == self.row_id + + def record_insert(self, change: PostgresChange) -> None: + if self.owns(change): + self.inserts.append(change) + + def record_wrong_table(self, change: PostgresChange) -> None: + if self.owns(change): + self.wrong_table.append(change) + + def record(self, change: PostgresChange) -> None: + if not self.owns(change): + return + self.events.append(change) + self.queue.put_nowait(change) + + async def next(self) -> PostgresChange: + return await asyncio.wait_for(self.queue.get(), timeout=10) + + def close(self) -> None: + for stop in self.stops: + stop() + + +def verify_change( + event: PostgresChange, + kind: str, + table: str, + row: Mapping[str, JSONValue], + *, + automatic: bool, +) -> None: + assert (event.type, event.schema, event.table) == (kind, "public", table) + datetime.fromisoformat(event.timestamp) + if automatic: + assert event.record == row + assert event.id is None + assert event.mode is None + else: + assert event.id == row["id"] + assert event.mode == "lightweight" + assert event.record is None + + +async def verify_postgres_changes(world: ContractWorld) -> list[str]: + table_name = world.fixture["realtime_table_name"] + row: dict[str, JSONValue] = { + "id": str(uuid4()), + "value": "inserted", + "owner_id": world.fixture["user_id"], + } + table = world.client.database(world.fixture["database_name"]).from_(table_name) + world.cleanup_callbacks.append(lambda: table.delete().eq("id", row["id"]).execute()) + channels = [] + for index, client in enumerate(world.realtime_clients): + client.realtime.set_database_name(world.fixture["database_name"]) + channels.append( + client.realtime.channel( + "public:" + table_name, + channel_type="postgres", + auto_fetch=index == 0, + ) + ) + observers = [ChangeObserver(channel, table_name, row["id"]) for channel in channels] + try: + await asyncio.gather(*(channel.subscribe() for channel in channels)) + for index, kind in enumerate(["INSERT", "UPDATE"]): + expected = {**row, "value": "inserted" if index == 0 else "updated"} + operation = ( + table.insert(expected) + if index == 0 + else table.update({"value": expected["value"]}).eq("id", row["id"]) + ) + assert await asyncio.to_thread(operation.execute) == [expected] + events = await asyncio.gather(*(observer.next() for observer in observers)) + for client_index, event in enumerate(events): + verify_change( + event, kind, table_name, expected, automatic=client_index == 0 + ) + for observer in observers: + assert [event.type for event in observer.events] == ["INSERT", "UPDATE"] + assert len(observer.inserts) == 1 + assert not observer.wrong_table + finally: + for observer in observers: + observer.close() + await asyncio.gather(*(channel.unsubscribe() for channel in channels)) + return ["INSERT", "UPDATE"] diff --git a/features/staged/realtime-postgres.feature b/features/staged/realtime-postgres.feature new file mode 100644 index 00000000..53553c03 --- /dev/null +++ b/features/staged/realtime-postgres.feature @@ -0,0 +1,8 @@ +Feature: SDK Postgres change delivery contract + + @realtime @SDK-REALTIME-004 + Scenario: Insert and update notifications support automatic rows and lightweight delivery + Given two authenticated realtime clients + When the clients observe an inserted and updated contract row + Then the SDK operation succeeds + And automatic and lightweight notifications retain metadata and row identity diff --git a/features/steps/sdk_contract_steps.py b/features/steps/sdk_contract_steps.py index 86089a5a..4abe65f8 100644 --- a/features/steps/sdk_contract_steps.py +++ b/features/steps/sdk_contract_steps.py @@ -17,6 +17,7 @@ classify_error, ) from logs_contract import LogContract +from postgres_changes import verify_postgres_changes from presence_membership import verify_presence_membership from volcano_sdk import NotFoundError, Session, VolcanoClient @@ -1178,3 +1179,18 @@ def verify_presence_rosters(context: Any) -> None: world = _world(context) assert world.last_outcome is not None assert world.last_outcome.value == [1, 2, 1] + + +@when("the clients observe an inserted and updated contract row") +def observe_postgres_changes(context: Any) -> None: + world = _world(context) + if world.last_outcome is not None and not world.last_outcome.ok: + return + world.record(lambda: world.run(verify_postgres_changes(world))) + + +@then("automatic and lightweight notifications retain metadata and row identity") +def verify_postgres_rows(context: Any) -> None: + world = _world(context) + assert world.last_outcome is not None + assert world.last_outcome.value == ["INSERT", "UPDATE"] diff --git a/tests/fixtures/sdk-contract-dry-run.json b/tests/fixtures/sdk-contract-dry-run.json index 7ddadfca..c3300d7f 100644 --- a/tests/fixtures/sdk-contract-dry-run.json +++ b/tests/fixtures/sdk-contract-dry-run.json @@ -38,5 +38,6 @@ "lock_key": "dry-run-lock", "function_name": "dry-run-function", "logs_access_token": "dry-run-project-token", - "function_id": "dry-run-function-id" + "function_id": "dry-run-function-id", + "realtime_table_name": "sdk_contract_dryrun_rt" } diff --git a/tests/unit/test_contract_bindings.py b/tests/unit/test_contract_bindings.py index 284885be..36a160e8 100644 --- a/tests/unit/test_contract_bindings.py +++ b/tests/unit/test_contract_bindings.py @@ -234,6 +234,8 @@ def test_every_contract_phrase_is_bound_verbatim() -> None: "the stored object path equals the contract path", "the subscriber receives the contract message within 10 seconds", "two authenticated realtime clients", + "the clients observe an inserted and updated contract row", + "automatic and lightweight notifications retain metadata and row identity", "a read-only project logs client", "the contract function emits three unique structured log events", "the contract function emits one unique structured log event", @@ -556,3 +558,59 @@ def test_presence_requires_original_handler_membership_sequence( module._observed_membership(snapshots, {"first"}, {"first", "second"}) is expected ) + + +def test_staged_postgres_feature_matches_proposed_shared_source() -> None: + staged = ROOT / "features" / "staged" / "realtime-postgres.feature" + expected = "794c2ecbb94fd262a37840f4c3fe3bd9f9ee58c22fda9df2a46de60f93e52c91" + assert hashlib.sha256(staged.read_bytes()).hexdigest() == expected + + +@pytest.mark.parametrize( + ("automatic", "wrong_field"), + [(True, "record"), (False, "id"), (True, "table"), (True, "id"), (True, "mode")], +) +def test_postgres_notification_checks_reject_wrong_identity( + *, automatic: bool, wrong_field: str +) -> None: + module = _load_module("postgres_changes", ROOT / "features" / "postgres_changes.py") + row = {"id": "row", "value": "inserted", "owner_id": "user"} + event = SimpleNamespace( + type="INSERT", + schema="public", + table="records", + timestamp="2026-09-18T12:00:00Z", + record=row if automatic else None, + id=None if automatic else "row", + mode=None if automatic else "lightweight", + ) + module.verify_change(event, "INSERT", "records", row, automatic=automatic) + setattr(event, wrong_field, "wrong-value") + with pytest.raises(AssertionError): + module.verify_change(event, "INSERT", "records", row, automatic=automatic) + + +@pytest.mark.parametrize("automatic", [True, False]) +def test_postgres_observer_ignores_other_rows(*, automatic: bool) -> None: + module = _load_module("postgres_changes", ROOT / "features" / "postgres_changes.py") + channel = Mock() + observer = module.ChangeObserver(channel, "records", "row") + callbacks = [ + entry.kwargs["callback"] for entry in channel.on_postgres_changes.call_args_list + ] + other = SimpleNamespace( + record={"id": "other"} if automatic else None, id=None if automatic else "other" + ) + for callback in callbacks: + callback(other) + assert not observer.events + assert not observer.inserts + assert not observer.wrong_table + own = SimpleNamespace( + record={"id": "row"} if automatic else None, id=None if automatic else "row" + ) + callbacks[0](own) + callbacks[1](own) + assert asyncio.run(observer.next()) is own + assert observer.inserts == [own] + observer.close()