perf(sse): reduce redundant connections and fanout overhead

- Share one visible-tab connection across a browser and harden retries.\n- Preframe hub events, disconnect slow clients, and reset idle heartbeats.\n- Add cross-tab, concurrency, race, and fanout benchmark coverage.
This commit is contained in:
boojack 2026-07-29 21:47:54 +08:00
parent c4221c6dfe
commit dd18002b12
6 changed files with 889 additions and 143 deletions

View file

@ -1,7 +1,7 @@
package v1
import (
"fmt"
"io"
"log/slog"
"net/http"
"time"
@ -15,6 +15,8 @@ import (
const (
// sseHeartbeatInterval is the interval between heartbeat pings to keep the connection alive.
sseHeartbeatInterval = 30 * time.Second
sseConnectedComment = ": connected\n\n"
sseHeartbeatComment = ": heartbeat\n\n"
)
type sseRouteRegistrar interface {
@ -51,64 +53,90 @@ func handleSSE(c *echo.Context, hub *SSEHub, authenticator *auth.Authenticator)
w.Header().Set("X-Accel-Buffering", "no") // Disable nginx buffering
w.WriteHeader(http.StatusOK)
// Flush headers immediately.
var flusher http.Flusher
if f, ok := w.(http.Flusher); ok {
f.Flush()
flusher = f
}
// Flush headers immediately.
if flusher != nil {
flusher.Flush()
}
// Subscribe to the hub.
client := hub.Subscribe(userID, role)
defer hub.Unsubscribe(client)
// Create a ticker for heartbeat pings.
heartbeat := time.NewTicker(sseHeartbeatInterval)
defer heartbeat.Stop()
ctx := c.Request().Context()
slog.Debug("SSE client connected", "userID", userID)
// Send an initial comment so clients and dev proxies observe the stream
// immediately instead of waiting for the first heartbeat or data event.
if _, err := fmt.Fprint(w, ": connected\n\n"); err != nil {
if _, err := io.WriteString(w, sseConnectedComment); err != nil {
return nil
}
if f, ok := w.(http.Flusher); ok {
f.Flush()
if flusher != nil {
flusher.Flush()
}
// Heartbeats are only needed after a period with no event data.
heartbeat := time.NewTimer(sseHeartbeatInterval)
defer heartbeat.Stop()
for {
// Prefer a hub-initiated disconnect over draining buffered events.
select {
case <-client.done:
return nil
default:
}
select {
case <-ctx.Done():
// Client disconnected.
slog.Debug("SSE client disconnected", "userID", userID)
return nil
case data, ok := <-client.events:
case <-client.done:
return nil
case frame, ok := <-client.events:
if !ok {
// Channel closed, client was unsubscribed.
return nil
}
// Write SSE event.
if _, err := fmt.Fprintf(w, "data: %s\n\n", data); err != nil {
if _, err := w.Write(frame); err != nil {
return nil
}
if f, ok := w.(http.Flusher); ok {
f.Flush()
if flusher != nil {
flusher.Flush()
}
resetSSETimer(heartbeat)
case <-heartbeat.C:
// Send a heartbeat comment to keep the connection alive.
if _, err := fmt.Fprint(w, ": heartbeat\n\n"); err != nil {
if _, err := io.WriteString(w, sseHeartbeatComment); err != nil {
return nil
}
if f, ok := w.(http.Flusher); ok {
f.Flush()
if flusher != nil {
flusher.Flush()
}
resetSSETimer(heartbeat)
}
}
}
func resetSSETimer(timer *time.Timer) {
if !timer.Stop() {
select {
case <-timer.C:
default:
}
}
timer.Reset(sseHeartbeatInterval)
}
func getSSEClientIdentity(result *auth.AuthResult) (int32, store.Role) {
if result == nil {
return 0, store.RoleUser

View file

@ -8,6 +8,11 @@ import (
"github.com/usememos/memos/store"
)
const (
sseDataPrefix = "data: "
sseClientEventBufferSize = 32
)
// SSEEventType represents the type of change event.
type SSEEventType string
@ -44,9 +49,23 @@ func (e *SSEEvent) JSON() []byte {
return data
}
// Frame returns the event encoded as a complete SSE data frame.
func (e *SSEEvent) Frame() []byte {
data := e.JSON()
if len(data) == 0 {
return nil
}
frame := make([]byte, 0, len(sseDataPrefix)+len(data)+2)
frame = append(frame, sseDataPrefix...)
frame = append(frame, data...)
frame = append(frame, '\n', '\n')
return frame
}
// SSEClient represents a single SSE connection.
type SSEClient struct {
events chan []byte
done chan struct{}
userID int32
role store.Role
}
@ -71,12 +90,14 @@ func NewSSEHub() *SSEHub {
func (h *SSEHub) Subscribe(userID int32, role store.Role) *SSEClient {
c := &SSEClient{
// Buffer a few events so a slow client doesn't block broadcasting.
events: make(chan []byte, 32),
events: make(chan []byte, sseClientEventBufferSize),
done: make(chan struct{}),
userID: userID,
role: role,
}
h.mu.Lock()
if h.closed {
close(c.done)
close(c.events)
} else {
h.clients[c] = struct{}{}
@ -85,11 +106,12 @@ func (h *SSEHub) Subscribe(userID int32, role store.Role) *SSEClient {
return c
}
// Unsubscribe removes a client and closes its channel.
// Unsubscribe removes a client and closes its channels.
func (h *SSEHub) Unsubscribe(c *SSEClient) {
h.mu.Lock()
if _, ok := h.clients[c]; ok {
delete(h.clients, c)
close(c.done)
close(c.events)
}
h.mu.Unlock()
@ -105,30 +127,50 @@ func (h *SSEHub) Close() {
h.closed = true
for c := range h.clients {
delete(h.clients, c)
close(c.done)
close(c.events)
}
}
// Broadcast sends an event to all connected clients.
// Slow clients that have a full buffer will have the event dropped
// to avoid blocking the broadcaster.
// Slow clients with a full buffer are disconnected so they can reconnect and
// resynchronize instead of silently missing an event.
func (h *SSEHub) Broadcast(event *SSEEvent) {
data := event.JSON()
if len(data) == 0 {
if event == nil || !event.hasKnownVisibility() {
return
}
frame := event.Frame()
if len(frame) == 0 {
return
}
var slowClients []*SSEClient
h.mu.RLock()
defer h.mu.RUnlock()
for c := range h.clients {
if !c.canReceive(event) {
continue
}
select {
case c.events <- data:
case c.events <- frame:
default:
// Drop event for slow client to avoid blocking.
slowClients = append(slowClients, c)
}
}
h.mu.RUnlock()
for _, c := range slowClients {
h.Unsubscribe(c)
}
}
func (e *SSEEvent) hasKnownVisibility() bool {
switch e.Visibility {
case store.Private, store.Public, store.Protected, "":
return true
default:
slog.Warn("SSE event has unknown visibility; denying broadcast", "visibility", string(e.Visibility))
return false
}
}
func (c *SSEClient) canReceive(event *SSEEvent) bool {
@ -138,7 +180,6 @@ func (c *SSEClient) canReceive(event *SSEEvent) bool {
case store.Public, store.Protected, "":
return true
default:
slog.Warn("SSE canReceive: unknown visibility type, denying event", "visibility", event.Visibility)
return false
}
}

View file

@ -1,6 +1,8 @@
package v1
import (
"fmt"
"sync"
"testing"
"time"
@ -45,6 +47,8 @@ func TestSSEHub_SubscribeUnsubscribe(t *testing.T) {
// Channel should be closed.
_, ok := <-client.events
assert.False(t, ok, "channel should be closed after Unsubscribe")
_, ok = <-client.done
assert.False(t, ok, "done channel should be closed after Unsubscribe")
}
func TestSSEHub_Close(t *testing.T) {
@ -59,6 +63,10 @@ func TestSSEHub_Close(t *testing.T) {
_, ok := <-ch
assert.False(t, ok, "channel should be closed after hub close")
}
for _, ch := range []chan struct{}{c1.done, c2.done} {
_, ok := <-ch
assert.False(t, ok, "done channel should be closed after hub close")
}
late := hub.Subscribe(3, store.RoleUser)
_, ok := <-late.events
@ -116,6 +124,11 @@ func TestSSEEvent_JSON(t *testing.T) {
assert.Contains(t, string(data), `"parent":"memos/123"`)
}
func TestSSEEvent_Frame(t *testing.T) {
e := &SSEEvent{Type: SSEEventMemoUpdated, Name: "memos/789"}
assert.Equal(t, "data: {\"type\":\"memo.updated\",\"name\":\"memos/789\"}\n\n", string(e.Frame()))
}
func TestSSEHub_PrivateEventsAreScoped(t *testing.T) {
hub := NewSSEHub()
owner := hub.Subscribe(1, store.RoleUser)
@ -151,7 +164,7 @@ func TestSSEHub_PrivateEventsAreScoped(t *testing.T) {
}
}
func TestSSEClient_CanReceive_UnknownVisibility(t *testing.T) {
func TestSSEHub_UnknownVisibilityDenied(t *testing.T) {
hub := NewSSEHub()
client := hub.Subscribe(1, store.RoleUser)
defer hub.Unsubscribe(client)
@ -166,20 +179,113 @@ func TestSSEClient_CanReceive_UnknownVisibility(t *testing.T) {
mustNotReceive(t, client.events, 100*time.Millisecond)
}
func TestSSEHub_SlowClientEventsDropped(t *testing.T) {
func TestSSEHub_SlowClientDisconnected(t *testing.T) {
hub := NewSSEHub()
// Subscribe but never read, so the channel fills up.
slow := hub.Subscribe(1, store.RoleUser)
defer hub.Unsubscribe(slow)
event := &SSEEvent{Type: SSEEventMemoCreated, Name: "memos/x"}
// Send more events than the buffer capacity (32).
for range 40 {
// Send more events than the buffer capacity.
for range sseClientEventBufferSize + 8 {
hub.Broadcast(event) // must not block
}
// At most 32 events should have been queued; the rest were silently dropped.
assert.LessOrEqual(t, len(slow.events), 32)
select {
case <-slow.done:
default:
t.Fatal("slow client should be disconnected after its event buffer fills")
}
received := 0
for range slow.events {
received++
}
assert.Equal(t, sseClientEventBufferSize, received)
}
func TestSSEHub_ConcurrentAccess(t *testing.T) {
hub := NewSSEHub()
const (
workers = 16
iterations = 100
)
var wg sync.WaitGroup
for workerID := range workers {
wg.Go(func() {
for iteration := range iterations {
client := hub.Subscribe(int32(workerID+1), store.RoleUser)
hub.Broadcast(&SSEEvent{
Type: SSEEventMemoUpdated,
Name: fmt.Sprintf("memos/%d-%d", workerID, iteration),
Visibility: store.Private,
CreatorID: int32(workerID + 1),
})
select {
case <-client.events:
default:
}
hub.Unsubscribe(client)
}
})
}
wg.Wait()
hub.mu.RLock()
defer hub.mu.RUnlock()
assert.Empty(t, hub.clients)
}
func BenchmarkSSEHubBroadcast(b *testing.B) {
for _, clientCount := range []int{100, 1_000, 10_000} {
b.Run(fmt.Sprintf("public/%d", clientCount), func(b *testing.B) {
hub, clients := newBenchmarkSSEHub(clientCount)
defer hub.Close()
event := &SSEEvent{Type: SSEEventMemoUpdated, Name: "memos/benchmark", Visibility: store.Public}
b.ReportAllocs()
b.ResetTimer()
for b.Loop() {
hub.Broadcast(event)
for _, client := range clients {
<-client.events
}
}
})
b.Run(fmt.Sprintf("private/%d", clientCount), func(b *testing.B) {
hub, clients := newBenchmarkSSEHub(clientCount)
defer hub.Close()
event := &SSEEvent{
Type: SSEEventMemoUpdated,
Name: "memos/benchmark",
Visibility: store.Private,
CreatorID: 1,
}
b.ReportAllocs()
b.ResetTimer()
for b.Loop() {
hub.Broadcast(event)
<-clients[0].events
<-clients[1].events
}
})
}
}
func newBenchmarkSSEHub(clientCount int) (*SSEHub, []*SSEClient) {
hub := NewSSEHub()
clients := make([]*SSEClient, 0, clientCount)
for i := range clientCount {
role := store.RoleUser
if i == 1 {
role = store.RoleAdmin
}
clients = append(clients, hub.Subscribe(int32(i+1), role))
}
return hub, clients
}
func TestResolveSSECreatorID(t *testing.T) {

View file

@ -79,7 +79,7 @@ func TestSSEHandler_Authentication(t *testing.T) {
require.Equal(t, http.StatusUnauthorized, rec.Code)
})
t.Run("valid token streams initial comment", func(t *testing.T) {
t.Run("valid token streams initial comment and event", func(t *testing.T) {
server := httptest.NewServer(e)
defer server.Close()
@ -95,9 +95,25 @@ func TestSSEHandler_Authentication(t *testing.T) {
require.Equal(t, http.StatusOK, resp.StatusCode)
require.Equal(t, "text/event-stream", resp.Header.Get("Content-Type"))
line, err := bufio.NewReader(resp.Body).ReadString('\n')
reader := bufio.NewReader(resp.Body)
line, err := reader.ReadString('\n')
require.NoError(t, err)
require.Equal(t, ": connected\n", line)
line, err = reader.ReadString('\n')
require.NoError(t, err)
require.Equal(t, "\n", line)
ts.Service.SSEHub.Broadcast(&apiv1.SSEEvent{
Type: apiv1.SSEEventMemoUpdated,
Name: "memos/streamed",
})
line, err = reader.ReadString('\n')
require.NoError(t, err)
require.Equal(t, "data: {\"type\":\"memo.updated\",\"name\":\"memos/streamed\"}\n", line)
line, err = reader.ReadString('\n')
require.NoError(t, err)
require.Equal(t, "\n", line)
})
t.Run("hub close disconnects stream", func(t *testing.T) {

View file

@ -11,6 +11,11 @@ import { userKeys } from "@/hooks/useUserQueries";
const INITIAL_RETRY_DELAY_MS = 1000;
const MAX_RETRY_DELAY_MS = 30000;
const RETRY_BACKOFF_MULTIPLIER = 2;
const STABLE_CONNECTION_THRESHOLD_MS = 60_000;
const HIDDEN_DISCONNECT_DELAY_MS = 30_000;
const SSE_CONNECTION_LOCK_NAME = "memos-sse-connection";
const SSE_SYNC_CHANNEL_NAME = "memos-sse-sync";
const SSE_EVENT_TYPES = {
memoCreated: "memo.created",
@ -56,6 +61,17 @@ export function useSSEConnectionStatus(): SSEConnectionStatus {
return useSyncExternalStore(subscribeSSEStatus, getSSEStatus, getSSEStatus);
}
interface SSEChangeEvent {
type: (typeof SSE_EVENT_TYPES)[keyof typeof SSE_EVENT_TYPES];
name: string;
parent?: string;
}
type SSESyncMessage =
| { kind: "event"; event: SSEChangeEvent }
| { kind: "status"; status: SSEConnectionStatus; targetID?: string }
| { kind: "status-request"; requesterID: string };
// ---------------------------------------------------------------------------
// Main hook
// ---------------------------------------------------------------------------
@ -71,7 +87,6 @@ export function useLiveMemoRefresh() {
const queryClient = useQueryClient();
const { currentUser } = useAuth();
const retryDelayRef = useRef(INITIAL_RETRY_DELAY_MS);
const abortControllerRef = useRef<AbortController | null>(null);
const hasConnectedOnceRef = useRef(false);
const currentUserName = currentUser?.name;
@ -79,118 +94,239 @@ export function useLiveMemoRefresh() {
useEffect(() => {
let mounted = true;
let retryTimeout: ReturnType<typeof setTimeout> | null = null;
const connect = async () => {
if (!mounted) return;
if (!currentUserName) {
setSSEStatus("disconnected");
return;
}
let token = await getRequestToken();
if (!token) {
setSSEStatus("disconnected");
// Not logged in; do not retry. Effect will re-run when currentUser is set.
return;
}
setSSEStatus("connecting");
const abortController = new AbortController();
abortControllerRef.current = abortController;
let isLeader = false;
let leadershipAbortController: AbortController | null = null;
let hiddenDisconnectTimeout: ReturnType<typeof setTimeout> | null = null;
const channel = createSSESyncChannel();
const channelClientID = createSSEChannelClientID();
const postSyncMessage = (message: SSESyncMessage) => {
try {
let response = await fetchSSEStream(token, abortController.signal);
if (response.status === 401) {
await refreshAccessToken();
token = await getRequestToken();
if (!token) {
throw new Error("SSE connection failed: missing token after refresh");
}
response = await fetchSSEStream(token, abortController.signal);
}
if (!response.ok || !response.body) {
throw new Error(`SSE connection failed: ${response.status}`);
}
// Successfully connected - reset retry delay.
retryDelayRef.current = INITIAL_RETRY_DELAY_MS;
setSSEStatus("connected");
if (hasConnectedOnceRef.current) {
// Resync active collaborative views after reconnect because the server may have
// dropped events while the client was disconnected or backpressured.
queryClient.invalidateQueries({ queryKey: memoKeys.all, refetchType: "active" });
queryClient.invalidateQueries({ queryKey: userKeys.stats(), refetchType: "active" });
}
hasConnectedOnceRef.current = true;
const reader = response.body.getReader();
const decoder = new TextDecoder();
let buffer = "";
while (mounted) {
const { done, value } = await reader.read();
if (done) break;
buffer += decoder.decode(value, { stream: true });
// Process complete SSE messages (separated by double newlines).
const messages = buffer.split("\n\n");
// Keep the last incomplete chunk in the buffer.
buffer = messages.pop() || "";
for (const message of messages) {
if (!message.trim()) continue;
// Parse SSE format: lines starting with "data: " contain JSON payload.
// Lines starting with ":" are comments (heartbeats).
for (const line of message.split("\n")) {
if (line.startsWith("data: ")) {
const jsonStr = line.slice(6);
try {
const event = JSON.parse(jsonStr) as SSEChangeEvent;
handleEvent(event);
} catch {
// Ignore malformed JSON.
}
}
}
}
}
} catch (err: unknown) {
if (err instanceof DOMException && err.name === "AbortError") {
// Intentional abort, don't reconnect.
setSSEStatus("disconnected");
return;
}
// Connection lost or failed - reconnect with backoff.
}
setSSEStatus("disconnected");
// Reconnect with exponential backoff.
if (mounted) {
const delay = retryDelayRef.current;
retryDelayRef.current = Math.min(delay * RETRY_BACKOFF_MULTIPLIER, MAX_RETRY_DELAY_MS);
retryTimeout = setTimeout(connect, delay);
channel?.postMessage(message);
} catch {
// Cross-tab synchronization is a performance optimization. The leader
// keeps its own live connection if the channel becomes unavailable.
}
};
connect();
const publishStatus = (status: SSEConnectionStatus) => {
setSSEStatus(status);
postSyncMessage({ kind: "status", status });
};
const publishEvent = (event: SSEChangeEvent) => {
handleEvent(event);
postSyncMessage({ kind: "event", event });
};
const runConnectionLoop = async (signal: AbortSignal) => {
while (mounted && !signal.aborted) {
let connectedAt: number | undefined;
try {
publishStatus("connecting");
let token = await getRequestToken();
if (!token) {
// Not logged in; do not retry. The effect will re-run if the
// authenticated user changes.
publishStatus("disconnected");
return;
}
let response = await fetchSSEStream(token, signal);
if (response.status === 401) {
await cancelResponseBody(response);
await refreshAccessToken();
token = await getRequestToken();
if (!token) {
throw new Error("SSE connection failed: missing token after refresh");
}
response = await fetchSSEStream(token, signal);
}
if (signal.aborted) {
await cancelResponseBody(response);
return;
}
if (!response.ok || !response.body) {
await cancelResponseBody(response);
throw new Error(`SSE connection failed: ${response.status}`);
}
connectedAt = Date.now();
publishStatus("connected");
if (hasConnectedOnceRef.current) {
// Resync active collaborative views after reconnect because the server may have
// dropped events while the client was disconnected or backpressured.
queryClient.invalidateQueries({ queryKey: memoKeys.all, refetchType: "active" });
queryClient.invalidateQueries({ queryKey: userKeys.stats(), refetchType: "active" });
}
hasConnectedOnceRef.current = true;
await consumeSSEStream(response.body, signal, publishEvent);
} catch (err: unknown) {
if (!isAbortError(err)) {
// Connection lost or failed; retry below with exponential backoff.
}
}
if (connectedAt !== undefined && Date.now() - connectedAt >= STABLE_CONNECTION_THRESHOLD_MS) {
retryDelayRef.current = INITIAL_RETRY_DELAY_MS;
}
if (!mounted || signal.aborted) {
return;
}
publishStatus("disconnected");
const retryDelay = retryDelayRef.current;
retryDelayRef.current = Math.min(retryDelay * RETRY_BACKOFF_MULTIPLIER, MAX_RETRY_DELAY_MS);
await waitForRetry(withFullJitter(retryDelay), signal);
}
};
const startLeadership = () => {
if (!mounted || leadershipAbortController || !currentUserName || !isPageEligibleForSSE()) {
return;
}
const abortController = new AbortController();
leadershipAbortController = abortController;
void runWithSSELeadership(
abortController.signal,
async () => {
// This tab may have become hidden while its lock request was queued.
if (!isPageEligibleForSSE()) {
return;
}
isLeader = true;
try {
await runConnectionLoop(abortController.signal);
} finally {
isLeader = false;
}
},
// Without BroadcastChannel, followers could not receive the leader's
// events, so each visible tab keeps its own connection as a fallback.
channel !== null,
)
.catch((err: unknown) => {
if (!isAbortError(err)) {
console.warn("SSE connection coordinator failed", err);
}
})
.finally(() => {
if (leadershipAbortController === abortController) {
leadershipAbortController = null;
}
});
postSyncMessage({ kind: "status-request", requesterID: channelClientID });
};
const stopLeadership = () => {
const abortController = leadershipAbortController;
leadershipAbortController = null;
abortController?.abort();
setSSEStatus("disconnected");
};
const clearHiddenDisconnectTimeout = () => {
if (hiddenDisconnectTimeout) {
clearTimeout(hiddenDisconnectTimeout);
hiddenDisconnectTimeout = null;
}
};
const reconcilePageState = () => {
if (document.visibilityState !== "visible") {
if (!hiddenDisconnectTimeout) {
hiddenDisconnectTimeout = setTimeout(() => {
hiddenDisconnectTimeout = null;
stopLeadership();
}, HIDDEN_DISCONNECT_DELAY_MS);
}
return;
}
clearHiddenDisconnectTimeout();
if (navigator.onLine === false) {
stopLeadership();
return;
}
startLeadership();
};
const handleVisibilityChange = () => reconcilePageState();
const handleOnline = () => reconcilePageState();
const handleOffline = () => stopLeadership();
const handlePageHide = () => {
clearHiddenDisconnectTimeout();
stopLeadership();
};
const handlePageShow = () => reconcilePageState();
const handleSyncMessage = (messageEvent: MessageEvent<SSESyncMessage>) => {
const message = messageEvent.data;
if (!message || typeof message !== "object") {
return;
}
switch (message.kind) {
case "event":
hasConnectedOnceRef.current = true;
handleEvent(message.event);
break;
case "status":
if (message.targetID && message.targetID !== channelClientID) {
break;
}
if (message.status === "connected") {
if (hasConnectedOnceRef.current) {
queryClient.invalidateQueries({ queryKey: memoKeys.all, refetchType: "active" });
queryClient.invalidateQueries({ queryKey: userKeys.stats(), refetchType: "active" });
}
hasConnectedOnceRef.current = true;
}
setSSEStatus(message.status);
break;
case "status-request":
if (isLeader) {
postSyncMessage({ kind: "status", status: getSSEStatus(), targetID: message.requesterID });
}
break;
}
};
channel?.addEventListener("message", handleSyncMessage);
document.addEventListener("visibilitychange", handleVisibilityChange);
window.addEventListener("online", handleOnline);
window.addEventListener("offline", handleOffline);
window.addEventListener("pagehide", handlePageHide);
window.addEventListener("pageshow", handlePageShow);
if (!currentUserName) {
setSSEStatus("disconnected");
} else {
reconcilePageState();
}
return () => {
mounted = false;
setSSEStatus("disconnected");
retryDelayRef.current = INITIAL_RETRY_DELAY_MS;
if (retryTimeout) {
clearTimeout(retryTimeout);
}
if (abortControllerRef.current) {
abortControllerRef.current.abort();
}
clearHiddenDisconnectTimeout();
stopLeadership();
channel?.removeEventListener("message", handleSyncMessage);
channel?.close();
document.removeEventListener("visibilitychange", handleVisibilityChange);
window.removeEventListener("online", handleOnline);
window.removeEventListener("offline", handleOffline);
window.removeEventListener("pagehide", handlePageHide);
window.removeEventListener("pageshow", handlePageShow);
};
}, [handleEvent, currentUserName]);
}
@ -199,6 +335,78 @@ export function useLiveMemoRefresh() {
// Event handling
// ---------------------------------------------------------------------------
function createSSESyncChannel(): BroadcastChannel | null {
// A channel without cross-tab locking would duplicate every event because
// each tab would still own a connection. Fall back to independent visible
// tab connections when either coordination primitive is unavailable.
if (!("locks" in navigator) || !navigator.locks) {
return null;
}
try {
return new BroadcastChannel(SSE_SYNC_CHANNEL_NAME);
} catch {
return null;
}
}
function createSSEChannelClientID(): string {
return `${Date.now().toString(36)}-${Math.random().toString(36).slice(2)}`;
}
function isPageEligibleForSSE(): boolean {
return document.visibilityState === "visible" && navigator.onLine !== false;
}
async function runWithSSELeadership(signal: AbortSignal, task: () => Promise<void>, coordinateAcrossTabs: boolean): Promise<void> {
const lockManager = "locks" in navigator ? navigator.locks : undefined;
if (!coordinateAcrossTabs || !lockManager) {
await task();
return;
}
await lockManager.request(SSE_CONNECTION_LOCK_NAME, { mode: "exclusive", signal }, async (lock) => {
if (!lock || signal.aborted) {
return;
}
await task();
});
}
function withFullJitter(delay: number): number {
return Math.floor(Math.random() * delay);
}
function waitForRetry(delay: number, signal: AbortSignal): Promise<void> {
if (signal.aborted) {
return Promise.resolve();
}
return new Promise((resolve) => {
const handleAbort = () => {
clearTimeout(timeout);
resolve();
};
const timeout = setTimeout(() => {
signal.removeEventListener("abort", handleAbort);
resolve();
}, delay);
signal.addEventListener("abort", handleAbort, { once: true });
});
}
function isAbortError(err: unknown): boolean {
return typeof err === "object" && err !== null && "name" in err && err.name === "AbortError";
}
async function cancelResponseBody(response: Response): Promise<void> {
try {
await response.body?.cancel();
} catch {
// The connection may already have been closed by the browser.
}
}
function fetchSSEStream(token: string, signal: AbortSignal): Promise<Response> {
return fetch("/api/v1/sse", {
headers: {
@ -210,10 +418,54 @@ function fetchSSEStream(token: string, signal: AbortSignal): Promise<Response> {
});
}
interface SSEChangeEvent {
type: (typeof SSE_EVENT_TYPES)[keyof typeof SSE_EVENT_TYPES];
name: string;
parent?: string;
async function consumeSSEStream(body: ReadableStream<Uint8Array>, signal: AbortSignal, onEvent: (event: SSEChangeEvent) => void) {
const reader = body.getReader();
const decoder = new TextDecoder();
let buffer = "";
const cancelReader = () => {
void reader.cancel().catch(() => {
// The stream may already have been closed by the server.
});
};
signal.addEventListener("abort", cancelReader, { once: true });
try {
while (!signal.aborted) {
const { done, value } = await reader.read();
if (done) {
break;
}
buffer += decoder.decode(value, { stream: true });
// Process complete SSE messages (separated by double newlines).
const messages = buffer.split("\n\n");
// Keep the last incomplete chunk in the buffer.
buffer = messages.pop() || "";
for (const message of messages) {
if (!message.trim()) continue;
// Parse SSE format: lines starting with "data: " contain JSON payload.
// Lines starting with ":" are comments (heartbeats).
for (const line of message.split("\n")) {
if (line.startsWith("data: ")) {
const jsonStr = line.slice(6);
try {
const event = JSON.parse(jsonStr) as SSEChangeEvent;
onEvent(event);
} catch {
// Ignore malformed JSON.
}
}
}
}
}
} finally {
signal.removeEventListener("abort", cancelReader);
reader.releaseLock();
}
}
function handleSSEEvent(event: SSEChangeEvent, queryClient: ReturnType<typeof useQueryClient>) {

View file

@ -0,0 +1,303 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { act, renderHook, waitFor } from "@testing-library/react";
import type { PropsWithChildren } from "react";
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
const mocks = vi.hoisted(() => ({
getRequestToken: vi.fn(),
refreshAccessToken: vi.fn(),
}));
vi.mock("@/connect", () => ({
getRequestToken: mocks.getRequestToken,
refreshAccessToken: mocks.refreshAccessToken,
}));
vi.mock("@/contexts/AuthContext", () => ({
useAuth: () => ({ currentUser: { name: "users/test" } }),
}));
vi.mock("@/hooks/useMemoQueries", () => ({
memoKeys: {
all: ["memos"],
lists: () => ["memos", "list"],
detail: (name: string) => ["memos", "detail", name],
comments: (name: string) => ["memos", "comments", name],
},
}));
vi.mock("@/hooks/useUserQueries", () => ({
userKeys: {
stats: () => ["users", "stats"],
},
}));
import { useLiveMemoRefresh } from "@/hooks/useLiveMemoRefresh";
type MessageListener = (event: MessageEvent) => void;
class TestBroadcastChannel {
static channels = new Map<string, Set<TestBroadcastChannel>>();
readonly name: string;
private readonly listeners = new Set<MessageListener>();
constructor(name: string) {
this.name = name;
const peers = TestBroadcastChannel.channels.get(name) ?? new Set<TestBroadcastChannel>();
peers.add(this);
TestBroadcastChannel.channels.set(name, peers);
}
postMessage(data: unknown): void {
for (const peer of TestBroadcastChannel.channels.get(this.name) ?? []) {
if (peer === this) continue;
queueMicrotask(() => {
const event = new MessageEvent("message", { data });
for (const listener of peer.listeners) {
listener(event);
}
});
}
}
addEventListener(type: string, listener: MessageListener): void {
if (type === "message") {
this.listeners.add(listener);
}
}
removeEventListener(type: string, listener: MessageListener): void {
if (type === "message") {
this.listeners.delete(listener);
}
}
close(): void {
TestBroadcastChannel.channels.get(this.name)?.delete(this);
this.listeners.clear();
}
}
interface PendingLock {
callback: (lock: Lock) => Promise<unknown> | unknown;
name: string;
onAbort: () => void;
reject: (reason: unknown) => void;
resolve: (value: unknown) => void;
signal?: AbortSignal;
}
class TestLockManager {
private held = false;
private readonly queue: PendingLock[] = [];
request(name: string, options: LockOptions, callback: (lock: Lock | null) => Promise<unknown> | unknown): Promise<unknown> {
return new Promise((resolve, reject) => {
const pending: PendingLock = {
callback: (lock) => callback(lock),
name,
onAbort: () => {
const index = this.queue.indexOf(pending);
if (index >= 0) {
this.queue.splice(index, 1);
reject(new DOMException("The lock request was aborted", "AbortError"));
}
},
reject,
resolve,
signal: options.signal,
};
options.signal?.addEventListener("abort", pending.onAbort, { once: true });
this.queue.push(pending);
this.drain();
});
}
private drain(): void {
if (this.held) return;
const pending = this.queue.shift();
if (!pending) return;
if (pending.signal?.aborted) {
pending.reject(new DOMException("The lock request was aborted", "AbortError"));
this.drain();
return;
}
this.held = true;
pending.signal?.removeEventListener("abort", pending.onAbort);
const lock = { mode: "exclusive", name: pending.name } as Lock;
void Promise.resolve(pending.callback(lock))
.then(pending.resolve, pending.reject)
.finally(() => {
this.held = false;
this.drain();
});
}
}
let visibilityState: DocumentVisibilityState;
function createQueryClient() {
return new QueryClient({
defaultOptions: {
queries: {
retry: false,
},
},
});
}
function createWrapper(queryClient: QueryClient) {
return function QueryWrapper({ children }: PropsWithChildren) {
return <QueryClientProvider client={queryClient}>{children}</QueryClientProvider>;
};
}
function installBrowserState(lockManager?: TestLockManager) {
vi.stubGlobal("BroadcastChannel", TestBroadcastChannel);
Object.defineProperty(document, "visibilityState", {
configurable: true,
get: () => visibilityState,
});
Object.defineProperty(navigator, "onLine", {
configurable: true,
value: true,
});
Object.defineProperty(navigator, "locks", {
configurable: true,
value: lockManager,
});
}
async function flushAsyncWork() {
await act(async () => {
for (let index = 0; index < 5; index++) {
await Promise.resolve();
}
});
}
describe("useLiveMemoRefresh", () => {
beforeEach(() => {
visibilityState = "visible";
TestBroadcastChannel.channels.clear();
mocks.getRequestToken.mockResolvedValue("access-token");
mocks.refreshAccessToken.mockResolvedValue(undefined);
});
afterEach(() => {
vi.useRealTimers();
vi.unstubAllGlobals();
Object.defineProperty(navigator, "locks", {
configurable: true,
value: undefined,
});
});
it("shares one SSE connection across tabs and broadcasts events to followers", async () => {
const lockManager = new TestLockManager();
installBrowserState(lockManager);
const streamControllers: ReadableStreamDefaultController<Uint8Array>[] = [];
const fetchMock = vi.fn().mockImplementation(
() =>
new Response(
new ReadableStream<Uint8Array>({
start(controller) {
streamControllers.push(controller);
},
}),
{ status: 200 },
),
);
vi.stubGlobal("fetch", fetchMock);
const firstClient = createQueryClient();
const secondClient = createQueryClient();
const firstInvalidate = vi.spyOn(firstClient, "invalidateQueries");
const secondInvalidate = vi.spyOn(secondClient, "invalidateQueries");
const first = renderHook(() => useLiveMemoRefresh(), { wrapper: createWrapper(firstClient) });
const second = renderHook(() => useLiveMemoRefresh(), { wrapper: createWrapper(secondClient) });
await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1));
act(() => {
streamControllers[0].enqueue(new TextEncoder().encode('data: {"type":"memo.created","name":"memos/1"}\n\n'));
});
await waitFor(() => {
expect(firstInvalidate).toHaveBeenCalledWith({ queryKey: ["memos", "list"] });
expect(secondInvalidate).toHaveBeenCalledWith({ queryKey: ["memos", "list"] });
});
first.unmount();
second.unmount();
});
it("keeps the connection during the hidden grace period, then reconnects when visible", async () => {
vi.useFakeTimers();
installBrowserState(new TestLockManager());
const fetchMock = vi.fn().mockImplementation(
() =>
new Response(
new ReadableStream<Uint8Array>({
start() {},
}),
{ status: 200 },
),
);
vi.stubGlobal("fetch", fetchMock);
const hook = renderHook(() => useLiveMemoRefresh(), { wrapper: createWrapper(createQueryClient()) });
await flushAsyncWork();
expect(fetchMock).toHaveBeenCalledTimes(1);
const firstSignal = (fetchMock.mock.calls[0][1] as RequestInit).signal as AbortSignal;
visibilityState = "hidden";
act(() => document.dispatchEvent(new Event("visibilitychange")));
act(() => vi.advanceTimersByTime(29_999));
expect(firstSignal.aborted).toBe(false);
act(() => vi.advanceTimersByTime(1));
expect(firstSignal.aborted).toBe(true);
visibilityState = "visible";
act(() => document.dispatchEvent(new Event("visibilitychange")));
await flushAsyncWork();
expect(fetchMock).toHaveBeenCalledTimes(2);
hook.unmount();
});
it("keeps backing off when a successful response closes before becoming stable", async () => {
vi.useFakeTimers();
installBrowserState();
vi.spyOn(Math, "random").mockReturnValue(1);
const fetchMock = vi.fn().mockImplementation(() => new Response(new Uint8Array(), { status: 200 }));
vi.stubGlobal("fetch", fetchMock);
const hook = renderHook(() => useLiveMemoRefresh(), { wrapper: createWrapper(createQueryClient()) });
await flushAsyncWork();
expect(fetchMock).toHaveBeenCalledTimes(1);
await act(async () => vi.advanceTimersByTimeAsync(1000));
await flushAsyncWork();
expect(fetchMock).toHaveBeenCalledTimes(2);
await act(async () => vi.advanceTimersByTimeAsync(1000));
await flushAsyncWork();
expect(fetchMock).toHaveBeenCalledTimes(2);
await act(async () => vi.advanceTimersByTimeAsync(1000));
await flushAsyncWork();
expect(fetchMock).toHaveBeenCalledTimes(3);
hook.unmount();
});
});