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
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 };
|
|
|