mirror of
https://github.com/ZhuLinsen/daily_stock_analysis
synced 2026-09-20 02:43:35 +08:00
* 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:
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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>
|
||||
);
|
||||
|
||||
@@ -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),
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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'),
|
||||
},
|
||||
}));
|
||||
|
||||
|
||||
@@ -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'
|
||||
);
|
||||
|
||||
// 无任务或不可见时不渲染
|
||||
|
||||
@@ -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
|
||||
|
||||
229
apps/dsa-web/src/hooks/__tests__/useRunFlowSnapshot.test.tsx
Normal file
229
apps/dsa-web/src/hooks/__tests__/useRunFlowSnapshot.test.tsx
Normal 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);
|
||||
});
|
||||
});
|
||||
@@ -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);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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 扫描误判为缺失。
|
||||
|
||||
@@ -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`,不影响报告读取。
|
||||
|
||||
@@ -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`。
|
||||
|
||||
## 运行流视图
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ A股自选股智能分析系统 - 异步任务队列
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import logging
|
||||
import threading
|
||||
import uuid
|
||||
@@ -53,6 +54,8 @@ class TaskStatus(str, Enum):
|
||||
PROCESSING = "processing" # In progress
|
||||
COMPLETED = "completed" # Completed
|
||||
FAILED = "failed" # Failed
|
||||
CANCEL_REQUESTED = "cancel_requested" # Cancellation requested
|
||||
CANCELLED = "cancelled" # Cancelled by user/system
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -82,6 +85,7 @@ class TaskInfo:
|
||||
skills: Optional[List[str]] = None
|
||||
report_language: Optional[str] = None
|
||||
trace_id: Optional[str] = None
|
||||
flow_events: List[Dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Convert task info into an API-friendly dictionary."""
|
||||
@@ -127,6 +131,7 @@ class TaskInfo:
|
||||
skills=list(self.skills) if self.skills is not None else None,
|
||||
report_language=self.report_language,
|
||||
trace_id=self.trace_id or self.task_id,
|
||||
flow_events=copy.deepcopy(self.flow_events),
|
||||
)
|
||||
|
||||
|
||||
@@ -190,6 +195,7 @@ class AnalysisTaskQueue:
|
||||
|
||||
# 任务历史保留数量(内存中)
|
||||
self._max_history = 100
|
||||
self._max_flow_events_per_task = 200
|
||||
|
||||
self._initialized = True
|
||||
logger.info(f"[TaskQueue] 初始化完成,最大并发: {max_workers}")
|
||||
@@ -524,6 +530,44 @@ class AnalysisTaskQueue:
|
||||
with self._data_lock:
|
||||
task = self._tasks.get(task_id)
|
||||
return task.copy() if task else None
|
||||
|
||||
def append_task_flow_event(
|
||||
self,
|
||||
task_id: str,
|
||||
flow_event: Dict[str, Any],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Append a recent run-flow event to an active task and broadcast it.
|
||||
|
||||
The event cache is deliberately bounded and fail-open; diagnostics must
|
||||
never affect the analysis pipeline.
|
||||
"""
|
||||
try:
|
||||
event_payload = copy.deepcopy(flow_event)
|
||||
except Exception:
|
||||
logger.debug("[TaskQueue] 忽略不可复制的运行流事件: task_id=%s", task_id)
|
||||
return None
|
||||
|
||||
with self._data_lock:
|
||||
task = self._tasks.get(task_id)
|
||||
if not task:
|
||||
return None
|
||||
task.flow_events.append(event_payload)
|
||||
if len(task.flow_events) > self._max_flow_events_per_task:
|
||||
task.flow_events = task.flow_events[-self._max_flow_events_per_task:]
|
||||
task_snapshot = task.copy()
|
||||
|
||||
payload = task_snapshot.to_dict()
|
||||
payload["flow_event"] = event_payload
|
||||
self._broadcast_event("task_progress", payload)
|
||||
return event_payload
|
||||
|
||||
def get_task_flow_events(self, task_id: str) -> List[Dict[str, Any]]:
|
||||
"""Return a copy of the recent run-flow events for a task."""
|
||||
with self._data_lock:
|
||||
task = self._tasks.get(task_id)
|
||||
if not task:
|
||||
return []
|
||||
return copy.deepcopy(task.flow_events)
|
||||
|
||||
def list_pending_tasks(self) -> List[TaskInfo]:
|
||||
"""
|
||||
@@ -535,7 +579,7 @@ class AnalysisTaskQueue:
|
||||
with self._data_lock:
|
||||
return [
|
||||
task.copy() for task in self._tasks.values()
|
||||
if task.status in (TaskStatus.PENDING, TaskStatus.PROCESSING)
|
||||
if task.status in (TaskStatus.PENDING, TaskStatus.PROCESSING, TaskStatus.CANCEL_REQUESTED)
|
||||
]
|
||||
|
||||
def list_all_tasks(self, limit: int = 50) -> List[TaskInfo]:
|
||||
@@ -669,6 +713,7 @@ class AnalysisTaskQueue:
|
||||
query_id=task_id,
|
||||
stock_code=stock_code,
|
||||
trigger_source=query_source,
|
||||
event_sink=lambda event: self.append_task_flow_event(task_id, event),
|
||||
)
|
||||
result = service.analyze_stock(
|
||||
stock_code=stock_code,
|
||||
@@ -777,6 +822,7 @@ class AnalysisTaskQueue:
|
||||
query_id=task_id,
|
||||
stock_code=task.stock_code,
|
||||
trigger_source="api",
|
||||
event_sink=lambda event: self.append_task_flow_event(task_id, event),
|
||||
)
|
||||
try:
|
||||
result = run_task()
|
||||
@@ -836,7 +882,7 @@ class AnalysisTaskQueue:
|
||||
# 按时间排序,删除旧的已完成任务
|
||||
completed_tasks = sorted(
|
||||
[t for t in self._tasks.values()
|
||||
if t.status in (TaskStatus.COMPLETED, TaskStatus.FAILED)],
|
||||
if t.status in (TaskStatus.COMPLETED, TaskStatus.FAILED, TaskStatus.CANCELLED)],
|
||||
key=lambda t: t.created_at
|
||||
)
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user