feat: [issue #1652] add live run-flow diagnostics (#1656)

* feat: add live run-flow updates and diagnostics

* feat: update task status filter to include 'cancel_requested'

* feat: enhance run-flow diagnostics and task management with cancel-requested status handling

* fix: stabilize run flow live updates

* fix: align live run-flow review contracts

* fix: sanitize live run-flow events
This commit is contained in:
LouisHong
2026-06-12 20:58:58 +08:00
committed by GitHub
parent 4082c0d012
commit efa41e0ada
28 changed files with 2357 additions and 254 deletions

View File

@@ -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:

View File

@@ -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,

View File

@@ -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<string, number>,
): 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<RunFlowGraphProps> = ({
lanes,
nodes,
@@ -81,7 +110,7 @@ export const RunFlowGraph: React.FC<RunFlowGraphProps> = ({
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<RunFlowGraphProps> = ({
const laneOrderByNode = new Map<string, number>();
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<RunFlowGraphProps> = ({
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<number>();
laneNodes.forEach((node) => {
@@ -405,6 +432,11 @@ export const RunFlowGraph: React.FC<RunFlowGraphProps> = ({
<span className="text-[11px] text-muted-text">{formatDuration(node.durationMs, t)}</span>
) : null}
</span>
{node.startedAt ? (
<span className="mt-1 block w-full truncate text-[11px] text-muted-text">
{t('runFlow.graph.startedAt')}: {formatDateTime(node.startedAt, language, t)}
</span>
) : null}
</button>
</Tooltip>
);

View File

@@ -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(
<RunFlowGraph
lanes={lanes}
nodes={timeOrderedNodes}
edges={[]}
/>,
);
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),
);
});
});

View File

@@ -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'),
},
}));

View File

