mirror of
https://github.com/ZhuLinsen/daily_stock_analysis.git
synced 2026-10-06 14:33:11 +08:00
feat: persist Agent Chat Skill selection by session (#2160)
* feat: persist Agent Chat skill selection by session * test: satisfy ChatPage mock immutability lint * fix: preserve implicit Skill state for legacy sessions * fix: make top-level Skill selection authoritative * fix: preserve session skills for invalid requests --------- Co-authored-by: zhulinsen <42829555+ZhuLinsen@users.noreply.github.com>
This commit is contained in:
co-authored by
zhulinsen
parent
ae19329d66
commit
ed848da6f0
@@ -19,6 +19,7 @@ from src.storage import DatabaseManager
|
||||
from src.config import get_config, Config
|
||||
from src.services.system_config_service import SystemConfigService
|
||||
from src.services.runtime_scheduler import RuntimeSchedulerService
|
||||
from src.services.agent_chat_session_service import AgentChatSessionService
|
||||
|
||||
|
||||
def get_db() -> Generator[Session, None, None]:
|
||||
@@ -63,6 +64,11 @@ def get_database_manager() -> DatabaseManager:
|
||||
return DatabaseManager.get_instance()
|
||||
|
||||
|
||||
def get_agent_chat_session_service() -> AgentChatSessionService:
|
||||
"""Build an Agent Chat session service for the current database manager."""
|
||||
return AgentChatSessionService(DatabaseManager.get_instance())
|
||||
|
||||
|
||||
def get_system_config_service(request: Request) -> SystemConfigService:
|
||||
"""Get app-lifecycle shared SystemConfigService instance."""
|
||||
service = getattr(request.app.state, "system_config_service", None)
|
||||
|
||||
+62
-22
@@ -10,12 +10,14 @@ import threading
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import AliasChoices, BaseModel, ConfigDict, Field
|
||||
|
||||
from api.deps import get_agent_chat_session_service
|
||||
from api.v1.schemas.system_config import AgentBackendStatusResponse
|
||||
from src.config import get_config
|
||||
from src.services.agent_chat_session_service import AgentChatSessionService
|
||||
from src.services.agent_model_service import list_agent_model_deployments
|
||||
|
||||
# Tool name -> Chinese display name mapping
|
||||
@@ -66,6 +68,8 @@ class ChatRequest(BaseModel):
|
||||
def _build_agent_chat_context(request: ChatRequest, config, skills: Optional[List[str]]) -> Dict[str, Any]:
|
||||
"""Build the shared context contract for regular and streaming Agent Chat."""
|
||||
context = dict(request.context or {})
|
||||
context.pop("skills", None)
|
||||
context.pop("strategies", None)
|
||||
if skills is not None:
|
||||
context["skills"] = skills
|
||||
report_language = context.get("report_language")
|
||||
@@ -191,7 +195,10 @@ async def get_strategies():
|
||||
)
|
||||
|
||||
@router.post("/chat", response_model=ChatResponse)
|
||||
async def agent_chat(request: ChatRequest):
|
||||
async def agent_chat(
|
||||
request: ChatRequest,
|
||||
session_service: AgentChatSessionService = Depends(get_agent_chat_session_service),
|
||||
):
|
||||
"""
|
||||
Chat with the AI Agent without progress events.
|
||||
|
||||
@@ -213,7 +220,13 @@ async def agent_chat(request: ChatRequest):
|
||||
session_id = request.session_id or str(uuid.uuid4())
|
||||
|
||||
try:
|
||||
skills = request.effective_skills
|
||||
skill_selection = session_service.resolve_skill_selection(
|
||||
config,
|
||||
session_id,
|
||||
request.effective_skills,
|
||||
)
|
||||
skills = skill_selection.effective_skill_ids
|
||||
selected_skill_ids = skill_selection.selected_skill_ids_update
|
||||
executor = _build_executor(config, skills or None)
|
||||
|
||||
ctx = _build_agent_chat_context(request, config, skills)
|
||||
@@ -223,7 +236,7 @@ async def agent_chat(request: ChatRequest):
|
||||
result = await loop.run_in_executor(
|
||||
None,
|
||||
lambda: executor.chat(message=request.message, session_id=session_id,
|
||||
context=ctx),
|
||||
context=ctx, selected_skill_ids=selected_skill_ids),
|
||||
)
|
||||
|
||||
return ChatResponse(
|
||||
@@ -249,13 +262,21 @@ class SessionItem(BaseModel):
|
||||
class SessionsResponse(BaseModel):
|
||||
sessions: List[SessionItem]
|
||||
|
||||
class SessionStateResponse(BaseModel):
|
||||
selected_skill_ids: Optional[List[str]]
|
||||
|
||||
class SessionMessagesResponse(BaseModel):
|
||||
session_id: str
|
||||
messages: List[Dict[str, Any]]
|
||||
session_state: SessionStateResponse
|
||||
|
||||
|
||||
@router.get("/chat/sessions", response_model=SessionsResponse)
|
||||
async def list_chat_sessions(limit: int = 50, user_id: Optional[str] = None):
|
||||
async def list_chat_sessions(
|
||||
limit: int = 50,
|
||||
user_id: Optional[str] = None,
|
||||
session_service: AgentChatSessionService = Depends(get_agent_chat_session_service),
|
||||
):
|
||||
"""获取聊天会话列表
|
||||
|
||||
Args:
|
||||
@@ -266,28 +287,37 @@ async def list_chat_sessions(limit: int = 50, user_id: Optional[str] = None):
|
||||
include the platform prefix, e.g. ``telegram_12345``,
|
||||
``feishu_ou_abc``.
|
||||
"""
|
||||
from src.storage import get_db
|
||||
sessions = get_db().get_chat_sessions(
|
||||
limit=limit,
|
||||
session_prefix=user_id,
|
||||
extra_session_ids=[user_id] if user_id else None,
|
||||
)
|
||||
sessions = session_service.list_sessions(limit, user_id)
|
||||
return SessionsResponse(sessions=sessions)
|
||||
|
||||
|
||||
@router.get("/chat/sessions/{session_id}", response_model=SessionMessagesResponse)
|
||||
async def get_chat_session_messages(session_id: str, limit: int = 100):
|
||||
async def get_chat_session_messages(
|
||||
session_id: str,
|
||||
limit: int = 100,
|
||||
session_service: AgentChatSessionService = Depends(get_agent_chat_session_service),
|
||||
):
|
||||
"""获取单个会话的完整消息"""
|
||||
from src.storage import get_db
|
||||
messages = get_db().get_conversation_messages(session_id, limit=limit)
|
||||
return SessionMessagesResponse(session_id=session_id, messages=messages)
|
||||
detail = session_service.get_session_detail(
|
||||
session_id,
|
||||
limit,
|
||||
)
|
||||
return SessionMessagesResponse(
|
||||
session_id=session_id,
|
||||
messages=detail.messages,
|
||||
session_state=SessionStateResponse(
|
||||
selected_skill_ids=detail.selected_skill_ids,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@router.delete("/chat/sessions/{session_id}")
|
||||
async def delete_chat_session(session_id: str):
|
||||
async def delete_chat_session(
|
||||
session_id: str,
|
||||
session_service: AgentChatSessionService = Depends(get_agent_chat_session_service),
|
||||
):
|
||||
"""删除指定会话"""
|
||||
from src.storage import get_db
|
||||
count = get_db().delete_conversation_session(session_id)
|
||||
count = session_service.delete_session(session_id)
|
||||
return {"deleted": count}
|
||||
|
||||
|
||||
@@ -444,7 +474,10 @@ async def agent_research(request: ResearchRequest):
|
||||
|
||||
|
||||
@router.post("/chat/stream")
|
||||
async def agent_chat_stream(request: ChatRequest):
|
||||
async def agent_chat_stream(
|
||||
request: ChatRequest,
|
||||
session_service: AgentChatSessionService = Depends(get_agent_chat_session_service),
|
||||
):
|
||||
"""
|
||||
Chat with the AI Agent, streaming progress via SSE.
|
||||
Each SSE event is a JSON object with a 'type' field:
|
||||
@@ -468,6 +501,15 @@ async def agent_chat_stream(request: ChatRequest):
|
||||
queue: asyncio.Queue = asyncio.Queue()
|
||||
cancel_event = threading.Event()
|
||||
request_id = request.request_id or str(uuid.uuid4())
|
||||
skill_selection = session_service.resolve_skill_selection(
|
||||
config,
|
||||
session_id,
|
||||
request.effective_skills,
|
||||
)
|
||||
skills = skill_selection.effective_skill_ids
|
||||
selected_skill_ids = skill_selection.selected_skill_ids_update
|
||||
stream_ctx = _build_agent_chat_context(request, config, skills)
|
||||
|
||||
if backend_id == "codex_app_server":
|
||||
with _ACTIVE_CODEX_STREAMS_LOCK:
|
||||
if request_id in _ACTIVE_CODEX_STREAMS:
|
||||
@@ -480,9 +522,6 @@ async def agent_chat_stream(request: ChatRequest):
|
||||
)
|
||||
_ACTIVE_CODEX_STREAMS[request_id] = cancel_event
|
||||
|
||||
skills = request.effective_skills
|
||||
stream_ctx = _build_agent_chat_context(request, config, skills)
|
||||
|
||||
def progress_callback(event: dict):
|
||||
if backend_id == "codex_app_server" and cancel_event.is_set():
|
||||
return
|
||||
@@ -538,6 +577,7 @@ async def agent_chat_stream(request: ChatRequest):
|
||||
message=request.message,
|
||||
session_id=session_id,
|
||||
context=stream_ctx,
|
||||
selected_skill_ids=selected_skill_ids,
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
|
||||
@@ -40,4 +40,41 @@ describe('agentApi', () => {
|
||||
message: 'Codex login is required',
|
||||
});
|
||||
});
|
||||
|
||||
it('returns session messages together with persisted Skill state', async () => {
|
||||
get.mockResolvedValueOnce({
|
||||
data: {
|
||||
session_id: 'session-1',
|
||||
messages: [
|
||||
{ id: '1', role: 'user', content: '分析 AAPL', created_at: null },
|
||||
],
|
||||
session_state: {
|
||||
selected_skill_ids: ['technical', 'risk'],
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const result = await agentApi.getChatSessionMessages('session-1');
|
||||
|
||||
expect(get).toHaveBeenCalledWith('/api/v1/agent/chat/sessions/session-1');
|
||||
expect(result.session_state.selected_skill_ids).toEqual(['technical', 'risk']);
|
||||
});
|
||||
|
||||
it('preserves null when a legacy session has no persisted Skill state', async () => {
|
||||
get.mockResolvedValueOnce({
|
||||
data: {
|
||||
session_id: 'legacy-session',
|
||||
messages: [
|
||||
{ id: '1', role: 'user', content: '继续分析', created_at: null },
|
||||
],
|
||||
session_state: {
|
||||
selected_skill_ids: null,
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const result = await agentApi.getChatSessionMessages('legacy-session');
|
||||
|
||||
expect(result.session_state.selected_skill_ids).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -66,6 +66,14 @@ export interface ChatSessionMessage {
|
||||
created_at: string | null;
|
||||
}
|
||||
|
||||
export interface ChatSessionDetail {
|
||||
session_id: string;
|
||||
messages: ChatSessionMessage[];
|
||||
session_state: {
|
||||
selected_skill_ids: string[] | null;
|
||||
};
|
||||
}
|
||||
|
||||
export const agentApi = {
|
||||
async chat(payload: ChatRequest): Promise<ChatResponse> {
|
||||
const response = await apiClient.post<ChatResponse>('/api/v1/agent/chat', payload, {
|
||||
@@ -85,9 +93,9 @@ export const agentApi = {
|
||||
const response = await apiClient.get<{ sessions: ChatSessionItem[] }>('/api/v1/agent/chat/sessions', { params: { limit } });
|
||||
return response.data.sessions;
|
||||
},
|
||||
async getChatSessionMessages(sessionId: string): Promise<ChatSessionMessage[]> {
|
||||
const response = await apiClient.get<{ messages: ChatSessionMessage[] }>(`/api/v1/agent/chat/sessions/${sessionId}`);
|
||||
return response.data.messages;
|
||||
async getChatSessionMessages(sessionId: string): Promise<ChatSessionDetail> {
|
||||
const response = await apiClient.get<ChatSessionDetail>(`/api/v1/agent/chat/sessions/${sessionId}`);
|
||||
return response.data;
|
||||
},
|
||||
async deleteChatSession(sessionId: string): Promise<void> {
|
||||
await apiClient.delete(`/api/v1/agent/chat/sessions/${sessionId}`);
|
||||
|
||||
@@ -205,7 +205,7 @@ const ChatPage: React.FC = () => {
|
||||
const [searchParams, setSearchParams] = useSearchParams();
|
||||
const [input, setInput] = useState('');
|
||||
const [skills, setSkills] = useState<SkillInfo[]>([]);
|
||||
const [selectedSkillIds, setSelectedSkillIds] = useState<string[]>([]);
|
||||
const [defaultSkillIds, setDefaultSkillIds] = useState<string[]>([]);
|
||||
const [showSkillDesc, setShowSkillDesc] = useState<string | null>(null);
|
||||
const [mobileSkillPickerOpen, setMobileSkillPickerOpen] = useState(false);
|
||||
const [expandedThinking, setExpandedThinking] = useState<Set<string>>(new Set());
|
||||
@@ -341,6 +341,7 @@ const ChatPage: React.FC = () => {
|
||||
|
||||
const {
|
||||
messages,
|
||||
selectedSkillIds: sessionSelectedSkillIds,
|
||||
loading,
|
||||
progressSteps,
|
||||
sessionId,
|
||||
@@ -350,6 +351,7 @@ const ChatPage: React.FC = () => {
|
||||
stopping,
|
||||
terminalStatus,
|
||||
stopError,
|
||||
setSelectedSkillIds,
|
||||
loadSessions,
|
||||
loadInitialSession,
|
||||
switchSession,
|
||||
@@ -357,6 +359,7 @@ const ChatPage: React.FC = () => {
|
||||
startStream,
|
||||
clearCompletionBadge,
|
||||
} = useAgentChatStore();
|
||||
const selectedSkillIds = sessionSelectedSkillIds ?? defaultSkillIds;
|
||||
|
||||
useEffect(() => {
|
||||
if (activeStockContext || messages.length === 0) {
|
||||
@@ -440,7 +443,7 @@ const ChatPage: React.FC = () => {
|
||||
res.default_skill_id ||
|
||||
res.skills[0]?.id ||
|
||||
'';
|
||||
setSelectedSkillIds(defaultId ? [defaultId] : []);
|
||||
setDefaultSkillIds(defaultId ? [defaultId] : []);
|
||||
})
|
||||
.catch((error) => {
|
||||
console.error('Failed to load chat skills:', error);
|
||||
@@ -581,16 +584,14 @@ const ChatPage: React.FC = () => {
|
||||
}, []);
|
||||
|
||||
const toggleSkillSelection = useCallback((skillId: string) => {
|
||||
setSelectedSkillIds((prev) => {
|
||||
if (prev.includes(skillId)) {
|
||||
return prev.filter((id) => id !== skillId);
|
||||
}
|
||||
if (prev.length >= MAX_SELECTED_SKILLS) {
|
||||
return prev;
|
||||
}
|
||||
return [...prev, skillId];
|
||||
});
|
||||
}, []);
|
||||
if (selectedSkillIds.includes(skillId)) {
|
||||
setSelectedSkillIds(selectedSkillIds.filter((id) => id !== skillId));
|
||||
return;
|
||||
}
|
||||
if (selectedSkillIds.length < MAX_SELECTED_SKILLS) {
|
||||
setSelectedSkillIds([...selectedSkillIds, skillId]);
|
||||
}
|
||||
}, [selectedSkillIds, setSelectedSkillIds]);
|
||||
|
||||
const handleStartNewChat = useCallback(() => {
|
||||
followUpContextRef.current = null;
|
||||
@@ -682,7 +683,10 @@ const ChatPage: React.FC = () => {
|
||||
if (overrideMessage !== undefined) {
|
||||
setInput(msgText);
|
||||
}
|
||||
const usedSkillIds = normalizeSelectedSkillIds(overrideSkillIds ?? selectedSkillIds);
|
||||
const requestedSkillIds = overrideSkillIds ?? sessionSelectedSkillIds;
|
||||
const usedSkillIds = normalizeSelectedSkillIds(
|
||||
requestedSkillIds ?? selectedSkillIds,
|
||||
);
|
||||
const usedSkillNames = usedSkillIds.length > 0 ? getSkillNames(usedSkillIds) : ['通用'];
|
||||
const codexStockContext = agentStatus?.backend === 'codex_app_server'
|
||||
? overrideStockContext
|
||||
@@ -714,7 +718,9 @@ const ChatPage: React.FC = () => {
|
||||
const payload = {
|
||||
message: msgText,
|
||||
session_id: sessionId,
|
||||
...(usedSkillIds.length > 0 ? { skills: usedSkillIds } : {}),
|
||||
...(requestedSkillIds !== null
|
||||
? { skills: normalizeSelectedSkillIds(requestedSkillIds) }
|
||||
: {}),
|
||||
context: contextForSend ?? undefined,
|
||||
};
|
||||
await startStream(payload, {
|
||||
@@ -734,7 +740,7 @@ const ChatPage: React.FC = () => {
|
||||
},
|
||||
});
|
||||
},
|
||||
[activeStockContext, agentAvailable, agentStatus, getSkillNames, input, loading, normalizeSelectedSkillIds, requestScrollToBottom, selectedSkillIds, sessionId, startStream, stockIndex],
|
||||
[activeStockContext, agentAvailable, agentStatus, getSkillNames, input, loading, normalizeSelectedSkillIds, requestScrollToBottom, selectedSkillIds, sessionId, sessionSelectedSkillIds, startStream, stockIndex],
|
||||
);
|
||||
|
||||
const handleKeyDown = (e: React.KeyboardEvent<HTMLTextAreaElement>) => {
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { act, fireEvent, render, screen, waitFor } from '@testing-library/react';
|
||||
import { StrictMode } from 'react';
|
||||
import { StrictMode, useState } from 'react';
|
||||
import { createMemoryRouter, MemoryRouter, RouterProvider } from 'react-router-dom';
|
||||
import { beforeAll, beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
import { createParsedApiError } from '../../api/error';
|
||||
@@ -63,6 +63,7 @@ const mockStartNewChat = vi.fn();
|
||||
|
||||
const mockStoreState = {
|
||||
messages: [] as Message[],
|
||||
selectedSkillIds: null as string[] | null,
|
||||
loading: false,
|
||||
progressSteps: [] as ProgressStep[],
|
||||
sessionId: 'session-1',
|
||||
@@ -129,9 +130,24 @@ vi.mock('../../hooks/useStockIndex', () => ({
|
||||
}));
|
||||
|
||||
vi.mock('../../stores/agentChatStore', () => {
|
||||
type MockStore = typeof mockStoreState & {
|
||||
setSelectedSkillIds: (skillIds: string[]) => void;
|
||||
};
|
||||
const useAgentChatStore = (
|
||||
selector?: (state: typeof mockStoreState) => unknown
|
||||
) => (typeof selector === 'function' ? selector(mockStoreState) : mockStoreState);
|
||||
selector?: (state: MockStore) => unknown
|
||||
) => {
|
||||
const [selectedSkillIds, setSelectedSkillIdsState] = useState(
|
||||
mockStoreState.selectedSkillIds,
|
||||
);
|
||||
const state: MockStore = {
|
||||
...mockStoreState,
|
||||
selectedSkillIds,
|
||||
setSelectedSkillIds: (skillIds) => {
|
||||
setSelectedSkillIdsState(skillIds);
|
||||
},
|
||||
};
|
||||
return typeof selector === 'function' ? selector(state) : state;
|
||||
};
|
||||
|
||||
useAgentChatStore.getState = () => ({
|
||||
startNewChat: mockStartNewChat,
|
||||
@@ -176,6 +192,7 @@ beforeEach(() => {
|
||||
window.localStorage.removeItem(UI_LANGUAGE_STORAGE_KEY);
|
||||
mockGetStatus.mockReset();
|
||||
mockStoreState.messages = [];
|
||||
mockStoreState.selectedSkillIds = null;
|
||||
mockStoreState.loading = false;
|
||||
mockStoreState.progressSteps = [];
|
||||
mockStoreState.chatError = null;
|
||||
@@ -813,6 +830,84 @@ describe('ChatPage', () => {
|
||||
expect(screen.getByRole('checkbox', { name: '通用分析' })).not.toBeChecked();
|
||||
});
|
||||
|
||||
it('keeps the restored session skills when the Skill catalog finishes loading', async () => {
|
||||
mockStoreState.selectedSkillIds = ['ma_golden_cross'];
|
||||
mockGetSkills.mockResolvedValue({
|
||||
skills: [
|
||||
{ id: 'bull_trend', name: '趋势分析', description: '默认趋势' },
|
||||
{ id: 'ma_golden_cross', name: '均线金叉', description: '均线交叉' },
|
||||
],
|
||||
default_skill_id: 'bull_trend',
|
||||
});
|
||||
|
||||
render(
|
||||
<MemoryRouter initialEntries={['/chat']}>
|
||||
<ChatPage />
|
||||
</MemoryRouter>
|
||||
);
|
||||
|
||||
expect(await screen.findByRole('checkbox', { name: '均线金叉' })).toBeChecked();
|
||||
expect(screen.getByRole('checkbox', { name: '趋势分析' })).not.toBeChecked();
|
||||
|
||||
fireEvent.change(screen.getByPlaceholderText(/分析 600519/), {
|
||||
target: { value: '继续分析' },
|
||||
});
|
||||
fireEvent.click(screen.getByRole('button', { name: '发送' }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockStartStream).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ skills: ['ma_golden_cross'] }),
|
||||
expect.any(Object),
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
it('omits skills for an untouched new session so the server resolves its default', async () => {
|
||||
render(
|
||||
<MemoryRouter initialEntries={['/chat']}>
|
||||
<ChatPage />
|
||||
</MemoryRouter>
|
||||
);
|
||||
|
||||
expect(await screen.findByRole('checkbox', { name: '趋势分析' })).toBeChecked();
|
||||
fireEvent.change(screen.getByPlaceholderText(/分析 600519/), {
|
||||
target: { value: '分析 AAPL' },
|
||||
});
|
||||
fireEvent.click(screen.getByRole('button', { name: '发送' }));
|
||||
|
||||
await waitFor(() => expect(mockStartStream).toHaveBeenCalled());
|
||||
expect(mockStartStream.mock.calls.at(-1)?.[0]).not.toHaveProperty('skills');
|
||||
});
|
||||
|
||||
it('omits skills when continuing a legacy session without persisted Skill state', async () => {
|
||||
mockStoreState.messages = [
|
||||
{ id: 'legacy-user', role: 'user', content: '分析 AAPL' },
|
||||
{ id: 'legacy-assistant', role: 'assistant', content: '历史分析结果' },
|
||||
];
|
||||
mockStoreState.selectedSkillIds = null;
|
||||
|
||||
render(
|
||||
<MemoryRouter initialEntries={['/chat']}>
|
||||
<ChatPage />
|
||||
</MemoryRouter>
|
||||
);
|
||||
|
||||
expect(await screen.findByRole('checkbox', { name: '趋势分析' })).toBeChecked();
|
||||
fireEvent.change(screen.getByPlaceholderText(/分析 600519/), {
|
||||
target: { value: '继续分析' },
|
||||
});
|
||||
fireEvent.click(screen.getByRole('button', { name: '发送' }));
|
||||
|
||||
await waitFor(() => expect(mockStartStream).toHaveBeenCalled());
|
||||
expect(mockStartStream.mock.calls.at(-1)?.[0]).toEqual(
|
||||
expect.objectContaining({
|
||||
message: '继续分析',
|
||||
session_id: 'session-1',
|
||||
}),
|
||||
);
|
||||
expect(mockStartStream.mock.calls.at(-1)?.[0]).not.toHaveProperty('skills');
|
||||
});
|
||||
|
||||
it('sends multiple selected skills in order', async () => {
|
||||
mockGetSkills.mockResolvedValue({
|
||||
skills: [
|
||||
@@ -926,7 +1021,7 @@ describe('ChatPage', () => {
|
||||
expect(skillPanel).toHaveClass('hidden');
|
||||
});
|
||||
|
||||
it('omits skills when all concrete skills are cleared', async () => {
|
||||
it('sends an explicit empty skills list when all concrete skills are cleared', async () => {
|
||||
render(
|
||||
<MemoryRouter initialEntries={['/chat']}>
|
||||
<ChatPage />
|
||||
@@ -945,8 +1040,10 @@ describe('ChatPage', () => {
|
||||
expect(mockStartStream).toHaveBeenCalled();
|
||||
});
|
||||
const lastCall = mockStartStream.mock.calls[mockStartStream.mock.calls.length - 1];
|
||||
expect(lastCall[0]).toEqual(expect.objectContaining({ message: '分析 AAPL' }));
|
||||
expect(lastCall[0]).not.toHaveProperty('skills');
|
||||
expect(lastCall[0]).toEqual(expect.objectContaining({
|
||||
message: '分析 AAPL',
|
||||
skills: [],
|
||||
}));
|
||||
expect(lastCall[1]).toEqual(expect.objectContaining({
|
||||
skillNames: ['通用'],
|
||||
skillName: '通用',
|
||||
|
||||
@@ -7,7 +7,11 @@ vi.mock('../../api/agent', async (importOriginal) => {
|
||||
...actual,
|
||||
agentApi: {
|
||||
getChatSessions: vi.fn(async () => []),
|
||||
getChatSessionMessages: vi.fn(async () => []),
|
||||
getChatSessionMessages: vi.fn(async (sessionId: string) => ({
|
||||
session_id: sessionId,
|
||||
messages: [],
|
||||
session_state: { selected_skill_ids: [] },
|
||||
})),
|
||||
chatStream: vi.fn(),
|
||||
cancelChatStream: vi.fn(),
|
||||
},
|
||||
@@ -59,6 +63,7 @@ beforeEach(() => {
|
||||
localStorage.clear();
|
||||
useAgentChatStore.setState({
|
||||
messages: [],
|
||||
selectedSkillIds: null,
|
||||
loading: false,
|
||||
progressSteps: [],
|
||||
sessionId: 'session-test',
|
||||
@@ -530,9 +535,13 @@ describe('agentChatStore.startStream', () => {
|
||||
describe('agentChatStore.switchSession', () => {
|
||||
it('clears transient loading state when switching sessions during a stream', async () => {
|
||||
const ac = new AbortController();
|
||||
vi.mocked(agentApi.getChatSessionMessages).mockResolvedValue([
|
||||
{ id: 'msg-2', role: 'assistant', content: '历史回复', created_at: null },
|
||||
]);
|
||||
vi.mocked(agentApi.getChatSessionMessages).mockResolvedValue({
|
||||
session_id: 'session-2',
|
||||
messages: [
|
||||
{ id: 'msg-2', role: 'assistant', content: '历史回复', created_at: null },
|
||||
],
|
||||
session_state: { selected_skill_ids: ['risk'] },
|
||||
});
|
||||
useAgentChatStore.setState({
|
||||
loading: true,
|
||||
progressSteps: [{ type: 'thinking', message: '正在制定分析路径...' }],
|
||||
@@ -557,28 +566,37 @@ describe('agentChatStore.switchSession', () => {
|
||||
expect(state.messages).toEqual([
|
||||
{ id: 'msg-2', role: 'assistant', content: '历史回复' },
|
||||
]);
|
||||
expect(state.selectedSkillIds).toEqual(['risk']);
|
||||
});
|
||||
|
||||
it('does not let a late session history response overwrite the current session', async () => {
|
||||
const sessionA = createDeferred<
|
||||
Array<{ id: string; role: 'user' | 'assistant'; content: string; created_at: string | null }>
|
||||
>();
|
||||
const sessionB = createDeferred<
|
||||
Array<{ id: string; role: 'user' | 'assistant'; content: string; created_at: string | null }>
|
||||
>();
|
||||
const sessionA = createDeferred<Awaited<ReturnType<typeof agentApi.getChatSessionMessages>>>();
|
||||
const sessionB = createDeferred<Awaited<ReturnType<typeof agentApi.getChatSessionMessages>>>();
|
||||
vi.mocked(agentApi.getChatSessionMessages).mockImplementation((targetSessionId: string) => {
|
||||
if (targetSessionId === 'session-a') return sessionA.promise;
|
||||
if (targetSessionId === 'session-b') return sessionB.promise;
|
||||
return Promise.resolve([]);
|
||||
return Promise.resolve({
|
||||
session_id: targetSessionId,
|
||||
messages: [],
|
||||
session_state: { selected_skill_ids: [] },
|
||||
});
|
||||
});
|
||||
|
||||
const switchToA = useAgentChatStore.getState().switchSession('session-a');
|
||||
const switchToB = useAgentChatStore.getState().switchSession('session-b');
|
||||
|
||||
sessionB.resolve([{ id: 'msg-b', role: 'assistant', content: 'B 回复', created_at: null }]);
|
||||
sessionB.resolve({
|
||||
session_id: 'session-b',
|
||||
messages: [{ id: 'msg-b', role: 'assistant', content: 'B 回复', created_at: null }],
|
||||
session_state: { selected_skill_ids: ['risk'] },
|
||||
});
|
||||
await switchToB;
|
||||
|
||||
sessionA.resolve([{ id: 'msg-a', role: 'assistant', content: 'A 回复', created_at: null }]);
|
||||
sessionA.resolve({
|
||||
session_id: 'session-a',
|
||||
messages: [{ id: 'msg-a', role: 'assistant', content: 'A 回复', created_at: null }],
|
||||
session_state: { selected_skill_ids: ['technical'] },
|
||||
});
|
||||
await switchToA;
|
||||
|
||||
const state = useAgentChatStore.getState();
|
||||
@@ -586,5 +604,69 @@ describe('agentChatStore.switchSession', () => {
|
||||
expect(state.messages).toEqual([
|
||||
{ id: 'msg-b', role: 'assistant', content: 'B 回复' },
|
||||
]);
|
||||
expect(state.selectedSkillIds).toEqual(['risk']);
|
||||
});
|
||||
});
|
||||
|
||||
describe('agentChatStore session Skill state', () => {
|
||||
it('restores the saved Skill selection during the initial session load', async () => {
|
||||
localStorage.setItem('dsa_chat_session_id', 'saved-session');
|
||||
useAgentChatStore.setState({ hasInitialLoad: false });
|
||||
vi.mocked(agentApi.getChatSessions).mockResolvedValue([
|
||||
{
|
||||
session_id: 'saved-session',
|
||||
title: 'saved',
|
||||
message_count: 1,
|
||||
created_at: null,
|
||||
last_active: null,
|
||||
},
|
||||
]);
|
||||
vi.mocked(agentApi.getChatSessionMessages).mockResolvedValue({
|
||||
session_id: 'saved-session',
|
||||
messages: [
|
||||
{ id: 'msg-1', role: 'user', content: '分析 AAPL', created_at: null },
|
||||
],
|
||||
session_state: { selected_skill_ids: ['technical', 'risk'] },
|
||||
});
|
||||
|
||||
await useAgentChatStore.getState().loadInitialSession();
|
||||
|
||||
expect(useAgentChatStore.getState().selectedSkillIds).toEqual([
|
||||
'technical',
|
||||
'risk',
|
||||
]);
|
||||
});
|
||||
|
||||
it('preserves null when an initial legacy session has no persisted Skill state', async () => {
|
||||
localStorage.setItem('dsa_chat_session_id', 'legacy-session');
|
||||
useAgentChatStore.setState({ hasInitialLoad: false });
|
||||
vi.mocked(agentApi.getChatSessions).mockResolvedValue([
|
||||
{
|
||||
session_id: 'legacy-session',
|
||||
title: 'legacy',
|
||||
message_count: 1,
|
||||
created_at: null,
|
||||
last_active: null,
|
||||
},
|
||||
]);
|
||||
vi.mocked(agentApi.getChatSessionMessages).mockResolvedValue({
|
||||
session_id: 'legacy-session',
|
||||
messages: [
|
||||
{ id: 'msg-1', role: 'user', content: '分析 AAPL', created_at: null },
|
||||
],
|
||||
session_state: { selected_skill_ids: null },
|
||||
});
|
||||
|
||||
await useAgentChatStore.getState().loadInitialSession();
|
||||
|
||||
expect(useAgentChatStore.getState().selectedSkillIds).toBeNull();
|
||||
});
|
||||
|
||||
it('clears the previous session Skill selection for a new chat', () => {
|
||||
useAgentChatStore.setState({ selectedSkillIds: ['risk'] });
|
||||
|
||||
useAgentChatStore.getState().startNewChat();
|
||||
|
||||
expect(useAgentChatStore.getState().selectedSkillIds).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -111,6 +111,7 @@ function getStreamFailureError(
|
||||
|
||||
interface AgentChatState {
|
||||
messages: Message[];
|
||||
selectedSkillIds: string[] | null;
|
||||
loading: boolean;
|
||||
progressSteps: ProgressStep[];
|
||||
sessionId: string;
|
||||
@@ -129,6 +130,7 @@ interface AgentChatState {
|
||||
}
|
||||
|
||||
interface AgentChatActions {
|
||||
setSelectedSkillIds: (skillIds: string[]) => void;
|
||||
setCurrentRoute: (path: string) => void;
|
||||
clearCompletionBadge: () => void;
|
||||
loadSessions: () => Promise<void>;
|
||||
@@ -158,6 +160,7 @@ export const useAgentChatStore = create<AgentChatState & AgentChatActions>((set,
|
||||
|
||||
return {
|
||||
messages: [],
|
||||
selectedSkillIds: null,
|
||||
loading: false,
|
||||
progressSteps: [],
|
||||
sessionId: getInitialSessionId(),
|
||||
@@ -174,6 +177,8 @@ export const useAgentChatStore = create<AgentChatState & AgentChatActions>((set,
|
||||
terminalStatus: null,
|
||||
stopError: false,
|
||||
|
||||
setSelectedSkillIds: (skillIds) => set({ selectedSkillIds: skillIds }),
|
||||
|
||||
setCurrentRoute: (path) => set({ currentRoute: path }),
|
||||
|
||||
clearCompletionBadge: () => set({ completionBadge: false }),
|
||||
@@ -203,19 +208,20 @@ export const useAgentChatStore = create<AgentChatState & AgentChatActions>((set,
|
||||
if (savedId) {
|
||||
const sessionExists = sessionList.some((s) => s.session_id === savedId);
|
||||
if (sessionExists) {
|
||||
const msgs = await agentApi.getChatSessionMessages(savedId);
|
||||
if (msgs.length > 0) {
|
||||
const detail = await agentApi.getChatSessionMessages(savedId);
|
||||
if (detail.messages.length > 0) {
|
||||
set({
|
||||
messages: msgs.map((m) => ({
|
||||
messages: detail.messages.map((m) => ({
|
||||
id: m.id,
|
||||
role: m.role,
|
||||
content: m.content,
|
||||
})),
|
||||
selectedSkillIds: detail.session_state.selected_skill_ids,
|
||||
});
|
||||
}
|
||||
} else {
|
||||
const newId = generateUUID();
|
||||
set({ sessionId: newId });
|
||||
set({ sessionId: newId, selectedSkillIds: null });
|
||||
localStorage.setItem(STORAGE_KEY_SESSION, newId);
|
||||
}
|
||||
} else {
|
||||
@@ -235,6 +241,7 @@ export const useAgentChatStore = create<AgentChatState & AgentChatActions>((set,
|
||||
abortController?.abort();
|
||||
set({
|
||||
messages: [],
|
||||
selectedSkillIds: null,
|
||||
sessionId: targetSessionId,
|
||||
loading: false,
|
||||
progressSteps: [],
|
||||
@@ -249,16 +256,17 @@ export const useAgentChatStore = create<AgentChatState & AgentChatActions>((set,
|
||||
localStorage.setItem(STORAGE_KEY_SESSION, targetSessionId);
|
||||
|
||||
try {
|
||||
const msgs = await agentApi.getChatSessionMessages(targetSessionId);
|
||||
const detail = await agentApi.getChatSessionMessages(targetSessionId);
|
||||
if (get().sessionId !== targetSessionId) {
|
||||
return;
|
||||
}
|
||||
set({
|
||||
messages: msgs.map((m) => ({
|
||||
messages: detail.messages.map((m) => ({
|
||||
id: m.id,
|
||||
role: m.role,
|
||||
content: m.content,
|
||||
})),
|
||||
selectedSkillIds: detail.session_state.selected_skill_ids,
|
||||
});
|
||||
} catch {
|
||||
// Ignore
|
||||
@@ -272,6 +280,7 @@ export const useAgentChatStore = create<AgentChatState & AgentChatActions>((set,
|
||||
set({
|
||||
sessionId: newId,
|
||||
messages: [],
|
||||
selectedSkillIds: null,
|
||||
loading: false,
|
||||
progressSteps: [],
|
||||
chatError: null,
|
||||
|
||||
@@ -9,6 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/).
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
- [新功能] Agent Chat 按会话持久化 Skill 选择,支持刷新和会话切换恢复,并区分省略 `skills`、显式空列表与非空选择;无持久化状态的历史会话继续使用运行时默认且不会被静默转为显式选择,复用分析 `context` 中残留的 legacy `skills` / `strategies` 也不会覆盖顶层三态或会话状态,非空但全部无效的 Skill 请求不会被当成显式空列表并清空既有选择
|
||||
- [改进] 后端 CI 在不跳过离线测试的前提下按完整测试文件分成三个独立 runner 并行执行,由单一 `backend-gate` 汇总门禁结果;实测文件耗时和首分片静态检查成本共同参与负载平衡,新测试文件自动纳入,现有 pip 安装和测试参数保持不变,避免 xdist 进程内并发的全局状态竞态。
|
||||
- [测试] 后端 CI 默认覆盖所有非 Web 改动,仅对已证明安全的纯 Web 路径跳过,并将整个 Web public 目录及前端渠道模板、设置帮助视为跨层运行合同;补充纯 Web、共享 Web 资产及 Web/非 Web 混合改动的过滤语义回归,明确 `predicate-quantifier: every` 按单文件匹配全部规则、再以任一匹配文件触发门禁。Docker CI 继续按构建输入过滤。离线测试保留稳定的串行执行与慢用例摘要,并移除重复用例和测试内真实等待。
|
||||
|
||||
|
||||
@@ -1578,6 +1578,7 @@ FastAPI 提供 RESTful API 服务,支持配置管理和触发分析。
|
||||
- 🧭 **市场位置卡片** - A 股普通分析报告会展示市场题材层和个股位置层,区分大盘主线、主关联题材、题材阶段、个股位置和缺失证据
|
||||
- 🧩 **输入数据块可见** - 普通分析报告会在历史详情、同步响应和 completed 任务状态中返回低敏 `AnalysisContextPack` overview,Web 报告页在策略点位和资讯之后默认折叠展示数据块状态、来源、缺失原因和降级摘要
|
||||
- 💬 **问股追问上下文** - 从历史报告进入问股后,后续追问会持续携带当前 `stock_code/stock_name`;切回或重载已有问股会话时,会从已加载的历史用户消息恢复基础当前标的;只有用户明确切换标的时才切换上下文,含比较/对比/vs/差异/相比等明确比较意图或多个非当前明确股票代码的问题不会污染当前标的
|
||||
- 🧠 **会话级 Skill 选择** - 问股会话会持久化当前 Skill 选择;刷新或切换会话时恢复各自选择,新会话继续使用服务器默认 Skill
|
||||
- 📈 **回测验证** - 评估历史分析准确率,查询方向胜率与模拟收益
|
||||
- 🔗 **API 文档** - 访问 `/docs` 查看 Swagger UI
|
||||
|
||||
@@ -1649,6 +1650,8 @@ FastAPI 提供 RESTful API 服务,支持配置管理和触发分析。
|
||||
> 说明(Issue #1520):列表中的模型名展示字段仅来源于历史快照中的 `model_used`,仅用于历史回溯展示,不影响运行时模型模型路由(`litellm_model`、`llm_model_list`)、Provider、Base URL 与配置迁移/清理语义。回退方式为回退本次提交,现网历史查询/抽屉/接口链路兼容性保持不变。
|
||||
> 说明:历史详情、同步分析响应和 completed 任务状态会在 `report.details.analysis_context_pack_overview` 返回低敏输入数据块 overview;其中同步分析响应依赖本次已持久化的 `analysis_history.context_snapshot`,`SAVE_CONTEXT_SNAPSHOT=false` 时新记录不保证返回 overview。`details.context_snapshot` 会剥离该顶层字段,不返回完整 `AnalysisContextPack` 或 Prompt summary。
|
||||
> 说明:`POST /api/v1/agent/chat` 与 `POST /api/v1/agent/chat/stream` 会把前端传入的 `context.stock_code` 作为问股当前标的基线,并在 `context.report_language` 缺失时使用全局 `REPORT_LANGUAGE`;调用方显式提供的 `context.report_language` 保持优先。服务端会先重新判定 stock scope。前端从历史报告进入问股后会持续发送 active stock context;切回或重载已有会话时,会根据已加载的历史用户消息恢复基础 `{stock_code, stock_name: null}`。服务端会在每轮消息中重新判定 `maintain` / `switch` / `compare`:未明确切换时,带 `stock_code` 的股票工具调用只能访问当前标的;显式切换会清理旧标的历史摘要和预取数据;含比较/对比/vs/差异/相比等明确比较意图或多个非当前明确股票代码的问题允许本轮明确出现的多个代码,但不改写当前标的。若模型误把 TTM、PE、MACD、KDJ 等金融缩写、移动均线语境下的 `MA` 指标词,或 SH/SZ/BJ/HK/SS 等交易所片段当成股票代码调用工具,后端会返回不可重试的 `stock_scope_violation` 工具结果,而不会执行对应股票工具。工具名只解析注册表中的精确名称;任何 provider namespace 或 suffix 都不会路由到已有工具。
|
||||
|
||||
> Skill 会话状态:上述两个 Chat 请求中的顶层 `skills` 为三态字段,也是请求 Skill 选择的唯一权威来源。省略或传 `null` 时沿用该 `session_id` 已保存的选择;会话尚无状态时使用服务器运行时默认 Skill。传 `[]` 表示清空显式选择并使用现有通用/服务器默认执行语义;传非空列表时会按现有 Skill catalog 规则清理、去重并保存。非空列表中有效项与无效项混合时保留有效项;若全部条目均无效,则不会把归一化后的空结果视为显式 `[]`,而是沿用会话状态或运行时默认且不写入空状态。用于复用分析数据的 `context` 中即使残留 `skills` 或 `strategies`,服务端也会移除这些 legacy 字段,不能覆盖顶层三态或会话状态。用户消息与本轮显式 Skill 更新在同一事务中写入,成功后流式接口才发送 `accepted`。`GET /api/v1/agent/chat/sessions/{session_id}` 会在消息之外返回 `session_state.selected_skill_ids`:没有持久化状态时为 `null`,显式清空时为 `[]`,否则为保存的 Skill 列表。Web 端使用该字段恢复持久化选择;对于 `null`,页面可以显示服务器默认 Skill,但未操作直接追问时仍省略 `skills`,不会把历史会话静默转换为显式 Skill 会话。删除会话会同时删除对应状态。
|
||||
> 说明:`POST /api/v1/backtest/run` 新增 `analysis_date_from` / `analysis_date_to`(`YYYY-MM-DD`)请求参数用于按历史分析日期筛选候选;若 `analysis_date_from > analysis_date_to`,接口返回 400 `invalid_params`。
|
||||
> 说明:回测执行成功但无新入库结果时,`BacktestRunResponse.message` 返回可读诊断说明,`diagnostics` 返回排查上下文(示例:`empty_reason`、`analysis_date_from`、`analysis_date_to`、`eval_window_days`、`min_age_days`、`limit`)。
|
||||
> 说明:`GET /api/v1/backtest/results`、`GET /api/v1/backtest/performance`、`GET /api/v1/backtest/performance/{code}` 同步支持 `analysis_date_from`、`analysis_date_to`;不传时保持历史行为。
|
||||
|
||||
@@ -1407,6 +1407,7 @@ FastAPI provides RESTful API service for configuration management and triggering
|
||||
- **UI Language Switch** - Toggle UI language (`zh`/`en`) on login page, shell/navigation, settings page, and shared controls; this switch is independent of `REPORT_LANGUAGE`.
|
||||
- **Quick Analysis** - Trigger stock analysis via API; the Home page also provides a one-run market selector next to Market Review, so Docker/server mode can use the server default or a temporary single/multi-market scope
|
||||
- **Strategy selection** - The Home page supports explicitly selecting analysis strategy skills; when `skills` is omitted, analysis uses the server default strategy so legacy clients keep existing behavior
|
||||
- **Session-level Skill selection** - Ask Stock sessions persist their current Skill selection; refresh and session switching restore each session independently, while new sessions keep the server default
|
||||
- **Today-state refresh safety** - Today and watchlist status loading uses history lookups with explicit timezone-aware date filtering and full pagination; a successful refresh from a newer stock-bar request is required to clear an unknown state, so stale in-flight responses cannot override completion refresh results
|
||||
- **First-run Setup Hint** - The Home page reads the read-only setup status and points users to Settings when required items such as the primary LLM channel or watchlist are missing
|
||||
- **Real-time Progress** - Analysis task status updates in real-time, supports parallel tasks; the regular stock-analysis path now prefers LiteLLM streaming during the LLM stage and pushes finer-grained `message/progress` updates through task SSE
|
||||
@@ -1484,6 +1485,8 @@ For this feature, the product behavior is:
|
||||
> Issue #1520 compatibility note: The `model`/`model_used` returned here is read-only historical snapshot metadata from each record, used only for trend drawer/history display. It does not alter runtime model/model-provider/base URL resolution, config migration, or cleanup semantics in the analysis path. Rollback is by reverting this commit; history query, API response shapes, and UI drawer consumption remain compatible.
|
||||
> Note: history detail, sync analysis responses, and completed task status responses expose a low-sensitivity input data-block overview at `report.details.analysis_context_pack_overview`; sync analysis responses depend on the just-persisted `analysis_history.context_snapshot`, so new records do not guarantee the overview when `SAVE_CONTEXT_SNAPSHOT=false`. `details.context_snapshot` strips that top-level field and does not return the full `AnalysisContextPack` or prompt summary.
|
||||
> Note: `POST /api/v1/agent/chat` and `POST /api/v1/agent/chat/stream` use the frontend-provided `context.stock_code` as the active Ask Stock baseline and fall back to global `REPORT_LANGUAGE` when `context.report_language` is absent; an explicitly supplied `context.report_language` keeps precedence. Stock scope is still resolved server-side. Each turn is classified as `maintain`, `switch`, or `compare`: unchanged follow-ups can call stock-scoped tools only for the current stock; explicit switches clear stale stock summaries and prefetched context; comparison prompts such as compare/vs/difference allow the explicitly mentioned codes for that turn without rewriting the current stock. If a model attempts to call a stock tool with financial abbreviations such as TTM, PE, MACD, KDJ, contextual indicator tokens such as `MA` in moving-average prompts, or exchange fragments such as SH/SZ/BJ/HK/SS, the backend returns a non-retriable `stock_scope_violation` tool result instead of executing that stock tool. Tool names are resolved only by exact registry name; provider namespaces or suffixes are not routed to existing tools.
|
||||
|
||||
> Session Skill state: the top-level `skills` field is tri-state on both Chat requests and is the only authoritative request source for Skill selection. Omitting it or sending `null` inherits the selection saved for that `session_id`; sessions without state use the runtime server default. Sending `[]` clears the explicit selection and preserves the existing general/server-default execution semantics. A non-empty list is cleaned, deduplicated, and saved with the existing Skill catalog rules. Mixed valid and invalid entries keep the valid entries; if every entry is invalid, the normalized empty result is not treated as an explicit `[]`, so the request inherits session state or the runtime default without persisting an empty state. Any legacy `skills` or `strategies` fields left in the analysis-reuse `context` are removed by the server and cannot override the top-level tri-state or session state. The user message and an explicit Skill update are written in one transaction before the streaming endpoint emits `accepted`. `GET /api/v1/agent/chat/sessions/{session_id}` returns `session_state.selected_skill_ids` alongside messages: `null` means no state has been persisted, `[]` means the selection was explicitly cleared, and a non-empty list is the saved selection. The Web client restores only persisted selections from this field. For `null`, it may display the server default Skill, but an untouched follow-up still omits `skills` so a legacy session is not silently converted into an explicit-Skill session. Deleting a session also deletes its state.
|
||||
> Note: `POST /api/v1/backtest/run` adds `analysis_date_from` / `analysis_date_to` (`YYYY-MM-DD`) to filter candidates by analysis date range. When `analysis_date_from > analysis_date_to`, it returns 400 `invalid_params`.
|
||||
> Note: When backtest runs successfully but yields no new persisted rows, `BacktestRunResponse.message` carries a readable diagnostic and `diagnostics` returns troubleshooting context (for example `empty_reason`, `analysis_date_from`, `analysis_date_to`, `eval_window_days`, `min_age_days`, `limit`).
|
||||
> Note: `GET /api/v1/backtest/results`, `GET /api/v1/backtest/performance`, and `GET /api/v1/backtest/performance/{code}` all support `analysis_date_from` and `analysis_date_to` consistently. Omitting them keeps historical default behavior.
|
||||
|
||||
@@ -5,7 +5,7 @@ from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
from src.agent.agent_backend import AgentBackend, AgentRunRequest
|
||||
from src.agent.conversation import conversation_manager
|
||||
@@ -56,11 +56,13 @@ class AgentChatExecutor:
|
||||
progress_callback: Optional[Callable[[Dict[str, Any]], None]] = None,
|
||||
context: Optional[Dict[str, Any]] = None,
|
||||
cancel_event=None,
|
||||
selected_skill_ids: Optional[List[str]] = None,
|
||||
) -> AgentResult:
|
||||
turn = self.prepare_turn(
|
||||
message=message,
|
||||
session_id=session_id,
|
||||
context=context,
|
||||
selected_skill_ids=selected_skill_ids,
|
||||
)
|
||||
return self.execute_turn(
|
||||
turn,
|
||||
@@ -74,6 +76,7 @@ class AgentChatExecutor:
|
||||
message: str,
|
||||
session_id: str,
|
||||
context: Optional[Dict[str, Any]] = None,
|
||||
selected_skill_ids: Optional[List[str]] = None,
|
||||
) -> PreparedAgentChatTurn:
|
||||
"""Prepare context and persist the user message without starting a backend."""
|
||||
conversation_manager.get_or_create(session_id)
|
||||
@@ -92,7 +95,11 @@ class AgentChatExecutor:
|
||||
)
|
||||
baseline_len = len(prepared.history_messages) + 2
|
||||
run_id = str(uuid.uuid4())
|
||||
user_message_id = conversation_manager.add_message(session_id, "user", message)
|
||||
user_message_id = conversation_manager.add_user_message(
|
||||
session_id,
|
||||
message,
|
||||
selected_skill_ids,
|
||||
)
|
||||
return PreparedAgentChatTurn(
|
||||
message=message,
|
||||
session_id=session_id,
|
||||
|
||||
@@ -29,6 +29,20 @@ class ConversationSession:
|
||||
self.last_active = datetime.now()
|
||||
return message_id
|
||||
|
||||
def add_user_message(
|
||||
self,
|
||||
content: str,
|
||||
selected_skill_ids: Optional[List[str]] = None,
|
||||
) -> int:
|
||||
"""Add a user message and optionally update the persisted Skill selection."""
|
||||
message_id = get_db().save_conversation_user_turn(
|
||||
self.session_id,
|
||||
content,
|
||||
selected_skill_ids,
|
||||
)
|
||||
self.last_active = datetime.now()
|
||||
return message_id
|
||||
|
||||
def update_context(self, key: str, value: Any):
|
||||
"""Update session context."""
|
||||
self.context[key] = value
|
||||
@@ -66,6 +80,16 @@ class ConversationManager:
|
||||
session = self.get_or_create(session_id)
|
||||
return session.add_message(role, content)
|
||||
|
||||
def add_user_message(
|
||||
self,
|
||||
session_id: str,
|
||||
content: str,
|
||||
selected_skill_ids: Optional[List[str]] = None,
|
||||
) -> int:
|
||||
"""Add a user message through the session-state transaction boundary."""
|
||||
session = self.get_or_create(session_id)
|
||||
return session.add_user_message(content, selected_skill_ids)
|
||||
|
||||
def get_history(self, session_id: str) -> List[Dict[str, Any]]:
|
||||
"""Get message history for a session."""
|
||||
session = self.get_or_create(session_id)
|
||||
|
||||
@@ -111,6 +111,23 @@ def _normalize_skill_ids(
|
||||
return normalized, unknown
|
||||
|
||||
|
||||
def normalize_requested_skill_ids(config, skill_ids: List[str]) -> List[str]:
|
||||
"""Normalize API-requested Skill ids with the AgentFactory catalog rules."""
|
||||
skill_manager = get_skill_manager(config)
|
||||
available_skill_ids = {
|
||||
str(getattr(skill, "name", "")).strip()
|
||||
for skill in skill_manager.list_skills()
|
||||
if str(getattr(skill, "name", "")).strip()
|
||||
}
|
||||
normalized, unknown = _normalize_skill_ids(
|
||||
skill_ids,
|
||||
available_skill_ids=available_skill_ids,
|
||||
)
|
||||
if unknown:
|
||||
logger.warning("[AgentFactory] Ignoring unknown request skill ids: %s", unknown)
|
||||
return normalized
|
||||
|
||||
|
||||
def _resolve_selected_skill_ids(
|
||||
*,
|
||||
requested_skills: Optional[List[str]],
|
||||
|
||||
@@ -379,6 +379,7 @@ class AgentOrchestrator:
|
||||
session_id: str,
|
||||
progress_callback: Optional[Callable] = None,
|
||||
context: Optional[Dict[str, Any]] = None,
|
||||
selected_skill_ids: Optional[List[str]] = None,
|
||||
) -> "AgentResult":
|
||||
"""Run the pipeline in chat mode (free-form answer, no dashboard parse).
|
||||
|
||||
@@ -390,6 +391,7 @@ class AgentOrchestrator:
|
||||
message=message,
|
||||
session_id=session_id,
|
||||
context=context,
|
||||
selected_skill_ids=selected_skill_ids,
|
||||
)
|
||||
return self.execute_turn(
|
||||
turn,
|
||||
@@ -402,6 +404,7 @@ class AgentOrchestrator:
|
||||
message: str,
|
||||
session_id: str,
|
||||
context: Optional[Dict[str, Any]] = None,
|
||||
selected_skill_ids: Optional[List[str]] = None,
|
||||
) -> PreparedOrchestratorChatTurn:
|
||||
"""Prepare context and persist the user turn before SSE acceptance."""
|
||||
from src.agent.conversation import conversation_manager
|
||||
@@ -420,7 +423,11 @@ class AgentOrchestrator:
|
||||
ctx.meta["conversation_history"] = history
|
||||
|
||||
# Persist user turn
|
||||
conversation_manager.add_message(session_id, "user", message)
|
||||
conversation_manager.add_user_message(
|
||||
session_id,
|
||||
message,
|
||||
selected_skill_ids,
|
||||
)
|
||||
|
||||
return PreparedOrchestratorChatTurn(
|
||||
session_id=session_id,
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Agent Chat session state service."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from src.agent.factory import normalize_requested_skill_ids
|
||||
from src.storage import DatabaseManager
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ChatSkillSelection:
|
||||
"""Effective Skill ids and the optional state update for one chat turn."""
|
||||
|
||||
effective_skill_ids: Optional[List[str]]
|
||||
selected_skill_ids_update: Optional[List[str]]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ChatSessionDetail:
|
||||
"""Visible messages and the persisted Skill selection for one session."""
|
||||
|
||||
messages: List[Dict[str, Any]]
|
||||
selected_skill_ids: Optional[List[str]]
|
||||
|
||||
|
||||
class AgentChatSessionService:
|
||||
"""Coordinate Agent Chat session state without exposing storage to HTTP handlers."""
|
||||
|
||||
def __init__(self, db_manager: Optional[DatabaseManager] = None):
|
||||
self.db = db_manager or DatabaseManager.get_instance()
|
||||
|
||||
def resolve_skill_selection(
|
||||
self,
|
||||
config,
|
||||
session_id: str,
|
||||
requested_skill_ids: Optional[List[str]],
|
||||
) -> ChatSkillSelection:
|
||||
if requested_skill_ids is None:
|
||||
return ChatSkillSelection(
|
||||
effective_skill_ids=(
|
||||
self.db.get_conversation_session_selected_skill_ids(session_id)
|
||||
),
|
||||
selected_skill_ids_update=None,
|
||||
)
|
||||
if not requested_skill_ids:
|
||||
return ChatSkillSelection(
|
||||
effective_skill_ids=[],
|
||||
selected_skill_ids_update=[],
|
||||
)
|
||||
|
||||
normalized = normalize_requested_skill_ids(config, requested_skill_ids)
|
||||
if not normalized:
|
||||
return ChatSkillSelection(
|
||||
effective_skill_ids=(
|
||||
self.db.get_conversation_session_selected_skill_ids(session_id)
|
||||
),
|
||||
selected_skill_ids_update=None,
|
||||
)
|
||||
return ChatSkillSelection(
|
||||
effective_skill_ids=normalized,
|
||||
selected_skill_ids_update=normalized,
|
||||
)
|
||||
|
||||
def list_sessions(
|
||||
self,
|
||||
limit: int,
|
||||
user_id: Optional[str],
|
||||
) -> List[Dict[str, Any]]:
|
||||
return self.db.get_chat_sessions(
|
||||
limit=limit,
|
||||
session_prefix=user_id,
|
||||
extra_session_ids=[user_id] if user_id else None,
|
||||
)
|
||||
|
||||
def get_session_detail(
|
||||
self,
|
||||
session_id: str,
|
||||
limit: int,
|
||||
) -> ChatSessionDetail:
|
||||
messages = self.db.get_conversation_messages(session_id, limit=limit)
|
||||
selected_skill_ids = self.db.get_conversation_session_selected_skill_ids(session_id)
|
||||
|
||||
return ChatSessionDetail(
|
||||
messages=messages,
|
||||
selected_skill_ids=selected_skill_ids,
|
||||
)
|
||||
|
||||
def delete_session(self, session_id: str) -> int:
|
||||
return self.db.delete_conversation_session(session_id)
|
||||
@@ -725,6 +725,17 @@ class ConversationMessage(Base):
|
||||
created_at = Column(DateTime, default=datetime.now, index=True)
|
||||
|
||||
|
||||
class ConversationSessionState(Base):
|
||||
"""Persisted user selections for an Agent chat session."""
|
||||
|
||||
__tablename__ = 'conversation_session_states'
|
||||
|
||||
session_id = Column(String(100), primary_key=True)
|
||||
selected_skill_ids_json = Column(Text, nullable=False)
|
||||
created_at = Column(DateTime, default=datetime.now, nullable=False)
|
||||
updated_at = Column(DateTime, default=datetime.now, onupdate=datetime.now, nullable=False)
|
||||
|
||||
|
||||
class ConversationSummary(Base):
|
||||
"""Rolling summary for visible Agent chat history."""
|
||||
|
||||
@@ -3251,6 +3262,54 @@ class DatabaseManager(metaclass=_DatabaseManagerMeta):
|
||||
session.flush()
|
||||
return int(msg.id)
|
||||
|
||||
def save_conversation_user_turn(
|
||||
self,
|
||||
session_id: str,
|
||||
content: str,
|
||||
selected_skill_ids: Optional[List[str]] = None,
|
||||
) -> int:
|
||||
"""Persist a user message and an optional session Skill selection atomically."""
|
||||
with self.session_scope() as session:
|
||||
msg = ConversationMessage(
|
||||
session_id=session_id,
|
||||
role="user",
|
||||
content=content,
|
||||
)
|
||||
session.add(msg)
|
||||
session.flush()
|
||||
|
||||
if selected_skill_ids is not None:
|
||||
now = datetime.now()
|
||||
values = {
|
||||
"session_id": session_id,
|
||||
"selected_skill_ids_json": json.dumps(selected_skill_ids, ensure_ascii=False),
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
stmt = sqlite_insert(ConversationSessionState).values(**values)
|
||||
session.execute(
|
||||
stmt.on_conflict_do_update(
|
||||
index_elements=["session_id"],
|
||||
set_={
|
||||
"selected_skill_ids_json": values["selected_skill_ids_json"],
|
||||
"updated_at": now,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
return int(msg.id)
|
||||
|
||||
def get_conversation_session_selected_skill_ids(
|
||||
self,
|
||||
session_id: str,
|
||||
) -> Optional[List[str]]:
|
||||
"""Return the saved Skill selection, or None when the session has no state row."""
|
||||
with self.session_scope() as session:
|
||||
state = session.get(ConversationSessionState, session_id)
|
||||
if state is None:
|
||||
return None
|
||||
return json.loads(state.selected_skill_ids_json)
|
||||
|
||||
def get_conversation_history(self, session_id: str, limit: int = 20) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
获取 Agent 对话历史
|
||||
@@ -3591,6 +3650,11 @@ class DatabaseManager(metaclass=_DatabaseManagerMeta):
|
||||
删除的消息数
|
||||
"""
|
||||
with self.session_scope() as session:
|
||||
session.execute(
|
||||
delete(ConversationSessionState).where(
|
||||
ConversationSessionState.session_id == session_id
|
||||
)
|
||||
)
|
||||
session.execute(
|
||||
delete(AgentProviderTurn).where(
|
||||
AgentProviderTurn.session_id == session_id
|
||||
|
||||
@@ -14,6 +14,7 @@ from fastapi.testclient import TestClient
|
||||
from api.app import create_app
|
||||
from api.v1.endpoints import agent as agent_endpoint
|
||||
from src.config import Config
|
||||
from src.services.agent_chat_session_service import AgentChatSessionService
|
||||
from src.storage import DatabaseManager
|
||||
|
||||
|
||||
@@ -75,7 +76,10 @@ def _sse_events(text: str) -> list[dict]:
|
||||
|
||||
|
||||
async def _collect_stream_events(request: "agent_endpoint.ChatRequest") -> list[dict]:
|
||||
response = await agent_endpoint.agent_chat_stream(request)
|
||||
response = await agent_endpoint.agent_chat_stream(
|
||||
request,
|
||||
session_service=AgentChatSessionService(),
|
||||
)
|
||||
return [
|
||||
json.loads(chunk.removeprefix("data: ").strip())
|
||||
async for chunk in response.body_iterator
|
||||
@@ -89,7 +93,11 @@ async def _immediate_to_thread(func, /, *args, **kwargs):
|
||||
def test_chat_session_messages_api_does_not_expose_provider_trace(tmp_path: Path) -> None:
|
||||
db = DatabaseManager(db_url=f"sqlite:///{tmp_path / 'trace.db'}")
|
||||
session_id = "api-trace-hidden"
|
||||
user_id = db.save_conversation_message(session_id, "user", "visible question")
|
||||
user_id = db.save_conversation_user_turn(
|
||||
session_id,
|
||||
"visible question",
|
||||
["technical"],
|
||||
)
|
||||
assistant_id = db.save_conversation_message(session_id, "assistant", "visible answer")
|
||||
db.save_agent_provider_turn(
|
||||
session_id=session_id,
|
||||
@@ -124,6 +132,9 @@ def test_chat_session_messages_api_does_not_expose_provider_trace(tmp_path: Path
|
||||
("user", "visible question"),
|
||||
("assistant", "visible answer"),
|
||||
]
|
||||
assert response.json()["session_state"] == {
|
||||
"selected_skill_ids": ["technical"],
|
||||
}
|
||||
assert "SECRET_REASONING" not in response.text
|
||||
assert "SECRET_TOOL_RESULT" not in response.text
|
||||
|
||||
@@ -243,6 +254,103 @@ def test_build_agent_chat_context_normalizes_default_report_language(
|
||||
assert context["report_language"] == expected_language
|
||||
|
||||
|
||||
def test_requested_skill_normalization_reuses_agent_factory_catalog_rules() -> None:
|
||||
from src.agent.factory import normalize_requested_skill_ids
|
||||
|
||||
skill_manager = MagicMock()
|
||||
skill_manager.list_skills.return_value = [
|
||||
SimpleNamespace(name="technical"),
|
||||
SimpleNamespace(name="risk"),
|
||||
]
|
||||
|
||||
with patch("src.agent.factory.get_skill_manager", return_value=skill_manager):
|
||||
normalized = normalize_requested_skill_ids(
|
||||
_litellm_config(),
|
||||
[" technical ", "technical", "unknown", "risk"],
|
||||
)
|
||||
|
||||
assert normalized == ["technical", "risk"]
|
||||
|
||||
|
||||
def test_agent_chat_inherits_saved_skills_without_rewriting_session_state(tmp_path: Path) -> None:
|
||||
db = DatabaseManager(db_url=f"sqlite:///{tmp_path / 'inherit.db'}")
|
||||
db.save_conversation_user_turn("saved-session", "first", ["technical"])
|
||||
config = _litellm_config()
|
||||
executor = MagicMock()
|
||||
executor.chat.return_value = _result()
|
||||
|
||||
with patch("api.middlewares.auth.is_auth_enabled", return_value=False), \
|
||||
patch("api.v1.endpoints.agent.get_config", return_value=config), \
|
||||
patch("api.v1.endpoints.agent._build_executor", return_value=executor) as build_executor:
|
||||
response = TestClient(create_app(static_dir=tmp_path / "static")).post(
|
||||
"/api/v1/agent/chat",
|
||||
json={
|
||||
"message": "follow up",
|
||||
"session_id": "saved-session",
|
||||
"context": {
|
||||
"stock_code": "600519",
|
||||
"skills": ["old_skill"],
|
||||
"strategies": ["older_strategy"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
build_executor.assert_called_once_with(config, ["technical"])
|
||||
context = executor.chat.call_args.kwargs["context"]
|
||||
assert context["stock_code"] == "600519"
|
||||
assert context["skills"] == ["technical"]
|
||||
assert "strategies" not in context
|
||||
assert executor.chat.call_args.kwargs["selected_skill_ids"] is None
|
||||
|
||||
|
||||
def test_agent_chat_all_invalid_skills_inherit_without_clearing_state(tmp_path: Path) -> None:
|
||||
db = DatabaseManager(db_url=f"sqlite:///{tmp_path / 'all-invalid.db'}")
|
||||
db.save_conversation_user_turn("saved-session", "first", ["technical"])
|
||||
config = _litellm_config()
|
||||
executor = MagicMock()
|
||||
executor.chat.return_value = _result()
|
||||
skill_manager = MagicMock()
|
||||
skill_manager.list_skills.return_value = [SimpleNamespace(name="technical")]
|
||||
|
||||
with patch("api.middlewares.auth.is_auth_enabled", return_value=False), \
|
||||
patch("api.v1.endpoints.agent.get_config", return_value=config), \
|
||||
patch("src.agent.factory.get_skill_manager", return_value=skill_manager), \
|
||||
patch("api.v1.endpoints.agent._build_executor", return_value=executor) as build_executor:
|
||||
response = TestClient(create_app(static_dir=tmp_path / "static")).post(
|
||||
"/api/v1/agent/chat",
|
||||
json={
|
||||
"message": "follow up",
|
||||
"session_id": "saved-session",
|
||||
"skills": ["old_technical"],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
build_executor.assert_called_once_with(config, ["technical"])
|
||||
assert executor.chat.call_args.kwargs["context"]["skills"] == ["technical"]
|
||||
assert executor.chat.call_args.kwargs["selected_skill_ids"] is None
|
||||
assert db.get_conversation_session_selected_skill_ids("saved-session") == [
|
||||
"technical"
|
||||
]
|
||||
|
||||
|
||||
def test_chat_session_messages_returns_null_when_state_is_missing(tmp_path: Path) -> None:
|
||||
db = DatabaseManager(db_url=f"sqlite:///{tmp_path / 'default-state.db'}")
|
||||
db.save_conversation_message("legacy-session", "user", "legacy question")
|
||||
|
||||
with patch("api.middlewares.auth.is_auth_enabled", return_value=False), \
|
||||
patch("api.v1.endpoints.agent.get_config", return_value=_litellm_config()):
|
||||
response = TestClient(create_app(static_dir=tmp_path / "static")).get(
|
||||
"/api/v1/agent/chat/sessions/legacy-session"
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["session_state"] == {
|
||||
"selected_skill_ids": None,
|
||||
}
|
||||
|
||||
|
||||
def test_codex_agent_chat_rejects_non_streaming_entrypoint(tmp_path: Path) -> None:
|
||||
with patch("api.middlewares.auth.is_auth_enabled", return_value=False), \
|
||||
patch("api.v1.endpoints.agent.get_config", return_value=_codex_config()), \
|
||||
@@ -329,7 +437,8 @@ def test_stream_prepares_and_persists_before_accepted_then_starts_backend() -> N
|
||||
session_id="accepted-session",
|
||||
request_id="accepted-request",
|
||||
context={"stock_code": "AAPL"},
|
||||
)
|
||||
),
|
||||
session_service=AgentChatSessionService(),
|
||||
)
|
||||
iterator = response.body_iterator
|
||||
first = json.loads((await anext(iterator)).removeprefix("data: ").strip())
|
||||
@@ -337,6 +446,7 @@ def test_stream_prepares_and_persists_before_accepted_then_starts_backend() -> N
|
||||
message="分析 AAPL",
|
||||
session_id="accepted-session",
|
||||
context={"stock_code": "AAPL", "report_language": "zh"},
|
||||
selected_skill_ids=None,
|
||||
)
|
||||
executor.execute_turn.assert_not_called()
|
||||
await iterator.aclose()
|
||||
@@ -352,6 +462,107 @@ def test_stream_prepares_and_persists_before_accepted_then_starts_backend() -> N
|
||||
executor.execute_turn.assert_not_called()
|
||||
|
||||
|
||||
def test_stream_forwards_normalized_skill_selection_to_prepare_turn() -> None:
|
||||
executor = _executor(_result(backend="litellm"))
|
||||
config = _litellm_config()
|
||||
|
||||
with patch("api.v1.endpoints.agent.asyncio.to_thread", side_effect=_immediate_to_thread), \
|
||||
patch("api.v1.endpoints.agent.get_config", return_value=config), \
|
||||
patch(
|
||||
"src.services.agent_chat_session_service.normalize_requested_skill_ids",
|
||||
return_value=["risk"],
|
||||
), \
|
||||
patch("api.v1.endpoints.agent._build_executor", return_value=executor) as build_executor:
|
||||
events = asyncio.run(
|
||||
_collect_stream_events(
|
||||
agent_endpoint.ChatRequest(
|
||||
message="check risk",
|
||||
session_id="risk-session",
|
||||
skills=[" risk ", "risk"],
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
assert [event["type"] for event in events] == ["accepted", "done"]
|
||||
build_executor.assert_called_once_with(config, ["risk"])
|
||||
executor.prepare_turn.assert_called_once_with(
|
||||
message="check risk",
|
||||
session_id="risk-session",
|
||||
context={"skills": ["risk"], "report_language": "zh"},
|
||||
selected_skill_ids=["risk"],
|
||||
)
|
||||
|
||||
|
||||
def test_stream_all_invalid_skills_inherit_without_clearing_state() -> None:
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
db.save_conversation_user_turn("saved-session", "first", ["technical"])
|
||||
session_service = AgentChatSessionService(db)
|
||||
executor = _executor(_result(backend="litellm"))
|
||||
config = _litellm_config()
|
||||
skill_manager = MagicMock()
|
||||
skill_manager.list_skills.return_value = [SimpleNamespace(name="technical")]
|
||||
|
||||
async def exercise() -> list[dict]:
|
||||
with patch("api.v1.endpoints.agent.asyncio.to_thread", side_effect=_immediate_to_thread), \
|
||||
patch("api.v1.endpoints.agent.get_config", return_value=config), \
|
||||
patch("src.agent.factory.get_skill_manager", return_value=skill_manager), \
|
||||
patch("api.v1.endpoints.agent._build_executor", return_value=executor) as build_executor:
|
||||
response = await agent_endpoint.agent_chat_stream(
|
||||
agent_endpoint.ChatRequest(
|
||||
message="follow up",
|
||||
session_id="saved-session",
|
||||
skills=["old_technical"],
|
||||
),
|
||||
session_service=session_service,
|
||||
)
|
||||
events = [
|
||||
json.loads(chunk.removeprefix("data: ").strip())
|
||||
async for chunk in response.body_iterator
|
||||
]
|
||||
|
||||
build_executor.assert_called_once_with(config, ["technical"])
|
||||
return events
|
||||
|
||||
events = asyncio.run(exercise())
|
||||
|
||||
assert [event["type"] for event in events] == ["accepted", "done"]
|
||||
executor.prepare_turn.assert_called_once_with(
|
||||
message="follow up",
|
||||
session_id="saved-session",
|
||||
context={"skills": ["technical"], "report_language": "zh"},
|
||||
selected_skill_ids=None,
|
||||
)
|
||||
assert db.get_conversation_session_selected_skill_ids("saved-session") == [
|
||||
"technical"
|
||||
]
|
||||
|
||||
|
||||
def test_codex_stream_skill_resolution_failure_does_not_register_request() -> None:
|
||||
request_id = "skill-resolution-failure"
|
||||
session_service = MagicMock(spec=AgentChatSessionService)
|
||||
session_service.resolve_skill_selection.side_effect = RuntimeError("database read failed")
|
||||
|
||||
try:
|
||||
with patch("api.v1.endpoints.agent.get_config", return_value=_codex_config()), \
|
||||
pytest.raises(RuntimeError, match="database read failed"):
|
||||
asyncio.run(
|
||||
agent_endpoint.agent_chat_stream(
|
||||
agent_endpoint.ChatRequest(
|
||||
message="question",
|
||||
session_id="failed-session",
|
||||
request_id=request_id,
|
||||
),
|
||||
session_service=session_service,
|
||||
)
|
||||
)
|
||||
|
||||
with agent_endpoint._ACTIVE_CODEX_STREAMS_LOCK:
|
||||
assert request_id not in agent_endpoint._ACTIVE_CODEX_STREAMS
|
||||
finally:
|
||||
with agent_endpoint._ACTIVE_CODEX_STREAMS_LOCK:
|
||||
agent_endpoint._ACTIVE_CODEX_STREAMS.pop(request_id, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure", ["context preparation failed", "database write failed"])
|
||||
def test_stream_preparation_failure_emits_no_accepted_and_never_starts_backend(failure: str) -> None:
|
||||
executor = _executor()
|
||||
@@ -362,7 +573,8 @@ def test_stream_preparation_failure_emits_no_accepted_and_never_starts_backend(f
|
||||
patch("api.v1.endpoints.agent.get_config", return_value=_codex_config()), \
|
||||
patch("api.v1.endpoints.agent._build_executor", return_value=executor):
|
||||
response = await agent_endpoint.agent_chat_stream(
|
||||
agent_endpoint.ChatRequest(message="question", session_id="failed-session")
|
||||
agent_endpoint.ChatRequest(message="question", session_id="failed-session"),
|
||||
session_service=AgentChatSessionService(),
|
||||
)
|
||||
return [
|
||||
json.loads(chunk.removeprefix("data: ").strip())
|
||||
@@ -382,7 +594,8 @@ def test_server_selects_actual_backend_for_stream() -> None:
|
||||
patch("api.v1.endpoints.agent._build_executor", return_value=executor):
|
||||
async def exercise() -> dict:
|
||||
response = await agent_endpoint.agent_chat_stream(
|
||||
agent_endpoint.ChatRequest(message="分析 AAPL", session_id="actual-backend")
|
||||
agent_endpoint.ChatRequest(message="分析 AAPL", session_id="actual-backend"),
|
||||
session_service=AgentChatSessionService(),
|
||||
)
|
||||
iterator = response.body_iterator
|
||||
first = json.loads((await anext(iterator)).removeprefix("data: ").strip())
|
||||
@@ -403,7 +616,8 @@ def test_agent_chat_stream_cancels_backend_when_generator_closes() -> None:
|
||||
patch("api.v1.endpoints.agent.get_config", return_value=_codex_config()), \
|
||||
patch("api.v1.endpoints.agent._build_executor", return_value=executor):
|
||||
response = await agent_endpoint.agent_chat_stream(
|
||||
agent_endpoint.ChatRequest(message="question", session_id="cancel-session")
|
||||
agent_endpoint.ChatRequest(message="question", session_id="cancel-session"),
|
||||
session_service=AgentChatSessionService(),
|
||||
)
|
||||
iterator = response.body_iterator
|
||||
accepted = json.loads((await anext(iterator)).removeprefix("data: ").strip())
|
||||
|
||||
@@ -67,12 +67,14 @@ def test_runtime_owned_backend_uses_visible_history_and_forwards_cancellation()
|
||||
)
|
||||
with patch("src.agent.chat_executor.prepare_agent_chat", return_value=prepared) as prepare, \
|
||||
patch("src.agent.chat_executor.conversation_manager.get_or_create"), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", side_effect=[1, 2]), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_user_message", return_value=1) as add_user_message, \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", return_value=2), \
|
||||
patch("src.agent.chat_executor.persist_provider_trace_turns") as persist_trace:
|
||||
result = _executor(backend).chat(
|
||||
"question",
|
||||
"session",
|
||||
cancel_event=cancel_event,
|
||||
selected_skill_ids=[],
|
||||
)
|
||||
|
||||
assert result.success is True
|
||||
@@ -83,6 +85,7 @@ def test_runtime_owned_backend_uses_visible_history_and_forwards_cancellation()
|
||||
assert backend.request.max_wall_clock_seconds == 45
|
||||
assert prepare.call_args.kwargs["include_provider_trace"] is False
|
||||
assert prepare.call_args.kwargs["strict_initial_stock_scope"] is True
|
||||
add_user_message.assert_called_once_with("session", "question", [])
|
||||
persist_trace.assert_not_called()
|
||||
|
||||
|
||||
@@ -95,7 +98,8 @@ def test_dsa_owned_backend_keeps_provider_trace_roundtrip() -> None:
|
||||
)
|
||||
with patch("src.agent.chat_executor.prepare_agent_chat", return_value=prepared) as prepare, \
|
||||
patch("src.agent.chat_executor.conversation_manager.get_or_create"), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", side_effect=[11, 12]), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_user_message", return_value=11), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", return_value=12), \
|
||||
patch("src.agent.chat_executor.persist_provider_trace_turns") as persist_trace:
|
||||
result = _executor(backend).chat("question", "session")
|
||||
|
||||
@@ -121,7 +125,8 @@ def test_cancelled_codex_turn_is_not_persisted_as_analysis_failure() -> None:
|
||||
)
|
||||
with patch("src.agent.chat_executor.prepare_agent_chat", return_value=prepared), \
|
||||
patch("src.agent.chat_executor.conversation_manager.get_or_create"), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", side_effect=[1, 2]) as add_message:
|
||||
patch("src.agent.chat_executor.conversation_manager.add_user_message", return_value=1), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", return_value=2) as add_message:
|
||||
result = _executor(backend).chat("question", "session", cancel_event=threading.Event())
|
||||
|
||||
assert result.error_code == "cancelled"
|
||||
@@ -145,7 +150,8 @@ def test_timed_out_codex_turn_uses_codex_terminal_note() -> None:
|
||||
)
|
||||
with patch("src.agent.chat_executor.prepare_agent_chat", return_value=prepared), \
|
||||
patch("src.agent.chat_executor.conversation_manager.get_or_create"), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", side_effect=[1, 2]) as add_message:
|
||||
patch("src.agent.chat_executor.conversation_manager.add_user_message", return_value=1), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", return_value=2) as add_message:
|
||||
result = _executor(backend).chat("question", "session")
|
||||
|
||||
assert result.error_code == "timeout"
|
||||
@@ -169,7 +175,8 @@ def test_timed_out_litellm_turn_keeps_existing_analysis_failure_note() -> None:
|
||||
)
|
||||
with patch("src.agent.chat_executor.prepare_agent_chat", return_value=prepared), \
|
||||
patch("src.agent.chat_executor.conversation_manager.get_or_create"), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", side_effect=[1, 2]) as add_message:
|
||||
patch("src.agent.chat_executor.conversation_manager.add_user_message", return_value=1), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", return_value=2) as add_message:
|
||||
result = _executor(backend).chat("question", "session")
|
||||
|
||||
assert result.error_code == "timeout"
|
||||
@@ -193,7 +200,8 @@ def test_failed_litellm_turn_keeps_existing_analysis_failure_note() -> None:
|
||||
)
|
||||
with patch("src.agent.chat_executor.prepare_agent_chat", return_value=prepared), \
|
||||
patch("src.agent.chat_executor.conversation_manager.get_or_create"), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", side_effect=[1, 2]) as add_message:
|
||||
patch("src.agent.chat_executor.conversation_manager.add_user_message", return_value=1), \
|
||||
patch("src.agent.chat_executor.conversation_manager.add_message", return_value=2) as add_message:
|
||||
result = _executor(backend).chat("question", "session")
|
||||
|
||||
assert result.error_code == "unknown_backend_error"
|
||||
@@ -210,8 +218,8 @@ def test_context_preparation_failure_does_not_persist_or_start_backend() -> None
|
||||
"src.agent.chat_executor.prepare_agent_chat",
|
||||
side_effect=RuntimeError("context preparation failed"),
|
||||
), patch("src.agent.chat_executor.conversation_manager.get_or_create"), patch(
|
||||
"src.agent.chat_executor.conversation_manager.add_message"
|
||||
) as add_message:
|
||||
"src.agent.chat_executor.conversation_manager.add_user_message"
|
||||
) as add_user_message:
|
||||
try:
|
||||
_executor(backend).prepare_turn(message="question", session_id="session")
|
||||
except RuntimeError as exc:
|
||||
@@ -219,7 +227,7 @@ def test_context_preparation_failure_does_not_persist_or_start_backend() -> None
|
||||
else:
|
||||
raise AssertionError("context preparation failure must propagate")
|
||||
|
||||
add_message.assert_not_called()
|
||||
add_user_message.assert_not_called()
|
||||
assert backend.request is None
|
||||
|
||||
|
||||
@@ -233,7 +241,7 @@ def test_user_message_persistence_failure_does_not_start_backend() -> None:
|
||||
with patch("src.agent.chat_executor.prepare_agent_chat", return_value=prepared), patch(
|
||||
"src.agent.chat_executor.conversation_manager.get_or_create"
|
||||
), patch(
|
||||
"src.agent.chat_executor.conversation_manager.add_message",
|
||||
"src.agent.chat_executor.conversation_manager.add_user_message",
|
||||
side_effect=RuntimeError("database write failed"),
|
||||
):
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Agent Chat session service tests."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from src.services.agent_chat_session_service import AgentChatSessionService
|
||||
from src.storage import DatabaseManager
|
||||
|
||||
|
||||
def test_skill_selection_distinguishes_inherit_clear_and_explicit() -> None:
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
service = AgentChatSessionService(db)
|
||||
config = SimpleNamespace()
|
||||
|
||||
new_selection = service.resolve_skill_selection(config, "new-session", None)
|
||||
assert new_selection.effective_skill_ids is None
|
||||
assert new_selection.selected_skill_ids_update is None
|
||||
|
||||
db.save_conversation_user_turn(
|
||||
"saved-session",
|
||||
"question",
|
||||
["technical", "risk"],
|
||||
)
|
||||
inherited = service.resolve_skill_selection(config, "saved-session", None)
|
||||
assert inherited.effective_skill_ids == ["technical", "risk"]
|
||||
assert inherited.selected_skill_ids_update is None
|
||||
|
||||
cleared = service.resolve_skill_selection(config, "saved-session", [])
|
||||
assert cleared.effective_skill_ids == []
|
||||
assert cleared.selected_skill_ids_update == []
|
||||
|
||||
with patch(
|
||||
"src.services.agent_chat_session_service.normalize_requested_skill_ids",
|
||||
return_value=["technical"],
|
||||
) as normalize:
|
||||
explicit = service.resolve_skill_selection(
|
||||
config,
|
||||
"saved-session",
|
||||
[" technical ", "technical", "unknown"],
|
||||
)
|
||||
|
||||
assert explicit.effective_skill_ids == ["technical"]
|
||||
assert explicit.selected_skill_ids_update == ["technical"]
|
||||
normalize.assert_called_once_with(
|
||||
config,
|
||||
[" technical ", "technical", "unknown"],
|
||||
)
|
||||
|
||||
|
||||
def test_all_invalid_nonempty_selection_inherits_without_clearing_state() -> None:
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
service = AgentChatSessionService(db)
|
||||
config = SimpleNamespace()
|
||||
db.save_conversation_user_turn(
|
||||
"saved-session",
|
||||
"first question",
|
||||
["technical"],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"src.services.agent_chat_session_service.normalize_requested_skill_ids",
|
||||
return_value=[],
|
||||
):
|
||||
inherited = service.resolve_skill_selection(
|
||||
config,
|
||||
"saved-session",
|
||||
["old_technical"],
|
||||
)
|
||||
|
||||
assert inherited.effective_skill_ids == ["technical"]
|
||||
assert inherited.selected_skill_ids_update is None
|
||||
|
||||
db.save_conversation_user_turn(
|
||||
"saved-session",
|
||||
"follow-up",
|
||||
inherited.selected_skill_ids_update,
|
||||
)
|
||||
assert db.get_conversation_session_selected_skill_ids("saved-session") == [
|
||||
"technical"
|
||||
]
|
||||
|
||||
|
||||
def test_all_invalid_nonempty_selection_uses_implicit_default_without_state() -> None:
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
service = AgentChatSessionService(db)
|
||||
|
||||
with patch(
|
||||
"src.services.agent_chat_session_service.normalize_requested_skill_ids",
|
||||
return_value=[],
|
||||
):
|
||||
inherited = service.resolve_skill_selection(
|
||||
SimpleNamespace(),
|
||||
"new-session",
|
||||
["unknown"],
|
||||
)
|
||||
|
||||
assert inherited.effective_skill_ids is None
|
||||
assert inherited.selected_skill_ids_update is None
|
||||
assert db.get_conversation_session_selected_skill_ids("new-session") is None
|
||||
|
||||
|
||||
def test_session_detail_preserves_missing_persisted_state() -> None:
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
service = AgentChatSessionService(db)
|
||||
db.save_conversation_message("legacy-session", "user", "legacy question")
|
||||
|
||||
detail = service.get_session_detail(
|
||||
"legacy-session",
|
||||
limit=100,
|
||||
)
|
||||
|
||||
assert [message["content"] for message in detail.messages] == ["legacy question"]
|
||||
assert detail.selected_skill_ids is None
|
||||
assert db.get_conversation_session_selected_skill_ids("legacy-session") is None
|
||||
@@ -342,6 +342,26 @@ class AgentSkillsEndpointTestCase(unittest.TestCase):
|
||||
],
|
||||
)
|
||||
|
||||
def test_chat_context_without_effective_skills_discards_legacy_selection_fields(self) -> None:
|
||||
request = agent.ChatRequest(
|
||||
message="hello",
|
||||
context={
|
||||
"stock_code": "600519",
|
||||
"skills": ["old_skill"],
|
||||
"strategies": ["older_strategy"],
|
||||
},
|
||||
)
|
||||
|
||||
context = agent._build_agent_chat_context(
|
||||
request,
|
||||
SimpleNamespace(report_language="zh"),
|
||||
skills=None,
|
||||
)
|
||||
|
||||
self.assertEqual(context["stock_code"], "600519")
|
||||
self.assertNotIn("skills", context)
|
||||
self.assertNotIn("strategies", context)
|
||||
|
||||
def test_chat_request_empty_skills_clears_context_without_triggering_activate_all(self) -> None:
|
||||
config = SimpleNamespace(
|
||||
is_agent_available=lambda: True,
|
||||
@@ -349,7 +369,11 @@ class AgentSkillsEndpointTestCase(unittest.TestCase):
|
||||
)
|
||||
executor = MagicMock()
|
||||
executor.chat.return_value = SimpleNamespace(success=True, content="ok", error=None)
|
||||
request = agent.ChatRequest(message="hello", skills=[], context={"skills": ["old_skill"]})
|
||||
request = agent.ChatRequest(
|
||||
message="hello",
|
||||
skills=[],
|
||||
context={"skills": ["old_skill"], "strategies": ["older_strategy"]},
|
||||
)
|
||||
real_get_running_loop = asyncio.get_running_loop
|
||||
|
||||
class _ImmediateLoop:
|
||||
@@ -368,11 +392,18 @@ class AgentSkillsEndpointTestCase(unittest.TestCase):
|
||||
"api.v1.endpoints.agent.asyncio.get_running_loop",
|
||||
side_effect=lambda: _ImmediateLoop(real_get_running_loop()),
|
||||
):
|
||||
payload = asyncio.run(agent.agent_chat(request)).model_dump()
|
||||
payload = asyncio.run(
|
||||
agent.agent_chat(
|
||||
request,
|
||||
session_service=agent.AgentChatSessionService(),
|
||||
)
|
||||
).model_dump()
|
||||
|
||||
mock_build_executor.assert_called_once_with(config, None)
|
||||
executor.chat.assert_called_once()
|
||||
self.assertEqual(executor.chat.call_args.kwargs["context"]["skills"], [])
|
||||
self.assertNotIn("strategies", executor.chat.call_args.kwargs["context"])
|
||||
self.assertEqual(executor.chat.call_args.kwargs["selected_skill_ids"], [])
|
||||
self.assertEqual(payload["content"], "ok")
|
||||
class AgentModelsSourceDetectionTestCase(unittest.TestCase):
|
||||
@patch("src.config.setup_env")
|
||||
|
||||
@@ -9,6 +9,24 @@ from src.agent.conversation import ConversationManager
|
||||
|
||||
|
||||
class ConversationManagerThreadSafetyTestCase(unittest.TestCase):
|
||||
@patch("src.agent.conversation.get_db")
|
||||
def test_add_user_message_uses_session_state_transaction(self, get_db):
|
||||
get_db.return_value.save_conversation_user_turn.return_value = 42
|
||||
manager = ConversationManager()
|
||||
|
||||
message_id = manager.add_user_message(
|
||||
"skill-session",
|
||||
"hello",
|
||||
[],
|
||||
)
|
||||
|
||||
self.assertEqual(message_id, 42)
|
||||
get_db.return_value.save_conversation_user_turn.assert_called_once_with(
|
||||
"skill-session",
|
||||
"hello",
|
||||
[],
|
||||
)
|
||||
|
||||
def test_add_message_is_safe_under_parallel_session_creation(self):
|
||||
manager = ConversationManager()
|
||||
errors = []
|
||||
|
||||
+24
-11
@@ -2003,7 +2003,8 @@ class TestOrchestratorExecution(unittest.TestCase):
|
||||
with patch.object(orch, "_execute_pipeline", side_effect=fake_execute):
|
||||
with patch("src.agent.orchestrator.build_visible_chat_history", return_value=history):
|
||||
with patch("src.agent.conversation.conversation_manager.get_or_create"):
|
||||
with patch("src.agent.conversation.conversation_manager.add_message"):
|
||||
with patch("src.agent.conversation.conversation_manager.add_user_message"), \
|
||||
patch("src.agent.conversation.conversation_manager.add_message"):
|
||||
orch.chat("hello", "session-1")
|
||||
|
||||
self.assertEqual(captured["history"], history)
|
||||
@@ -2016,7 +2017,8 @@ class TestOrchestratorExecution(unittest.TestCase):
|
||||
with patch.object(orch, "_execute_pipeline", return_value=OrchestratorResult(success=True, content="ok")):
|
||||
with patch("src.agent.orchestrator.build_visible_chat_history", return_value=[]) as build_history:
|
||||
with patch("src.agent.conversation.conversation_manager.get_or_create"):
|
||||
with patch("src.agent.conversation.conversation_manager.add_message"):
|
||||
with patch("src.agent.conversation.conversation_manager.add_user_message"), \
|
||||
patch("src.agent.conversation.conversation_manager.add_message"):
|
||||
orch.chat("hello", "session-1")
|
||||
|
||||
build_history.assert_called_once()
|
||||
@@ -2036,7 +2038,8 @@ class TestOrchestratorExecution(unittest.TestCase):
|
||||
with patch.object(orch, "_execute_pipeline", side_effect=fake_execute):
|
||||
with patch("src.agent.orchestrator.build_visible_chat_history", return_value=[]):
|
||||
with patch("src.agent.conversation.conversation_manager.get_or_create"):
|
||||
with patch("src.agent.conversation.conversation_manager.add_message"):
|
||||
with patch("src.agent.conversation.conversation_manager.add_user_message"), \
|
||||
patch("src.agent.conversation.conversation_manager.add_message"):
|
||||
orch.chat(
|
||||
"换成 AAPL 看看",
|
||||
"session-1",
|
||||
@@ -2109,13 +2112,13 @@ class TestOrchestratorExecution(unittest.TestCase):
|
||||
fake_result = OrchestratorResult(success=True, content="assistant reply")
|
||||
|
||||
with patch.object(orch, "_execute_pipeline", return_value=fake_result):
|
||||
with patch("src.agent.conversation.conversation_manager.add_message") as add_message:
|
||||
with patch("src.agent.conversation.conversation_manager.add_user_message") as add_user_message, \
|
||||
patch("src.agent.conversation.conversation_manager.add_message") as add_message:
|
||||
result = orch.chat("hello", "session-1")
|
||||
|
||||
self.assertTrue(result.success)
|
||||
self.assertEqual(add_message.call_count, 2)
|
||||
add_message.assert_any_call("session-1", "user", "hello")
|
||||
add_message.assert_any_call("session-1", "assistant", "assistant reply")
|
||||
add_user_message.assert_called_once_with("session-1", "hello", None)
|
||||
add_message.assert_called_once_with("session-1", "assistant", "assistant reply")
|
||||
|
||||
def test_chat_transaction_persists_user_before_multi_agent_execution(self):
|
||||
"""SSE acceptance can occur after persistence but before the pipeline starts."""
|
||||
@@ -2127,10 +2130,19 @@ class TestOrchestratorExecution(unittest.TestCase):
|
||||
with patch.object(orch, "_execute_pipeline", return_value=fake_result) as execute_pipeline:
|
||||
with patch("src.agent.orchestrator.build_visible_chat_history", return_value=[]):
|
||||
with patch("src.agent.conversation.conversation_manager.get_or_create"):
|
||||
with patch("src.agent.conversation.conversation_manager.add_message") as add_message:
|
||||
turn = orch.prepare_turn(message="hello", session_id="session-accepted")
|
||||
with patch("src.agent.conversation.conversation_manager.add_user_message") as add_user_message, \
|
||||
patch("src.agent.conversation.conversation_manager.add_message") as add_message:
|
||||
turn = orch.prepare_turn(
|
||||
message="hello",
|
||||
session_id="session-accepted",
|
||||
selected_skill_ids=["technical"],
|
||||
)
|
||||
|
||||
add_message.assert_called_once_with("session-accepted", "user", "hello")
|
||||
add_user_message.assert_called_once_with(
|
||||
"session-accepted",
|
||||
"hello",
|
||||
["technical"],
|
||||
)
|
||||
execute_pipeline.assert_not_called()
|
||||
|
||||
result = orch.execute_turn(turn)
|
||||
@@ -2154,7 +2166,8 @@ class TestOrchestratorExecution(unittest.TestCase):
|
||||
fake_result = OrchestratorResult(success=False, error="boom")
|
||||
|
||||
with patch.object(orch, "_execute_pipeline", return_value=fake_result):
|
||||
with patch("src.agent.conversation.conversation_manager.add_message") as add_message:
|
||||
with patch("src.agent.conversation.conversation_manager.add_user_message"), \
|
||||
patch("src.agent.conversation.conversation_manager.add_message") as add_message:
|
||||
result = orch.chat("hello", "session-2")
|
||||
|
||||
self.assertFalse(result.success)
|
||||
|
||||
@@ -601,6 +601,71 @@ class TestStorage(unittest.TestCase):
|
||||
|
||||
DatabaseManager.reset_instance()
|
||||
|
||||
def test_conversation_user_turn_persists_and_updates_session_skills(self):
|
||||
DatabaseManager.reset_instance()
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
|
||||
first_id = db.save_conversation_user_turn(
|
||||
"skill-session",
|
||||
"first question",
|
||||
["technical", "risk"],
|
||||
)
|
||||
second_id = db.save_conversation_user_turn(
|
||||
"skill-session",
|
||||
"use general analysis",
|
||||
[],
|
||||
)
|
||||
|
||||
self.assertGreater(first_id, 0)
|
||||
self.assertGreater(second_id, first_id)
|
||||
self.assertEqual(
|
||||
[message["content"] for message in db.get_conversation_messages("skill-session")],
|
||||
["first question", "use general analysis"],
|
||||
)
|
||||
self.assertEqual(
|
||||
db.get_conversation_session_selected_skill_ids("skill-session"),
|
||||
[],
|
||||
)
|
||||
|
||||
deleted = db.delete_conversation_session("skill-session")
|
||||
|
||||
self.assertEqual(deleted, 2)
|
||||
self.assertIsNone(
|
||||
db.get_conversation_session_selected_skill_ids("skill-session")
|
||||
)
|
||||
DatabaseManager.reset_instance()
|
||||
|
||||
def test_conversation_user_turn_without_skill_update_keeps_session_state(self):
|
||||
DatabaseManager.reset_instance()
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
|
||||
db.save_conversation_user_turn("skill-session", "first", ["technical"])
|
||||
db.save_conversation_user_turn("skill-session", "follow up")
|
||||
|
||||
self.assertEqual(
|
||||
db.get_conversation_session_selected_skill_ids("skill-session"),
|
||||
["technical"],
|
||||
)
|
||||
DatabaseManager.reset_instance()
|
||||
|
||||
def test_conversation_user_turn_rolls_back_message_when_state_write_fails(self):
|
||||
DatabaseManager.reset_instance()
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
|
||||
with patch("src.storage.sqlite_insert", side_effect=RuntimeError("state write failed")):
|
||||
with self.assertRaisesRegex(RuntimeError, "state write failed"):
|
||||
db.save_conversation_user_turn(
|
||||
"skill-session",
|
||||
"not accepted",
|
||||
["technical"],
|
||||
)
|
||||
|
||||
self.assertEqual(db.get_conversation_messages("skill-session"), [])
|
||||
self.assertIsNone(
|
||||
db.get_conversation_session_selected_skill_ids("skill-session")
|
||||
)
|
||||
DatabaseManager.reset_instance()
|
||||
|
||||
def test_provider_turn_round_trip_preserves_protocol_fields_and_flags(self):
|
||||
DatabaseManager.reset_instance()
|
||||
db = DatabaseManager(db_url="sqlite:///:memory:")
|
||||
|
||||
Reference in New Issue
Block a user