From 1643babcb0576553a0621d0ef4d8b80070266177 Mon Sep 17 00:00:00 2001 From: Guilherme Souza Date: Wed, 9 Sep 2026 16:41:02 -0300 Subject: [PATCH] fix(realtime): tie channel join_ref to the actual join push ref RealtimeChannel.subscribe() sent the phx_join message with join_ref hardcoded to None, then set self.join_ref to a brand-new, unrelated ref generated *after* the join succeeded. Every later push on the channel (leave, broadcast, presence track/untrack) therefore carried a join_ref that never matched any ref the server had actually seen, defeating its purpose: letting the server detect stale messages after a rejoin. Newer Phoenix versions rely on this and silently drop pushes that don't line up. Fix: generate the join ref once, use it as both `ref` and `join_ref` on the phx_join message, and store that same value as channel.join_ref. Fixes SDK-1526 --- src/realtime/src/realtime/channel.py | 7 +- src/realtime/tests/test_join_ref.py | 132 +++++++++++++++++++++++++++ 2 files changed, 136 insertions(+), 3 deletions(-) create mode 100644 src/realtime/tests/test_join_ref.py diff --git a/src/realtime/src/realtime/channel.py b/src/realtime/src/realtime/channel.py index 449bee13..e7b2e6dc 100644 --- a/src/realtime/src/realtime/channel.py +++ b/src/realtime/src/realtime/channel.py @@ -199,12 +199,13 @@ async def subscribe(self) -> ReplyMessage: if self.socket.last_token is not None: self.params.token = self.socket.last_token payload = self.params.to_payload() + join_ref = self.socket._make_ref() message = Message( topic=self.topic, event=ChannelEvents.join, payload=payload.model_dump(exclude_none=True), - ref=self.socket._make_ref(), - join_ref=None, + ref=join_ref, + join_ref=join_ref, ) msg = await self.socket.send(message) logger.info(f"Subscribe reply: {msg!r}") @@ -212,7 +213,7 @@ async def subscribe(self) -> ReplyMessage: raise Exception( f"error while subscribing to channel: {msg.payload.response!r}" ) - self.join_ref = self.socket._make_ref() + self.join_ref = join_ref self.joined = True return msg diff --git a/src/realtime/tests/test_join_ref.py b/src/realtime/tests/test_join_ref.py new file mode 100644 index 00000000..156de35e --- /dev/null +++ b/src/realtime/tests/test_join_ref.py @@ -0,0 +1,132 @@ +from typing import AsyncIterator + +import pytest +from pydantic import TypeAdapter +from websockets.asyncio.server import ServerConnection, serve + +from realtime.channel import RealtimeChannelOptions +from realtime.client import connect_once +from realtime.message import ( + ClientMessage, + Message, + ReplyMessage, + ReplyPostgresChanges, + SuccessReplyMessage, +) +from realtime.types import ChannelEvents + +# SDK-1526: the Realtime protocol requires every client-sent event (join, leave, +# broadcast, presence track/untrack) to carry the channel's current `join_ref` so +# the server can detect stale messages after a rejoin. +# See: https://supabase.com/docs/guides/realtime/protocol#client-sent-events + +PORT = 55565 +URL = "localhost" +PUBLISHABLE_KEY = "my-publishable-key" +MessageParser: TypeAdapter[ClientMessage] = TypeAdapter(ClientMessage) + + +async def get_token() -> str: + return PUBLISHABLE_KEY + + +def reply_ack(topic: str, ref: str) -> str: + reply_msg = ReplyMessage( + event=ChannelEvents.reply, + topic=topic, + ref=ref, + payload=SuccessReplyMessage( + status="ok", response=ReplyPostgresChanges(postgres_changes=[]) + ), + ) + return reply_msg.model_dump_json() + + +class RecordingServer: + """A local websocket server that acks every ref'd message and records what it received.""" + + def __init__(self, url: str, port: int): + self.url = url + self.port = port + self.received: list[Message] = [] + + async def start(self) -> None: + self.ws_server = await serve(self.handler, self.url, self.port) + await self.ws_server.start_serving() + + async def __aenter__(self) -> "RecordingServer": + await self.start() + return self + + async def __aexit__(self, *exc) -> None: + self.ws_server.close() + await self.ws_server.wait_closed() + + async def handler(self, connection: ServerConnection) -> None: + while True: + raw = await connection.recv(decode=False) + message = Message.model_validate_json(raw) + self.received.append(message) + if message.ref is not None: + await connection.send(reply_ack(message.topic, message.ref)) + + def events(self, event: str) -> list[Message]: + return [m for m in self.received if m.event == event] + + +@pytest.fixture +async def server() -> AsyncIterator[RecordingServer]: + async with RecordingServer(URL, PORT) as server: + yield server + + +@pytest.mark.asyncio +async def test_join_carries_join_ref_equal_to_its_own_ref(server: RecordingServer): + async with connect_once(f"http://{URL}:{PORT}", get_token) as client: + async with client.channel("test-join-ref") as channel: + [join_message] = server.events("phx_join") + assert join_message.join_ref == join_message.ref + assert join_message.join_ref == channel.join_ref + + +@pytest.mark.asyncio +async def test_broadcast_push_carries_channel_join_ref(server: RecordingServer): + options = RealtimeChannelOptions().broadcast(ack=True) + async with connect_once(f"http://{URL}:{PORT}", get_token) as client: + async with client.channel("test-join-ref-broadcast", params=options) as channel: + await channel.send_broadcast("cursor", {"x": 1}) + + [join_message] = server.events("phx_join") + [broadcast_message] = server.events("broadcast") + assert broadcast_message.join_ref == join_message.ref + assert broadcast_message.join_ref is not None + + +@pytest.mark.asyncio +async def test_presence_track_and_untrack_carry_channel_join_ref( + server: RecordingServer, +): + options = RealtimeChannelOptions().presence(enabled=True) + async with connect_once(f"http://{URL}:{PORT}", get_token) as client: + async with client.channel("test-join-ref-presence", params=options) as channel: + await channel.track({"id": 123}) + await channel.untrack() + + [join_message] = server.events("phx_join") + track_message, untrack_message = server.events("presence") + assert track_message.payload["event"] == "track" + assert track_message.join_ref == join_message.ref + + assert untrack_message.payload["event"] == "untrack" + assert untrack_message.join_ref == join_message.ref + + +@pytest.mark.asyncio +async def test_leave_push_carries_channel_join_ref(server: RecordingServer): + async with connect_once(f"http://{URL}:{PORT}", get_token) as client: + async with client.channel("test-join-ref-leave"): + pass + + [join_message] = server.events("phx_join") + [leave_message] = server.events("phx_leave") + assert leave_message.join_ref == join_message.ref