Skip to content

Commit a135d55

Browse files
authored
feat(realtime): order Postgres deliveries (#88)
1 parent aa55d98 commit a135d55

2 files changed

Lines changed: 178 additions & 23 deletions

File tree

‎src/volcano_sdk/realtime.py‎

Lines changed: 104 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,11 @@
1111
from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeAlias, cast
1212
from urllib.parse import quote, urlsplit, urlunsplit
1313

14+
from ._realtime_fetch_worker import (
15+
PostgresFetchJob,
16+
PostgresFetchOutcome,
17+
PostgresFetchWorker,
18+
)
1419
from .database import Database
1520
from .models import JSONValue, _freeze_json
1621

@@ -34,6 +39,7 @@
3439
importlib.import_module("centrifuge").CentrifugeError,
3540
)
3641
CALLBACK_QUEUE_LIMIT = 128
42+
POSTGRES_QUEUE_LIMIT = 128
3743
NO_PENDING_CALLBACK = object()
3844
CALLBACK_QUEUE_FULL_MESSAGE = (
3945
"Volcano realtime callback queue is full; publication dropped"
@@ -154,6 +160,19 @@ class _PostgresDeliveryIdentity:
154160
subscription_epoch: int
155161

156162

163+
@dataclass(frozen=True, slots=True)
164+
class _PostgresDelivery:
165+
change: PostgresChange
166+
identity: _PostgresDeliveryIdentity
167+
168+
169+
@dataclass(frozen=True, slots=True)
170+
class _CallbackDelivery:
171+
event: str
172+
data: Any
173+
postgres_identity: _PostgresDeliveryIdentity | None = None
174+
175+
157176
@dataclass(slots=True)
158177
class _SessionBoundDatabaseContext:
159178
_transport: Transport
@@ -397,20 +416,20 @@ async def on_publication(self, ctx: PublicationContext) -> None:
397416
async def on_subscribing(self, ctx: Any) -> None:
398417
del ctx
399418
self._channel._subscribed = False
400-
self._channel._end_postgres_epoch()
419+
await self._channel._end_postgres_epoch()
401420
await self._channel._presence_unsubscribed()
402421

403422
async def on_subscribed(self, ctx: Any) -> None:
404423
del ctx
405424
self._channel._subscribed = True
406-
self._channel._begin_postgres_epoch()
425+
await self._channel._begin_postgres_epoch()
407426
if self._channel._type == "presence":
408427
self._channel._schedule_presence_sync()
409428

410429
async def on_unsubscribed(self, ctx: Any) -> None:
411430
del ctx
412431
self._channel._subscribed = False
413-
self._channel._end_postgres_epoch()
432+
await self._channel._end_postgres_epoch()
414433
await self._channel._presence_unsubscribed()
415434

416435
async def on_join(self, ctx: Any) -> None:
@@ -498,7 +517,7 @@ def __init__(
498517
self._presence_lock = asyncio.Lock()
499518
self._presence_sync_task: asyncio.Task[None] | None = None
500519
self._presence_sync_pending = False
501-
self._callback_queue: asyncio.Queue[tuple[str, Any]] = asyncio.Queue(
520+
self._callback_queue: asyncio.Queue[_CallbackDelivery] = asyncio.Queue(
502521
maxsize=CALLBACK_QUEUE_LIMIT
503522
)
504523
self._callback_task: asyncio.Task[None] | None = None
@@ -507,6 +526,8 @@ def __init__(
507526
self._pending_presence_sync: Any = NO_PENDING_CALLBACK
508527
self._postgres_epoch = 0
509528
self._postgres_session_lineage = 0
529+
self._postgres_lock = asyncio.Lock()
530+
self._postgres_worker: PostgresFetchWorker[_PostgresDelivery] | None = None
510531

511532
@property
512533
def name(self) -> str:
@@ -598,14 +619,25 @@ def _capture_postgres_delivery_identity(self) -> _PostgresDeliveryIdentity:
598619
subscription_epoch=self._postgres_epoch,
599620
)
600621

601-
def _begin_postgres_epoch(self) -> None:
602-
if self._type == "postgres":
603-
self._postgres_epoch += 1
604-
self._postgres_session_lineage = self._realtime._connection_lineage()
622+
async def _begin_postgres_epoch(self) -> None:
623+
if self._type != "postgres":
624+
return
625+
await self._stop_postgres_worker()
626+
self._postgres_epoch += 1
627+
self._postgres_session_lineage = self._realtime._connection_lineage()
605628

606-
def _end_postgres_epoch(self) -> None:
607-
if self._type == "postgres":
608-
self._postgres_epoch += 1
629+
async def _end_postgres_epoch(self) -> None:
630+
if self._type != "postgres":
631+
return
632+
self._postgres_epoch += 1
633+
await self._stop_postgres_worker()
634+
635+
async def _stop_postgres_worker(self) -> None:
636+
async with self._postgres_lock:
637+
worker = self._postgres_worker
638+
self._postgres_worker = None
639+
if worker is not None:
640+
await worker.close()
609641

610642
def _postgres_delivery_is_current(
611643
self,
@@ -622,14 +654,41 @@ def _postgres_delivery_is_current(
622654

623655
async def _receive_postgres_change(self, data: Any) -> None:
624656
change = _postgres_change(data)
625-
if change is None:
657+
if change is None or not self._callbacks.get("*"):
626658
return
627659
if change.mode == "lightweight" and change.type == "DELETE":
628660
old_record = change.old_record
629661
if old_record is None and change.id is not None:
630662
old_record = {"id": change.id}
631663
change = replace(change, old_record=old_record, id=None, mode=None)
632-
await self._emit("*", change)
664+
identity = self._capture_postgres_delivery_identity()
665+
if not self._postgres_delivery_is_current(identity):
666+
return
667+
delivery = _PostgresDelivery(change=change, identity=identity)
668+
async with self._postgres_lock:
669+
if not self._postgres_delivery_is_current(identity):
670+
return
671+
worker = self._postgres_worker
672+
if worker is None:
673+
worker = PostgresFetchWorker(
674+
self._realtime._fetch_postgres_row,
675+
self._deliver_postgres,
676+
queue_limit=POSTGRES_QUEUE_LIMIT,
677+
)
678+
self._postgres_worker = worker
679+
await worker.enqueue(PostgresFetchJob(request=None, fallback=delivery))
680+
681+
async def _deliver_postgres(
682+
self,
683+
outcome: PostgresFetchOutcome[_PostgresDelivery],
684+
) -> None:
685+
delivery = outcome.job.fallback
686+
if self._postgres_delivery_is_current(delivery.identity):
687+
await self._emit(
688+
"*",
689+
delivery.change,
690+
postgres_identity=delivery.identity,
691+
)
633692

634693
async def subscribe(self) -> None:
635694
"""Subscribe to this channel."""
@@ -643,7 +702,13 @@ async def unsubscribe(self) -> None:
643702
"""Unsubscribe from this channel."""
644703
await self._realtime._unsubscribe(self)
645704

646-
async def _emit(self, event: str, data: Any) -> None:
705+
async def _emit(
706+
self,
707+
event: str,
708+
data: Any,
709+
*,
710+
postgres_identity: _PostgresDeliveryIdentity | None = None,
711+
) -> None:
647712
if not self._callbacks.get(event):
648713
return
649714
if event == "presence_sync":
@@ -657,7 +722,9 @@ async def _emit(self, event: str, data: Any) -> None:
657722
)
658723
self._callback_stop = None
659724
try:
660-
self._callback_queue.put_nowait((event, data))
725+
self._callback_queue.put_nowait(
726+
_CallbackDelivery(event, data, postgres_identity)
727+
)
661728
except asyncio.QueueFull:
662729
if event == "presence_sync":
663730
self._pending_presence_sync = data
@@ -682,11 +749,15 @@ async def _restart_callback_dispatcher(self, previous: asyncio.Task[None]) -> No
682749

683750
async def _dispatch_callbacks(self, stop: asyncio.Event) -> None:
684751
while not stop.is_set():
685-
event, data = await self._callback_queue.get()
752+
delivery = await self._callback_queue.get()
686753
try:
687-
for callback in tuple(self._callbacks.get(event, [])):
754+
if not self._callback_delivery_is_current(delivery):
755+
continue
756+
for callback in tuple(self._callbacks.get(delivery.event, [])):
757+
if not self._callback_delivery_is_current(delivery):
758+
break
688759
active_task = asyncio.create_task(
689-
self._run_callback(callback, data)
760+
self._run_callback(callback, delivery)
690761
)
691762
self._active_callback_task = active_task
692763
try:
@@ -707,15 +778,25 @@ async def _dispatch_callbacks(self, stop: asyncio.Event) -> None:
707778
self._callback_queue.task_done()
708779
self._enqueue_pending_presence_sync()
709780

781+
def _callback_delivery_is_current(self, delivery: _CallbackDelivery) -> bool:
782+
identity = delivery.postgres_identity
783+
return identity is None or self._postgres_delivery_is_current(identity)
784+
710785
def _enqueue_pending_presence_sync(self) -> None:
711786
pending = self._pending_presence_sync
712787
if pending is NO_PENDING_CALLBACK or self._callback_queue.full():
713788
return
714789
self._pending_presence_sync = NO_PENDING_CALLBACK
715-
self._callback_queue.put_nowait(("presence_sync", pending))
790+
self._callback_queue.put_nowait(_CallbackDelivery("presence_sync", pending))
716791

717-
async def _run_callback(self, callback: MessageCallback, data: Any) -> None:
718-
result = callback(data)
792+
async def _run_callback(
793+
self,
794+
callback: MessageCallback,
795+
delivery: _CallbackDelivery,
796+
) -> None:
797+
if not self._callback_delivery_is_current(delivery):
798+
return
799+
result = callback(delivery.data)
719800
if inspect.isawaitable(result):
720801
await result
721802

@@ -850,6 +931,8 @@ async def _cancel_presence_sync(self) -> None:
850931

851932
async def _reset(self) -> None:
852933
self._subscription = None
934+
self._subscribed = False
935+
await self._end_postgres_epoch()
853936
await self._cancel_presence_sync()
854937
task = self._callback_task
855938
active_task = self._active_callback_task

‎tests/unit/test_realtime.py‎

Lines changed: 74 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -503,6 +503,73 @@ async def scenario() -> None:
503503
asyncio.run(scenario())
504504

505505

506+
def test_realtime_drops_queued_postgres_callbacks_from_an_old_epoch() -> None:
507+
official = FakeCentrifugeClient()
508+
client = VolcanoClient(
509+
anon_key="anon-key",
510+
_transport=AuthTransport(),
511+
_realtime_client_factory=FakeCentrifugeFactory(official),
512+
)
513+
client.auth.sign_in(email="user@example.com", password="secret")
514+
515+
async def scenario() -> None:
516+
def insert(record_id: int) -> dict[str, Any]:
517+
return {
518+
"type": "INSERT",
519+
"schema": "public",
520+
"table": "messages",
521+
"record": {"id": record_id},
522+
"timestamp": "2026-09-03T12:00:00Z",
523+
}
524+
525+
received: list[int] = []
526+
first_started = asyncio.Event()
527+
release_first = asyncio.Event()
528+
third_received = asyncio.Event()
529+
channel = client.realtime.channel(
530+
"public:messages",
531+
channel_type="postgres",
532+
)
533+
534+
async def on_insert(change: Any) -> None:
535+
record_id = change.record["id"]
536+
if record_id == 1:
537+
first_started.set()
538+
await release_first.wait()
539+
received.append(record_id)
540+
if record_id == 3:
541+
third_received.set()
542+
543+
channel.on_postgres_changes(
544+
"INSERT",
545+
schema="public",
546+
table="messages",
547+
callback=on_insert,
548+
)
549+
await channel.subscribe()
550+
subscription = official.subscription
551+
assert subscription is not None
552+
553+
await subscription.emit(insert(1))
554+
await first_started.wait()
555+
first_worker = channel._postgres_worker
556+
assert first_worker is not None
557+
await subscription.emit(insert(2))
558+
559+
await subscription.emit_subscribing()
560+
assert channel._postgres_worker is None
561+
release_first.set()
562+
await subscription.emit_subscribed()
563+
await subscription.emit(insert(3))
564+
await asyncio.wait_for(third_received.wait(), timeout=0.2)
565+
566+
assert received == [1, 3]
567+
assert channel._postgres_worker is not first_worker
568+
await client.realtime.disconnect()
569+
570+
asyncio.run(scenario())
571+
572+
506573
def test_realtime_routes_immutable_rls_scoped_postgres_changes() -> None:
507574
official = FakeCentrifugeClient()
508575
client = VolcanoClient(
@@ -1082,15 +1149,20 @@ async def wait_forever() -> None:
10821149
blocker = asyncio.create_task(wait_forever())
10831150
channel._callback_task = blocker
10841151
for _ in range(channel._callback_queue.maxsize):
1085-
channel._callback_queue.put_nowait(("presence_sync", {"version": 0}))
1152+
channel._callback_queue.put_nowait(
1153+
realtime_module._CallbackDelivery(
1154+
"presence_sync",
1155+
{"version": 0},
1156+
)
1157+
)
10861158

10871159
await channel._emit("presence_sync", {"version": 1})
10881160
channel._callback_queue.get_nowait()
10891161
channel._callback_queue.task_done()
10901162
await channel._emit("presence_sync", {"version": 2})
10911163
queued: list[Any] = []
10921164
while not channel._callback_queue.empty():
1093-
queued.append(channel._callback_queue.get_nowait()[1])
1165+
queued.append(channel._callback_queue.get_nowait().data)
10941166
channel._callback_queue.task_done()
10951167
channel._enqueue_pending_presence_sync()
10961168
blocker.cancel()

0 commit comments

Comments
 (0)