Files
memind/agent-run-routes.mjs
T
2026-07-02 10:47:22 +08:00

223 lines
7.3 KiB
JavaScript

import { normalizeAgentRunToolMode } from './agent-run-gateway.mjs';
function envFlag(value) {
return ['1', 'true', 'yes', 'on'].includes(String(value ?? '').trim().toLowerCase());
}
function parseUserIdSet(value) {
return new Set(
String(value ?? '')
.split(',')
.map((item) => item.trim())
.filter(Boolean),
);
}
function parseTaskTypeSet(value) {
return new Set(
String(value ?? '')
.split(',')
.map((item) => item.trim().toLowerCase())
.filter(Boolean),
);
}
function hasExpectedFileValidation(userMessage) {
const metadata = userMessage?.metadata;
const runMetadata = metadata?.memindRun ?? metadata?.agentRun ?? {};
const validation = runMetadata.validation ?? metadata?.toolGatewayValidation;
if (!validation || typeof validation !== 'object' || Array.isArray(validation)) return false;
const expectedFile = validation.expectedFile ?? validation.expectedPath;
if (typeof expectedFile === 'string' && expectedFile.trim()) return true;
if (expectedFile && typeof expectedFile === 'object' && !Array.isArray(expectedFile)) {
const filePath = String(expectedFile.path ?? expectedFile.file ?? expectedFile.relativePath ?? '').trim();
if (filePath) return true;
}
if (!Array.isArray(validation.expectedFiles)) return false;
return validation.expectedFiles.some((item) => {
if (typeof item === 'string') return Boolean(item.trim());
if (!item || typeof item !== 'object' || Array.isArray(item)) return false;
return Boolean(String(item.path ?? item.file ?? item.relativePath ?? '').trim());
});
}
export function createPostAgentRunsHandler({
userAuth,
agentRunGateway,
codeRunsEnabled = envFlag(process.env.MEMIND_AGENT_CODE_RUNS_ENABLED),
codeRunUserIds = parseUserIdSet(process.env.MEMIND_AGENT_CODE_RUNS_USER_IDS),
codeRunTaskTypes = parseTaskTypeSet(process.env.MEMIND_AGENT_CODE_RUN_TASK_TYPES),
requireCodeRunValidation = envFlag(process.env.MEMIND_AGENT_CODE_RUNS_REQUIRE_VALIDATION),
}) {
return async function postAgentRuns(request, response) {
try {
const sessionId = String(request.body?.session_id ?? '').trim() || null;
const requestId = String(request.body?.request_id ?? '').trim();
const userMessage = request.body?.user_message ?? null;
const rawToolMode = request.body?.tool_mode ?? request.body?.toolMode ?? 'chat';
const taskType = String(request.body?.task_type ?? request.body?.taskType ?? '').trim() || null;
if (!requestId) {
response.status(400).json({ message: '缺少 request_id' });
return;
}
if (!userMessage) {
response.status(400).json({ message: '缺少 user_message' });
return;
}
let toolMode = 'chat';
try {
toolMode = normalizeAgentRunToolMode(rawToolMode);
} catch (err) {
response.status(400).json({
message: err instanceof Error ? err.message : '不支持的 tool_mode',
});
return;
}
if (toolMode === 'code' && !codeRunsEnabled) {
response.status(403).json({ message: '代码任务灰度未开启' });
return;
}
if (
toolMode === 'code' &&
codeRunUserIds.size > 0 &&
!codeRunUserIds.has(request.currentUser.id)
) {
response.status(403).json({ message: '当前用户未开启代码任务灰度' });
return;
}
if (
toolMode === 'code' &&
codeRunTaskTypes.size > 0 &&
(!taskType || !codeRunTaskTypes.has(taskType.toLowerCase()))
) {
response.status(403).json({ message: '当前代码任务类型未开启灰度' });
return;
}
if (toolMode === 'code' && requireCodeRunValidation && !hasExpectedFileValidation(userMessage)) {
response.status(400).json({ message: '代码任务必须声明产物校验规则' });
return;
}
if (sessionId) {
const owns = await userAuth.ownsSession(request.currentUser.id, sessionId);
if (!owns) {
response.status(403).json({ message: '无权访问该会话' });
return;
}
}
const run = await agentRunGateway.createRun(request.currentUser.id, {
sessionId,
requestId,
userMessage,
toolMode,
taskType,
});
response.status(202).json({ run });
} catch (err) {
response.status(500).json({
message: err instanceof Error ? err.message : '创建任务失败',
});
}
};
}
export function createGetAgentRunHandler({ agentRunGateway }) {
return async function getAgentRun(request, response) {
try {
const run = await agentRunGateway.getRunForUser(
request.currentUser.id,
request.params.runId,
);
if (!run) {
response.status(404).json({ message: '任务不存在' });
return;
}
if (run.status !== 'succeeded' && run.status !== 'failed') {
agentRunGateway.dispatchRun(run.id);
}
response.json({ run });
} catch (err) {
response.status(500).json({
message: err instanceof Error ? err.message : '读取任务失败',
});
}
};
}
export function createAgentRunEventsHandler({
agentRunGateway,
pollIntervalMs = 1000,
keepaliveIntervalMs = 20000,
}) {
return async function getAgentRunEvents(request, response) {
const userId = request.currentUser.id;
const runId = request.params.runId;
let closed = false;
let lastPayload = null;
let pollTimer = null;
let keepaliveTimer = null;
const cleanup = () => {
closed = true;
if (pollTimer) clearInterval(pollTimer);
if (keepaliveTimer) clearInterval(keepaliveTimer);
pollTimer = null;
keepaliveTimer = null;
};
const sendEvent = (event, data) => {
response.write(`event: ${event}\n`);
response.write(`data: ${JSON.stringify(data)}\n\n`);
};
const publish = async (prefetchedRun = null) => {
if (closed) return;
try {
const run = prefetchedRun ?? await agentRunGateway.getRunForUser(userId, runId);
if (!run) {
sendEvent('error', { message: '任务不存在' });
cleanup();
response.end();
return;
}
const nextPayload = JSON.stringify(run);
if (nextPayload !== lastPayload) {
lastPayload = nextPayload;
sendEvent('run', { run });
}
if (run.status !== 'succeeded' && run.status !== 'failed') {
agentRunGateway.dispatchRun(run.id);
return;
}
cleanup();
response.end();
} catch (err) {
sendEvent('error', {
message: err instanceof Error ? err.message : '读取任务失败',
});
cleanup();
response.end();
}
};
const firstRun = await agentRunGateway.getRunForUser(userId, runId);
if (!firstRun) {
response.status(404).json({ message: '任务不存在' });
return;
}
response.status(200);
response.setHeader('Content-Type', 'text/event-stream');
response.setHeader('Cache-Control', 'no-cache');
response.setHeader('Connection', 'keep-alive');
request.on('close', cleanup);
keepaliveTimer = setInterval(() => {
if (!closed) response.write(': keepalive\n\n');
}, keepaliveIntervalMs);
pollTimer = setInterval(() => {
void publish();
}, pollIntervalMs);
await publish(firstRun);
};
}