From d928ddfb779e8204270861c64b53bd22f241d1a6 Mon Sep 17 00:00:00 2001 From: Henry Su Date: Sat, 19 Sep 2026 13:42:52 -0500 Subject: [PATCH] fix(realtime): include connection params in WebSocket URL --- src/realtime/src/realtime/_async/client.py | 10 +++++---- src/realtime/tests/test_connection.py | 25 ++++++++++++++++++++++ 2 files changed, 31 insertions(+), 4 deletions(-) diff --git a/src/realtime/src/realtime/_async/client.py b/src/realtime/src/realtime/_async/client.py index a9b98850..da289d62 100644 --- a/src/realtime/src/realtime/_async/client.py +++ b/src/realtime/src/realtime/_async/client.py @@ -5,7 +5,7 @@ import sys from functools import wraps from typing import Any, Callable, Dict, List, Optional, Union -from urllib.parse import urlencode, urlparse, urlunparse +from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse from warnings import warn import websockets @@ -163,11 +163,12 @@ async def connect(self) -> None: retries = 0 backoff = self.initial_backoff - logger.debug(f"Attempting to connect to WebSocket at {self.url}") + endpoint_url = self.endpoint_url() + logger.debug(f"Attempting to connect to WebSocket at {endpoint_url}") while retries < self.max_retries: try: - ws = await connect(self.url) + ws = await connect(endpoint_url) self._ws_connection = ws logger.debug("WebSocket connection established successfully") return await self._on_connect() @@ -393,7 +394,8 @@ async def _leave_open_topic(self, topic: str): def endpoint_url(self) -> str: parsed_url = urlparse(self.url) - query = urlencode({**self.params, "vsn": VSN}, doseq=True) + current_params = dict(parse_qsl(parsed_url.query, keep_blank_values=True)) + query = urlencode({**self.params, **current_params, "vsn": VSN}, doseq=True) return urlunparse( ( parsed_url.scheme, diff --git a/src/realtime/tests/test_connection.py b/src/realtime/tests/test_connection.py index 4d1c6209..a304ca98 100644 --- a/src/realtime/tests/test_connection.py +++ b/src/realtime/tests/test_connection.py @@ -1,6 +1,8 @@ import asyncio import datetime import os +from unittest.mock import AsyncMock, patch +from urllib.parse import parse_qs, urlparse import aiohttp import pytest @@ -80,6 +82,29 @@ def test_init_client(): assert client.timeout == DEFAULT_TIMEOUT +@pytest.mark.asyncio +async def test_connect_includes_connection_params() -> None: + client = AsyncRealtimeClient( + "https://project.supabase.co/realtime/v1", + "publishable-key", + params={"log_level": "info"}, + ) + + with ( + patch("realtime._async.client.connect", new_callable=AsyncMock) as connect, + patch.object(client, "_on_connect", new_callable=AsyncMock), + ): + await client.connect() + + assert connect.await_args is not None + websocket_url = urlparse(connect.await_args.args[0]) + assert parse_qs(websocket_url.query) == { + "apikey": ["publishable-key"], + "log_level": ["info"], + "vsn": ["1.0.0"], + } + + @pytest.mark.asyncio async def test_broadcast_events(socket: AsyncRealtimeClient): await socket.connect()