feat(auth): add generic OIDC login with mandatory groups
This commit is contained in:
@@ -14,6 +14,7 @@ import type { LoadedAuthConfig } from "./auth/types.js";
|
|||||||
import { createCurrentLocalUserRegistryResolver, type LocalUserRegistry } from "./auth/local-registry.js";
|
import { createCurrentLocalUserRegistryResolver, type LocalUserRegistry } from "./auth/local-registry.js";
|
||||||
import { AuthSessionOperationalError, createFileAuthSessionStore, type AuthSessionStore, type AuthSessionValidity } from "./auth/session-store.js";
|
import { AuthSessionOperationalError, createFileAuthSessionStore, type AuthSessionStore, type AuthSessionValidity } from "./auth/session-store.js";
|
||||||
import { registerAuthRoutes } from "./auth/routes.js";
|
import { registerAuthRoutes } from "./auth/routes.js";
|
||||||
|
import { createOidcProtocol, type OidcProtocol } from "./auth/oidc-client.js";
|
||||||
import { sessionRoutes } from "./routes/sessions.js";
|
import { sessionRoutes } from "./routes/sessions.js";
|
||||||
import { sqlRoutes } from "./routes/sql.js";
|
import { sqlRoutes } from "./routes/sql.js";
|
||||||
import { metaRoutes, type ListModelsFn } from "./routes/meta.js";
|
import { metaRoutes, type ListModelsFn } from "./routes/meta.js";
|
||||||
@@ -48,6 +49,7 @@ export interface BuildAppDeps {
|
|||||||
piManagement?: PiManagementService;
|
piManagement?: PiManagementService;
|
||||||
localUserRegistry?: LocalUserRegistry;
|
localUserRegistry?: LocalUserRegistry;
|
||||||
authSessionStore?: AuthSessionStore;
|
authSessionStore?: AuthSessionStore;
|
||||||
|
oidcProtocol?: OidcProtocol;
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface AppWithAuthSessionStore extends FastifyInstance {
|
export interface AppWithAuthSessionStore extends FastifyInstance {
|
||||||
@@ -192,6 +194,26 @@ export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstanc
|
|||||||
currentAuthConfigRevision: () => loaded.revision,
|
currentAuthConfigRevision: () => loaded.revision,
|
||||||
currentLocalUser: (subject) => localUserForSnapshot(loaded, subject),
|
currentLocalUser: (subject) => localUserForSnapshot(loaded, subject),
|
||||||
});
|
});
|
||||||
|
const resolveOidcProtocol = (loaded: LoadedAuthConfig): OidcProtocol | undefined => {
|
||||||
|
if (deps?.oidcProtocol) return deps.oidcProtocol;
|
||||||
|
if (loaded.value.mode !== "oidc") return undefined;
|
||||||
|
const clientSecret = process.env.THT_OIDC_CLIENT_SECRET;
|
||||||
|
if (typeof clientSecret !== "string" || clientSecret.length === 0 || clientSecret.length > 4096 || /\p{Cc}/u.test(clientSecret)) {
|
||||||
|
return undefined;
|
||||||
|
}
|
||||||
|
try {
|
||||||
|
return createOidcProtocol({
|
||||||
|
issuer: loaded.value.oidc.issuer,
|
||||||
|
clientId: loaded.value.oidc.clientId,
|
||||||
|
clientSecret,
|
||||||
|
callbackUrl: new URL("/api/auth/oidc/callback", loaded.value.publicUrl).href,
|
||||||
|
scopes: loaded.value.oidc.scopes,
|
||||||
|
groupsClaim: loaded.value.oidc.groupsClaim,
|
||||||
|
});
|
||||||
|
} catch {
|
||||||
|
return undefined;
|
||||||
|
}
|
||||||
|
};
|
||||||
const authSessionStore = deps?.authSessionStore ?? (config.authMode === "local" || config.authMode === "oidc"
|
const authSessionStore = deps?.authSessionStore ?? (config.authMode === "local" || config.authMode === "oidc"
|
||||||
? createFileAuthSessionStore(config.authStateRoot, {
|
? createFileAuthSessionStore(config.authStateRoot, {
|
||||||
currentAuthConfigRevision: () => {
|
currentAuthConfigRevision: () => {
|
||||||
@@ -251,6 +273,7 @@ export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstanc
|
|||||||
sessionStore: authSessionStore,
|
sessionStore: authSessionStore,
|
||||||
localUserRegistry: deps?.localUserRegistry,
|
localUserRegistry: deps?.localUserRegistry,
|
||||||
resolveLocalUserRegistry,
|
resolveLocalUserRegistry,
|
||||||
|
resolveOidcProtocol,
|
||||||
});
|
});
|
||||||
sessionRoutes(app, {
|
sessionRoutes(app, {
|
||||||
mgr, tht: tht as ThtRunner, hub, getSettings, readiness, listModels, workspaceRegistry,
|
mgr, tht: tht as ThtRunner, hub, getSettings, readiness, listModels, workspaceRegistry,
|
||||||
|
|||||||
@@ -0,0 +1,294 @@
|
|||||||
|
import {
|
||||||
|
authorizationCodeGrant,
|
||||||
|
buildAuthorizationUrl,
|
||||||
|
calculatePKCECodeChallenge,
|
||||||
|
customFetch,
|
||||||
|
discovery,
|
||||||
|
type Configuration,
|
||||||
|
} from "openid-client";
|
||||||
|
import { constants, createPublicKey, verify as verifySignature } from "node:crypto";
|
||||||
|
|
||||||
|
export interface OidcIdentity {
|
||||||
|
issuer: string;
|
||||||
|
subject: string;
|
||||||
|
displayName?: string;
|
||||||
|
groups: readonly string[];
|
||||||
|
tokenExpiresAt: Date;
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface OidcProtocol {
|
||||||
|
authorizationUrl(input: { state: string; nonce: string; codeVerifier: string }): Promise<URL>;
|
||||||
|
callback(input: { currentUrl: URL; state: string; nonce: string; codeVerifier: string }): Promise<OidcIdentity>;
|
||||||
|
diagnose(signal: AbortSignal): Promise<void>;
|
||||||
|
verifyDeviceFlow?(signal: AbortSignal, present: (uri: string, code: string) => void): Promise<OidcIdentity>;
|
||||||
|
}
|
||||||
|
|
||||||
|
export class OidcProtocolError extends Error {
|
||||||
|
constructor() {
|
||||||
|
super("oidc_protocol_invalid");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface OidcProtocolOptions {
|
||||||
|
issuer: string;
|
||||||
|
clientId: string;
|
||||||
|
clientSecret: string;
|
||||||
|
callbackUrl: string;
|
||||||
|
scopes: readonly string[];
|
||||||
|
groupsClaim: string;
|
||||||
|
fetch?: typeof globalThis.fetch;
|
||||||
|
}
|
||||||
|
|
||||||
|
const MAX_GROUPS = 128;
|
||||||
|
const MAX_GROUP_LENGTH = 256;
|
||||||
|
const MAX_ID_TOKEN_LENGTH = 16 * 1024;
|
||||||
|
const MAX_JWKS_BYTES = 1024 * 1024;
|
||||||
|
const text = (value: unknown, maximum = 2048): value is string =>
|
||||||
|
typeof value === "string" && value.length > 0 && value.length <= maximum && !/\p{Cc}/u.test(value);
|
||||||
|
|
||||||
|
function configuredUrl(value: string): URL {
|
||||||
|
let url: URL;
|
||||||
|
try {
|
||||||
|
url = new URL(value);
|
||||||
|
} catch {
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
if (url.protocol !== "https:" || url.username || url.password || url.search || url.hash) throw new OidcProtocolError();
|
||||||
|
return url;
|
||||||
|
}
|
||||||
|
|
||||||
|
function httpsEndpoint(value: unknown): URL {
|
||||||
|
if (!text(value, 2048)) throw new OidcProtocolError();
|
||||||
|
let url: URL;
|
||||||
|
try {
|
||||||
|
url = new URL(value);
|
||||||
|
} catch {
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
if (url.protocol !== "https:" || url.username || url.password || url.hash) throw new OidcProtocolError();
|
||||||
|
return url;
|
||||||
|
}
|
||||||
|
|
||||||
|
function groupsFromClaims(claims: Record<string, unknown>, name: string): string[] {
|
||||||
|
const indirect = claims._claim_names;
|
||||||
|
if ((indirect && typeof indirect === "object" && !Array.isArray(indirect)
|
||||||
|
&& Object.prototype.hasOwnProperty.call(indirect, name))
|
||||||
|
|| claims.hasgroups === true) throw new OidcProtocolError();
|
||||||
|
const raw = claims[name];
|
||||||
|
if (!Array.isArray(raw) || raw.length > MAX_GROUPS) throw new OidcProtocolError();
|
||||||
|
const groups: string[] = [];
|
||||||
|
const unique = new Set<string>();
|
||||||
|
for (const group of raw) {
|
||||||
|
if (!text(group, MAX_GROUP_LENGTH) || group.trim().length === 0 || unique.has(group)) throw new OidcProtocolError();
|
||||||
|
unique.add(group);
|
||||||
|
groups.push(group);
|
||||||
|
}
|
||||||
|
return groups;
|
||||||
|
}
|
||||||
|
|
||||||
|
function identityFromClaims(claims: Record<string, unknown>, options: OidcProtocolOptions): OidcIdentity {
|
||||||
|
if (claims.iss !== options.issuer || !text(claims.sub, 512)) throw new OidcProtocolError();
|
||||||
|
const audience = claims.aud;
|
||||||
|
if (!(audience === options.clientId || (Array.isArray(audience) && audience.includes(options.clientId)))) {
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
if (typeof claims.exp !== "number" || !Number.isSafeInteger(claims.exp) || claims.exp * 1000 <= Date.now()) {
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
const tokenExpiresAt = new Date(claims.exp * 1000);
|
||||||
|
if (Number.isNaN(tokenExpiresAt.getTime())) throw new OidcProtocolError();
|
||||||
|
return {
|
||||||
|
issuer: options.issuer,
|
||||||
|
subject: claims.sub,
|
||||||
|
...(text(claims.name, 256) ? { displayName: claims.name } : {}),
|
||||||
|
groups: groupsFromClaims(claims, options.groupsClaim),
|
||||||
|
tokenExpiresAt,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
function jsonPart(part: string): Record<string, unknown> {
|
||||||
|
if (!/^[A-Za-z0-9_-]+$/.test(part) || part.length > MAX_ID_TOKEN_LENGTH) throw new OidcProtocolError();
|
||||||
|
try {
|
||||||
|
const value = JSON.parse(new TextDecoder("utf-8", { fatal: true }).decode(Buffer.from(part, "base64url")));
|
||||||
|
if (!value || typeof value !== "object" || Array.isArray(value)) throw new OidcProtocolError();
|
||||||
|
return value as Record<string, unknown>;
|
||||||
|
} catch {
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function signatureAlgorithm(algorithm: string): {
|
||||||
|
digest: string | null;
|
||||||
|
pss?: boolean;
|
||||||
|
saltLength?: number;
|
||||||
|
ecdsaPartLength?: number;
|
||||||
|
} {
|
||||||
|
switch (algorithm) {
|
||||||
|
case "RS256": return { digest: "RSA-SHA256" };
|
||||||
|
case "RS384": return { digest: "RSA-SHA384" };
|
||||||
|
case "RS512": return { digest: "RSA-SHA512" };
|
||||||
|
case "PS256": return { digest: "sha256", pss: true, saltLength: 32 };
|
||||||
|
case "PS384": return { digest: "sha384", pss: true, saltLength: 48 };
|
||||||
|
case "PS512": return { digest: "sha512", pss: true, saltLength: 64 };
|
||||||
|
case "ES256": return { digest: "sha256", ecdsaPartLength: 32 };
|
||||||
|
case "ES384": return { digest: "sha384", ecdsaPartLength: 48 };
|
||||||
|
case "ES512": return { digest: "sha512", ecdsaPartLength: 66 };
|
||||||
|
case "EdDSA": return { digest: null };
|
||||||
|
default: throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function derLength(length: number): Buffer {
|
||||||
|
if (length < 128) return Buffer.from([length]);
|
||||||
|
if (length < 256) return Buffer.from([0x81, length]);
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
|
||||||
|
function derInteger(raw: Buffer): Buffer {
|
||||||
|
let start = 0;
|
||||||
|
while (start < raw.length - 1 && raw[start] === 0) start += 1;
|
||||||
|
let value = raw.subarray(start);
|
||||||
|
if ((value[0] & 0x80) !== 0) value = Buffer.concat([Buffer.from([0]), value]);
|
||||||
|
return Buffer.concat([Buffer.from([0x02]), derLength(value.length), value]);
|
||||||
|
}
|
||||||
|
|
||||||
|
function joseEcdsaSignatureToDer(signature: Buffer, partLength: number): Buffer {
|
||||||
|
if (signature.length !== partLength * 2) throw new OidcProtocolError();
|
||||||
|
const sequence = Buffer.concat([
|
||||||
|
derInteger(signature.subarray(0, partLength)),
|
||||||
|
derInteger(signature.subarray(partLength)),
|
||||||
|
]);
|
||||||
|
return Buffer.concat([Buffer.from([0x30]), derLength(sequence.length), sequence]);
|
||||||
|
}
|
||||||
|
|
||||||
|
async function verifyIdTokenSignature(
|
||||||
|
idToken: unknown,
|
||||||
|
config: Configuration,
|
||||||
|
options: OidcProtocolOptions,
|
||||||
|
): Promise<void> {
|
||||||
|
if (!text(idToken, MAX_ID_TOKEN_LENGTH)) throw new OidcProtocolError();
|
||||||
|
const [protectedPart, payloadPart, signaturePart, extra] = idToken.split(".");
|
||||||
|
if (!protectedPart || !payloadPart || !signaturePart || extra) throw new OidcProtocolError();
|
||||||
|
const header = jsonPart(protectedPart);
|
||||||
|
if (!text(header.alg, 16) || !text(header.kid, 256)) throw new OidcProtocolError();
|
||||||
|
const metadata = config.serverMetadata();
|
||||||
|
if (!Array.isArray(metadata.id_token_signing_alg_values_supported)
|
||||||
|
|| !metadata.id_token_signing_alg_values_supported.includes(header.alg)
|
||||||
|
|| !text(metadata.jwks_uri, 2048)) throw new OidcProtocolError();
|
||||||
|
const jwksUrl = httpsEndpoint(metadata.jwks_uri);
|
||||||
|
let response: Response;
|
||||||
|
try {
|
||||||
|
response = await (options.fetch ?? globalThis.fetch)(jwksUrl, {
|
||||||
|
headers: { accept: "application/json" }, redirect: "error",
|
||||||
|
});
|
||||||
|
if (!response.ok) throw new OidcProtocolError();
|
||||||
|
const body = await response.text();
|
||||||
|
if (Buffer.byteLength(body, "utf8") > MAX_JWKS_BYTES) throw new OidcProtocolError();
|
||||||
|
const parsed = JSON.parse(body) as { keys?: unknown };
|
||||||
|
if (!Array.isArray(parsed.keys) || parsed.keys.length === 0 || parsed.keys.length > 16) throw new OidcProtocolError();
|
||||||
|
const matching = parsed.keys.filter((key): key is Record<string, unknown> =>
|
||||||
|
Boolean(key) && typeof key === "object" && !Array.isArray(key) && key.kid === header.kid);
|
||||||
|
if (matching.length !== 1) throw new OidcProtocolError();
|
||||||
|
const key = matching[0];
|
||||||
|
if (key.use !== undefined && key.use !== "sig") throw new OidcProtocolError();
|
||||||
|
if (key.alg !== undefined && key.alg !== header.alg) throw new OidcProtocolError();
|
||||||
|
const algorithm = signatureAlgorithm(header.alg);
|
||||||
|
if (!/^[A-Za-z0-9_-]+$/.test(signaturePart)) throw new OidcProtocolError();
|
||||||
|
const signature = Buffer.from(signaturePart, "base64url");
|
||||||
|
if (signature.length === 0) throw new OidcProtocolError();
|
||||||
|
const publicKey = createPublicKey({ key: key as never, format: "jwk" });
|
||||||
|
const normalizedSignature = algorithm.ecdsaPartLength
|
||||||
|
? joseEcdsaSignatureToDer(signature, algorithm.ecdsaPartLength)
|
||||||
|
: signature;
|
||||||
|
const verified = algorithm.pss
|
||||||
|
? verifySignature(algorithm.digest, Buffer.from(`${protectedPart}.${payloadPart}`), {
|
||||||
|
key: publicKey, padding: constants.RSA_PKCS1_PSS_PADDING, saltLength: algorithm.saltLength,
|
||||||
|
}, normalizedSignature)
|
||||||
|
: verifySignature(algorithm.digest, Buffer.from(`${protectedPart}.${payloadPart}`), publicKey, normalizedSignature);
|
||||||
|
if (!verified) throw new OidcProtocolError();
|
||||||
|
} catch (error) {
|
||||||
|
if (error instanceof OidcProtocolError) throw error;
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
export function createOidcProtocol(options: OidcProtocolOptions): OidcProtocol {
|
||||||
|
const issuerUrl = configuredUrl(options.issuer);
|
||||||
|
const callbackUrl = configuredUrl(options.callbackUrl);
|
||||||
|
if (!text(options.clientId, 512) || !text(options.clientSecret, 4096)
|
||||||
|
|| !text(options.groupsClaim, 128) || options.scopes.length === 0 || options.scopes.length > 16
|
||||||
|
|| options.scopes.some((scope) => !text(scope, 128))) throw new OidcProtocolError();
|
||||||
|
|
||||||
|
let discovered: Promise<Configuration> | undefined;
|
||||||
|
const configuration = async (): Promise<Configuration> => {
|
||||||
|
if (!discovered) {
|
||||||
|
discovered = (async () => {
|
||||||
|
try {
|
||||||
|
const config = await discovery(
|
||||||
|
issuerUrl,
|
||||||
|
options.clientId,
|
||||||
|
{ client_secret: options.clientSecret, redirect_uris: [callbackUrl.href], response_types: ["code"] },
|
||||||
|
undefined,
|
||||||
|
options.fetch ? { [customFetch]: options.fetch as never } : undefined,
|
||||||
|
);
|
||||||
|
const metadata = config.serverMetadata();
|
||||||
|
if (metadata.issuer !== options.issuer) throw new OidcProtocolError();
|
||||||
|
httpsEndpoint(metadata.authorization_endpoint);
|
||||||
|
httpsEndpoint(metadata.token_endpoint);
|
||||||
|
httpsEndpoint(metadata.jwks_uri);
|
||||||
|
return config;
|
||||||
|
} catch (error) {
|
||||||
|
if (error instanceof OidcProtocolError) throw error;
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
})();
|
||||||
|
}
|
||||||
|
return await discovered;
|
||||||
|
};
|
||||||
|
|
||||||
|
return {
|
||||||
|
async authorizationUrl(input) {
|
||||||
|
try {
|
||||||
|
const challenge = await calculatePKCECodeChallenge(input.codeVerifier);
|
||||||
|
return buildAuthorizationUrl(await configuration(), {
|
||||||
|
response_type: "code",
|
||||||
|
redirect_uri: callbackUrl.href,
|
||||||
|
scope: options.scopes.join(" "),
|
||||||
|
state: input.state,
|
||||||
|
nonce: input.nonce,
|
||||||
|
code_challenge: challenge,
|
||||||
|
code_challenge_method: "S256",
|
||||||
|
});
|
||||||
|
} catch (error) {
|
||||||
|
if (error instanceof OidcProtocolError) throw error;
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
},
|
||||||
|
async callback(input) {
|
||||||
|
if (input.currentUrl.origin !== callbackUrl.origin || input.currentUrl.pathname !== callbackUrl.pathname) {
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
try {
|
||||||
|
const config = await configuration();
|
||||||
|
const tokens = await authorizationCodeGrant(config, input.currentUrl, {
|
||||||
|
expectedState: input.state,
|
||||||
|
expectedNonce: input.nonce,
|
||||||
|
pkceCodeVerifier: input.codeVerifier,
|
||||||
|
idTokenExpected: true,
|
||||||
|
});
|
||||||
|
await verifyIdTokenSignature(tokens.id_token, config, options);
|
||||||
|
const claims = tokens.claims();
|
||||||
|
if (!claims || Array.isArray(claims)) throw new OidcProtocolError();
|
||||||
|
return identityFromClaims(claims as Record<string, unknown>, options);
|
||||||
|
} catch (error) {
|
||||||
|
if (error instanceof OidcProtocolError) throw error;
|
||||||
|
throw new OidcProtocolError();
|
||||||
|
}
|
||||||
|
},
|
||||||
|
async diagnose(signal) {
|
||||||
|
signal.throwIfAborted();
|
||||||
|
await configuration();
|
||||||
|
signal.throwIfAborted();
|
||||||
|
},
|
||||||
|
};
|
||||||
|
}
|
||||||
+132
-8
@@ -1,6 +1,6 @@
|
|||||||
import type { FastifyInstance, FastifyReply, FastifyRequest } from "fastify";
|
import type { FastifyInstance, FastifyReply, FastifyRequest } from "fastify";
|
||||||
import { randomBytes } from "node:crypto";
|
import { randomBytes } from "node:crypto";
|
||||||
import type { AuthenticationConfigProvider, LoadedAuthConfig } from "./types.js";
|
import type { AuthenticationConfigProvider, LoadedAuthConfig, OidcAuthenticationConfig, Role } from "./types.js";
|
||||||
import type { LocalUserRecord, LocalUserRegistry } from "./local-registry.js";
|
import type { LocalUserRecord, LocalUserRegistry } from "./local-registry.js";
|
||||||
import type { AuthSessionStore } from "./session-store.js";
|
import type { AuthSessionStore } from "./session-store.js";
|
||||||
import { rolesToPermissions } from "./config.js";
|
import { rolesToPermissions } from "./config.js";
|
||||||
@@ -8,12 +8,15 @@ import { captureAuthConfigSnapshot, getPrincipal, requireExactOrigin, sessionCoo
|
|||||||
import { requirePermission, isPrincipalContext } from "./authorization.js";
|
import { requirePermission, isPrincipalContext } from "./authorization.js";
|
||||||
import { deriveCsrfToken } from "./csrf.js";
|
import { deriveCsrfToken } from "./csrf.js";
|
||||||
import { verifyWithDummy } from "./password.js";
|
import { verifyWithDummy } from "./password.js";
|
||||||
|
import type { OidcProtocol } from "./oidc-client.js";
|
||||||
|
|
||||||
const TEN_MINUTES_MS = 10 * 60 * 1000;
|
const TEN_MINUTES_MS = 10 * 60 * 1000;
|
||||||
const REMEMBER_COOKIE_SECONDS = 2_592_000;
|
const REMEMBER_COOKIE_SECONDS = 2_592_000;
|
||||||
const MAX_USERNAME_LENGTH = 64;
|
const MAX_USERNAME_LENGTH = 64;
|
||||||
const MAX_PASSWORD_LENGTH = 1024;
|
const MAX_PASSWORD_LENGTH = 1024;
|
||||||
const MAX_LIMIT_ENTRIES = 10_000;
|
const MAX_LIMIT_ENTRIES = 10_000;
|
||||||
|
const MAX_OIDC_CALLBACK_QUERY_LENGTH = 4096;
|
||||||
|
const OIDC_CALLBACK_PATH = "/api/auth/oidc/callback";
|
||||||
|
|
||||||
export interface AuthRouteDependencies {
|
export interface AuthRouteDependencies {
|
||||||
authMode: "local" | "oidc" | "upstream" | "none" | "mock";
|
authMode: "local" | "oidc" | "upstream" | "none" | "mock";
|
||||||
@@ -22,6 +25,8 @@ export interface AuthRouteDependencies {
|
|||||||
/** Test-only compatibility seam; production resolves from each loaded config snapshot. */
|
/** Test-only compatibility seam; production resolves from each loaded config snapshot. */
|
||||||
localUserRegistry?: LocalUserRegistry;
|
localUserRegistry?: LocalUserRegistry;
|
||||||
resolveLocalUserRegistry?: (loaded: LoadedAuthConfig) => LocalUserRegistry | undefined;
|
resolveLocalUserRegistry?: (loaded: LoadedAuthConfig) => LocalUserRegistry | undefined;
|
||||||
|
oidcProtocol?: OidcProtocol;
|
||||||
|
resolveOidcProtocol?: (loaded: LoadedAuthConfig) => OidcProtocol | undefined;
|
||||||
}
|
}
|
||||||
|
|
||||||
interface LoginPayload {
|
interface LoginPayload {
|
||||||
@@ -118,7 +123,7 @@ export function registerAuthRoutes(app: FastifyInstance, deps: AuthRouteDependen
|
|||||||
const snapshot = captureAuthConfigSnapshot(request, deps.authentication);
|
const snapshot = captureAuthConfigSnapshot(request, deps.authentication);
|
||||||
if (!snapshot) return unavailable(reply);
|
if (!snapshot) return unavailable(reply);
|
||||||
const mode = snapshot.value.mode;
|
const mode = snapshot.value.mode;
|
||||||
return reply.send({ mode, localLogin: mode === "local", oidcLogin: false });
|
return reply.send({ mode, localLogin: mode === "local", oidcLogin: mode === "oidc" });
|
||||||
});
|
});
|
||||||
|
|
||||||
app.post("/auth/local/login", async (request, reply) => {
|
app.post("/auth/local/login", async (request, reply) => {
|
||||||
@@ -180,8 +185,83 @@ export function registerAuthRoutes(app: FastifyInstance, deps: AuthRouteDependen
|
|||||||
}
|
}
|
||||||
});
|
});
|
||||||
|
|
||||||
app.get("/auth/oidc/login", async (_request, reply) => notImplemented(reply));
|
app.get("/auth/oidc/login", async (request, reply) => {
|
||||||
app.get("/auth/oidc/callback", async (_request, reply) => notImplemented(reply));
|
const loaded = captureAuthConfigSnapshot(request, deps.authentication);
|
||||||
|
const configured = currentOidcConfig(loaded, deps);
|
||||||
|
if (!configured || !deps.sessionStore) return unavailable(reply);
|
||||||
|
const nonce = randomOidcValue();
|
||||||
|
const codeVerifier = randomOidcValue();
|
||||||
|
try {
|
||||||
|
const created = await deps.sessionStore.createOidcState({
|
||||||
|
nonce,
|
||||||
|
codeVerifier,
|
||||||
|
returnTo: "/",
|
||||||
|
authConfigRevision: configured.loaded.revision,
|
||||||
|
issuer: configured.config.oidc.issuer,
|
||||||
|
});
|
||||||
|
try {
|
||||||
|
const location = await configured.protocol.authorizationUrl({ state: created.state, nonce, codeVerifier });
|
||||||
|
return reply.redirect(location.href);
|
||||||
|
} catch {
|
||||||
|
await deps.sessionStore.consumeOidcState(created.state).catch(() => undefined);
|
||||||
|
return unavailable(reply);
|
||||||
|
}
|
||||||
|
} catch {
|
||||||
|
return unavailable(reply);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
app.get("/auth/oidc/callback", async (request, reply) => {
|
||||||
|
const loaded = captureAuthConfigSnapshot(request, deps.authentication);
|
||||||
|
if (!loaded || !deps.sessionStore) return oidcCallbackFailed(reply);
|
||||||
|
const callback = oidcCallbackUrl(request, loaded.value.publicUrl);
|
||||||
|
if (!callback) return oidcCallbackFailed(reply);
|
||||||
|
let state;
|
||||||
|
try {
|
||||||
|
state = await deps.sessionStore.consumeOidcState(callback.state);
|
||||||
|
} catch {
|
||||||
|
return oidcCallbackFailed(reply);
|
||||||
|
}
|
||||||
|
const configured = currentOidcConfig(loaded, deps);
|
||||||
|
if (!configured || !state || state.returnTo !== "/" || state.authConfigRevision !== configured.loaded.revision
|
||||||
|
|| state.issuer !== configured.config.oidc.issuer) {
|
||||||
|
return oidcCallbackFailed(reply);
|
||||||
|
}
|
||||||
|
try {
|
||||||
|
const identity = await configured.protocol.callback({
|
||||||
|
currentUrl: callback.currentUrl,
|
||||||
|
state: callback.state,
|
||||||
|
nonce: state.nonce,
|
||||||
|
codeVerifier: state.codeVerifier,
|
||||||
|
});
|
||||||
|
if (identity.issuer !== configured.config.oidc.issuer) return oidcCallbackFailed(reply);
|
||||||
|
const roles = oidcRoles(identity.groups, configured.config);
|
||||||
|
const absoluteTtlMs = Math.min(
|
||||||
|
configured.config.session.oidcTtlSeconds * 1000,
|
||||||
|
identity.tokenExpiresAt.getTime() - Date.now(),
|
||||||
|
);
|
||||||
|
if (!Number.isSafeInteger(absoluteTtlMs) || absoluteTtlMs <= 0) return oidcCallbackFailed(reply);
|
||||||
|
const created = await deps.sessionStore.create({
|
||||||
|
principal: {
|
||||||
|
issuer: identity.issuer,
|
||||||
|
subject: identity.subject,
|
||||||
|
...(identity.displayName === undefined ? {} : { displayName: identity.displayName }),
|
||||||
|
roles,
|
||||||
|
permissions: rolesToPermissions(roles),
|
||||||
|
isAdmin: roles.includes("admin"),
|
||||||
|
},
|
||||||
|
method: "oidc",
|
||||||
|
remembered: false,
|
||||||
|
authConfigRevision: configured.loaded.revision,
|
||||||
|
idleTtlMs: configured.config.session.regularIdleSeconds * 1000,
|
||||||
|
absoluteTtlMs,
|
||||||
|
});
|
||||||
|
reply.setCookie(sessionCookieName(), created.token, cookieOptions(configured.loaded, false));
|
||||||
|
return reply.redirect(state.returnTo);
|
||||||
|
} catch {
|
||||||
|
return oidcCallbackFailed(reply);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
app.post("/auth/logout", async (request, reply) => {
|
app.post("/auth/logout", async (request, reply) => {
|
||||||
const token = request.authSessionToken;
|
const token = request.authSessionToken;
|
||||||
@@ -278,6 +358,54 @@ function cookieOptions(snapshot: LoadedAuthConfig | undefined, remembered: boole
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function currentOidcConfig(loaded: LoadedAuthConfig | undefined, deps: AuthRouteDependencies):
|
||||||
|
| { loaded: LoadedAuthConfig; config: OidcAuthenticationConfig; protocol: OidcProtocol }
|
||||||
|
| undefined {
|
||||||
|
if (!loaded || loaded.value.mode !== "oidc") return undefined;
|
||||||
|
const protocol = deps.resolveOidcProtocol?.(loaded) ?? deps.oidcProtocol;
|
||||||
|
return protocol ? { loaded, config: loaded.value, protocol } : undefined;
|
||||||
|
}
|
||||||
|
|
||||||
|
function randomOidcValue(): string {
|
||||||
|
return randomBytes(32).toString("base64url");
|
||||||
|
}
|
||||||
|
|
||||||
|
function oidcCallbackUrl(request: FastifyRequest, publicUrl: string): { currentUrl: URL; state: string } | undefined {
|
||||||
|
if (request.url.length > MAX_OIDC_CALLBACK_QUERY_LENGTH) return undefined;
|
||||||
|
let supplied: URL;
|
||||||
|
let target: URL;
|
||||||
|
try {
|
||||||
|
supplied = new URL(request.url, "http://callback.invalid");
|
||||||
|
target = new URL(OIDC_CALLBACK_PATH, publicUrl);
|
||||||
|
} catch {
|
||||||
|
return undefined;
|
||||||
|
}
|
||||||
|
if (supplied.pathname !== "/auth/oidc/callback") return undefined;
|
||||||
|
const allowed = new Set(["code", "state", "error", "error_description", "error_uri", "iss"]);
|
||||||
|
const copied = new URLSearchParams();
|
||||||
|
let state: string | undefined;
|
||||||
|
for (const [key, value] of supplied.searchParams) {
|
||||||
|
if (!allowed.has(key) || value.length > 2048 || copied.has(key)) return undefined;
|
||||||
|
copied.set(key, value);
|
||||||
|
if (key === "state") state = value;
|
||||||
|
}
|
||||||
|
if (!state || !/^[A-Za-z0-9_-]{43}$/.test(state)) return undefined;
|
||||||
|
target.search = copied.toString();
|
||||||
|
return { currentUrl: target, state };
|
||||||
|
}
|
||||||
|
|
||||||
|
function oidcCallbackFailed(reply: FastifyReply) {
|
||||||
|
return reply.code(401).send({ code: "oidc_callback_failed", error: "OIDC sign-in could not be completed" });
|
||||||
|
}
|
||||||
|
|
||||||
|
function oidcRoles(groups: readonly string[], config: OidcAuthenticationConfig): Role[] {
|
||||||
|
const roles = new Set<Role>();
|
||||||
|
for (const group of groups) {
|
||||||
|
for (const role of config.authorization.groupRoles[group] ?? []) roles.add(role);
|
||||||
|
}
|
||||||
|
return [...roles];
|
||||||
|
}
|
||||||
|
|
||||||
function loginPayload(request: FastifyRequest): LoginPayload {
|
function loginPayload(request: FastifyRequest): LoginPayload {
|
||||||
const body = request.body;
|
const body = request.body;
|
||||||
if (!body || typeof body !== "object" || Array.isArray(body)) return { username: "", password: "", remember: false };
|
if (!body || typeof body !== "object" || Array.isArray(body)) return { username: "", password: "", remember: false };
|
||||||
@@ -316,7 +444,3 @@ function loginLimited(reply: FastifyReply): FastifyReply {
|
|||||||
function unavailable(reply: FastifyReply): FastifyReply {
|
function unavailable(reply: FastifyReply): FastifyReply {
|
||||||
return reply.code(503).send({ code: "auth_unavailable", error: "Authentication is unavailable" });
|
return reply.code(503).send({ code: "auth_unavailable", error: "Authentication is unavailable" });
|
||||||
}
|
}
|
||||||
|
|
||||||
function notImplemented(reply: FastifyReply): FastifyReply {
|
|
||||||
return reply.code(501).send({ code: "auth_not_implemented", error: "OIDC login is not implemented" });
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -68,6 +68,8 @@ export interface OidcStateCreateInput {
|
|||||||
nonce: string;
|
nonce: string;
|
||||||
codeVerifier: string;
|
codeVerifier: string;
|
||||||
returnTo: "/";
|
returnTo: "/";
|
||||||
|
authConfigRevision: string;
|
||||||
|
issuer: string;
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface CreatedOidcState {
|
export interface CreatedOidcState {
|
||||||
@@ -188,6 +190,8 @@ const oidcStateRecordSchema = z.strictObject({
|
|||||||
nonce: z.string().min(16).max(512).regex(/^[A-Za-z0-9_-]+$/),
|
nonce: z.string().min(16).max(512).regex(/^[A-Za-z0-9_-]+$/),
|
||||||
codeVerifier: z.string().min(43).max(128).regex(/^[A-Za-z0-9._~-]+$/),
|
codeVerifier: z.string().min(43).max(128).regex(/^[A-Za-z0-9._~-]+$/),
|
||||||
returnTo: z.literal("/"),
|
returnTo: z.literal("/"),
|
||||||
|
authConfigRevision: z.string().regex(/^[a-f0-9]{64}$/),
|
||||||
|
issuer: z.string().min(1).max(2048).refine((value) => !/\p{Cc}/u.test(value)),
|
||||||
createdAt: timestamp,
|
createdAt: timestamp,
|
||||||
expiresAt: timestamp,
|
expiresAt: timestamp,
|
||||||
}).superRefine((record, context) => {
|
}).superRefine((record, context) => {
|
||||||
@@ -226,6 +230,8 @@ const oidcStateInputSchema = z.strictObject({
|
|||||||
nonce: z.string().min(16).max(512).regex(/^[A-Za-z0-9_-]+$/),
|
nonce: z.string().min(16).max(512).regex(/^[A-Za-z0-9_-]+$/),
|
||||||
codeVerifier: z.string().min(43).max(128).regex(/^[A-Za-z0-9._~-]+$/),
|
codeVerifier: z.string().min(43).max(128).regex(/^[A-Za-z0-9._~-]+$/),
|
||||||
returnTo: z.literal("/"),
|
returnTo: z.literal("/"),
|
||||||
|
authConfigRevision: z.string().regex(/^[a-f0-9]{64}$/),
|
||||||
|
issuer: z.string().min(1).max(2048).refine((value) => !/\p{Cc}/u.test(value)),
|
||||||
});
|
});
|
||||||
|
|
||||||
function sameFileIdentity(left: FileIdentity, right: FileIdentity): boolean {
|
function sameFileIdentity(left: FileIdentity, right: FileIdentity): boolean {
|
||||||
@@ -927,6 +933,8 @@ export function createFileAuthSessionStore(
|
|||||||
nonce: validated.nonce,
|
nonce: validated.nonce,
|
||||||
codeVerifier: validated.codeVerifier,
|
codeVerifier: validated.codeVerifier,
|
||||||
returnTo: validated.returnTo,
|
returnTo: validated.returnTo,
|
||||||
|
authConfigRevision: validated.authConfigRevision,
|
||||||
|
issuer: validated.issuer,
|
||||||
createdAt: isoAt(nowMs),
|
createdAt: isoAt(nowMs),
|
||||||
expiresAt: isoAt(expiresMs),
|
expiresAt: isoAt(expiresMs),
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -79,6 +79,8 @@ export interface OidcStateRecord {
|
|||||||
nonce: string;
|
nonce: string;
|
||||||
codeVerifier: string;
|
codeVerifier: string;
|
||||||
returnTo: "/";
|
returnTo: "/";
|
||||||
|
authConfigRevision: string;
|
||||||
|
issuer: string;
|
||||||
createdAt: string;
|
createdAt: string;
|
||||||
expiresAt: string;
|
expiresAt: string;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ import { stringify } from "yaml";
|
|||||||
import { buildApp } from "../src/app.js";
|
import { buildApp } from "../src/app.js";
|
||||||
import { loadConfig } from "../src/config.js";
|
import { loadConfig } from "../src/config.js";
|
||||||
|
|
||||||
test("configured OIDC starts with provider-neutral protocol placeholders that fail closed", async () => {
|
test("configured OIDC advertises login but fails closed without its runtime client secret", async () => {
|
||||||
const directory = mkdtempSync(join(tmpdir(), "thothii-app-oidc-mode-"));
|
const directory = mkdtempSync(join(tmpdir(), "thothii-app-oidc-mode-"));
|
||||||
const file = join(directory, "auth.yaml");
|
const file = join(directory, "auth.yaml");
|
||||||
writeFileSync(file, stringify({
|
writeFileSync(file, stringify({
|
||||||
@@ -22,10 +22,10 @@ test("configured OIDC starts with provider-neutral protocol placeholders that fa
|
|||||||
const app = buildApp(loadConfig({ THT_AUTH_CONFIG_FILE: file, THT_AUTH_STATE_ROOT: join(directory, "auth-state") }));
|
const app = buildApp(loadConfig({ THT_AUTH_CONFIG_FILE: file, THT_AUTH_STATE_ROOT: join(directory, "auth-state") }));
|
||||||
try {
|
try {
|
||||||
expect((await app.inject({ method: "GET", url: "/auth/config" })).json())
|
expect((await app.inject({ method: "GET", url: "/auth/config" })).json())
|
||||||
.toEqual({ mode: "oidc", localLogin: false, oidcLogin: false });
|
.toEqual({ mode: "oidc", localLogin: false, oidcLogin: true });
|
||||||
const placeholder = await app.inject({ method: "GET", url: "/auth/oidc/login" });
|
const placeholder = await app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
expect(placeholder.statusCode).toBe(501);
|
expect(placeholder.statusCode).toBe(503);
|
||||||
expect(placeholder.json()).toEqual({ code: "auth_not_implemented", error: "OIDC login is not implemented" });
|
expect(placeholder.json()).toEqual({ code: "auth_unavailable", error: "Authentication is unavailable" });
|
||||||
} finally {
|
} finally {
|
||||||
await app.close();
|
await app.close();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -500,7 +500,7 @@ test("operational config failures return 503 and never consume login-failure cap
|
|||||||
expect((await login(app, { password: `${password}!` })).statusCode).toBe(429);
|
expect((await login(app, { password: `${password}!` })).statusCode).toBe(429);
|
||||||
});
|
});
|
||||||
|
|
||||||
test("public auth configuration is safe and OIDC protocol placeholders fail closed", async () => {
|
test("public auth configuration is safe and unavailable OIDC login fails closed", async () => {
|
||||||
const { app } = await createLocalApp();
|
const { app } = await createLocalApp();
|
||||||
const configuration = await app.inject({ method: "GET", url: "/auth/config" });
|
const configuration = await app.inject({ method: "GET", url: "/auth/config" });
|
||||||
expect(configuration.statusCode).toBe(200);
|
expect(configuration.statusCode).toBe(200);
|
||||||
@@ -508,6 +508,6 @@ test("public auth configuration is safe and OIDC protocol placeholders fail clos
|
|||||||
expect(JSON.stringify(configuration.json())).not.toContain("users.yaml");
|
expect(JSON.stringify(configuration.json())).not.toContain("users.yaml");
|
||||||
|
|
||||||
const placeholder = await app.inject({ method: "GET", url: "/auth/oidc/login" });
|
const placeholder = await app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
expect(placeholder.statusCode).toBe(501);
|
expect(placeholder.statusCode).toBe(503);
|
||||||
expect(placeholder.json()).toEqual({ code: "auth_not_implemented", error: "OIDC login is not implemented" });
|
expect(placeholder.json()).toEqual({ code: "auth_unavailable", error: "Authentication is unavailable" });
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -0,0 +1,215 @@
|
|||||||
|
import Fastify from "fastify";
|
||||||
|
import cookie from "@fastify/cookie";
|
||||||
|
import { afterEach, expect, test } from "vitest";
|
||||||
|
import { registerAuthRoutes } from "../src/auth/routes.js";
|
||||||
|
import type { LoadedAuthConfig, OidcStateRecord } from "../src/auth/types.js";
|
||||||
|
import type { AuthSessionStore } from "../src/auth/session-store.js";
|
||||||
|
import type { OidcProtocol } from "../src/auth/oidc-client.js";
|
||||||
|
|
||||||
|
const revision = "a".repeat(64);
|
||||||
|
const issuer = "https://issuer.example.test";
|
||||||
|
const state = "s".repeat(43);
|
||||||
|
const nonce = "n".repeat(43);
|
||||||
|
const verifier = "v".repeat(43);
|
||||||
|
const createdApps: Array<ReturnType<typeof Fastify>> = [];
|
||||||
|
|
||||||
|
function config(overrides: Partial<LoadedAuthConfig["value"]> = {}): LoadedAuthConfig {
|
||||||
|
return {
|
||||||
|
revision,
|
||||||
|
sourcePath: "/private/auth.yaml",
|
||||||
|
value: {
|
||||||
|
version: 1,
|
||||||
|
mode: "oidc",
|
||||||
|
publicUrl: "https://thothii.example.test",
|
||||||
|
session: {
|
||||||
|
regularTtlSeconds: 3600, regularIdleSeconds: 300,
|
||||||
|
rememberTtlSeconds: 3600, rememberIdleSeconds: 300, oidcTtlSeconds: 3600,
|
||||||
|
},
|
||||||
|
oidc: {
|
||||||
|
issuer, clientId: "thothii", clientSecretRef: "THT_OIDC_CLIENT_SECRET",
|
||||||
|
scopes: ["openid", "profile"], groupsClaim: "groups",
|
||||||
|
},
|
||||||
|
groupCatalog: { driver: "authentik", baseUrl: issuer, apiTokenRef: "THT_AUTHENTIK_API_TOKEN" },
|
||||||
|
authorization: { groupRoles: { Users: ["user"], Admins: ["admin"] } },
|
||||||
|
...overrides,
|
||||||
|
},
|
||||||
|
} as LoadedAuthConfig;
|
||||||
|
}
|
||||||
|
|
||||||
|
function stateRecord(extra: Partial<OidcStateRecord> = {}): OidcStateRecord {
|
||||||
|
return {
|
||||||
|
version: 1, nonce, codeVerifier: verifier, returnTo: "/",
|
||||||
|
authConfigRevision: revision, issuer,
|
||||||
|
createdAt: "2030-01-01T00:00:00.000Z", expiresAt: "2030-01-01T00:10:00.000Z",
|
||||||
|
...extra,
|
||||||
|
} as OidcStateRecord;
|
||||||
|
}
|
||||||
|
|
||||||
|
function fixture(options: {
|
||||||
|
loaded?: LoadedAuthConfig;
|
||||||
|
identity?: Awaited<ReturnType<OidcProtocol["callback"]>>;
|
||||||
|
callbackFailure?: boolean;
|
||||||
|
stateReturnTo?: string;
|
||||||
|
} = {}) {
|
||||||
|
let loaded = options.loaded ?? config();
|
||||||
|
let protocolAvailable = true;
|
||||||
|
let storedState: OidcStateRecord | undefined;
|
||||||
|
const stateInputs: Array<Record<string, unknown>> = [];
|
||||||
|
const creates: Array<Record<string, unknown>> = [];
|
||||||
|
const callbacks: URL[] = [];
|
||||||
|
const protocol: OidcProtocol = {
|
||||||
|
authorizationUrl: async ({ state: received, nonce: receivedNonce, codeVerifier }) => {
|
||||||
|
expect(received).toBe(state);
|
||||||
|
expect(receivedNonce).toHaveLength(43);
|
||||||
|
expect(codeVerifier).toHaveLength(43);
|
||||||
|
return new URL(`https://issuer.example.test/authorize?state=${received}`);
|
||||||
|
},
|
||||||
|
callback: async ({ currentUrl }) => {
|
||||||
|
callbacks.push(currentUrl);
|
||||||
|
if (options.callbackFailure) throw new Error("provider failure with access-token-must-not-leak");
|
||||||
|
return options.identity ?? {
|
||||||
|
issuer, subject: "user-123", displayName: "Ada", groups: ["Users", "Admins", "Unmapped"],
|
||||||
|
tokenExpiresAt: new Date(Date.now() + 120_000),
|
||||||
|
};
|
||||||
|
},
|
||||||
|
diagnose: async () => undefined,
|
||||||
|
};
|
||||||
|
const store = {
|
||||||
|
createOidcState: async (input: Record<string, unknown>) => {
|
||||||
|
stateInputs.push(input);
|
||||||
|
const record = stateRecord({
|
||||||
|
nonce: input.nonce as string,
|
||||||
|
codeVerifier: input.codeVerifier as string,
|
||||||
|
authConfigRevision: input.authConfigRevision as string,
|
||||||
|
issuer: input.issuer as string,
|
||||||
|
});
|
||||||
|
if (options.stateReturnTo) (record as { returnTo: string }).returnTo = options.stateReturnTo;
|
||||||
|
storedState = record;
|
||||||
|
return { state, record: storedState };
|
||||||
|
},
|
||||||
|
consumeOidcState: async (received: string) => {
|
||||||
|
if (received !== state) return undefined;
|
||||||
|
const consumed = storedState;
|
||||||
|
storedState = undefined;
|
||||||
|
return consumed;
|
||||||
|
},
|
||||||
|
create: async (input: Record<string, unknown>) => {
|
||||||
|
creates.push(input);
|
||||||
|
return { token: "opaque-session-token", csrfToken: "c".repeat(43), record: {} };
|
||||||
|
},
|
||||||
|
} as unknown as AuthSessionStore;
|
||||||
|
const app = Fastify();
|
||||||
|
app.decorateRequest("authConfigSnapshot", undefined);
|
||||||
|
app.decorateRequest("authConfigSnapshotCaptured", false);
|
||||||
|
app.decorateRequest("authConfigSnapshotUnavailable", false);
|
||||||
|
app.register(cookie);
|
||||||
|
registerAuthRoutes(app, {
|
||||||
|
authMode: "oidc",
|
||||||
|
authentication: { current: () => loaded },
|
||||||
|
sessionStore: store,
|
||||||
|
resolveOidcProtocol: () => protocolAvailable ? protocol : undefined,
|
||||||
|
});
|
||||||
|
createdApps.push(app);
|
||||||
|
return {
|
||||||
|
app, creates, callbacks, stateInputs,
|
||||||
|
setConfig(next: LoadedAuthConfig) { loaded = next; },
|
||||||
|
setProtocolAvailable(available: boolean) { protocolAvailable = available; },
|
||||||
|
stateWasConsumed: () => storedState === undefined,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
afterEach(async () => {
|
||||||
|
await Promise.all(createdApps.splice(0).map((app) => app.close()));
|
||||||
|
});
|
||||||
|
|
||||||
|
test("creates digest-only bound state, maps exact groups, creates a cookie session, and redirects safely", async () => {
|
||||||
|
const subject = fixture();
|
||||||
|
const start = await subject.app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
|
expect(start.statusCode).toBe(302);
|
||||||
|
expect(new URL(start.headers.location ?? "").searchParams.get("state")).toBe(state);
|
||||||
|
expect(subject.stateInputs[0]).toMatchObject({ returnTo: "/", authConfigRevision: revision, issuer });
|
||||||
|
|
||||||
|
const callback = await subject.app.inject({
|
||||||
|
method: "GET",
|
||||||
|
url: `/auth/oidc/callback?code=good&state=${state}`,
|
||||||
|
headers: { host: "attacker.example.test" },
|
||||||
|
});
|
||||||
|
expect(callback.statusCode).toBe(302);
|
||||||
|
expect(callback.headers.location).toBe("/");
|
||||||
|
expect(callback.headers["set-cookie"]).toContain("HttpOnly");
|
||||||
|
expect(callback.headers["set-cookie"]).toContain("SameSite=Lax");
|
||||||
|
expect(callback.headers["set-cookie"]).toContain("Secure");
|
||||||
|
expect(subject.callbacks[0]?.href).toBe(`https://thothii.example.test/api/auth/oidc/callback?code=good&state=${state}`);
|
||||||
|
expect(subject.creates).toHaveLength(1);
|
||||||
|
expect(subject.creates[0]).toMatchObject({
|
||||||
|
method: "oidc", remembered: false, authConfigRevision: revision,
|
||||||
|
principal: { issuer, subject: "user-123", roles: ["user", "admin"] },
|
||||||
|
idleTtlMs: 300_000,
|
||||||
|
});
|
||||||
|
const absoluteTtlMs = subject.creates[0]?.absoluteTtlMs;
|
||||||
|
expect(typeof absoluteTtlMs).toBe("number");
|
||||||
|
expect(absoluteTtlMs as number).toBeGreaterThan(0);
|
||||||
|
expect(absoluteTtlMs as number).toBeLessThanOrEqual(120_000);
|
||||||
|
expect(JSON.stringify(subject.creates)).not.toContain("access-token-must-not-leak");
|
||||||
|
expect(JSON.stringify(subject.creates)).not.toContain("refresh-token-must-not-leak");
|
||||||
|
});
|
||||||
|
|
||||||
|
test("consumes state on callback failure and refuses replay", async () => {
|
||||||
|
const subject = fixture({ callbackFailure: true });
|
||||||
|
await subject.app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
|
const failed = await subject.app.inject({ method: "GET", url: `/auth/oidc/callback?code=good&state=${state}` });
|
||||||
|
expect(failed.statusCode).toBe(401);
|
||||||
|
const replay = await subject.app.inject({ method: "GET", url: `/auth/oidc/callback?code=good&state=${state}` });
|
||||||
|
expect(replay.statusCode).toBe(401);
|
||||||
|
expect(subject.callbacks).toHaveLength(1);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("consumes state when the protocol becomes unavailable before callback", async () => {
|
||||||
|
const subject = fixture();
|
||||||
|
await subject.app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
|
subject.setProtocolAvailable(false);
|
||||||
|
const failed = await subject.app.inject({ method: "GET", url: `/auth/oidc/callback?code=good&state=${state}` });
|
||||||
|
expect(failed.statusCode).toBe(401);
|
||||||
|
expect(subject.stateWasConsumed()).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("rejects a consumed state with a non-root return target", async () => {
|
||||||
|
const subject = fixture({ stateReturnTo: "https://attacker.example.test" });
|
||||||
|
await subject.app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
|
const callback = await subject.app.inject({ method: "GET", url: `/auth/oidc/callback?code=good&state=${state}` });
|
||||||
|
expect(callback.statusCode).toBe(401);
|
||||||
|
expect(subject.callbacks).toEqual([]);
|
||||||
|
expect(subject.creates).toEqual([]);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("rejects an OIDC state when its configuration revision changes before callback", async () => {
|
||||||
|
const subject = fixture();
|
||||||
|
await subject.app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
|
subject.setConfig({ ...config(), revision: "b".repeat(64) });
|
||||||
|
const callback = await subject.app.inject({ method: "GET", url: `/auth/oidc/callback?code=good&state=${state}` });
|
||||||
|
expect(callback.statusCode).toBe(401);
|
||||||
|
expect(subject.callbacks).toEqual([]);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("rejects an OIDC state when its issuer changes before callback", async () => {
|
||||||
|
const subject = fixture();
|
||||||
|
await subject.app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
|
const previous = config();
|
||||||
|
subject.setConfig({
|
||||||
|
...previous,
|
||||||
|
value: { ...previous.value, oidc: { ...previous.value.oidc, issuer: "https://other.example.test" } },
|
||||||
|
} as LoadedAuthConfig);
|
||||||
|
const callback = await subject.app.inject({ method: "GET", url: `/auth/oidc/callback?code=good&state=${state}` });
|
||||||
|
expect(callback.statusCode).toBe(401);
|
||||||
|
expect(subject.callbacks).toEqual([]);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("creates an authenticated but forbidden principal for extra unmapped groups", async () => {
|
||||||
|
const subject = fixture({ identity: {
|
||||||
|
issuer, subject: "user-123", groups: ["Unmapped"], tokenExpiresAt: new Date(Date.now() + 60_000),
|
||||||
|
} });
|
||||||
|
await subject.app.inject({ method: "GET", url: "/auth/oidc/login" });
|
||||||
|
const callback = await subject.app.inject({ method: "GET", url: `/auth/oidc/callback?code=good&state=${state}` });
|
||||||
|
expect(callback.statusCode).toBe(302);
|
||||||
|
expect(subject.creates[0]).toMatchObject({ principal: { roles: [], permissions: [], isAdmin: false } });
|
||||||
|
});
|
||||||
@@ -69,6 +69,10 @@ const base = new Date("2030-01-02T03:04:05.000Z");
|
|||||||
const revision = "a".repeat(64);
|
const revision = "a".repeat(64);
|
||||||
const validLocalUser = { enabled: true, authRevision: 7, roles: ["admin"] as const };
|
const validLocalUser = { enabled: true, authRevision: 7, roles: ["admin"] as const };
|
||||||
|
|
||||||
|
function oidcInput(nonce: string, codeVerifier: string) {
|
||||||
|
return { nonce, codeVerifier, returnTo: "/" as const, authConfigRevision: revision, issuer: "https://issuer.example.test" };
|
||||||
|
}
|
||||||
|
|
||||||
afterEach(() => {
|
afterEach(() => {
|
||||||
fsHooks.afterRead = undefined;
|
fsHooks.afterRead = undefined;
|
||||||
fsHooks.afterWrite = undefined;
|
fsHooks.afterWrite = undefined;
|
||||||
@@ -258,7 +262,7 @@ describe("file-backed auth session store", () => {
|
|||||||
const store = validStore(storageRoot);
|
const store = validStore(storageRoot);
|
||||||
const revoked = await create(store);
|
const revoked = await create(store);
|
||||||
const expired = await create(store, { idleTtlMs: 60_000, absoluteTtlMs: 60_000 });
|
const expired = await create(store, { idleTtlMs: 60_000, absoluteTtlMs: 60_000 });
|
||||||
const oidc = await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base);
|
const oidc = await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
|
|
||||||
await store.revoke(revoked.token);
|
await store.revoke(revoked.token);
|
||||||
await expect(store.resolve(revoked.token)).resolves.toBeUndefined();
|
await expect(store.resolve(revoked.token)).resolves.toBeUndefined();
|
||||||
@@ -272,20 +276,19 @@ describe("file-backed auth session store", () => {
|
|||||||
test("creates bounded OIDC state records that expire and are single-use", async () => {
|
test("creates bounded OIDC state records that expire and are single-use", async () => {
|
||||||
const storageRoot = root();
|
const storageRoot = root();
|
||||||
const store = validStore(storageRoot);
|
const store = validStore(storageRoot);
|
||||||
const created = await store.createOidcState({
|
const created = await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
nonce: "n".repeat(43),
|
|
||||||
codeVerifier: "v".repeat(43),
|
|
||||||
returnTo: "/",
|
|
||||||
}, base);
|
|
||||||
const path = digestPath(storageRoot, "oidc", created.state);
|
const path = digestPath(storageRoot, "oidc", created.state);
|
||||||
|
|
||||||
expect(created.state).toMatch(/^[A-Za-z0-9_-]{43}$/);
|
expect(created.state).toMatch(/^[A-Za-z0-9_-]{43}$/);
|
||||||
expect(readFileSync(path, "utf8")).not.toContain(created.state);
|
expect(readFileSync(path, "utf8")).not.toContain(created.state);
|
||||||
await expect(store.consumeOidcState(created.state, new Date(base.getTime() + 9 * 60_000)))
|
await expect(store.consumeOidcState(created.state, new Date(base.getTime() + 9 * 60_000)))
|
||||||
.resolves.toMatchObject({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" });
|
.resolves.toMatchObject({
|
||||||
|
nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/",
|
||||||
|
authConfigRevision: revision, issuer: "https://issuer.example.test",
|
||||||
|
});
|
||||||
await expect(store.consumeOidcState(created.state)).resolves.toBeUndefined();
|
await expect(store.consumeOidcState(created.state)).resolves.toBeUndefined();
|
||||||
|
|
||||||
const expired = await store.createOidcState({ nonce: "x".repeat(43), codeVerifier: "y".repeat(43), returnTo: "/" }, base);
|
const expired = await store.createOidcState(oidcInput("x".repeat(43), "y".repeat(43)), base);
|
||||||
await expect(store.consumeOidcState(expired.state, new Date(base.getTime() + 10 * 60_000)))
|
await expect(store.consumeOidcState(expired.state, new Date(base.getTime() + 10 * 60_000)))
|
||||||
.resolves.toBeUndefined();
|
.resolves.toBeUndefined();
|
||||||
});
|
});
|
||||||
@@ -293,7 +296,7 @@ describe("file-backed auth session store", () => {
|
|||||||
test("fails closed when an OIDC state already has an atomic filesystem claim", async () => {
|
test("fails closed when an OIDC state already has an atomic filesystem claim", async () => {
|
||||||
const storageRoot = root();
|
const storageRoot = root();
|
||||||
const store = validStore(storageRoot);
|
const store = validStore(storageRoot);
|
||||||
const created = await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base);
|
const created = await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
const statePath = digestPath(storageRoot, "oidc", created.state);
|
const statePath = digestPath(storageRoot, "oidc", created.state);
|
||||||
linkSync(statePath, claimPath(storageRoot, created.state));
|
linkSync(statePath, claimPath(storageRoot, created.state));
|
||||||
|
|
||||||
@@ -304,7 +307,7 @@ describe("file-backed auth session store", () => {
|
|||||||
test("treats a competing OIDC claim installed between availability and state checks as unavailable", async () => {
|
test("treats a competing OIDC claim installed between availability and state checks as unavailable", async () => {
|
||||||
const storageRoot = root();
|
const storageRoot = root();
|
||||||
const store = validStore(storageRoot);
|
const store = validStore(storageRoot);
|
||||||
const created = await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base);
|
const created = await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
const statePath = digestPath(storageRoot, "oidc", created.state);
|
const statePath = digestPath(storageRoot, "oidc", created.state);
|
||||||
const stateClaimPath = claimPath(storageRoot, created.state);
|
const stateClaimPath = claimPath(storageRoot, created.state);
|
||||||
fsHooks.beforeLstat = (observed) => {
|
fsHooks.beforeLstat = (observed) => {
|
||||||
@@ -321,7 +324,7 @@ describe("file-backed auth session store", () => {
|
|||||||
test("prunes an expired OIDC state abandoned after an atomic claim", async () => {
|
test("prunes an expired OIDC state abandoned after an atomic claim", async () => {
|
||||||
const storageRoot = root();
|
const storageRoot = root();
|
||||||
const store = validStore(storageRoot);
|
const store = validStore(storageRoot);
|
||||||
const created = await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base);
|
const created = await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
const statePath = digestPath(storageRoot, "oidc", created.state);
|
const statePath = digestPath(storageRoot, "oidc", created.state);
|
||||||
const stateClaimPath = claimPath(storageRoot, created.state);
|
const stateClaimPath = claimPath(storageRoot, created.state);
|
||||||
linkSync(statePath, stateClaimPath);
|
linkSync(statePath, stateClaimPath);
|
||||||
@@ -334,7 +337,7 @@ describe("file-backed auth session store", () => {
|
|||||||
test("retains an in-flight orphan claim but removes it after the bounded recovery window", async () => {
|
test("retains an in-flight orphan claim but removes it after the bounded recovery window", async () => {
|
||||||
const storageRoot = root();
|
const storageRoot = root();
|
||||||
const store = validStore(storageRoot);
|
const store = validStore(storageRoot);
|
||||||
const created = await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base);
|
const created = await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
const statePath = digestPath(storageRoot, "oidc", created.state);
|
const statePath = digestPath(storageRoot, "oidc", created.state);
|
||||||
const stateClaimPath = claimPath(storageRoot, created.state);
|
const stateClaimPath = claimPath(storageRoot, created.state);
|
||||||
linkSync(statePath, stateClaimPath);
|
linkSync(statePath, stateClaimPath);
|
||||||
@@ -349,7 +352,7 @@ describe("file-backed auth session store", () => {
|
|||||||
test("allows exactly one separate Node isolate to consume an OIDC state", async () => {
|
test("allows exactly one separate Node isolate to consume an OIDC state", async () => {
|
||||||
const storageRoot = root();
|
const storageRoot = root();
|
||||||
const store = validStore(storageRoot);
|
const store = validStore(storageRoot);
|
||||||
const created = await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base);
|
const created = await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
const [first, second] = await Promise.all([
|
const [first, second] = await Promise.all([
|
||||||
isolatedOidcConsumer(storageRoot, created.state),
|
isolatedOidcConsumer(storageRoot, created.state),
|
||||||
isolatedOidcConsumer(storageRoot, created.state),
|
isolatedOidcConsumer(storageRoot, created.state),
|
||||||
@@ -587,12 +590,12 @@ describe("file-backed auth session store", () => {
|
|||||||
await store.revoke(session.token);
|
await store.revoke(session.token);
|
||||||
await expect(store.resolve(session.token)).resolves.toBeUndefined();
|
await expect(store.resolve(session.token)).resolves.toBeUndefined();
|
||||||
|
|
||||||
const oidc = await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base);
|
const oidc = await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
await expect(store.consumeOidcState(oidc.state, new Date(base.getTime() + 9 * 60_000)))
|
await expect(store.consumeOidcState(oidc.state, new Date(base.getTime() + 9 * 60_000)))
|
||||||
.resolves.toMatchObject({ nonce: "n".repeat(43) });
|
.resolves.toMatchObject({ nonce: "n".repeat(43) });
|
||||||
await expect(store.consumeOidcState(oidc.state)).resolves.toBeUndefined();
|
await expect(store.consumeOidcState(oidc.state)).resolves.toBeUndefined();
|
||||||
await create(store, { idleTtlMs: 60_000, absoluteTtlMs: 60_000 });
|
await create(store, { idleTtlMs: 60_000, absoluteTtlMs: 60_000 });
|
||||||
await store.createOidcState({ nonce: "x".repeat(43), codeVerifier: "y".repeat(43), returnTo: "/" }, base);
|
await store.createOidcState(oidcInput("x".repeat(43), "y".repeat(43)), base);
|
||||||
await expect(store.prune(new Date(base.getTime() + 11 * 60_000))).resolves.toBe(2);
|
await expect(store.prune(new Date(base.getTime() + 11 * 60_000))).resolves.toBe(2);
|
||||||
expect(calls).toEqual(expect.arrayContaining(["create", "read", "replace", "remove", "claim-consume", "list"]));
|
expect(calls).toEqual(expect.arrayContaining(["create", "read", "replace", "remove", "claim-consume", "list"]));
|
||||||
expect(records).toHaveLength(0);
|
expect(records).toHaveLength(0);
|
||||||
@@ -651,7 +654,7 @@ describe("file-backed auth session store", () => {
|
|||||||
findLocalUser: async () => validLocalUser,
|
findLocalUser: async () => validLocalUser,
|
||||||
}, { windowsStorageBridge: bridge });
|
}, { windowsStorageBridge: bridge });
|
||||||
await create(store, { idleTtlMs: 60_000, absoluteTtlMs: 60_000 });
|
await create(store, { idleTtlMs: 60_000, absoluteTtlMs: 60_000 });
|
||||||
await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base);
|
await store.createOidcState(oidcInput("n".repeat(43), "v".repeat(43)), base);
|
||||||
|
|
||||||
await expect(store.prune(new Date(base.getTime() + 11 * 60_000))).resolves.toBe(2);
|
await expect(store.prune(new Date(base.getTime() + 11 * 60_000))).resolves.toBe(2);
|
||||||
expect(records).toHaveLength(0);
|
expect(records).toHaveLength(0);
|
||||||
|
|||||||
@@ -0,0 +1,142 @@
|
|||||||
|
import { createSign, generateKeyPairSync } from "node:crypto";
|
||||||
|
import { expect, test } from "vitest";
|
||||||
|
import { createOidcProtocol, OidcProtocolError } from "../src/auth/oidc-client.js";
|
||||||
|
|
||||||
|
const issuer = "https://issuer.example.test";
|
||||||
|
const clientId = "thothii";
|
||||||
|
const callbackUrl = "https://thothii.example.test/api/auth/oidc/callback";
|
||||||
|
const verifier = "v".repeat(43);
|
||||||
|
const nonce = "n".repeat(43);
|
||||||
|
const state = "s".repeat(43);
|
||||||
|
|
||||||
|
const keys = generateKeyPairSync("rsa", { modulusLength: 2048 });
|
||||||
|
const jwk = { ...keys.publicKey.export({ format: "jwk" }), kid: "test-key", use: "sig", alg: "RS256" };
|
||||||
|
|
||||||
|
function token(claims: Record<string, unknown>, invalidSignature = false): string {
|
||||||
|
const encode = (value: unknown) => Buffer.from(JSON.stringify(value)).toString("base64url");
|
||||||
|
const input = `${encode({ alg: "RS256", kid: "test-key", typ: "JWT" })}.${encode(claims)}`;
|
||||||
|
const signer = createSign("RSA-SHA256");
|
||||||
|
signer.update(input);
|
||||||
|
signer.end();
|
||||||
|
const signature = signer.sign(keys.privateKey).toString("base64url");
|
||||||
|
const corruptedSignature = signature.startsWith("A") ? `B${signature.slice(1)}` : `A${signature.slice(1)}`;
|
||||||
|
return `${input}.${invalidSignature ? corruptedSignature : signature}`;
|
||||||
|
}
|
||||||
|
|
||||||
|
function protocol(options: {
|
||||||
|
claims?: Record<string, unknown>;
|
||||||
|
discoveryIssuer?: string;
|
||||||
|
invalidSignature?: boolean;
|
||||||
|
seen?: URL[];
|
||||||
|
} = {}) {
|
||||||
|
const now = Math.floor(Date.now() / 1000);
|
||||||
|
const claims = {
|
||||||
|
iss: issuer,
|
||||||
|
sub: "user-123",
|
||||||
|
aud: clientId,
|
||||||
|
exp: now + 300,
|
||||||
|
iat: now,
|
||||||
|
nonce,
|
||||||
|
name: "Ada Lovelace",
|
||||||
|
groups: ["TOT Users", "Unmapped group"],
|
||||||
|
...options.claims,
|
||||||
|
};
|
||||||
|
const fetch = async (input: RequestInfo | URL) => {
|
||||||
|
const url = new URL(input instanceof Request ? input.url : typeof input === "string" ? input : input.toString());
|
||||||
|
options.seen?.push(url);
|
||||||
|
if (url.pathname.includes(".well-known/")) {
|
||||||
|
return Response.json({
|
||||||
|
issuer: options.discoveryIssuer ?? issuer,
|
||||||
|
authorization_endpoint: `${issuer}/authorize`,
|
||||||
|
token_endpoint: `${issuer}/token`,
|
||||||
|
jwks_uri: `${issuer}/jwks`,
|
||||||
|
response_types_supported: ["code"],
|
||||||
|
grant_types_supported: ["authorization_code"],
|
||||||
|
subject_types_supported: ["public"],
|
||||||
|
id_token_signing_alg_values_supported: ["RS256"],
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if (url.pathname === "/jwks") return Response.json({ keys: [jwk] });
|
||||||
|
if (url.pathname === "/token") {
|
||||||
|
return Response.json({
|
||||||
|
token_type: "Bearer",
|
||||||
|
access_token: "access-token-must-not-be-persisted",
|
||||||
|
refresh_token: "refresh-token-must-not-be-persisted",
|
||||||
|
id_token: token(claims, options.invalidSignature),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
return new Response(null, { status: 404 });
|
||||||
|
};
|
||||||
|
return createOidcProtocol({
|
||||||
|
issuer,
|
||||||
|
clientId,
|
||||||
|
clientSecret: "client-secret-must-not-be-persisted",
|
||||||
|
callbackUrl,
|
||||||
|
scopes: ["openid", "profile"],
|
||||||
|
groupsClaim: "groups",
|
||||||
|
fetch,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
async function callback(subject = protocol()) {
|
||||||
|
return subject.callback({
|
||||||
|
currentUrl: new URL(`${callbackUrl}?code=good&state=${state}`), state, nonce, codeVerifier: verifier,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
test("uses HTTPS discovery, Authorization Code, and PKCE S256 without external network", async () => {
|
||||||
|
const seen: URL[] = [];
|
||||||
|
const subject = protocol({ seen });
|
||||||
|
const authorization = await subject.authorizationUrl({ state, nonce, codeVerifier: verifier });
|
||||||
|
|
||||||
|
expect(authorization.origin).toBe(issuer);
|
||||||
|
expect(authorization.pathname).toBe("/authorize");
|
||||||
|
expect(Object.fromEntries(authorization.searchParams)).toMatchObject({
|
||||||
|
response_type: "code", client_id: clientId, redirect_uri: callbackUrl, state, nonce,
|
||||||
|
code_challenge_method: "S256", scope: "openid profile",
|
||||||
|
});
|
||||||
|
expect(authorization.searchParams.get("code_challenge")).not.toBe(verifier);
|
||||||
|
await expect(callback(subject)).resolves.toEqual({
|
||||||
|
issuer, subject: "user-123", displayName: "Ada Lovelace",
|
||||||
|
groups: ["TOT Users", "Unmapped group"], tokenExpiresAt: expect.any(Date),
|
||||||
|
});
|
||||||
|
expect(seen.map((url) => url.origin)).toEqual([issuer, issuer, issuer]);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("rejects non-HTTPS issuer configuration and a discovery issuer mismatch", async () => {
|
||||||
|
expect(() => createOidcProtocol({
|
||||||
|
issuer: "http://issuer.example.test", clientId, clientSecret: "secret", callbackUrl,
|
||||||
|
scopes: ["openid"], groupsClaim: "groups",
|
||||||
|
})).toThrow(OidcProtocolError);
|
||||||
|
await expect(protocol({ discoveryIssuer: "https://other.example.test" }).authorizationUrl({ state, nonce, codeVerifier: verifier }))
|
||||||
|
.rejects.toThrow(OidcProtocolError);
|
||||||
|
});
|
||||||
|
|
||||||
|
test.each([
|
||||||
|
["state", new URL(`${callbackUrl}?code=good&state=wrong`), {}],
|
||||||
|
["nonce", new URL(`${callbackUrl}?code=good&state=${state}`), { nonce: "wrong" }],
|
||||||
|
["audience", new URL(`${callbackUrl}?code=good&state=${state}`), { aud: "someone-else" }],
|
||||||
|
["issuer", new URL(`${callbackUrl}?code=good&state=${state}`), { iss: "https://other.example.test" }],
|
||||||
|
["expiry", new URL(`${callbackUrl}?code=good&state=${state}`), { exp: Math.floor(Date.now() / 1000) - 1 }],
|
||||||
|
["subject", new URL(`${callbackUrl}?code=good&state=${state}`), { sub: undefined }],
|
||||||
|
])("rejects invalid %s claims or callback bindings", async (_label, currentUrl, claims) => {
|
||||||
|
const subject = protocol({ claims });
|
||||||
|
await expect(subject.callback({ currentUrl, state, nonce, codeVerifier: verifier })).rejects.toThrow(OidcProtocolError);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("rejects invalid ID-token signatures", async () => {
|
||||||
|
await expect(callback(protocol({ invalidSignature: true }))).rejects.toThrow(OidcProtocolError);
|
||||||
|
});
|
||||||
|
|
||||||
|
test.each([
|
||||||
|
["absent", { groups: undefined }],
|
||||||
|
["non-array", { groups: "TOT Users" }],
|
||||||
|
["empty", { groups: [""] }],
|
||||||
|
["duplicate", { groups: ["TOT Users", "TOT Users"] }],
|
||||||
|
["control", { groups: ["TOT\u0000Users"] }],
|
||||||
|
["oversized", { groups: ["x".repeat(257)] }],
|
||||||
|
["distributed", { _claim_names: { groups: "source" }, _claim_sources: { source: { endpoint: "https://issuer.example.test/claims" } } }],
|
||||||
|
["overage", { hasgroups: true }],
|
||||||
|
])("rejects %s mandatory groups claims", async (_label, claims) => {
|
||||||
|
await expect(callback(protocol({ claims }))).rejects.toThrow(OidcProtocolError);
|
||||||
|
});
|
||||||
@@ -0,0 +1,64 @@
|
|||||||
|
import { beforeEach, expect, test, vi } from "vitest";
|
||||||
|
import { http, HttpResponse } from "msw";
|
||||||
|
import { beginOidcLogin, logout } from "./auth";
|
||||||
|
import { server } from "../test/msw";
|
||||||
|
import { clearAuthState, setAuthState } from "../auth/authState";
|
||||||
|
|
||||||
|
const user = {
|
||||||
|
issuer: "local", subject: "user-a", roles: ["user"] as const,
|
||||||
|
permissions: ["session.use"], isAdmin: false, csrfToken: "a".repeat(43), session: null,
|
||||||
|
};
|
||||||
|
|
||||||
|
function deferred() {
|
||||||
|
let resolve!: () => void;
|
||||||
|
const promise = new Promise<void>((onResolve) => { resolve = onResolve; });
|
||||||
|
return { promise, resolve };
|
||||||
|
}
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
clearAuthState();
|
||||||
|
setAuthState(user);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("OIDC navigation waits for a successful pending logout response", async () => {
|
||||||
|
const gate = deferred();
|
||||||
|
let logoutStarted!: () => void;
|
||||||
|
const started = new Promise<void>((resolve) => { logoutStarted = resolve; });
|
||||||
|
server.use(http.post("/api/auth/logout", async () => {
|
||||||
|
logoutStarted();
|
||||||
|
await gate.promise;
|
||||||
|
return new HttpResponse(null, { status: 204 });
|
||||||
|
}));
|
||||||
|
const logoutPromise = logout();
|
||||||
|
await started;
|
||||||
|
const navigate = vi.fn();
|
||||||
|
const oidc = beginOidcLogin(navigate);
|
||||||
|
|
||||||
|
await Promise.resolve();
|
||||||
|
expect(navigate).not.toHaveBeenCalled();
|
||||||
|
gate.resolve();
|
||||||
|
await Promise.all([logoutPromise, oidc]);
|
||||||
|
expect(navigate).toHaveBeenCalledOnce();
|
||||||
|
expect(navigate).toHaveBeenCalledWith("/api/auth/oidc/login");
|
||||||
|
});
|
||||||
|
|
||||||
|
test("OIDC navigation waits for a failed pending logout response before continuing", async () => {
|
||||||
|
const gate = deferred();
|
||||||
|
let logoutStarted!: () => void;
|
||||||
|
const started = new Promise<void>((resolve) => { logoutStarted = resolve; });
|
||||||
|
server.use(http.post("/api/auth/logout", async () => {
|
||||||
|
logoutStarted();
|
||||||
|
await gate.promise;
|
||||||
|
return HttpResponse.json({ code: "auth_unavailable" }, { status: 503 });
|
||||||
|
}));
|
||||||
|
const logoutPromise = logout().catch(() => undefined);
|
||||||
|
await started;
|
||||||
|
const navigate = vi.fn();
|
||||||
|
const oidc = beginOidcLogin(navigate);
|
||||||
|
|
||||||
|
await Promise.resolve();
|
||||||
|
expect(navigate).not.toHaveBeenCalled();
|
||||||
|
gate.resolve();
|
||||||
|
await Promise.all([logoutPromise, oidc]);
|
||||||
|
expect(navigate).toHaveBeenCalledWith("/api/auth/oidc/login");
|
||||||
|
});
|
||||||
@@ -6,6 +6,7 @@ const authModes = new Set<AuthPublicConfig["mode"]>(["local", "oidc", "upstream"
|
|||||||
const roles = new Set<AuthRole>(["user", "admin"]);
|
const roles = new Set<AuthRole>(["user", "admin"]);
|
||||||
const sessionMethods = new Set<AuthSessionInfo["method"]>(["local", "oidc", "upstream"]);
|
const sessionMethods = new Set<AuthSessionInfo["method"]>(["local", "oidc", "upstream"]);
|
||||||
let pendingLogoutResponse: Promise<void> | null = null;
|
let pendingLogoutResponse: Promise<void> | null = null;
|
||||||
|
const oidcLoginPath = "/api/auth/oidc/login";
|
||||||
|
|
||||||
function record(value: unknown): Record<string, unknown> | undefined {
|
function record(value: unknown): Record<string, unknown> | undefined {
|
||||||
return value && typeof value === "object" && !Array.isArray(value)
|
return value && typeof value === "object" && !Array.isArray(value)
|
||||||
@@ -86,6 +87,23 @@ export async function loginLocal(username: string, password: string, remember: b
|
|||||||
return user;
|
return user;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* OIDC creates its browser session through a top-level same-origin navigation. A previous logout
|
||||||
|
* response may still carry a clearing Set-Cookie, so it must settle before this navigation begins.
|
||||||
|
*/
|
||||||
|
export async function beginOidcLogin(navigate: (path: string) => void = (path) => window.location.assign(path)): Promise<void> {
|
||||||
|
const pending = pendingLogoutResponse;
|
||||||
|
if (pending) {
|
||||||
|
try {
|
||||||
|
await pending;
|
||||||
|
} catch {
|
||||||
|
// A failed logout must release the coordinator; the current browser cookie remains the
|
||||||
|
// backend's authority during the following OIDC handshake.
|
||||||
|
}
|
||||||
|
}
|
||||||
|
navigate(oidcLoginPath);
|
||||||
|
}
|
||||||
|
|
||||||
export async function logout(): Promise<boolean> {
|
export async function logout(): Promise<boolean> {
|
||||||
const logoutGeneration = getAuthGeneration();
|
const logoutGeneration = getAuthGeneration();
|
||||||
if (pendingLogoutResponse) {
|
if (pendingLogoutResponse) {
|
||||||
|
|||||||
@@ -76,14 +76,12 @@ describe("LoginPage", () => {
|
|||||||
expect(screen.getByLabelText(/password/i)).toHaveValue("");
|
expect(screen.getByLabelText(/password/i)).toHaveValue("");
|
||||||
});
|
});
|
||||||
|
|
||||||
test("shows OIDC only when public configuration enables it and uses same-origin navigation", () => {
|
test("shows OIDC only when public configuration enables it", () => {
|
||||||
const { rerender } = render(<LoginPage config={localConfig} onAuthenticated={vi.fn()} />);
|
const { rerender } = render(<LoginPage config={localConfig} onAuthenticated={vi.fn()} />);
|
||||||
expect(screen.queryByRole("link", { name: /single sign-on/i })).not.toBeInTheDocument();
|
expect(screen.queryByRole("link", { name: /single sign-on/i })).not.toBeInTheDocument();
|
||||||
|
|
||||||
rerender(<LoginPage config={oidcConfig} onAuthenticated={vi.fn()} />);
|
rerender(<LoginPage config={oidcConfig} onAuthenticated={vi.fn()} />);
|
||||||
expect(screen.getByRole("link", { name: /single sign-on/i })).toHaveAttribute(
|
expect(screen.getByRole("button", { name: /single sign-on/i })).toBeEnabled();
|
||||||
"href", "/api/auth/oidc/login",
|
|
||||||
);
|
|
||||||
});
|
});
|
||||||
|
|
||||||
test("does not dispatch local login until an in-flight logout response settles", async () => {
|
test("does not dispatch local login until an in-flight logout response settles", async () => {
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import { useEffect, useRef, useState } from "react";
|
|||||||
import type { FormEvent } from "react";
|
import type { FormEvent } from "react";
|
||||||
import { AlertTriangle, ArrowRight, LockKeyhole } from "lucide-react";
|
import { AlertTriangle, ArrowRight, LockKeyhole } from "lucide-react";
|
||||||
import { ApiError } from "../api/client";
|
import { ApiError } from "../api/client";
|
||||||
import { loginLocal } from "../api/auth";
|
import { beginOidcLogin, loginLocal } from "../api/auth";
|
||||||
import type { AuthenticatedUser, AuthPublicConfig } from "../api/types";
|
import type { AuthenticatedUser, AuthPublicConfig } from "../api/types";
|
||||||
import { Button } from "../components/ui/button";
|
import { Button } from "../components/ui/button";
|
||||||
|
|
||||||
@@ -24,6 +24,7 @@ function loginError(error: unknown): { message: string; retry: boolean } {
|
|||||||
|
|
||||||
export function LoginPage({ config, onAuthenticated, onRetry }: LoginPageProps) {
|
export function LoginPage({ config, onAuthenticated, onRetry }: LoginPageProps) {
|
||||||
const localLogin = config.mode === "local" && config.localLogin;
|
const localLogin = config.mode === "local" && config.localLogin;
|
||||||
|
const oidcLogin = config.mode === "oidc" && config.oidcLogin;
|
||||||
const formRef = useRef<HTMLFormElement>(null);
|
const formRef = useRef<HTMLFormElement>(null);
|
||||||
const passwordRef = useRef<HTMLInputElement>(null);
|
const passwordRef = useRef<HTMLInputElement>(null);
|
||||||
const mountedRef = useRef(true);
|
const mountedRef = useRef(true);
|
||||||
@@ -65,6 +66,10 @@ export function LoginPage({ config, onAuthenticated, onRetry }: LoginPageProps)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function startOidcLogin() {
|
||||||
|
void beginOidcLogin();
|
||||||
|
}
|
||||||
|
|
||||||
return (
|
return (
|
||||||
<main className="min-h-screen bg-background px-5 py-8 text-foreground sm:px-8 sm:py-12">
|
<main className="min-h-screen bg-background px-5 py-8 text-foreground sm:px-8 sm:py-12">
|
||||||
<div className="mx-auto grid min-h-[calc(100vh-4rem)] max-w-5xl items-center gap-12 lg:grid-cols-[minmax(0,1fr)_26rem]">
|
<div className="mx-auto grid min-h-[calc(100vh-4rem)] max-w-5xl items-center gap-12 lg:grid-cols-[minmax(0,1fr)_26rem]">
|
||||||
@@ -138,6 +143,12 @@ export function LoginPage({ config, onAuthenticated, onRetry }: LoginPageProps)
|
|||||||
</Button>
|
</Button>
|
||||||
</form>
|
</form>
|
||||||
)}
|
)}
|
||||||
|
{oidcLogin && (
|
||||||
|
<Button type="button" size="lg" className="w-full" onClick={startOidcLogin}>
|
||||||
|
Continue with single sign-on
|
||||||
|
<ArrowRight aria-hidden="true" />
|
||||||
|
</Button>
|
||||||
|
)}
|
||||||
|
|
||||||
{config.oidcLogin && (
|
{config.oidcLogin && (
|
||||||
<a
|
<a
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
import { act, render, screen, waitFor } from "@testing-library/react";
|
||||||
|
import userEvent from "@testing-library/user-event";
|
||||||
|
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||||
|
import { beforeEach, expect, test, vi } from "vitest";
|
||||||
|
import { http, HttpResponse } from "msw";
|
||||||
|
import { AppShell } from "./AppShell";
|
||||||
|
import { server } from "../test/msw";
|
||||||
|
import { FakeEventSource } from "../test/fakeEventSource";
|
||||||
|
import { clearAuthState, setAuthState } from "../auth/authState";
|
||||||
|
import { useSessionStore } from "../store/sessionStore";
|
||||||
|
|
||||||
|
function renderShell() {
|
||||||
|
const client = new QueryClient({ defaultOptions: { queries: { retry: false } } });
|
||||||
|
return render(<QueryClientProvider client={client}><AppShell /></QueryClientProvider>);
|
||||||
|
}
|
||||||
|
|
||||||
|
function deferred() {
|
||||||
|
let resolve!: () => void;
|
||||||
|
const promise = new Promise<void>((onResolve) => { resolve = onResolve; });
|
||||||
|
return { promise, resolve };
|
||||||
|
}
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
clearAuthState();
|
||||||
|
setAuthState({
|
||||||
|
issuer: "local", subject: "user-a", roles: ["user"], permissions: ["session.use"],
|
||||||
|
isAdmin: false, csrfToken: null, session: null,
|
||||||
|
});
|
||||||
|
FakeEventSource.instances = [];
|
||||||
|
(globalThis as { EventSource: typeof EventSource }).EventSource = FakeEventSource as unknown as typeof EventSource;
|
||||||
|
useSessionStore.getState().resetSession();
|
||||||
|
server.use(
|
||||||
|
http.get("/api/me", () => HttpResponse.json({ issuer: "local", subject: "user-a", isAdmin: false })),
|
||||||
|
http.get("/api/settings", () => HttpResponse.json({ workspace: "default", provider: "test", model: "test", thinking: "low" })),
|
||||||
|
http.get("/api/workspaces", () => HttpResponse.json([])),
|
||||||
|
http.get("/api/models", () => HttpResponse.json({ models: [] })),
|
||||||
|
http.post("/api/runtime/prewarm", () => new HttpResponse(null, { status: 202 })),
|
||||||
|
);
|
||||||
|
});
|
||||||
|
|
||||||
|
test("a held stop for s1 cannot reset the newer active s2 session", async () => {
|
||||||
|
const closeGate = deferred();
|
||||||
|
const closeStarted = deferred();
|
||||||
|
const active = (id: string, question: string) => ({
|
||||||
|
id, status: "open", question, summary: null, created_at: "2026-01-02T00:00:00Z",
|
||||||
|
updated_at: null, author: null, name: null, group: null, archived: false, active: true,
|
||||||
|
});
|
||||||
|
server.use(
|
||||||
|
http.get("/api/sessions", () => HttpResponse.json([active("s1", "Active one"), active("s2", "Active two")])),
|
||||||
|
http.post("/api/sessions/:id/resume", ({ params }) => HttpResponse.json({ id: params.id, alreadyActive: false })),
|
||||||
|
http.get("/api/sessions/:id", ({ params }) => HttpResponse.json({ id: params.id, status: "open", phase: 1 })),
|
||||||
|
http.post("/api/sessions/s1/close", async () => {
|
||||||
|
closeStarted.resolve();
|
||||||
|
await closeGate.promise;
|
||||||
|
return new HttpResponse(null, { status: 204 });
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
renderShell();
|
||||||
|
await userEvent.click(await screen.findByTestId("session-item-s1"));
|
||||||
|
await waitFor(() => expect(FakeEventSource.instances.at(-1)?.url).toContain("/sessions/s1/events"));
|
||||||
|
await userEvent.click(screen.getByRole("button", { name: /stop and save session/i }));
|
||||||
|
await userEvent.click(await screen.findByRole("button", { name: "Stop & save" }));
|
||||||
|
await closeStarted.promise;
|
||||||
|
|
||||||
|
await userEvent.click(screen.getByTestId("session-item-s2"));
|
||||||
|
await waitFor(() => expect(FakeEventSource.instances.at(-1)?.url).toContain("/sessions/s2/events"));
|
||||||
|
act(() => useSessionStore.setState({ currentPhase: "F2" }));
|
||||||
|
closeGate.resolve();
|
||||||
|
await new Promise((resolve) => setImmediate(resolve));
|
||||||
|
|
||||||
|
expect(FakeEventSource.instances.at(-1)?.url).toContain("/sessions/s2/events");
|
||||||
|
expect(FakeEventSource.instances.at(-1)?.closed).toBe(false);
|
||||||
|
expect(useSessionStore.getState().currentPhase).toBe("F2");
|
||||||
|
});
|
||||||
|
|
||||||
|
test("a held new-session completion cannot replace the newer active s2 target", async () => {
|
||||||
|
const createGate = deferred();
|
||||||
|
const createStarted = deferred();
|
||||||
|
server.use(
|
||||||
|
http.get("/api/sessions", () => HttpResponse.json([{
|
||||||
|
id: "s2", status: "open", question: "Current session", summary: null,
|
||||||
|
created_at: "2026-01-02T00:00:00Z", updated_at: null, author: null, name: null,
|
||||||
|
group: null, archived: false, active: true,
|
||||||
|
}])),
|
||||||
|
http.post("/api/sessions", async () => {
|
||||||
|
createStarted.resolve();
|
||||||
|
await createGate.promise;
|
||||||
|
return HttpResponse.json({ id: "s3" });
|
||||||
|
}),
|
||||||
|
http.post("/api/sessions/:id/resume", ({ params }) => HttpResponse.json({ id: params.id, alreadyActive: false })),
|
||||||
|
http.get("/api/sessions/:id", ({ params }) => HttpResponse.json({ id: params.id, status: "open", phase: 1 })),
|
||||||
|
);
|
||||||
|
renderShell();
|
||||||
|
const composer = screen.getByRole("textbox", { name: /new question/i });
|
||||||
|
await userEvent.type(composer, "Held new question");
|
||||||
|
await userEvent.click(screen.getByRole("button", { name: /send/i }));
|
||||||
|
await createStarted.promise;
|
||||||
|
|
||||||
|
await userEvent.click(screen.getByTestId("session-item-s2"));
|
||||||
|
await waitFor(() => expect(FakeEventSource.instances.at(-1)?.url).toContain("/sessions/s2/events"));
|
||||||
|
act(() => useSessionStore.setState({ currentPhase: "F2" }));
|
||||||
|
createGate.resolve();
|
||||||
|
await new Promise((resolve) => setImmediate(resolve));
|
||||||
|
|
||||||
|
expect(FakeEventSource.instances.at(-1)?.url).toContain("/sessions/s2/events");
|
||||||
|
expect(FakeEventSource.instances.at(-1)?.closed).toBe(false);
|
||||||
|
expect(useSessionStore.getState().currentPhase).toBe("F2");
|
||||||
|
});
|
||||||
@@ -80,6 +80,8 @@ export function AppShell() {
|
|||||||
} as CSSProperties;
|
} as CSSProperties;
|
||||||
const [activeSessionId, setActiveSessionId] = useState<string | null>(null);
|
const [activeSessionId, setActiveSessionId] = useState<string | null>(null);
|
||||||
const activeSessionIdRef = useRef<string | null>(null);
|
const activeSessionIdRef = useRef<string | null>(null);
|
||||||
|
const activeSessionEpochRef = useRef(0);
|
||||||
|
const newSessionOperationRef = useRef<{ target: string | null; epoch: number } | null>(null);
|
||||||
const resumeInvocationRef = useRef(0);
|
const resumeInvocationRef = useRef(0);
|
||||||
const latestResumeIntentRef = useRef<{ token: number; id: string } | null>(null);
|
const latestResumeIntentRef = useRef<{ token: number; id: string } | null>(null);
|
||||||
const resumeInFlightRef = useRef(new Map<string, {
|
const resumeInFlightRef = useRef(new Map<string, {
|
||||||
@@ -165,6 +167,7 @@ export function AppShell() {
|
|||||||
|
|
||||||
function selectActiveSession(id: string | null) {
|
function selectActiveSession(id: string | null) {
|
||||||
// Keep async Resume completions synchronized before React commits the state update.
|
// Keep async Resume completions synchronized before React commits the state update.
|
||||||
|
if (activeSessionIdRef.current !== id) activeSessionEpochRef.current += 1;
|
||||||
activeSessionIdRef.current = id;
|
activeSessionIdRef.current = id;
|
||||||
setActiveSessionId(id);
|
setActiveSessionId(id);
|
||||||
}
|
}
|
||||||
@@ -505,6 +508,7 @@ export function AppShell() {
|
|||||||
|
|
||||||
function startNewSession() {
|
function startNewSession() {
|
||||||
invalidateResumeIntent();
|
invalidateResumeIntent();
|
||||||
|
newSessionOperationRef.current = null;
|
||||||
resetSession();
|
resetSession();
|
||||||
// Starting a new question closes any open session detail panel: the reader is
|
// Starting a new question closes any open session detail panel: the reader is
|
||||||
// moving away from that session, so its left-hand box must not linger.
|
// moving away from that session, so its left-hand box must not linger.
|
||||||
@@ -519,11 +523,18 @@ export function AppShell() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
function beginSessionCreation() {
|
function beginSessionCreation() {
|
||||||
|
newSessionOperationRef.current = {
|
||||||
|
target: activeSessionIdRef.current,
|
||||||
|
epoch: activeSessionEpochRef.current,
|
||||||
|
};
|
||||||
setAwaitingQuestion(false);
|
setAwaitingQuestion(false);
|
||||||
setCreatingSession(true);
|
setCreatingSession(true);
|
||||||
}
|
}
|
||||||
|
|
||||||
function finishSessionCreation(id: string) {
|
function finishSessionCreation(id: string) {
|
||||||
|
const operation = newSessionOperationRef.current;
|
||||||
|
newSessionOperationRef.current = null;
|
||||||
|
if (!operation || operation.target !== activeSessionIdRef.current || operation.epoch !== activeSessionEpochRef.current) return;
|
||||||
// React batches these updates, preserving the provisional session view
|
// React batches these updates, preserving the provisional session view
|
||||||
// while useSessionStream opens the durable session's SSE channel.
|
// while useSessionStream opens the durable session's SSE channel.
|
||||||
selectActiveSession(id);
|
selectActiveSession(id);
|
||||||
@@ -533,21 +544,25 @@ export function AppShell() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
function failSessionCreation(message?: string) {
|
function failSessionCreation(message?: string) {
|
||||||
|
const operation = newSessionOperationRef.current;
|
||||||
|
newSessionOperationRef.current = null;
|
||||||
|
if (!operation || operation.target !== activeSessionIdRef.current || operation.epoch !== activeSessionEpochRef.current) return;
|
||||||
setCreatingSession(false);
|
setCreatingSession(false);
|
||||||
resetSession();
|
resetSession();
|
||||||
toast.error(message ?? "Failed to create session. Your question is ready to retry.");
|
toast.error(message ?? "Failed to create session. Your question is ready to retry.");
|
||||||
}
|
}
|
||||||
|
|
||||||
async function stopSession() {
|
async function stopSession() {
|
||||||
if (!activeSessionId) return;
|
const id = activeSessionIdRef.current;
|
||||||
const id = activeSessionId;
|
if (!id) return;
|
||||||
const guard = captureAuthOperation({ sessionId: id, disposalEpoch: operationEpochRef.current });
|
const guard = captureAuthOperation({ sessionId: id, disposalEpoch: operationEpochRef.current });
|
||||||
if (!guard) return;
|
if (!guard) return;
|
||||||
invalidateResumeIntent();
|
invalidateResumeIntent();
|
||||||
try {
|
try {
|
||||||
await closeSession(id);
|
await closeSession(id);
|
||||||
} finally {
|
} finally {
|
||||||
if (!isAuthOperationCurrent(guard, { sessionId: id, disposalEpoch: operationEpochRef.current })) return;
|
if (!isAuthOperationCurrent(guard, { sessionId: id, disposalEpoch: operationEpochRef.current })
|
||||||
|
|| activeSessionIdRef.current !== id) return;
|
||||||
resetSession();
|
resetSession();
|
||||||
selectActiveSession(null);
|
selectActiveSession(null);
|
||||||
setAwaitingQuestion(false);
|
setAwaitingQuestion(false);
|
||||||
@@ -643,7 +658,7 @@ export function AppShell() {
|
|||||||
</header>
|
</header>
|
||||||
)}
|
)}
|
||||||
<CentralStatus working={working} />
|
<CentralStatus working={working} />
|
||||||
{activeSessionId && <WidgetHost sessionId={activeSessionId} />}
|
{activeSessionId && <WidgetHost key={`widget:${activeSessionId}:${activeSessionEpochRef.current}`} sessionId={activeSessionId} />}
|
||||||
{finalized && !agentActive && (
|
{finalized && !agentActive && (
|
||||||
<div className="rounded-2xl border border-border/80 bg-card p-5 text-center shadow-md">
|
<div className="rounded-2xl border border-border/80 bg-card p-5 text-center shadow-md">
|
||||||
<p className="text-sm text-muted-foreground">
|
<p className="text-sm text-muted-foreground">
|
||||||
@@ -668,6 +683,7 @@ export function AppShell() {
|
|||||||
<div className="rounded-2xl border border-border/80 bg-card shadow-md transition-colors focus-within:border-primary/50 focus-within:ring-3 focus-within:ring-ring/15">
|
<div className="rounded-2xl border border-border/80 bg-card shadow-md transition-colors focus-within:border-primary/50 focus-within:ring-3 focus-within:ring-ring/15">
|
||||||
<div className="px-2.5 py-2">
|
<div className="px-2.5 py-2">
|
||||||
<SteerInput
|
<SteerInput
|
||||||
|
key={`steer:${activeSessionId ?? "new"}:${activeSessionEpochRef.current}`}
|
||||||
sessionId={activeSessionId}
|
sessionId={activeSessionId}
|
||||||
onSessionCreating={beginSessionCreation}
|
onSessionCreating={beginSessionCreation}
|
||||||
onSessionCreated={finishSessionCreation}
|
onSessionCreated={finishSessionCreation}
|
||||||
|
|||||||
@@ -113,6 +113,38 @@ test("a delayed steer from user A cannot mutate user B's store or composer", asy
|
|||||||
view.unmount();
|
view.unmount();
|
||||||
});
|
});
|
||||||
|
|
||||||
|
test("a held steer for s1 cannot complete into the active s2 operation scope", async () => {
|
||||||
|
let release!: () => void;
|
||||||
|
let started!: () => void;
|
||||||
|
let settled!: () => void;
|
||||||
|
const held = new Promise<void>((resolve) => { release = resolve; });
|
||||||
|
const requestStarted = new Promise<void>((resolve) => { started = resolve; });
|
||||||
|
const requestSettled = new Promise<void>((resolve) => { settled = resolve; });
|
||||||
|
server.use(http.post("/api/sessions/s1/steer", async () => {
|
||||||
|
started();
|
||||||
|
try {
|
||||||
|
await held;
|
||||||
|
return new HttpResponse(null, { status: 204 });
|
||||||
|
} finally {
|
||||||
|
settled();
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
|
||||||
|
const view = render(<SteerInput sessionId="s1" />);
|
||||||
|
const input = screen.getByRole("textbox");
|
||||||
|
await userEvent.type(input, "Keep s1 isolated");
|
||||||
|
await userEvent.click(screen.getByRole("button", { name: /invia|send|steer/i }));
|
||||||
|
await requestStarted;
|
||||||
|
|
||||||
|
view.rerender(<SteerInput sessionId="s2" />);
|
||||||
|
act(() => useSessionStore.getState().setLastUserEntry({ kind: "input", text: "s2 owned" }));
|
||||||
|
release();
|
||||||
|
await act(async () => { await requestSettled; });
|
||||||
|
|
||||||
|
expect(useSessionStore.getState().lastUserEntry).toEqual({ kind: "input", text: "s2 owned" });
|
||||||
|
expect(input).toHaveValue("Keep s1 isolated");
|
||||||
|
});
|
||||||
|
|
||||||
test("a settings preflight from user A prevents session POST after user B logs in", async () => {
|
test("a settings preflight from user A prevents session POST after user B logs in", async () => {
|
||||||
let releaseSettings!: () => void;
|
let releaseSettings!: () => void;
|
||||||
let settingsStarted!: () => void;
|
let settingsStarted!: () => void;
|
||||||
|
|||||||
@@ -51,6 +51,12 @@ export function SteerInput({
|
|||||||
const taRef = useRef<HTMLTextAreaElement | null>(null);
|
const taRef = useRef<HTMLTextAreaElement | null>(null);
|
||||||
const operationEpochRef = useRef(0);
|
const operationEpochRef = useRef(0);
|
||||||
useEffect(() => () => { operationEpochRef.current += 1; }, []);
|
useEffect(() => () => { operationEpochRef.current += 1; }, []);
|
||||||
|
useEffect(() => {
|
||||||
|
// A new active session owns a new composer operation scope. Invalidate any held request
|
||||||
|
// before it can write the old session's completion into the new target.
|
||||||
|
operationEpochRef.current += 1;
|
||||||
|
setBusy(false);
|
||||||
|
}, [sessionId]);
|
||||||
|
|
||||||
// Merge our own ref (for autosizing) with the forwarded one (parent focus).
|
// Merge our own ref (for autosizing) with the forwarded one (parent focus).
|
||||||
function attachRef(el: HTMLTextAreaElement | null) {
|
function attachRef(el: HTMLTextAreaElement | null) {
|
||||||
|
|||||||
@@ -92,3 +92,37 @@ test("a delayed gate response from user A cannot clear user B's pending gate", a
|
|||||||
expect(useSessionStore.getState().pendingWidget?.id).toBe("gate-b");
|
expect(useSessionStore.getState().pendingWidget?.id).toBe("gate-b");
|
||||||
expect(useSessionStore.getState().lastUserEntry).toBeNull();
|
expect(useSessionStore.getState().lastUserEntry).toBeNull();
|
||||||
});
|
});
|
||||||
|
|
||||||
|
test("a held s1 gate response cannot clear the pending s2 gate", async () => {
|
||||||
|
let release!: () => void;
|
||||||
|
let started!: () => void;
|
||||||
|
let settled!: () => void;
|
||||||
|
const held = new Promise<void>((resolve) => { release = resolve; });
|
||||||
|
const requestStarted = new Promise<void>((resolve) => { started = resolve; });
|
||||||
|
const requestSettled = new Promise<void>((resolve) => { settled = resolve; });
|
||||||
|
server.use(http.post("/api/sessions/s1/response", async () => {
|
||||||
|
started();
|
||||||
|
try {
|
||||||
|
await held;
|
||||||
|
return new HttpResponse(null, { status: 204 });
|
||||||
|
} finally {
|
||||||
|
settled();
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
useSessionStore.setState({
|
||||||
|
pendingWidget: { id: "gate-s1", widget: "select", options: [{ id: "s1", label: "Answer s1" }] },
|
||||||
|
});
|
||||||
|
const view = render(<WidgetHost sessionId="s1" />);
|
||||||
|
await userEvent.click(screen.getByRole("button", { name: "Answer s1" }));
|
||||||
|
await requestStarted;
|
||||||
|
|
||||||
|
view.rerender(<WidgetHost sessionId="s2" />);
|
||||||
|
act(() => useSessionStore.setState({
|
||||||
|
pendingWidget: { id: "gate-s2", widget: "select", options: [{ id: "s2", label: "Answer s2" }] },
|
||||||
|
}));
|
||||||
|
release();
|
||||||
|
await act(async () => { await requestSettled; });
|
||||||
|
|
||||||
|
expect(useSessionStore.getState().pendingWidget?.id).toBe("gate-s2");
|
||||||
|
expect(useSessionStore.getState().lastUserEntry).toBeNull();
|
||||||
|
});
|
||||||
|
|||||||
@@ -15,6 +15,12 @@ export function WidgetHost({ sessionId }: { sessionId: string | null }) {
|
|||||||
const responseInFlight = useRef(false);
|
const responseInFlight = useRef(false);
|
||||||
const operationEpochRef = useRef(0);
|
const operationEpochRef = useRef(0);
|
||||||
useEffect(() => () => { operationEpochRef.current += 1; }, []);
|
useEffect(() => () => { operationEpochRef.current += 1; }, []);
|
||||||
|
useEffect(() => {
|
||||||
|
// A gate belongs to its active session target, not merely the authenticated user.
|
||||||
|
operationEpochRef.current += 1;
|
||||||
|
responseInFlight.current = false;
|
||||||
|
setResponding(false);
|
||||||
|
}, [sessionId]);
|
||||||
if (!pending) return null;
|
if (!pending) return null;
|
||||||
const Renderer = resolve(pending.widget);
|
const Renderer = resolve(pending.widget);
|
||||||
const onRespond = async (r: UiResponse) => {
|
const onRespond = async (r: UiResponse) => {
|
||||||
|
|||||||
Reference in New Issue
Block a user