diff --git a/server/router/api/v1/sse_handler.go b/server/router/api/v1/sse_handler.go index dedfbec6..ded2fcd4 100644 --- a/server/router/api/v1/sse_handler.go +++ b/server/router/api/v1/sse_handler.go @@ -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 diff --git a/server/router/api/v1/sse_hub.go b/server/router/api/v1/sse_hub.go index 88887c4a..e66f7f3f 100644 --- a/server/router/api/v1/sse_hub.go +++ b/server/router/api/v1/sse_hub.go @@ -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 } } diff --git a/server/router/api/v1/sse_hub_test.go b/server/router/api/v1/sse_hub_test.go index 2e2005ff..e3c0da7f 100644 --- a/server/router/api/v1/sse_hub_test.go +++ b/server/router/api/v1/sse_hub_test.go @@ -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) { diff --git a/server/router/api/v1/test/sse_handler_test.go b/server/router/api/v1/test/sse_handler_test.go index 711abfc1..ff01e1aa 100644 --- a/server/router/api/v1/test/sse_handler_test.go +++ b/server/router/api/v1/test/sse_handler_test.go @@ -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) { diff --git a/web/src/hooks/useLiveMemoRefresh.ts b/web/src/hooks/useLiveMemoRefresh.ts index c9df6a60..c3e09696 100644 --- a/web/src/hooks/useLiveMemoRefresh.ts +++ b/web/src/hooks/useLiveMemoRefresh.ts @@ -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(null); const hasConnectedOnceRef = useRef(false); const currentUserName = currentUser?.name; @@ -79,118 +94,239 @@ export function useLiveMemoRefresh() { useEffect(() => { let mounted = true; - let retryTimeout: ReturnType | 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 | 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) => { + 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, coordinateAcrossTabs: boolean): Promise { + 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 { + 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 { + try { + await response.body?.cancel(); + } catch { + // The connection may already have been closed by the browser. + } +} + function fetchSSEStream(token: string, signal: AbortSignal): Promise { return fetch("/api/v1/sse", { headers: { @@ -210,10 +418,54 @@ function fetchSSEStream(token: string, signal: AbortSignal): Promise { }); } -interface SSEChangeEvent { - type: (typeof SSE_EVENT_TYPES)[keyof typeof SSE_EVENT_TYPES]; - name: string; - parent?: string; +async function consumeSSEStream(body: ReadableStream, 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) { diff --git a/web/tests/live-memo-refresh.test.tsx b/web/tests/live-memo-refresh.test.tsx new file mode 100644 index 00000000..194b2251 --- /dev/null +++ b/web/tests/live-memo-refresh.test.tsx @@ -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>(); + + readonly name: string; + private readonly listeners = new Set(); + + constructor(name: string) { + this.name = name; + const peers = TestBroadcastChannel.channels.get(name) ?? new Set(); + 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; + 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): Promise { + 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 {children}; + }; +} + +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[] = []; + const fetchMock = vi.fn().mockImplementation( + () => + new Response( + new ReadableStream({ + 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({ + 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(); + }); +});