import type Database from 'better-sqlite3'; import { encrypt, decrypt, loadKeyFromEnv } from './crypto.js'; import { McpNotConnectedError, McpTokenExpiredError, type TokenEndpointResponse } from './types.js'; import { logger } from '../logger.js'; export interface SaveTokensInput { userId: string; serverId: string; accessToken: string; refreshToken: string | null; expiresAt: string | null; scope: string | null; } export interface DoRefreshFn { (serverId: string, refreshToken: string): Promise; } interface Row { user_id: string; server_id: string; access_token_enc: Buffer; refresh_token_enc: Buffer | null; expires_at: string | null; scope: string | null; scope_type: string; scope_id: string | null; connected_at: string; updated_at: string; } export function createTokenManager( db: Database.Database, deps: { doRefresh: DoRefreshFn }, ) { // In-process mutex per (userId, serverId) const mutexes = new Map>(); async function withMutex(key: string, fn: () => Promise): Promise { const prev = mutexes.get(key) ?? Promise.resolve(); const next = prev.catch(() => undefined).then(fn); mutexes.set(key, next); try { return await next; } finally { if (mutexes.get(key) === next) mutexes.delete(key); } } function getRow(userId: string, serverId: string): Row | null { return (db .prepare('SELECT * FROM user_mcp_tokens WHERE user_id=? AND server_id=?') .get(userId, serverId) as Row | undefined) ?? null; } interface ServerAuthRow { auth_kind: string; static_token_enc: Buffer | null; } function getServerAuthRow(serverId: string): ServerAuthRow | null { return (db .prepare('SELECT auth_kind, static_token_enc FROM mcp_servers WHERE id = ?') .get(serverId) as ServerAuthRow | undefined) ?? null; } return { saveTokens(input: SaveTokensInput): void { const key = loadKeyFromEnv(); db.prepare( `INSERT INTO user_mcp_tokens (user_id, server_id, access_token_enc, refresh_token_enc, expires_at, scope, scope_type, scope_id, connected_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, 'user', NULL, datetime('now'), datetime('now')) ON CONFLICT(user_id, server_id) DO UPDATE SET access_token_enc=excluded.access_token_enc, refresh_token_enc=excluded.refresh_token_enc, expires_at=excluded.expires_at, scope=excluded.scope, updated_at=datetime('now')`, ).run( input.userId, input.serverId, encrypt(input.accessToken, key), input.refreshToken ? encrypt(input.refreshToken, key) : null, input.expiresAt, input.scope, ); }, hasToken(userId: string, serverId: string): boolean { const server = getServerAuthRow(serverId); if (!server) return false; if (server.auth_kind === 'api_key') { // For api_key servers, a token is "present" iff static_token_enc is non-null. // (NULL means the server was stored without a token, which shouldn't happen // due to registry validation, but we guard defensively.) return server.static_token_enc !== null; } // oauth path: check user_mcp_tokens row return getRow(userId, serverId) !== null; }, deleteToken(userId: string, serverId: string): void { db.prepare('DELETE FROM user_mcp_tokens WHERE user_id=? AND server_id=?').run(userId, serverId); }, async getValidToken(userId: string, serverId: string): Promise { const key = loadKeyFromEnv(); // Check auth_kind first — api_key servers skip user_mcp_tokens entirely. const server = getServerAuthRow(serverId); if (!server) throw new McpNotConnectedError(serverId); if (server.auth_kind === 'api_key') { if (!server.static_token_enc) throw new McpNotConnectedError(serverId); return decrypt(server.static_token_enc, key); } const row = getRow(userId, serverId); if (!row) throw new McpNotConnectedError(serverId); if (row.expires_at && new Date(row.expires_at).getTime() > Date.now() + 30_000) { return decrypt(row.access_token_enc, key); } if (!row.refresh_token_enc) { db.prepare('DELETE FROM user_mcp_tokens WHERE user_id=? AND server_id=?').run(userId, serverId); throw new McpTokenExpiredError(serverId); } return withMutex(`${userId}:${serverId}`, async () => { const latest = getRow(userId, serverId); if (!latest) throw new McpNotConnectedError(serverId); if (latest.expires_at && new Date(latest.expires_at).getTime() > Date.now() + 30_000) { return decrypt(latest.access_token_enc, key); } if (!latest.refresh_token_enc) { db.prepare('DELETE FROM user_mcp_tokens WHERE user_id=? AND server_id=?').run(userId, serverId); throw new McpTokenExpiredError(serverId); } const refreshPlain = decrypt(latest.refresh_token_enc, key); let resp: TokenEndpointResponse; try { resp = await deps.doRefresh(serverId, refreshPlain); } catch (err) { const code = (err as { code?: string }).code; if (code === 'invalid_grant') { db.prepare('DELETE FROM user_mcp_tokens WHERE user_id=? AND server_id=?').run(userId, serverId); logger.warn(`[mcp:token] invalid_grant, cleared tokens user=${userId} server=${serverId}`); } throw err; } const newExpiresAt = resp.expires_in ? new Date(Date.now() + resp.expires_in * 1000).toISOString() : null; const newRefresh = resp.refresh_token ?? refreshPlain; const result = db.prepare( `UPDATE user_mcp_tokens SET access_token_enc=?, refresh_token_enc=?, expires_at=?, updated_at=datetime('now') WHERE user_id=? AND server_id=? AND access_token_enc=?`, ).run( encrypt(resp.access_token, key), encrypt(newRefresh, key), newExpiresAt, userId, serverId, latest.access_token_enc, ); if (result.changes === 0) { // Another worker updated first — re-read. const fresh = getRow(userId, serverId); if (!fresh) throw new McpNotConnectedError(serverId); return decrypt(fresh.access_token_enc, key); } return resp.access_token; }); }, }; } export type McpTokenManager = ReturnType;