Files
maestro/src/mcp/token-manager.ts
T

183 lines
6.5 KiB
TypeScript

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<TokenEndpointResponse>;
}
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<string, Promise<unknown>>();
async function withMutex<T>(key: string, fn: () => Promise<T>): Promise<T> {
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<string> {
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<typeof createTokenManager>;