You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 

341 lines
10 KiB

/**
* 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<Awaited<ReturnType<typeof getAgentBySlug>>> }> {
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<A2ATask> {
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<A2ATask | null> {
const row = await getTaskRow(taskId);
if (!row) return null;
return rowToTask(row);
}
// ============ cancelTask ============
export async function cancelTask(taskId: string): Promise<boolean> {
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<void> {
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<void> {
const updateData: Record<string, unknown> = {
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 };