From a72ae49e2ce1e66d4604eb0ac0444e3086117d4b Mon Sep 17 00:00:00 2001 From: Wxw-Gu Date: Tue, 18 Aug 2026 17:46:14 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E6=8A=BD=E7=A6=BB=20AI=20=E6=90=9C?= =?UTF-8?q?=E7=B4=A2=E8=BF=90=E8=A1=8C=E7=94=9F=E5=91=BD=E5=91=A8=E6=9C=9F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../components/search/AISearchWorkspace.tsx | 75 ++-- .../components/search/hooks/useAiSearchRun.ts | 145 ++++++++ tests/component/ai-search-run-hook.test.tsx | 339 ++++++++++++++++++ 3 files changed, 513 insertions(+), 46 deletions(-) create mode 100644 src/renderer/src/components/search/hooks/useAiSearchRun.ts create mode 100644 tests/component/ai-search-run-hook.test.tsx diff --git a/src/renderer/src/components/search/AISearchWorkspace.tsx b/src/renderer/src/components/search/AISearchWorkspace.tsx index 5cff5b0..3440745 100644 --- a/src/renderer/src/components/search/AISearchWorkspace.tsx +++ b/src/renderer/src/components/search/AISearchWorkspace.tsx @@ -1,15 +1,10 @@ import React, { useMemo, useRef, useState } from 'react' import * as Popover from '@radix-ui/react-popover' import { aiSearchIntentLabel, aiSearchRangeStart } from '../../../../shared/ai-search' -import type { - AiSearchAgentRun, - AiSearchProgressEvent, - AiSearchTimeRange -} from '../../../../shared/ai-search' +import type { AiSearchProgressEvent, AiSearchTimeRange } from '../../../../shared/ai-search' import type { AISearchWorkspaceProps, - SearchProgressByStage, SearchRange, SearchScope, SearchStage, @@ -42,6 +37,7 @@ import { useSearchHistory } from './hooks/useSearchHistory' import { useKnowledgeStatus } from './hooks/useKnowledgeStatus' import { useExternalProviderConsent } from './hooks/useExternalProviderConsent' import { EVIDENCE_PAGE_SIZE, useEvidenceCollection } from './hooks/useEvidenceCollection' +import { useAiSearchRun } from './hooks/useAiSearchRun' export function AISearchWorkspace({ contacts, @@ -67,15 +63,12 @@ export function AISearchWorkspace({ const [senderNames, setSenderNames] = useState>({}) const [cachedAt, setCachedAt] = useState(0) const [searchTrace, setSearchTrace] = useState(null) - const [searchProgress, setSearchProgress] = useState({}) - const [agentTrace, setAgentTrace] = useState([]) const [searchDetailsOpen, setSearchDetailsOpen] = useState(false) const [historyOpen, setHistoryOpen] = useState(false) const [debugEnabled, setDebugEnabled] = useState(false) const [debugPanelOpen, setDebugPanelOpen] = useState(false) const [debugEntries, setDebugEntries] = useState([]) const [appLogPath, setAppLogPath] = useState('') - const searchRequestIdRef = useRef('') const composerRef = useRef(null) const { evidence, @@ -150,6 +143,15 @@ export function AISearchWorkspace({ settleExternalProviderConsent, clearExternalProviderConsent } = useExternalProviderConsent() + const { + requestId: searchRunRequestId, + progress: searchProgress, + agentTrace, + createRequestId, + startSearch, + cancelSearch, + resetSearchRun + } = useAiSearchRun() const resetSearchResult = (): void => { const reset = createSearchResultResetState() @@ -158,8 +160,7 @@ export function AISearchWorkspace({ clearEvidenceCollection() setCachedAt(reset.cachedAt) setSearchTrace(reset.searchTrace) - setSearchProgress(reset.searchProgress) - setAgentTrace(reset.agentTrace) + resetSearchRun() setSearchDetailsOpen(reset.searchDetailsOpen) } @@ -172,22 +173,6 @@ export function AISearchWorkspace({ ) }, []) - React.useEffect( - () => - window.api.onAiSearchProgress((progress) => { - if (progress.requestId !== searchRequestIdRef.current) return - setSearchProgress((current) => ({ ...current, [progress.stage]: progress })) - if (progress.agentTrace) { - setAgentTrace((current) => - current.some((item) => item.sequence === progress.agentTrace?.sequence) - ? current - : [...current, progress.agentTrace as AiSearchAgentRun['trace'][number]] - ) - } - }), - [] - ) - const addDebugEntry = (message: string, details: Record = {}): void => { const entry = `${new Date().toLocaleTimeString('zh-CN')} ${message} ${JSON.stringify(details)}` setDebugEntries((current) => [entry, ...current].slice(0, 80)) @@ -235,18 +220,15 @@ export function AISearchWorkspace({ const cancelAnalysis = async (): Promise => { clearExternalProviderConsent() - const requestId = searchRequestIdRef.current + const requestId = searchRunRequestId if (!requestId) return - searchRequestIdRef.current = '' setStage('idle') setAnalysisError('') - setSearchProgress({}) - setAgentTrace([]) setSearchDetailsOpen(false) onNotice('已取消本次分析') composerRef.current?.focus() try { - await window.api.cancelAiSearch(requestId) + await cancelSearch() } catch (error) { addDebugEntry('取消检索请求失败', { requestId, @@ -284,7 +266,6 @@ export function AISearchWorkspace({ effectiveRange, normalizedQuery ) - let requestId = '' try { const cached = consumeCacheBypass() ? null : readCachedResult(cacheKey) if (cached) { @@ -298,7 +279,7 @@ export function AISearchWorkspace({ onNotice('已使用最近的检索缓存,可点击刷新数据读取最新消息') return } - requestId = globalThis.crypto?.randomUUID?.() || `search-${Date.now()}` + const requestId = createRequestId() try { if (!(await ensureAiSearchDataConsent(requestId))) { onNotice('已取消本次 AI Search,未执行检索,也未向远程 AI 服务发送聊天内容') @@ -314,8 +295,7 @@ export function AISearchWorkspace({ } setStage('loading') resetSearchResult() - searchRequestIdRef.current = requestId - const searchResult = await window.api.runAiSearch({ + const outcome = await startSearch({ requestId, text: normalizedQuery, scope, @@ -323,7 +303,19 @@ export function AISearchWorkspace({ conversationId: scope === 'conversation' ? activeContact?.md5 : undefined, timeRangeOverride: effectiveTimeRangeOverride }) - if (searchRequestIdRef.current !== requestId) return + if (outcome.kind === 'stale') return + if (outcome.kind === 'cancelled') { + onNotice('已取消本次分析') + setStage('idle') + return + } + if (outcome.kind === 'failed') { + addDebugEntry('检索失败', { error: outcome.error }) + setAnalysisError(outcome.error) + setStage('insufficient') + return + } + const searchResult = outcome.result addDebugEntry('主进程搜索任务完成', { status: searchResult.status, candidateEvidenceCount: searchResult.candidateEvidenceCount, @@ -331,18 +323,12 @@ export function AISearchWorkspace({ elapsedMs: searchResult.elapsedMs, errorStage: searchResult.errorStage }) - if (searchResult.status === 'cancelled') { - onNotice('已取消本次分析') - setStage('idle') - return - } const evidenceItems = mapPipelineEvidence(searchResult.evidence, allContacts) const collectionItems = mapPipelineEvidence( searchResult.evidenceCollection || searchResult.evidence, allContacts ) setSearchTrace(mapSearchResultToTrace(searchResult, evidenceItems.length)) - setAgentTrace(searchResult.agent.trace) setEvidenceResult(evidenceItems, collectionItems) const nextSenderNames = mapEvidenceSenderNames(evidenceItems) setSenderNames(nextSenderNames) @@ -381,13 +367,10 @@ export function AISearchWorkspace({ }) setStage('result') } catch (error) { - if (requestId && searchRequestIdRef.current !== requestId) return const errorMessage = error instanceof Error ? error.message : '读取聊天记录失败' addDebugEntry('检索失败', { error: errorMessage }) setAnalysisError(errorMessage) setStage('insufficient') - } finally { - if (requestId && searchRequestIdRef.current === requestId) searchRequestIdRef.current = '' } } diff --git a/src/renderer/src/components/search/hooks/useAiSearchRun.ts b/src/renderer/src/components/search/hooks/useAiSearchRun.ts new file mode 100644 index 0000000..8d85a99 --- /dev/null +++ b/src/renderer/src/components/search/hooks/useAiSearchRun.ts @@ -0,0 +1,145 @@ +import { useEffect, useRef, useState } from 'react' +import type { + AiSearchAgentRun, + AiSearchPipelineRequest, + AiSearchPipelineResult, + AiSearchProgressEvent +} from '../../../../../shared/ai-search' +import type { SearchProgressByStage } from '../searchTypes' + +export type AiSearchRunStatus = + | 'idle' + | 'starting' + | 'running' + | 'completed' + | 'failed' + | 'cancelled' + +export type AiSearchRunOutcome = + | { kind: 'completed'; requestId: string; result: AiSearchPipelineResult } + | { kind: 'cancelled'; requestId: string; result?: AiSearchPipelineResult } + | { kind: 'failed'; requestId: string; error: string } + | { kind: 'stale'; requestId: string } + +const errorMessage = (error: unknown): string => + error instanceof Error ? error.message : '读取聊天记录失败' + +export function useAiSearchRun(): { + status: AiSearchRunStatus + requestId: string + result: AiSearchPipelineResult | null + error: string + progress: SearchProgressByStage + agentTrace: AiSearchAgentRun['trace'] + createRequestId: () => string + startSearch: (request: AiSearchPipelineRequest) => Promise + cancelSearch: () => Promise + resetSearchRun: () => void + isCurrentRequest: (requestId: string) => boolean +} { + const [status, setStatus] = useState('idle') + const [requestId, setRequestId] = useState('') + const [result, setResult] = useState(null) + const [error, setError] = useState('') + const [progress, setProgress] = useState({}) + const [agentTrace, setAgentTrace] = useState([]) + const requestIdRef = useRef('') + + const createRequestId = (): string => globalThis.crypto?.randomUUID?.() || `search-${Date.now()}` + + const isCurrentRequest = (currentRequestId: string): boolean => + Boolean(currentRequestId) && requestIdRef.current === currentRequestId + + useEffect(() => { + const unsubscribe = window.api.onAiSearchProgress((event: AiSearchProgressEvent) => { + if (!isCurrentRequest(event.requestId)) return + setProgress((current) => ({ ...current, [event.stage]: event })) + if (event.agentTrace) { + setAgentTrace((current) => + current.some((item) => item.sequence === event.agentTrace?.sequence) + ? current + : [...current, event.agentTrace as AiSearchAgentRun['trace'][number]] + ) + } + }) + return () => { + requestIdRef.current = '' + unsubscribe() + } + }, []) + + const startSearch = async (request: AiSearchPipelineRequest): Promise => { + requestIdRef.current = request.requestId + setRequestId(request.requestId) + setStatus('starting') + setError('') + setResult(null) + setProgress({}) + setAgentTrace([]) + setStatus('running') + try { + const searchResult = await window.api.runAiSearch(request) + if (!isCurrentRequest(request.requestId)) + return { kind: 'stale', requestId: request.requestId } + requestIdRef.current = '' + setRequestId('') + setResult(searchResult) + setAgentTrace(searchResult.agent.trace) + if (searchResult.status === 'cancelled') { + setStatus('cancelled') + return { kind: 'cancelled', requestId: request.requestId, result: searchResult } + } + if (searchResult.status === 'failed' || searchResult.status === 'ai_failed') { + setStatus('failed') + } else { + setStatus('completed') + } + return { kind: 'completed', requestId: request.requestId, result: searchResult } + } catch (caughtError) { + if (!isCurrentRequest(request.requestId)) + return { kind: 'stale', requestId: request.requestId } + requestIdRef.current = '' + setRequestId('') + const message = errorMessage(caughtError) + setError(message) + setStatus('failed') + return { kind: 'failed', requestId: request.requestId, error: message } + } + } + + const cancelSearch = async (): Promise => { + const currentRequestId = requestIdRef.current + if (!currentRequestId) return + requestIdRef.current = '' + setRequestId('') + setStatus('cancelled') + setError('') + setProgress({}) + setAgentTrace([]) + await window.api.cancelAiSearch(currentRequestId) + } + + const resetSearchRun = (): void => { + requestIdRef.current = '' + setRequestId('') + setStatus('idle') + setResult(null) + setError('') + setProgress({}) + setAgentTrace([]) + } + + return { + status, + requestId, + result, + error, + progress, + agentTrace, + createRequestId, + startSearch, + cancelSearch, + resetSearchRun, + isCurrentRequest + } +} diff --git a/tests/component/ai-search-run-hook.test.tsx b/tests/component/ai-search-run-hook.test.tsx new file mode 100644 index 0000000..7aa0907 --- /dev/null +++ b/tests/component/ai-search-run-hook.test.tsx @@ -0,0 +1,339 @@ +import { act, renderHook, waitFor } from '@testing-library/react' +import { beforeEach, describe, expect, it, vi } from 'vitest' +import type { + AiSearchAgentTraceItem, + AiSearchPipelineRequest, + AiSearchPipelineResult, + AiSearchProgressEvent, + AiSearchProgressStage +} from '../../src/shared/ai-search' +import { useAiSearchRun } from '../../src/renderer/src/components/search/hooks/useAiSearchRun' +import { makeSearchResult } from './support/ai-search-fixtures' + +const api = { + runAiSearch: vi.fn(), + cancelAiSearch: vi.fn(), + onAiSearchProgress: vi.fn() +} + +const request = (requestId: string, text = '测试问题'): AiSearchPipelineRequest => ({ + requestId, + text, + scope: 'global', + range: '30d' +}) + +const resultFor = ( + requestId: string, + status: 'completed' | 'failed' | 'ai_failed' | 'cancelled' = 'completed' +): AiSearchPipelineResult => makeSearchResult({ requestId, status }) + +const progressFor = ( + requestId: string, + stage: AiSearchProgressStage, + trace?: AiSearchAgentTraceItem +): AiSearchProgressEvent => ({ + requestId, + stage, + status: 'running', + message: `${stage} ${requestId}`, + agentTrace: trace +}) + +const traceFor = (sequence: number): AiSearchAgentTraceItem => ({ + sequence, + event: 'agentDecision', + label: `decision-${sequence}` +}) + +const 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 } +} + +let progressListener: ((event: AiSearchProgressEvent) => void) | undefined +let unsubscribe: ReturnType + +beforeEach(() => { + vi.clearAllMocks() + progressListener = undefined + unsubscribe = vi.fn() + Object.defineProperty(window, 'api', { configurable: true, value: api }) + api.onAiSearchProgress.mockImplementation((listener) => { + progressListener = listener + return unsubscribe + }) + api.cancelAiSearch.mockResolvedValue({ cancelled: true }) +}) + +describe('useAiSearchRun', () => { + it('starts idle and registers exactly one progress listener', () => { + const { result, rerender } = renderHook(() => useAiSearchRun()) + + expect(result.current.status).toBe('idle') + expect(result.current.requestId).toBe('') + expect(result.current.progress).toEqual({}) + expect(result.current.agentTrace).toEqual([]) + rerender() + expect(api.onAiSearchProgress).toHaveBeenCalledOnce() + }) + + it('starts a run with the exact request payload and exposes the running requestId', async () => { + const run = deferred() + api.runAiSearch.mockReturnValue(run.promise) + const { result } = renderHook(() => useAiSearchRun()) + + let startPromise!: ReturnType + act(() => { + startPromise = result.current.startSearch(request('request-a', '原始问题')) + }) + + expect(result.current.status).toBe('running') + expect(result.current.requestId).toBe('request-a') + expect(api.runAiSearch).toHaveBeenCalledWith(request('request-a', '原始问题')) + + await act(async () => { + run.resolve(resultFor('request-a')) + await startPromise + }) + }) + + it('applies the current request result and converges to completed', async () => { + api.runAiSearch.mockResolvedValue(resultFor('request-a')) + const { result } = renderHook(() => useAiSearchRun()) + + let outcome!: Awaited> + await act(async () => { + outcome = await result.current.startSearch(request('request-a')) + }) + + expect(outcome).toMatchObject({ kind: 'completed', requestId: 'request-a' }) + expect(result.current.status).toBe('completed') + expect(result.current.requestId).toBe('') + expect(result.current.result?.requestId).toBe('request-a') + }) + + it('converges to failed when the pipeline returns a failed status', async () => { + api.runAiSearch.mockResolvedValue(resultFor('request-failed', 'failed')) + const { result } = renderHook(() => useAiSearchRun()) + + let outcome!: Awaited> + await act(async () => { + outcome = await result.current.startSearch(request('request-failed')) + }) + + expect(outcome.kind).toBe('completed') + expect(result.current.status).toBe('failed') + expect(result.current.result?.status).toBe('failed') + }) + + it('converges to failed with an error when runAiSearch rejects', async () => { + api.runAiSearch.mockRejectedValue(new Error('Worker failed')) + const { result } = renderHook(() => useAiSearchRun()) + + let outcome!: Awaited> + await act(async () => { + outcome = await result.current.startSearch(request('request-error')) + }) + + expect(outcome).toEqual({ kind: 'failed', requestId: 'request-error', error: 'Worker failed' }) + expect(result.current.status).toBe('failed') + expect(result.current.error).toBe('Worker failed') + expect(result.current.requestId).toBe('') + }) + + it('applies current progress and deduplicates current Agent Trace events', async () => { + const run = deferred() + api.runAiSearch.mockReturnValue(run.promise) + const { result } = renderHook(() => useAiSearchRun()) + act(() => { + void result.current.startSearch(request('request-a')) + }) + + act(() => progressListener?.(progressFor('request-a', 'agent_tool', traceFor(1)))) + act(() => progressListener?.(progressFor('request-a', 'agent_tool', traceFor(1)))) + + expect(result.current.progress.agent_tool?.requestId).toBe('request-a') + expect(result.current.agentTrace).toHaveLength(1) + expect(result.current.agentTrace[0].sequence).toBe(1) + await act(async () => { + run.resolve(resultFor('request-a')) + }) + }) + + it('ignores stale progress and Agent Trace after a newer request starts', async () => { + const runA = deferred() + const runB = deferred() + api.runAiSearch.mockReturnValueOnce(runA.promise).mockReturnValueOnce(runB.promise) + const { result } = renderHook(() => useAiSearchRun()) + let promiseA!: ReturnType + let promiseB!: ReturnType + act(() => { + promiseA = result.current.startSearch(request('request-a')) + promiseB = result.current.startSearch(request('request-b')) + }) + + act(() => { + progressListener?.(progressFor('request-a', 'query_understanding', traceFor(1))) + progressListener?.(progressFor('request-b', 'agent_tool', traceFor(2))) + }) + + expect(result.current.progress.query_understanding).toBeUndefined() + expect(result.current.progress.agent_tool?.requestId).toBe('request-b') + expect(result.current.agentTrace.map((item) => item.sequence)).toEqual([2]) + + await act(async () => { + runA.resolve(resultFor('request-a')) + expect(await promiseA).toEqual({ kind: 'stale', requestId: 'request-a' }) + runB.resolve(resultFor('request-b')) + expect((await promiseB).kind).toBe('completed') + }) + expect(result.current.result?.requestId).toBe('request-b') + }) + + it('ignores a stale result so it cannot overwrite the current request', async () => { + const runA = deferred() + const runB = deferred() + api.runAiSearch.mockReturnValueOnce(runA.promise).mockReturnValueOnce(runB.promise) + const { result } = renderHook(() => useAiSearchRun()) + let promiseA!: ReturnType + let promiseB!: ReturnType + act(() => { + promiseA = result.current.startSearch(request('request-a')) + promiseB = result.current.startSearch(request('request-b')) + }) + + await act(async () => { + runA.resolve(resultFor('request-a')) + expect(await promiseA).toEqual({ kind: 'stale', requestId: 'request-a' }) + expect(result.current.result).toBeNull() + runB.resolve(resultFor('request-b')) + await promiseB + }) + + expect(result.current.result?.requestId).toBe('request-b') + expect(result.current.status).toBe('completed') + }) + + it('ignores a stale error so it cannot overwrite the current request', async () => { + const runA = deferred() + const runB = deferred() + api.runAiSearch.mockReturnValueOnce(runA.promise).mockReturnValueOnce(runB.promise) + const { result } = renderHook(() => useAiSearchRun()) + let promiseA!: ReturnType + let promiseB!: ReturnType + act(() => { + promiseA = result.current.startSearch(request('request-a')) + promiseB = result.current.startSearch(request('request-b')) + }) + + await act(async () => { + runA.reject(new Error('late A error')) + expect(await promiseA).toEqual({ kind: 'stale', requestId: 'request-a' }) + expect(result.current.error).toBe('') + runB.resolve(resultFor('request-b')) + await promiseB + }) + + expect(result.current.error).toBe('') + expect(result.current.result?.requestId).toBe('request-b') + }) + + it('cancels the current request through the existing IPC and converges state', async () => { + const run = deferred() + api.runAiSearch.mockReturnValue(run.promise) + const { result } = renderHook(() => useAiSearchRun()) + act(() => { + void result.current.startSearch(request('request-cancel')) + }) + + await act(async () => result.current.cancelSearch()) + + expect(api.cancelAiSearch).toHaveBeenCalledWith('request-cancel') + expect(result.current.status).toBe('cancelled') + expect(result.current.requestId).toBe('') + expect(result.current.progress).toEqual({}) + expect(result.current.agentTrace).toEqual([]) + }) + + it('ignores a late result after cancellation', async () => { + const run = deferred() + api.runAiSearch.mockReturnValue(run.promise) + const { result } = renderHook(() => useAiSearchRun()) + let startPromise!: ReturnType + act(() => { + startPromise = result.current.startSearch(request('request-cancel')) + }) + await act(async () => result.current.cancelSearch()) + + await act(async () => { + run.resolve(resultFor('request-cancel')) + expect(await startPromise).toEqual({ kind: 'stale', requestId: 'request-cancel' }) + }) + expect(result.current.result).toBeNull() + expect(result.current.status).toBe('cancelled') + }) + + it('does not call cancel IPC when no request is active', async () => { + const { result } = renderHook(() => useAiSearchRun()) + + await act(async () => result.current.cancelSearch()) + + expect(api.cancelAiSearch).not.toHaveBeenCalled() + expect(result.current.status).toBe('idle') + }) + + it('ignores progress after the request has ended', async () => { + api.runAiSearch.mockResolvedValue(resultFor('request-ended')) + const { result } = renderHook(() => useAiSearchRun()) + await act(async () => result.current.startSearch(request('request-ended'))) + const progressBefore = result.current.progress + + act(() => progressListener?.(progressFor('request-ended', 'error', traceFor(9)))) + + expect(result.current.progress).toBe(progressBefore) + expect(result.current.agentTrace).toEqual([]) + }) + + it('resets the run state without cancelling an IPC request', async () => { + const run = deferred() + api.runAiSearch.mockReturnValue(run.promise) + const { result } = renderHook(() => useAiSearchRun()) + act(() => { + void result.current.startSearch(request('request-reset')) + }) + act(() => result.current.resetSearchRun()) + + expect(result.current.status).toBe('idle') + expect(result.current.requestId).toBe('') + expect(result.current.result).toBeNull() + expect(result.current.progress).toEqual({}) + expect(result.current.agentTrace).toEqual([]) + expect(api.cancelAiSearch).not.toHaveBeenCalled() + await act(async () => { + run.resolve(resultFor('request-reset')) + }) + }) + + it('cleans up the progress listener and invalidates an active request on unmount', async () => { + const run = deferred() + api.runAiSearch.mockReturnValue(run.promise) + const { result, unmount } = renderHook(() => useAiSearchRun()) + let startPromise!: ReturnType + act(() => { + startPromise = result.current.startSearch(request('request-unmount')) + }) + + unmount() + expect(unsubscribe).toHaveBeenCalledOnce() + await act(async () => { + run.resolve(resultFor('request-unmount')) + expect(await startPromise).toEqual({ kind: 'stale', requestId: 'request-unmount' }) + }) + }) +})