返回 AiToEarn
agent.methods.ts
根目录 / project / aitoearn-web / src / store / agent / agent.methods.ts
1 /**
2 * Agent Store - 核心方法
3 * 包含创建任务、继续任务等核心逻辑
4 * 支持按任务ID隔离消息数据
5 *
6 * 使用 TaskInstance 架构:每个任务有独立的实例,消除多任务竞态条件
7 */
8
9 import type {
10 IActionContext,
11 IAgentState,
12 ICreateTaskParams,
13 IPendingTask,
14 ISSEMessage,
15 ITaskMessageData,
16 } from './agent.types'
17 import type { ITaskInstanceContext, ISSECallbacks as ITaskSSECallbacks } from './task-instance'
18 import type { MessageUtils } from './utils/message'
19 import type { IAgentRefs } from './utils/refs'
20 import { agentApi } from '@/api/ai/ai.api'
21 import { useUserStore } from '@/store/user'
22 import { toast } from '@/utils/ui/toast'
23 import { getDefaultTaskData, getInitialState } from './agent.state'
24 import { TaskInstance } from './task-instance'
25 import { buildPromptForAPI } from './utils/buildPrompt'
26
27 // ============ 常量配置 ============
28
29 /** 最大缓存任务数量 */
30 const MAX_CACHED_TASKS = 10
31
32 // ============ TaskInstance 管理 ============
33
34 /** 任务实例映射表 */
35 const taskInstances = new Map<string, TaskInstance>()
36
37 // ============ 方法工厂上下文 ============
38
39 export interface IMethodsContext {
40 refs: IAgentRefs
41 set: (partial: Partial<IAgentState> | ((state: IAgentState) => Partial<IAgentState>)) => void
42 get: () => IAgentState
43 messageUtils: MessageUtils
44 resetRefs: () => void
45 }
46
47 // ============ 创建 Store 方法 ============
48
49 export function createStoreMethods(ctx: IMethodsContext) {
50 const { refs, set, get, messageUtils, resetRefs } = ctx
51
52 // ============ 任务数据操作辅助方法 ============
53
54 /**
55 * 获取指定任务的消息数据
56 */
57 function getTaskData(taskId: string): ITaskMessageData {
58 const state = get()
59 return state.taskMessages[taskId] || getDefaultTaskData()
60 }
61
62 /**
63 * 更新指定任务的消息数据
64 */
65 function updateTaskData(
66 taskId: string,
67 updater: (data: ITaskMessageData) => Partial<ITaskMessageData>,
68 ) {
69 if (!taskId) {
70 console.warn('[AgentStore] updateTaskData called without taskId')
71 return
72 }
73 set((state) => {
74 const currentData = state.taskMessages[taskId] || getDefaultTaskData()
75 const updates = updater(currentData)
76 return {
77 taskMessages: {
78 ...state.taskMessages,
79 [taskId]: {
80 ...currentData,
81 ...updates,
82 lastUpdated: Date.now(),
83 },
84 },
85 }
86 })
87 }
88
89 /**
90 * 初始化任务数据
91 */
92 function initTaskData(taskId: string, initialData?: Partial<ITaskMessageData>) {
93 set(state => ({
94 taskMessages: {
95 ...state.taskMessages,
96 [taskId]: {
97 ...getDefaultTaskData(),
98 ...initialData,
99 lastUpdated: Date.now(),
100 },
101 },
102 }))
103 }
104
105 /**
106 * 清理过期任务缓存
107 */
108 function cleanupTaskCache() {
109 set((state) => {
110 const entries = Object.entries(state.taskMessages)
111
112 // 如果未超过最大数量,不清理
113 if (entries.length <= MAX_CACHED_TASKS)
114 return {}
115
116 // 按最后更新时间排序
117 const sorted = entries.sort((a, b) => (b[1].lastUpdated || 0) - (a[1].lastUpdated || 0))
118
119 // 保留最近的 N 个任务
120 const tasksToKeep = sorted.slice(0, MAX_CACHED_TASKS)
121 const newTaskMessages: Record<string, ITaskMessageData> = {}
122
123 for (const [taskId, data] of tasksToKeep) {
124 newTaskMessages[taskId] = data
125 }
126
127 // 确保当前任务不被清理
128 if (state.currentTaskId && !newTaskMessages[state.currentTaskId]) {
129 newTaskMessages[state.currentTaskId] = state.taskMessages[state.currentTaskId]
130 }
131
132 return { taskMessages: newTaskMessages }
133 })
134 }
135
136 // ============ 返回 Store 方法 ============
137
138 return {
139 // ============ 任务数据 Getters ============
140
141 /** 获取指定任务的数据 */
142 getTaskData,
143
144 /** 更新指定任务的数据 */
145 updateTaskData,
146
147 /** 初始化任务数据 */
148 initTaskData,
149
150 /** 清理任务缓存 */
151 cleanupTaskCache,
152
153 // ============ 核心方法:创建任务 ============
154
155 /**
156 * 创建 AI 生成任务
157 * 使用 TaskInstance 架构:创建独立实例,SSE 回调绑定到实例
158 */
159 async createTask(params: ICreateTaskParams): Promise<string | null> {
160 const { prompt, medias = [], t, onTaskIdReady, onLoginRequired } = params
161
162 if (!prompt.trim()) {
163 return null
164 }
165
166 refs.t.value = t
167
168 // 检查登录状态
169 const currentToken = useUserStore.getState().token
170 if (!currentToken) {
171 onLoginRequired?.()
172 return null
173 }
174
175 try {
176 // 生成临时任务ID(用于在获取真实ID之前存储消息)
177 const tempTaskId = `temp-${Date.now()}`
178
179 // 创建 TaskInstance 上下文
180 const instanceContext: ITaskInstanceContext = {
181 syncToStore: (taskId, updater) => updateTaskData(taskId, updater),
182 getData: taskId => getTaskData(taskId),
183 migrateTaskData: (fromTaskId, toTaskId) => {
184 set((state) => {
185 const tempData = state.taskMessages[fromTaskId]
186 if (!tempData)
187 return {}
188
189 const { [fromTaskId]: _, ...restTaskMessages } = state.taskMessages
190
191 return {
192 currentTaskId: toTaskId,
193 taskMessages: {
194 ...restTaskMessages,
195 [toTaskId]: {
196 ...tempData,
197 lastUpdated: Date.now(),
198 },
199 },
200 }
201 })
202 },
203 setCurrentTaskId: (taskId) => {
204 set({ currentTaskId: taskId })
205 },
206 }
207
208 // 创建 TaskInstance(绑定到临时任务ID)
209 const instance = new TaskInstance(tempTaskId, instanceContext)
210 taskInstances.set(tempTaskId, instance)
211
212 // 设置翻译函数和 Action 上下文
213 instance.setTranslation(t)
214 if (refs.actionContext.value) {
215 instance.setActionContext(refs.actionContext.value)
216 }
217
218 // 重置全局状态,使用临时任务ID
219 set({
220 currentTaskId: tempTaskId,
221 currentCost: 0,
222 })
223 resetRefs()
224
225 // 设置 SSE 任务ID(确保后续 SSE 消息写入正确的任务)
226 refs.currentSSETaskId.value = tempTaskId
227
228 // 初始化临时任务的数据
229 initTaskData(tempTaskId, {
230 isGenerating: true,
231 progress: 0,
232 messages: [],
233 markdownMessages: [],
234 workflowSteps: [],
235 streamingText: '',
236 })
237
238 // 添加用户消息(通过 TaskInstance)
239 const userMessage = instance.createUserMessage(prompt, medias)
240 instance.addMessage(userMessage)
241 instance.addMarkdownMessage(`👤 ${prompt}`)
242
243 // 构建 Claude Prompt 格式
244 const apiPrompt = buildPromptForAPI(prompt, medias)
245
246 // 添加 AI 待回复消息(通过 TaskInstance)
247 const assistantMessage = instance.createAssistantMessage()
248 instance.addMessage(assistantMessage)
249
250 // 同步 refs(为了兼容旧的 SSE handler)
251 refs.currentAssistantMessageId.value = assistantMessage.id
252
253 // 创建 SSE 回调,绑定到 TaskInstance
254 const taskSSECallbacks: ITaskSSECallbacks = {
255 onTaskIdReady: (realTaskId: string) => {
256 // 迁移 TaskInstance(从临时ID到真实ID)
257 taskInstances.delete(tempTaskId)
258 taskInstances.set(realTaskId, instance)
259
260 // 调用外部回调
261 onTaskIdReady?.(realTaskId)
262
263 // 清理过期缓存
264 cleanupTaskCache()
265 },
266 onError: (error) => {
267 console.error('[AgentStore] TaskInstance SSE Error:', error)
268 },
269 onComplete: () => {
270 useUserStore.getState().fetchCreditsBalance()
271 },
272 }
273
274 // 创建任务(SSE)- SSE 消息通过 TaskInstance 处理
275 const abortFn = await agentApi.createTaskWithSSE(
276 { prompt: apiPrompt, includePartialMessages: true },
277 (sseMessage: ISSEMessage) => {
278 // 使用 TaskInstance 处理 SSE 消息(消息会自动写入实例的 taskId)
279 instance.handleSSEMessage(sseMessage, taskSSECallbacks)
280 },
281 (error) => {
282 console.error('[AgentStore] SSE Error:', error)
283 const errorMsg = refs.t.value
284 ? `${refs.t.value('aiGeneration.createTaskFailed' as any)}: ${error.message || refs.t.value('aiGeneration.unknownError' as any)}`
285 : `Create task failed: ${error.message}`
286 toast.error(errorMsg)
287
288 instance.markMessageError(error.message)
289 instance.setIsGenerating(false)
290 instance.setProgress(0)
291 },
292 async () => {
293 instance.markMessageDone()
294 instance.setIsGenerating(false)
295 instance.clearWorkflowSteps()
296 refs.sseAbort.value = null
297 useUserStore.getState().fetchCreditsBalance()
298 },
299 )
300
301 // 保存 abort 函数到实例和全局 refs
302 instance.setAbort(abortFn)
303 refs.sseAbort.value = abortFn
304
305 // 等待获取 taskId
306 let waitTime = 0
307 const maxWaitTime = 30000
308 const checkInterval = 100
309
310 while (get().currentTaskId.startsWith('temp-') && waitTime < maxWaitTime) {
311 await new Promise(resolve => setTimeout(resolve, checkInterval))
312 waitTime += checkInterval
313 }
314
315 const finalTaskId = get().currentTaskId
316 return finalTaskId.startsWith('temp-') ? null : finalTaskId
317 }
318 catch (error: any) {
319 console.error('[AgentStore] Create task error:', error)
320 const errorMsg = refs.t.value
321 ? `${refs.t.value('aiGeneration.createTaskFailed' as any)}: ${error.message || refs.t.value('aiGeneration.unknownError' as any)}`
322 : `Create task failed: ${error.message}`
323 toast.error(errorMsg)
324
325 const currentTaskId = get().currentTaskId
326 if (currentTaskId) {
327 updateTaskData(currentTaskId, () => ({
328 isGenerating: false,
329 progress: 0,
330 }))
331 }
332 refs.sseAbort.value = null
333 return null
334 }
335 },
336
337 /**
338 * 继续对话
339 * 使用 TaskInstance 架构:获取或创建实例,SSE 回调绑定到实例
340 */
341 async continueTask(params: ICreateTaskParams & { taskId: string }): Promise<void> {
342 const { prompt, medias = [], t, taskId } = params
343
344 if (!prompt.trim() || !taskId) {
345 return
346 }
347
348 refs.t.value = t
349
350 try {
351 // 创建 TaskInstance 上下文
352 const instanceContext: ITaskInstanceContext = {
353 syncToStore: (tid, updater) => updateTaskData(tid, updater),
354 getData: tid => getTaskData(tid),
355 migrateTaskData: (fromTaskId, toTaskId) => {
356 // continueTask 不需要迁移,taskId 已知
357 set((state) => {
358 const data = state.taskMessages[fromTaskId]
359 if (!data || fromTaskId === toTaskId)
360 return {}
361
362 const { [fromTaskId]: _, ...restTaskMessages } = state.taskMessages
363
364 return {
365 currentTaskId: toTaskId,
366 taskMessages: {
367 ...restTaskMessages,
368 [toTaskId]: {
369 ...data,
370 lastUpdated: Date.now(),
371 },
372 },
373 }
374 })
375 },
376 setCurrentTaskId: (tid) => {
377 set({ currentTaskId: tid })
378 },
379 }
380
381 // 获取或创建 TaskInstance
382 let instance = taskInstances.get(taskId)
383 if (!instance) {
384 instance = new TaskInstance(taskId, instanceContext)
385 taskInstances.set(taskId, instance)
386 }
387
388 // 重置实例状态(新一轮对话)
389 instance.resetForNewRound()
390
391 // 设置翻译函数和 Action 上下文
392 instance.setTranslation(t)
393 if (refs.actionContext.value) {
394 instance.setActionContext(refs.actionContext.value)
395 }
396
397 // 设置当前任务ID
398 set({ currentTaskId: taskId })
399 resetRefs()
400
401 // 设置 SSE 任务ID(确保后续 SSE 消息写入正确的任务)
402 refs.currentSSETaskId.value = taskId
403
404 // 确保任务数据存在,更新状态
405 const existingData = getTaskData(taskId)
406 updateTaskData(taskId, () => ({
407 isGenerating: true,
408 progress: 10,
409 workflowSteps: [],
410 // 保留现有消息
411 messages: existingData.messages,
412 markdownMessages: existingData.markdownMessages,
413 }))
414
415 // 添加用户消息(通过 TaskInstance)
416 const userMessage = instance.createUserMessage(prompt, medias)
417 instance.addMessage(userMessage)
418 instance.addMarkdownMessage(`👤 ${prompt}`)
419
420 // 构建 Claude Prompt 格式
421 const apiPrompt = buildPromptForAPI(prompt, medias)
422
423 // 添加 AI 待回复消息(通过 TaskInstance)
424 const assistantMessage = instance.createAssistantMessage()
425 instance.addMessage(assistantMessage)
426
427 // 同步 refs(为了兼容旧的 SSE handler)
428 refs.currentAssistantMessageId.value = assistantMessage.id
429
430 // 创建 SSE 回调,绑定到 TaskInstance
431 const taskSSECallbacks: ITaskSSECallbacks = {
432 onTaskIdReady: (_receivedTaskId: string) => {
433 // continueTask 时 taskId 已知,只需确认
434 },
435 onError: (error) => {
436 console.error('[AgentStore] TaskInstance SSE Error:', error)
437 },
438 onComplete: () => {
439 useUserStore.getState().fetchCreditsBalance()
440 },
441 }
442
443 // 创建任务(SSE)- SSE 消息通过 TaskInstance 处理
444 const abortFn = await agentApi.createTaskWithSSE(
445 { prompt: apiPrompt, taskId, includePartialMessages: true },
446 (sseMessage: ISSEMessage) => {
447 // 使用 TaskInstance 处理 SSE 消息
448 instance!.handleSSEMessage(sseMessage, taskSSECallbacks)
449 },
450 (error) => {
451 console.error('[AgentStore] SSE Error:', error)
452 toast.error(error.message || 'Generation failed')
453 instance!.markMessageError(error.message)
454 instance!.setIsGenerating(false)
455 instance!.setProgress(0)
456 },
457 async () => {
458 instance!.markMessageDone()
459 instance!.setIsGenerating(false)
460 instance!.clearWorkflowSteps()
461 refs.sseAbort.value = null
462 useUserStore.getState().fetchCreditsBalance()
463 },
464 )
465
466 // 保存 abort 函数到实例和全局 refs
467 instance.setAbort(abortFn)
468 refs.sseAbort.value = abortFn
469 }
470 catch (error: any) {
471 console.error('[AgentStore] Continue task error:', error)
472 toast.error(error.message || 'Continue task failed')
473 updateTaskData(taskId, () => ({
474 isGenerating: false,
475 progress: 0,
476 }))
477 refs.sseAbort.value = null
478 }
479 },
480
481 // ============ 任务控制 ============
482
483 /** 停止当前任务 */
484 stopTask() {
485 if (refs.sseAbort.value) {
486 refs.sseAbort.value()
487 refs.sseAbort.value = null
488 }
489
490 const taskId = get().currentTaskId
491 if (taskId) {
492 updateTaskData(taskId, () => ({
493 isGenerating: false,
494 progress: 0,
495 workflowSteps: [],
496 }))
497 }
498
499 messageUtils.markMessageDone()
500 // 移除 toast 显示,改为由调用方处理
501 },
502
503 /** 重置状态 */
504 reset() {
505 if (refs.sseAbort.value) {
506 refs.sseAbort.value()
507 refs.sseAbort.value = null
508 }
509 resetRefs()
510 refs.t.value = null
511 refs.actionContext.value = null
512 set(getInitialState())
513 },
514
515 // ============ 消息管理 ============
516
517 setMessages: messageUtils.setMessages.bind(messageUtils),
518 appendMessage: messageUtils.addMessage.bind(messageUtils),
519
520 // ============ 待处理任务管理 ============
521
522 /** 设置待处理任务(从首页跳转时使用) */
523 setPendingTask(task: IPendingTask) {
524 set({ pendingTask: task })
525 },
526
527 /** 获取并清除待处理任务 */
528 consumePendingTask(): IPendingTask | null {
529 const task = get().pendingTask
530 if (task) {
531 set({ pendingTask: null })
532 }
533 return task
534 },
535
536 // ============ Action 上下文管理 ============
537
538 /** 设置 Action 上下文 */
539 setActionContext(context: IActionContext) {
540 refs.actionContext.value = context
541 },
542
543 /** 获取 Action 上下文 */
544 getActionContext(): IActionContext | null {
545 return refs.actionContext.value
546 },
547
548 // ============ Debug 模式管理 ============
549
550 /**
551 * 设置 debug 文件列表
552 * @param files debug 文件名数组(如 ['sse1.txt', 'sse2.txt'])
553 */
554 setDebugFiles(files: string[]) {
555 set({
556 debugFiles: files,
557 debugMessageIndex: 0,
558 })
559 },
560
561 /**
562 * 获取下一个 debug 文件路径并递增索引
563 * @returns 文件路径(如 '/en/debug/sse1.txt')或 null(没有更多文件)
564 */
565 consumeDebugFile(): string | null {
566 const state = get()
567 const { debugFiles, debugMessageIndex } = state
568
569 if (debugMessageIndex >= debugFiles.length) {
570 return null
571 }
572
573 const fileName = debugFiles[debugMessageIndex]
574 const filePath = `/en/debug/${fileName}`
575
576 // 递增索引
577 set({ debugMessageIndex: debugMessageIndex + 1 })
578
579 return filePath
580 },
581
582 /**
583 * 检查是否处于 debug 模式
584 */
585 isDebugMode(): boolean {
586 const state = get()
587 return state.debugFiles.length > 0
588 },
589
590 /**
591 * 检查是否还有更多 debug 文件可用
592 */
593 hasMoreDebugFiles(): boolean {
594 const state = get()
595 return state.debugMessageIndex < state.debugFiles.length
596 },
597
598 /**
599 * 清除 debug 模式
600 */
601 clearDebugMode() {
602 set({
603 debugFiles: [],
604 debugMessageIndex: 0,
605 })
606 },
607 }
608 }
609
609 lines TYPESCRIPT