feat: configure network access policies through environment
This commit is contained in:
1 parent
91365ee315
commit
265f28e16d
17 files changed
+456
-58
No files matched your search
+12
-15
@@ -1,3 +1,4 @@
|
||||
import { networkConfig, isNetworkOriginAllowed } from './network';
|
||||
import {
|
||||
Injectable,
|
||||
Controller,
|
||||
@@ -36,30 +37,26 @@ export function allowedOrigin(
|
||||
configured: string | undefined,
|
||||
production = false,
|
||||
) {
|
||||
if (configured !== '*') return !!origin && origin === configured;
|
||||
if (production || !origin) return false;
|
||||
try {
|
||||
const url = new URL(origin);
|
||||
return ['http:', 'https:'].includes(url.protocol) && url.origin === origin;
|
||||
} catch {
|
||||
return false;
|
||||
}
|
||||
return isNetworkOriginAllowed(origin, configured, !production);
|
||||
}
|
||||
@Injectable()
|
||||
export class AuthService {
|
||||
private attempts = new Map<string, { count: number; until: number }>();
|
||||
constructor(private db: Database) {}
|
||||
limit(req: Request) {
|
||||
const network = networkConfig();
|
||||
if (!network.rateLimitEnabled) return;
|
||||
const key =
|
||||
(req.ip || 'local') +
|
||||
('userId' in req && typeof req.userId === 'string' ? ':' + req.userId : ''),
|
||||
now = Date.now();
|
||||
let v = this.attempts.get(key);
|
||||
if (!v || v.until < now) {
|
||||
v = { count: 0, until: now + 900000 };
|
||||
v = { count: 0, until: now + network.rateLimitWindowMs };
|
||||
this.attempts.set(key, v);
|
||||
}
|
||||
if (++v.count > 30) throw new HttpException('尝试过于频繁,请 15 分钟后重试', 429);
|
||||
if (++v.count > network.authRateLimitMax)
|
||||
throw new HttpException('尝试过于频繁,请稍后重试', 429);
|
||||
if (this.attempts.size > 10000) {
|
||||
for (const [k, v] of this.attempts) if (v.until < now) this.attempts.delete(k);
|
||||
if (this.attempts.size > 10000) throw new HttpException('服务繁忙,请稍后重试', 429);
|
||||
@@ -75,8 +72,8 @@ export class AuthService {
|
||||
cookie(token: string, expiresAt: Date, res: Response) {
|
||||
res.cookie('wp_session', token, {
|
||||
httpOnly: true,
|
||||
sameSite: 'strict',
|
||||
secure: process.env.COOKIE_SECURE === 'true',
|
||||
sameSite: networkConfig().sameSite,
|
||||
secure: networkConfig().cookieSecure,
|
||||
expires: expiresAt,
|
||||
path: '/api',
|
||||
});
|
||||
@@ -100,8 +97,8 @@ export class AuthService {
|
||||
await this.db.session.deleteMany({ where: { id: digest(req.cookies.wp_session) } });
|
||||
res.clearCookie('wp_session', {
|
||||
path: '/api',
|
||||
sameSite: 'strict',
|
||||
secure: process.env.COOKIE_SECURE === 'true',
|
||||
sameSite: networkConfig().sameSite,
|
||||
secure: networkConfig().cookieSecure,
|
||||
httpOnly: true,
|
||||
});
|
||||
}
|
||||
@@ -119,7 +116,7 @@ export class AuthGuard implements CanActivate {
|
||||
!allowedOrigin(
|
||||
req.headers.origin,
|
||||
process.env.WEB_ORIGIN,
|
||||
process.env.NODE_ENV === 'production',
|
||||
!networkConfig().allowWildcardOrigins,
|
||||
)
|
||||
)
|
||||
throw new ForbiddenException('请求来源不受信任');
|
||||
|
||||
+26
-4
@@ -1,3 +1,4 @@
|
||||
import { networkConfig, isNetworkHostAllowed } from './network';
|
||||
import 'reflect-metadata';
|
||||
import 'dotenv/config';
|
||||
import { setupOpenApi } from './openapi';
|
||||
@@ -92,12 +93,33 @@ class AppModule {}
|
||||
async function bootstrap() {
|
||||
if (!process.env.DATABASE_URL || !process.env.WEB_ORIGIN)
|
||||
throw Error('Missing local environment configuration');
|
||||
if (process.env.NODE_ENV === 'production' && process.env.WEB_ORIGIN === '*')
|
||||
const network = networkConfig();
|
||||
if (
|
||||
!network.allowWildcardOrigins &&
|
||||
process.env.WEB_ORIGIN.split(',').some((v) => v.trim() === '*')
|
||||
)
|
||||
throw Error('Production requires an explicit web origin');
|
||||
if (process.env.NODE_ENV === 'production' && process.env.COOKIE_SECURE !== 'true')
|
||||
if (network.requireSecureCookie && !network.cookieSecure)
|
||||
throw Error('Production requires secure cookies');
|
||||
const app = await NestFactory.create(AppModule, { logger: false, bodyParser: false });
|
||||
app.use(helmet());
|
||||
app.use(
|
||||
helmet({
|
||||
strictTransportSecurity: network.hsts ? undefined : false,
|
||||
contentSecurityPolicy: {
|
||||
directives: { 'upgrade-insecure-requests': network.upgradeInsecureRequests ? [] : null },
|
||||
},
|
||||
}),
|
||||
);
|
||||
app.use((req: { headers: { host?: string } }, res: any, next: () => void) => {
|
||||
if (!isNetworkHostAllowed(req.headers.host, network.apiAllowedHosts))
|
||||
return res.status(403).json({ message: '请求 Host 不受信任' });
|
||||
next();
|
||||
});
|
||||
const origins = process.env.WEB_ORIGIN.split(',').map((v) => v.trim());
|
||||
app.enableCors({
|
||||
origin: network.allowWildcardOrigins && origins.includes('*') ? true : origins,
|
||||
credentials: true,
|
||||
});
|
||||
app.use(json({ limit: '8mb' }));
|
||||
app.use(cookieParser());
|
||||
app.use((_req: unknown, res: { setHeader: (k: string, v: string) => void }, next: () => void) => {
|
||||
@@ -108,7 +130,7 @@ async function bootstrap() {
|
||||
app.useGlobalFilters(new SafeErrors());
|
||||
setupOpenApi(app);
|
||||
app.enableShutdownHooks();
|
||||
await app.listen(Number(process.env.PORT || 3100), '0.0.0.0');
|
||||
await app.listen(Number(process.env.PORT || 3100), network.apiHost);
|
||||
console.log('WorthPath API ready');
|
||||
}
|
||||
void bootstrap().catch(() => {
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
/** Explicit entries match Host including its port; an omitted list preserves defaults. */
|
||||
export function isAllowedMcpHost(
|
||||
host: string | undefined,
|
||||
resource: URL,
|
||||
configured: string | undefined,
|
||||
production: boolean,
|
||||
): boolean {
|
||||
if (!host) return false;
|
||||
const defaults = [resource.host];
|
||||
if (!production)
|
||||
defaults.push('127.0.0.1:' + resource.port, 'localhost:' + resource.port);
|
||||
const allowed = configured?.trim()
|
||||
? configured.split(',').map((value) => value.trim().toLowerCase()).filter(Boolean)
|
||||
: defaults.map((value) => value.toLowerCase());
|
||||
return allowed.includes('*') || allowed.includes(host.toLowerCase());
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
import { networkConfig } from '../network';
|
||||
import { Injectable, BadRequestException, ForbiddenException } from '@nestjs/common';
|
||||
import { randomBytes, randomUUID, createHash } from 'node:crypto';
|
||||
import { Response } from 'express';
|
||||
@@ -55,17 +56,11 @@ export function urls() {
|
||||
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)
|
||||
)
|
||||
!(networkConfig().allowHttp && resource.protocol === 'http:')
|
||||
)
|
||||
throw Error('MCP requires HTTPS except local development');
|
||||
throw Error('MCP HTTP is disabled; enable NETWORK_ALLOW_HTTP or use HTTPS');
|
||||
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))
|
||||
)
|
||||
if (web.protocol !== 'https:' && !(networkConfig().allowHttp && web.protocol === 'http:'))
|
||||
throw Error('MCP web confirmation requires HTTPS');
|
||||
return { resource, issuer: new URL(resource.origin), web };
|
||||
}
|
||||
@@ -102,7 +97,9 @@ export class AgentOAuth implements OAuthServerProvider {
|
||||
u.password ||
|
||||
!(
|
||||
u.protocol === 'https:' ||
|
||||
(u.protocol === 'http:' && ['127.0.0.1', 'localhost', '[::1]'].includes(u.hostname))
|
||||
(u.protocol === 'http:' &&
|
||||
(networkConfig().allowHttpRedirects ||
|
||||
['127.0.0.1', 'localhost', '[::1]'].includes(u.hostname)))
|
||||
)
|
||||
)
|
||||
throw new InvalidClientMetadataError('HTTPS or loopback redirect required');
|
||||
|
||||
@@ -1,12 +1,9 @@
|
||||
import { networkConfig, isNetworkOriginAllowed } from '../network';
|
||||
import { Injectable } from '@nestjs/common';
|
||||
import { Express } from 'express';
|
||||
import { z } from 'zod';
|
||||
import { McpServer } from '@modelcontextprotocol/sdk/server/mcp.js';
|
||||
import { StreamableHTTPServerTransport } from '@modelcontextprotocol/sdk/server/streamableHttp.js';
|
||||
import {
|
||||
mcpAuthRouter,
|
||||
getOAuthProtectedResourceMetadataUrl,
|
||||
} from '@modelcontextprotocol/sdk/server/auth/router.js';
|
||||
import { requireBearerAuth } from '@modelcontextprotocol/sdk/server/auth/middleware/bearerAuth.js';
|
||||
import { AgentOAuth, urls, scopes } from './oauth';
|
||||
import { AgentOperations, writing, plain } from './operations';
|
||||
@@ -14,6 +11,7 @@ import { AgentFiles } from './files';
|
||||
import { Database } from '../database';
|
||||
import { HttpException } from '@nestjs/common';
|
||||
import { ZodError } from 'zod';
|
||||
import { isAllowedMcpHost } from './hosts';
|
||||
|
||||
export function toolResult(value: unknown) {
|
||||
const data = plain(value);
|
||||
@@ -40,6 +38,15 @@ export class AgentTransport {
|
||||
) {}
|
||||
install(app: Express) {
|
||||
const { issuer, resource } = urls();
|
||||
// The SDK reads this flag when its router module is first loaded.
|
||||
process.env.MCP_DANGEROUSLY_ALLOW_INSECURE_ISSUER_URL =
|
||||
networkConfig().allowHttp && issuer.protocol === 'http:' ? 'true' : 'false';
|
||||
const { mcpAuthRouter, getOAuthProtectedResourceMetadataUrl } =
|
||||
require('@modelcontextprotocol/sdk/server/auth/router.js') as typeof import('@modelcontextprotocol/sdk/server/auth/router.js');
|
||||
const network = networkConfig();
|
||||
const rateLimit = network.rateLimitEnabled
|
||||
? { windowMs: network.rateLimitWindowMs, limit: network.mcpAuthRateLimitMax }
|
||||
: (false as const);
|
||||
app.use(
|
||||
mcpAuthRouter({
|
||||
provider: this.oauth,
|
||||
@@ -47,23 +54,32 @@ export class AgentTransport {
|
||||
resourceServerUrl: resource,
|
||||
scopesSupported: [...scopes],
|
||||
resourceName: 'WorthPath',
|
||||
authorizationOptions: { rateLimit },
|
||||
tokenOptions: { rateLimit },
|
||||
revocationOptions: { rateLimit },
|
||||
clientRegistrationOptions: { rateLimit },
|
||||
}),
|
||||
);
|
||||
this.files.install(app);
|
||||
app.all(
|
||||
'/mcp',
|
||||
(req, res, next) => {
|
||||
const allowed = (process.env.MCP_ALLOWED_ORIGINS || urls().web.origin)
|
||||
.split(',')
|
||||
.map((s) => s.trim());
|
||||
if (req.headers.origin && !allowed.includes(req.headers.origin)) {
|
||||
const allowed = process.env.MCP_ALLOWED_ORIGINS || urls().web.origin;
|
||||
if (
|
||||
req.headers.origin &&
|
||||
!isNetworkOriginAllowed(req.headers.origin, allowed, networkConfig().allowWildcardOrigins)
|
||||
) {
|
||||
res.status(403).json({ error: 'Untrusted origin' });
|
||||
return;
|
||||
}
|
||||
const validHosts = [resource.host];
|
||||
if (process.env.NODE_ENV !== 'production')
|
||||
validHosts.push('127.0.0.1:' + resource.port, 'localhost:' + resource.port);
|
||||
if (!validHosts.includes(req.headers.host || '')) {
|
||||
if (
|
||||
!isAllowedMcpHost(
|
||||
req.headers.host,
|
||||
resource,
|
||||
process.env.MCP_ALLOWED_HOSTS,
|
||||
!networkConfig().allowLoopbackHosts,
|
||||
)
|
||||
) {
|
||||
res.status(403).json({ error: 'Untrusted host' });
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
export function networkFlag(name: string, fallback: boolean): boolean {
|
||||
const value = process.env[name]?.trim().toLowerCase();
|
||||
if (!value) return fallback;
|
||||
if (value === 'true' || value === '1') return true;
|
||||
if (value === 'false' || value === '0') return false;
|
||||
throw Error(name + ' must be true or false');
|
||||
}
|
||||
export function networkNumber(name: string, fallback: number) {
|
||||
const value = Number(process.env[name] || fallback);
|
||||
if (!Number.isSafeInteger(value) || value < 1) throw Error(name + ' must be a positive integer');
|
||||
return value;
|
||||
}
|
||||
export function networkConfig() {
|
||||
const production = process.env.NODE_ENV === 'production';
|
||||
const sameSite = process.env.COOKIE_SAME_SITE || 'strict';
|
||||
if (!['strict', 'lax', 'none'].includes(sameSite)) throw Error('Invalid COOKIE_SAME_SITE');
|
||||
const cookieSecure = networkFlag('COOKIE_SECURE', production);
|
||||
if (sameSite === 'none' && !cookieSecure)
|
||||
throw Error('SameSite=None requires COOKIE_SECURE=true');
|
||||
return {
|
||||
rateLimitEnabled: networkFlag('NETWORK_RATE_LIMIT_ENABLED', true),
|
||||
rateLimitWindowMs: networkNumber('NETWORK_RATE_LIMIT_WINDOW_MS', 900000),
|
||||
authRateLimitMax: networkNumber('NETWORK_AUTH_RATE_LIMIT_MAX', 30),
|
||||
mcpAuthRateLimitMax: networkNumber('MCP_AUTH_RATE_LIMIT_MAX', 100),
|
||||
allowHttp: networkFlag('NETWORK_ALLOW_HTTP', !production),
|
||||
allowWildcardOrigins: networkFlag('NETWORK_ALLOW_WILDCARD_ORIGINS', !production),
|
||||
allowHttpRedirects: networkFlag('NETWORK_ALLOW_HTTP_REDIRECTS', !production),
|
||||
requireSecureCookie: networkFlag('NETWORK_REQUIRE_SECURE_COOKIE', production),
|
||||
hsts: networkFlag('NETWORK_HSTS', production),
|
||||
upgradeInsecureRequests: networkFlag('NETWORK_UPGRADE_INSECURE_REQUESTS', production),
|
||||
allowLoopbackHosts: networkFlag('MCP_ALLOW_LOOPBACK_HOSTS', !production),
|
||||
apiHost: process.env.API_HOST || '0.0.0.0',
|
||||
apiAllowedHosts: process.env.API_ALLOWED_HOSTS || '*',
|
||||
cookieSecure,
|
||||
sameSite: sameSite as 'strict' | 'lax' | 'none',
|
||||
};
|
||||
}
|
||||
export function isNetworkOriginAllowed(
|
||||
origin: string | undefined,
|
||||
configured: string | undefined,
|
||||
wildcard: boolean,
|
||||
): boolean {
|
||||
if (!origin) return false;
|
||||
try {
|
||||
const url = new URL(origin);
|
||||
if (!['http:', 'https:'].includes(url.protocol) || url.origin !== origin) return false;
|
||||
const list = (configured || '').split(',').map((s) => s.trim());
|
||||
return list.includes(origin) || (wildcard && list.includes('*'));
|
||||
} catch {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
export function isNetworkHostAllowed(host: string | undefined, configured: string): boolean {
|
||||
if (!host) return false;
|
||||
const list = configured.split(',').map((s) => s.trim().toLowerCase());
|
||||
return list.includes('*') || list.includes(host.toLowerCase());
|
||||
}
|
||||
Reference in new issue
Block a user