feat: improve selfhosted login (#14502)

fix #13397
fix #14011

#### PR Dependency Tree


* **PR #14502** 👈

This tree was auto-generated by
[Charcoal](https://github.com/danerwilliams/charcoal)

<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Centralized CORS policy with dynamic origin validation applied to
server and realtime connections
* Improved sign-in flows with contextual, localized error hints and
toast notifications
* Centralized network-error normalization and conditional OAuth provider
fetching

* **Bug Fixes**
* Better feedback for self-hosted connection failures and clearer
authentication error handling
* More robust handling of network-related failures with user-friendly
messages
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
This commit is contained in:
DarkSky
2026-02-23 21:23:01 +08:00
committed by GitHub
parent 6d805b302c
commit 91c5869053
7 changed files with 294 additions and 41 deletions
+99
View File
@@ -0,0 +1,99 @@
import { URLHelper } from './helpers';
const DEV_LOOPBACK_PROTOCOLS = new Set(['http:', 'https:']);
const DEV_LOOPBACK_HOSTS = new Set(['localhost', '127.0.0.1', '::1']);
const MOBILE_CLIENT_ORIGINS = new Set([
'https://localhost',
'capacitor://localhost',
'ionic://localhost',
]);
const DESKTOP_CLIENT_ORIGINS = new Set(['assets://.', 'assets://another-host']);
export const CORS_ALLOWED_METHODS = [
'GET',
'HEAD',
'PUT',
'PATCH',
'POST',
'DELETE',
'OPTIONS',
];
export const CORS_ALLOWED_HEADERS = [
'accept',
'authorization',
'content-type',
'x-affine-version',
'x-operation-name',
'x-request-id',
'x-captcha-token',
'x-captcha-challenge',
'x-affine-csrf-token',
'x-requested-with',
'range',
];
export const CORS_EXPOSED_HEADERS = [
'content-length',
'content-range',
'x-request-id',
];
function normalizeHostname(hostname: string) {
return hostname.toLowerCase().replace(/^\[/, '').replace(/\]$/, '');
}
function isDevLoopbackOrigin(origin: string) {
try {
const parsed = new URL(origin);
return (
DEV_LOOPBACK_PROTOCOLS.has(parsed.protocol) &&
DEV_LOOPBACK_HOSTS.has(normalizeHostname(parsed.hostname))
);
} catch {
return false;
}
}
export function buildCorsAllowedOrigins(url: URLHelper) {
return new Set<string>([
...url.allowedOrigins,
...MOBILE_CLIENT_ORIGINS,
...DESKTOP_CLIENT_ORIGINS,
]);
}
export function isCorsOriginAllowed(
origin: string | undefined | null,
allowedOrigins: Set<string>
) {
if (!origin) {
return true;
}
if (allowedOrigins.has(origin)) {
return true;
}
if ((env.dev || env.testing) && isDevLoopbackOrigin(origin)) {
return true;
}
return false;
}
export function corsOriginCallback(
origin: string | undefined,
allowedOrigins: Set<string>,
onBlocked: (origin: string) => void,
callback: (error: Error | null, allow?: boolean) => void
) {
if (isCorsOriginAllowed(origin, allowedOrigins)) {
callback(null, true);
return;
}
const blockedOrigin = origin ?? '<empty>';
onBlocked(blockedOrigin);
callback(null, false);
}
@@ -11,6 +11,7 @@ export {
defineModuleConfig,
type JSONSchema,
} from './config';
export * from './cors';
export * from './error';
export { EventBus, OnEvent } from './event';
export {
@@ -4,7 +4,15 @@ import { createAdapter } from '@socket.io/redis-adapter';
import { Server, Socket } from 'socket.io';
import { Config } from '../config';
import {
buildCorsAllowedOrigins,
CORS_ALLOWED_HEADERS,
CORS_ALLOWED_METHODS,
corsOriginCallback,
} from '../cors';
import { AuthenticationRequired } from '../error';
import { URLHelper } from '../helpers';
import { AFFiNELogger } from '../logger';
import { SocketIoRedis } from '../redis';
import { WEBSOCKET_OPTIONS } from './options';
@@ -14,17 +22,34 @@ export class SocketIoAdapter extends IoAdapter {
}
override createIOServer(port: number, options?: any): Server {
const logger = this.app.get(AFFiNELogger);
const config = this.app.get(WEBSOCKET_OPTIONS) as Config['websocket'] & {
canActivate: (socket: Socket) => Promise<boolean>;
};
const url = this.app.get(URLHelper);
const allowedOrigins = buildCorsAllowedOrigins(url);
const server: Server = super.createIOServer(port, {
...config,
...options,
// Enable CORS for Socket.IO
cors: {
origin: true, // Allow all origins
credentials: true, // Allow credentials (cookies, auth headers)
methods: ['GET', 'POST'],
origin: (
origin: string | undefined,
callback: (error: Error | null, allow?: boolean) => void
) => {
corsOriginCallback(
origin,
allowedOrigins,
blockedOrigin =>
logger.warn(
`Blocked WebSocket CORS request from origin: ${blockedOrigin}`
),
callback
);
},
credentials: true,
methods: CORS_ALLOWED_METHODS,
allowedHeaders: CORS_ALLOWED_HEADERS,
},
});
+27 -4
View File
@@ -5,9 +5,14 @@ import graphqlUploadExpress from 'graphql-upload/graphqlUploadExpress.mjs';
import {
AFFiNELogger,
buildCorsAllowedOrigins,
CacheInterceptor,
CloudThrottlerGuard,
Config,
CORS_ALLOWED_HEADERS,
CORS_ALLOWED_METHODS,
CORS_EXPOSED_HEADERS,
corsOriginCallback,
GlobalExceptionFilter,
URLHelper,
} from './base';
@@ -16,12 +21,11 @@ import { AuthGuard } from './core/auth';
import { serverTimingAndCache } from './middleware/timing';
const OneMB = 1024 * 1024;
export async function run() {
const { AppModule } = await import('./app.module');
const app = await NestFactory.create<NestExpressApplication>(AppModule, {
cors: true,
cors: false,
rawBody: true,
bodyParser: true,
bufferLogs: true,
@@ -32,6 +36,27 @@ export async function run() {
const logger = app.get(AFFiNELogger);
app.useLogger(logger);
const config = app.get(Config);
const url = app.get(URLHelper);
const allowedOrigins = buildCorsAllowedOrigins(url);
app.enableCors({
origin: (origin, callback) => {
corsOriginCallback(
origin,
allowedOrigins,
blockedOrigin =>
logger.warn(`Blocked CORS request from origin: ${blockedOrigin}`),
callback
);
},
credentials: true,
methods: CORS_ALLOWED_METHODS,
allowedHeaders: CORS_ALLOWED_HEADERS,
exposedHeaders: CORS_EXPOSED_HEADERS,
maxAge: 86400,
optionsSuccessStatus: 204,
});
if (config.server.path) {
app.setGlobalPrefix(config.server.path);
@@ -74,8 +99,6 @@ export async function run() {
});
}
const url = app.get(URLHelper);
await app.listen(config.server.port, config.server.listenAddr);
const formattedAddr = config.server.listenAddr.includes(':')