import { Readable } from 'node:stream'; import { Agent, fetch as undiciFetch } from 'undici'; import { appendBalanceEvent, createSseBillingTransform } from './sse-billing.mjs'; import { developerToolsFromPolicy } from './capabilities.mjs'; import { evaluateProxyRequest, isNativeH5ApiPath } from './policies.mjs'; import { buildSandboxSessionConstraints } from './user-publish.mjs'; import { reconcileAgentSession } from './session-reconcile.mjs'; const insecureDispatcher = new Agent({ connect: { rejectUnauthorized: false }, }); function isHttpsTarget(target) { return target.startsWith('https://'); } async function apiFetch(target, apiSecret, pathname, init = {}) { const url = new URL(pathname, target); const headers = { ...(init.headers ?? {}), 'X-Secret-Key': apiSecret, }; if (init.body && !headers['Content-Type']) { headers['Content-Type'] = 'application/json'; } return undiciFetch(url, { ...init, headers, dispatcher: isHttpsTarget(target) ? insecureDispatcher : undefined, }); } async function readJsonBody(req) { if (req.method === 'GET' || req.method === 'HEAD') return null; const chunks = []; for await (const chunk of req) { chunks.push(chunk); } if (chunks.length === 0) return null; const raw = Buffer.concat(chunks).toString('utf8'); if (!raw.trim()) return null; return JSON.parse(raw); } function sendProxyResponse(res, upstream) { res.status(upstream.status); upstream.headers.forEach((value, key) => { if (key === 'transfer-encoding') return; res.setHeader(key, value); }); if (!upstream.body) { res.end(); return; } Readable.fromWeb(upstream.body).pipe(res); } function extractSessionId(req, body) { const fromParams = req.params?.sessionId ?? req.params?.id; if (fromParams) return fromParams; if (body?.session_id) return body.session_id; if (body?.sessionId) return body.sessionId; return null; } export function createTkmindProxy({ apiTarget, apiTargets, apiSecret, userAuth, llmProviderService }) { const targets = apiTargets?.length ? apiTargets : apiTarget ? [apiTarget] : []; const primaryTarget = targets[0] ?? apiTarget ?? ''; let rrIdx = 0; async function targetHealthy(target) { try { const upstream = await apiFetch(target, apiSecret, '/status', { method: 'GET', signal: AbortSignal.timeout(1500), }); return upstream.ok; } catch { return false; } } async function pickTarget() { if (targets.length <= 1) return primaryTarget; for (let i = 0; i < targets.length; i += 1) { const target = targets[rrIdx]; rrIdx = (rrIdx + 1) % targets.length; if (await targetHealthy(target)) return target; } return primaryTarget; } async function resolveTarget(sessionId) { if (targets.length <= 1 || !sessionId) return primaryTarget; try { const node = await userAuth.getSessionNode(sessionId); return targets[node] ?? primaryTarget; } catch { return primaryTarget; } } async function applySessionLlmProvider(sessionId) { if (!llmProviderService || !sessionId) return null; try { const target = await resolveTarget(sessionId); return await llmProviderService.applyBestProviderForSession( sessionId, (url, init) => apiFetch(target, apiSecret, `${url.pathname}${url.search}`, init), ); } catch (err) { console.warn( 'LLM provider apply skipped:', err instanceof Error ? err.message : err, ); return null; } } async function applyLocalFallbackForSession(sessionId) { if (!llmProviderService || !sessionId) return null; const target = await resolveTarget(sessionId); return llmProviderService.applyLocalFallbackForSession( sessionId, (url, init) => apiFetch(target, apiSecret, `${url.pathname}${url.search}`, init), ); } const requireUser = async (req, res, next) => { try { const session = req.userSession; if (!session) { res.status(401).json({ message: '未登录' }); return; } const me = await userAuth.getMe(req.userToken); if (!me) { res.status(401).json({ message: '登录已过期' }); return; } req.currentUser = me; next(); } catch (err) { res.status(500).json({ message: err instanceof Error ? err.message : '认证失败' }); } }; const ensureChatAllowed = async (req, res, next) => { const gate = await userAuth.canUseChat(req.currentUser.id); if (!gate.ok) { res.status(402).json({ message: gate.message, code: gate.code, balanceCents: gate.balanceCents, minRechargeCents: gate.minRechargeCents, suggestedTiers: gate.suggestedTiers, }); return; } next(); }; const handlers = { 'POST /agent/start': [ requireUser, ensureChatAllowed, async (req, res) => { try { const workingDir = await userAuth.resolveWorkingDir(req.currentUser.id); const sessionPolicy = await userAuth.getAgentSessionPolicy(req.currentUser.id); const startTarget = await pickTarget(); const upstream = await apiFetch(startTarget, apiSecret, '/agent/start', { method: 'POST', body: JSON.stringify({ working_dir: workingDir, enable_context_memory: sessionPolicy.enableContextMemory, ...(sessionPolicy.extensionOverrides ? { extension_overrides: sessionPolicy.extensionOverrides } : {}), ...(req.body?.recipe ? { recipe: req.body.recipe } : {}), }), }); const text = await upstream.text(); if (!upstream.ok) { res.status(upstream.status).send(text); return; } const session = JSON.parse(text); if (session?.id) { await userAuth.registerAgentSession( req.currentUser.id, session.id, Math.max(0, targets.indexOf(startTarget)), ); if (sessionPolicy.gooseMode) { const modeRes = await apiFetch(startTarget, apiSecret, '/agent/update_session', { method: 'POST', body: JSON.stringify({ session_id: session.id, goose_mode: sessionPolicy.gooseMode, }), }); if (!modeRes.ok) { const modeText = await modeRes.text().catch(() => ''); res.status(modeRes.status).send(modeText || '设置会话模式失败'); return; } } const publishLayout = await userAuth.getUserPublishLayout(req.currentUser.id); const api = (pathname, init) => apiFetch(startTarget, apiSecret, pathname, init); try { await reconcileAgentSession(api, session.id, { workingDir, sessionPolicy, sandboxConstraints: publishLayout?.constraints ?? null, userContext: publishLayout ? { userId: req.currentUser.id, displayName: publishLayout.displayName, username: publishLayout.username, slug: publishLayout.slug, } : null, }); } catch (reconcileErr) { res.status(500).json({ message: reconcileErr instanceof Error ? `会话策略同步失败:${reconcileErr.message}` : '会话策略同步失败', }); return; } await applySessionLlmProvider(session.id); } res.status(upstream.status).json(session); } catch (err) { res.status(500).json({ message: err instanceof Error ? err.message : '启动会话失败' }); } }, ], 'POST /agent/resume': [ requireUser, ensureChatAllowed, async (req, res) => { try { const sessionId = req.body?.session_id; if (!sessionId) { res.status(400).json({ message: '缺少 session_id' }); return; } const owns = await userAuth.ownsSession(req.currentUser.id, sessionId); if (!owns) { res.status(403).json({ message: '无权访问该会话' }); return; } const resumeTarget = await resolveTarget(sessionId); const upstream = await apiFetch(resumeTarget, apiSecret, '/agent/resume', { method: 'POST', body: JSON.stringify(req.body ?? {}), }); const text = await upstream.text(); if (!upstream.ok) { res.status(upstream.status).send(text); return; } const payload = JSON.parse(text); const skipReconcile = req.body?.skip_reconcile === true; if (!skipReconcile) { const workingDir = await userAuth.resolveWorkingDir(req.currentUser.id); const sessionPolicy = await userAuth.getAgentSessionPolicy(req.currentUser.id); const publishLayout = await userAuth.getUserPublishLayout(req.currentUser.id); try { await reconcileAgentSession( (pathname, init) => apiFetch(resumeTarget, apiSecret, pathname, init), sessionId, { workingDir, sessionPolicy, sandboxConstraints: publishLayout?.constraints ?? null, tolerateInvalidWorkingDir: true, userContext: publishLayout ? { userId: req.currentUser.id, displayName: publishLayout.displayName, username: publishLayout.username, slug: publishLayout.slug, } : null, }, ); } catch (reconcileErr) { res.status(500).json({ message: reconcileErr instanceof Error ? `会话恢复后策略同步失败:${reconcileErr.message}` : '会话恢复后策略同步失败', }); return; } } await applySessionLlmProvider(sessionId); res.status(upstream.status).json(payload); } catch (err) { res.status(500).json({ message: err instanceof Error ? err.message : '恢复会话失败' }); } }, ], 'GET /sessions': [ requireUser, async (req, res) => { try { const owned = await userAuth.listOwnedSessionIds(req.currentUser.id); const sessionsById = new Map(); let healthyTargets = 0; let lastFailure = null; for (const target of targets) { try { const upstream = await apiFetch(target, apiSecret, '/sessions', { method: 'GET', signal: AbortSignal.timeout(3000), }); const text = await upstream.text(); if (!upstream.ok) { lastFailure = text || `upstream ${upstream.status}`; continue; } healthyTargets += 1; const payload = JSON.parse(text); for (const item of payload.sessions ?? []) { if (owned.has(item.id)) sessionsById.set(item.id, item); } } catch (err) { lastFailure = err instanceof Error ? err.message : '读取会话失败'; } } if (healthyTargets === 0) { res.status(502).json({ message: lastFailure ?? '所有 goose 服务不可用' }); return; } if (healthyTargets < targets.length) { res.setHeader('X-TKMind-Degraded', '1'); } res.json({ sessions: [...sessionsById.values()] }); } catch (err) { res.status(500).json({ message: err instanceof Error ? err.message : '读取会话失败' }); } }, ], }; const sessionScoped = (build) => [ requireUser, async (req, res, next) => { try { const body = req.body ?? (await readJsonBody(req)); req.body = body; const sessionId = extractSessionId(req, body); if (!sessionId) { res.status(400).json({ message: '缺少 session_id' }); return; } const owns = await userAuth.ownsSession(req.currentUser.id, sessionId); if (!owns) { res.status(403).json({ message: '无权访问该会话' }); return; } req.agentSessionId = sessionId; await build(req, res, next); } catch (err) { res.status(500).json({ message: err instanceof Error ? err.message : '请求失败' }); } }, ]; const proxySessionEvents = async (req, res, sessionId) => { try { const pathname = `/sessions/${sessionId}/events`; const sessionTarget = await resolveTarget(sessionId); const upstream = await apiFetch(sessionTarget, apiSecret, pathname, { method: 'GET', headers: { Accept: 'text/event-stream', 'Last-Event-ID': req.get('last-event-id') ?? '', }, }); if (!upstream.ok || !upstream.body) { const text = await upstream.text().catch(() => ''); res.status(upstream.status).send(text); return; } res.status(upstream.status); res.setHeader('Content-Type', 'text/event-stream'); res.setHeader('Cache-Control', 'no-cache'); res.setHeader('Connection', 'keep-alive'); let pendingBalance = null; const billingTransform = createSseBillingTransform({ onFinish: async (event) => { const result = await userAuth.billSessionUsage( req.currentUser.id, sessionId, event.token_state, null, ); if (result.ok && result.costCents > 0 && result.balanceCents != null) { pendingBalance = { balanceCents: result.balanceCents, tokensUsed: result.tokensUsed ?? undefined, lastUsage: { inputTokens: result.deltaInputTokens ?? 0, outputTokens: result.deltaOutputTokens ?? 0, costCents: result.costCents, }, }; } }, }); const source = Readable.fromWeb(upstream.body); billingTransform.on('data', (chunk) => { res.write(chunk); if (pendingBalance != null) { res.write(appendBalanceEvent(pendingBalance)); pendingBalance = null; } }); billingTransform.on('end', () => res.end()); billingTransform.on('error', () => res.end()); source.on('error', () => res.end()); source.pipe(billingTransform); } catch (err) { res.status(502).json({ message: err instanceof Error ? err.message : 'SSE 代理失败' }); } }; const proxyFallback = async (req, res) => { try { const pathname = req.originalUrl.replace(/^\/api/, '') || '/'; if (isNativeH5ApiPath(pathname)) { res.status(404).json({ message: 'H5 本地接口未找到,请确认服务端已更新并重启', code: 'not_found', }); return; } const policyState = await userAuth.resolveUserPolicies(req.currentUser); const capabilityState = await userAuth.resolveUserCapabilities(req.currentUser); const gate = evaluateProxyRequest(req.method, pathname, policyState.policies, { unrestricted: policyState.unrestricted, }); if (!gate.allowed) { res.status(403).json({ message: gate.reason ?? '该 API 已被策略禁止' }); return; } if ( /^\/agent\/harness_(bootstrap|remember)$/.test(pathname) && !capabilityState.unrestricted && !capabilityState.capabilities.context_memory ) { res.status(403).json({ message: '当前账户未开通项目记忆,无法访问该 API' }); return; } const body = req.method === 'GET' || req.method === 'HEAD' ? undefined : req.body ? JSON.stringify(req.body) : undefined; const sessionMatch = pathname.match(/^\/sessions\/([^/]+)/); const fallbackTarget = req.goosedTarget ?? (sessionMatch ? await resolveTarget(sessionMatch[1]) : primaryTarget); const upstream = await apiFetch(fallbackTarget, apiSecret, pathname, { method: req.method, body, headers: { Accept: req.get('accept') ?? '*/*', 'Last-Event-ID': req.get('last-event-id') ?? '', }, }); sendProxyResponse(res, upstream); } catch (err) { res.status(502).json({ message: err instanceof Error ? err.message : '代理失败' }); } }; return { requireUser, ensureChatAllowed, applySessionLlmProvider, applyLocalFallbackForSession, handlers, sessionScoped, proxyFallback, proxySessionEvents, resolveTarget, apiFetch: async (pathname, init) => apiFetch(await pickTarget(), apiSecret, pathname, init), apiFetchTo: (target, pathname, init) => apiFetch(target, apiSecret, pathname, init), }; } export function matchHandler(handlers, method, path) { const key = `${method.toUpperCase()} ${path}`; return handlers[key] ?? null; }