feat: configure network access policies through environment

This commit is contained in:
陈煜 committed 2026-10-04 11:27:34 +08:00
1 parent 91365ee315
commit 265f28e16d
17 files changed
+456 -58

No files matched your search

+12 -15
View File
@@ -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
View File
@@ -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(() => {
+16
View File
@@ -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());
}
+7 -10
View File
@@ -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');
+28 -12
View File
@@ -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;
}
+57
View File
@@ -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());
}