diff --git a/app/src/hooks/useThreadQueries.test.ts b/app/src/hooks/useThreadQueries.test.ts new file mode 100644 index 000000000..2e088dda7 --- /dev/null +++ b/app/src/hooks/useThreadQueries.test.ts @@ -0,0 +1,191 @@ +import { act, renderHook, waitFor } from '@testing-library/react'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; + +import type { Thread, ThreadMessage } from '../types/thread'; + +const mockGetThreads = vi.fn(); +const mockGetThreadMessages = vi.fn(); + +vi.mock('../services/api/threadApi', () => ({ + threadApi: { + getThreads: () => mockGetThreads(), + getThreadMessages: (threadId: string) => mockGetThreadMessages(threadId), + }, +})); + +function deferred() { + let resolve!: (value: T) => void; + let reject!: (error: unknown) => void; + const promise = new Promise((resolvePromise, rejectPromise) => { + resolve = resolvePromise; + reject = rejectPromise; + }); + return { promise, resolve, reject }; +} + +const thread: Thread = { + id: 'thread-1', + title: 'Planning', + chatId: null, + isActive: true, + messageCount: 1, + lastMessageAt: '2026-05-06T00:00:00.000Z', + createdAt: '2026-05-06T00:00:00.000Z', + labels: [], +}; + +const message: ThreadMessage = { + id: 'message-1', + content: 'hello', + type: 'text', + extraMetadata: {}, + sender: 'user', + createdAt: '2026-05-06T00:00:00.000Z', +}; + +describe('useThreadQueries', () => { + beforeEach(() => { + mockGetThreads.mockReset(); + mockGetThreadMessages.mockReset(); + }); + + it('loads threads with loading and success state', async () => { + mockGetThreads.mockResolvedValue({ threads: [thread], count: 1 }); + const { useThreads } = await import('./useThreadQueries'); + + const { result } = renderHook(() => useThreads()); + + expect(result.current.loading).toBe(true); + await waitFor(() => expect(result.current.loading).toBe(false)); + expect(result.current.data?.threads).toEqual([thread]); + expect(result.current.data?.count).toBe(1); + expect(result.current.error).toBeNull(); + expect(mockGetThreads).toHaveBeenCalledTimes(1); + }); + + it('surfaces RPC errors without throwing from render', async () => { + mockGetThreads.mockRejectedValue(new Error('rpc failed')); + const { useThreads } = await import('./useThreadQueries'); + + const { result } = renderHook(() => useThreads()); + + await waitFor(() => expect(result.current.loading).toBe(false)); + expect(result.current.data).toBeNull(); + expect(result.current.error?.message).toBe('rpc failed'); + }); + + it('does not load messages when no thread id is available', async () => { + const { useThreadMessages } = await import('./useThreadQueries'); + + const { result, rerender } = renderHook(({ threadId }) => useThreadMessages(threadId), { + initialProps: { threadId: null as string | null }, + }); + + expect(result.current.loading).toBe(false); + expect(result.current.data).toBeNull(); + expect(result.current.error).toBeNull(); + expect(mockGetThreadMessages).not.toHaveBeenCalled(); + + rerender({ threadId: ' ' }); + expect(mockGetThreadMessages).not.toHaveBeenCalled(); + }); + + it('loads and refetches thread messages', async () => { + mockGetThreadMessages + .mockResolvedValueOnce({ messages: [message], count: 1 }) + .mockResolvedValueOnce({ + messages: [{ ...message, id: 'message-2', content: 'updated' }], + count: 1, + }); + const { useThreadMessages } = await import('./useThreadQueries'); + + const { result } = renderHook(() => useThreadMessages('thread-1')); + + await waitFor(() => expect(result.current.data?.messages[0].id).toBe('message-1')); + await act(async () => { + await result.current.refetch(); + }); + + expect(result.current.data?.messages[0].id).toBe('message-2'); + expect(mockGetThreadMessages).toHaveBeenNthCalledWith(1, 'thread-1'); + expect(mockGetThreadMessages).toHaveBeenNthCalledWith(2, 'thread-1'); + }); + + it('exposes refetching state while keeping previous data', async () => { + const nextMessages = deferred<{ messages: ThreadMessage[]; count: number }>(); + mockGetThreadMessages + .mockResolvedValueOnce({ messages: [message], count: 1 }) + .mockReturnValueOnce(nextMessages.promise); + const { useThreadMessages } = await import('./useThreadQueries'); + + const { result } = renderHook(() => useThreadMessages('thread-1')); + + await waitFor(() => expect(result.current.data?.messages[0].id).toBe('message-1')); + + let refetchPromise: Promise; + act(() => { + refetchPromise = result.current.refetch(); + }); + + expect(result.current.isRefetching).toBe(true); + expect(result.current.data?.messages[0].id).toBe('message-1'); + + nextMessages.resolve({ messages: [{ ...message, id: 'message-2' }], count: 1 }); + await act(async () => { + await refetchPromise; + }); + + expect(result.current.isRefetching).toBe(false); + expect(result.current.data?.messages[0].id).toBe('message-2'); + }); + + it('clears previous messages while loading a different thread id', async () => { + const nextMessages = deferred<{ messages: ThreadMessage[]; count: number }>(); + mockGetThreadMessages + .mockResolvedValueOnce({ messages: [message], count: 1 }) + .mockReturnValueOnce(nextMessages.promise); + const { useThreadMessages } = await import('./useThreadQueries'); + + const { result, rerender } = renderHook(({ threadId }) => useThreadMessages(threadId), { + initialProps: { threadId: 'thread-1' }, + }); + + await waitFor(() => expect(result.current.data?.messages[0].id).toBe('message-1')); + + rerender({ threadId: 'thread-2' }); + + await waitFor(() => expect(result.current.data).toBeNull()); + expect(result.current.loading).toBe(true); + expect(result.current.isRefetching).toBe(false); + + nextMessages.resolve({ messages: [{ ...message, id: 'message-2' }], count: 1 }); + await waitFor(() => expect(result.current.data?.messages[0].id).toBe('message-2')); + expect(mockGetThreadMessages).toHaveBeenNthCalledWith(1, 'thread-1'); + expect(mockGetThreadMessages).toHaveBeenNthCalledWith(2, 'thread-2'); + }); + + it('ignores stale message responses after switching thread ids', async () => { + const firstMessages = deferred<{ messages: ThreadMessage[]; count: number }>(); + const secondMessages = deferred<{ messages: ThreadMessage[]; count: number }>(); + mockGetThreadMessages + .mockReturnValueOnce(firstMessages.promise) + .mockReturnValueOnce(secondMessages.promise); + const { useThreadMessages } = await import('./useThreadQueries'); + + const { result, rerender } = renderHook(({ threadId }) => useThreadMessages(threadId), { + initialProps: { threadId: 'thread-1' }, + }); + + rerender({ threadId: 'thread-2' }); + + secondMessages.resolve({ messages: [{ ...message, id: 'message-2' }], count: 1 }); + await waitFor(() => expect(result.current.data?.messages[0].id).toBe('message-2')); + + firstMessages.resolve({ messages: [{ ...message, id: 'message-1' }], count: 1 }); + await act(async () => { + await firstMessages.promise; + }); + + expect(result.current.data?.messages[0].id).toBe('message-2'); + }); +}); diff --git a/app/src/hooks/useThreadQueries.ts b/app/src/hooks/useThreadQueries.ts new file mode 100644 index 000000000..e05b64ee6 --- /dev/null +++ b/app/src/hooks/useThreadQueries.ts @@ -0,0 +1,134 @@ +import debug from 'debug'; +import { useCallback, useEffect, useRef, useState } from 'react'; + +import { threadApi } from '../services/api/threadApi'; +import type { ThreadMessagesData, ThreadsListData } from '../types/thread'; + +const log = debug('hooks:threadQueries'); + +export interface ThreadQueryState { + data: T | null; + loading: boolean; + error: Error | null; + isRefetching: boolean; + refetch: () => Promise; +} + +function normalizeError(error: unknown): Error { + return error instanceof Error ? error : new Error(String(error)); +} + +function useThreadQuery( + queryName: string, + load: () => Promise, + enabled = true, + queryKey = queryName +): ThreadQueryState { + const [data, setData] = useState(null); + const [loading, setLoading] = useState(enabled); + const [error, setError] = useState(null); + const [isRefetching, setIsRefetching] = useState(false); + const requestIdRef = useRef(0); + const dataRef = useRef(null); + const queryKeyRef = useRef(queryKey); + + const execute = useCallback( + async (reason: 'initial' | 'refetch'): Promise => { + if (!enabled) { + log('%s skip disabled reason=%s', queryName, reason); + return undefined; + } + + const requestId = requestIdRef.current + 1; + requestIdRef.current = requestId; + const hasData = dataRef.current !== null; + log('%s start requestId=%d reason=%s hasData=%s', queryName, requestId, reason, hasData); + + setError(null); + if (hasData || reason === 'refetch') { + setIsRefetching(true); + } else { + setLoading(true); + } + + try { + const nextData = await load(); + if (requestIdRef.current !== requestId) { + log('%s ignore stale success requestId=%d', queryName, requestId); + return nextData; + } + dataRef.current = nextData; + setData(nextData); + log('%s success requestId=%d', queryName, requestId); + return nextData; + } catch (caught) { + const nextError = normalizeError(caught); + if (requestIdRef.current !== requestId) { + log('%s ignore stale error requestId=%d error=%o', queryName, requestId, nextError); + return undefined; + } + setError(nextError); + log('%s error requestId=%d error=%o', queryName, requestId, nextError); + return undefined; + } finally { + if (requestIdRef.current === requestId) { + setLoading(false); + setIsRefetching(false); + } + } + }, + [enabled, load, queryName] + ); + + useEffect(() => { + if (queryKeyRef.current !== queryKey) { + requestIdRef.current += 1; + queryKeyRef.current = queryKey; + dataRef.current = null; + setData(null); + setError(null); + setIsRefetching(false); + } + + if (!enabled) { + requestIdRef.current += 1; + dataRef.current = null; + setData(null); + setError(null); + setLoading(false); + setIsRefetching(false); + return; + } + void execute('initial'); + }, [enabled, execute, queryKey]); + + useEffect( + () => () => { + requestIdRef.current += 1; + }, + [] + ); + + const refetch = useCallback(() => execute('refetch'), [execute]); + + return { data, loading, error, isRefetching, refetch }; +} + +export function useThreads(): ThreadQueryState { + const load = useCallback(() => threadApi.getThreads(), []); + return useThreadQuery('threads.list', load); +} + +export function useThreadMessages(threadId?: string | null): ThreadQueryState { + const normalizedThreadId = threadId?.trim() || null; + const load = useCallback( + () => threadApi.getThreadMessages(normalizedThreadId ?? ''), + [normalizedThreadId] + ); + return useThreadQuery( + 'threads.messages', + load, + normalizedThreadId !== null, + `threads.messages:${normalizedThreadId ?? 'disabled'}` + ); +}