diff --git a/api/v1/endpoints/analysis.py b/api/v1/endpoints/analysis.py index 3d48fc945..6a31e6a5c 100644 --- a/api/v1/endpoints/analysis.py +++ b/api/v1/endpoints/analysis.py @@ -552,7 +552,7 @@ def trigger_market_review( def get_task_list( status: Optional[str] = Query( None, - description="筛选状态:pending, processing, completed, failed(支持逗号分隔多个)" + description="筛选状态:pending, processing, completed, failed, cancel_requested, cancelled(支持逗号分隔多个)" ), limit: int = Query(20, description="返回数量限制", ge=1, le=100), ) -> TaskListResponse: diff --git a/api/v1/schemas/analysis.py b/api/v1/schemas/analysis.py index dd40ee9e8..bdb6e8e43 100644 --- a/api/v1/schemas/analysis.py +++ b/api/v1/schemas/analysis.py @@ -23,6 +23,8 @@ class TaskStatusEnum(str, Enum): PROCESSING = "processing" COMPLETED = "completed" FAILED = "failed" + CANCEL_REQUESTED = "cancel_requested" + CANCELLED = "cancelled" AnalysisPhase = Literal["auto", "premarket", "intraday", "postmarket"] @@ -263,10 +265,9 @@ class TaskStatus(BaseModel): task_id: str = Field(..., description="任务 ID") trace_id: Optional[str] = Field(None, description="诊断 trace ID") - status: str = Field( + status: TaskStatusEnum = Field( ..., description="任务状态", - pattern="^(pending|processing|completed|failed)$" ) progress: Optional[int] = Field( None, diff --git a/apps/dsa-web/src/components/run-flow/RunFlowGraph.tsx b/apps/dsa-web/src/components/run-flow/RunFlowGraph.tsx index c554bb4df..fcb0ab2d3 100644 --- a/apps/dsa-web/src/components/run-flow/RunFlowGraph.tsx +++ b/apps/dsa-web/src/components/run-flow/RunFlowGraph.tsx @@ -5,6 +5,7 @@ import { useUiLanguage } from '../../contexts/UiLanguageContext'; import type { RunFlowEdge, RunFlowLane, RunFlowNode, RunFlowStatus } from '../../types/runFlow'; import { compactText, + formatDateTime, formatDuration, getNodeDisplayOrder, getRunFlowEdgeKindLabel, @@ -31,9 +32,9 @@ type PositionedNode = RunFlowNode & { const LANE_WIDTH = 292; const NODE_WIDTH = 244; -const NODE_HEIGHT = 108; +const NODE_HEIGHT = 124; const HEADER_HEIGHT = 42; -const ROW_HEIGHT = 126; +const ROW_HEIGHT = 144; const LEFT_PADDING = 20; const TOP_PADDING = 18; const BOTTOM_PADDING = 30; @@ -73,6 +74,34 @@ const getCenteredTrackOffset = (total: number, index: number, step = 12): number (index - (total - 1) / 2) * step ); +const nodeTimeOrder = (node: RunFlowNode): number | null => { + const rawTime = node.startedAt || node.endedAt; + if (!rawTime) { + return null; + } + const parsed = Date.parse(rawTime); + return Number.isFinite(parsed) ? parsed : null; +}; + +const compareLaneNodes = ( + laneId: string, + left: RunFlowNode, + right: RunFlowNode, + originalIndex: Map, +): number => { + const leftOriginal = originalIndex.get(left.id) ?? 0; + const rightOriginal = originalIndex.get(right.id) ?? 0; + if (laneId === 'data_source') { + const leftTime = nodeTimeOrder(left); + const rightTime = nodeTimeOrder(right); + if (leftTime !== null || rightTime !== null) { + return (leftTime ?? Number.MAX_SAFE_INTEGER) - (rightTime ?? Number.MAX_SAFE_INTEGER) + || getNodeDisplayOrder(left, leftOriginal) - getNodeDisplayOrder(right, rightOriginal); + } + } + return getNodeDisplayOrder(left, leftOriginal) - getNodeDisplayOrder(right, rightOriginal); +}; + export const RunFlowGraph: React.FC = ({ lanes, nodes, @@ -81,7 +110,7 @@ export const RunFlowGraph: React.FC = ({ onSelectNode, }) => { const arrowId = useId().replace(/:/g, '-'); - const { t } = useUiLanguage(); + const { language, t } = useUiLanguage(); const laneList = useMemo(() => { const sortedLanes = [...lanes].sort((left, right) => left.order - right.order); const knownLaneIds = new Set(sortedLanes.map((lane) => lane.id)); @@ -122,8 +151,7 @@ export const RunFlowGraph: React.FC = ({ const laneOrderByNode = new Map(); laneList.forEach((lane) => { const laneNodes = [...(grouped.get(lane.id) || [])].sort((left, right) => ( - getNodeDisplayOrder(left, originalIndex.get(left.id) ?? 0) - - getNodeDisplayOrder(right, originalIndex.get(right.id) ?? 0) + compareLaneNodes(lane.id, left, right, originalIndex) )); laneNodes.forEach((node, index) => { laneOrderByNode.set(node.id, index); @@ -169,8 +197,7 @@ export const RunFlowGraph: React.FC = ({ laneList.forEach((lane, lanePosition) => { const laneNodes = [...(grouped.get(lane.id) || [])].sort((left, right) => ( resolvePreferredRow(left.id) - resolvePreferredRow(right.id) - || getNodeDisplayOrder(left, originalIndex.get(left.id) ?? 0) - - getNodeDisplayOrder(right, originalIndex.get(right.id) ?? 0) + || compareLaneNodes(lane.id, left, right, originalIndex) )); const occupiedRows = new Set(); laneNodes.forEach((node) => { @@ -405,6 +432,11 @@ export const RunFlowGraph: React.FC = ({ {formatDuration(node.durationMs, t)} ) : null} + {node.startedAt ? ( + + {t('runFlow.graph.startedAt')}: {formatDateTime(node.startedAt, language, t)} + + ) : null} ); diff --git a/apps/dsa-web/src/components/run-flow/__tests__/RunFlowGraph.test.tsx b/apps/dsa-web/src/components/run-flow/__tests__/RunFlowGraph.test.tsx index 5022a8b8e..2ed7a6472 100644 --- a/apps/dsa-web/src/components/run-flow/__tests__/RunFlowGraph.test.tsx +++ b/apps/dsa-web/src/components/run-flow/__tests__/RunFlowGraph.test.tsx @@ -24,6 +24,7 @@ const nodes: RunFlowNode[] = [ label: '新闻舆情', status: 'fallback', provider: 'AkShare', + startedAt: '2026-06-08T10:00:00', }, ]; @@ -53,6 +54,8 @@ describe('RunFlowGraph', () => { expect(screen.getByText('入口')).toBeInTheDocument(); expect(screen.getByText('数据来源')).toBeInTheDocument(); expect(screen.getByText('降级')).toBeInTheDocument(); + expect(screen.getByTestId('run-flow-node-news')).toHaveTextContent('开始'); + expect(screen.getByTestId('run-flow-node-news')).toHaveTextContent('2026'); expect(screen.getByRole('button', { name: '新闻舆情 节点,状态 Fallback' })).toBeInTheDocument(); fireEvent.click(screen.getByRole('button', { name: '新闻舆情 节点,状态 Fallback' })); @@ -189,4 +192,48 @@ describe('RunFlowGraph', () => { expect(startY).toBe(dailyBottom); expect(endY).toBe(quoteTop); }); + + it('orders data-source lane cards by their observed timestamps', () => { + const timeOrderedNodes: RunFlowNode[] = [ + { + id: 'late-news', + lane: 'data_source', + kind: 'data_source', + label: '新闻舆情', + status: 'success', + startedAt: '2026-06-08T10:00:05', + }, + { + id: 'early-quote', + lane: 'data_source', + kind: 'data_source', + label: '实时行情', + status: 'success', + startedAt: '2026-06-08T10:00:01', + }, + { + id: 'middle-daily', + lane: 'data_source', + kind: 'data_source', + label: '日线K线', + status: 'success', + endedAt: '2026-06-08T10:00:03', + }, + ]; + + render( + , + ); + + expect(Number(screen.getByTestId('run-flow-node-early-quote').dataset.layoutRow)).toBeLessThan( + Number(screen.getByTestId('run-flow-node-middle-daily').dataset.layoutRow), + ); + expect(Number(screen.getByTestId('run-flow-node-middle-daily').dataset.layoutRow)).toBeLessThan( + Number(screen.getByTestId('run-flow-node-late-news').dataset.layoutRow), + ); + }); }); diff --git a/apps/dsa-web/src/components/run-flow/__tests__/RunFlowPanel.test.tsx b/apps/dsa-web/src/components/run-flow/__tests__/RunFlowPanel.test.tsx index 6232cac00..2ea3e600d 100644 --- a/apps/dsa-web/src/components/run-flow/__tests__/RunFlowPanel.test.tsx +++ b/apps/dsa-web/src/components/run-flow/__tests__/RunFlowPanel.test.tsx @@ -8,6 +8,7 @@ import { RunFlowPanel } from '../RunFlowPanel'; vi.mock('../../../api/analysis', () => ({ analysisApi: { getTaskFlow: vi.fn(), + getTaskStreamUrl: vi.fn(() => 'http://localhost/api/v1/analysis/tasks/stream'), }, })); diff --git a/apps/dsa-web/src/components/tasks/TaskPanel.tsx b/apps/dsa-web/src/components/tasks/TaskPanel.tsx index 5c1805082..5d0d33a3f 100644 --- a/apps/dsa-web/src/components/tasks/TaskPanel.tsx +++ b/apps/dsa-web/src/components/tasks/TaskPanel.tsx @@ -21,9 +21,15 @@ const TaskItem: React.FC = ({ task, onOpenRunFlow }) => { const { language, t } = useUiLanguage(); const isPending = task.status === 'pending'; const isProcessing = task.status === 'processing'; - const statusLabel = isProcessing ? t('taskPanel.processing') : t('taskPanel.pending'); - const statusVariant = isProcessing ? 'info' : 'default'; - const statusTone = isProcessing ? 'info' : 'neutral'; + const isCancelRequested = task.status === 'cancel_requested'; + const isCancelled = task.status === 'cancelled'; + const statusLabel = isCancelRequested + ? t('taskPanel.cancelRequested') + : isCancelled + ? t('taskPanel.cancelled') + : isProcessing ? t('taskPanel.processing') : t('taskPanel.pending'); + const statusVariant = isCancelRequested ? 'warning' : isProcessing ? 'info' : 'default'; + const statusTone = isCancelRequested ? 'warning' : isProcessing ? 'info' : 'neutral'; const progress = Math.max(0, Math.min(100, task.progress || 0)); const traceId = (task.traceId || '').trim(); const requestedPhaseLabel = getRequestedPhaseLabel(task.analysisPhase, language); @@ -35,13 +41,15 @@ const TaskItem: React.FC = ({ task, onOpenRunFlow }) => {
{isProcessing ? ( + ) : isCancelRequested ? ( + ) : isPending ? ( ) : null}
{/* 任务信息 */} -
+
{task.stockName || task.stockCode} @@ -93,7 +101,7 @@ const TaskItem: React.FC = ({ task, onOpenRunFlow }) => {
{/* 状态标签 */} -
+
{onOpenRunFlow ? ( @@ -120,7 +128,7 @@ const TaskItem: React.FC = ({ task, onOpenRunFlow }) => { className="min-w-[4.75rem] justify-center gap-1.5 shadow-none" aria-label={t('taskPanel.statusAria', { status: statusLabel })} > - + {statusLabel}
@@ -156,9 +164,9 @@ export const TaskPanel: React.FC = ({ onOpenRunFlow, }) => { const { t } = useUiLanguage(); - // 筛选活跃任务(pending 和 processing) + // 筛选活跃任务(pending / processing / cancel requested) const activeTasks = tasks.filter( - (t) => t.status === 'pending' || t.status === 'processing' + (t) => t.status === 'pending' || t.status === 'processing' || t.status === 'cancel_requested' ); // 无任务或不可见时不渲染 diff --git a/apps/dsa-web/src/components/tasks/__tests__/TaskPanel.test.tsx b/apps/dsa-web/src/components/tasks/__tests__/TaskPanel.test.tsx index 3fb018207..31ed23bd8 100644 --- a/apps/dsa-web/src/components/tasks/__tests__/TaskPanel.test.tsx +++ b/apps/dsa-web/src/components/tasks/__tests__/TaskPanel.test.tsx @@ -86,6 +86,39 @@ describe('TaskPanel', () => { expect(onOpenRunFlow).toHaveBeenCalledWith(baseTask); }); + it('keeps cancel-requested tasks visible without rendering them as failed', () => { + render( + , + ); + + expect(screen.getByText('贵州茅台')).toBeInTheDocument(); + expect(screen.getByLabelText('任务状态:请求取消')).toBeInTheDocument(); + expect(screen.queryByText('失败')).not.toBeInTheDocument(); + }); + + it('does not keep cancelled terminal tasks in the active task panel', () => { + const { container } = render( + , + ); + + expect(container).toBeEmptyDOMElement(); + }); + it('does not render when there are no active tasks', () => { const { container } = render( ({ + analysisApi: { + getTaskFlow: vi.fn(), + }, +})); + +vi.mock('../../api/history', () => ({ + historyApi: { + getRecordFlow: vi.fn(), + }, +})); + +const taskStreamCalls: UseTaskStreamOptions[] = []; + +vi.mock('../useTaskStream', () => ({ + useTaskStream: (options: UseTaskStreamOptions) => { + taskStreamCalls.push(options); + return { + isConnected: true, + reconnect: vi.fn(), + disconnect: vi.fn(), + }; + }, +})); + +const snapshot: RunFlowSnapshot = { + taskId: 'task-1', + traceId: 'trace-1', + stockCode: '600519', + stockName: '贵州茅台', + status: 'running', + generatedAt: '2026-06-08T08:00:00Z', + summary: { + elapsedMs: null, + failedAttempts: 0, + fallbackCount: 0, + model: null, + dataSourceCount: 0, + eventCount: 1, + }, + lanes: [ + { id: 'entry', label: '入口', order: 1 }, + { id: 'data_source', label: '数据来源', order: 2 }, + ], + nodes: [ + { + id: 'task_queue', + lane: 'entry', + kind: 'queue', + label: '任务队列', + status: 'running', + }, + ], + edges: [], + events: [ + { + id: 'evt-1', + timestamp: '2026-06-08T08:00:00Z', + severity: 'info', + type: 'task_started', + nodeId: 'task_queue', + title: '任务开始执行', + }, + ], +}; + +function createDeferred() { + let resolve!: (value: T) => void; + let reject!: (reason?: unknown) => void; + const promise = new Promise((promiseResolve, promiseReject) => { + resolve = promiseResolve; + reject = promiseReject; + }); + return { promise, resolve, reject }; +} + +describe('useRunFlowSnapshot', () => { + beforeEach(() => { + vi.clearAllMocks(); + taskStreamCalls.length = 0; + }); + + it('merges active task flow events, strips node metadata, and refetches after stream errors', async () => { + vi.mocked(analysisApi.getTaskFlow).mockResolvedValue(snapshot); + + const { result } = renderHook(() => useRunFlowSnapshot({ + source: { type: 'task', taskId: 'task-1' }, + enabled: true, + })); + + await waitFor(() => expect(result.current.snapshot).not.toBeNull()); + + act(() => { + taskStreamCalls.at(-1)?.onTaskFlowEvent?.( + { + taskId: 'task-1', + stockCode: '600519', + status: 'processing', + progress: 30, + reportType: 'detailed', + createdAt: '2026-06-08T08:00:00Z', + }, + { + id: 'flow-1', + timestamp: '2026-06-08T08:00:01Z', + severity: 'success', + type: 'provider_run', + nodeId: 'provider_daily_1', + title: '日线K线成功', + metadata: { + provider: 'DailyFetcher', + node: { + id: 'provider_daily_1', + lane: 'data_source', + kind: 'data_source', + label: '日线K线', + status: 'success', + }, + }, + }, + ); + }); + + expect(result.current.snapshot?.events).toHaveLength(2); + expect(result.current.snapshot?.events[1].metadata).not.toHaveProperty('node'); + expect(result.current.snapshot?.nodes.some((node) => node.id === 'provider_daily_1')).toBe(true); + + act(() => { + taskStreamCalls.at(-1)?.onError?.(new Event('error')); + }); + + await waitFor(() => expect(analysisApi.getTaskFlow).toHaveBeenCalledTimes(2)); + }); + + it('does not enable task stream for history snapshots', async () => { + vi.mocked(historyApi.getRecordFlow).mockResolvedValue({ ...snapshot, status: 'success' }); + + renderHook(() => useRunFlowSnapshot({ + source: { type: 'history', recordId: 7 }, + enabled: true, + })); + + await waitFor(() => expect(historyApi.getRecordFlow).toHaveBeenCalledWith(7)); + expect(taskStreamCalls.at(-1)?.enabled).toBe(false); + }); + + it('replays buffered flow events into refetched task snapshots', async () => { + const initialRequest = createDeferred(); + const refreshedRequest = createDeferred(); + vi.mocked(analysisApi.getTaskFlow) + .mockReturnValueOnce(initialRequest.promise) + .mockReturnValueOnce(refreshedRequest.promise); + + const { result } = renderHook(() => useRunFlowSnapshot({ + source: { type: 'task', taskId: 'task-1' }, + enabled: true, + })); + + act(() => { + initialRequest.resolve(snapshot); + }); + + await waitFor(() => expect(result.current.snapshot?.events).toHaveLength(1)); + + act(() => { + taskStreamCalls.at(-1)?.onTaskCompleted?.({ + taskId: 'task-1', + stockCode: '600519', + status: 'completed', + progress: 100, + reportType: 'detailed', + createdAt: '2026-06-08T08:00:00Z', + }); + taskStreamCalls.at(-1)?.onTaskFlowEvent?.( + { + taskId: 'task-1', + stockCode: '600519', + status: 'processing', + progress: 99, + reportType: 'detailed', + createdAt: '2026-06-08T08:00:00Z', + }, + { + id: 'flow-late', + timestamp: '2026-06-08T08:00:02Z', + severity: 'success', + type: 'provider_run', + nodeId: 'provider_news_1', + title: '新闻检索成功', + metadata: { + provider: 'NewsFetcher', + node: { + id: 'provider_news_1', + lane: 'data_source', + kind: 'data_source', + label: '新闻 · NewsFetcher', + status: 'success', + }, + }, + }, + ); + }); + + await waitFor(() => expect(analysisApi.getTaskFlow).toHaveBeenCalledTimes(2)); + + act(() => { + refreshedRequest.resolve({ + ...snapshot, + status: 'success', + events: snapshot.events, + nodes: snapshot.nodes, + }); + }); + + await waitFor(() => expect(result.current.snapshot?.status).toBe('success')); + const lateEvent = result.current.snapshot?.events.find((event) => event.id === 'flow-late'); + expect(lateEvent).toBeDefined(); + expect(lateEvent?.metadata).not.toHaveProperty('node'); + expect(result.current.snapshot?.nodes.some((node) => node.id === 'provider_news_1')).toBe(true); + }); +}); diff --git a/apps/dsa-web/src/hooks/__tests__/useTaskStream.test.tsx b/apps/dsa-web/src/hooks/__tests__/useTaskStream.test.tsx index 25fc778fe..ec5c33aa1 100644 --- a/apps/dsa-web/src/hooks/__tests__/useTaskStream.test.tsx +++ b/apps/dsa-web/src/hooks/__tests__/useTaskStream.test.tsx @@ -1,5 +1,5 @@ -import { renderHook } from '@testing-library/react'; -import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, renderHook, waitFor } from '@testing-library/react'; +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; import { useTaskStream } from '../useTaskStream'; const { getTaskStreamUrl } = vi.hoisted(() => ({ @@ -21,26 +21,43 @@ type MockEventSourceInstance = { describe('useTaskStream', () => { let eventSourceInstance: MockEventSourceInstance; + let eventSourceInstances: MockEventSourceInstance[]; beforeEach(() => { vi.clearAllMocks(); + eventSourceInstances = []; + eventSourceInstance = createEventSourceInstance(); - eventSourceInstance = { - listeners: {}, - addEventListener: vi.fn((type: string, listener: (event: MessageEvent) => void) => { - eventSourceInstance.listeners[type] = listener; - }), - close: vi.fn(), - onerror: null, - }; + function createEventSourceInstance(): MockEventSourceInstance { + const instance: MockEventSourceInstance = { + listeners: {}, + addEventListener: vi.fn((type: string, listener: (event: MessageEvent) => void) => { + instance.listeners[type] = listener; + }), + close: vi.fn(), + onerror: null, + }; + return instance; + } class MockEventSource { - addEventListener = eventSourceInstance.addEventListener; - close = eventSourceInstance.close; - onerror = eventSourceInstance.onerror; + addEventListener: MockEventSourceInstance['addEventListener']; + close: MockEventSourceInstance['close']; constructor(...args: unknown[]) { void args; + const instance = createEventSourceInstance(); + eventSourceInstance = instance; + eventSourceInstances.push(instance); + this.addEventListener = instance.addEventListener; + this.close = instance.close; + Object.defineProperty(this, 'onerror', { + configurable: true, + get: () => instance.onerror, + set: (handler: ((event: Event) => void) | null) => { + instance.onerror = handler; + }, + }); } } @@ -51,20 +68,27 @@ describe('useTaskStream', () => { }); }); - it('closes the SSE connection when the hook unmounts', () => { + afterEach(() => { + cleanup(); + vi.useRealTimers(); + }); + + it('closes the SSE connection when the hook unmounts', async () => { const { unmount } = renderHook(() => useTaskStream({ enabled: true })); - expect(getTaskStreamUrl).toHaveBeenCalledTimes(1); + await waitFor(() => expect(getTaskStreamUrl).toHaveBeenCalledTimes(1)); unmount(); expect(eventSourceInstance.close).toHaveBeenCalled(); }); - it('parses task_progress events and forwards the updated task payload', () => { + it('parses task_progress events and forwards the updated task payload', async () => { const onTaskProgress = vi.fn(); + const onTaskFlowEvent = vi.fn(); - renderHook(() => useTaskStream({ enabled: true, onTaskProgress })); + renderHook(() => useTaskStream({ enabled: true, onTaskProgress, onTaskFlowEvent })); + await waitFor(() => expect(eventSourceInstance.listeners.task_progress).toBeDefined()); eventSourceInstance.listeners.task_progress?.( new MessageEvent('task_progress', { @@ -80,6 +104,23 @@ describe('useTaskStream', () => { analysis_phase: 'intraday', created_at: '2026-03-29T08:00:00Z', skills: ['growth_quality'], + flow_event: { + id: 'flow-1', + timestamp: '2026-03-29T08:00:01Z', + severity: 'success', + type: 'provider_run', + node_id: 'provider_daily_1', + title: '日线K线成功', + metadata: { + node: { + id: 'provider_daily_1', + lane: 'data_source', + kind: 'data_source', + label: '日线K线', + status: 'success', + }, + }, + }, }), }), ); @@ -102,5 +143,88 @@ describe('useTaskStream', () => { analysisPhase: 'intraday', skills: ['growth_quality'], }); + expect(onTaskFlowEvent).toHaveBeenCalledWith( + expect.objectContaining({ taskId: 'task-1' }), + expect.objectContaining({ + id: 'flow-1', + nodeId: 'provider_daily_1', + metadata: expect.objectContaining({ + node: expect.objectContaining({ lane: 'data_source' }), + }), + }), + ); + }); + + it('shares one SSE connection across multiple hook instances', async () => { + const firstConnected = vi.fn(); + const secondConnected = vi.fn(); + const firstProgress = vi.fn(); + const secondProgress = vi.fn(); + + const first = renderHook(() => useTaskStream({ + enabled: true, + onConnected: firstConnected, + onTaskProgress: firstProgress, + })); + const second = renderHook(() => useTaskStream({ + enabled: true, + onConnected: secondConnected, + onTaskProgress: secondProgress, + })); + + await waitFor(() => expect(eventSourceInstances).toHaveLength(1)); + expect(getTaskStreamUrl).toHaveBeenCalledTimes(1); + + eventSourceInstance.listeners.connected?.(new MessageEvent('connected')); + + await waitFor(() => expect(first.result.current.isConnected).toBe(true)); + expect(second.result.current.isConnected).toBe(true); + expect(firstConnected).toHaveBeenCalledTimes(1); + expect(secondConnected).toHaveBeenCalledTimes(1); + + eventSourceInstance.listeners.task_progress?.( + new MessageEvent('task_progress', { + data: JSON.stringify({ + task_id: 'task-1', + stock_code: '600519', + status: 'processing', + progress: 50, + report_type: 'detailed', + created_at: '2026-03-29T08:00:00Z', + }), + }), + ); + + expect(firstProgress).toHaveBeenCalledTimes(1); + expect(secondProgress).toHaveBeenCalledTimes(1); + + first.unmount(); + expect(eventSourceInstance.close).not.toHaveBeenCalled(); + + second.unmount(); + expect(eventSourceInstance.close).toHaveBeenCalledTimes(1); + }); + + it('reconnects the shared stream once after errors', async () => { + vi.useFakeTimers(); + const firstError = vi.fn(); + const secondError = vi.fn(); + + renderHook(() => useTaskStream({ enabled: true, onError: firstError })); + renderHook(() => useTaskStream({ enabled: true, onError: secondError })); + + await vi.runOnlyPendingTimersAsync(); + expect(eventSourceInstances).toHaveLength(1); + + eventSourceInstance.onerror?.(new Event('error')); + + expect(firstError).toHaveBeenCalledTimes(1); + expect(secondError).toHaveBeenCalledTimes(1); + expect(eventSourceInstances[0].close).toHaveBeenCalledTimes(1); + + await vi.advanceTimersByTimeAsync(3000); + + expect(eventSourceInstances).toHaveLength(2); + expect(getTaskStreamUrl).toHaveBeenCalledTimes(2); }); }); diff --git a/apps/dsa-web/src/hooks/useRunFlowSnapshot.ts b/apps/dsa-web/src/hooks/useRunFlowSnapshot.ts index 992aff812..981aeb308 100644 --- a/apps/dsa-web/src/hooks/useRunFlowSnapshot.ts +++ b/apps/dsa-web/src/hooks/useRunFlowSnapshot.ts @@ -1,8 +1,9 @@ -import { useCallback, useEffect, useMemo, useState } from 'react'; +import { useCallback, useEffect, useMemo, useRef, useState } from 'react'; import { analysisApi } from '../api/analysis'; import { getParsedApiError, type ParsedApiError } from '../api/error'; import { historyApi } from '../api/history'; -import type { RunFlowSnapshot, RunFlowSnapshotSource } from '../types/runFlow'; +import type { RunFlowEvent, RunFlowNode, RunFlowSnapshot, RunFlowSnapshotSource } from '../types/runFlow'; +import { useTaskStream } from './useTaskStream'; interface UseRunFlowSnapshotOptions { source?: RunFlowSnapshotSource | null; @@ -22,6 +23,8 @@ type RunFlowRequestState = { error: ParsedApiError | null; }; +const MAX_BUFFERED_FLOW_EVENTS = 50; + const getSourceKey = (source?: RunFlowSnapshotSource | null): string => { if (!source) { return 'none'; @@ -41,6 +44,74 @@ const isUsableSource = (source?: RunFlowSnapshotSource | null): source is RunFlo return Number.isFinite(source.recordId); }; +const eventTime = (event: RunFlowEvent): number => ( + event.timestamp ? Date.parse(event.timestamp) || 0 : 0 +); + +const mergeEvents = (events: RunFlowEvent[], incoming: RunFlowEvent): RunFlowEvent[] => { + const byId = new Map(); + [...events, incoming].forEach((event, index) => { + byId.set(event.id || `event-${index}`, event); + }); + return Array.from(byId.values()).sort((left, right) => eventTime(left) - eventTime(right)); +}; + +const isRunFlowNode = (value: unknown): value is RunFlowNode => { + if (!value || typeof value !== 'object') { + return false; + } + const node = value as Partial; + return Boolean(node.id && node.lane && node.kind && node.label && node.status); +}; + +const mergeFlowEventIntoSnapshot = ( + snapshot: RunFlowSnapshot, + flowEvent: RunFlowEvent, +): RunFlowSnapshot => { + const nodeCandidate = flowEvent.metadata?.node; + const eventMetadata = { ...(flowEvent.metadata || {}) }; + delete eventMetadata.node; + const displayEvent: RunFlowEvent = { + ...flowEvent, + metadata: eventMetadata, + }; + const events = mergeEvents(snapshot.events, displayEvent); + const shouldMergeNode = isRunFlowNode(nodeCandidate) + && !snapshot.nodes.some((node) => node.id === nodeCandidate.id); + const nodes = shouldMergeNode + ? [...snapshot.nodes, nodeCandidate] + : snapshot.nodes; + + return { + ...snapshot, + nodes, + events, + summary: { + ...snapshot.summary, + eventCount: events.length, + }, + generatedAt: flowEvent.timestamp || snapshot.generatedAt, + }; +}; + +const rememberFlowEvent = (events: RunFlowEvent[], flowEvent: RunFlowEvent): RunFlowEvent[] => { + const byId = new Map(); + [...events, flowEvent].forEach((event, index) => { + byId.set(event.id || `${event.type}:${event.nodeId || 'none'}:${event.timestamp || index}`, event); + }); + return Array.from(byId.values()) + .sort((left, right) => eventTime(left) - eventTime(right)) + .slice(-MAX_BUFFERED_FLOW_EVENTS); +}; + +const replayFlowEvents = ( + snapshot: RunFlowSnapshot, + flowEvents: RunFlowEvent[], +): RunFlowSnapshot => flowEvents.reduce( + (currentSnapshot, flowEvent) => mergeFlowEventIntoSnapshot(currentSnapshot, flowEvent), + snapshot, +); + export function useRunFlowSnapshot({ source, enabled = true, @@ -57,11 +128,51 @@ export function useRunFlowSnapshot({ const recordId = source?.type === 'history' ? source.recordId : null; const requestKey = `${sourceKey}:${reloadToken}`; const shouldLoad = enabled && isUsableSource(source); + const flowEventBufferRef = useRef([]); const refetch = useCallback(async () => { setReloadToken((value) => value + 1); }, []); + useEffect(() => { + flowEventBufferRef.current = []; + }, [sourceKey]); + + useTaskStream({ + enabled: shouldLoad && sourceType === 'task', + onTaskFlowEvent: (task, flowEvent) => { + if (task.taskId !== taskId) { + return; + } + flowEventBufferRef.current = rememberFlowEvent(flowEventBufferRef.current, flowEvent); + setRequestState((current) => { + const hasFreshState = current.requestKey === requestKey && current.snapshot; + if (!hasFreshState) { + return current; + } + return { + ...current, + snapshot: mergeFlowEventIntoSnapshot(current.snapshot as RunFlowSnapshot, flowEvent), + }; + }); + }, + onTaskCompleted: (task) => { + if (task.taskId === taskId) { + void refetch(); + } + }, + onTaskFailed: (task) => { + if (task.taskId === taskId) { + void refetch(); + } + }, + onError: () => { + if (sourceType === 'task') { + void refetch(); + } + }, + }); + useEffect(() => { if (!shouldLoad || !sourceType) { return undefined; @@ -76,9 +187,12 @@ export function useRunFlowSnapshot({ request .then((result) => { if (active) { + const snapshot = sourceType === 'task' + ? replayFlowEvents(result, flowEventBufferRef.current) + : result; setRequestState({ requestKey, - snapshot: result, + snapshot, error: null, }); } diff --git a/apps/dsa-web/src/hooks/useTaskStream.ts b/apps/dsa-web/src/hooks/useTaskStream.ts index 3b0debaa0..0074a3f5a 100644 --- a/apps/dsa-web/src/hooks/useTaskStream.ts +++ b/apps/dsa-web/src/hooks/useTaskStream.ts @@ -1,6 +1,8 @@ -import { useEffect, useRef, useCallback, useState } from 'react'; +import { useEffect, useRef, useCallback, useState, type MutableRefObject } from 'react'; import { analysisApi } from '../api/analysis'; +import { toCamelCase } from '../api/utils'; import type { TaskInfo } from '../types/analysis'; +import type { RunFlowEvent } from '../types/runFlow'; /** * SSE event types. @@ -20,6 +22,7 @@ export type SSEEventType = export interface SSEEvent { type: SSEEventType; task?: TaskInfo; + flowEvent?: RunFlowEvent; timestamp?: string; } @@ -37,6 +40,8 @@ export interface UseTaskStreamOptions { onTaskProgress?: (task: TaskInfo) => void; /** Task failed callback */ onTaskFailed?: (task: TaskInfo) => void; + /** Incremental run-flow event callback carried by task_progress */ + onTaskFlowEvent?: (task: TaskInfo, event: RunFlowEvent) => void; /** Connected callback */ onConnected?: () => void; /** Connection error callback */ @@ -61,6 +66,198 @@ export interface UseTaskStreamResult { disconnect: () => void; } +type TaskStreamCallbacks = Pick< + UseTaskStreamOptions, + | 'onTaskCreated' + | 'onTaskStarted' + | 'onTaskCompleted' + | 'onTaskProgress' + | 'onTaskFailed' + | 'onTaskFlowEvent' + | 'onConnected' + | 'onError' +>; + +type ParsedTaskStreamPayload = { + task: TaskInfo; + flowEvent?: RunFlowEvent; +}; + +type TaskStreamSubscriber = { + callbacksRef: MutableRefObject; + setIsConnected: (value: boolean) => void; + autoReconnect: boolean; + reconnectDelay: number; +}; + +let sharedEventSource: EventSource | null = null; +let sharedReconnectTimeout: ReturnType | null = null; +let sharedConnected = false; +let nextSubscriberId = 1; +const subscribers = new Map(); + +// Convert snake_case payloads into camelCase TaskInfo objects. +const toTaskInfo = (data: Record): TaskInfo => { + const task: TaskInfo = { + taskId: data.task_id as string, + stockCode: data.stock_code as string, + stockName: data.stock_name as string | undefined, + status: data.status as TaskInfo['status'], + progress: data.progress as number, + message: data.message as string | undefined, + reportType: data.report_type as string, + createdAt: data.created_at as string, + startedAt: data.started_at as string | undefined, + completedAt: data.completed_at as string | undefined, + error: data.error as string | undefined, + originalQuery: data.original_query as string | undefined, + selectionSource: data.selection_source as string | undefined, + analysisPhase: data.analysis_phase as TaskInfo['analysisPhase'], + skills: Array.isArray(data.skills) ? data.skills.map(String) : undefined, + }; + + if (typeof data.trace_id === 'string' && data.trace_id.trim()) { + task.traceId = data.trace_id; + } + + return task; +}; + +const parseEventData = (eventData: string): ParsedTaskStreamPayload | null => { + try { + const data = JSON.parse(eventData); + const task = toTaskInfo(data); + const flowEvent = data.flow_event + ? toCamelCase(data.flow_event) + : undefined; + return { task, flowEvent }; + } catch (e) { + console.error('Failed to parse SSE event data:', e); + return null; + } +}; + +const notifyConnectionState = (connected: boolean) => { + sharedConnected = connected; + subscribers.forEach((subscriber) => subscriber.setIsConnected(connected)); +}; + +const forEachSubscriber = (notify: (callbacks: TaskStreamCallbacks) => void) => { + subscribers.forEach((subscriber) => notify(subscriber.callbacksRef.current)); +}; + +const clearSharedReconnect = () => { + if (sharedReconnectTimeout) { + clearTimeout(sharedReconnectTimeout); + sharedReconnectTimeout = null; + } +}; + +const closeSharedConnection = () => { + clearSharedReconnect(); + if (sharedEventSource) { + sharedEventSource.close(); + sharedEventSource = null; + } + notifyConnectionState(false); +}; + +const scheduleSharedReconnect = () => { + if (sharedReconnectTimeout || subscribers.size === 0) { + return; + } + const reconnectDelays = Array.from(subscribers.values()) + .filter((subscriber) => subscriber.autoReconnect) + .map((subscriber) => subscriber.reconnectDelay); + if (reconnectDelays.length === 0) { + return; + } + const reconnectDelay = Math.min(...reconnectDelays); + sharedReconnectTimeout = setTimeout(() => { + sharedReconnectTimeout = null; + connectSharedStream(); + }, reconnectDelay); +}; + +function connectSharedStream() { + if (sharedEventSource || subscribers.size === 0) { + return; + } + + if (typeof window.EventSource !== 'function') { + notifyConnectionState(false); + return; + } + + const url = analysisApi.getTaskStreamUrl(); + const eventSource = new window.EventSource(url, { withCredentials: true }); + sharedEventSource = eventSource; + + eventSource.addEventListener('connected', () => { + notifyConnectionState(true); + forEachSubscriber((callbacks) => callbacks.onConnected?.()); + }); + + eventSource.addEventListener('task_created', (e) => { + const payload = parseEventData((e as MessageEvent).data); + if (payload) { + forEachSubscriber((callbacks) => callbacks.onTaskCreated?.(payload.task)); + } + }); + + eventSource.addEventListener('task_started', (e) => { + const payload = parseEventData((e as MessageEvent).data); + if (payload) { + forEachSubscriber((callbacks) => callbacks.onTaskStarted?.(payload.task)); + } + }); + + eventSource.addEventListener('task_progress', (e) => { + const payload = parseEventData((e as MessageEvent).data); + if (payload) { + forEachSubscriber((callbacks) => { + callbacks.onTaskProgress?.(payload.task); + if (payload.flowEvent) { + callbacks.onTaskFlowEvent?.(payload.task, payload.flowEvent); + } + }); + } + }); + + eventSource.addEventListener('task_completed', (e) => { + const payload = parseEventData((e as MessageEvent).data); + if (payload) { + forEachSubscriber((callbacks) => callbacks.onTaskCompleted?.(payload.task)); + } + }); + + eventSource.addEventListener('task_failed', (e) => { + const payload = parseEventData((e as MessageEvent).data); + if (payload) { + forEachSubscriber((callbacks) => callbacks.onTaskFailed?.(payload.task)); + } + }); + + eventSource.addEventListener('heartbeat', () => { + // Optional place to record the latest heartbeat timestamp. + }); + + eventSource.onerror = (error) => { + notifyConnectionState(false); + forEachSubscriber((callbacks) => callbacks.onError?.(error)); + if (sharedEventSource === eventSource) { + eventSource.close(); + sharedEventSource = null; + } + scheduleSharedReconnect(); + }; +} + +const reconnectSharedStream = () => { + closeSharedConnection(); + connectSharedStream(); +}; + /** * Task-stream SSE hook for realtime task status updates. */ @@ -71,6 +268,7 @@ export function useTaskStream(options: UseTaskStreamOptions = {}): UseTaskStream onTaskCompleted, onTaskProgress, onTaskFailed, + onTaskFlowEvent, onConnected, onError, autoReconnect = true, @@ -78,18 +276,18 @@ export function useTaskStream(options: UseTaskStreamOptions = {}): UseTaskStream enabled = true, } = options; - const eventSourceRef = useRef(null); const [isConnected, setIsConnected] = useState(false); - const reconnectTimeoutRef = useRef | null>(null); - const connectRef = useRef<() => void>(() => {}); + const subscriberIdRef = useRef(null); + const connectTimerRef = useRef | null>(null); // Store callbacks in a ref to avoid reconnecting on every render. - const callbacksRef = useRef({ + const callbacksRef = useRef({ onTaskCreated, onTaskStarted, onTaskCompleted, onTaskProgress, onTaskFailed, + onTaskFlowEvent, onConnected, onError, }); @@ -102,154 +300,69 @@ export function useTaskStream(options: UseTaskStreamOptions = {}): UseTaskStream onTaskCompleted, onTaskProgress, onTaskFailed, + onTaskFlowEvent, onConnected, onError, }; }); - // Convert snake_case payloads into camelCase TaskInfo objects. - const toCamelCase = (data: Record): TaskInfo => { - const task: TaskInfo = { - taskId: data.task_id as string, - stockCode: data.stock_code as string, - stockName: data.stock_name as string | undefined, - status: data.status as TaskInfo['status'], - progress: data.progress as number, - message: data.message as string | undefined, - reportType: data.report_type as string, - createdAt: data.created_at as string, - startedAt: data.started_at as string | undefined, - completedAt: data.completed_at as string | undefined, - error: data.error as string | undefined, - originalQuery: data.original_query as string | undefined, - selectionSource: data.selection_source as string | undefined, - analysisPhase: data.analysis_phase as TaskInfo['analysisPhase'], - skills: Array.isArray(data.skills) ? data.skills.map(String) : undefined, - }; - - if (typeof data.trace_id === 'string' && data.trace_id.trim()) { - task.traceId = data.trace_id; - } - - return task; - }; - - // Parse an SSE payload. - const parseEventData = useCallback((eventData: string): TaskInfo | null => { - try { - const data = JSON.parse(eventData); - return toCamelCase(data); - } catch (e) { - console.error('Failed to parse SSE event data:', e); - return null; - } - }, []); - - // Create an EventSource connection. - const connect = useCallback(() => { - if (eventSourceRef.current) { - eventSourceRef.current.close(); - } - - const url = analysisApi.getTaskStreamUrl(); - const eventSource = new EventSource(url, { withCredentials: true }); - eventSourceRef.current = eventSource; - - // Connected event - eventSource.addEventListener('connected', () => { - setIsConnected(true); - callbacksRef.current.onConnected?.(); - }); - - // Task created event - eventSource.addEventListener('task_created', (e) => { - const task = parseEventData(e.data); - if (task) callbacksRef.current.onTaskCreated?.(task); - }); - - // Task started event - eventSource.addEventListener('task_started', (e) => { - const task = parseEventData(e.data); - if (task) callbacksRef.current.onTaskStarted?.(task); - }); - - eventSource.addEventListener('task_progress', (e) => { - const task = parseEventData(e.data); - if (task) callbacksRef.current.onTaskProgress?.(task); - }); - - // Task completed event - eventSource.addEventListener('task_completed', (e) => { - const task = parseEventData(e.data); - if (task) callbacksRef.current.onTaskCompleted?.(task); - }); - - // Task failed event - eventSource.addEventListener('task_failed', (e) => { - const task = parseEventData(e.data); - if (task) callbacksRef.current.onTaskFailed?.(task); - }); - - // Heartbeat event used to keep the connection alive. - eventSource.addEventListener('heartbeat', () => { - // Optional place to record the latest heartbeat timestamp. - }); - - // Connection error handling - eventSource.onerror = (error) => { - setIsConnected(false); - callbacksRef.current.onError?.(error); - - // Auto-reconnect via ref to avoid stale closure issues. - if (autoReconnect && enabled) { - eventSource.close(); - reconnectTimeoutRef.current = setTimeout(() => { - connectRef.current(); - }, reconnectDelay); - } - }; - }, [ - autoReconnect, - reconnectDelay, - enabled, - parseEventData, - ]); - - useEffect(() => { - connectRef.current = connect; - }, [connect]); - // Disconnect and defer the state update to avoid nested renders. const disconnect = useCallback(() => { - if (reconnectTimeoutRef.current) { - clearTimeout(reconnectTimeoutRef.current); - reconnectTimeoutRef.current = null; + if (connectTimerRef.current) { + window.clearTimeout(connectTimerRef.current); + connectTimerRef.current = null; } - if (eventSourceRef.current) { - eventSourceRef.current.close(); - eventSourceRef.current = null; + if (subscriberIdRef.current !== null) { + subscribers.delete(subscriberIdRef.current); + subscriberIdRef.current = null; + } + if (subscribers.size === 0) { + closeSharedConnection(); } queueMicrotask(() => setIsConnected(false)); }, []); // Reconnect const reconnect = useCallback(() => { - disconnect(); - connect(); - }, [disconnect, connect]); + if (subscriberIdRef.current === null) { + const subscriberId = nextSubscriberId++; + subscriberIdRef.current = subscriberId; + subscribers.set(subscriberId, { + callbacksRef, + setIsConnected, + autoReconnect, + reconnectDelay, + }); + } + reconnectSharedStream(); + }, [autoReconnect, reconnectDelay]); // Connect or disconnect when the hook is enabled or disabled. useEffect(() => { if (enabled) { - connect(); - } else { - disconnect(); + const subscriberId = nextSubscriberId++; + subscriberIdRef.current = subscriberId; + subscribers.set(subscriberId, { + callbacksRef, + setIsConnected, + autoReconnect, + reconnectDelay, + }); + setIsConnected(sharedConnected); + connectTimerRef.current = window.setTimeout(() => { + connectTimerRef.current = null; + connectSharedStream(); + }, 0); + return () => { + disconnect(); + }; } + disconnect(); return () => { disconnect(); }; - }, [enabled, connect, disconnect]); + }, [autoReconnect, disconnect, enabled, reconnectDelay]); return { isConnected, diff --git a/apps/dsa-web/src/i18n/uiText.ts b/apps/dsa-web/src/i18n/uiText.ts index 8384f4aac..0c347df14 100644 --- a/apps/dsa-web/src/i18n/uiText.ts +++ b/apps/dsa-web/src/i18n/uiText.ts @@ -210,6 +210,9 @@ const zh = { 'taskPanel.processing': '分析中', 'taskPanel.processingAria': '任务进行中', 'taskPanel.processingTasks': '{count} 进行中', + 'taskPanel.cancelRequested': '请求取消', + 'taskPanel.cancelRequestedAria': '任务请求取消', + 'taskPanel.cancelled': '已取消', 'taskPanel.pendingAria': '任务等待中', 'taskPanel.openRunFlow': '查看运行流', 'taskPanel.openRunFlowAria': '查看 {stock} 运行流', @@ -274,6 +277,7 @@ const zh = { 'runFlow.graph.title': '运行拓扑', 'runFlow.graph.description': '自动分层展示入口、数据来源、分析引擎和产物链路。', 'runFlow.graph.nodeAria': '{label} 节点,状态 {status}', + 'runFlow.graph.startedAt': '开始', 'runFlow.events.title': '事件流', 'runFlow.events.count': '{count} 条事件', 'runFlow.events.filters': '事件筛选', @@ -688,6 +692,9 @@ const en: Record = { 'taskPanel.processing': 'Processing', 'taskPanel.processingAria': 'Task processing', 'taskPanel.processingTasks': '{count} processing', + 'taskPanel.cancelRequested': 'Cancel requested', + 'taskPanel.cancelRequestedAria': 'Task cancel requested', + 'taskPanel.cancelled': 'Cancelled', 'taskPanel.pendingAria': 'Task pending', 'taskPanel.openRunFlow': 'View run flow', 'taskPanel.openRunFlowAria': 'View {stock} run flow', @@ -752,6 +759,7 @@ const en: Record = { 'runFlow.graph.title': 'Run topology', 'runFlow.graph.description': 'Auto-layered lanes show entry, data sources, analysis engines, and artifact paths.', 'runFlow.graph.nodeAria': '{label} node, status {status}', + 'runFlow.graph.startedAt': 'Start', 'runFlow.events.title': 'Event stream', 'runFlow.events.count': '{count} events', 'runFlow.events.filters': 'Event filters', diff --git a/apps/dsa-web/src/stores/__tests__/stockPoolStore.test.ts b/apps/dsa-web/src/stores/__tests__/stockPoolStore.test.ts index c9a1a4e52..8cf7bc55f 100644 --- a/apps/dsa-web/src/stores/__tests__/stockPoolStore.test.ts +++ b/apps/dsa-web/src/stores/__tests__/stockPoolStore.test.ts @@ -711,7 +711,7 @@ describe('stockPoolStore', () => { await useStockPoolStore.getState().refreshActiveTasks(); expect(analysisApi.getTasks).toHaveBeenCalledWith({ - status: 'pending,processing', + status: 'pending,processing,cancel_requested', limit: 100, }); expect(useStockPoolStore.getState().activeTasks).toHaveLength(0); @@ -813,6 +813,24 @@ describe('stockPoolStore', () => { expect(useStockPoolStore.getState().activeTasks).toEqual([localTask, remoteTask]); }); + it('prunes stale local tasks when a complete backend snapshot contains cancel-requested tasks', async () => { + const staleTask = createTask({ taskId: 'task-stale', status: 'processing' }); + const cancelRequestedTask = createTask({ + taskId: 'task-cancel-requested', + status: 'cancel_requested', + progress: 60, + message: '正在取消任务', + }); + useStockPoolStore.getState().syncTaskCreated(staleTask); + vi.mocked(analysisApi.getTasks).mockResolvedValue( + createTaskListResponse([cancelRequestedTask]), + ); + + await useStockPoolStore.getState().refreshActiveTasks(); + + expect(useStockPoolStore.getState().activeTasks).toEqual([cancelRequestedTask]); + }); + it('keeps active tasks unchanged when backend reconciliation fails', async () => { const activeTask = createTask(); useStockPoolStore.getState().syncTaskCreated(activeTask); diff --git a/apps/dsa-web/src/stores/stockPoolStore.ts b/apps/dsa-web/src/stores/stockPoolStore.ts index 990454a7e..b961dba6e 100644 --- a/apps/dsa-web/src/stores/stockPoolStore.ts +++ b/apps/dsa-web/src/stores/stockPoolStore.ts @@ -856,7 +856,7 @@ export const useStockPoolStore = create((set, get) => ({ const localRevisionAtRequest = activeTaskLocalRevision; try { const response = await analysisApi.getTasks({ - status: 'pending,processing', + status: 'pending,processing,cancel_requested', limit: 100, }); if (requestId !== activeTaskRequestSeq) { @@ -868,7 +868,10 @@ export const useStockPoolStore = create((set, get) => ({ ); const remoteTaskIds = new Set(remoteTasks.map((task) => task.taskId)); const remoteTaskById = new Map(remoteTasks.map((task) => [task.taskId, task])); - const isCompleteSnapshot = response.tasks.length === response.pending + response.processing; + const activeTaskCount = response.pending + + response.processing + + response.tasks.filter((task) => task.status === 'cancel_requested').length; + const isCompleteSnapshot = response.tasks.length === activeTaskCount; const canPruneLocalTasks = isCompleteSnapshot && activeTaskLocalRevision === localRevisionAtRequest; const currentTasks = get().activeTasks; diff --git a/apps/dsa-web/src/types/analysis.ts b/apps/dsa-web/src/types/analysis.ts index 970db2c73..9a8c332f2 100644 --- a/apps/dsa-web/src/types/analysis.ts +++ b/apps/dsa-web/src/types/analysis.ts @@ -344,7 +344,7 @@ export type AnalyzeResponse = AnalysisResult | AnalyzeAsyncResponse; export interface TaskStatus { taskId: string; traceId?: string; - status: 'pending' | 'processing' | 'completed' | 'failed'; + status: 'pending' | 'processing' | 'completed' | 'failed' | 'cancel_requested' | 'cancelled'; progress?: number; result?: AnalysisResult; marketReviewReport?: string; @@ -363,7 +363,7 @@ export interface TaskInfo { traceId?: string; stockCode: string; stockName?: string; - status: 'pending' | 'processing' | 'completed' | 'failed'; + status: 'pending' | 'processing' | 'completed' | 'failed' | 'cancel_requested' | 'cancelled'; progress: number; message?: string; reportType: string; diff --git a/docs/CHANGELOG.md b/docs/CHANGELOG.md index 19df5ff13..8fb8a9b95 100644 --- a/docs/CHANGELOG.md +++ b/docs/CHANGELOG.md @@ -20,6 +20,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/). - [新功能] Web 端为活跃任务、历史报告和大盘复盘报告补充运行流视图入口,支持查看运行摘要、拓扑节点、事件流和基础排障详情。 - [修复] 修复历史报告运行流快照在混合时区事件时间戳下返回 500 的问题。 - [改进] #1459 持仓管理页新增持仓账户删除入口,复用现有账户软删除接口,误建账户会从默认列表、快照、风险、录入入口和事件列表隐藏且不物理清理历史流水。 +- [修复] 修复运行流 live SSE 事件未复用快照层递归脱敏规则的问题,避免本地路径、prompt/raw response、代理头等敏感诊断字段在 refetch 前短暂暴露。 - [修复] 桌面发布打包改用冻结可执行文件运行时探针校验 `alphasift.dsa_adapter`,避免 macOS PyInstaller 将模块内嵌进可执行文件时被文件系统/zip 扫描误判为缺失。 diff --git a/docs/full-guide.md b/docs/full-guide.md index a9d62d3cd..f41e07b1c 100644 --- a/docs/full-guide.md +++ b/docs/full-guide.md @@ -1418,12 +1418,14 @@ FastAPI 提供 RESTful API 服务,支持配置管理和触发分析。 | `/api/v1/analysis/analyze` | POST | 触发股票分析 | | `/api/v1/analysis/market-review` | POST | 后台触发大盘复盘;请求体可传 `{"send_notification": true}`;与 `main.py --market-review` 与 `bot` 复用同一套 `GeminiAnalyzer/SearchService/NotificationService` 组装语义 | | `/api/v1/analysis/tasks` | GET | 查询任务列表 | -| `/api/v1/analysis/tasks/stream` | GET (SSE) | 订阅任务实时状态流 | +| `/api/v1/analysis/tasks/stream` | GET (SSE) | 订阅任务实时状态流;`task_progress` 可选携带 `flow_event` 增量运行流事件 | +| `/api/v1/analysis/tasks/{task_id}/flow` | GET | 查询 active task 的运行流快照 | | `/api/v1/analysis/status/{task_id}` | GET | 查询任务状态 | | `/api/v1/alphasift/screen/tasks` | POST | 后台提交 AlphaSift 选股任务(需先开启 `ALPHASIFT_ENABLED`) | | `/api/v1/alphasift/screen/tasks/{task_id}` | GET | 查询 AlphaSift 选股任务状态与完成结果 | | `/api/v1/history` | GET | 查询分析历史 | | `/api/v1/history/{record_id}/diagnostics` | GET | 查询历史报告运行诊断摘要与脱敏复制文本 | +| `/api/v1/history/{record_id}/flow` | GET | 查询历史报告运行流快照,普通个股和 `MARKET/market_review` 大盘复盘复用同一契约 | | `/api/v1/decision-signals` | POST | 显式创建或按同源键去重决策信号,返回 `{ item, created }` | | `/api/v1/decision-signals` | GET | 分页查询决策信号,支持股票、市场、动作、阶段、来源、状态、时间范围和 cache-only 持仓过滤 | | `/api/v1/decision-signals/{signal_id}` | GET | 查询单条决策信号,读取前执行懒过期 | @@ -1451,6 +1453,7 @@ FastAPI 提供 RESTful API 服务,支持配置管理和触发分析。 > 说明:`POST /api/v1/analysis/market-review` 触发后,报告会以 `report_type=market_review` 写入历史库;你可直接查询 `/api/v1/history` 或 `/api/v1/history/{record_id}` 获取历史 Markdown,避免再次触发分析重算。 > 说明:历史列表新增 `report_type` 查询参数;通过 `stock_code=MARKET&report_type=market_review` 可单独读取大盘复盘历史集合,与普通个股历史逻辑完全隔离。 > 说明:`POST /api/v1/analysis/market-review` 的返回与历史持久化都会包含 `market_review_payload`:`market_scope`、`sections`、`sectors`、`news`、`market_light`、`indices` 等结构化字段。Web 端 Markdown 渲染与历史详情会复用该结构化字段;若结构化字段为空则回退到原始 Markdown。 +> 说明:运行流快照接口返回 `lanes/nodes/edges/events/summary` 统一契约。active task 缺少 diagnostics 时返回 skeleton flow;若任务 SSE 已收到真实 `flow_event`,快照会包含最近增量事件。completed history 优先使用 `context_snapshot.diagnostics` 与 `analysis_context_pack_overview` 构建完整拓扑。`cancel_requested/cancelled` 是合法状态,不会映射为 failed。 > 说明:`market_review_payload` 中的 `breadth` 仅在行情宽度数据真实可用时下发;当美股/港股或接口暂不可用时不下发该字段。前端显示层需按“字段缺失”降级为“暂无数据”而不是展示 0。 > 说明:该端点若返回 `task_id`,WebUI 会轮询 `GET /api/v1/analysis/status/{task_id}` 展示状态。状态为 `completed` 时给出完成提示(报告已生成并按配置推送),状态为 `failed` 时在前端错误区域显示 `error` 原因。 > 说明:`GET /api/v1/history/{record_id}/diagnostics` 支持历史记录主键 ID 或 `query_id`,返回 `normal/degraded/failed/unknown` 摘要、关键链路组件和可复制的脱敏 `copy_text`;旧报告缺少诊断快照时返回 `unknown`,不影响报告读取。 diff --git a/docs/run-diagnostics-p3.md b/docs/run-diagnostics-p3.md index 3e1617de0..b1f1cf169 100644 --- a/docs/run-diagnostics-p3.md +++ b/docs/run-diagnostics-p3.md @@ -15,6 +15,65 @@ GET /api/v1/history/{record_id}/diagnostics - 同步分析响应若已经带有 `diagnostic_summary`,前端可直接展示,不额外请求历史接口。 - 诊断面板支持复制后端生成的脱敏 `copy_text`,用于 issue 或部署排障。 - 分析链路在保存历史后会补齐任务/Provider/LLM/通知诊断到 `context_snapshot.diagnostics`,历史诊断接口统一聚合为用户可读摘要。 +- 首页运行流面板复用同一 RunFlowSnapshot 契约展示 active task、completed report 与大盘复盘;active task 通过任务 SSE 的可选增量事件实时追加事件流,完成或断线后再 refetch 快照保证最终一致。 + +## 运行流实时增量 + +运行流增量不新增独立 SSE endpoint,继续复用: + +```http +GET /api/v1/analysis/tasks/stream +``` + +兼容契约: + +- 事件类型仍为 `task_progress`。 +- 原有 task payload 字段保持不变。 +- 当本次进度更新来自运行诊断时,可选追加 `flow_event` 字段;旧客户端忽略该字段即可。 +- `flow_event` 使用与 `RunFlowSnapshot.events[]` 相同的脱敏事件结构:`id`、`timestamp`、`severity`、`type`、`node_id`、`title`、`message`、`metadata`。 +- 后端 TaskQueue 只为每个 active task 保留最近 N 条运行流事件,避免内存无限增长;完整历史仍以 `context_snapshot.diagnostics` 和历史 RunFlowSnapshot 为准。 + +示例: + +```json +{ + "task_id": "3f87...", + "trace_id": "3f87...", + "stock_code": "600519", + "status": "processing", + "progress": 64, + "message": "LLM 正在生成分析结果", + "flow_event": { + "id": "flow_0002", + "timestamp": "2026-06-08T22:30:24", + "severity": "success", + "type": "llm_run", + "node_id": "llm_analysis_1", + "title": "LLM 成功", + "message": "LLM deepseek-chat 成功" + } +} +``` + +运行诊断记录函数会在 provider、LLM、历史保存、通知记录成功写入内存诊断后 fail-open 触发 event sink。sink 失败只记录 warning,不改变分析、保存或通知的成功/失败判定。 + +新闻情报搜索也纳入同一 provider 诊断语义:`SearchService.search_stock_news()` 会以 `data_type=news_search` 记录 Tavily、SearXNG、Bocha、Brave 等搜索 provider 的尝试、过滤后结果数、缓存命中和失败原因。多个搜索 provider 连续尝试时,运行流拓扑会将它们展示为“新闻舆情”节点,并通过 fallback / retry 边表达降级过程。 + +运行流拓扑的数据来源泳道优先按节点开始时间排序;provider / LLM 节点若只有完成时间和耗时,会以 `ended_at - duration_ms` 推导 `started_at`,并在卡片上展示开始时间。无可用时间的节点保留原展示顺序作为兜底。 + +## 运行流 API + +```http +GET /api/v1/analysis/tasks/{task_id}/flow +GET /api/v1/history/{record_id}/flow +``` + +- 两个接口返回同一 `RunFlowSnapshot` 契约。 +- active task 缺少 diagnostics 时返回 skeleton flow,不伪造 provider / LLM 事件。 +- active task 若已有 recent `flow_event`,snapshot 会返回这些真实事件,并可根据事件中的节点元数据补出临时节点。 +- completed history 优先从 `context_snapshot.diagnostics` 与 `analysis_context_pack_overview` 构建完整拓扑。 +- 大盘复盘历史记录使用 `code=MARKET`、`report_type=market_review`,同样走 `/history/{record_id}/flow` 与 Web 运行流面板,不提供单独 UI 分叉。 +- `cancel_requested` 与 `cancelled` 是合法运行流状态;用户取消不应映射为 `failed`。 ## 运行流视图 diff --git a/src/core/market_review.py b/src/core/market_review.py index 4dacc721d..96618c28e 100644 --- a/src/core/market_review.py +++ b/src/core/market_review.py @@ -22,6 +22,11 @@ from src.market_analyzer import MarketAnalyzer from src.report_language import normalize_report_language from src.search_service import SearchService from src.analyzer import AnalysisResult, GeminiAnalyzer +from src.services.run_diagnostics import ( + current_diagnostic_snapshot, + record_history_run, + record_notification_run, +) logger = logging.getLogger(__name__) @@ -45,6 +50,46 @@ class MarketReviewRunResult: market_review_payload: Dict[str, Any] = field(default_factory=dict) +def _refresh_market_review_history_diagnostics(*, query_id: str) -> None: + """Refresh persisted market-review diagnostics after late flow events are recorded.""" + diagnostic_snapshot = current_diagnostic_snapshot() + if diagnostic_snapshot is None: + return + + try: + from src.storage import DatabaseManager + + db = DatabaseManager.get_instance() + updater = getattr(db, "update_analysis_history_diagnostics", None) + if callable(updater): + updater( + query_id=query_id, + code=MARKET_REVIEW_HISTORY_CODE, + diagnostics=diagnostic_snapshot, + ) + except Exception as exc: + logger.warning("回写大盘复盘运行诊断失败(fail-open): %s", exc) + + +def _record_market_review_notification_run( + *, + query_id: str, + channel: str, + status: str, + success: bool, + attempts: int = 1, + error_message: Optional[Any] = None, +) -> None: + record_notification_run( + channel=channel, + status=status, + success=success, + attempts=attempts, + error_message=error_message, + ) + _refresh_market_review_history_diagnostics(query_id=query_id) + + def _get_market_review_text(language: str) -> dict[str, str]: normalized = normalize_report_language(language) if normalized == "en": @@ -114,6 +159,7 @@ def run_market_review( 复盘报告文本 """ runtime_config = config or get_config() + history_query_id = query_id or f"market_review_{uuid.uuid4().hex}" review_text = _get_market_review_text(getattr(runtime_config, "report_language", "zh")) raw_region = ( override_region @@ -125,7 +171,7 @@ def run_market_review( logger.info( "[MarketReview] component=market_review action=start trigger_source=%s query_id=%s region=%s", trigger_source, - query_id or "-", + history_query_id, persist_region, ) @@ -142,7 +188,7 @@ def run_market_review( "[MarketReview] component=market_review action=build_report " "trigger_source=%s query_id=%s region=%s label=%s", trigger_source, - query_id or "-", + history_query_id, mkt, label, ) @@ -176,7 +222,7 @@ def run_market_review( "[MarketReview] component=market_review action=build_report " "trigger_source=%s query_id=%s region=%s label=%s", trigger_source, - query_id or "-", + history_query_id, run_region, label, ) @@ -220,7 +266,7 @@ def run_market_review( "[MarketReview] component=market_review action=save_report " "trigger_source=%s query_id=%s region=%s path=%s", trigger_source, - query_id or "-", + history_query_id, persist_region, filepath, ) @@ -230,7 +276,7 @@ def run_market_review( markdown_report=markdown_report, region=persist_region, config=runtime_config, - query_id=query_id, + query_id=history_query_id, market_light_snapshots=market_light_snapshots, market_review_payload=market_review_payload, ) @@ -241,9 +287,16 @@ def run_market_review( "[MarketReview] component=market_review action=skip_standalone_notification " "trigger_source=%s query_id=%s region=%s", trigger_source, - query_id or "-", + history_query_id, persist_region, ) + _record_market_review_notification_run( + query_id=history_query_id, + channel="report", + status="skipped", + success=False, + attempts=0, + ) elif send_notification and notifier.is_available(): # 添加标题 report_content = _render_market_review_payload_markdown( @@ -252,12 +305,18 @@ def run_market_review( ) success = notifier.send(report_content, email_send_to_all=True, route_type="report") + _record_market_review_notification_run( + query_id=history_query_id, + channel="report", + status="success" if success else "failed", + success=success, + ) if success: logger.info( "[MarketReview] component=market_review action=send_notification " "status=success trigger_source=%s query_id=%s region=%s", trigger_source, - query_id or "-", + history_query_id, persist_region, ) else: @@ -265,7 +324,7 @@ def run_market_review( "[MarketReview] component=market_review action=send_notification " "status=failed trigger_source=%s query_id=%s region=%s", trigger_source, - query_id or "-", + history_query_id, persist_region, ) elif not send_notification: @@ -273,9 +332,31 @@ def run_market_review( "[MarketReview] component=market_review action=skip_notification " "reason=no_notify trigger_source=%s query_id=%s region=%s", trigger_source, - query_id or "-", + history_query_id, persist_region, ) + _record_market_review_notification_run( + query_id=history_query_id, + channel="report", + status="skipped", + success=False, + attempts=0, + ) + else: + logger.info( + "[MarketReview] component=market_review action=skip_notification " + "reason=not_configured trigger_source=%s query_id=%s region=%s", + trigger_source, + history_query_id, + persist_region, + ) + _record_market_review_notification_run( + query_id=history_query_id, + channel="report", + status="not_configured", + success=False, + attempts=0, + ) if return_structured: return MarketReviewRunResult( @@ -289,7 +370,7 @@ def run_market_review( "[MarketReview] component=market_review action=failed " "trigger_source=%s query_id=%s region=%s", trigger_source, - query_id or "-", + history_query_id, persist_region, ) @@ -456,8 +537,12 @@ def _persist_market_review_history( context_snapshot["market_light_snapshots"] = market_light_snapshots if market_review_payload: context_snapshot["market_review_payload"] = market_review_payload + diagnostic_snapshot = current_diagnostic_snapshot() + if diagnostic_snapshot is not None: + context_snapshot["diagnostics"] = diagnostic_snapshot - saved = DatabaseManager.get_instance().save_analysis_history( + db = DatabaseManager.get_instance() + saved = db.save_analysis_history( result=result, query_id=history_query_id, report_type=MARKET_REVIEW_REPORT_TYPE, @@ -465,12 +550,28 @@ def _persist_market_review_history( context_snapshot=context_snapshot, save_snapshot=True, ) + saved_history_id = ( + saved + if isinstance(saved, int) and not isinstance(saved, bool) and saved > 0 + else None + ) + record_history_run( + report_saved=bool(saved), + metadata_saved=bool(saved), + analysis_history_id=saved_history_id, + ) + _refresh_market_review_history_diagnostics(query_id=history_query_id) if saved: logger.info("大盘复盘历史记录已保存: query_id=%s", history_query_id) else: logger.warning("大盘复盘历史记录保存失败: query_id=%s", history_query_id) return saved except Exception as exc: + record_history_run( + report_saved=False, + metadata_saved=False, + error_message=exc, + ) logger.warning("大盘复盘历史记录保存异常,报告文件与推送流程继续: %s", exc, exc_info=True) return 0 diff --git a/src/search_service.py b/src/search_service.py index a11632534..eef8c9b47 100644 --- a/src/search_service.py +++ b/src/search_service.py @@ -38,6 +38,7 @@ from src.config import ( normalize_news_strategy_profile, resolve_news_window_days, ) +from src.services.run_diagnostics import record_provider_run logger = logging.getLogger(__name__) @@ -3165,6 +3166,34 @@ class SearchService: search_time=response.search_time, ) + @staticmethod + def _elapsed_ms(started_at: float) -> int: + return max(0, int((time.monotonic() - started_at) * 1000)) + + @staticmethod + def _record_news_search_run( + *, + provider: str, + operation: str, + success: bool, + latency_ms: Optional[int] = None, + record_count: Optional[int] = None, + cache_hit: Optional[bool] = None, + error_type: Optional[str] = None, + error_message: Optional[Any] = None, + ) -> None: + record_provider_run( + data_type="news_search", + provider=provider, + operation=operation, + success=success, + latency_ms=latency_ms, + error_type=error_type, + error_message=error_message, + cache_hit=cache_hit, + record_count=record_count, + ) + def search_stock_news( self, stock_code: str, @@ -3235,16 +3264,43 @@ class SearchService: cached, cache_owner, cache_event = self._get_cached_or_reserve(cache_key) if cached is not None: logger.info(f"使用缓存搜索结果: {stock_name}({stock_code})") + self._record_news_search_run( + provider=cached.provider or "SearchCache", + operation="search_stock_news_cache", + success=bool(cached.success), + latency_ms=0, + record_count=len(cached.results or []), + cache_hit=True, + error_message=cached.error_message, + ) return cached if not cache_owner and cache_event is not None: cached = self._wait_for_cached(cache_key, cache_event) if cached is not None: logger.info(f"使用并发填充后的缓存搜索结果: {stock_name}({stock_code})") + self._record_news_search_run( + provider=cached.provider or "SearchCache", + operation="search_stock_news_cache_wait", + success=bool(cached.success), + latency_ms=0, + record_count=len(cached.results or []), + cache_hit=True, + error_message=cached.error_message, + ) return cached cached, cache_owner, cache_event = self._get_cached_or_reserve(cache_key) if cached is not None: logger.info(f"使用等待后命中的缓存搜索结果: {stock_name}({stock_code})") + self._record_news_search_run( + provider=cached.provider or "SearchCache", + operation="search_stock_news_cache_retry", + success=bool(cached.success), + latency_ms=0, + record_count=len(cached.results or []), + cache_hit=True, + error_message=cached.error_message, + ) return cached try: @@ -3267,7 +3323,19 @@ class SearchService: ) ) - response = provider.search(query, provider_max_results, days=search_days, **search_kwargs) + started_at = time.monotonic() + try: + response = provider.search(query, provider_max_results, days=search_days, **search_kwargs) + except Exception as exc: + self._record_news_search_run( + provider=provider.name, + operation="search_stock_news", + success=False, + latency_ms=self._elapsed_ms(started_at), + error_type=type(exc).__name__, + error_message=exc, + ) + raise filtered_response = self._filter_news_response( response, search_days=search_days, @@ -3275,6 +3343,16 @@ class SearchService: log_scope=f"{stock_code}:{provider.name}:stock_news", ) had_provider_success = had_provider_success or bool(response.success) + filtered_count = len(filtered_response.results or []) if filtered_response.success else 0 + self._record_news_search_run( + provider=provider.name, + operation="search_stock_news", + success=bool(filtered_response.success and filtered_response.results), + latency_ms=self._elapsed_ms(started_at), + record_count=filtered_count, + error_type=None if filtered_count else "NoUsableNews", + error_message=None if filtered_count else (response.error_message or "过滤后无有效新闻"), + ) if filtered_response.success and filtered_response.results: language_response, _preferred_count = self._prioritize_news_language( diff --git a/src/services/run_diagnostics.py b/src/services/run_diagnostics.py index ac5408b0e..af3f10b33 100644 --- a/src/services/run_diagnostics.py +++ b/src/services/run_diagnostics.py @@ -11,10 +11,11 @@ from __future__ import annotations import logging import re import uuid +from collections.abc import Mapping from contextvars import ContextVar, Token from dataclasses import dataclass, field -from datetime import datetime -from typing import Any, Dict, List, Optional +from datetime import datetime, timedelta +from typing import Any, Callable, Dict, List, Optional logger = logging.getLogger(__name__) @@ -63,6 +64,18 @@ _SECRET_REDACTIONS = ( "Bearer ", ), ) +_SENSITIVE_KEY_RE = re.compile( + r"(?i)(authorization|api[_-]?key|access[_-]?token|(?:^|[_-])(?:auth|refresh|session|bearer)?[_-]?token$|secret|password|passwd|cookie|" + r"webhook|sendkey|prompt|raw[_-]?prompt|raw[_-]?response|headers?|proxy)" +) +_WEBHOOK_URL_RE = re.compile(r"https?://[^\s]+?(?:webhook|token|key|secret|sendkey)[^\s]*", re.IGNORECASE) +_LOCAL_ABSOLUTE_PATH_RE = re.compile( + r"(? str: @@ -71,7 +84,7 @@ def build_trace_id() -> str: def sanitize_diagnostic_text(value: Any, *, max_length: int = 300) -> Optional[str]: - """Return a short diagnostic string with obvious credentials redacted.""" + """Return a short diagnostic string with sensitive details redacted.""" if value is None: return None @@ -81,12 +94,51 @@ def sanitize_diagnostic_text(value: Any, *, max_length: int = 300) -> Optional[s for pattern, replacement in _SECRET_REDACTIONS: text = pattern.sub(replacement, text) + text = _WEBHOOK_URL_RE.sub("", text) + text = _LOCAL_ABSOLUTE_PATH_RE.sub("", text) + text = _SENSITIVE_ASSIGNMENT_RE.sub(lambda match: f"{match.group(1)}=", text) if len(text) > max_length: return f"{text[:max_length].rstrip()}..." return text +def safe_diagnostic_key(value: Any) -> str: + """Normalize a diagnostic object key after applying text redaction.""" + text = sanitize_diagnostic_text(value, max_length=80) or "" + return re.sub(r"[^A-Za-z0-9_]+", "_", text.strip().lower()).strip("_")[:80] + + +def sanitize_diagnostic_metadata(value: Any, *, depth: int = 0) -> Any: + """Recursively redact diagnostic metadata before it reaches API/SSE payloads.""" + if depth > 3: + return "" + if isinstance(value, Mapping): + sanitized: Dict[str, Any] = {} + for index, (key, item) in enumerate(value.items()): + if index >= 20: + sanitized["truncated"] = True + break + safe_key = safe_diagnostic_key(key) + if not safe_key: + continue + if _SENSITIVE_KEY_RE.search(str(key)): + sanitized[safe_key] = "" + continue + safe_value = sanitize_diagnostic_metadata(item, depth=depth + 1) + if safe_value not in (None, "", [], {}): + sanitized[safe_key] = safe_value + return sanitized + if isinstance(value, list): + items = [sanitize_diagnostic_metadata(item, depth=depth + 1) for item in value[:8]] + return [item for item in items if item not in (None, "", [], {})] + if isinstance(value, tuple): + return sanitize_diagnostic_metadata(list(value), depth=depth) + if isinstance(value, (int, float, bool)): + return value + return sanitize_diagnostic_text(value, max_length=160) + + @dataclass class ProviderRun: """One provider attempt in a trace.""" @@ -274,18 +326,40 @@ class RunDiagnosticContext: llm_runs: List[LLMRun] = field(default_factory=list) notification_runs: List[NotificationRun] = field(default_factory=list) history_runs: List[HistoryRun] = field(default_factory=list) + event_sink: Optional[Callable[[Dict[str, Any]], None]] = None + flow_event_index: int = 0 + provider_attempt_index_by_type: Dict[str, int] = field(default_factory=dict) def record_provider_run(self, provider_run: ProviderRun) -> None: self.provider_runs.append(provider_run) + data_type_key = _safe_event_key(provider_run.data_type) or "provider" + attempt_index = self.provider_attempt_index_by_type.get(data_type_key, 0) + 1 + self.provider_attempt_index_by_type[data_type_key] = attempt_index + self._emit_flow_event(_provider_flow_event(self, provider_run, attempt_index)) def record_llm_run(self, llm_run: LLMRun) -> None: self.llm_runs.append(llm_run) + self._emit_flow_event(_llm_flow_event(self, llm_run, len(self.llm_runs))) def record_notification_run(self, notification_run: NotificationRun) -> None: self.notification_runs.append(notification_run) + self._emit_flow_event(_notification_flow_event(self, notification_run, len(self.notification_runs))) def record_history_run(self, history_run: HistoryRun) -> None: self.history_runs.append(history_run) + self._emit_flow_event(_history_flow_event(self, history_run, len(self.history_runs))) + + def _emit_flow_event(self, event: Dict[str, Any]) -> None: + if self.event_sink is None: + return + try: + self.flow_event_index += 1 + event_payload = sanitize_diagnostic_metadata(event) + event_payload = dict(event_payload) if isinstance(event_payload, Mapping) else {} + event_payload["id"] = event_payload.get("id") or f"flow_{self.flow_event_index:04d}" + self.event_sink(event_payload) + except Exception as exc: # pragma: no cover - defensive fail-open guard + logger.warning("run-flow event sink failed: %s", exc) def snapshot(self) -> Dict[str, Any]: return { @@ -312,6 +386,7 @@ def activate_run_diagnostic_context( query_id: Optional[str] = None, stock_code: Optional[str] = None, trigger_source: Optional[str] = None, + event_sink: Optional[Callable[[Dict[str, Any]], None]] = None, ) -> Token: """Activate a diagnostic context and return its reset token.""" context = RunDiagnosticContext( @@ -320,6 +395,7 @@ def activate_run_diagnostic_context( query_id=query_id, stock_code=stock_code, trigger_source=trigger_source, + event_sink=event_sink, ) return _CURRENT_CONTEXT.set(context) @@ -344,6 +420,236 @@ def current_diagnostic_snapshot() -> Optional[Dict[str, Any]]: return None +_DATA_TYPE_LABELS = { + "realtime_quote": "实时行情", + "daily_data": "日线K线", + "daily_bars": "日线K线", + "technical": "技术指标", + "news": "新闻舆情", + "news_search": "新闻舆情", + "fundamental": "基本面", + "fundamentals": "基本面", + "chip": "筹码结构", +} + + +def _safe_event_key(value: Any) -> str: + return safe_diagnostic_key(value) + + +def _clean_metadata(value: Dict[str, Any]) -> Dict[str, Any]: + return { + str(key): item + for key, item in value.items() + if item not in (None, "", [], {}) + } + + +def _flow_status_for_success(success: bool, *, fallback: bool = False, skipped: bool = False) -> str: + if skipped: + return "skipped" + if success: + return "fallback" if fallback else "success" + return "failed" + + +def _started_at_from_end_and_duration(end: Any, duration_ms: Optional[int]) -> Optional[str]: + if duration_ms is None or duration_ms < 0: + return None + if isinstance(end, datetime): + parsed = end + elif isinstance(end, str) and "T" in end: + normalized = end[:-1] + "+00:00" if end.endswith("Z") else end + try: + parsed = datetime.fromisoformat(normalized) + except ValueError: + return None + else: + return None + return (parsed - timedelta(milliseconds=duration_ms)).isoformat() + + +def _provider_flow_event( + context: RunDiagnosticContext, + run: ProviderRun, + index: int, +) -> Dict[str, Any]: + data_type = _safe_event_key(run.data_type) or "provider" + provider_key = _safe_event_key(run.provider) or "unknown" + label = _DATA_TYPE_LABELS.get(data_type, data_type) + fallback = bool(run.fallback_from or run.fallback_to) + status = _flow_status_for_success(run.success, fallback=fallback) + node_id = f"provider_{data_type}_{provider_key}_{index}" + started_at = _started_at_from_end_and_duration(run.created_at, run.latency_ms) + message = ( + f"{label} {run.provider} 成功" + if run.success + else f"{label} {run.provider} 失败:{run.error_message_sanitized or run.error_type or '未知错误'}" + ) + return { + "timestamp": run.created_at, + "severity": "success" if run.success else "warning", + "type": "provider_run", + "node_id": node_id, + "title": f"{label}{'成功' if run.success else '失败'}", + "message": sanitize_diagnostic_text(message, max_length=220), + "metadata": _clean_metadata( + { + "trace_id": context.trace_id, + "provider": run.provider, + "data_type": run.data_type, + "operation": run.operation, + "duration_ms": run.latency_ms, + "record_count": run.record_count, + "fallback_from": run.fallback_from, + "fallback_to": run.fallback_to, + "error_type": run.error_type, + "node": { + "id": node_id, + "lane": "data_source", + "kind": "data_source", + "label": f"{label} · {run.provider}", + "status": status, + "provider": run.provider, + "started_at": started_at, + "ended_at": run.created_at, + "duration_ms": run.latency_ms, + "record_count": run.record_count, + "message": message, + }, + } + ), + } + + +def _llm_flow_event( + context: RunDiagnosticContext, + run: LLMRun, + index: int, +) -> Dict[str, Any]: + call_type = _safe_event_key(run.call_type) or "analysis" + model = run.model or run.provider or "unknown" + status = _flow_status_for_success(run.success, fallback=bool(run.fallback_model or index > 1)) + node_id = f"llm_{call_type}_{index}" + started_at = _started_at_from_end_and_duration(run.created_at, run.duration_ms) + message = ( + f"LLM {model} 成功" + if run.success + else f"LLM {model} 失败:{run.error_message_sanitized or run.error_type or '未知错误'}" + ) + return { + "timestamp": run.created_at, + "severity": "success" if run.success else "danger", + "type": "llm_run", + "node_id": node_id, + "title": f"LLM {'成功' if run.success else '失败'}", + "message": sanitize_diagnostic_text(message, max_length=220), + "metadata": _clean_metadata( + { + "trace_id": context.trace_id, + "provider": run.provider, + "model": run.model, + "call_type": run.call_type, + "duration_ms": run.duration_ms, + "fallback_model": run.fallback_model, + "error_type": run.error_type, + "node": { + "id": node_id, + "lane": "analysis", + "kind": "model", + "label": "LLM 生成", + "status": status, + "provider": model, + "started_at": started_at, + "ended_at": run.created_at, + "duration_ms": run.duration_ms, + "message": message, + }, + } + ), + } + + +def _history_flow_event( + context: RunDiagnosticContext, + run: HistoryRun, + index: int, +) -> Dict[str, Any]: + node_id = "history_save" if index == 1 else f"history_save_{index}" + status = "success" if run.report_saved else "failed" + message = "报告历史已保存" if run.report_saved else f"报告历史保存失败:{run.error_message_sanitized or '未知错误'}" + return { + "timestamp": run.created_at, + "severity": "success" if run.report_saved else "danger", + "type": "history_run", + "node_id": node_id, + "title": "历史保存成功" if run.report_saved else "历史保存失败", + "message": sanitize_diagnostic_text(message, max_length=220), + "metadata": _clean_metadata( + { + "trace_id": context.trace_id, + "metadata_saved": run.metadata_saved, + "analysis_history_id": run.analysis_history_id, + "node": { + "id": node_id, + "lane": "artifact", + "kind": "artifact", + "label": "保存报告", + "status": status, + "message": message, + }, + } + ), + } + + +def _notification_flow_event( + context: RunDiagnosticContext, + run: NotificationRun, + index: int, +) -> Dict[str, Any]: + channel = run.channel or "unknown" + channel_key = _safe_event_key(channel) or "unknown" + skipped = run.status in {"skipped", "not_configured"} + status = _flow_status_for_success(run.success, skipped=skipped) + node_id = f"notification_{channel_key}_{index}" + if status == "success": + title = "通知发送成功" + message = f"{channel} 通知发送成功" + elif status == "skipped": + title = "通知跳过" + message = f"{channel} 通知跳过" + else: + title = "通知失败" + message = f"{channel} 通知失败:{run.error_message_sanitized or run.status or '未知错误'}" + return { + "timestamp": run.created_at, + "severity": "success" if status == "success" else ("warning" if status == "skipped" else "danger"), + "type": "notification_run", + "node_id": node_id, + "title": title, + "message": sanitize_diagnostic_text(message, max_length=220), + "metadata": _clean_metadata( + { + "trace_id": context.trace_id, + "channel": channel, + "status": run.status, + "attempts": run.attempts, + "node": { + "id": node_id, + "lane": "artifact", + "kind": "notification", + "label": f"推送通知 · {channel}", + "status": status, + "provider": channel, + "attempts": run.attempts, + "message": message, + }, + } + ), + } + + def record_provider_run( *, data_type: str, diff --git a/src/services/run_flow.py b/src/services/run_flow.py index 7b2b52a2e..6ac2e2057 100644 --- a/src/services/run_flow.py +++ b/src/services/run_flow.py @@ -4,15 +4,18 @@ from __future__ import annotations import json -import re from collections import defaultdict from collections.abc import Mapping -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from typing import Any, Dict, Iterable, List, Optional, Tuple from api.v1.schemas.run_flow import RunFlowSnapshot from src.analysis_context_pack_overview import extract_analysis_context_pack_overview -from src.services.run_diagnostics import sanitize_diagnostic_text +from src.services.run_diagnostics import ( + safe_diagnostic_key, + sanitize_diagnostic_metadata, + sanitize_diagnostic_text, +) from src.utils.data_processing import normalize_model_used, parse_json_field @@ -69,20 +72,6 @@ _CONTEXT_STATUS_TO_FLOW = { "fetch_failed": "failed", } -_SENSITIVE_KEY_RE = re.compile( - r"(?i)(authorization|api[_-]?key|access[_-]?token|(?:^|[_-])(?:auth|refresh|session|bearer)?[_-]?token$|secret|password|passwd|cookie|" - r"webhook|sendkey|prompt|raw[_-]?prompt|raw[_-]?response|headers?|proxy)" -) -_WEBHOOK_URL_RE = re.compile(r"https?://[^\s]+?(?:webhook|token|key|secret|sendkey)[^\s]*", re.IGNORECASE) -_LOCAL_ABSOLUTE_PATH_RE = re.compile( - r"(? 1: @@ -563,6 +565,7 @@ def _append_llm_runs( label="LLM 生成", status=status, provider=model or provider, + started_at=started_at, ended_at=timestamp, duration_ms=duration_ms, attempts=1, @@ -820,6 +823,117 @@ def _append_task_events(events: List[Dict[str, Any]], task: Any, flow_status: st ) +def _append_active_flow_events( + nodes: Dict[str, Dict[str, Any]], + edges: List[Dict[str, Any]], + events: List[Dict[str, Any]], + flow_events: List[Any], + *, + flow_status: str, +) -> None: + if not flow_events: + return + + known_node_ids = set(nodes) + last_provider_node_by_type: Dict[str, Tuple[str, Dict[str, Any]]] = {} + last_llm_node: Optional[str] = None + last_history_node: Optional[str] = None + + for raw_event in flow_events: + event = _as_mapping(raw_event) + if not event: + continue + metadata = _sanitize_metadata(event.get("metadata") or {}) + node_payload = metadata.get("node") if isinstance(metadata, Mapping) else None + node_id = _safe_key(event.get("node_id")) + + if isinstance(node_payload, Mapping): + raw_node_id = _safe_text(node_payload.get("id"), max_length=120) or node_id + if raw_node_id: + node_id = raw_node_id + _put_node( + nodes, + node_id, + lane=str(node_payload.get("lane") or "analysis"), + kind=str(node_payload.get("kind") or "analysis"), + label=str(node_payload.get("label") or node_id), + status=str(node_payload.get("status") or flow_status), + provider=node_payload.get("provider"), + started_at=node_payload.get("started_at") + or _started_at_from_end_and_duration( + node_payload.get("ended_at") or event.get("timestamp"), + node_payload.get("duration_ms"), + ), + ended_at=node_payload.get("ended_at") or event.get("timestamp"), + duration_ms=node_payload.get("duration_ms"), + attempts=node_payload.get("attempts"), + record_count=node_payload.get("record_count"), + message=node_payload.get("message") or event.get("message"), + metadata={key: value for key, value in metadata.items() if key != "node"}, + ) + + event_type = _safe_key(event.get("type")) or "event" + if node_id and node_id in nodes and node_id not in known_node_ids: + if event_type == "provider_run": + provider_data_type = _safe_key(metadata.get("data_type") or "provider") + provider_run = { + "provider": metadata.get("provider") or nodes[node_id].get("provider"), + "success": event.get("severity") == "success" or nodes[node_id].get("status") in {"success", "fallback"}, + "fallback_from": metadata.get("fallback_from"), + "fallback_to": metadata.get("fallback_to"), + } + previous_provider = last_provider_node_by_type.get(provider_data_type) + if previous_provider: + previous_provider_node, previous_provider_run = previous_provider + edge_kind = _provider_transition_kind(previous_provider_run, provider_run) + _append_edge( + edges, + previous_provider_node, + node_id, + edge_kind, + nodes[node_id].get("status", "unknown"), + label="降级" if edge_kind == "fallback" else ("重试" if edge_kind == "retry" else "调用"), + ) + else: + _append_edge(edges, "task_queue", node_id, "control", nodes[node_id].get("status", "unknown"), label="调用") + last_provider_node_by_type[provider_data_type] = (node_id, provider_run) + elif event_type == "llm_run": + anchor = "analysis_pipeline" if "analysis_pipeline" in nodes else "task_queue" + _append_edge(edges, anchor, node_id, "data", nodes[node_id].get("status", "unknown"), label="生成") + last_llm_node = node_id + elif event_type == "history_run": + anchor = last_llm_node or ("analysis_pipeline" if "analysis_pipeline" in nodes else "task_queue") + _append_edge(edges, anchor, node_id, "data", nodes[node_id].get("status", "unknown"), label="保存") + last_history_node = node_id + elif event_type == "notification_run": + anchor = last_history_node or last_llm_node or ("analysis_pipeline" if "analysis_pipeline" in nodes else "task_queue") + _append_edge(edges, anchor, node_id, "control", nodes[node_id].get("status", "unknown"), label="通知") + known_node_ids.add(node_id) + + _append_external_event(events, event) + + +def _append_external_event(events: List[Dict[str, Any]], event: Dict[str, Any]) -> None: + event_id = _safe_text(event.get("id"), max_length=96) or f"flow_{len(events) + 1:04d}" + if any(existing.get("id") == event_id for existing in events): + return + metadata = _sanitize_metadata(event.get("metadata") or {}) + if isinstance(metadata, Mapping) and "node" in metadata: + metadata = {key: value for key, value in metadata.items() if key != "node"} + events.append( + { + "id": event_id, + "timestamp": _datetime_to_iso(event.get("timestamp")), + "severity": event.get("severity") if event.get("severity") in {"info", "success", "warning", "danger"} else "info", + "type": _safe_key(event.get("type")) or "event", + "node_id": _safe_text(event.get("node_id"), max_length=120), + "title": _safe_text(event.get("title"), max_length=100) or "运行事件", + "message": _safe_text(event.get("message"), max_length=220), + "metadata": metadata, + } + ) + + def _group_provider_runs(provider_runs: List[Any]) -> Dict[str, List[Dict[str, Any]]]: grouped: Dict[str, List[Dict[str, Any]]] = defaultdict(list) for run in provider_runs: @@ -1136,23 +1250,11 @@ def _valid_status(value: Any) -> str: def _safe_text(value: Any, *, max_length: int = 300) -> Optional[str]: - if value is None: - return None - text = sanitize_diagnostic_text(value, max_length=max_length) - if not text: - return None - text = _WEBHOOK_URL_RE.sub("", text) - text = _LOCAL_ABSOLUTE_PATH_RE.sub("", text) - text = _SENSITIVE_ASSIGNMENT_RE.sub(lambda match: f"{match.group(1)}=", text) - if len(text) > max_length: - return f"{text[:max_length].rstrip()}..." - return text + return sanitize_diagnostic_text(value, max_length=max_length) def _safe_key(value: Any) -> str: - text = _safe_text(value, max_length=80) or "" - text = re.sub(r"[^A-Za-z0-9_]+", "_", text.strip().lower()) - return text.strip("_")[:80] + return safe_diagnostic_key(value) def _safe_int(value: Any) -> Optional[int]: @@ -1166,32 +1268,7 @@ def _safe_int(value: Any) -> Optional[int]: def _sanitize_metadata(value: Any, *, depth: int = 0) -> Any: - if depth > 3: - return "" - if isinstance(value, Mapping): - sanitized: Dict[str, Any] = {} - for index, (key, item) in enumerate(value.items()): - if index >= 20: - sanitized["truncated"] = True - break - safe_key = _safe_key(key) - if not safe_key: - continue - if _SENSITIVE_KEY_RE.search(str(key)): - sanitized[safe_key] = "" - continue - safe_value = _sanitize_metadata(item, depth=depth + 1) - if safe_value not in (None, "", [], {}): - sanitized[safe_key] = safe_value - return sanitized - if isinstance(value, list): - items = [_sanitize_metadata(item, depth=depth + 1) for item in value[:8]] - return [item for item in items if item not in (None, "", [], {})] - if isinstance(value, tuple): - return _sanitize_metadata(list(value), depth=depth) - if isinstance(value, (int, float, bool)): - return value - return _safe_text(value, max_length=160) + return sanitize_diagnostic_metadata(value, depth=depth) def _as_mapping(value: Any) -> Dict[str, Any]: @@ -1230,6 +1307,23 @@ def _elapsed_ms(start: Any, end: Any) -> Optional[int]: return int(seconds * 1000) +def _started_at_from_end_and_duration(end: Any, duration_ms: Any) -> Optional[str]: + duration = _safe_int(duration_ms) + if duration is None: + return None + if isinstance(end, datetime): + parsed = end + elif isinstance(end, str) and "T" in end: + normalized = end[:-1] + "+00:00" if end.endswith("Z") else end + try: + parsed = datetime.fromisoformat(normalized) + except ValueError: + return None + else: + return None + return (parsed - timedelta(milliseconds=duration)).isoformat() + + def _local_timezone(): return datetime.now().astimezone().tzinfo or timezone.utc diff --git a/src/services/task_queue.py b/src/services/task_queue.py index 83df64d61..b92865160 100644 --- a/src/services/task_queue.py +++ b/src/services/task_queue.py @@ -14,6 +14,7 @@ A股自选股智能分析系统 - 异步任务队列 from __future__ import annotations import asyncio +import copy import logging import threading import uuid @@ -53,6 +54,8 @@ class TaskStatus(str, Enum): PROCESSING = "processing" # In progress COMPLETED = "completed" # Completed FAILED = "failed" # Failed + CANCEL_REQUESTED = "cancel_requested" # Cancellation requested + CANCELLED = "cancelled" # Cancelled by user/system @dataclass @@ -82,6 +85,7 @@ class TaskInfo: skills: Optional[List[str]] = None report_language: Optional[str] = None trace_id: Optional[str] = None + flow_events: List[Dict[str, Any]] = field(default_factory=list) def to_dict(self) -> Dict[str, Any]: """Convert task info into an API-friendly dictionary.""" @@ -127,6 +131,7 @@ class TaskInfo: skills=list(self.skills) if self.skills is not None else None, report_language=self.report_language, trace_id=self.trace_id or self.task_id, + flow_events=copy.deepcopy(self.flow_events), ) @@ -190,6 +195,7 @@ class AnalysisTaskQueue: # 任务历史保留数量(内存中) self._max_history = 100 + self._max_flow_events_per_task = 200 self._initialized = True logger.info(f"[TaskQueue] 初始化完成,最大并发: {max_workers}") @@ -524,6 +530,44 @@ class AnalysisTaskQueue: with self._data_lock: task = self._tasks.get(task_id) return task.copy() if task else None + + def append_task_flow_event( + self, + task_id: str, + flow_event: Dict[str, Any], + ) -> Optional[Dict[str, Any]]: + """Append a recent run-flow event to an active task and broadcast it. + + The event cache is deliberately bounded and fail-open; diagnostics must + never affect the analysis pipeline. + """ + try: + event_payload = copy.deepcopy(flow_event) + except Exception: + logger.debug("[TaskQueue] 忽略不可复制的运行流事件: task_id=%s", task_id) + return None + + with self._data_lock: + task = self._tasks.get(task_id) + if not task: + return None + task.flow_events.append(event_payload) + if len(task.flow_events) > self._max_flow_events_per_task: + task.flow_events = task.flow_events[-self._max_flow_events_per_task:] + task_snapshot = task.copy() + + payload = task_snapshot.to_dict() + payload["flow_event"] = event_payload + self._broadcast_event("task_progress", payload) + return event_payload + + def get_task_flow_events(self, task_id: str) -> List[Dict[str, Any]]: + """Return a copy of the recent run-flow events for a task.""" + with self._data_lock: + task = self._tasks.get(task_id) + if not task: + return [] + return copy.deepcopy(task.flow_events) def list_pending_tasks(self) -> List[TaskInfo]: """ @@ -535,7 +579,7 @@ class AnalysisTaskQueue: with self._data_lock: return [ task.copy() for task in self._tasks.values() - if task.status in (TaskStatus.PENDING, TaskStatus.PROCESSING) + if task.status in (TaskStatus.PENDING, TaskStatus.PROCESSING, TaskStatus.CANCEL_REQUESTED) ] def list_all_tasks(self, limit: int = 50) -> List[TaskInfo]: @@ -669,6 +713,7 @@ class AnalysisTaskQueue: query_id=task_id, stock_code=stock_code, trigger_source=query_source, + event_sink=lambda event: self.append_task_flow_event(task_id, event), ) result = service.analyze_stock( stock_code=stock_code, @@ -777,6 +822,7 @@ class AnalysisTaskQueue: query_id=task_id, stock_code=task.stock_code, trigger_source="api", + event_sink=lambda event: self.append_task_flow_event(task_id, event), ) try: result = run_task() @@ -836,7 +882,7 @@ class AnalysisTaskQueue: # 按时间排序,删除旧的已完成任务 completed_tasks = sorted( [t for t in self._tasks.values() - if t.status in (TaskStatus.COMPLETED, TaskStatus.FAILED)], + if t.status in (TaskStatus.COMPLETED, TaskStatus.FAILED, TaskStatus.CANCELLED)], key=lambda t: t.created_at ) diff --git a/tests/test_analysis_api_contract.py b/tests/test_analysis_api_contract.py index d57779377..28d1d5f55 100644 --- a/tests/test_analysis_api_contract.py +++ b/tests/test_analysis_api_contract.py @@ -502,6 +502,38 @@ class AnalysisApiContractTestCase(unittest.TestCase): self.assertEqual(status.market_review_payload["kind"], "market_review") self.assertIsNone(status.result) + def test_get_analysis_status_accepts_cancel_states_from_queue(self) -> None: + if get_analysis_status is None or analysis_endpoint_module is None: + self.skipTest("analysis endpoint helpers unavailable in this environment") + + for task_status in ( + analysis_endpoint_module.TaskStatusEnum.CANCEL_REQUESTED, + analysis_endpoint_module.TaskStatusEnum.CANCELLED, + ): + with self.subTest(task_status=task_status.value): + queue = MagicMock() + queue.get_task.return_value = SimpleNamespace( + task_id=f"task-{task_status.value}", + trace_id=f"trace-{task_status.value}", + stock_code="600519", + stock_name="贵州茅台", + status=task_status, + progress=42, + result=None, + error=None, + original_query=None, + selection_source=None, + analysis_phase="auto", + skills=[], + ) + + with patch("api.v1.endpoints.analysis.get_task_queue", return_value=queue): + status = get_analysis_status(f"task-{task_status.value}") + + self.assertEqual(status.status, task_status.value) + self.assertEqual(status.progress, 42) + self.assertIsNone(status.result) + def test_get_analysis_status_normalizes_completed_queue_result_contract(self) -> None: if get_analysis_status is None or analysis_endpoint_module is None: self.skipTest("analysis endpoint helpers unavailable in this environment") diff --git a/tests/test_market_review.py b/tests/test_market_review.py index 9105f1335..6b4f748e9 100644 --- a/tests/test_market_review.py +++ b/tests/test_market_review.py @@ -2,6 +2,7 @@ """Tests for localized market review wrappers.""" import importlib +import json import os import sys import tempfile @@ -40,6 +41,7 @@ def _build_optional_module_stubs() -> dict[str, ModuleType]: sys.modules.update(_build_optional_module_stubs()) import src.core.market_review as market_review_module from src.config import Config +from src.services.run_diagnostics import activate_run_diagnostic_context, reset_run_diagnostic_context from src.storage import AnalysisHistory, DatabaseManager run_market_review = market_review_module.run_market_review @@ -100,7 +102,7 @@ class MarketReviewLocalizationTestCase(unittest.TestCase): self.assertTrue(notifier.send.call_args.kwargs["email_send_to_all"]) self.assertEqual(notifier.send.call_args.kwargs["route_type"], "report") persist_history.assert_called_once() - self.assertEqual(persist_history.call_args.kwargs["query_id"], None) + self.assertTrue(persist_history.call_args.kwargs["query_id"].startswith("market_review_")) def test_run_market_review_passes_request_config_to_generation(self) -> None: notifier = self._make_notifier() @@ -422,6 +424,103 @@ class MarketReviewLocalizationTestCase(unittest.TestCase): else: os.environ["DATABASE_PATH"] = old_db_path + def test_run_market_review_persists_notification_diagnostics_after_history_save(self) -> None: + with tempfile.TemporaryDirectory() as temp_dir: + old_db_path = os.environ.get("DATABASE_PATH") + os.environ["DATABASE_PATH"] = os.path.join(temp_dir, "market_review_notification.db") + Config._instance = None + DatabaseManager.reset_instance() + query_id = "market-task-notification" + notifier = self._make_notifier() + market_analyzer = MagicMock() + market_analyzer.run_daily_review_with_snapshot.return_value = SimpleNamespace( + report="## 今日大盘\n\n复盘正文", + market_light_snapshot={"region": "cn", "trade_date": "2026-03-06", "score": 60}, + ) + token = activate_run_diagnostic_context( + trace_id="trace-market-notification", + task_id=query_id, + query_id=query_id, + stock_code=market_review_module.MARKET_REVIEW_HISTORY_CODE, + trigger_source="api", + ) + try: + with patch.object(market_review_module, "MarketAnalyzer", return_value=market_analyzer): + result = run_market_review( + notifier, + config=SimpleNamespace(report_language="zh", market_review_region="cn"), + send_notification=True, + query_id=query_id, + trigger_source="api", + ) + + self.assertEqual(result, "## 今日大盘\n\n复盘正文") + db = DatabaseManager.get_instance() + with db.get_session() as session: + row = session.query(AnalysisHistory).filter( + AnalysisHistory.query_id == query_id + ).first() + self.assertIsNotNone(row) + context_snapshot = json.loads(row.context_snapshot) + notification_runs = context_snapshot["diagnostics"]["notification_runs"] + self.assertEqual(notification_runs[-1]["status"], "success") + self.assertTrue(notification_runs[-1]["success"]) + finally: + reset_run_diagnostic_context(token) + DatabaseManager.reset_instance() + Config._instance = None + if old_db_path is None: + os.environ.pop("DATABASE_PATH", None) + else: + os.environ["DATABASE_PATH"] = old_db_path + + def test_run_market_review_reuses_generated_query_id_for_notification_diagnostics(self) -> None: + with tempfile.TemporaryDirectory() as temp_dir: + old_db_path = os.environ.get("DATABASE_PATH") + os.environ["DATABASE_PATH"] = os.path.join(temp_dir, "market_review_generated_query.db") + Config._instance = None + DatabaseManager.reset_instance() + notifier = self._make_notifier() + market_analyzer = MagicMock() + market_analyzer.run_daily_review_with_snapshot.return_value = SimpleNamespace( + report="## 今日大盘\n\n复盘正文", + market_light_snapshot={"region": "cn", "trade_date": "2026-03-06", "score": 60}, + ) + token = activate_run_diagnostic_context( + trace_id="trace-market-generated", + task_id="task-market-generated", + stock_code=market_review_module.MARKET_REVIEW_HISTORY_CODE, + trigger_source="cli", + ) + try: + with patch.object(market_review_module, "MarketAnalyzer", return_value=market_analyzer): + result = run_market_review( + notifier, + config=SimpleNamespace(report_language="zh", market_review_region="cn"), + send_notification=True, + trigger_source="cli", + ) + + self.assertEqual(result, "## 今日大盘\n\n复盘正文") + db = DatabaseManager.get_instance() + with db.get_session() as session: + rows = session.query(AnalysisHistory).all() + self.assertEqual(len(rows), 1) + row = rows[0] + self.assertTrue(row.query_id.startswith("market_review_")) + context_snapshot = json.loads(row.context_snapshot) + notification_runs = context_snapshot["diagnostics"]["notification_runs"] + self.assertEqual(notification_runs[-1]["status"], "success") + self.assertTrue(notification_runs[-1]["success"]) + finally: + reset_run_diagnostic_context(token) + DatabaseManager.reset_instance() + Config._instance = None + if old_db_path is None: + os.environ.pop("DATABASE_PATH", None) + else: + os.environ["DATABASE_PATH"] = old_db_path + if __name__ == "__main__": unittest.main() diff --git a/tests/test_run_diagnostics_p1.py b/tests/test_run_diagnostics_p1.py index 8f58e1c83..f95a85fb8 100644 --- a/tests/test_run_diagnostics_p1.py +++ b/tests/test_run_diagnostics_p1.py @@ -3,10 +3,12 @@ from __future__ import annotations +import json import os import sys import unittest from concurrent.futures import Future +from datetime import datetime from types import SimpleNamespace from unittest.mock import patch @@ -16,6 +18,7 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..") from data_provider.base import BaseFetcher, DataFetcherManager from src.services.run_diagnostics import ( + RunDiagnosticContext, activate_run_diagnostic_context, current_diagnostic_snapshot, record_provider_run, @@ -213,6 +216,146 @@ class RunDiagnosticsP1TestCase(unittest.TestCase): self.assertNotIn("secret", message) self.assertNotIn("example.com/webhook", message) + def test_diagnostic_event_sink_receives_provider_llm_history_and_notification_events(self) -> None: + events = [] + token = activate_run_diagnostic_context( + trace_id="trace-flow", + task_id="task-flow", + query_id="task-flow", + stock_code="600519", + event_sink=events.append, + ) + try: + record_provider_run( + data_type="daily_data", + provider="UnitFetcher", + operation="get_daily_data", + success=True, + latency_ms=12, + record_count=3, + ) + from src.services.run_diagnostics import record_history_run, record_llm_run, record_notification_run + + record_llm_run(success=True, model="deepseek-chat", duration_ms=34) + record_history_run(report_saved=True, metadata_saved=True, analysis_history_id=7) + record_notification_run(channel="report", status="success", success=True) + finally: + reset_run_diagnostic_context(token) + + self.assertEqual( + [event["type"] for event in events], + ["provider_run", "llm_run", "history_run", "notification_run"], + ) + self.assertEqual(events[0]["metadata"]["node"]["lane"], "data_source") + provider_node = events[0]["metadata"]["node"] + provider_started = datetime.fromisoformat(provider_node["started_at"]) + provider_ended = datetime.fromisoformat(provider_node["ended_at"]) + self.assertEqual(int((provider_ended - provider_started).total_seconds() * 1000), 12) + self.assertEqual(events[1]["metadata"]["node"]["kind"], "model") + llm_node = events[1]["metadata"]["node"] + llm_started = datetime.fromisoformat(llm_node["started_at"]) + llm_ended = datetime.fromisoformat(llm_node["ended_at"]) + self.assertEqual(int((llm_ended - llm_started).total_seconds() * 1000), 34) + + def test_provider_flow_event_attempt_index_is_scoped_by_data_type(self) -> None: + events = [] + token = activate_run_diagnostic_context( + trace_id="trace-provider-attempts", + task_id="task-provider-attempts", + query_id="query-provider-attempts", + stock_code="600519", + event_sink=events.append, + ) + try: + record_provider_run( + data_type="daily_data", + provider="DailyFetcher", + operation="get_daily_data", + success=True, + ) + record_provider_run( + data_type="news_search", + provider="NewsFetcher", + operation="search_stock_news", + success=True, + ) + record_provider_run( + data_type="daily_data", + provider="BackupDailyFetcher", + operation="get_daily_data", + success=True, + ) + finally: + reset_run_diagnostic_context(token) + + provider_node_ids = [ + event["metadata"]["node"]["id"] + for event in events + if event["type"] == "provider_run" + ] + self.assertEqual( + provider_node_ids, + [ + "provider_daily_data_dailyfetcher_1", + "provider_news_search_newsfetcher_1", + "provider_daily_data_backupdailyfetcher_2", + ], + ) + + def test_live_flow_event_sink_redacts_paths_and_sensitive_metadata(self) -> None: + events = [] + context = RunDiagnosticContext(trace_id="trace-live-redaction", event_sink=events.append) + + context._emit_flow_event( + { + "timestamp": "2026-06-08T10:00:01", + "severity": "danger", + "type": "provider_run", + "node_id": "provider_daily_data_unsafe_1", + "title": "Provider failed", + "message": ( + r"failed /home/activer/private/.env C:\Users\activer\.env " + "prompt=full-user-prompt raw_response=full-raw-response " + "https://hooks.example.com/webhook?key=secret" + ), + "metadata": { + "trace_id": "trace-live-redaction", + "operation": "/home/activer/project/.env", + "prompt": "full prompt body", + "raw_response": "full raw body", + "headers": {"Authorization": "Bearer sk-live-secret"}, + "proxy": "http://proxy_user:proxy_pass@proxy.internal", + "node": { + "id": "provider_daily_data_unsafe_1", + "lane": "data_source", + "kind": "data_source", + "label": "日线K线 · UnsafeFetcher", + "status": "failed", + "message": r"failed in C:\Users\activer\.env raw_response=full-raw-response", + }, + }, + } + ) + + payload = json.dumps(events, ensure_ascii=False) + for leaked in ( + "/home/activer", + "Users", + "full-user-prompt", + "full-raw-response", + "full prompt body", + "full raw body", + "hooks.example.com/webhook", + "sk-live-secret", + "proxy_user", + "proxy_pass", + ): + self.assertNotIn(leaked, payload) + self.assertIn("", payload) + self.assertIn("", payload) + self.assertIn('"prompt": ""', payload) + self.assertIn('"raw_response": ""', payload) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_run_flow.py b/tests/test_run_flow.py index 40823f797..f4fec5ecd 100644 --- a/tests/test_run_flow.py +++ b/tests/test_run_flow.py @@ -21,7 +21,13 @@ from src.services.run_flow import ( build_history_run_flow_snapshot, build_task_run_flow_snapshot, ) -from src.services.task_queue import TaskInfo, TaskStatus +from src.services.run_diagnostics import ( + activate_run_diagnostic_context, + current_diagnostic_snapshot, + record_provider_run, + reset_run_diagnostic_context, +) +from src.services.task_queue import AnalysisTaskQueue, TaskInfo, TaskStatus def _overview(*, blocks: list[dict]) -> dict: @@ -158,13 +164,20 @@ def _diagnostics(*, with_fallback: bool = False, unsafe: bool = False) -> dict: } -def _history_record(*, context_snapshot: dict | None, raw_result: dict | None = None) -> SimpleNamespace: +def _history_record( + *, + context_snapshot: dict | None, + raw_result: dict | None = None, + code: str = "600519", + name: str = "贵州茅台", + report_type: str = "detailed", +) -> SimpleNamespace: return SimpleNamespace( id=7, query_id="query-flow", - code="600519", - name="贵州茅台", - report_type="detailed", + code=code, + name=name, + report_type=report_type, created_at=datetime(2026, 6, 8, 10, 0, 6), raw_result=json.dumps(raw_result or {"success": True, "model_used": "deepseek-chat"}, ensure_ascii=False), context_snapshot=json.dumps(context_snapshot, ensure_ascii=False) if context_snapshot is not None else None, @@ -182,7 +195,29 @@ class _FakeHistoryDb: return self.record if self.record is not None and query_id == self.record.query_id else None +class _FakeMarketReviewDb: + def __init__(self, save_result): + self.save_result = save_result + self.saved_context_snapshot = None + self.updated_diagnostics = None + + def save_analysis_history(self, **kwargs): + self.saved_context_snapshot = kwargs.get("context_snapshot") + return self.save_result + + def update_analysis_history_diagnostics(self, *, query_id: str, code: str, diagnostics: dict) -> None: + _ = (query_id, code) + self.updated_diagnostics = diagnostics + + class RunFlowTestCase(unittest.TestCase): + def setUp(self) -> None: + self._original_queue = AnalysisTaskQueue._instance + AnalysisTaskQueue._instance = None + + def tearDown(self) -> None: + AnalysisTaskQueue._instance = self._original_queue + def test_active_task_missing_diagnostics_returns_skeleton_flow(self) -> None: task = TaskInfo( task_id="task-active", @@ -204,6 +239,203 @@ class RunFlowTestCase(unittest.TestCase): self.assertNotIn("provider_run", {event.type for event in snapshot.events}) self.assertNotIn("llm_run", {event.type for event in snapshot.events}) + def test_active_task_snapshot_includes_recent_flow_events_without_faking_missing_diagnostics(self) -> None: + task = TaskInfo( + task_id="task-active", + trace_id="trace-active", + stock_code="600519", + stock_name="贵州茅台", + status=TaskStatus.PROCESSING, + message="正在分析中", + created_at=datetime(2026, 6, 8, 10, 0, 0), + started_at=datetime(2026, 6, 8, 10, 0, 1), + flow_events=[ + { + "id": "flow-1", + "timestamp": "2026-06-08T10:00:02", + "severity": "success", + "type": "provider_run", + "node_id": "provider_daily_unit_1", + "title": "日线K线成功", + "message": "日线K线 UnitFetcher 成功", + "metadata": { + "provider": "UnitFetcher", + "node": { + "id": "provider_daily_unit_1", + "lane": "data_source", + "kind": "data_source", + "label": "日线K线 · UnitFetcher", + "status": "success", + "provider": "UnitFetcher", + "record_count": 30, + }, + }, + } + ], + ) + + snapshot = build_task_run_flow_snapshot(task) + + self.assertIn("provider_run", {event.type for event in snapshot.events}) + self.assertIn("provider_daily_unit_1", {node.id for node in snapshot.nodes}) + self.assertNotIn("llm_run", {event.type for event in snapshot.events}) + + def test_active_provider_events_only_link_fallbacks_within_same_data_type(self) -> None: + task = TaskInfo( + task_id="task-active-providers", + trace_id="trace-active-providers", + stock_code="600519", + stock_name="贵州茅台", + status=TaskStatus.PROCESSING, + created_at=datetime(2026, 6, 8, 10, 0, 0), + flow_events=[ + { + "id": "flow-daily", + "timestamp": "2026-06-08T10:00:02", + "severity": "success", + "type": "provider_run", + "node_id": "provider_daily_unit_1", + "title": "日线K线成功", + "metadata": { + "provider": "DailyFetcher", + "data_type": "daily_data", + "node": { + "id": "provider_daily_unit_1", + "lane": "data_source", + "kind": "data_source", + "label": "日线K线 · DailyFetcher", + "status": "success", + "provider": "DailyFetcher", + }, + }, + }, + { + "id": "flow-news", + "timestamp": "2026-06-08T10:00:03", + "severity": "success", + "type": "provider_run", + "node_id": "provider_news_unit_1", + "title": "新闻舆情成功", + "metadata": { + "provider": "NewsFetcher", + "data_type": "news_search", + "node": { + "id": "provider_news_unit_1", + "lane": "data_source", + "kind": "data_source", + "label": "新闻舆情 · NewsFetcher", + "status": "success", + "provider": "NewsFetcher", + }, + }, + }, + ], + ) + + snapshot = build_task_run_flow_snapshot(task) + edge_payload = [edge.model_dump(by_alias=True) for edge in snapshot.edges] + + self.assertEqual(snapshot.summary.fallback_count, 0) + self.assertFalse(any(edge["kind"] in {"fallback", "retry"} for edge in edge_payload)) + + def test_active_and_history_provider_nodes_share_id_and_core_fields(self) -> None: + flow_events: list[dict] = [] + token = activate_run_diagnostic_context( + trace_id="trace-provider-contract", + task_id="task-provider-contract", + query_id="query-provider-contract", + stock_code="600519", + trigger_source="api", + event_sink=flow_events.append, + ) + try: + record_provider_run( + data_type="daily_data", + provider="DailyFetcher", + operation="get_daily_data", + success=True, + latency_ms=120, + record_count=30, + ) + record_provider_run( + data_type="news_search", + provider="NewsFetcher", + operation="search_stock_news", + success=True, + latency_ms=80, + record_count=5, + ) + record_provider_run( + data_type="daily_data", + provider="BackupDailyFetcher", + operation="get_daily_data", + success=True, + latency_ms=90, + record_count=28, + ) + diagnostics = current_diagnostic_snapshot() + finally: + reset_run_diagnostic_context(token) + + self.assertIsNotNone(diagnostics) + active_snapshot = build_task_run_flow_snapshot( + TaskInfo( + task_id="task-provider-contract", + trace_id="trace-provider-contract", + stock_code="600519", + stock_name="贵州茅台", + status=TaskStatus.PROCESSING, + created_at=datetime(2026, 6, 8, 10, 0, 0), + flow_events=flow_events, + ) + ) + history_snapshot = build_history_run_flow_snapshot( + _history_record(context_snapshot={"diagnostics": diagnostics}) + ) + + expected_provider_ids = [ + "provider_daily_data_dailyfetcher_1", + "provider_news_search_newsfetcher_1", + "provider_daily_data_backupdailyfetcher_2", + ] + active_providers = { + node.id: node for node in active_snapshot.nodes if node.id in expected_provider_ids + } + history_providers = { + node.id: node for node in history_snapshot.nodes if node.id in expected_provider_ids + } + self.assertEqual(list(active_providers), expected_provider_ids) + self.assertEqual(list(history_providers), expected_provider_ids) + for node_id in expected_provider_ids: + active_provider = active_providers[node_id] + history_provider = history_providers[node_id] + for field in ("id", "label", "provider", "status", "record_count", "duration_ms"): + self.assertEqual( + getattr(active_provider, field), + getattr(history_provider, field), + f"{node_id}.{field}", + ) + + def test_task_queue_stores_bounded_flow_events_and_broadcasts_task_progress(self) -> None: + queue = AnalysisTaskQueue(max_workers=1) + queue._max_flow_events_per_task = 2 + task = TaskInfo( + task_id="task-flow", + stock_code="600519", + status=TaskStatus.PROCESSING, + ) + queue._tasks[task.task_id] = task + events = [] + queue._broadcast_event = lambda event_type, data: events.append((event_type, data)) + + queue.append_task_flow_event("task-flow", {"id": "evt-1", "type": "provider_run"}) + queue.append_task_flow_event("task-flow", {"id": "evt-2", "type": "llm_run"}) + queue.append_task_flow_event("task-flow", {"id": "evt-3", "type": "history_run"}) + + self.assertEqual([event["id"] for event in queue.get_task_flow_events("task-flow")], ["evt-2", "evt-3"]) + self.assertEqual(events[-1][0], "task_progress") + self.assertEqual(events[-1][1]["flow_event"]["id"], "evt-3") + def test_completed_history_uses_diagnostics_and_context_pack_overview(self) -> None: context_snapshot = { "diagnostics": _diagnostics(), @@ -245,6 +477,9 @@ class RunFlowTestCase(unittest.TestCase): node_ids = {node.id for node in snapshot.nodes} self.assertIn("context_pack", node_ids) self.assertTrue(any(node.kind == "model" and node.status == "success" for node in snapshot.nodes)) + quote_node = next(node for node in snapshot.nodes if node.id == "provider_realtime_quote_quotefetcher_1") + self.assertEqual(quote_node.started_at, "2026-06-08T10:00:00.880000") + self.assertEqual(quote_node.ended_at, "2026-06-08T10:00:01") self.assertIn("history_run", {event.type for event in snapshot.events}) self.assertIn("notification_run", {event.type for event in snapshot.events}) @@ -276,6 +511,52 @@ class RunFlowTestCase(unittest.TestCase): any(event.type == "provider_run" and event.severity == "warning" for event in snapshot.events) ) + def test_news_search_provider_runs_map_to_run_flow_nodes(self) -> None: + context_snapshot = { + "diagnostics": { + "trace_id": "trace-news", + "task_id": "task-news", + "query_id": "query-news", + "stock_code": "600519", + "trigger_source": "api", + "provider_runs": [ + { + "trace_id": "trace-news", + "data_type": "news_search", + "provider": "Tavily", + "operation": "search_stock_news", + "success": False, + "latency_ms": 500, + "error_type": "NoUsableNews", + "error_message_sanitized": "过滤后无有效新闻", + "created_at": "2026-06-08T10:00:01", + }, + { + "trace_id": "trace-news", + "data_type": "news_search", + "provider": "SearXNG", + "operation": "search_stock_news", + "success": True, + "latency_ms": 700, + "record_count": 3, + "created_at": "2026-06-08T10:00:02", + }, + ], + "llm_runs": [], + "history_runs": [], + "notification_runs": [], + } + } + + snapshot = build_history_run_flow_snapshot(_history_record(context_snapshot=context_snapshot)) + node_labels = {node.label for node in snapshot.nodes} + edge_payload = [edge.model_dump(by_alias=True) for edge in snapshot.edges] + + self.assertIn("新闻舆情 · Tavily", node_labels) + self.assertIn("新闻舆情 · SearXNG", node_labels) + self.assertTrue(any(edge["kind"] == "fallback" for edge in edge_payload)) + self.assertTrue(any(event.type == "provider_run" and event.node_id.endswith("searxng_2") for event in snapshot.events)) + def test_degraded_context_blocks_do_not_increment_fallback_count(self) -> None: diagnostics = _diagnostics() context_snapshot = { @@ -386,6 +667,86 @@ class RunFlowTestCase(unittest.TestCase): ): self.assertNotIn(leaked, payload) + def test_market_review_history_uses_same_run_flow_contract(self) -> None: + context_snapshot = { + "report_kind": "market_review", + "market_review_region": "cn", + "diagnostics": { + "trace_id": "trace-market", + "task_id": "task-market", + "query_id": "query-flow", + "stock_code": "MARKET", + "trigger_source": "api", + "provider_runs": [], + "llm_runs": [], + "history_runs": [ + { + "trace_id": "trace-market", + "report_saved": True, + "metadata_saved": True, + "analysis_history_id": 7, + "created_at": "2026-06-08T10:00:02", + } + ], + "notification_runs": [ + { + "trace_id": "trace-market", + "channel": "report", + "status": "skipped", + "success": False, + "attempts": 0, + "created_at": "2026-06-08T10:00:03", + } + ], + }, + } + + snapshot = build_history_run_flow_snapshot( + _history_record( + context_snapshot=context_snapshot, + code="MARKET", + name="大盘复盘", + report_type="market_review", + ) + ) + + self.assertEqual(snapshot.stock_code, "MARKET") + self.assertEqual(snapshot.task_id, "task-market") + self.assertIn("history_run", {event.type for event in snapshot.events}) + self.assertTrue(snapshot.lanes) + + def test_market_review_persist_records_diagnostics_without_bool_history_id(self) -> None: + from src.core.market_review import _persist_market_review_history + + fake_db = _FakeMarketReviewDb(save_result=True) + config = SimpleNamespace(report_language="zh") + token = activate_run_diagnostic_context( + trace_id="trace-market", + task_id="task-market", + query_id="query-flow", + stock_code="MARKET", + trigger_source="api", + ) + try: + with patch("src.storage.DatabaseManager.get_instance", return_value=fake_db): + saved = _persist_market_review_history( + review_report="大盘复盘报告", + markdown_report="# 大盘复盘报告", + region="cn", + config=config, + query_id="query-flow", + ) + finally: + reset_run_diagnostic_context(token) + + self.assertTrue(saved) + self.assertIsNotNone(fake_db.saved_context_snapshot) + self.assertIn("diagnostics", fake_db.saved_context_snapshot) + self.assertIsNotNone(fake_db.updated_diagnostics) + history_runs = fake_db.updated_diagnostics["history_runs"] + self.assertTrue(history_runs) + self.assertNotEqual(history_runs[-1].get("analysis_history_id"), True) + def test_flow_endpoints_return_404_for_missing_records(self) -> None: with self.assertRaises(HTTPException) as history_ctx: get_history_run_flow("404", db_manager=_FakeHistoryDb(None)) diff --git a/tests/test_search_news_freshness.py b/tests/test_search_news_freshness.py index 813e35bdd..877b6cf95 100644 --- a/tests/test_search_news_freshness.py +++ b/tests/test_search_news_freshness.py @@ -17,6 +17,11 @@ if "newspaper" not in sys.modules: sys.modules["newspaper"] = mock_np from src.search_service import SearchResponse, SearchResult, SearchService +from src.services.run_diagnostics import ( + activate_run_diagnostic_context, + current_diagnostic_snapshot, + reset_run_diagnostic_context, +) def _result( @@ -159,6 +164,50 @@ class SearchNewsFreshnessTestCase(unittest.TestCase): p1.search.assert_called_once() p2.search.assert_called_once() + def test_search_stock_news_records_provider_diagnostics_for_fallback(self) -> None: + """News search provider attempts should appear in run-flow diagnostics.""" + today = datetime.now().date() + old = (today - timedelta(days=90)).isoformat() + fresh = today.isoformat() + service = SearchService( + bocha_keys=["dummy_key"], + searxng_public_instances_enabled=False, + news_max_age_days=3, + news_strategy_profile="short", + ) + tavily = SimpleNamespace( + is_available=True, + name="Tavily", + search=MagicMock(return_value=_response([_result("too_old", old)])), + ) + searxng = SimpleNamespace( + is_available=True, + name="SearXNG", + search=MagicMock(return_value=_response([_result("贵州茅台 600519 最新公告", fresh)])), + ) + service._providers = [tavily, searxng] + + token = activate_run_diagnostic_context( + trace_id="trace-news", + task_id="task-news", + query_id="query-news", + stock_code="600519", + trigger_source="api", + ) + try: + response = service.search_stock_news("600519", "贵州茅台", max_results=3) + diagnostics = current_diagnostic_snapshot() + finally: + reset_run_diagnostic_context(token) + + self.assertEqual([item.title for item in response.results], ["贵州茅台 600519 最新公告"]) + provider_runs = diagnostics["provider_runs"] + self.assertEqual([run["data_type"] for run in provider_runs], ["news_search", "news_search"]) + self.assertEqual([run["provider"] for run in provider_runs], ["Tavily", "SearXNG"]) + self.assertFalse(provider_runs[0]["success"]) + self.assertTrue(provider_runs[1]["success"]) + self.assertEqual(provider_runs[1]["record_count"], 1) + def test_search_stock_news_tries_next_provider_when_chinese_context_is_english_only(self) -> None: """Chinese-preferred queries should not stop on English-only provider results.""" fresh = datetime.now().date().isoformat()