/** * A2A Service 核心层 * * 负责 A2A task 的生命周期管理,与 chat-engine 解耦: * - createTask: 创建 task(创建 session + 映射记录) * - sendTask: 同步执行 task(调用 generateText) * - getTask: 查询 task 状态 * - cancelTask: 取消 task(对接 abort-manager) * - subscribeTask: streaming(预留) * * 复用现有模块: * - agent/agent.ts: getAgentBySlug, getAgentById * - agent/session.ts: createSession, getSessionById * - agent-tool/index.ts: getAgentToolsForChatByAgentId * - agent/abort-manager.ts: registerAbortController, unregisterAbortController * - llm/model-resolver: resolveModelForUser, resolveModelAny, toLanguageModel * - utils/context: getConfigGlobal */ import { generateText, stepCountIs } from "ai"; import { dbGlobal } from "drizzle-pkg/lib/db"; import { a2aTaskSessions } from "drizzle-pkg/lib/schema/a2a"; import { eq, and } from "drizzle-orm"; import log4js from "logger"; import { getAgentBySlug, getAgentById } from "../agent/agent"; import { createSession, getSessionById } from "../agent/session"; import { registerAbortController, unregisterAbortController, abortSessionStream } from "../agent/abort-manager"; import { getAgentToolsForChatByAgentId } from "#server/service/agent-tool"; import { resolveModelForUser, resolveModelAny, toLanguageModel } from "#server/service/llm/model-resolver"; import { getConfigGlobal } from "#server/utils/context"; import { a2aMessageToText, textToA2AMessage } from "./converter"; import { DEFAULT_A2A_CONFIG, A2A_PROTOCOL_VERSION, JSONRPC_ERROR_CODES, } from "./types"; import type { A2ATask, A2ATaskState, A2AMessage, A2AInvocationContext, A2AServiceConfig, } from "./types"; const logger = log4js.getLogger("APP"); // ============ 辅助:生成 Task ID ============ function generateTaskId(): string { return `a2a_${Date.now().toString(36)}_${Math.random().toString(36).slice(2, 10)}`; } // ============ 辅助:row → A2ATask ============ function rowToTask(row: typeof a2aTaskSessions.$inferSelect, message?: A2AMessage): A2ATask { return { id: row.taskId, sessionId: row.sessionId, state: row.state as A2ATaskState, message, createdAt: row.createdAt.getTime(), updatedAt: row.updatedAt.getTime(), }; } // ============ createTask ============ export async function createTask(params: { calleeAgentSlug: string; message: A2AMessage; callerAgentId?: number | null; userId?: number | null; tempToken?: string | null; existingSessionId?: string; }): Promise<{ task: A2ATask; agent: NonNullable>> }> { const { calleeAgentSlug, message, callerAgentId = null, userId = null, tempToken = null } = params; const agent = await getAgentBySlug(calleeAgentSlug); if (!agent) { throw createA2AError(JSONRPC_ERROR_CODES.AGENT_NOT_CALLABLE, `agent not found: ${calleeAgentSlug}`); } let sessionId: string; if (params.existingSessionId) { const existing = await getSessionById(params.existingSessionId); if (!existing) { throw createA2AError(JSONRPC_ERROR_CODES.INTERNAL_ERROR, `session not found: ${params.existingSessionId}`); } sessionId = existing.id; } else { const session = await createSession({ userId, tempToken, agentId: agent.id, modelId: agent.defaultModelId, enableThinking: agent.enableThinking, enableTools: agent.enableTools, systemPrompt: agent.systemPrompt, }); sessionId = session.id; } const inputText = a2aMessageToText(message); const taskId = generateTaskId(); const now = new Date(); const [row] = await dbGlobal .insert(a2aTaskSessions) .values({ taskId, sessionId, callerAgentId, calleeAgentId: agent.id, state: "submitted", callerContext: inputText, createdAt: now, updatedAt: now, }) .returning(); const task = rowToTask(row!, message); return { task, agent }; } // ============ sendTask ============ export async function sendTask( taskId: string, context: A2AInvocationContext, ): Promise { const { calleeAgentId, calleeAgentSlug, userId, recursionDepth, maxRecursionDepth, timeoutMs, maxOutputTokens } = context; const taskRow = await getTaskRow(taskId); if (!taskRow) { throw createA2AError(JSONRPC_ERROR_CODES.TASK_NOT_FOUND, `task not found: ${taskId}`); } if (calleeAgentId === null) { throw createA2AError(JSONRPC_ERROR_CODES.AGENT_NOT_CALLABLE, `calleeAgentId is null for task ${taskId}`); } const agent = await getAgentById(calleeAgentId); if (!agent) { throw createA2AError(JSONRPC_ERROR_CODES.AGENT_NOT_CALLABLE, `agent not found by id: ${calleeAgentId}`); } if (!agent.isCallable) { await updateTaskState(taskId, "failed"); throw createA2AError(JSONRPC_ERROR_CODES.AGENT_NOT_CALLABLE, `agent not callable: ${calleeAgentSlug}`); } if (context.callerAgentId !== null && context.callerAgentId === calleeAgentId) { await updateTaskState(taskId, "failed"); throw createA2AError(JSONRPC_ERROR_CODES.SELF_INVOCATION_BLOCKED, `self-invocation blocked: agent ${calleeAgentSlug} cannot invoke itself`); } if (recursionDepth >= maxRecursionDepth) { await updateTaskState(taskId, "failed"); throw createA2AError(JSONRPC_ERROR_CODES.MAX_RECURSION_REACHED, `max recursion depth (${maxRecursionDepth}) reached`); } const inputText = taskRow.callerContext ?? ""; if (!inputText) { throw createA2AError(JSONRPC_ERROR_CODES.INVALID_PARAMS, `task ${taskId} has no input message`); } await updateTaskState(taskId, "working"); const globalDefaultModelId = (await getConfigGlobal("agentDefaultModelId")) ?? null; const targetModelId = agent.defaultModelId ?? globalDefaultModelId; if (!targetModelId) { await updateTaskState(taskId, "failed"); throw createA2AError(JSONRPC_ERROR_CODES.INTERNAL_ERROR, `agent ${calleeAgentSlug} has no defaultModelId and no global default model configured`); } const resolved = userId ? await resolveModelForUser(targetModelId, userId) : await resolveModelAny(targetModelId); if (!resolved) { await updateTaskState(taskId, "failed"); throw createA2AError(JSONRPC_ERROR_CODES.INTERNAL_ERROR, `model ${targetModelId} not resolvable for agent ${calleeAgentSlug}`); } const languageModel = toLanguageModel(resolved); const enableTools = agent.enableTools === 1; const { tools } = await getAgentToolsForChatByAgentId({ agentId: agent.id, userId, userRole: null, enableTools, recursionDepth, }); const maxSteps = agent.maxStepCount ?? 8; const abortController = new AbortController(); const timeoutSignal = AbortSignal.timeout(timeoutMs); const combinedSignal = anySignal([abortController.signal, timeoutSignal]); registerAbortController(taskRow.sessionId, abortController); logger.info( "[A2A] sendTask taskId=%s agentSlug=%s agentId=%d callerAgentId=%s recursionDepth=%d tools=%d maxSteps=%d", taskId, calleeAgentSlug, agent.id, context.callerAgentId, recursionDepth, Object.keys(tools).length, maxSteps, ); try { const result = await generateText({ model: languageModel, system: agent.systemPrompt, prompt: inputText, ...(Object.keys(tools).length > 0 ? { tools, stopWhen: stepCountIs(maxSteps), } : {}), maxOutputTokens, abortSignal: combinedSignal, }); const output = result.text ?? ""; const responseMessage = textToA2AMessage(output, "agent"); await updateTaskState(taskId, "completed"); logger.info( "[A2A] sendTask done taskId=%s agentSlug=%s outputLen=%d usage=%j", taskId, calleeAgentSlug, output.length, result.usage, ); const updatedRow = await getTaskRow(taskId); return rowToTask(updatedRow!, responseMessage); } catch (err) { const errorMsg = err instanceof Error ? err.message : String(err); logger.error("[A2A] sendTask error taskId=%s agentSlug=%s error=%s", taskId, calleeAgentSlug, errorMsg); const isTimeout = errorMsg.includes("timeout") || errorMsg.includes("abort"); await updateTaskState(taskId, isTimeout ? "failed" : "failed"); throw createA2AError( isTimeout ? JSONRPC_ERROR_CODES.TIMEOUT : JSONRPC_ERROR_CODES.INTERNAL_ERROR, errorMsg, ); } finally { unregisterAbortController(taskRow.sessionId); } } // ============ getTask ============ export async function getTask(taskId: string): Promise { const row = await getTaskRow(taskId); if (!row) return null; return rowToTask(row); } // ============ cancelTask ============ export async function cancelTask(taskId: string): Promise { const row = await getTaskRow(taskId); if (!row) return false; const cancelableStates: A2ATaskState[] = ["submitted", "working", "input-required"]; if (!cancelableStates.includes(row.state as A2ATaskState)) { return false; } const aborted = abortSessionStream(row.sessionId); await updateTaskState(taskId, "canceled"); return aborted; } // ============ subscribeTask(预留) ============ export async function subscribeTask( _taskId: string, _onChunk: (chunk: string) => void, _onDone: (task: A2ATask) => void, _onError: (error: Error) => void, ): Promise { throw createA2AError(JSONRPC_ERROR_CODES.NOT_IMPLEMENTED, "subscribeTask not implemented yet"); } // ============ 内部辅助函数 ============ async function getTaskRow(taskId: string) { const [row] = await dbGlobal .select() .from(a2aTaskSessions) .where(eq(a2aTaskSessions.taskId, taskId)) .limit(1); return row ?? null; } async function updateTaskState(taskId: string, state: A2ATaskState): Promise { const updateData: Record = { state, updatedAt: new Date(), }; if (state === "completed" || state === "failed" || state === "canceled") { updateData.completedAt = new Date(); } await dbGlobal .update(a2aTaskSessions) .set(updateData) .where(eq(a2aTaskSessions.taskId, taskId)); } function createA2AError(code: number, message: string): Error { const err = new Error(message); (err as any).code = code; return err; } function anySignal(signals: AbortSignal[]): AbortSignal { const controller = new AbortController(); for (const signal of signals) { if (signal.aborted) { controller.abort(); break; } signal.addEventListener("abort", () => controller.abort(), { once: true }); } return controller.signal; } // ============ 导出配置 ============ export { DEFAULT_A2A_CONFIG, A2A_PROTOCOL_VERSION }; export type { A2AServiceConfig };