@@ -21,9 +21,15 @@ const TaskItem: React.FC<TaskItemProps> = ({ 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<TaskItemProps> = ({ task, onOpenRunFlow }) => {
<div className="shrink-0">
{isProcessing ? (
<StatusDot tone="info" pulse className="h-2.5 w-2.5" aria-label={t('taskPanel.processingAria')} />
) : isCancelRequested ? (
<StatusDot tone="warning" pulse className="h-2.5 w-2.5" aria-label={t('taskPanel.cancelRequestedAria')} />
) : isPending ? (
<StatusDot tone="neutral" className="h-2.5 w-2.5" aria-label={t('taskPanel.pendingAria')} />
) : null}
</div>
{/* 任务信息 */}
<div className="flex-1 min-w-0">
<div className="min-w-0 flex-1 overflow-hidden">
<div className="flex items-center gap-2">
<span className="text-sm font-medium text-foreground truncate">
{task.stockName || task.stockCode}
@@ -93,7 +101,7 @@ const TaskItem: React.FC<TaskItemProps> = ({ task, onOpenRunFlow }) => {
</div>
{/* 状态标签 */}
<div className="flex flex-shrink-0 items-center gap-2">
<div className="relative z-10 flex flex-shrink-0 items-center gap-2">
{onOpenRunFlow ? (
<Tooltip content={t('taskPanel.openRunFlow')}>
<span className="inline-flex">
@@ -120,7 +128,7 @@ const TaskItem: React.FC<TaskItemProps> = ({ task, onOpenRunFlow }) => {
className="min-w-[4.75rem] justify-center gap-1.5 shadow-none"
aria-label={t('taskPanel.statusAria', { status: statusLabel })}
>
<StatusDot tone={statusTone} pulse={isProcessing} className="h-1.5 w-1.5" />
<StatusDot tone={statusTone} pulse={isProcessing || isCancelRequested} className="h-1.5 w-1.5" />
{statusLabel}
</Badge>
</div>
@@ -156,9 +164,9 @@ export const TaskPanel: React.FC<TaskPanelProps> = ({
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'
);
// 无任务或不可见时不渲染

View File

@@ -86,6 +86,39 @@ describe('TaskPanel', () => {
expect(onOpenRunFlow).toHaveBeenCalledWith(baseTask);
});
it('keeps cancel-requested tasks visible without rendering them as failed', () => {
render(
<TaskPanel
tasks={[
{
...baseTask,
status: 'cancel_requested',
message: '正在请求取消',
},
]}
/>,
);
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(
<TaskPanel
tasks={[
{
...baseTask,
status: 'cancelled',
},
]}
/>,
);
expect(container).toBeEmptyDOMElement();
});
it('does not render when there are no active tasks', () => {
const { container } = render(
<TaskPanel

View File

@@ -0,0 +1,229 @@
import { act, renderHook, waitFor } from '@testing-library/react';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { analysisApi } from '../../api/analysis';
import { historyApi } from '../../api/history';
import type { RunFlowSnapshot } from '../../types/runFlow';
import type { UseTaskStreamOptions } from '../useTaskStream';
import { useRunFlowSnapshot } from '../useRunFlowSnapshot';
vi.mock('../../api/analysis', () => ({
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<T>() {
let resolve!: (value: T) => void;
let reject!: (reason?: unknown) => void;
const promise = new Promise<T>((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<RunFlowSnapshot>();
const refreshedRequest = createDeferred<RunFlowSnapshot>();
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);
});
});

View File

@@ -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<string>) => 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<string>) => 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);
});
});

View File

@@ -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<string, RunFlowEvent>();
[...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<RunFlowNode>;
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<string, RunFlowEvent>();
[...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<RunFlowEvent[]>([]);
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,
});
}

View File

@@ -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<TaskStreamCallbacks>;
setIsConnected: (value: boolean) => void;
autoReconnect: boolean;
reconnectDelay: number;
};
let sharedEventSource: EventSource | null = null;
let sharedReconnectTimeout: ReturnType<typeof setTimeout> | null = null;
let sharedConnected = false;
let nextSubscriberId = 1;
const subscribers = new Map<number, TaskStreamSubscriber>();
// Convert snake_case payloads into camelCase TaskInfo objects.
const toTaskInfo = (data: Record<string, unknown>): 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<RunFlowEvent>(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<string>).data);
if (payload) {
forEachSubscriber((callbacks) => callbacks.onTaskCreated?.(payload.task));
}
});
eventSource.addEventListener('task_started', (e) => {
const payload = parseEventData((e as MessageEvent<string>).data);
if (payload) {
forEachSubscriber((callbacks) => callbacks.onTaskStarted?.(payload.task));
}
});
eventSource.addEventListener('task_progress', (e) => {
const payload = parseEventData((e as MessageEvent<string>).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<string>).data);
if (payload) {
forEachSubscriber((callbacks) => callbacks.onTaskCompleted?.(payload.task));
}
});
eventSource.addEventListener('task_failed', (e) => {
const payload = parseEventData((e as MessageEvent<string>).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<EventSource | null>(null);
const [isConnected, setIsConnected] = useState(false);
const reconnectTimeoutRef = useRef<ReturnType<typeof setTimeout> | null>(null);
const connectRef = useRef<() => void>(() => {});
const subscriberIdRef = useRef<number | null>(null);
const connectTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null);
// Store callbacks in a ref to avoid reconnecting on every render.
const callbacksRef = useRef({
const callbacksRef = useRef<TaskStreamCallbacks>({
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<string, unknown>): 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,

View File

@@ -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<UiTextKey, string> = {
'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<UiTextKey, string> = {
'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',

View File

@@ -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);

View File

@@ -856,7 +856,7 @@ export const useStockPoolStore = create<StockPoolState>((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<StockPoolState>((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;

View File

@@ -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;

View File

@@ -20,6 +20,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/).
- [新功能] Web 端为活跃任务、历史报告和大盘复盘报告补充运行流视图入口,支持查看运行摘要、拓扑节点、事件流和基础排障详情。
- [修复] 修复历史报告运行流快照在混合时区事件时间戳下返回 500 的问题。
- [改进] #1459 持仓管理页新增持仓账户删除入口,复用现有账户软删除接口,误建账户会从默认列表、快照、风险、录入入口和事件列表隐藏且不物理清理历史流水。
- [修复] 修复运行流 live SSE 事件未复用快照层递归脱敏规则的问题避免本地路径、prompt/raw response、代理头等敏感诊断字段在 refetch 前短暂暴露。
<!-- 新条目格式:- [类型] 描述(类型取值:新功能/改进/修复/文档/测试/chore-->
<!-- 每条独立一行追加到本段末尾,无需分类标题,合并时冲突最小 -->
- [修复] 桌面发布打包改用冻结可执行文件运行时探针校验 `alphasift.dsa_adapter`,避免 macOS PyInstaller 将模块内嵌进可执行文件时被文件系统/zip 扫描误判为缺失。

View File

@@ -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`,不影响报告读取。

View File

@@ -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`
## 运行流视图

View File

@@ -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

View File

@@ -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(

View File

@@ -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 <redacted>",
),
)
_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"(?<![\w:/.-])(?:/(?:home|Users|root|var|tmp|opt|etc)/[^\s,;]+|[A-Za-z]:\\[^\s,;]+)"
)
_SENSITIVE_ASSIGNMENT_RE = re.compile(
r"(?i)\b(api[_-]?key|access[_-]?token|token|secret|password|passwd|cookie|webhook|sendkey|"
r"prompt|raw[_-]?prompt|raw[_-]?response)\s*[:=]\s*([^\s,&;]+)"
)
def build_trace_id() -> 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("<redacted-url>", text)
text = _LOCAL_ABSOLUTE_PATH_RE.sub("<redacted-path>", text)
text = _SENSITIVE_ASSIGNMENT_RE.sub(lambda match: f"{match.group(1)}=<redacted>", 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 "<truncated>"
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] = "<redacted>"
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,

View File

@@ -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"(?<![\w:/.-])(?:/(?:home|Users|root|var|tmp|opt|etc)/[^\s,;]+|[A-Za-z]:\\[^\s,;]+)"
)
_SENSITIVE_ASSIGNMENT_RE = re.compile(
r"(?i)\b(api[_-]?key|access[_-]?token|token|secret|password|passwd|cookie|webhook|sendkey|"
r"prompt|raw[_-]?prompt|raw[_-]?response)\s*[:=]\s*([^\s,&;]+)"
)
def build_task_run_flow_snapshot(
task: Any,
*,
@@ -157,6 +146,13 @@ def build_task_run_flow_snapshot(
_put_skeleton_tail(nodes, edges, anchor_node_id="task_queue", status=flow_status)
_append_task_events(events, task, flow_status)
_append_active_flow_events(
nodes,
edges,
events,
_as_list(getattr(task, "flow_events", None)),
flow_status=flow_status,
)
summary = _build_summary(
nodes,
@@ -397,6 +393,7 @@ def _append_provider_runs(
status = _provider_run_status(run, had_previous_failure=had_previous_failure)
duration_ms = _safe_int(run.get("latency_ms"))
timestamp = _datetime_to_iso(run.get("created_at"))
started_at = _started_at_from_end_and_duration(timestamp, duration_ms)
message = _provider_run_message(label, provider, run, success=success)
block_key = _DATA_TYPE_TO_BLOCK_KEY.get(data_type, data_type)
@@ -408,6 +405,7 @@ def _append_provider_runs(
label=f"{label} · {provider}",
status=status,
provider=provider,
started_at=started_at,
ended_at=timestamp,
duration_ms=duration_ms,
attempts=1,
@@ -479,6 +477,7 @@ def _append_context_blocks(
if not overview:
return
metadata = overview.get("metadata") if isinstance(overview.get("metadata"), Mapping) else {}
overview_timestamp = overview.get("created_at")
for block in _as_list(overview.get("blocks")):
block_map = _as_mapping(block)
key = _safe_key(block_map.get("key"))
@@ -495,6 +494,8 @@ def _append_context_blocks(
label=_safe_text(block_map.get("label"), max_length=80) or key,
status=status,
provider=block_map.get("source"),
started_at=overview_timestamp,
ended_at=overview_timestamp,
record_count=_safe_int(record_count),
message=_context_block_message(block_map),
metadata={
@@ -551,6 +552,7 @@ def _append_llm_runs(
status = "fallback"
timestamp = _datetime_to_iso(run.get("created_at"))
duration_ms = _safe_int(run.get("duration_ms"))
started_at = _started_at_from_end_and_duration(timestamp, duration_ms)
message = _llm_run_message(model, run, success=success)
edge_kind = "data"
if index > 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("<redacted-url>", text)
text = _LOCAL_ABSOLUTE_PATH_RE.sub("<redacted-path>", text)
text = _SENSITIVE_ASSIGNMENT_RE.sub(lambda match: f"{match.group(1)}=<redacted>", 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 "<truncated>"
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] = "<redacted>"
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

View File

@@ -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}")
@@ -525,6 +531,44 @@ class AnalysisTaskQueue:
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]:
"""
获取所有进行中的任务pending + processing
@@ -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
)

View File

@@ -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")

View File

@@ -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()

View File

@@ -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("<redacted-path>", payload)
self.assertIn("<redacted>", payload)
self.assertIn('"prompt": "<redacted>"', payload)
self.assertIn('"raw_response": "<redacted>"', payload)
if __name__ == "__main__":
unittest.main()

View File

@@ -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))

View File

@@ -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()