183 lines
6.5 KiB
TypeScript
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>;
|