import { Injectable, BadRequestException, ForbiddenException } from '@nestjs/common'; import { randomBytes, randomUUID, createHash } from 'node:crypto'; import { Response } from 'express'; import { z } from 'zod'; import { Database } from '../database'; import { OAuthServerProvider, AuthorizationParams, } from '@modelcontextprotocol/sdk/server/auth/provider.js'; import { OAuthClientInformationFull, OAuthTokens, OAuthTokenRevocationRequest, } from '@modelcontextprotocol/sdk/shared/auth.js'; import { InvalidClientMetadataError, InvalidGrantError, InvalidScopeError, InvalidTokenError, InvalidTargetError, } from '@modelcontextprotocol/sdk/server/auth/errors.js'; export const scopes = ['read', 'draft', 'write', 'hidden_read', 'hidden_write'] as const; export const scopeInput = z .array(z.enum(scopes)) .min(1) .max(4) .refine( (v) => v.includes('read') && new Set(v).size === v.length && !(v.includes('draft') && v.includes('write')) && (!v.includes('hidden_write') || (v.includes('hidden_read') && (v.includes('draft') || v.includes('write')))), ); export const oauthDays = z.union([ z.literal(1), z.literal(3), z.literal(7), z.literal(30), z.literal(365), z.literal(null), ]); export const digest = (s: string) => createHash('sha256').update(s).digest('hex'); const secret = () => randomBytes(32).toString('base64url'); export function urls() { const resource = new URL(process.env.MCP_PUBLIC_URL || 'http://localhost:3100/mcp'); if ( resource.pathname !== '/mcp' || resource.search || resource.hash || resource.username || resource.password ) throw Error('MCP_PUBLIC_URL must be the canonical /mcp URL'); if ( resource.protocol !== 'https:' && !( process.env.NODE_ENV !== 'production' && ['localhost', '127.0.0.1', '[::1]'].includes(resource.hostname) ) ) throw Error('MCP requires HTTPS except local development'); const web = new URL(process.env.MCP_WEB_URL || 'http://localhost:5173'); if ( web.protocol !== 'https:' && !(process.env.NODE_ENV !== 'production' && ['localhost', '127.0.0.1'].includes(web.hostname)) ) throw Error('MCP web confirmation requires HTTPS'); return { resource, issuer: new URL(resource.origin), web }; } export function webLink(key: string, id: string) { const u = new URL(urls().web); u.pathname = u.pathname.replace(/\/$/, '') + (key === 'agent_authorization' ? '/agent/authorize' : '/agent/operation'); u.searchParams.set(key, id); return u.toString(); } @Injectable() export class AgentOAuth implements OAuthServerProvider { constructor(private db: Database) {} get clientsStore() { return { getClient: async (id: string) => { const row = await this.db.agentClient.findUnique({ where: { id } }); return row?.metadata as OAuthClientInformationFull | undefined; }, registerClient: async ( input: Omit, ) => { if (input.token_endpoint_auth_method !== 'none') throw new InvalidClientMetadataError('Only public PKCE clients are supported'); if (!input.redirect_uris.length || input.redirect_uris.length > 10) throw new InvalidClientMetadataError('Invalid redirect URIs'); for (const value of input.redirect_uris) { const u = new URL(value); if ( u.hash || u.username || u.password || !( u.protocol === 'https:' || (u.protocol === 'http:' && ['127.0.0.1', 'localhost', '[::1]'].includes(u.hostname)) ) ) throw new InvalidClientMetadataError('HTTPS or loopback redirect required'); } if ((await this.db.agentClient.count()) >= 10000) throw new InvalidClientMetadataError('Client registration limit reached'); const client = { ...input, client_name: z .string() .min(1) .max(100) .parse(input.client_name || 'MCP client'), client_id: randomUUID(), client_id_issued_at: Math.floor(Date.now() / 1000), }; await this.db.agentClient.create({ data: { id: client.client_id, metadata: JSON.parse(JSON.stringify(client)) }, }); return client; }, }; } private resource(resource?: URL) { if (resource?.toString() !== urls().resource.toString()) throw new InvalidTargetError('WorthPath resource is required'); } async authorize(client: OAuthClientInformationFull, params: AuthorizationParams, res: Response) { this.resource(params.resource); const selected = params.scopes?.length ? params.scopes : ['read']; if ( !z .array(z.enum(scopes)) .min(1) .max(scopes.length) .refine((v) => v.includes('read') && new Set(v).size === v.length) .safeParse(selected).success ) throw new InvalidScopeError('Unsupported scope'); const row = await this.db.agentAuthorization.create({ data: { clientId: client.client_id, parameters: JSON.parse( JSON.stringify({ ...params, resource: params.resource!.toString(), scopes: selected }), ), expiresAt: new Date(Date.now() + 600000), }, }); res.redirect(webLink('agent_authorization', row.id)); } async pending(id: string) { const row = await this.db.agentAuthorization.findUnique({ where: { id } }); if (!row || row.status !== 'pending' || row.expiresAt <= new Date()) throw new BadRequestException('授权请求已失效'); const client = await this.clientsStore.getClient(row.clientId); const p = row.parameters as any; return { id, name: client?.client_name, redirectUri: p.redirectUri, scopes: p.scopes, resource: p.resource, }; } async consent( userId: string, id: string, approved: boolean, selected?: string[], days: number | null = 30, ) { return this.db.atomic(async () => { await this.pending(id); const row = await this.db.agentAuthorization.findUniqueOrThrow({ where: { id } }); const parameters = row.parameters as any; const authorizationDays = oauthDays.parse(days); const allowed = selected || ['read']; scopeInput.parse(allowed); if (allowed.some((scope: string) => !parameters.scopes.includes(scope))) throw new BadRequestException('不能授予客户端未请求的权限'); const code = secret(); const changed = await this.db.agentAuthorization.updateMany({ where: { id, status: 'pending', expiresAt: { gt: new Date() } }, data: { userId, status: approved ? 'approved' : 'denied', codeDigest: approved ? digest(code) : null, parameters: { ...parameters, scopes: allowed, authorizationDays }, }, }); if (!changed.count) throw new BadRequestException('授权请求已处理'); const p = row.parameters as any, callback = new URL(p.redirectUri); callback.searchParams.set(approved ? 'code' : 'error', approved ? code : 'access_denied'); if (p.state) callback.searchParams.set('state', p.state); return { redirect: callback.toString() }; }); } async challengeForAuthorizationCode(client: OAuthClientInformationFull, code: string) { const row = await this.db.agentAuthorization.findUnique({ where: { codeDigest: digest(code) }, }); if ( !row || row.clientId !== client.client_id || row.status !== 'approved' || row.expiresAt <= new Date() ) throw new InvalidGrantError('Invalid authorization code'); return (row.parameters as any).codeChallenge as string; } async issue( userId: string, name: string, selected: string[], days: number | null, clientId?: string, authorizationDays: number | null = 30, ) { if (clientId && days === null) throw new BadRequestException('OAuth 连接必须有期限'); const access = secret(), refresh = clientId ? secret() : undefined, sessionId = digest(secret()); const expiresAt = days === null ? null : new Date(Date.now() + days * 86400000); const lifetime = oauthDays.parse(authorizationDays); const refreshExpiresAt = clientId && lifetime !== null ? new Date(Date.now() + lifetime * 86400000) : null; await this.db.session.create({ data: { id: sessionId, userId, expiresAt: clientId ? refreshExpiresAt || new Date(Date.now() + 30 * 86400000) : expiresAt || new Date(Date.now() + 86400000), }, }); const grant = await this.db.agentGrant.create({ data: { userId, name, clientId, scopes: selected, resource: urls().resource.toString(), accessDigest: digest(access), refreshDigest: refresh ? digest(refresh) : null, expiresAt, refreshExpiresAt, sessionId, }, }); return { grant, tokens: { access_token: access, token_type: 'Bearer', ...(days === null ? {} : { expires_in: Math.floor(days * 86400) }), scope: selected.join(' '), ...(refresh ? { refresh_token: refresh } : {}), } as OAuthTokens, }; } async exchangeAuthorizationCode( client: OAuthClientInformationFull, code: string, verifier?: string, redirectUri?: string, resource?: URL, ) { this.resource(resource); return this.db.atomic(async () => { await this.challengeForAuthorizationCode(client, code); const row = await this.db.agentAuthorization.findUniqueOrThrow({ where: { codeDigest: digest(code) }, }), p = row.parameters as any; // SDK tokenHandler validates S256 PKCE before invoking this provider, and passes // undefined for verifier after successful local validation (skipLocalPkceValidation=false). if ( redirectUri !== p.redirectUri || (verifier && createHash('sha256').update(verifier).digest('base64url') !== p.codeChallenge) ) throw new InvalidGrantError('PKCE or redirect mismatch'); const changed = await this.db.agentAuthorization.updateMany({ where: { id: row.id, status: 'approved' }, data: { status: 'used', codeDigest: null }, }); if (!changed.count) throw new InvalidGrantError('Code already used'); return ( await this.issue( row.userId!, client.client_name || 'MCP client', p.scopes, 1 / 24, client.client_id, p.authorizationDays === undefined ? 30 : p.authorizationDays, ) ).tokens; }); } async exchangeRefreshToken( client: OAuthClientInformationFull, token: string, selected?: string[], resource?: URL, ) { this.resource(resource); return this.db.atomic(async () => { const row = await this.db.agentGrant.findUnique({ where: { refreshDigest: digest(token) } }); if ( !row || row.clientId !== client.client_id || row.revokedAt || (row.refreshExpiresAt && row.refreshExpiresAt <= new Date()) ) throw new InvalidGrantError('Invalid refresh token'); const current = row.scopes as string[]; if ( selected && (!scopeInput.safeParse(selected).success || selected.some((s) => !current.includes(s))) ) throw new InvalidScopeError('Scope escalation rejected'); const access = secret(), refresh = secret(); const accessExpiresAt = new Date( Math.min(Date.now() + 3600000, row.refreshExpiresAt ? +row.refreshExpiresAt : Infinity), ); const changed = await this.db.agentGrant.updateMany({ where: { id: row.id, refreshDigest: digest(token), revokedAt: null }, data: { accessDigest: digest(access), refreshDigest: digest(refresh), expiresAt: accessExpiresAt, scopes: selected || current, }, }); if (!changed.count) throw new InvalidGrantError('Refresh token already used'); // Permanent OAuth keeps a bounded business session, renewed only after a valid refresh. const sessionExpiresAt = row.refreshExpiresAt || new Date(Date.now() + 30 * 86400000); await this.db.session.upsert({ where: { id: row.sessionId }, create: { id: row.sessionId, userId: row.userId, expiresAt: sessionExpiresAt }, update: { expiresAt: sessionExpiresAt, revealUntil: null, backupDigest: null, backupExpiresAt: null, }, }); return { access_token: access, refresh_token: refresh, token_type: 'Bearer', expires_in: Math.max(0, Math.floor((+accessExpiresAt - Date.now()) / 1000)), scope: (selected || current).join(' '), }; }); } async verifyAccessToken(token: string) { if (!/^[A-Za-z0-9_-]{43}$/.test(token)) throw new InvalidTokenError('Invalid token'); const row = await this.db.agentGrant.findUnique({ where: { accessDigest: digest(token) } }); if ( !row || row.revokedAt || (row.clientId && row.refreshExpiresAt && row.refreshExpiresAt <= new Date()) || (row.expiresAt ? row.expiresAt <= new Date() : !!row.clientId) || row.resource !== urls().resource.toString() ) throw new InvalidTokenError('Expired, revoked or invalid resource token'); return { token, clientId: row.clientId || row.id, scopes: row.scopes as string[], // SDK bearer middleware requires a finite verified-authentication expiry. // A permanent PAT stays expiry-free in storage; every request rechecks revocation. expiresAt: Math.floor((row.expiresAt?.getTime() ?? Date.now() + 3600000) / 1000), resource: new URL(row.resource), extra: { grantId: row.id, userId: row.userId }, }; } async revokeToken(client: OAuthClientInformationFull, request: OAuthTokenRevocationRequest) { await this.db.agentGrant.updateMany({ where: { clientId: client.client_id, OR: [{ accessDigest: digest(request.token) }, { refreshDigest: digest(request.token) }], }, data: { revokedAt: new Date() }, }); } async grant(id: string, userId?: string) { const row = await this.db.agentGrant.findFirst({ where: { id, ...(userId ? { userId } : {}), revokedAt: null, AND: [ { OR: [ { clientId: null }, { refreshExpiresAt: null }, { refreshExpiresAt: { gt: new Date() } }, ], }, ], OR: [{ expiresAt: { gt: new Date() } }, { expiresAt: null, clientId: null }], }, }); if (!row) throw new ForbiddenException('Agent 连接已过期或撤销'); return row; } }