Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 20 additions & 2 deletions holmes/core/conversations_worker/realtime_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,9 +23,11 @@
import os
import ssl
import threading
import time
import urllib.parse
from typing import Any, Callable, Dict, Optional, TYPE_CHECKING

import jwt
import realtime._async.client as rt_client
from realtime._async.channel import ChannelStates
from realtime._async.client import AsyncRealtimeClient
Expand Down Expand Up @@ -53,6 +55,15 @@
# reconnect loop can never be stalled indefinitely by a hung auth call.
_RECONNECT_SIGN_IN_TIMEOUT_SECONDS = 90

# Must exceed the auth refresh interval so a tick lands inside it.
_AUTH_REFRESH_LEEWAY_SECONDS = 300
Comment thread
Avi-Robusta marked this conversation as resolved.


def _expires_within(token: str, seconds: float) -> bool:
# Signature is irrelevant; exp is our own claim.
exp = jwt.decode(token, options={"verify_signature": False})["exp"]
return exp - time.time() <= seconds


# ---- channel topic helpers ----

Expand Down Expand Up @@ -328,7 +339,9 @@ async def _run(self) -> None:
# own full teardown/reconnect on any failure signal.
unhealthy_reason = self._channel_unhealthy()
if unhealthy_reason is not None:
logging.warning(
# A first reconnect is routine and self-healing.
logging.log(
logging.INFO if reconnect_attempts == 0 else logging.WARNING,
"Realtime channel unhealthy (%s), reconnecting",
unhealthy_reason,
)
Expand Down Expand Up @@ -488,14 +501,19 @@ async def _full_reconnect(self) -> None:
await self._connect_and_subscribe()

async def _maybe_refresh_auth(self) -> None:
"""Re-push the Supabase JWT to the realtime client if it rotated."""
"""Re-push the Supabase JWT, re-signing in first if it is near expiry."""
if not self._client:
return
try:
session = self.dal.client.auth.get_session() # type: ignore[attr-defined]
if session is None:
return
new_jwt = session.access_token
if new_jwt and _expires_within(new_jwt, _AUTH_REFRESH_LEEWAY_SECONDS):
# Nothing else refreshes the JWT on a realtime-only path.
await asyncio.to_thread(self.dal.sign_in)
session = self.dal.client.auth.get_session() # type: ignore[attr-defined]
new_jwt = session.access_token if session is not None else None
if not new_jwt or new_jwt == self._last_auth_jwt:
return
await self._client.set_auth(new_jwt)
Expand Down
118 changes: 117 additions & 1 deletion tests/core/conversations_worker/test_realtime_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,16 +3,19 @@
import logging
import os
import ssl as _ssl
from unittest.mock import MagicMock
import time
from unittest.mock import AsyncMock, MagicMock

import certifi
import jwt
import pytest
import realtime._async.client as rt_client
from realtime._async.channel import ChannelStates

from holmes.core.conversations_worker.realtime_manager import (
RealtimeWorker,
_build_ssl_context,
_expires_within,
_install_realtime_log_filter_if_needed,
_install_ssl_patch_if_needed,
_RealtimeConnectivityWarningFilter,
Expand Down Expand Up @@ -470,3 +473,116 @@ async def boom_connect():
m._connect_and_subscribe = boom_connect # type: ignore[method-assign]
with pytest.raises(type(exc)):
asyncio.run(m._full_reconnect())


# ---- proactive near-expiry auth refresh ----


def _token(expires_in: float) -> str:
return jwt.encode({"exp": int(time.time() + expires_in)}, "k" * 32, algorithm="HS256")


def _manager_with_session(token):
m = _make_manager()
m._client = MagicMock()
m._client.set_auth = AsyncMock()
session = MagicMock()
session.access_token = token
m.dal.client.auth.get_session = MagicMock(return_value=session)
return m


def test_expires_within():
assert _expires_within(_token(60), 300) is True
assert _expires_within(_token(-10), 300) is True
assert _expires_within(_token(3600), 300) is False


def test_refresh_auth_re_signs_in_when_token_near_expiry():
"""The bug: nothing refreshed the JWT on a realtime-only path, so it lapsed
and Supabase closed the socket with InvalidJWTToken."""
fresh = _token(3600)
m = _manager_with_session(_token(30))

def sign_in():
rotated = MagicMock()
rotated.access_token = fresh
m.dal.client.auth.get_session = MagicMock(return_value=rotated)

m.dal.sign_in = MagicMock(side_effect=sign_in)

asyncio.run(m._maybe_refresh_auth())

m.dal.sign_in.assert_called_once()
m._client.set_auth.assert_awaited_once_with(fresh)
assert m._last_auth_jwt == fresh


def test_refresh_auth_leaves_fresh_token_alone():
fresh = _token(3600)
m = _manager_with_session(fresh)
m.dal.sign_in = MagicMock()
m._last_auth_jwt = fresh

asyncio.run(m._maybe_refresh_auth())

m.dal.sign_in.assert_not_called()
m._client.set_auth.assert_not_awaited()


def test_refresh_auth_still_pushes_externally_rotated_token():
fresh = _token(3600)
m = _manager_with_session(fresh)
m.dal.sign_in = MagicMock()
m._last_auth_jwt = "older-token"

asyncio.run(m._maybe_refresh_auth())

m.dal.sign_in.assert_not_called()
m._client.set_auth.assert_awaited_once_with(fresh)


def test_refresh_auth_survives_sign_in_failure():
"""A failure must not escape into _run and kill the thread; the reconnect
path stays the fallback."""
m = _manager_with_session(_token(30))
m.dal.sign_in = MagicMock(side_effect=ConnectionError("network unreachable"))

asyncio.run(m._maybe_refresh_auth())

m._client.set_auth.assert_not_awaited()


# ---- reconnect log level ----


def test_first_reconnect_logs_info_and_repeat_logs_warning(caplog):
async def _scenario():
m = _make_manager() # channel is None -> unhealthy on first check
m._async_stop = asyncio.Event()
attempts = []

async def fake_reconnect():
attempts.append(1)
if len(attempts) == 1:
raise ConnectionError("first connect fails")
m._async_stop.set()
m._stop_event.set()

m._full_reconnect = fake_reconnect # type: ignore[method-assign]

import holmes.core.conversations_worker.realtime_manager as _rm
original = _rm.CONVERSATION_WORKER_REALTIME_RECONNECT_MAX_SECONDS
_rm.CONVERSATION_WORKER_REALTIME_RECONNECT_MAX_SECONDS = 0
try:
await asyncio.wait_for(m._run(), timeout=5.0)
finally:
_rm.CONVERSATION_WORKER_REALTIME_RECONNECT_MAX_SECONDS = original

with caplog.at_level(logging.INFO):
asyncio.run(_scenario())

unhealthy = [r for r in caplog.records if "Realtime channel unhealthy" in r.getMessage()]
assert len(unhealthy) >= 2
assert unhealthy[0].levelno == logging.INFO
assert unhealthy[1].levelno == logging.WARNING
Loading