diff --git a/backend/src/app.ts b/backend/src/app.ts index a6339f17..a027fcec 100644 --- a/backend/src/app.ts +++ b/backend/src/app.ts @@ -13,9 +13,11 @@ import type { PrincipalContext } from "./auth/principal.js"; import type { LoadedAuthConfig } from "./auth/types.js"; import { createCurrentLocalUserRegistryResolver, type LocalUserRegistry } from "./auth/local-registry.js"; import { AuthSessionOperationalError, createFileAuthSessionStore, type AuthSessionStore, type AuthSessionValidity } from "./auth/session-store.js"; +import type { WindowsAuthStorageBridge } from "./auth/windows-auth-storage.js"; import { registerAuthRoutes } from "./auth/routes.js"; -import { createOidcProtocol, type OidcProtocol } from "./auth/oidc-client.js"; +import { createOidcProtocol, type OidcProtocol, type OidcProtocolOptions } from "./auth/oidc-client.js"; import { isUsableAuthenticationSecret } from "./auth/secret-policy.js"; +import { secretValue } from "./config/secret-bundle.js"; import { sessionRoutes } from "./routes/sessions.js"; import { sqlRoutes } from "./routes/sql.js"; import { metaRoutes, type ListModelsFn } from "./routes/meta.js"; @@ -50,7 +52,11 @@ export interface BuildAppDeps { piManagement?: PiManagementService; localUserRegistry?: LocalUserRegistry; authSessionStore?: AuthSessionStore; + /** Explicit test-only transport seam; production always invokes the hidden tht bridge. */ + authStorageBridgeForTest?: WindowsAuthStorageBridge; oidcProtocol?: OidcProtocol; + /** Explicit test seam; production uses the provider-neutral OIDC constructor. */ + oidcProtocolFactory?: (options: OidcProtocolOptions) => OidcProtocol; } export interface AppWithAuthSessionStore extends FastifyInstance { @@ -198,12 +204,10 @@ export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstanc 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 (!isUsableAuthenticationSecret("THT_OIDC_CLIENT_SECRET", clientSecret)) { - return undefined; - } try { - return createOidcProtocol({ + const clientSecret = secretValue(config, loaded.value.oidc.clientSecretRef); + if (!isUsableAuthenticationSecret("THT_OIDC_CLIENT_SECRET", clientSecret)) return undefined; + return (deps?.oidcProtocolFactory ?? createOidcProtocol)({ issuer: loaded.value.oidc.issuer, clientId: loaded.value.oidc.clientId, clientSecret, @@ -234,7 +238,11 @@ export function buildApp(config: AppConfig, deps?: BuildAppDeps): FastifyInstanc throw new AuthSessionOperationalError(); } }, - }) + }, deps?.authStorageBridgeForTest === undefined + ? undefined + : process.platform === "win32" + ? { windowsStorageBridge: deps.authStorageBridgeForTest } + : { posixStorageBridge: deps.authStorageBridgeForTest }) : undefined); (app as AppWithAuthSessionStore).thothiiAuthSessionStore = authSessionStore; const authenticate = authenticateSession({ diff --git a/backend/src/auth/local-registry.ts b/backend/src/auth/local-registry.ts index 364e981f..8478b369 100644 --- a/backend/src/auth/local-registry.ts +++ b/backend/src/auth/local-registry.ts @@ -13,6 +13,7 @@ import { parseDocument } from "yaml"; import { z } from "zod"; import { isValidPasswordHash, verifyPassword, verifyWithDummy } from "./password.js"; import type { LoadedAuthConfig, Role } from "./types.js"; +import { createWindowsAuthStorageBridge, type WindowsAuthStorageBridge } from "./windows-auth-storage.js"; const MAX_USERS_YAML_BYTES = 1 << 20; const USERNAME_PATTERN = /^[A-Za-z0-9][A-Za-z0-9._@-]{2,63}$/; @@ -44,6 +45,11 @@ export interface CurrentLocalUserRegistryResolver { resolve(loaded: LoadedAuthConfig): LocalUserRegistry | undefined; } +/** Native Windows obtains protected registry bytes only from the hidden tht bridge. */ +export interface LocalUserRegistryOptions { + windowsStorageBridge?: Pick; +} + interface FileIdentity { dev: number; ino: number; @@ -220,10 +226,13 @@ function load(path: string): { records: LocalUserRecord[]; identity: RegistryIde return { records: parseRegistry(read.source), identity: read.identity }; } -export function createLocalUserRegistry(usersPath: string): LocalUserRegistry { +export function createLocalUserRegistry(usersPath: string, options: LocalUserRegistryOptions = {}): LocalUserRegistry { let cached: { records: LocalUserRecord[]; identity: RegistryIdentity } | undefined; + const windowsStorage = process.platform === "win32" + ? options.windowsStorageBridge ?? createWindowsAuthStorageBridge() + : undefined; - function current(): LocalUserRecord[] { + function currentPosix(): LocalUserRecord[] { try { const before = registryIdentity(usersPath); if (cached && sameIdentity(cached.identity, before)) return cached.records; @@ -240,22 +249,34 @@ export function createLocalUserRegistry(usersPath: string): LocalUserRegistry { throw invalid(); } - function operationalRecords(): LocalUserRecord[] { - const records = current(); + async function current(): Promise { + if (process.platform !== "win32") return currentPosix(); + try { + if (!windowsStorage) throw invalid(); + const contents = await windowsStorage.readLocalUsers(usersPath); + if (!Buffer.isBuffer(contents) || contents.length === 0 || contents.length > MAX_USERS_YAML_BYTES) throw invalid(); + return parseRegistry(new TextDecoder("utf-8", { fatal: true }).decode(contents)); + } catch { + throw invalid(); + } + } + + async function operationalRecords(): Promise { + const records = await current(); if (!records.some((user) => user.enabled && user.roles.includes("admin"))) throw invalid(); return records; } return { async hasEnabledAdmin(): Promise { - return current().some((user) => user.enabled && user.roles.includes("admin")); + return (await current()).some((user) => user.enabled && user.roles.includes("admin")); }, async findByUsername(username: string): Promise { const normalized = normalizeUsername(username); - return operationalRecords().find((user) => user.normalizedUsername === normalized); + return (await operationalRecords()).find((user) => user.normalizedUsername === normalized); }, async findBySubject(id: string): Promise { - return operationalRecords().find((user) => user.id === id); + return (await operationalRecords()).find((user) => user.id === id); }, async verify(user: LocalUserRecord | undefined, password: string): Promise { if (!user || !user.enabled) { @@ -267,14 +288,14 @@ export function createLocalUserRegistry(usersPath: string): LocalUserRegistry { }; } -export function createCurrentLocalUserRegistryResolver(): CurrentLocalUserRegistryResolver { +export function createCurrentLocalUserRegistryResolver(options: LocalUserRegistryOptions = {}): CurrentLocalUserRegistryResolver { let current: { usersPath: string; registry: LocalUserRegistry } | undefined; return { resolve(loaded: LoadedAuthConfig): LocalUserRegistry | undefined { if (loaded.value.mode !== "local") return undefined; const usersPath = join(dirname(loaded.sourcePath), loaded.value.local.usersFile); if (current?.usersPath === usersPath) return current.registry; - const registry = createLocalUserRegistry(usersPath); + const registry = createLocalUserRegistry(usersPath, options); current = { usersPath, registry }; return registry; }, diff --git a/backend/src/auth/session-store.ts b/backend/src/auth/session-store.ts index b5c338b6..6e5fa08a 100644 --- a/backend/src/auth/session-store.ts +++ b/backend/src/auth/session-store.ts @@ -1,23 +1,4 @@ import { createHash, hkdfSync, randomBytes } from "node:crypto"; -import { - accessSync, - closeSync, - constants, - fchmodSync, - fstatSync, - fsyncSync, - lstatSync, - linkSync, - openSync, - opendirSync, - readSync, - realpathSync, - renameSync, - unlinkSync, - writeSync, -} from "node:fs"; -import type { Stats } from "node:fs"; -import { dirname, isAbsolute, join, normalize } from "node:path"; import { z } from "zod"; import type { PrincipalContext } from "./principal.js"; import type { @@ -38,8 +19,6 @@ const TOKEN_PATTERN = /^[A-Za-z0-9_-]{43}$/; const DIGEST_FILENAME_PATTERN = /^[a-f0-9]{64}\.json$/; const CLAIM_FILENAME_PATTERN = /^[a-f0-9]{64}\.claim$/; const OIDC_SLOT_FILENAME_PATTERN = /^slot-(\d{2})\.json$/; -const PRIVATE_DIRECTORY_MODE = 0o700; -const PRIVATE_FILE_MODE = 0o600; const MAX_SESSION_RECORD_BYTES = 16 * 1024; const MAX_OIDC_STATE_RECORD_BYTES = 8 * 1024; const MAX_OIDC_SLOT_RECORD_BYTES = 512; @@ -48,9 +27,6 @@ const OIDC_STATE_TTL_MS = 10 * 60 * 1000; const OIDC_STATE_CAPACITY = 64; const MAX_OIDC_STORAGE_ENTRIES = OIDC_STATE_CAPACITY * 3; const MAX_SESSION_PRUNE_ENTRIES = 512; -// A separate scan ceiling bounds directory-walk CPU and the duplicate-detection set without -// reviving the former 512-record availability ceiling. -const MAX_SESSION_DIRECTORY_PAGE_SCAN_ENTRIES = 16_384; const TOUCH_INTERVAL_MS = 5 * 60 * 1000; const CSRF_CONTEXT = Buffer.from("thothii-csrf-v1", "utf8"); const EMPTY_HKDF_SALT = Buffer.alloc(0); @@ -140,43 +116,12 @@ export interface AuthSessionStore { /** Narrow test seams for the native tht-backed storage adaptors. */ export interface FileAuthSessionStoreOptions { windowsStorageBridge?: WindowsAuthStorageBridge; - /** Test seam for POSIX layout creation; production uses the bounded hidden tht bridge. */ - posixStorageBridge?: Pick; + /** Test seam; production uses the bounded hidden tht bridge for every POSIX record operation. */ + posixStorageBridge?: WindowsAuthStorageBridge; /** Test-only capacity seam; production always uses the fixed 64-state bound. */ oidcStateCapacity?: number; } -interface FileIdentity { - dev: number; - ino: number; - uid: number; - size: number; - mtimeMs: number; -} - -interface DirectoryIdentity { - dev: number; - ino: number; - uid: number; - mode?: number; -} - -interface DirectoryScanIdentity extends DirectoryIdentity { - mtimeMs: number; - ctimeMs: number; -} - -interface TrustedFile { - value: T; - identity: FileIdentity; -} - -interface StorageDirectories { - root: string; - sessions: string; - oidc: string; -} - interface SessionDirectoryPage { entries: string[]; more: boolean; @@ -282,23 +227,6 @@ const oidcStateInputSchema = z.strictObject({ browserTransactionTransport: z.enum(["https", "loopback_http"]), }); -function sameFileIdentity(left: FileIdentity, right: FileIdentity): boolean { - return left.dev === right.dev && left.ino === right.ino && left.uid === right.uid && left.size === right.size - && left.mtimeMs === right.mtimeMs; -} - -function sameDirectoryIdentity(left: DirectoryIdentity, right: DirectoryIdentity): boolean { - return left.dev === right.dev && left.ino === right.ino && left.uid === right.uid && left.mode === right.mode; -} - -function sameDirectoryScanIdentity(left: DirectoryScanIdentity, right: DirectoryScanIdentity): boolean { - return sameDirectoryIdentity(left, right) && left.mtimeMs === right.mtimeMs && left.ctimeMs === right.ctimeMs; -} - -function isNotFound(error: unknown): boolean { - return (error as NodeJS.ErrnoException | undefined)?.code === "ENOENT"; -} - function canonicalRawValue(value: string): boolean { if (typeof value !== "string" || !TOKEN_PATTERN.test(value)) return false; try { @@ -330,447 +258,6 @@ function oidcSlotIndex(filename: string): number | undefined { return Number.isInteger(index) && index >= 0 && index < OIDC_STATE_CAPACITY ? index : undefined; } -function assertFilename(filename: string): void { - if (!DIGEST_FILENAME_PATTERN.test(filename) && !CLAIM_FILENAME_PATTERN.test(filename) - && oidcSlotIndex(filename) === undefined) throw invalid(); -} - -function filePath(directory: string, filename: string): string { - assertFilename(filename); - const path = join(directory, filename); - if (dirname(path) !== directory) throw invalid(); - return path; -} - -function ownerId(): number { - const getEffectiveUserId = process.geteuid; - if (typeof getEffectiveUserId !== "function") throw invalid(); - const euid = getEffectiveUserId(); - if (!Number.isSafeInteger(euid) || euid < 0) throw invalid(); - return euid; -} - -function fileIdentity(info: Stats, expectedLinks = 1): FileIdentity { - if (!info.isFile() || info.isSymbolicLink() || info.nlink !== expectedLinks - || info.uid !== ownerId() || (info.mode & 0o7777) !== PRIVATE_FILE_MODE || info.size < 0) { - throw invalid(); - } - return { dev: info.dev, ino: info.ino, uid: info.uid, size: info.size, mtimeMs: info.mtimeMs }; -} - -function directoryIdentity(path: string): DirectoryIdentity { - const info = lstatSync(path) as Stats; - if (!info.isDirectory() || info.isSymbolicLink() || realpathSync(path) !== path) throw invalid(); - if (info.uid !== ownerId() || (info.mode & 0o7777) !== PRIVATE_DIRECTORY_MODE) throw invalid(); - return { - dev: info.dev, - ino: info.ino, - uid: info.uid, - mode: info.mode & 0o7777, - }; -} - -function directoryScanIdentity(path: string): DirectoryScanIdentity { - const identity = directoryIdentity(path); - const info = lstatSync(path) as Stats; - if (!info.isDirectory() || info.isSymbolicLink() || info.dev !== identity.dev || info.ino !== identity.ino - || info.uid !== identity.uid || (info.mode & 0o7777) !== identity.mode - || !Number.isFinite(info.mtimeMs) || !Number.isFinite(info.ctimeMs)) throw invalid(); - return { ...identity, mtimeMs: info.mtimeMs, ctimeMs: info.ctimeMs }; -} - -interface SessionRootAncestor { - path: string; - descriptor: number; - dev: number; - ino: number; - uid: number; - mode: number; -} - -function validateSessionRootSyntax(root: string): void { - if (process.platform === "win32" || typeof root !== "string" || root.length === 0 - || root.includes("\0") || /\p{Cc}/u.test(root) || !isAbsolute(root) || normalize(root) !== root) throw invalid(); -} - -function sameAncestor(ancestor: SessionRootAncestor, info: Stats): boolean { - return info.isDirectory() && !info.isSymbolicLink() && ancestor.dev === info.dev - && ancestor.ino === info.ino && ancestor.uid === info.uid && ancestor.mode === info.mode; -} - -function withSessionRootPreflight(root: string, use: (exists: boolean) => T): T { - validateSessionRootSyntax(root); - const ancestors: SessionRootAncestor[] = []; - try { - const components = root.split("/").filter((component) => component.length > 0); - let current = "/"; - for (let index = -1; index < components.length; index += 1) { - if (index >= 0) current = join(current, components[index]!); - let info: Stats; - try { - info = lstatSync(current) as Stats; - } catch (error) { - if (!isNotFound(error) || index !== components.length - 1) throw invalid(); - accessSync(dirname(current), constants.W_OK | constants.X_OK); - for (const ancestor of ancestors) { - const observed = lstatSync(ancestor.path) as Stats; - if (!sameAncestor(ancestor, observed) || realpathSync(ancestor.path) !== ancestor.path) throw invalid(); - } - const result = use(false); - for (const ancestor of ancestors) { - const observed = lstatSync(ancestor.path) as Stats; - if (!sameAncestor(ancestor, observed) || realpathSync(ancestor.path) !== ancestor.path) throw invalid(); - } - return result; - } - if (!info.isDirectory() || info.isSymbolicLink() || realpathSync(current) !== current) throw invalid(); - const descriptor = openSync(current, constants.O_RDONLY | (constants.O_DIRECTORY ?? 0) - | (constants.O_NOFOLLOW ?? 0) | (constants.O_NONBLOCK ?? 0)); - const opened = fstatSync(descriptor) as Stats; - if (!opened.isDirectory() || opened.dev !== info.dev || opened.ino !== info.ino - || opened.uid !== info.uid || opened.mode !== info.mode) { - closeSync(descriptor); - throw invalid(); - } - ancestors.push({ - path: current, - descriptor, - dev: info.dev, - ino: info.ino, - uid: info.uid, - mode: info.mode, - }); - } - directoryIdentity(root); - const result = use(true); - for (const ancestor of ancestors) { - const observed = lstatSync(ancestor.path) as Stats; - if (!sameAncestor(ancestor, observed) || realpathSync(ancestor.path) !== ancestor.path) throw invalid(); - } - return result; - } catch { - throw invalid(); - } finally { - for (const ancestor of ancestors.reverse()) { - try { closeSync(ancestor.descriptor); } catch { /* preflight has already failed closed */ } - } - } -} - -/** Side-effect-free POSIX validator shared by runtime storage and static diagnostics. */ -export function validateAuthSessionRoot(root: string): void { - try { - const exists = withSessionRootPreflight(root, (currentExists) => currentExists); - if (!exists) return; - for (const child of ["sessions", "oidc"]) { - const path = join(root, child); - if (dirname(path) !== root) throw invalid(); - withSessionRootPreflight(path, () => undefined); - } - } catch { - throw invalid(); - } -} - -function storageDirectories(root: string): StorageDirectories { - // Native Windows calls must dispatch to the tht DACL-capable bridge before reaching this - // POSIX-only helper. Keep this guard so an un-routed caller cannot fall back to chmod. - validateSessionRootSyntax(root); - validateAuthSessionRoot(root); - const sessions = join(root, "sessions"); - const oidc = join(root, "oidc"); - directoryIdentity(sessions); - directoryIdentity(oidc); - return { root, sessions, oidc }; -} - -function boundedDirectoryNames(directory: string, maximumEntries: number): string[] { - if (!Number.isInteger(maximumEntries) || maximumEntries < 1) throw invalid(); - let handle: ReturnType | undefined; - try { - const before = directoryIdentity(directory); - handle = opendirSync(directory); - const names: string[] = []; - while (names.length <= maximumEntries) { - const entry = handle.readSync(); - if (entry === null) { - handle.closeSync(); - handle = undefined; - if (!sameDirectoryIdentity(before, directoryIdentity(directory))) throw invalid(); - return names; - } - names.push(entry.name); - } - throw invalid(); - } catch { - throw invalid(); - } finally { - if (handle !== undefined) { - try { handle.closeSync(); } catch { /* the operation is already fail-closed */ } - } - } -} - -/** - * Return the next lexical page without retaining a complete directory listing. Every direct - * child is structurally validated during the bounded scan, so an unsafe entry cannot hide after - * a full page of ordinary sessions. - */ -function boundedSessionDirectoryPage( - directory: string, - after: string | undefined, - maximumEntries: number, -): SessionDirectoryPage { - if (!Number.isInteger(maximumEntries) || maximumEntries < 1 - || (after !== undefined && !DIGEST_FILENAME_PATTERN.test(after))) throw invalid(); - let handle: ReturnType | undefined; - try { - const before = directoryScanIdentity(directory); - handle = opendirSync(directory); - const seen = new Set(); - const selected: string[] = []; - let scanned = 0; - while (true) { - const entry = handle.readSync(); - if (entry === null) { - handle.closeSync(); - handle = undefined; - if (!sameDirectoryScanIdentity(before, directoryScanIdentity(directory))) throw invalid(); - selected.sort((left, right) => left < right ? -1 : left > right ? 1 : 0); - const more = selected.length > maximumEntries; - return { entries: more ? selected.slice(0, maximumEntries) : selected, more }; - } - scanned += 1; - if (scanned > MAX_SESSION_DIRECTORY_PAGE_SCAN_ENTRIES) throw invalid(); - const filename = entry.name; - if (!DIGEST_FILENAME_PATTERN.test(filename) || seen.has(filename)) throw invalid(); - seen.add(filename); - let info: Stats; - try { - info = lstatSync(filePath(directory, filename)) as Stats; - } catch (error) { - throw invalid(); - } - fileIdentity(info); - if (after !== undefined && filename <= after) continue; - appendBoundedSessionFilename(selected, filename, maximumEntries + 1); - } - } catch { - throw invalid(); - } finally { - if (handle !== undefined) { - try { handle.closeSync(); } catch { /* the operation is already fail-closed */ } - } - } -} - -function appendBoundedSessionFilename(names: string[], filename: string, maximumEntries: number): void { - if (names.length < maximumEntries) { - names.push(filename); - return; - } - let greatest = 0; - for (let index = 1; index < names.length; index += 1) { - if (names[index] > names[greatest]) greatest = index; - } - if (filename < names[greatest]) names[greatest] = filename; -} - -function openDirectory(directory: string): number | undefined { - if (process.platform === "win32") return undefined; - return openSync(directory, constants.O_RDONLY | (constants.O_DIRECTORY ?? 0) - | (constants.O_NOFOLLOW ?? 0) | (constants.O_NONBLOCK ?? 0)); -} - -function syncDirectory(directory: string): void { - if (process.platform === "win32") return; - let descriptor: number | undefined; - try { - descriptor = openDirectory(directory); - if (descriptor === undefined) throw invalid(); - fsyncSync(descriptor); - } catch { - throw invalid(); - } finally { - if (descriptor !== undefined) { - try { closeSync(descriptor); } catch { /* converted to a sanitized failure above */ } - } - } -} - -function writeFully(descriptor: number, contents: Buffer): void { - let offset = 0; - while (offset < contents.length) { - const written = writeSync(descriptor, contents, offset, contents.length - offset); - if (written <= 0) throw invalid(); - offset += written; - } -} - -function readTrusted( - directory: string, - filename: string, - maximumBytes: number, - parse: (source: string) => T, - expectedLinks = 1, -): TrustedFile | undefined { - const path = filePath(directory, filename); - let directoryDescriptor: number | undefined; - let descriptor: number | undefined; - try { - const beforeDirectory = directoryIdentity(directory); - const beforePath = lstatSync(path) as Stats; - const before = fileIdentity(beforePath, expectedLinks); - if (before.size > maximumBytes) throw invalid(); - - directoryDescriptor = openDirectory(directory); - const openedDirectory = directoryDescriptor === undefined - ? beforeDirectory - : directoryIdentityFromDescriptor(directoryDescriptor); - if (!sameDirectoryIdentity(beforeDirectory, openedDirectory)) throw invalid(); - descriptor = openSync(path, constants.O_RDONLY | (constants.O_NOFOLLOW ?? 0) | (constants.O_NONBLOCK ?? 0)); - const opened = fileIdentity(fstatSync(descriptor) as Stats, expectedLinks); - if (!sameFileIdentity(before, opened) || opened.size > maximumBytes) throw invalid(); - - const contents = Buffer.allocUnsafe(maximumBytes + 1); - let offset = 0; - while (offset < contents.length) { - const read = readSync(descriptor, contents, offset, contents.length - offset, null); - if (read === 0) break; - offset += read; - } - if (offset > maximumBytes) throw invalid(); - - const after = fileIdentity(fstatSync(descriptor) as Stats, expectedLinks); - const afterPath = fileIdentity(lstatSync(path) as Stats, expectedLinks); - const afterDirectory = directoryIdentity(directory); - const afterOpenedDirectory = directoryDescriptor === undefined - ? afterDirectory - : directoryIdentityFromDescriptor(directoryDescriptor); - if (!sameFileIdentity(opened, after) || !sameFileIdentity(after, afterPath) - || !sameDirectoryIdentity(beforeDirectory, afterDirectory) - || !sameDirectoryIdentity(openedDirectory, afterOpenedDirectory)) throw invalid(); - - const source = new TextDecoder("utf-8", { fatal: true }).decode(contents.subarray(0, offset)); - return { value: parse(source), identity: after }; - } catch (error) { - if (isNotFound(error)) return undefined; - throw invalid(); - } finally { - if (descriptor !== undefined) { - try { closeSync(descriptor); } catch { /* descriptor is no longer trusted */ } - } - if (directoryDescriptor !== undefined) { - try { closeSync(directoryDescriptor); } catch { /* descriptor is no longer trusted */ } - } - } -} - -function directoryIdentityFromDescriptor(descriptor: number): DirectoryIdentity { - const info = fstatSync(descriptor) as Stats; - if (!info.isDirectory() || info.isSymbolicLink()) throw invalid(); - if (info.uid !== ownerId() || (info.mode & 0o7777) !== PRIVATE_DIRECTORY_MODE) throw invalid(); - return { - dev: info.dev, - ino: info.ino, - uid: info.uid, - mode: info.mode & 0o7777, - }; -} - -function writeExclusive(directory: string, filename: string, contents: Buffer): boolean { - const path = filePath(directory, filename); - let descriptor: number | undefined; - try { - descriptor = openSync( - path, - constants.O_WRONLY | constants.O_CREAT | constants.O_EXCL | (constants.O_NOFOLLOW ?? 0), - PRIVATE_FILE_MODE, - ); - fchmodSync(descriptor, PRIVATE_FILE_MODE); - writeFully(descriptor, contents); - fsyncSync(descriptor); - const written = fileIdentity(fstatSync(descriptor) as Stats); - if (written.size !== contents.length) throw invalid(); - closeSync(descriptor); - descriptor = undefined; - syncDirectory(directory); - return true; - } catch (error) { - if (isNotFound(error)) throw invalid(); - if ((error as NodeJS.ErrnoException | undefined)?.code === "EEXIST") return false; - throw invalid(); - } finally { - if (descriptor !== undefined) { - try { closeSync(descriptor); } catch { /* best effort only */ } - try { unlinkSync(path); } catch { /* only our exclusive temporary record can remain */ } - } - } -} - -function replaceTrusted( - directory: string, - filename: string, - expected: FileIdentity, - contents: Buffer, -): void { - const path = filePath(directory, filename); - const temporary = join(directory, `.${filename}.${randomBytes(12).toString("hex")}.tmp`); - let descriptor: number | undefined; - try { - descriptor = openSync( - temporary, - constants.O_WRONLY | constants.O_CREAT | constants.O_EXCL | (constants.O_NOFOLLOW ?? 0), - PRIVATE_FILE_MODE, - ); - fchmodSync(descriptor, PRIVATE_FILE_MODE); - writeFully(descriptor, contents); - fsyncSync(descriptor); - const temporaryInfo = fileIdentity(fstatSync(descriptor) as Stats); - if (temporaryInfo.size !== contents.length) throw invalid(); - closeSync(descriptor); - descriptor = undefined; - - const current = fileIdentity(lstatSync(path) as Stats); - if (!sameFileIdentity(expected, current)) throw invalid(); - renameSync(temporary, path); - syncDirectory(directory); - } catch { - throw invalid(); - } finally { - if (descriptor !== undefined) { - try { closeSync(descriptor); } catch { /* the failing write remains untrusted */ } - } - try { unlinkSync(temporary); } catch { /* rename or no creation: nothing to remove */ } - } -} - -function removeTrusted( - directory: string, - filename: string, - expected?: FileIdentity, - expectedLinks = 1, -): boolean { - const path = filePath(directory, filename); - try { - const beforeDirectory = directoryIdentity(directory); - const current = fileIdentity(lstatSync(path) as Stats, expectedLinks); - if (expected && !sameFileIdentity(expected, current)) throw invalid(); - // Revalidate both names immediately before unlink. On POSIX unlink never follows a final - // symlink, and this closes the observable replacement window before that operation. - const finalDirectory = directoryIdentity(directory); - const final = fileIdentity(lstatSync(path) as Stats, expectedLinks); - if (!sameDirectoryIdentity(beforeDirectory, finalDirectory) || !sameFileIdentity(current, final) - || (expected !== undefined && !sameFileIdentity(expected, final))) throw invalid(); - unlinkSync(path); - syncDirectory(directory); - return true; - } catch (error) { - if (isNotFound(error)) return false; - throw invalid(); - } -} - function parseSessionRecord(source: string): AuthSessionRecord { try { return sessionRecordSchema.parse(JSON.parse(source)) as AuthSessionRecord; @@ -806,107 +293,10 @@ function parseWindowsRecord(contents: Buffer, maximumBytes: number, parse: (s } } -interface OidcStateClaim { - state: TrustedFile; - claimIdentity: FileIdentity; -} - interface StoredOidcSlot { filename: string; index: number; record: OidcSlotRecord; - identity?: FileIdentity; -} - -function inspectOidcStateClaim( - directory: string, - filename: string, -): OidcStateClaim | "orphan" | undefined { - const claimedFilename = claimFilename(filename); - let claimInfo: Stats; - try { - claimInfo = lstatSync(filePath(directory, claimedFilename)) as Stats; - } catch (error) { - if (isNotFound(error)) return undefined; - throw invalid(); - } - let sourceInfo: Stats; - try { - sourceInfo = lstatSync(filePath(directory, filename)) as Stats; - } catch (error) { - if (!isNotFound(error)) throw invalid(); - // A consumer may have just unlinked the source and not yet removed its claim. Do not - // remove that orphan here: doing so could make the winning consumer fail closed after it - // has read the state. prune() removes abandoned orphan claims after the state lifetime. - fileIdentity(claimInfo, 1); - return "orphan"; - } - const sourceIdentity = fileIdentity(sourceInfo, 2); - const claimIdentity = fileIdentity(claimInfo, 2); - if (!sameFileIdentity(sourceIdentity, claimIdentity)) throw invalid(); - const state = readTrusted(directory, filename, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord, 2); - if (!state || !sameFileIdentity(sourceIdentity, state.identity)) throw invalid(); - return { state, claimIdentity }; -} - -function oidcClaimExists(directory: string, filename: string): boolean { - const claimedFilename = claimFilename(filename); - try { - const info = lstatSync(filePath(directory, claimedFilename)) as Stats; - if (info.nlink !== 1 && info.nlink !== 2) throw invalid(); - fileIdentity(info, info.nlink); - return true; - } catch (error) { - if (isNotFound(error)) return false; - throw invalid(); - } -} - -/** - * Claim a state by creating a hard link with a deterministic digest-only name. link(2) / NTFS - * CreateHardLink fails if another process has already installed that name, unlike rename which - * can overwrite the previous claimant. A crash leaves the pair unavailable until expiry/prune. - */ -function claimOidcState(directory: string, filename: string): OidcStateClaim | undefined { - // A competing process may be partway through consumption. It has already won and must be - // the only process allowed to deserialize the record; this process simply treats it as used. - if (oidcClaimExists(directory, filename)) return undefined; - - const statePath = filePath(directory, filename); - const claimedFilename = claimFilename(filename); - const claimedPath = filePath(directory, claimedFilename); - let before: FileIdentity; - try { - before = fileIdentity(lstatSync(statePath) as Stats); - } catch (error) { - if (isNotFound(error)) return undefined; - // A winner can install its hard-link after the first claim check and before this state - // identity check. That process owns the state; this contender must fail closed as used. - if (oidcClaimExists(directory, filename)) return undefined; - throw invalid(); - } - try { - linkSync(statePath, claimedPath); - } catch (error) { - if (isNotFound(error)) return undefined; - if ((error as NodeJS.ErrnoException | undefined)?.code === "EEXIST") return undefined; - throw invalid(); - } - try { - const sourceIdentity = fileIdentity(lstatSync(statePath) as Stats, 2); - const claimIdentity = fileIdentity(lstatSync(claimedPath) as Stats, 2); - if (!sameFileIdentity(before, sourceIdentity) || !sameFileIdentity(sourceIdentity, claimIdentity)) throw invalid(); - const state = readTrusted(directory, filename, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord, 2); - if (!state || !sameFileIdentity(state.identity, sourceIdentity)) throw invalid(); - return { state, claimIdentity }; - } catch { - throw invalid(); - } -} - -function removeClaimedOidcState(directory: string, filename: string, claim: OidcStateClaim): void { - if (!removeTrusted(directory, filename, claim.state.identity, 2)) throw invalid(); - if (!removeTrusted(directory, claimFilename(filename), claim.claimIdentity)) throw invalid(); } function serialize(record: AuthSessionRecord | OidcStateRecord | OidcSlotRecord, maximumBytes: number): Buffer { @@ -1005,40 +395,35 @@ export function createFileAuthSessionStore( if (!Number.isInteger(oidcStateCapacity) || oidcStateCapacity < 1 || oidcStateCapacity > OIDC_STATE_CAPACITY) { throw invalid(); } - const windowsStorage = process.platform === "win32" + const rawStorage = process.platform === "win32" ? options.windowsStorageBridge ?? createWindowsAuthStorageBridge() - : undefined; - const posixStorage = process.platform === "win32" - ? undefined : options.posixStorageBridge ?? createPosixAuthStorageBridge(); + // A malformed or failed helper must be indistinguishable from any other storage failure to + // callers. This also keeps narrow test seams from accidentally exposing transport details. + const storage: WindowsAuthStorageBridge = { + validateRoot: async (value) => { try { await rawStorage.validateRoot(value); } catch { throw invalid(); } }, + ensureLayout: async (value) => { try { await rawStorage.ensureLayout(value); } catch { throw invalid(); } }, + readAuthConfig: (value) => { try { return rawStorage.readAuthConfig(value); } catch { throw invalid(); } }, + readLocalUsers: async (value) => { try { return await rawStorage.readLocalUsers(value); } catch { throw invalid(); } }, + create: async (...args) => { try { return await rawStorage.create(...args); } catch { throw invalid(); } }, + read: async (...args) => { try { return await rawStorage.read(...args); } catch { throw invalid(); } }, + replace: async (...args) => { try { await rawStorage.replace(...args); } catch { throw invalid(); } }, + remove: async (...args) => { try { return await rawStorage.remove(...args); } catch { throw invalid(); } }, + list: async (...args) => { try { return await rawStorage.list(...args); } catch { throw invalid(); } }, + listPage: async (...args) => { try { return await rawStorage.listPage(...args); } catch { throw invalid(); } }, + claimConsume: async (...args) => { try { return await rawStorage.claimConsume(...args); } catch { throw invalid(); } }, + readClaim: async (...args) => { try { return await rawStorage.readClaim(...args); } catch { throw invalid(); } }, + removeClaim: async (...args) => { try { return await rawStorage.removeClaim(...args); } catch { throw invalid(); } }, + }; let sessionPruneCursor: string | undefined; - function requiredWindowsStorage(): WindowsAuthStorageBridge { - if (windowsStorage === undefined) throw invalid(); - return windowsStorage; - } - - async function posixStorageDirectories(): Promise { - if (process.platform === "win32" || posixStorage === undefined) throw invalid(); - try { - return storageDirectories(root); - } catch { - // A missing safe layout is the only case the helper can repair. Unsafe layouts are - // rejected by the same native primitive without path-based fallback in this process. - } - try { - await posixStorage.ensureLayout(root); - return storageDirectories(root); - } catch { - throw invalid(); - } + function requiredStorage(): WindowsAuthStorageBridge { + if (!storage) throw invalid(); + return storage; } async function ordinarySessionPage(after: string | undefined): Promise { - if (process.platform !== "win32") { - return boundedSessionDirectoryPage((await posixStorageDirectories()).sessions, after, MAX_SESSION_PRUNE_ENTRIES); - } - const page = await requiredWindowsStorage().listPage(root, "sessions", after, MAX_SESSION_PRUNE_ENTRIES); + const page = await requiredStorage().listPage(root, "sessions", after, MAX_SESSION_PRUNE_ENTRIES); if (!page || !Array.isArray(page.entries) || typeof page.more !== "boolean") throw invalid(); return { entries: page.entries.map((entry) => entry.name), more: page.more }; } @@ -1064,32 +449,19 @@ export function createFileAuthSessionStore( const page = await ordinarySessionPage(after); const next = nextSessionPruneCursor(page, after); let removed = 0; - if (process.platform === "win32") { - const bridge = requiredWindowsStorage(); - for (const filename of page.entries) { - const contents = await bridge.read(root, "sessions", filename); - if (!contents) continue; - const record = parseWindowsRecord(contents, MAX_SESSION_RECORD_BYTES, parseSessionRecord); - if (sessionExpired(record, nowMs) && await bridge.remove(root, "sessions", filename)) removed += 1; - } - } else { - const directory = (await posixStorageDirectories()).sessions; - for (const filename of page.entries) { - await withLock(lockKey(root, "sessions", filename), async () => { - const trusted = readTrusted(directory, filename, MAX_SESSION_RECORD_BYTES, parseSessionRecord); - if (trusted && sessionExpired(trusted.value, nowMs) - && removeTrusted(directory, filename, trusted.identity)) removed += 1; - }); - } + const bridge = requiredStorage(); + for (const filename of page.entries) { + const contents = await bridge.read(root, "sessions", filename); + if (!contents) continue; + const record = parseWindowsRecord(contents, MAX_SESSION_RECORD_BYTES, parseSessionRecord); + if (sessionExpired(record, nowMs) && await bridge.remove(root, "sessions", filename)) removed += 1; } sessionPruneCursor = next; return removed; } async function oidcStorageEntries(): Promise { - const entries = process.platform === "win32" - ? (await requiredWindowsStorage().list(root, "oidc", MAX_OIDC_STORAGE_ENTRIES)).map((entry) => entry.name) - : boundedDirectoryNames((await posixStorageDirectories()).oidc, MAX_OIDC_STORAGE_ENTRIES); + const entries = (await requiredStorage().list(root, "oidc", MAX_OIDC_STORAGE_ENTRIES)).map((entry) => entry.name); if (entries.length > MAX_OIDC_STORAGE_ENTRIES) throw invalid(); if (entries.some((entry) => !DIGEST_FILENAME_PATTERN.test(entry) && !CLAIM_FILENAME_PATTERN.test(entry) && oidcSlotIndex(entry) === undefined)) throw invalid(); @@ -1103,51 +475,22 @@ export function createFileAuthSessionStore( for (const filename of entries) { const index = oidcSlotIndex(filename); if (index === undefined) continue; - if (process.platform === "win32") { - const contents = await requiredWindowsStorage().read(root, "oidc", filename); - if (!contents) continue; - const record = parseWindowsRecord(contents, MAX_OIDC_SLOT_RECORD_BYTES, parseOidcSlotRecord); - if (stateFilenames.has(record.stateFilename)) throw invalid(); - stateFilenames.add(record.stateFilename); - slots.push({ filename, index, record }); - continue; - } - let trusted: TrustedFile | undefined; - let lastError: unknown; - for (const retryDelayMs of [0, 1, 2, 4, 8, 16, 32]) { - if (retryDelayMs > 0) await new Promise((resolve) => setTimeout(resolve, retryDelayMs)); - try { - trusted = readTrusted( - (await posixStorageDirectories()).oidc, - filename, - MAX_OIDC_SLOT_RECORD_BYTES, - parseOidcSlotRecord, - ); - lastError = undefined; - break; - } catch (error) { - lastError = error; - } - } - if (lastError !== undefined) throw invalid(); - if (!trusted) continue; - if (stateFilenames.has(trusted.value.stateFilename)) throw invalid(); - stateFilenames.add(trusted.value.stateFilename); - slots.push({ filename, index, record: trusted.value, identity: trusted.identity }); + const contents = await requiredStorage().read(root, "oidc", filename); + if (!contents) continue; + const record = parseWindowsRecord(contents, MAX_OIDC_SLOT_RECORD_BYTES, parseOidcSlotRecord); + if (stateFilenames.has(record.stateFilename)) throw invalid(); + stateFilenames.add(record.stateFilename); + slots.push({ filename, index, record }); } return slots; } async function removeOidcSlot(slot: StoredOidcSlot): Promise { - if (process.platform === "win32") { - const current = await requiredWindowsStorage().read(root, "oidc", slot.filename); - if (!current) return; - const record = parseWindowsRecord(current, MAX_OIDC_SLOT_RECORD_BYTES, parseOidcSlotRecord); - if (record.stateFilename !== slot.record.stateFilename || record.expiresAt !== slot.record.expiresAt - || !await requiredWindowsStorage().remove(root, "oidc", slot.filename)) throw invalid(); - return; - } - if (!slot.identity || !removeTrusted((await posixStorageDirectories()).oidc, slot.filename, slot.identity)) throw invalid(); + const current = await requiredStorage().read(root, "oidc", slot.filename); + if (!current) return; + const record = parseWindowsRecord(current, MAX_OIDC_SLOT_RECORD_BYTES, parseOidcSlotRecord); + if (record.stateFilename !== slot.record.stateFilename || record.expiresAt !== slot.record.expiresAt + || !await requiredStorage().remove(root, "oidc", slot.filename)) throw invalid(); } async function releaseOidcSlot(index: number | undefined, stateFilename: string): Promise { @@ -1177,9 +520,7 @@ export function createFileAuthSessionStore( for (let index = 0; index < availableSlotCount; index += 1) { if (occupied.has(index)) continue; const filename = oidcSlotFilename(index); - const created = process.platform === "win32" - ? await requiredWindowsStorage().create(root, "oidc", filename, contents) - : writeExclusive((await posixStorageDirectories()).oidc, filename, contents); + const created = await requiredStorage().create(root, "oidc", filename, contents); if (created) return index; } throw new OidcStateCapacityError(); @@ -1213,21 +554,11 @@ export function createFileAuthSessionStore( absoluteExpiresAt: isoAt(absoluteExpiresMs), }; const contents = serialize(record, MAX_SESSION_RECORD_BYTES); - if (process.platform === "win32") { - const bridge = requiredWindowsStorage(); - for (let attempt = 0; attempt < 8; attempt += 1) { - const token = randomBytes(TOKEN_BYTES).toString("base64url"); - if (await bridge.create(root, "sessions", digestFilename(token), contents)) { - return { token, csrfToken: deriveCsrfToken(token), record }; - } - } - throw invalid(); - } - const directories = await posixStorageDirectories(); + const bridge = requiredStorage(); for (let attempt = 0; attempt < 8; attempt += 1) { const token = randomBytes(TOKEN_BYTES).toString("base64url"); const filename = digestFilename(token); - if (writeExclusive(directories.sessions, filename, contents)) { + if (await bridge.create(root, "sessions", filename, contents)) { return { token, csrfToken: deriveCsrfToken(token), record }; } } @@ -1243,40 +574,22 @@ export function createFileAuthSessionStore( const nowMs = dateMilliseconds(now); const filename = digestFilename(token); return withLock(lockKey(root, "sessions", filename), async () => { - if (process.platform === "win32") { - const bridge = requiredWindowsStorage(); - const contents = await bridge.read(root, "sessions", filename); - if (!contents) return undefined; - const record = parseWindowsRecord(contents, MAX_SESSION_RECORD_BYTES, parseSessionRecord); - if (sessionExpired(record, nowMs)) { - await bridge.remove(root, "sessions", filename); - return undefined; - } - try { - if (await recordIsCurrent(record, requestValidity)) return record; - } catch (error) { - if (error instanceof AuthSessionOperationalError) throw error; - await bridge.remove(root, "sessions", filename); - throw invalid(); - } + const bridge = requiredStorage(); + const contents = await bridge.read(root, "sessions", filename); + if (!contents) return undefined; + const record = parseWindowsRecord(contents, MAX_SESSION_RECORD_BYTES, parseSessionRecord); + if (sessionExpired(record, nowMs)) { await bridge.remove(root, "sessions", filename); return undefined; } - const directories = await posixStorageDirectories(); - const trusted = readTrusted(directories.sessions, filename, MAX_SESSION_RECORD_BYTES, parseSessionRecord); - if (!trusted) return undefined; - if (sessionExpired(trusted.value, nowMs)) { - removeTrusted(directories.sessions, filename, trusted.identity); - return undefined; - } try { - if (await recordIsCurrent(trusted.value, requestValidity)) return trusted.value; + if (await recordIsCurrent(record, requestValidity)) return record; } catch (error) { if (error instanceof AuthSessionOperationalError) throw error; - removeTrusted(directories.sessions, filename, trusted.identity); + await bridge.remove(root, "sessions", filename); throw invalid(); } - removeTrusted(directories.sessions, filename, trusted.identity); + await bridge.remove(root, "sessions", filename); return undefined; }); } @@ -1286,49 +599,24 @@ export function createFileAuthSessionStore( const nowMs = dateMilliseconds(now); const filename = digestFilename(token); await withLock(lockKey(root, "sessions", filename), async () => { - if (process.platform === "win32") { - const bridge = requiredWindowsStorage(); - const contents = await bridge.read(root, "sessions", filename); - if (!contents) return; - const record = parseWindowsRecord(contents, MAX_SESSION_RECORD_BYTES, parseSessionRecord); - if (sessionExpired(record, nowMs)) { - await bridge.remove(root, "sessions", filename); - return; - } - const lastSeenMs = Date.parse(record.lastSeenAt); - if (nowMs <= lastSeenMs || nowMs - lastSeenMs < TOUCH_INTERVAL_MS) return; - const idleWindowMs = Date.parse(record.idleExpiresAt) - lastSeenMs; - if (idleWindowMs <= 0 || idleWindowMs > MAX_TTL_MS) throw invalid(); - const touched: AuthSessionRecord = { - ...record, - lastSeenAt: isoAt(nowMs), - idleExpiresAt: isoAt(Math.min(nowMs + idleWindowMs, Date.parse(record.absoluteExpiresAt))), - }; - await bridge.replace(root, "sessions", filename, serialize(touched, MAX_SESSION_RECORD_BYTES)); - return; + const bridge = requiredStorage(); + const contents = await bridge.read(root, "sessions", filename); + if (!contents) return; + const record = parseWindowsRecord(contents, MAX_SESSION_RECORD_BYTES, parseSessionRecord); + if (sessionExpired(record, nowMs)) { + await bridge.remove(root, "sessions", filename); + return undefined; } - const directories = await posixStorageDirectories(); - const trusted = readTrusted(directories.sessions, filename, MAX_SESSION_RECORD_BYTES, parseSessionRecord); - if (!trusted) return; - if (sessionExpired(trusted.value, nowMs)) { - removeTrusted(directories.sessions, filename, trusted.identity); - return; - } - const lastSeenMs = Date.parse(trusted.value.lastSeenAt); + const lastSeenMs = Date.parse(record.lastSeenAt); if (nowMs <= lastSeenMs || nowMs - lastSeenMs < TOUCH_INTERVAL_MS) return; - const idleWindowMs = Date.parse(trusted.value.idleExpiresAt) - lastSeenMs; + const idleWindowMs = Date.parse(record.idleExpiresAt) - lastSeenMs; if (idleWindowMs <= 0 || idleWindowMs > MAX_TTL_MS) throw invalid(); const touched: AuthSessionRecord = { - ...trusted.value, + ...record, lastSeenAt: isoAt(nowMs), - idleExpiresAt: isoAt(Math.min(nowMs + idleWindowMs, Date.parse(trusted.value.absoluteExpiresAt))), + idleExpiresAt: isoAt(Math.min(nowMs + idleWindowMs, Date.parse(record.absoluteExpiresAt))), }; - replaceTrusted( - directories.sessions, - filename, - trusted.identity, - serialize(touched, MAX_SESSION_RECORD_BYTES), - ); + await bridge.replace(root, "sessions", filename, serialize(touched, MAX_SESSION_RECORD_BYTES)); }); } @@ -1336,126 +624,61 @@ export function createFileAuthSessionStore( if (!canonicalRawValue(token)) return; const filename = digestFilename(token); await withLock(lockKey(root, "sessions", filename), async () => { - if (process.platform === "win32") { - await requiredWindowsStorage().remove(root, "sessions", filename); - return; - } - const directories = await posixStorageDirectories(); - removeTrusted(directories.sessions, filename); + await requiredStorage().remove(root, "sessions", filename); }); } async function pruneOidcStates(nowMs: number): Promise { - if (process.platform === "win32") { - const bridge = requiredWindowsStorage(); - const oidcEntries = await bridge.list(root, "oidc", MAX_OIDC_STORAGE_ENTRIES); - if (oidcEntries.length > MAX_OIDC_STORAGE_ENTRIES) throw invalid(); - const stateNames = new Set(oidcEntries - .filter((entry) => DIGEST_FILENAME_PATTERN.test(entry.name)) - .map((entry) => entry.name)); - const claimEntries = new Map(oidcEntries - .filter((entry) => CLAIM_FILENAME_PATTERN.test(entry.name)) - .map((entry) => [entry.name, entry])); - const slotEntries = oidcEntries.filter((entry) => oidcSlotIndex(entry.name) !== undefined); - if (stateNames.size + claimEntries.size + slotEntries.length !== oidcEntries.length) throw invalid(); - let removed = 0; - for (const filename of stateNames) { - const claim = claimFilename(filename); - const contents = claimEntries.has(claim) - ? await bridge.readClaim(root, filename) - : await bridge.read(root, "oidc", filename); - if (!contents) continue; - const record = parseWindowsRecord(contents, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord); - if (oidcStateExpired(record, nowMs)) { - const didRemove = claimEntries.has(claim) - ? await bridge.removeClaim(root, filename) - : await bridge.remove(root, "oidc", filename); - if (didRemove) removed += 1; - } - } - for (const [claim, entry] of claimEntries) { - const filename = `${claim.slice(0, -".claim".length)}.json`; - if (stateNames.has(filename)) continue; - if (nowMs >= entry.modifiedUnixMs + OIDC_STATE_TTL_MS - && await bridge.remove(root, "oidc", claim)) removed += 1; - } - for (const entry of slotEntries) { - const contents = await bridge.read(root, "oidc", entry.name); - if (!contents) continue; - const slot = parseWindowsRecord(contents, MAX_OIDC_SLOT_RECORD_BYTES, parseOidcSlotRecord); - if (nowMs < Date.parse(slot.expiresAt)) continue; - const claim = claimFilename(slot.stateFilename); - const stateContents = claimEntries.has(claim) - ? await bridge.readClaim(root, slot.stateFilename) - : await bridge.read(root, "oidc", slot.stateFilename); - if (stateContents) { - const state = parseWindowsRecord(stateContents, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord); - if (!oidcStateExpired(state, nowMs)) throw invalid(); - const didRemove = claimEntries.has(claim) - ? await bridge.removeClaim(root, slot.stateFilename) - : await bridge.remove(root, "oidc", slot.stateFilename); - if (didRemove) removed += 1; - } - if (!await bridge.remove(root, "oidc", entry.name)) throw invalid(); - } - return removed; - } - - const directory = (await posixStorageDirectories()).oidc; - const oidcEntries = boundedDirectoryNames(directory, MAX_OIDC_STORAGE_ENTRIES); - const stateFilenames = new Set(oidcEntries.filter((entry) => DIGEST_FILENAME_PATTERN.test(entry))); + const bridge = requiredStorage(); + const oidcEntries = await bridge.list(root, "oidc", MAX_OIDC_STORAGE_ENTRIES); + if (oidcEntries.length > MAX_OIDC_STORAGE_ENTRIES) throw invalid(); + const stateNames = new Set(oidcEntries + .filter((entry) => DIGEST_FILENAME_PATTERN.test(entry.name)) + .map((entry) => entry.name)); + const claimEntries = new Map(oidcEntries + .filter((entry) => CLAIM_FILENAME_PATTERN.test(entry.name)) + .map((entry) => [entry.name, entry])); + const slotEntries = oidcEntries.filter((entry) => oidcSlotIndex(entry.name) !== undefined); + if (stateNames.size + claimEntries.size + slotEntries.length !== oidcEntries.length) throw invalid(); let removed = 0; - for (const filename of stateFilenames) { - await withLock(lockKey(root, "oidc", filename), async () => { - const claim = inspectOidcStateClaim(directory, filename); - if (claim === "orphan") return; - if (claim) { - if (oidcStateExpired(claim.state.value, nowMs)) { - removeClaimedOidcState(directory, filename, claim); - removed += 1; - } - return; - } - const trusted = readTrusted(directory, filename, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord); - if (trusted && oidcStateExpired(trusted.value, nowMs) - && removeTrusted(directory, filename, trusted.identity)) removed += 1; - }); + for (const filename of stateNames) { + const claim = claimFilename(filename); + const contents = claimEntries.has(claim) + ? await bridge.readClaim(root, filename) + : await bridge.read(root, "oidc", filename); + if (!contents) continue; + const record = parseWindowsRecord(contents, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord); + if (oidcStateExpired(record, nowMs)) { + const didRemove = claimEntries.has(claim) + ? await bridge.removeClaim(root, filename) + : await bridge.remove(root, "oidc", filename); + if (didRemove) removed += 1; + } } - for (const claimedFilename of oidcEntries) { - if (!CLAIM_FILENAME_PATTERN.test(claimedFilename)) continue; - const filename = `${claimedFilename.slice(0, -".claim".length)}.json`; - if (stateFilenames.has(filename)) continue; - await withLock(lockKey(root, "oidc", filename), async () => { - const claim = inspectOidcStateClaim(directory, filename); - if (claim !== "orphan") return; - const claimIdentity = fileIdentity(lstatSync(filePath(directory, claimedFilename)) as Stats); - if (nowMs >= claimIdentity.mtimeMs + OIDC_STATE_TTL_MS - && removeTrusted(directory, claimedFilename, claimIdentity)) removed += 1; - }); + for (const [claim, entry] of claimEntries) { + const filename = `${claim.slice(0, -".claim".length)}.json`; + if (stateNames.has(filename)) continue; + if (nowMs >= entry.modifiedUnixMs + OIDC_STATE_TTL_MS + && await bridge.remove(root, "oidc", claim)) removed += 1; } - const slots = await storedOidcSlots(oidcEntries.filter((entry) => oidcSlotIndex(entry) !== undefined)); - for (const slot of slots) { - if (nowMs < Date.parse(slot.record.expiresAt)) continue; - await withLock(lockKey(root, "oidc", slot.record.stateFilename), async () => { - const claim = inspectOidcStateClaim(directory, slot.record.stateFilename); - if (claim && claim !== "orphan") { - if (!oidcStateExpired(claim.state.value, nowMs)) throw invalid(); - removeClaimedOidcState(directory, slot.record.stateFilename, claim); - removed += 1; - } else if (!claim) { - const trusted = readTrusted( - directory, - slot.record.stateFilename, - MAX_OIDC_STATE_RECORD_BYTES, - parseOidcStateRecord, - ); - if (trusted) { - if (!oidcStateExpired(trusted.value, nowMs)) throw invalid(); - if (removeTrusted(directory, slot.record.stateFilename, trusted.identity)) removed += 1; - } - } - await removeOidcSlot(slot); - }); + for (const entry of slotEntries) { + const contents = await bridge.read(root, "oidc", entry.name); + if (!contents) continue; + const slot = parseWindowsRecord(contents, MAX_OIDC_SLOT_RECORD_BYTES, parseOidcSlotRecord); + if (nowMs < Date.parse(slot.expiresAt)) continue; + const claim = claimFilename(slot.stateFilename); + const stateContents = claimEntries.has(claim) + ? await bridge.readClaim(root, slot.stateFilename) + : await bridge.read(root, "oidc", slot.stateFilename); + if (stateContents) { + const state = parseWindowsRecord(stateContents, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord); + if (!oidcStateExpired(state, nowMs)) throw invalid(); + const didRemove = claimEntries.has(claim) + ? await bridge.removeClaim(root, slot.stateFilename) + : await bridge.remove(root, "oidc", slot.stateFilename); + if (didRemove) removed += 1; + } + if (!await bridge.remove(root, "oidc", entry.name)) throw invalid(); } return removed; } @@ -1490,9 +713,7 @@ export function createFileAuthSessionStore( expiresAt: isoAt(expiresMs), }; const contents = serialize(record, MAX_OIDC_STATE_RECORD_BYTES); - const created = process.platform === "win32" - ? await requiredWindowsStorage().create(root, "oidc", filename, contents) - : writeExclusive((await posixStorageDirectories()).oidc, filename, contents); + const created = await requiredStorage().create(root, "oidc", filename, contents); if (created) return { state, record }; await releaseOidcSlot(capacitySlot, filename); } @@ -1505,26 +726,11 @@ export function createFileAuthSessionStore( const nowMs = dateMilliseconds(now); const filename = digestFilename(state); return withLock(lockKey(root, "oidc", filename), async () => { - if (process.platform === "win32") { - const contents = await requiredWindowsStorage().claimConsume(root, filename); - if (!contents) return undefined; - const record = parseWindowsRecord(contents, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord); - await releaseOidcSlot(record.capacitySlot, filename); - return oidcStateExpired(record, nowMs) ? undefined : record; - } - const directories = await posixStorageDirectories(); - const claim = claimOidcState(directories.oidc, filename); - // An installed claim belongs to another process/store instance. Only the process which - // created the hard link is allowed to receive the record. - if (!claim) return undefined; - if (oidcStateExpired(claim.state.value, nowMs)) { - removeClaimedOidcState(directories.oidc, filename, claim); - await releaseOidcSlot(claim.state.value.capacitySlot, filename); - return undefined; - } - removeClaimedOidcState(directories.oidc, filename, claim); - await releaseOidcSlot(claim.state.value.capacitySlot, filename); - return claim.state.value; + const contents = await requiredStorage().claimConsume(root, filename); + if (!contents) return undefined; + const record = parseWindowsRecord(contents, MAX_OIDC_STATE_RECORD_BYTES, parseOidcStateRecord); + await releaseOidcSlot(record.capacitySlot, filename); + return oidcStateExpired(record, nowMs) ? undefined : record; }); } diff --git a/backend/src/auth/windows-auth-storage.ts b/backend/src/auth/windows-auth-storage.ts index bfbb3487..4de5f24e 100644 --- a/backend/src/auth/windows-auth-storage.ts +++ b/backend/src/auth/windows-auth-storage.ts @@ -14,6 +14,8 @@ const MAX_OIDC_BYTES = 8 * 1024; const DEFAULT_MAX_ENTRIES = 256; const MAX_ENTRIES = 512; const TIMEOUT_MS = 5_000; +const TERMINATION_GRACE_MS = 100; +const FINAL_SETTLEMENT_MS = 750; const DIGEST_FILENAME = /^[a-f0-9]{64}\.json$/; const CLAIM_FILENAME = /^[a-f0-9]{64}\.claim$/; const OIDC_SLOT_FILENAME = /^slot-(?:[0-5][0-9]|6[0-3])\.json$/; @@ -38,6 +40,7 @@ export interface WindowsAuthStorageBridge { validateRoot(root: string): Promise; ensureLayout(root: string): Promise; readAuthConfig(path: string): Buffer; + readLocalUsers(path: string): Promise; create(root: string, directory: WindowsAuthStorageDirectory, filename: string, contents: Buffer): Promise; read(root: string, directory: WindowsAuthStorageDirectory, filename: string): Promise; replace(root: string, directory: WindowsAuthStorageDirectory, filename: string, contents: Buffer): Promise; @@ -77,9 +80,11 @@ interface WindowsAuthStorageChild { readonly stdout: Readable | null; readonly stderr: Readable | null; kill(signal?: NodeJS.Signals | number): boolean; + unref?(): void; on(event: "error", listener: (error: Error) => void): this; once(event: "error", listener: (error: Error) => void): this; once(event: "close", listener: (code: number | null, signal: NodeJS.Signals | null) => void): this; + removeListener?(event: "error" | "close", listener: (...args: any[]) => void): this; } type WindowsAuthStorageSpawn = ( @@ -99,6 +104,12 @@ export interface WindowsAuthStorageBridgeOptions { spawnChild?: WindowsAuthStorageSpawn; /** Test-only input scheduling seam for real child-process lifecycle tests. */ beforeInputForTest?: () => Promise; + /** Test-only bounded lifecycle timings. Production always uses the fixed deadlines below. */ + deadlinesForTest?: { + timeoutMs?: number; + terminationGraceMs?: number; + finalSettlementMs?: number; + }; } type AuthStoragePathStyle = "posix" | "windows"; @@ -124,7 +135,7 @@ type BridgeResponse = z.infer; interface BridgeRequest { version: typeof PROTOCOL_VERSION; - operation: "validate-root" | "ensure-layout" | "read-auth-config" | "create" | "read" | "replace" | "remove" | "list" | "claim-consume" | "read-claim" | "remove-claim"; + operation: "validate-root" | "ensure-layout" | "read-auth-config" | "read-local-users" | "create" | "read" | "replace" | "remove" | "list" | "claim-consume" | "read-claim" | "remove-claim"; root: string; directory?: WindowsAuthStorageDirectory; filename?: string; @@ -190,7 +201,7 @@ function encodedRequest(request: BridgeRequest, pathStyle: AuthStoragePathStyle) if (request.operation === "validate-root" || request.operation === "ensure-layout") { if (request.directory !== undefined || request.filename !== undefined || request.contentBase64 !== undefined || request.maximumEntries !== undefined || request.afterName !== undefined || request.continuation !== undefined) throw invalid(); - } else if (request.operation === "read-auth-config") { + } else if (request.operation === "read-auth-config" || request.operation === "read-local-users") { if (request.directory !== undefined || request.contentBase64 !== undefined || request.maximumEntries !== undefined || request.afterName !== undefined || request.continuation !== undefined || request.filename === undefined || !AUTH_CONFIG_FILENAME.test(request.filename)) throw invalid(); @@ -261,52 +272,127 @@ async function invokeTht( invocation: WindowsAuthStorageInvocation, spawnChild: WindowsAuthStorageSpawn = spawnTht, beforeInputForTest?: () => Promise, + terminationGraceMs = TERMINATION_GRACE_MS, + finalSettlementMs = FINAL_SETTLEMENT_MS, ): Promise { return new Promise((resolve, reject) => { let settled = false; - let closeSeen = false; let aborted = false; - let closeCode: number | null = null; - let closeSignal: NodeJS.Signals | null = null; let timeout: NodeJS.Timeout | undefined; + let terminationTimer: NodeJS.Timeout | undefined; + let finalSettlementTimer: NodeJS.Timeout | undefined; + let lateErrorReleaseTimer: NodeJS.Timeout | undefined; const stdout: Buffer[] = []; const stderr: Buffer[] = []; let stdoutBytes = 0; let stderrBytes = 0; let child: WindowsAuthStorageChild | undefined; - const settle = (callback: () => void): void => { - if (settled) return; - settled = true; - if (timeout !== undefined) clearTimeout(timeout); - callback(); + let stdin: Writable | undefined; + let stdoutStream: Readable | undefined; + let stderrStream: Readable | undefined; + const swallowChildError = (): void => undefined; + const swallowStreamError = (): void => undefined; + const releaseQuarantine = (): void => { + if (lateErrorReleaseTimer !== undefined) clearTimeout(lateErrorReleaseTimer); + lateErrorReleaseTimer = undefined; + removeChildListener("error", swallowChildError); + removeChildListener("close", releaseQuarantine); + removeStreamListener(stdin, "error", swallowStreamError); + removeStreamListener(stdoutStream, "error", swallowStreamError); + removeStreamListener(stderrStream, "error", swallowStreamError); + }; + const quarantineLateErrors = (): void => { + // A final-deadline settlement can precede a broken ChildProcess object's terminal events. + // Keep only no-capture listeners for a bounded grace period so a late EventEmitter error + // cannot become uncaught, including after an already-observed close event. + child?.on("error", swallowChildError); + child?.once("close", releaseQuarantine); + stdin?.on("error", swallowStreamError); + stdoutStream?.on("error", swallowStreamError); + stderrStream?.on("error", swallowStreamError); + lateErrorReleaseTimer = setTimeout(releaseQuarantine, finalSettlementMs); + lateErrorReleaseTimer.unref?.(); + }; + const removeChildListener = (event: "error" | "close", listener: (...args: any[]) => void): void => { + try { child?.removeListener?.(event, listener); } catch { /* the helper is already terminal */ } + }; + const removeStreamListener = (stream: Writable | Readable | undefined, event: "data" | "error", listener: (...args: any[]) => void): void => { + try { stream?.removeListener(event, listener); } catch { /* the helper is already terminal */ } }; const stopStream = (stream: Writable | Readable | null | undefined): void => { try { stream?.destroy(); } catch { /* abort is already fail-closed */ } }; - const finishAfterClose = (): void => { - if (!closeSeen || settled) return; - if (aborted || typeof closeCode !== "number" || !Number.isInteger(closeCode) || closeSignal !== null) { - settle(() => reject(invalid())); + const onChildError = (): void => abort(); + const onStdinError = (): void => abort(); + const onStdoutError = (): void => abort(); + const onStderrError = (): void => abort(); + const onStdoutData = (chunk: Buffer): void => { + if (aborted || settled) return; + stdoutBytes += chunk.length; + if (stdoutBytes > invocation.maximumOutputBytes) { + abort(); + return; + } + stdout.push(Buffer.from(chunk)); + }; + const onStderrData = (chunk: Buffer): void => { + if (aborted || settled) return; + stderrBytes += chunk.length; + if (stderrBytes > MAX_RESPONSE_BYTES) { + abort(); + return; + } + stderr.push(Buffer.from(chunk)); + }; + const onClose = (code: number | null, signal: NodeJS.Signals | null): void => { + if (settled) return; + if (aborted || code === null || !Number.isInteger(code) || signal !== null) { + settle(() => reject(invalid()), true); return; } - const code = closeCode; settle(() => resolve({ code, stdout: Buffer.concat(stdout), stderr: Buffer.concat(stderr), })); }; + const cleanup = (quarantine = false): void => { + if (timeout !== undefined) clearTimeout(timeout); + if (terminationTimer !== undefined) clearTimeout(terminationTimer); + if (finalSettlementTimer !== undefined) clearTimeout(finalSettlementTimer); + removeChildListener("error", onChildError); + removeChildListener("close", onClose as (...args: any[]) => void); + removeStreamListener(stdin, "error", onStdinError); + removeStreamListener(stdoutStream, "data", onStdoutData); + removeStreamListener(stdoutStream, "error", onStdoutError); + removeStreamListener(stderrStream, "data", onStderrData); + removeStreamListener(stderrStream, "error", onStderrError); + if (quarantine) quarantineLateErrors(); + }; + const settle = (callback: () => void, quarantine = false): void => { + if (settled) return; + settled = true; + cleanup(quarantine); + callback(); + }; const abort = (): void => { if (aborted || settled) return; aborted = true; if (timeout !== undefined) clearTimeout(timeout); if (child !== undefined) { - stopStream(child.stdin); - stopStream(child.stdout); - stopStream(child.stderr); - try { child.kill(); } catch { /* the close handler still owns settlement */ } + stopStream(stdin); + stopStream(stdoutStream); + stopStream(stderrStream); + try { child.kill("SIGTERM"); } catch { /* final settlement still owns completion */ } + try { child.unref?.(); } catch { /* the bounded timers still own completion */ } } - finishAfterClose(); + terminationTimer = setTimeout(() => { + if (settled || child === undefined) return; + try { child.kill("SIGKILL"); } catch { /* final settlement still owns completion */ } + }, terminationGraceMs); + finalSettlementTimer = setTimeout(() => { + settle(() => reject(invalid()), true); + }, finalSettlementMs); }; try { child = spawnChild(invocation.executable, invocation.args, { @@ -319,43 +405,23 @@ async function invokeTht( settle(() => reject(invalid())); return; } - child.once("close", (code, signal) => { - closeSeen = true; - closeCode = code; - closeSignal = signal; - if (code === null || signal !== null) aborted = true; - finishAfterClose(); - }); - child.on("error", abort); + child.once("close", onClose); + child.on("error", onChildError); if (!child.stdin || !child.stdout || !child.stderr) { abort(); return; } - const stdin = child.stdin; - const stdoutStream = child.stdout; - const stderrStream = child.stderr; + stdin = child.stdin; + stdoutStream = child.stdout; + stderrStream = child.stderr; timeout = setTimeout(() => { abort(); }, invocation.timeoutMs); - stdoutStream.on("data", (chunk: Buffer) => { - if (aborted) return; - stdoutBytes += chunk.length; - if (stdoutBytes > MAX_RESPONSE_BYTES) { - abort(); - return; - } - stdout.push(Buffer.from(chunk)); - }); - stderrStream.on("data", (chunk: Buffer) => { - if (aborted) return; - stderrBytes += chunk.length; - if (stderrBytes > MAX_RESPONSE_BYTES) { - abort(); - return; - } - stderr.push(Buffer.from(chunk)); - }); - stdin.once("error", abort); + stdoutStream.on("data", onStdoutData); + stdoutStream.once("error", onStdoutError); + stderrStream.on("data", onStderrData); + stderrStream.once("error", onStderrError); + stdin.once("error", onStdinError); const writeInput = (): void => { if (aborted || settled) return; try { @@ -402,20 +468,33 @@ function createAuthStorageBridge( options: WindowsAuthStorageBridgeOptions = {}, ): WindowsAuthStorageBridge { const executable = safeThtExecutable(options.thtExecutable, pathStyle); + const testDeadlines = options.deadlinesForTest; + const timeoutMs = testDeadlines?.timeoutMs ?? TIMEOUT_MS; + const terminationGraceMs = testDeadlines?.terminationGraceMs ?? TERMINATION_GRACE_MS; + const finalSettlementMs = testDeadlines?.finalSettlementMs ?? FINAL_SETTLEMENT_MS; + if (!Number.isSafeInteger(timeoutMs) || timeoutMs < 1 || timeoutMs > TIMEOUT_MS + || !Number.isSafeInteger(terminationGraceMs) || terminationGraceMs < 1 || terminationGraceMs > TIMEOUT_MS + || !Number.isSafeInteger(finalSettlementMs) || finalSettlementMs <= terminationGraceMs || finalSettlementMs > TIMEOUT_MS) { + throw invalid(); + } const invoke = options.invoke ?? ((invocation: WindowsAuthStorageInvocation) => invokeTht( invocation, options.spawnChild, options.beforeInputForTest, + terminationGraceMs, + finalSettlementMs, )); const invokeSync = options.invokeSync ?? invokeThtSync; const request = async (value: BridgeRequest): Promise => { try { - const maximumOutputBytes = MAX_RESPONSE_BYTES; + const maximumOutputBytes = value.operation === "read-local-users" + ? MAX_AUTH_CONFIG_RESPONSE_BYTES + : MAX_RESPONSE_BYTES; const response = await invoke({ executable, args: ["_auth-storage"], input: encodedRequest(value, pathStyle), - timeoutMs: TIMEOUT_MS, + timeoutMs, maximumOutputBytes, }); return parseResponse(response, maximumOutputBytes); @@ -429,7 +508,7 @@ function createAuthStorageBridge( executable, args: ["_auth-storage"], input: encodedRequest(value, pathStyle), - timeoutMs: TIMEOUT_MS, + timeoutMs, maximumOutputBytes: MAX_AUTH_CONFIG_RESPONSE_BYTES, }); return parseResponse(response, MAX_AUTH_CONFIG_RESPONSE_BYTES); @@ -470,6 +549,19 @@ function createAuthStorageBridge( if (contents === undefined) throw invalid(); return contents; }, + async readLocalUsers(path) { + const paths = pathStyle === "windows" ? win32 : posix; + if (typeof path !== "string" || path.length === 0 || /[\u0000-\u001f\u007f]/.test(path) + || !paths.isAbsolute(path) || paths.normalize(path) !== path) throw invalid(); + const root = paths.dirname(path); + const filename = paths.basename(path); + if (!AUTH_CONFIG_FILENAME.test(filename) || paths.join(root, filename) !== path) throw invalid(); + const response = await request({ version: PROTOCOL_VERSION, operation: "read-local-users", root, filename }); + if (Object.keys(response).some((key) => !["version", "ok", "found", "contentBase64"].includes(key))) throw invalid(); + const contents = contentFrom(response, MAX_AUTH_CONFIG_BYTES); + if (contents === undefined) throw invalid(); + return contents; + }, async create(root, directory, filename, contents) { if (!Buffer.isBuffer(contents) || contents.length === 0 || contents.length > directoryMaximum(directory)) throw invalid(); const response = await request(recordRequest("create", root, directory, filename, contents)); diff --git a/backend/test/app-auth-mode.test.ts b/backend/test/app-auth-mode.test.ts index 31d03fa3..7b4515dd 100644 --- a/backend/test/app-auth-mode.test.ts +++ b/backend/test/app-auth-mode.test.ts @@ -1,4 +1,4 @@ -import { expect, test } from "vitest"; +import { expect, test, vi } from "vitest"; import { chmodSync, mkdtempSync, realpathSync, rmSync, writeFileSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; @@ -35,3 +35,56 @@ test("configured OIDC advertises login but fails closed without its runtime clie rmSync(directory, { recursive: true, force: true }); } }); + +test("configured OIDC initializes login from the literal secret bundle without an environment duplicate", async () => { + const directory = mkdtempSync(join(realpathSync(tmpdir()), "thothii-app-oidc-bundle-")); + chmodSync(directory, 0o700); + const file = join(directory, "auth.yaml"); + const bundle = join(directory, "thothii.secrets"); + const clientSecret = "bundle-only-oidc-client-secret"; + writeFileSync(file, stringify({ + version: 1, mode: "oidc", publicUrl: "https://thothii.example.org", + oidc: { + issuer: "https://authentik.example.org/application/o/thothii/", clientId: "thothii", + clientSecretRef: "THT_OIDC_CLIENT_SECRET", scopes: ["openid"], groupsClaim: "groups", + }, + groupCatalog: { driver: "authentik", baseUrl: "https://authentik.example.org", apiTokenRef: "THT_AUTHENTIK_API_TOKEN" }, + authorization: { groupRoles: { "TOT Users": ["user"], "TOT Admin": ["admin"] } }, + }), { encoding: "utf8", mode: 0o600 }); + writeFileSync(bundle, `THT_OIDC_CLIENT_SECRET=${clientSecret}\nTHT_AUTHENTIK_API_TOKEN=bundle-only-authentik-token\n`, { + encoding: "utf8", mode: 0o600, + }); + chmodSync(file, 0o600); + chmodSync(bundle, 0o600); + const original = process.env.THT_OIDC_CLIENT_SECRET; + delete process.env.THT_OIDC_CLIENT_SECRET; + const oidcProtocolFactory = vi.fn((input: { clientSecret: string }) => ({ + authorizationUrl: async ({ state }: { state: string }) => new URL(`https://authentik.example.org/authorize?state=${state}`), + callback: async () => { throw new Error("callback is outside this login-start regression"); }, + diagnose: async () => undefined, + })); + const authSessionStore = { + createOidcState: async () => ({ + state: "s".repeat(43), + record: { version: 1 }, + }), + }; + try { + const app = buildApp(loadConfig({ + NODE_ENV: "test", THT_AUTH_CONFIG_FILE: file, THT_AUTH_STATE_ROOT: join(directory, "auth-state"), + THT_SECRETS_FILE: bundle, + }), { oidcProtocolFactory, authSessionStore } as never); + try { + const response = await app.inject({ method: "GET", url: "/auth/oidc/login" }); + expect(response.statusCode).toBe(302); + expect(oidcProtocolFactory).toHaveBeenCalledWith(expect.objectContaining({ clientSecret })); + expect(process.env.THT_OIDC_CLIENT_SECRET).toBeUndefined(); + } finally { + await app.close(); + } + } finally { + if (original === undefined) delete process.env.THT_OIDC_CLIENT_SECRET; + else process.env.THT_OIDC_CLIENT_SECRET = original; + rmSync(directory, { recursive: true, force: true }); + } +}); diff --git a/backend/test/auth-diagnostics.test.ts b/backend/test/auth-diagnostics.test.ts index 58a74292..894b589f 100644 --- a/backend/test/auth-diagnostics.test.ts +++ b/backend/test/auth-diagnostics.test.ts @@ -7,7 +7,6 @@ import { createAuthenticationConfigProvider } from "../src/auth/config.js"; import { createAuthentikGroupCatalog } from "../src/auth/authentik-group-catalog.js"; import { createLocalUserRegistry } from "../src/auth/local-registry.js"; import { createOidcProtocol, OidcJwksUnavailableError } from "../src/auth/oidc-client.js"; -import { validateAuthSessionRoot } from "../src/auth/session-store.js"; import type { LoadedAuthConfig } from "../src/auth/types.js"; const sentinels = [ @@ -152,7 +151,7 @@ test("distinguishes a valid registry without an enabled admin from a malformed r chmodSync(validUsers, 0o600); const validReport = await createAuthDiagnoser({ authMode: "local", authStateRoot: validRoot, - sessionRootValidator: validateAuthSessionRoot, + sessionRootValidator: acceptSessionRoot, authentication: { current: () => localConfig(join(validRoot, "auth.yaml")) }, localUserRegistry: createLocalUserRegistry(validUsers), }).inspect({ live: false }); @@ -164,7 +163,7 @@ test("distinguishes a valid registry without an enabled admin from a malformed r chmodSync(malformedUsers, 0o600); const malformedReport = await createAuthDiagnoser({ authMode: "local", authStateRoot: malformedRoot, - sessionRootValidator: validateAuthSessionRoot, + sessionRootValidator: acceptSessionRoot, authentication: { current: () => localConfig(join(malformedRoot, "auth.yaml")) }, localUserRegistry: createLocalUserRegistry(malformedUsers), }).inspect({ live: false }); @@ -182,7 +181,7 @@ test("maps unsafe auth.yaml storage from the real provider to a redacted config chmodSync(unsafePath, 0o640); const report = await createAuthDiagnoser({ authMode: "local", authStateRoot: root, - sessionRootValidator: validateAuthSessionRoot, + sessionRootValidator: acceptSessionRoot, authentication: createAuthenticationConfigProvider(unsafePath), }).inspect({ live: false }); @@ -442,73 +441,6 @@ test("redacts exceptional configuration, registry, protocol, and catalog errors" expect(report.checks.every((check) => check.level === "error" || check.level === "info")).toBe(true); }); -test.skipIf(process.platform === "win32")("uses the runtime validator for canonical, private session roots", async () => { - const dependencies = (authStateRoot: string) => ({ - authMode: "none" as const, - authStateRoot, - sessionRootValidator: validateAuthSessionRoot, - }); - const valid = privateRoot(); - await expect(createAuthDiagnoser(dependencies(valid)).inspect({ live: false })) - .resolves.toMatchObject({ ready: true, checks: [expect.objectContaining({ code: "auth_ready" })] }); - - const realRoot = join(privateRoot(), "real-auth"); - mkdirSync(realRoot, { mode: 0o700 }); - chmodSync(realRoot, 0o700); - const linkedRoot = join(privateRoot(), "linked-auth"); - symlinkSync(realRoot, linkedRoot); - const absent = join(privateRoot(), "absent-auth"); - const absentParent = join(privateRoot(), "absent-parent"); - const absentNested = join(absentParent, "auth"); - const blockedParent = join(privateRoot(), "not-a-directory"); - writeFileSync(blockedParent, "blocked", { mode: 0o600 }); - const traversal = `${valid}/../${basename(valid)}`; - - const missingReport = await createAuthDiagnoser(dependencies(absent)).inspect({ live: false }); - expect(missingReport).toMatchObject({ ready: true, checks: [expect.objectContaining({ code: "auth_ready" })] }); - expect(existsSync(absent)).toBe(false); - - for (const unsafe of [traversal, linkedRoot, absentNested, join(blockedParent, "auth")]) { - const report = await createAuthDiagnoser(dependencies(unsafe)).inspect({ live: false }); - expect(report).toMatchObject({ ready: false, checks: [expect.objectContaining({ code: "auth_session_store_invalid" })] }); - expect(JSON.stringify(report)).not.toContain(unsafe); - } - expect(existsSync(absentParent)).toBe(false); - - chmodSync(valid, 0o750); - await expect(createAuthDiagnoser(dependencies(valid)).inspect({ live: false })) - .resolves.toMatchObject({ ready: false, checks: [expect.objectContaining({ code: "auth_session_store_invalid" })] }); -}); - -test.skipIf(process.platform === "win32")("diagnoses unsafe existing session-store children without creating missing children", async () => { - const root = privateRoot(); - const outside = privateRoot(); - symlinkSync(outside, join(root, "sessions")); - - const linked = await createAuthDiagnoser({ - authMode: "none", authStateRoot: root, sessionRootValidator: validateAuthSessionRoot, - }).inspect({ live: false }); - expect(linked).toMatchObject({ - ready: false, - checks: [expect.objectContaining({ code: "auth_session_store_invalid" })], - }); - expect(existsSync(join(root, "oidc"))).toBe(false); - expect(JSON.stringify(linked)).not.toContain(root); - - rmSync(join(root, "sessions")); - mkdirSync(join(root, "sessions"), { mode: 0o700 }); - chmodSync(join(root, "sessions"), 0o700); - mkdirSync(join(root, "oidc"), { mode: 0o700 }); - chmodSync(join(root, "oidc"), 0o750); - const nonPrivate = await createAuthDiagnoser({ - authMode: "none", authStateRoot: root, sessionRootValidator: validateAuthSessionRoot, - }).inspect({ live: false }); - expect(nonPrivate).toMatchObject({ - ready: false, - checks: [expect.objectContaining({ code: "auth_session_store_invalid" })], - }); -}); - test.skipIf(process.platform === "win32")("routes production POSIX static validation through the native auth-storage bridge", async () => { const validateRoot = vi.fn(async () => undefined); const report = await createAuthDiagnoser({ diff --git a/backend/test/auth-dynamic-registry.test.ts b/backend/test/auth-dynamic-registry.test.ts index 854e912a..67eb01fa 100644 --- a/backend/test/auth-dynamic-registry.test.ts +++ b/backend/test/auth-dynamic-registry.test.ts @@ -6,7 +6,7 @@ import { stringify } from "yaml"; import { buildApp, type AppWithAuthSessionStore } from "../src/app.js"; import { loadAuthenticationConfig } from "../src/auth/config.js"; import { loadConfig } from "../src/config.js"; -import { prepareAuthStateRoot } from "./auth-test-fixtures.js"; +import { createFixtureAuthStorageBridge, prepareAuthStateRoot } from "./auth-test-fixtures.js"; const password = "correct horse battery staple"; const passwordHash = "$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4"; @@ -54,7 +54,7 @@ test("each login and session resolve uses the current config snapshot users file THT_AUTH_CONFIG_FILE: authFile, THT_AUTH_STATE_ROOT: authStateRoot, THT_HARNESS_DIR: "/tmp/h", - })); + }), { authStorageBridgeForTest: createFixtureAuthStorageBridge() }); cleanups.push(async () => { await app.close(); rmSync(directory, { recursive: true, force: true }); diff --git a/backend/test/auth-request-snapshot.test.ts b/backend/test/auth-request-snapshot.test.ts index 836344b5..548723a9 100644 --- a/backend/test/auth-request-snapshot.test.ts +++ b/backend/test/auth-request-snapshot.test.ts @@ -7,7 +7,7 @@ import { buildApp, type AppWithAuthSessionStore } from "../src/app.js"; import { loadAuthenticationConfig } from "../src/auth/config.js"; import type { AuthenticationConfigProvider, LoadedAuthConfig } from "../src/auth/types.js"; import { loadConfig } from "../src/config.js"; -import { prepareAuthStateRoot } from "./auth-test-fixtures.js"; +import { createFixtureAuthStorageBridge, prepareAuthStateRoot } from "./auth-test-fixtures.js"; const password = "correct horse battery staple"; const passwordHash = "$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4"; @@ -87,7 +87,7 @@ async function createFixture(first: "A" | "B", later: "A" | "B" | "oidc") { THT_HARNESS_DIR: "/tmp/h", }); config.authentication = provider; - const app = buildApp(config) as AppWithAuthSessionStore; + const app = buildApp(config, { authStorageBridgeForTest: createFixtureAuthStorageBridge() }) as AppWithAuthSessionStore; cleanups.push(async () => { await app.close(); rmSync(directory, { recursive: true, force: true }); diff --git a/backend/test/auth-routes-local.test.ts b/backend/test/auth-routes-local.test.ts index 1dc40e15..8ff97319 100644 --- a/backend/test/auth-routes-local.test.ts +++ b/backend/test/auth-routes-local.test.ts @@ -6,7 +6,7 @@ import { stringify } from "yaml"; import { buildApp } from "../src/app.js"; import { loadConfig } from "../src/config.js"; import { LoginFailureLimiter } from "../src/auth/routes.js"; -import { prepareAuthStateRoot } from "./auth-test-fixtures.js"; +import { createFixtureAuthStorageBridge, prepareAuthStateRoot } from "./auth-test-fixtures.js"; const password = "correct horse battery staple"; const passwordHash = "$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4"; @@ -88,7 +88,10 @@ async function createLocalApp(options: { THT_AUTH_CONFIG_FILE: authConfigFile, THT_AUTH_STATE_ROOT: authStateRoot, THT_HARNESS_DIR: "/tmp/h", - }), options.registry === undefined ? undefined : { localUserRegistry: options.registry } as any); + }), { + ...(options.registry === undefined ? {} : { localUserRegistry: options.registry }), + authStorageBridgeForTest: createFixtureAuthStorageBridge(), + } as any); cleanups.push(async () => { await app.close(); rmSync(directory, { recursive: true, force: true }); @@ -155,7 +158,7 @@ test("remembered login uses a persistent secure cookie under an HTTPS public URL chmodSync(usersFile, 0o600); prepareAuthStateRoot(authStateRoot); const config = () => loadConfig({ THT_AUTH_CONFIG_FILE: authConfigFile, THT_AUTH_STATE_ROOT: authStateRoot, THT_HARNESS_DIR: "/tmp/h" }); - const first = buildApp(config()); + const first = buildApp(config(), { authStorageBridgeForTest: createFixtureAuthStorageBridge() }); try { const signedIn = await first.inject({ method: "POST", @@ -169,7 +172,7 @@ test("remembered login uses a persistent secure cookie under an HTTPS public URL expect(setCookie).toContain("Secure"); await first.close(); - const restarted = buildApp(config()); + const restarted = buildApp(config(), { authStorageBridgeForTest: createFixtureAuthStorageBridge() }); cleanups.push(async () => { await restarted.close(); rmSync(directory, { recursive: true, force: true }); diff --git a/backend/test/auth-session-store.test.ts b/backend/test/auth-session-store.test.ts index 1fa8c652..7531ef2f 100644 --- a/backend/test/auth-session-store.test.ts +++ b/backend/test/auth-session-store.test.ts @@ -7,18 +7,20 @@ import { lstatSync, mkdirSync, mkdtempSync, + opendirSync, readFileSync, readdirSync, realpathSync, renameSync, rmSync, + statSync, symlinkSync, utimesSync, unlinkSync, writeFileSync, } from "node:fs"; import { tmpdir } from "node:os"; -import { basename, join } from "node:path"; +import { basename, dirname, join } from "node:path"; import { afterEach, describe, expect, test, vi } from "vitest"; const fsHooks = vi.hoisted(() => ({ @@ -66,12 +68,15 @@ vi.mock("node:fs", async (importOriginal) => { import { createFileAuthSessionStore, deriveCsrfToken, - validateAuthSessionRoot, type AuthSessionStore, type FileAuthSessionStoreOptions, type SessionCreateInput, } from "../src/auth/session-store.js"; -import { createWindowsAuthStorageBridge } from "../src/auth/windows-auth-storage.js"; +import { + createWindowsAuthStorageBridge, + type WindowsAuthStorageBridge, + type WindowsAuthStorageDirectory, +} from "../src/auth/windows-auth-storage.js"; const roots: string[] = []; const base = new Date("2030-01-02T03:04:05.000Z"); @@ -123,11 +128,203 @@ function claimPath(rootPath: string, rawState: string): string { return join(rootPath, "oidc", `${createHash("sha256").update(rawState).digest("hex")}.claim`); } +// This fixture adapter keeps session-store behavior tests self-contained. Production has no +// Node filesystem fallback: it always uses the hidden Go bridge. The Go authstorage suite owns +// descriptor-pinning/race assertions; this adapter only supplies ordinary fixture semantics. +function testPosixStorageBridge( + ensureOverride?: (rootPath: string) => Promise, +): WindowsAuthStorageBridge { + let overrideUsed = false; + const ensure = async (rootPath: string): Promise => { + if (ensureOverride && !overrideUsed) { + overrideUsed = true; + await ensureOverride(rootPath); + return; + } + if (!existsSync(rootPath)) { + if (realpathSync(dirname(rootPath)) !== dirname(rootPath)) throw new Error("invalid"); + mkdirSync(rootPath, { mode: 0o700 }); + chmodSync(rootPath, 0o700); + } + const rootInfo = lstatSync(rootPath); + if (!rootInfo.isDirectory() || rootInfo.isSymbolicLink() || (rootInfo.mode & 0o7777) !== 0o700) throw new Error("invalid"); + for (const child of ["sessions", "oidc"]) { + const path = join(rootPath, child); + if (!existsSync(path)) { + mkdirSync(path, { mode: 0o700 }); + chmodSync(path, 0o700); + } + const info = lstatSync(path); + if (!info.isDirectory() || info.isSymbolicLink() || (info.mode & 0o7777) !== 0o700) throw new Error("invalid"); + } + }; + const directory = async (rootPath: string, name: WindowsAuthStorageDirectory): Promise => { + await ensure(rootPath); + return join(rootPath, name); + }; + const record = (path: string, links: readonly number[]): import("node:fs").Stats => { + const info = lstatSync(path); + if (!info.isFile() || info.isSymbolicLink() || !links.includes(info.nlink) || (info.mode & 0o7777) !== 0o600 + || (typeof process.geteuid === "function" && info.uid !== process.geteuid())) { + throw new Error("invalid"); + } + return info; + }; + const same = (left: import("node:fs").Stats, right: import("node:fs").Stats): boolean => + left.dev === right.dev && left.ino === right.ino && left.nlink === right.nlink; + const paths = async (rootPath: string, name: WindowsAuthStorageDirectory, filename: string): Promise => + join(await directory(rootPath, name), filename); + const listNames = async ( + rootPath: string, + name: WindowsAuthStorageDirectory, + maximumEntries: number, + ): Promise<{ name: string; modifiedUnixMs: number }[]> => { + const handle = opendirSync(await directory(rootPath, name)); + const entries: string[] = []; + try { + for (;;) { + const entry = handle.readSync(); + if (entry === null) break; + entries.push(entry.name); + if (entries.length > maximumEntries) throw new Error("invalid"); + } + } finally { + handle.closeSync(); + } + return entries.sort().map((filename) => { + const links = name === "oidc" ? [1, 2] : [1]; + const info = record(join(rootPath, name, filename), links); + return { name: filename, modifiedUnixMs: info.mtimeMs }; + }); + }; + return { + validateRoot: async (rootPath) => { await ensure(rootPath); }, + ensureLayout: ensure, + readAuthConfig: (path) => readFileSync(path), + readLocalUsers: async (path) => readFileSync(path), + create: async (rootPath, name, filename, contents) => { + const path = await paths(rootPath, name, filename); + try { + writeFileSync(path, contents, { flag: "wx", mode: 0o600 }); + chmodSync(path, 0o600); + return true; + } catch (error: any) { + if (error?.code !== "EEXIST") throw error; + record(path, [1]); + return false; + } + }, + read: async (rootPath, name, filename) => { + const path = await paths(rootPath, name, filename); + try { + record(path, [1]); + } catch (error: any) { + if (error?.code === "ENOENT") return undefined; + throw error; + } + return readFileSync(path); + }, + replace: async (rootPath, name, filename, contents) => { + const path = await paths(rootPath, name, filename); + record(path, [1]); + const temporary = `${path}.test-replacement`; + writeFileSync(temporary, contents, { flag: "wx", mode: 0o600 }); + chmodSync(temporary, 0o600); + renameSync(temporary, path); + }, + remove: async (rootPath, name, filename) => { + const path = await paths(rootPath, name, filename); + try { + record(path, [1]); + } catch (error: any) { + if (error?.code === "ENOENT") return false; + throw error; + } + unlinkSync(path); + return true; + }, + list: async (rootPath, name, maximumEntries = 256) => await listNames(rootPath, name, maximumEntries), + listPage: async (rootPath, name, afterName, maximumEntries) => { + if (name !== "sessions") throw new Error("invalid"); + const handle = opendirSync(await directory(rootPath, name)); + const all: string[] = []; + try { + for (;;) { + const entry = handle.readSync(); + if (entry === null) break; + all.push(entry.name); + } + } finally { + handle.closeSync(); + } + all.sort(); + for (const filename of all) record(join(rootPath, name, filename), [1]); + const selected = all.filter((filename) => afterName === undefined || filename > afterName).slice(0, maximumEntries + 1); + return { + entries: selected.slice(0, maximumEntries).map((filename) => { + const info = statSync(join(rootPath, name, filename)); + return { name: filename, modifiedUnixMs: info.mtimeMs }; + }), + more: selected.length > maximumEntries, + }; + }, + claimConsume: async (rootPath, filename) => { + const source = await paths(rootPath, "oidc", filename); + const claim = source.replace(/\.json$/, ".claim"); + try { + record(source, [1]); + } catch (error: any) { + if (error?.code === "ENOENT") return undefined; + return undefined; + } + try { + linkSync(source, claim); + } catch (error: any) { + if (error?.code === "EEXIST") return undefined; + throw error; + } + const contents = readFileSync(source); + unlinkSync(source); + unlinkSync(claim); + return contents; + }, + readClaim: async (rootPath, filename) => { + const source = await paths(rootPath, "oidc", filename); + const claim = source.replace(/\.json$/, ".claim"); + try { + const sourceInfo = record(source, [2]); + const claimInfo = record(claim, [2]); + if (!same(sourceInfo, claimInfo)) throw new Error("invalid"); + return readFileSync(source); + } catch (error: any) { + if (error?.code === "ENOENT") return undefined; + return undefined; + } + }, + removeClaim: async (rootPath, filename) => { + const source = await paths(rootPath, "oidc", filename); + const claim = source.replace(/\.json$/, ".claim"); + try { + const sourceInfo = record(source, [2]); + const claimInfo = record(claim, [2]); + if (!same(sourceInfo, claimInfo)) throw new Error("invalid"); + } catch (error: any) { + if (error?.code === "ENOENT") return false; + return false; + } + unlinkSync(source); + unlinkSync(claim); + return true; + }, + }; +} + function validStore(storageRoot: string, options: FileAuthSessionStoreOptions = {}): AuthSessionStore { + const posixStorageBridge = options.posixStorageBridge ?? testPosixStorageBridge(); return createFileAuthSessionStore(storageRoot, { currentAuthConfigRevision: () => revision, findLocalUser: async () => validLocalUser, - }, options); + }, { ...options, posixStorageBridge }); } async function create( @@ -246,56 +443,7 @@ async function isolatedOidcCreator(storageRoot: string, attempts: number): Promi } describe("file-backed auth session store", () => { - test.skipIf(process.platform === "win32")("exports its side-effect-free canonical session-root validator", () => { - const valid = root(); - expect(() => validateAuthSessionRoot(valid)).not.toThrow(); - - const outer = root(); - const realRoot = join(outer, "real-auth"); - mkdirSync(realRoot, { mode: 0o700 }); - chmodSync(realRoot, 0o700); - const linkedRoot = join(outer, "linked-auth"); - symlinkSync(realRoot, linkedRoot); - const traversal = `${valid}/../${basename(valid)}`; - const absent = join(outer, "absent-auth"); - const absentNested = join(outer, "absent-parent", "auth"); - const fileParent = join(outer, "not-a-directory"); - writeFileSync(fileParent, "blocked", { mode: 0o600 }); - - expect(() => validateAuthSessionRoot(absent)).not.toThrow(); - expect(existsSync(absent)).toBe(false); - for (const unsafe of [traversal, linkedRoot, absentNested, join(fileParent, "auth")]) { - expect(() => validateAuthSessionRoot(unsafe)).toThrow("auth_session_store_invalid"); - } - expect(existsSync(join(outer, "absent-parent"))).toBe(false); - - for (const child of ["sessions", "oidc"] as const) { - const childRoot = join(outer, `child-${child}`); - mkdirSync(childRoot, { mode: 0o700 }); - chmodSync(childRoot, 0o700); - const childPath = join(childRoot, child); - const outside = root(); - symlinkSync(outside, childPath); - expect(() => validateAuthSessionRoot(childRoot)).toThrow("auth_session_store_invalid"); - expect(readdirSync(outside).sort()).toEqual(["oidc", "sessions"]); - unlinkSync(childPath); - mkdirSync(childPath, { mode: 0o700 }); - chmodSync(childPath, 0o750); - expect(() => validateAuthSessionRoot(childRoot)).toThrow("auth_session_store_invalid"); - } - - const missingChildren = join(outer, "missing-children"); - mkdirSync(missingChildren, { mode: 0o700 }); - chmodSync(missingChildren, 0o700); - expect(() => validateAuthSessionRoot(missingChildren)).not.toThrow(); - expect(existsSync(join(missingChildren, "sessions"))).toBe(false); - expect(existsSync(join(missingChildren, "oidc"))).toBe(false); - - chmodSync(valid, 0o750); - expect(() => validateAuthSessionRoot(valid)).toThrow("auth_session_store_invalid"); - }); - - test.skipIf(process.platform === "win32")("delegates missing layout creation and then enforces static/runtime parity", async () => { + test.skipIf(process.platform === "win32")("delegates missing layout creation to the retained native bridge", async () => { const storageRoot = join(root(), "auth"); const ensureLayout = vi.fn(async (requestedRoot: string) => { expect(requestedRoot).toBe(storageRoot); @@ -306,14 +454,12 @@ describe("file-backed auth session store", () => { chmodSync(join(requestedRoot, child), 0o700); } }); - const store = validStore(storageRoot, { posixStorageBridge: { ensureLayout } }); + const store = validStore(storageRoot, { posixStorageBridge: testPosixStorageBridge(ensureLayout) }); await expect(create(store)).resolves.toMatchObject({ record: { method: "local" } }); expect(ensureLayout).toHaveBeenCalledOnce(); - expect(() => validateAuthSessionRoot(storageRoot)).not.toThrow(); chmodSync(join(storageRoot, "oidc"), 0o750); - expect(() => validateAuthSessionRoot(storageRoot)).toThrow("auth_session_store_invalid"); await expectStoreInvalid(store.createOidcState(oidcInput("n".repeat(16), "v".repeat(43)), base)); }); @@ -330,7 +476,7 @@ describe("file-backed auth session store", () => { symlinkSync(outside, parent); throw new Error(`${storageRoot} rejected`); }); - const store = validStore(storageRoot, { posixStorageBridge: { ensureLayout } }); + const store = validStore(storageRoot, { posixStorageBridge: testPosixStorageBridge(ensureLayout) }); await expectStoreInvalid(create(store)); expect(ensureLayout).toHaveBeenCalledOnce(); @@ -349,14 +495,13 @@ describe("file-backed auth session store", () => { expect(existsSync(join(outside, "auth"))).toBe(false); }); - test.skipIf(process.platform === "win32")("rejects an uncreatable missing root without side effects in static and runtime paths", async () => { + test.skipIf(process.platform === "win32")("rejects an uncreatable missing root without side effects", async () => { const outer = root(); const lockedParent = join(outer, "locked-parent"); mkdirSync(lockedParent, { mode: 0o700 }); chmodSync(lockedParent, 0o500); const storageRoot = join(lockedParent, "auth"); try { - expect(() => validateAuthSessionRoot(storageRoot)).toThrow("auth_session_store_invalid"); await expectStoreInvalid(create(validStore(storageRoot))); expect(existsSync(storageRoot)).toBe(false); } finally { @@ -364,43 +509,20 @@ describe("file-backed auth session store", () => { } }); - test.skipIf(process.platform === "win32")("detects an ancestor replacement before creating any session directory", async () => { - const outer = root(); - const outside = root(); - const parent = join(outer, "parent"); - const movedParent = join(outer, "parent-original"); - mkdirSync(parent, { mode: 0o700 }); - chmodSync(parent, 0o700); - let replaced = false; - fsHooks.afterLstat = (observed) => { - if (observed !== parent) return false; - renameSync(parent, movedParent); - symlinkSync(outside, parent); - replaced = true; - return true; - }; + test("does not fall back to Node paths when the retained POSIX bridge rejects creation", async () => { + const storageRoot = join(root(), "auth"); + const bridge = { + ...testPosixStorageBridge(), + create: async () => { throw new Error("native create rejected"); }, + } as WindowsAuthStorageBridge; - await expectStoreInvalid(create(validStore(join(parent, "auth")))); - expect(replaced).toBe(true); - expect(existsSync(join(outside, "auth"))).toBe(false); - expect(existsSync(join(movedParent, "auth"))).toBe(false); - }); - - test.skipIf(process.platform === "win32")("rejects a session root owned by another identity", () => { - const storageRoot = root(); - fsHooks.transformLstat = (observed, info) => { - if (observed !== storageRoot) return info; - const foreign = Object.create(info) as import("node:fs").Stats; - Object.defineProperty(foreign, "uid", { value: info.uid + 1 }); - return foreign; - }; - - expect(() => validateAuthSessionRoot(storageRoot)).toThrow("auth_session_store_invalid"); + await expectStoreInvalid(create(validStore(storageRoot, { posixStorageBridge: bridge }))); + expect(existsSync(storageRoot)).toBe(false); }); test("fails closed and revokes a session when constructed without validity dependencies", async () => { const storageRoot = root(); - const store = createFileAuthSessionStore(storageRoot); + const store = createFileAuthSessionStore(storageRoot, undefined, { posixStorageBridge: testPosixStorageBridge() }); const created = await create(store); await expect(store.resolve(created.token)).resolves.toBeUndefined(); @@ -412,7 +534,7 @@ describe("file-backed auth session store", () => { const store = createFileAuthSessionStore(storageRoot, { currentAuthConfigRevision: () => { throw new Error("dependency unavailable"); }, findLocalUser: async () => validLocalUser, - }); + }, { posixStorageBridge: testPosixStorageBridge() }); const created = await create(store); await expectStoreInvalid(store.resolve(created.token)); @@ -562,21 +684,13 @@ describe("file-backed auth session store", () => { .rejects.toThrow("auth_session_store_invalid"); }); - test.skipIf(process.platform === "win32")("fails closed if the ordinary-session directory changes during a page scan", async () => { + test("fails closed when the retained POSIX bridge rejects a session continuation page", async () => { const storageRoot = root(); - const store = validStore(storageRoot); - const seed = await create(store, { idleTtlMs: 30 * 60_000, absoluteTtlMs: 30 * 60_000 }); - const contents = readFileSync(digestPath(storageRoot, "sessions", seed.token)); - unlinkSync(digestPath(storageRoot, "sessions", seed.token)); - const filename = sessionFilename(0); - const sessionsDirectory = join(storageRoot, "sessions"); - writeFileSync(join(sessionsDirectory, filename), contents, { mode: 0o600 }); - fsHooks.beforeLstat = (path) => { - if (path !== join(sessionsDirectory, filename)) return false; - const changed = new Date(base.getTime() + 60 * 60_000); - utimesSync(sessionsDirectory, changed, changed); - return true; - }; + const bridge = { + ...testPosixStorageBridge(), + listPage: async () => { throw new Error("native page rejected"); }, + } as WindowsAuthStorageBridge; + const store = validStore(storageRoot, { posixStorageBridge: bridge }); await expect(store.prune(new Date(base.getTime() + 2 * 60_000))) .rejects.toThrow("auth_session_store_invalid"); @@ -853,46 +967,38 @@ describe("file-backed auth session store", () => { await expectStoreInvalid(create(validStore(linkedRoot))); }); - test.skipIf(process.platform === "win32")("refuses storage owned by a different identity", async () => { + test("fails closed when the retained POSIX bridge rejects a session read", async () => { const storageRoot = root(); - const store = validStore(storageRoot); + const bridge = { + ...testPosixStorageBridge(), + read: async () => { throw new Error("native read rejected"); }, + } as WindowsAuthStorageBridge; + const store = validStore(storageRoot, { posixStorageBridge: bridge }); const created = await create(store); - const sessions = join(storageRoot, "sessions"); - fsHooks.transformLstat = (observed, info) => { - if (observed !== sessions) return info; - const foreign = Object.create(info) as import("node:fs").Stats; - Object.defineProperty(foreign, "uid", { value: info.uid + 1 }); - return foreign; - }; await expectStoreInvalid(store.resolve(created.token)); }); - test.skipIf(process.platform === "win32")("refuses directory replacement during a session read", async () => { + test("does not bypass a retained POSIX bridge read failure", async () => { const storageRoot = root(); - const store = validStore(storageRoot); + const bridge = { + ...testPosixStorageBridge(), + read: async () => { throw new Error("native root replacement rejected"); }, + } as WindowsAuthStorageBridge; + const store = validStore(storageRoot, { posixStorageBridge: bridge }); const created = await create(store); - const sessions = join(storageRoot, "sessions"); - const replacement = join(storageRoot, "sessions-replacement"); - fsHooks.afterRead = () => { - renameSync(sessions, replacement); - symlinkSync(replacement, sessions); - }; await expectStoreInvalid(store.resolve(created.token)); }); - test("refuses file replacement during a touch", async () => { + test("fails closed when the retained POSIX bridge rejects a session replacement", async () => { const storageRoot = root(); - const store = validStore(storageRoot); + const bridge = { + ...testPosixStorageBridge(), + replace: async () => { throw new Error("native replace rejected"); }, + } as WindowsAuthStorageBridge; + const store = validStore(storageRoot, { posixStorageBridge: bridge }); const created = await create(store); - const path = digestPath(storageRoot, "sessions", created.token); - const replacement = `${path}.replacement`; - fsHooks.afterWrite = () => { - writeFileSync(replacement, "{}", { encoding: "utf8", mode: 0o600 }); - chmodSync(replacement, 0o600); - renameSync(replacement, path); - }; await expectStoreInvalid(store.touch(created.token, new Date(base.getTime() + 5 * 60_000))); }); @@ -914,37 +1020,26 @@ describe("file-backed auth session store", () => { await expectStoreInvalid(store.touch(created.token, new Date(base.getTime() + 5 * 60_000))); }); - test("refuses file replacement during revoke", async () => { + test("fails closed when the retained POSIX bridge rejects a session removal", async () => { const storageRoot = root(); - const store = validStore(storageRoot); + const bridge = { + ...testPosixStorageBridge(), + remove: async () => { throw new Error("native remove rejected"); }, + } as WindowsAuthStorageBridge; + const store = validStore(storageRoot, { posixStorageBridge: bridge }); const created = await create(store); - const path = digestPath(storageRoot, "sessions", created.token); - const replacement = `${path}.replacement`; - writeFileSync(replacement, "{}", { encoding: "utf8", mode: 0o600 }); - chmodSync(replacement, 0o600); - fsHooks.afterLstat = (observed) => { - if (observed !== path) return false; - renameSync(replacement, path); - return true; - }; await expectStoreInvalid(store.revoke(created.token)); - expect(readFileSync(path, "utf8")).toBe("{}"); }); - test.skipIf(process.platform === "win32")("refuses directory replacement during revoke", async () => { + test("does not bypass a retained POSIX bridge removal failure", async () => { const storageRoot = root(); - const store = validStore(storageRoot); + const bridge = { + ...testPosixStorageBridge(), + remove: async () => { throw new Error("native root replacement rejected"); }, + } as WindowsAuthStorageBridge; + const store = validStore(storageRoot, { posixStorageBridge: bridge }); const created = await create(store); - const path = digestPath(storageRoot, "sessions", created.token); - const sessions = join(storageRoot, "sessions"); - const replacement = join(storageRoot, "sessions-replacement"); - fsHooks.afterLstat = (observed) => { - if (observed !== path) return false; - renameSync(sessions, replacement); - symlinkSync(replacement, sessions); - return true; - }; await expectStoreInvalid(store.revoke(created.token)); }); @@ -1354,7 +1449,7 @@ describe("file-backed auth session store", () => { const store = createFileAuthSessionStore(storageRoot, { currentAuthConfigRevision: () => currentRevision, findLocalUser: async () => localUser, - }); + }, { posixStorageBridge: testPosixStorageBridge() }); const configChanged = await create(store); currentRevision = "b".repeat(64); diff --git a/backend/test/auth-test-fixtures.ts b/backend/test/auth-test-fixtures.ts index d9a4cf15..60ce878b 100644 --- a/backend/test/auth-test-fixtures.ts +++ b/backend/test/auth-test-fixtures.ts @@ -5,6 +5,7 @@ import type { FastifyInstance } from "fastify"; import { stringify } from "yaml"; import { buildApp, type BuildAppDeps } from "../src/app.js"; import { loadConfig } from "../src/config.js"; +import { createFixturePosixAuthStorageBridge } from "./fixtures/posix-auth-storage-bridge.mjs"; export const localPassword = "correct horse battery staple"; export const localPasswordHash = "$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4"; @@ -47,6 +48,11 @@ export function prepareAuthStateRoot(root: string): void { } } +/** Explicit test-only substitute for the production tht-backed POSIX storage bridge. */ +export function createFixtureAuthStorageBridge() { + return createFixturePosixAuthStorageBridge(); +} + /** Creates a production-local app and authenticates through the real login/session boundary. */ export async function createLocalAuthFixture( deps?: BuildAppDeps, @@ -82,11 +88,17 @@ export async function createLocalAuthFixture( const authStateRoot = join(directory, "auth-state"); prepareAuthStateRoot(authStateRoot); - const app = buildApp(loadConfig({ - THT_AUTH_CONFIG_FILE: authConfigFile, - THT_AUTH_STATE_ROOT: authStateRoot, - THT_HARNESS_DIR: "/tmp/h", - }), deps); + const app = buildApp( + loadConfig({ + THT_AUTH_CONFIG_FILE: authConfigFile, + THT_AUTH_STATE_ROOT: authStateRoot, + THT_HARNESS_DIR: "/tmp/h", + }), + { + ...deps, + authStorageBridgeForTest: deps?.authStorageBridgeForTest ?? createFixtureAuthStorageBridge(), + }, + ); let downstream = 0; app.addHook("preHandler", async () => { downstream += 1; }); diff --git a/backend/test/fixtures/oidc-state-consumer.mts b/backend/test/fixtures/oidc-state-consumer.mts index a196a0c8..f1eb0ab0 100644 --- a/backend/test/fixtures/oidc-state-consumer.mts +++ b/backend/test/fixtures/oidc-state-consumer.mts @@ -1,4 +1,5 @@ import { createFileAuthSessionStore } from "../../src/auth/session-store.js"; +import { createFixturePosixAuthStorageBridge } from "./posix-auth-storage-bridge.mjs"; const root = process.env.THT_TEST_SESSION_ROOT; const state = process.env.THT_TEST_OIDC_STATE; @@ -9,7 +10,9 @@ process.send("ready"); process.once("message", async (message) => { if (message !== "consume") process.exit(3); try { - const record = await createFileAuthSessionStore(root).consumeOidcState(state); + const record = await createFileAuthSessionStore(root, undefined, { + posixStorageBridge: createFixturePosixAuthStorageBridge(), + }).consumeOidcState(state); process.send?.({ consumed: record !== undefined }); process.exit(0); } catch { diff --git a/backend/test/fixtures/oidc-state-creator.mts b/backend/test/fixtures/oidc-state-creator.mts index 895cdfaf..ab3848bf 100644 --- a/backend/test/fixtures/oidc-state-creator.mts +++ b/backend/test/fixtures/oidc-state-creator.mts @@ -1,4 +1,5 @@ import { createFileAuthSessionStore } from "../../src/auth/session-store.js"; +import { createFixturePosixAuthStorageBridge } from "./posix-auth-storage-bridge.mjs"; const root = process.env.THT_TEST_SESSION_ROOT; const attempts = Number(process.env.THT_TEST_OIDC_ATTEMPTS); @@ -7,7 +8,10 @@ const capacity = Number(process.env.THT_TEST_OIDC_CAPACITY); if (!root || !Number.isSafeInteger(attempts) || attempts < 1 || Number.isNaN(now.getTime()) || !Number.isSafeInteger(capacity) || capacity < 1) process.exit(2); -const store = createFileAuthSessionStore(root, undefined, { oidcStateCapacity: capacity }); +const store = createFileAuthSessionStore(root, undefined, { + oidcStateCapacity: capacity, + posixStorageBridge: createFixturePosixAuthStorageBridge(), +}); process.send?.("ready"); process.once("message", async (message) => { if (message !== "create") process.exit(3); diff --git a/backend/test/fixtures/posix-auth-storage-bridge.mts b/backend/test/fixtures/posix-auth-storage-bridge.mts new file mode 100644 index 00000000..3dda1545 --- /dev/null +++ b/backend/test/fixtures/posix-auth-storage-bridge.mts @@ -0,0 +1,150 @@ +import { + chmodSync, + existsSync, + linkSync, + lstatSync, + mkdirSync, + opendirSync, + readFileSync, + renameSync, + unlinkSync, + writeFileSync, +} from "node:fs"; +import { join } from "node:path"; +import type { + WindowsAuthStorageBridge, + WindowsAuthStorageDirectory, +} from "../../src/auth/windows-auth-storage.js"; + +// Test-process fixture only. The production POSIX implementation is the Go auth-storage bridge; +// the Go package covers retained-descriptor adversarial races separately. +export function createFixturePosixAuthStorageBridge(): WindowsAuthStorageBridge { + const ensure = async (root: string): Promise => { + if (!existsSync(root)) { + mkdirSync(root, { mode: 0o700 }); + chmodSync(root, 0o700); + } + for (const name of ["sessions", "oidc"]) { + const path = join(root, name); + if (!existsSync(path)) { + mkdirSync(path, { mode: 0o700 }); + chmodSync(path, 0o700); + } + } + }; + const path = async (root: string, directory: WindowsAuthStorageDirectory, filename: string): Promise => { + await ensure(root); + return join(root, directory, filename); + }; + const names = async (root: string, directory: WindowsAuthStorageDirectory): Promise => { + await ensure(root); + const handle = opendirSync(join(root, directory)); + const result: string[] = []; + try { + for (;;) { + const entry = handle.readSync(); + if (entry === null) break; + result.push(entry.name); + } + } finally { + handle.closeSync(); + } + return result.sort(); + }; + const missing = (error: unknown): boolean => (error as { code?: unknown })?.code === "ENOENT"; + return { + validateRoot: ensure, + ensureLayout: ensure, + readAuthConfig: (value) => readFileSync(value), + readLocalUsers: async (value) => readFileSync(value), + create: async (root, directory, filename, contents) => { + try { + const value = await path(root, directory, filename); + writeFileSync(value, contents, { flag: "wx", mode: 0o600 }); + chmodSync(value, 0o600); + return true; + } catch (error) { + if (missing(error) || (error as { code?: unknown })?.code === "EEXIST") return false; + throw error; + } + }, + read: async (root, directory, filename) => { + try { return readFileSync(await path(root, directory, filename)); } catch (error) { + if (missing(error)) return undefined; + throw error; + } + }, + replace: async (root, directory, filename, contents) => { + const value = await path(root, directory, filename); + const temporary = `${value}.fixture-replacement`; + writeFileSync(temporary, contents, { flag: "wx", mode: 0o600 }); + renameSync(temporary, value); + }, + remove: async (root, directory, filename) => { + try { + unlinkSync(await path(root, directory, filename)); + return true; + } catch (error) { + if (missing(error)) return false; + throw error; + } + }, + list: async (root, directory, maximumEntries = 256) => { + const result = await names(root, directory); + if (result.length > maximumEntries) throw new Error("fixture list overflow"); + return result.map((name) => ({ name, modifiedUnixMs: lstatSync(join(root, directory, name)).mtimeMs })); + }, + listPage: async (root, directory, afterName, maximumEntries) => { + if (directory !== "sessions") throw new Error("fixture directory invalid"); + const selected = (await names(root, directory)).filter((name) => afterName === undefined || name > afterName); + return { + entries: selected.slice(0, maximumEntries).map((name) => ({ + name, modifiedUnixMs: lstatSync(join(root, directory, name)).mtimeMs, + })), + more: selected.length > maximumEntries, + }; + }, + claimConsume: async (root, filename) => { + const source = await path(root, "oidc", filename); + const claim = source.replace(/\.json$/, ".claim"); + try { + linkSync(source, claim); + } catch (error) { + if (missing(error) || (error as { code?: unknown })?.code === "EEXIST") return undefined; + throw error; + } + const contents = readFileSync(source); + unlinkSync(source); + unlinkSync(claim); + return contents; + }, + readClaim: async (root, filename) => { + const source = await path(root, "oidc", filename); + const claim = source.replace(/\.json$/, ".claim"); + try { + const left = lstatSync(source); + const right = lstatSync(claim); + if (left.ino !== right.ino || left.dev !== right.dev || left.nlink !== 2 || right.nlink !== 2) return undefined; + return readFileSync(source); + } catch (error) { + if (missing(error)) return undefined; + throw error; + } + }, + removeClaim: async (root, filename) => { + const source = await path(root, "oidc", filename); + const claim = source.replace(/\.json$/, ".claim"); + try { + const left = lstatSync(source); + const right = lstatSync(claim); + if (left.ino !== right.ino || left.dev !== right.dev || left.nlink !== 2 || right.nlink !== 2) return false; + unlinkSync(source); + unlinkSync(claim); + return true; + } catch (error) { + if (missing(error)) return false; + throw error; + } + }, + }; +} diff --git a/backend/test/fixtures/windows-auth-storage-real-child.mjs b/backend/test/fixtures/windows-auth-storage-real-child.mjs index 394666f8..0fe8adc0 100644 --- a/backend/test/fixtures/windows-auth-storage-real-child.mjs +++ b/backend/test/fixtures/windows-auth-storage-real-child.mjs @@ -6,7 +6,7 @@ writeFileSync(marker, `started:${process.pid}\n`); process.on("exit", () => appendFileSync(marker, "exited\n")); process.on("SIGTERM", () => { appendFileSync(marker, "terminated\n"); - process.exit(0); + if (mode !== "timeout") process.exit(0); }); if (mode === "stdin") { diff --git a/backend/test/local-registry.test.ts b/backend/test/local-registry.test.ts index d1c74c6e..649fc099 100644 --- a/backend/test/local-registry.test.ts +++ b/backend/test/local-registry.test.ts @@ -14,7 +14,7 @@ import { import { mkdtempSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; -import { afterEach, describe, expect, test } from "vitest"; +import { afterEach, describe, expect, test, vi } from "vitest"; import { createLocalUserRegistry } from "../src/auth/local-registry.js"; const password = "correct horse battery staple"; @@ -204,4 +204,28 @@ describe("local user registry", () => { await expect(registry.findByUsername("admin")).resolves.toMatchObject({ displayName: "Owner" }); }); + + test("routes native Windows users.yaml loading only through the bounded auth-storage bridge", async () => { + const usersPath = "C:\\ProgramData\\ThothII\\auth\\users.yaml"; + const readLocalUsers = vi.fn(async (path: string) => { + expect(path).toBe(usersPath); + return Buffer.from(registryYaml(userYaml({ displayName: "Bridge administrator" })), "utf8"); + }); + const originalPlatform = Object.getOwnPropertyDescriptor(process, "platform"); + if (!originalPlatform) throw new Error("platform descriptor unavailable"); + Object.defineProperty(process, "platform", { configurable: true, value: "win32" }); + try { + const registry = createLocalUserRegistry(usersPath, { windowsStorageBridge: { readLocalUsers } } as never); + await expect(registry.findByUsername("ADMIN")).resolves.toMatchObject({ + id: adminId, + displayName: "Bridge administrator", + }); + await expect(registry.hasEnabledAdmin()).resolves.toBe(true); + // Windows reloads from the bridge on every registry observation so an atomic host + // replacement cannot be missed between authorization checks. + expect(readLocalUsers).toHaveBeenCalledTimes(2); + } finally { + Object.defineProperty(process, "platform", originalPlatform); + } + }); }); diff --git a/backend/test/windows-auth-storage.test.ts b/backend/test/windows-auth-storage.test.ts index a08be585..fcd4546e 100644 --- a/backend/test/windows-auth-storage.test.ts +++ b/backend/test/windows-auth-storage.test.ts @@ -12,6 +12,7 @@ import { } from "../src/auth/windows-auth-storage.js"; const root = "C:\\ProgramData\\ThothII\\auth"; +const posixRoot = "/var/lib/thothii/auth"; const filename = "a".repeat(64) + ".json"; const realChildFixture = fileURLToPath(new URL("./fixtures/windows-auth-storage-real-child.mjs", import.meta.url)); const fixtureRoots: string[] = []; @@ -29,18 +30,25 @@ async function waitForMarker(marker: string, expected: string): Promise { throw new Error(`real helper marker did not contain ${expected}`); } -function realChildBridge(mode: "timeout" | "stdout" | "stderr" | "stdin") { +function realChildBridge( + mode: "timeout" | "stdout" | "stderr" | "stdin", + pathStyle: "windows" | "posix", +) { const directory = mkdtempSync(join(tmpdir(), "thothii-auth-bridge-child-")); fixtureRoots.push(directory); const marker = join(directory, "marker.txt"); const launcher = join(directory, "tht.exe"); writeFileSync(launcher, `#!/bin/sh\nexec ${shellQuote(process.execPath)} ${shellQuote(realChildFixture)} ${shellQuote(mode)} ${shellQuote(marker)} "$@"\n`, { mode: 0o700 }); chmodSync(launcher, 0o700); + const factory = pathStyle === "windows" ? createWindowsAuthStorageBridge : createPosixAuthStorageBridge; return { marker, - bridge: createWindowsAuthStorageBridge({ - thtExecutable: "C:\\tht.exe", + bridge: factory({ + thtExecutable: pathStyle === "windows" ? "C:\\tht.exe" : launcher, spawnChild: (_executable, args, options) => spawn(launcher, [...args], options), + // Leave enough startup headroom for a real child under a busy CI host while retaining a + // sub-1.5-second bound from request start through final settlement. + deadlinesForTest: { timeoutMs: 750, terminationGraceMs: 50, finalSettlementMs: 500 }, ...(mode === "stdin" ? { beforeInputForTest: async () => { await waitForMarker(marker, "stdin-closed"); @@ -67,16 +75,18 @@ class FakeBridgeChild extends EventEmitter { readonly stdout = new PassThrough(); readonly stderr = new PassThrough(); readonly kill = vi.fn(() => true); + readonly unref = vi.fn(); close(code = 0, signal: NodeJS.Signals | null = null): void { this.emit("close", code, signal); } } -function bridgeForChild(child: FakeBridgeChild) { +function bridgeForChild(child: FakeBridgeChild, pathStyle: "windows" | "posix" = "windows") { const spawnChild = vi.fn(() => child); - const bridge = createWindowsAuthStorageBridge({ - thtExecutable: "C:\\tht.exe", + const factory = pathStyle === "windows" ? createWindowsAuthStorageBridge : createPosixAuthStorageBridge; + const bridge = factory({ + thtExecutable: pathStyle === "windows" ? "C:\\tht.exe" : "/opt/thothii/bin/tht", spawnChild, } as never); return { bridge, spawnChild }; @@ -214,6 +224,34 @@ describe("Windows auth-storage bridge", () => { expect(JSON.stringify(syncCalls[0]!.args)).not.toContain(config.toString("utf8")); }); + test("reads native Windows users.yaml only through a bounded hidden bridge request", async () => { + const users = Buffer.from("version: 1\nusers:\n - passwordHash: not-in-argv\n", "utf8"); + const calls: Array<{ args: readonly string[]; input: Buffer; maximumOutputBytes: number }> = []; + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\Program Files\\ThothII\\tht.exe", + invoke: async (call) => { + calls.push(call); + return { + code: 0, + stdout: Buffer.from(`${JSON.stringify({ + version: 1, ok: true, found: true, contentBase64: users.toString("base64"), + })}\n`), + stderr: Buffer.alloc(0), + }; + }, + }); + + await expect(bridge.readLocalUsers(`${root}\\users.yaml`)).resolves.toEqual(users); + expect(calls).toHaveLength(1); + expect(calls[0]!.args).toEqual(["_auth-storage"]); + expect(calls[0]!.maximumOutputBytes).toBeGreaterThan(1024 * 1024); + expect(JSON.parse(calls[0]!.input.toString("utf8"))).toEqual({ + version: 1, operation: "read-local-users", root, filename: "users.yaml", + }); + expect(JSON.stringify(calls[0]!.args)).not.toContain("not-in-argv"); + expect(calls[0]!.input.toString("utf8")).not.toContain("not-in-argv"); + }); + test.each([ { label: "nonzero", result: { code: 1, stdout: Buffer.from('{"version":1,"ok":true}\n'), stderr: Buffer.from("secret") } }, { label: "malformed stdout", result: { code: 0, stdout: Buffer.from("not-json"), stderr: Buffer.alloc(0) } }, @@ -400,7 +438,7 @@ describe("Windows auth-storage bridge", () => { await expect(bridge.list(root, "sessions")).rejects.toThrow("auth_session_store_invalid"); }); - test("aborts a stdin-closed looping helper on timeout and waits for close", async () => { + test("settles a stdin-closed looping helper by a final deadline when close never arrives", async () => { vi.useFakeTimers(); const child = new FakeBridgeChild(); const { bridge, spawnChild } = bridgeForChild(child); @@ -410,17 +448,44 @@ describe("Windows auth-storage bridge", () => { await vi.advanceTimersByTimeAsync(5_000); expect(spawnChild).toHaveBeenCalledOnce(); expect(child.kill).toHaveBeenCalledOnce(); + expect(child.unref).toHaveBeenCalledOnce(); expect(child.stdin.destroyed).toBe(true); await expect(Promise.race([outcome, Promise.resolve("pending")])).resolves.toBe("pending"); - child.close(); + await vi.advanceTimersByTimeAsync(1_000); + expect(child.kill).toHaveBeenCalledTimes(2); await expect(outcome).resolves.toBe("rejected"); + expect(() => child.emit("error", new Error("late helper failure"))).not.toThrow(); + expect(() => child.stdin.emit("error", new Error("late stdin failure"))).not.toThrow(); } finally { child.close(); vi.useRealTimers(); } }); + test("releases bounded late-error guards when a finally-dead helper closes", async () => { + vi.useFakeTimers(); + const child = new FakeBridgeChild(); + const { bridge } = bridgeForChild(child); + const outcome = bridge.list(root, "sessions").then(() => "resolved", () => "rejected"); + try { + await vi.advanceTimersByTimeAsync(6_000); + await expect(outcome).resolves.toBe("rejected"); + expect(child.listenerCount("error")).toBe(1); + expect(child.stdin.listenerCount("error")).toBe(1); + expect(child.stdout.listenerCount("error")).toBe(1); + expect(child.stderr.listenerCount("error")).toBe(1); + + child.close(); + expect(child.listenerCount("error")).toBe(0); + expect(child.stdin.listenerCount("error")).toBe(0); + expect(child.stdout.listenerCount("error")).toBe(0); + expect(child.stderr.listenerCount("error")).toBe(0); + } finally { + vi.useRealTimers(); + } + }); + test.each(["stdout", "stderr"] as const)("aborts a %s-flooding helper and waits for close", async (stream) => { const child = new FakeBridgeChild(); const { bridge, spawnChild } = bridgeForChild(child); @@ -469,6 +534,10 @@ describe("Windows auth-storage bridge", () => { await expect(Promise.race([outcome, Promise.resolve("pending")])).resolves.toBe("pending"); child.close(); await expect(outcome).resolves.toBe("rejected"); + expect(() => child.emit("error", new Error("late helper failure"))).not.toThrow(); + expect(() => child.stdin.emit("error", new Error("late stdin failure"))).not.toThrow(); + expect(() => child.stdout.emit("error", new Error("late stdout failure"))).not.toThrow(); + expect(() => child.stderr.emit("error", new Error("late stderr failure"))).not.toThrow(); }); test("fails closed when launching the helper throws before a child exists", async () => { @@ -480,16 +549,19 @@ describe("Windows auth-storage bridge", () => { await expect(bridge.list(root, "sessions")).rejects.toThrow("auth_session_store_invalid"); }); - test.each(["timeout", "stdout", "stderr", "stdin"] as const)("kills a real %s helper process through the production spawn path", async (mode) => { - const { bridge, marker } = realChildBridge(mode); + test.each([ + ...(["windows", "posix"] as const).flatMap((pathStyle) => + (["timeout", "stdout", "stderr", "stdin"] as const).map((mode) => ({ pathStyle, mode }))), + ])("settles a real $pathStyle $mode helper within the production deadline", async ({ pathStyle, mode }) => { + const { bridge, marker } = realChildBridge(mode, pathStyle); const startedAt = Date.now(); - const pending = bridge.list(root, "sessions"); + const pending = bridge.list(pathStyle === "windows" ? root : posixRoot, "sessions"); const outcome = pending.then(() => undefined, (error: unknown) => error); await waitForMarker(marker, "started"); if (mode === "stdin") await waitForMarker(marker, "before-input"); await expect(outcome).resolves.toMatchObject({ message: "auth_session_store_invalid" }); await waitForMarker(marker, "terminated"); - if (mode === "stdin") expect(Date.now() - startedAt).toBeLessThan(2_000); - }, 10_000); + expect(Date.now() - startedAt).toBeLessThan(1_500); + }, 5_000); }); diff --git a/tools/tht/internal/authstorage/storage.go b/tools/tht/internal/authstorage/storage.go index 77f1a9c7..85ec4a47 100644 --- a/tools/tht/internal/authstorage/storage.go +++ b/tools/tht/internal/authstorage/storage.go @@ -9,8 +9,6 @@ import ( "encoding/json" "errors" "io" - "os" - "path/filepath" "regexp" "strings" "unicode" @@ -123,21 +121,21 @@ func execute(input request) (response, error) { return response{Version: protocolVersion, OK: true, Prepared: true}, nil } if input.Operation == "read-auth-config" { - root, err := existingPrivateRoot(input.Root) - if err != nil { - return response{}, errInvalid - } - contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(root, input.Filename), maximumAuthConfigBytes) - if err != nil { - return response{}, errInvalid - } - return contentResponse(true, contents), nil + return readRootPrivateRegular(input.Root, input.Filename) + } + if input.Operation == "read-local-users" { + return readRootPrivateRegular(input.Root, input.Filename) } if !validDirectory(input.Directory) { return response{}, errInvalid } - directory, err := storageDirectory(input.Root, input.Directory) - if err != nil { + layout, err := openStorageLayout(input.Root, true) + if err != nil || layout == nil { + return response{}, errInvalid + } + defer layout.Close() + directory := layout.directory(input.Directory) + if directory == nil { return response{}, errInvalid } switch input.Operation { @@ -162,7 +160,7 @@ func execute(input request) (response, error) { if err != nil { return response{}, errInvalid } - if err := safeio.ReplaceCanonicalRegular(filepath.Join(directory, input.Filename), contents, 0o600); err != nil { + if err := directory.ReplaceRegular(input.Filename, contents); err != nil { return response{}, errInvalid } return response{Version: protocolVersion, OK: true, Replaced: true}, nil @@ -177,46 +175,32 @@ func execute(input request) (response, error) { if limit == 0 { limit = defaultMaximumEntries } + afterName := "" if input.Continuation { - page, err := safeio.ListCanonicalPrivateDirectoryPage( - directory, - limit, - input.AfterName, - digestFilename.MatchString, - ) - if err != nil { - return response{}, errInvalid - } - more := page.More - return response{Version: protocolVersion, OK: true, Entries: &page.Entries, More: &more}, nil + afterName = input.AfterName } - entries, err := safeio.ListCanonicalPrivateDirectory(directory, limit) + page, err := directory.ListPage(limit, afterName, recordListName(input.Directory), recordListLinks(input.Directory)) if err != nil { return response{}, errInvalid } - for _, entry := range entries { - if !digestFilename.MatchString(entry.Name) && !(input.Directory == "oidc" && (claimFilename.MatchString(entry.Name) || oidcSlotFilename.MatchString(entry.Name))) { - return response{}, errInvalid - } + if input.Continuation { + more := page.More + return response{Version: protocolVersion, OK: true, Entries: &page.Entries, More: &more}, nil } - return response{Version: protocolVersion, OK: true, Entries: &entries}, nil + if page.More { + return response{}, errInvalid + } + return response{Version: protocolVersion, OK: true, Entries: &page.Entries}, nil case "claim-consume": return claimConsume(directory, input.Filename) case "read-claim": - contents, found, err := safeio.ReadCanonicalPrivateClaim( - filepath.Join(directory, input.Filename), - filepath.Join(directory, asClaimFilename(input.Filename)), - maximumOIDCStateBytes, - ) + contents, found, err := directory.ReadClaim(input.Filename, asClaimFilename(input.Filename), maximumOIDCStateBytes) if err != nil { return response{}, errInvalid } return contentResponse(found, contents), nil case "remove-claim": - removed, err := safeio.RemoveCanonicalPrivateClaim( - filepath.Join(directory, input.Filename), - filepath.Join(directory, asClaimFilename(input.Filename)), - ) + removed, err := directory.RemoveClaim(input.Filename, asClaimFilename(input.Filename)) if err != nil { return response{}, errInvalid } @@ -234,7 +218,7 @@ func validOperationShape(input request) bool { switch input.Operation { case "validate-root", "ensure-layout": return input.Directory == "" && input.Filename == "" && noContents && noMaximumEntries && noAfterName && noContinuation - case "read-auth-config": + case "read-auth-config", "read-local-users": return input.Directory == "" && authConfigFilename.MatchString(input.Filename) && noContents && noMaximumEntries && noAfterName && noContinuation case "create", "replace": return noMaximumEntries && noAfterName && noContinuation && (digestFilename.MatchString(input.Filename) || (input.Operation == "create" && input.Directory == "oidc" && oidcSlotFilename.MatchString(input.Filename))) @@ -252,78 +236,170 @@ func validOperationShape(input request) bool { } } -func preflightRoot(root string) (bool, error) { - if !filepath.IsAbs(root) || filepath.Clean(root) != root || strings.IndexFunc(root, unicode.IsControl) >= 0 { - return false, errInvalid - } - return safeio.PreflightPrivateDirectory(root) +type storageLayout struct { + root safeio.PrivateDirectoryHandle + sessions safeio.PrivateDirectoryHandle + oidc safeio.PrivateDirectoryHandle } -func existingPrivateRoot(root string) (string, error) { - exists, err := preflightRoot(root) - if err != nil || !exists || safeio.ValidatePrivateDirectory(root) != nil { - return "", errInvalid +func (layout *storageLayout) Close() { + if layout == nil { + return } - return root, nil + if layout.oidc != nil { + _ = layout.oidc.Close() + layout.oidc = nil + } + if layout.sessions != nil { + _ = layout.sessions.Close() + layout.sessions = nil + } + if layout.root != nil { + _ = layout.root.Close() + layout.root = nil + } +} + +func (layout *storageLayout) closeChildren() { + if layout == nil { + return + } + if layout.oidc != nil { + _ = layout.oidc.Close() + layout.oidc = nil + } + if layout.sessions != nil { + _ = layout.sessions.Close() + layout.sessions = nil + } +} + +func (layout *storageLayout) directory(name string) safeio.PrivateDirectoryHandle { + if layout == nil { + return nil + } + if name == "sessions" { + return layout.sessions + } + if name == "oidc" { + return layout.oidc + } + return nil +} + +func validStorageRoot(root string) bool { + return strings.IndexFunc(root, unicode.IsControl) < 0 +} + +// openStorageLayout keeps the root descriptor/handle open from its initial canonical validation +// through every child observation. The side-effect-free path performs a second child pass after +// the testable race boundary and retains those final handles; it never validates an earlier, +// discarded exists result. +func openStorageLayout(root string, ensure bool) (*storageLayout, error) { + if !validStorageRoot(root) { + return nil, errInvalid + } + rootHandle, found, err := safeio.OpenPrivateDirectory(root, ensure) + if err != nil { + return nil, errInvalid + } + if !found { + return nil, nil + } + layout := &storageLayout{root: rootHandle} + failed := true + defer func() { + if failed { + layout.Close() + } + }() + safeio.NotifyPrivateDirectoryTestHookForTest("after-auth-root-open") + if err := layout.openChildren(ensure); err != nil { + return nil, errInvalid + } + safeio.NotifyPrivateDirectoryTestHookForTest("after-auth-layout-first-pass") + if !ensure { + layout.closeChildren() + if err := layout.openChildren(false); err != nil { + return nil, errInvalid + } + } + if layout.root.Validate() != nil || (layout.sessions != nil && layout.sessions.Validate() != nil) || + (layout.oidc != nil && layout.oidc.Validate() != nil) { + return nil, errInvalid + } + failed = false + return layout, nil +} + +func (layout *storageLayout) openChildren(ensure bool) error { + if layout == nil || layout.root == nil || layout.root.Validate() != nil { + return errInvalid + } + for _, name := range []string{"sessions", "oidc"} { + child, found, err := layout.root.OpenChild(name, ensure) + if err != nil || (ensure && !found) { + if child != nil { + _ = child.Close() + } + return errInvalid + } + if !found { + continue + } + if child.Validate() != nil { + _ = child.Close() + return errInvalid + } + if name == "sessions" { + layout.sessions = child + } else { + layout.oidc = child + } + } + if layout.root.Validate() != nil { + return errInvalid + } + return nil } func validateStorageLayout(root string) error { - exists, err := preflightRoot(root) + layout, err := openStorageLayout(root, false) if err != nil { return errInvalid } - if !exists { - return nil - } - existingChildren := make([]string, 0, 2) - for _, directory := range []string{"sessions", "oidc"} { - path := filepath.Join(root, directory) - if filepath.Dir(path) != root { - return errInvalid - } - childExists, err := safeio.PreflightPrivateDirectory(path) - if err != nil { - return errInvalid - } - if childExists { - existingChildren = append(existingChildren, path) - } - } - // Close permission/identity races between the individual side-effect-free preflights. - if safeio.ValidatePrivateDirectory(root) != nil { - return errInvalid - } - for _, directory := range []string{"sessions", "oidc"} { - if _, err := safeio.PreflightPrivateDirectory(filepath.Join(root, directory)); err != nil { - return errInvalid - } - } - for _, path := range existingChildren { - if safeio.ValidatePrivateDirectory(path) != nil { - return errInvalid - } + if layout != nil { + layout.Close() } return nil } func ensureStorageLayout(root string) error { - if _, err := preflightRoot(root); err != nil || safeio.EnsurePrivateDirectory(root) != nil { - return errInvalid - } - for _, directory := range []string{"sessions", "oidc"} { - path := filepath.Join(root, directory) - if filepath.Dir(path) != root || safeio.EnsurePrivateDirectory(path) != nil { - return errInvalid - } - } - if safeio.ValidatePrivateDirectory(root) != nil || - safeio.ValidatePrivateDirectory(filepath.Join(root, "sessions")) != nil || - safeio.ValidatePrivateDirectory(filepath.Join(root, "oidc")) != nil { + layout, err := openStorageLayout(root, true) + if err != nil || layout == nil { return errInvalid } + layout.Close() return nil } +func readRootPrivateRegular(root, filename string) (response, error) { + if !validStorageRoot(root) { + return response{}, errInvalid + } + directory, found, err := safeio.OpenPrivateDirectory(root, false) + if err != nil || !found || directory == nil { + return response{}, errInvalid + } + defer directory.Close() + safeio.NotifyPrivateDirectoryTestHookForTest("after-auth-root-open") + contents, found, err := directory.ReadRegular(filename, maximumAuthConfigBytes) + if err != nil || !found || directory.Validate() != nil { + return response{}, errInvalid + } + return contentResponse(true, contents), nil +} + func contentResponse(found bool, contents []byte) response { if !found { return response{Version: protocolVersion, OK: true} @@ -331,17 +407,6 @@ func contentResponse(found bool, contents []byte) response { return response{Version: protocolVersion, OK: true, Found: true, ContentBase64: base64.StdEncoding.EncodeToString(contents)} } -func storageDirectory(root, directory string) (string, error) { - if _, err := preflightRoot(root); err != nil || safeio.EnsurePrivateDirectory(root) != nil { - return "", errInvalid - } - path := filepath.Join(root, directory) - if filepath.Dir(path) != root || safeio.EnsurePrivateDirectory(path) != nil { - return "", errInvalid - } - return path, nil -} - func validDirectory(value string) bool { return value == "sessions" || value == "oidc" } @@ -361,67 +426,67 @@ func decodeContents(input request) ([]byte, error) { return contents, nil } -func createPrivate(directory, filename string, contents []byte) (bool, error) { - path := filepath.Join(directory, filename) - if _, err := os.Lstat(path); err == nil { - if safeio.ValidatePrivateRegular(path) != nil { - return false, errInvalid - } - return false, nil - } else if !errors.Is(err, os.ErrNotExist) { +func createPrivate(directory safeio.PrivateDirectoryHandle, filename string, contents []byte) (bool, error) { + created, err := directory.CreateRegular(filename, contents) + if err != nil { return false, errInvalid } - if err := safeio.WriteCanonicalNewPrivateFile(path, contents, 0o600); err == nil { - return true, nil - } - if safeio.ValidatePrivateRegular(path) == nil { - return false, nil - } - return false, errInvalid + return created, nil } -func readPrivate(directory, filename string, maximum int64) ([]byte, bool, error) { - path := filepath.Join(directory, filename) - if _, err := os.Lstat(path); errors.Is(err, os.ErrNotExist) { - return nil, false, nil - } else if err != nil { - return nil, false, errInvalid - } - contents, err := safeio.ReadCanonicalPrivateRegular(path, maximum) +func readPrivate(directory safeio.PrivateDirectoryHandle, filename string, maximum int64) ([]byte, bool, error) { + contents, found, err := directory.ReadRegular(filename, maximum) if err != nil { return nil, false, errInvalid } - return contents, true, nil + return contents, found, nil } -func removePrivate(directory, filename string) (bool, error) { - path := filepath.Join(directory, filename) - if _, err := os.Lstat(path); errors.Is(err, os.ErrNotExist) { - return false, nil - } else if err != nil { +func removePrivate(directory safeio.PrivateDirectoryHandle, filename string) (bool, error) { + removed, err := directory.RemoveRegular(filename) + if err != nil { return false, errInvalid } - if err := safeio.RemoveCanonicalPrivateRegular(path); err != nil { - return false, errInvalid - } - return true, nil + return removed, nil } -func claimConsume(directory, filename string) (response, error) { - source := filepath.Join(directory, filename) - claim := filepath.Join(directory, asClaimFilename(filename)) - claimed, err := safeio.ClaimCanonicalPrivateRegular(source, claim) +func recordListName(directory string) func(string) bool { + if directory == "sessions" { + return digestFilename.MatchString + } + return func(name string) bool { + return digestFilename.MatchString(name) || claimFilename.MatchString(name) || oidcSlotFilename.MatchString(name) + } +} + +func recordListLinks(directory string) func(string, uint64) bool { + if directory == "sessions" { + return func(name string, links uint64) bool { + return digestFilename.MatchString(name) && links == 1 + } + } + return func(name string, links uint64) bool { + if digestFilename.MatchString(name) || claimFilename.MatchString(name) { + return links == 1 || links == 2 + } + return oidcSlotFilename.MatchString(name) && links == 1 + } +} + +func claimConsume(directory safeio.PrivateDirectoryHandle, filename string) (response, error) { + claim := asClaimFilename(filename) + claimed, err := directory.ClaimRegular(filename, claim) if err != nil { return response{}, errInvalid } if !claimed { return response{Version: protocolVersion, OK: true}, nil } - contents, found, err := safeio.ReadCanonicalPrivateClaim(source, claim, maximumOIDCStateBytes) + contents, found, err := directory.ReadClaim(filename, claim, maximumOIDCStateBytes) if err != nil || !found { return response{}, errInvalid } - removed, err := safeio.RemoveCanonicalPrivateClaim(source, claim) + removed, err := directory.RemoveClaim(filename, claim) if err != nil || !removed { return response{}, errInvalid } diff --git a/tools/tht/internal/authstorage/storage_test.go b/tools/tht/internal/authstorage/storage_test.go index 0181a832..4017e20b 100644 --- a/tools/tht/internal/authstorage/storage_test.go +++ b/tools/tht/internal/authstorage/storage_test.go @@ -167,6 +167,369 @@ func TestProtocolEnsuresTheCompletePrivateSessionLayout(t *testing.T) { runRejected(t, request{Version: 1, Operation: "ensure-layout", Root: root, Directory: "sessions"}) } +func TestProtocolPinsTheOriginalPrivateRootBeforeEveryRecordMutation(t *testing.T) { + parent := privateTestRoot(t) + root := filepath.Join(parent, "auth") + replacement := filepath.Join(parent, "replacement") + moved := filepath.Join(parent, "auth-original") + filename := "eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee.json" + for _, directory := range []string{root, replacement} { + if err := safeio.EnsurePrivateDirectory(directory); err != nil { + t.Fatal(err) + } + for _, child := range []string{"sessions", "oidc"} { + if err := safeio.EnsurePrivateDirectory(filepath.Join(directory, child)); err != nil { + t.Fatal(err) + } + } + } + + swapped := false + blocked := false + restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) { + if stage != "after-auth-root-open" || swapped || blocked { + return + } + if runtime.GOOS == "windows" { + if err := os.Rename(root, moved); err == nil { + t.Fatal("retained Windows root handle permitted rename") + } + blocked = true + return + } + if err := os.Rename(root, moved); err != nil { + t.Fatal(err) + } + if err := os.Rename(replacement, root); err != nil { + t.Fatal(err) + } + swapped = true + }) + t.Cleanup(restore) + + runRequest(t, request{ + Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, + ContentBase64: base64.StdEncoding.EncodeToString([]byte("pinned-root-record")), + }) + if !swapped && !blocked { + t.Fatal("record operation did not expose the retained-root test hook") + } + if runtime.GOOS == "windows" { + if _, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(root, "sessions", filename), maximumSessionBytes); err != nil { + t.Fatalf("record missing from retained Windows root: %v", err) + } + return + } + contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(moved, "sessions", filename), maximumSessionBytes) + if err != nil || string(contents) != "pinned-root-record" { + t.Fatalf("pinned root contents = %q error = %v", contents, err) + } + if _, err := os.Lstat(filepath.Join(root, "sessions", filename)); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("replacement root was mutated: %v", err) + } +} + +func TestProtocolPinsTheOriginalPrivateRootForEveryRecordOperation(t *testing.T) { + filename := "eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee.json" + otherFilename := "ffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff.json" + for _, operation := range []string{ + "create", "read", "replace", "remove", "list", "claim-consume", "read-claim", "remove-claim", + } { + t.Run(operation, func(t *testing.T) { + parent := privateTestRoot(t) + root := filepath.Join(parent, "auth") + replacement := filepath.Join(parent, "replacement") + moved := filepath.Join(parent, "auth-original") + for _, directory := range []string{root, replacement} { + if err := safeio.EnsurePrivateDirectory(directory); err != nil { + t.Fatal(err) + } + for _, child := range []string{"sessions", "oidc"} { + if err := safeio.EnsurePrivateDirectory(filepath.Join(directory, child)); err != nil { + t.Fatal(err) + } + } + } + write := func(base, directory, name, contents string) { + t.Helper() + if err := safeio.WriteCanonicalNewPrivateFile(filepath.Join(base, directory, name), []byte(contents), 0o600); err != nil { + t.Fatal(err) + } + } + pair := func(base, name, contents string) { + t.Helper() + write(base, "oidc", name, contents) + if err := os.Link(filepath.Join(base, "oidc", name), filepath.Join(base, "oidc", asClaimFilename(name))); err != nil { + t.Fatal(err) + } + } + input := request{Version: 1, Operation: operation, Root: root} + switch operation { + case "create": + input.Directory, input.Filename = "sessions", filename + input.ContentBase64 = base64.StdEncoding.EncodeToString([]byte("original-created")) + case "read": + write(root, "sessions", filename, "original-read") + write(replacement, "sessions", filename, "replacement-read") + input.Directory, input.Filename = "sessions", filename + case "replace": + write(root, "sessions", filename, "original-before") + write(replacement, "sessions", filename, "replacement-before") + input.Directory, input.Filename = "sessions", filename + input.ContentBase64 = base64.StdEncoding.EncodeToString([]byte("original-replaced")) + case "remove": + write(root, "sessions", filename, "original-remove") + write(replacement, "sessions", filename, "replacement-remove") + input.Directory, input.Filename = "sessions", filename + case "list": + write(root, "sessions", filename, "original-list") + write(replacement, "sessions", otherFilename, "replacement-list") + input.Directory = "sessions" + case "claim-consume": + write(root, "oidc", filename, "original-claim") + write(replacement, "oidc", filename, "replacement-claim") + input.Directory, input.Filename = "oidc", filename + case "read-claim", "remove-claim": + pair(root, filename, "original-pair") + pair(replacement, filename, "replacement-pair") + input.Directory, input.Filename = "oidc", filename + } + + swapped := false + blocked := false + restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) { + if stage != "after-auth-root-open" || swapped || blocked { + return + } + if runtime.GOOS == "windows" { + if err := os.Rename(root, moved); err == nil { + t.Fatal("retained Windows root handle permitted rename") + } + blocked = true + return + } + if err := os.Rename(root, moved); err != nil { + t.Fatal(err) + } + if err := os.Rename(replacement, root); err != nil { + t.Fatal(err) + } + swapped = true + }) + t.Cleanup(restore) + + output := runRequest(t, input) + if !swapped && !blocked { + t.Fatal("record operation did not expose the retained-root test hook") + } + original := root + replacementRoot := replacement + if swapped { + original = moved + replacementRoot = root + } + switch operation { + case "create": + contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(original, "sessions", filename), maximumSessionBytes) + if err != nil || string(contents) != "original-created" { + t.Fatalf("pinned create contents = %q error = %v", contents, err) + } + if _, err := os.Lstat(filepath.Join(replacementRoot, "sessions", filename)); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("replacement root was mutated: %v", err) + } + case "read": + if !output.Found || decodeContent(t, output) != "original-read" { + t.Fatalf("pinned read = %#v", output) + } + case "replace": + contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(original, "sessions", filename), maximumSessionBytes) + if err != nil || string(contents) != "original-replaced" { + t.Fatalf("pinned replace contents = %q error = %v", contents, err) + } + replacementContents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(replacementRoot, "sessions", filename), maximumSessionBytes) + if err != nil || string(replacementContents) != "replacement-before" { + t.Fatalf("replacement contents = %q error = %v", replacementContents, err) + } + case "remove": + if _, err := os.Lstat(filepath.Join(original, "sessions", filename)); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("pinned remove left original record: %v", err) + } + if _, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(replacementRoot, "sessions", filename), maximumSessionBytes); err != nil { + t.Fatalf("replacement record was mutated: %v", err) + } + case "list": + if output.Entries == nil || len(*output.Entries) != 1 || (*output.Entries)[0].Name != filename { + t.Fatalf("pinned list = %#v", output) + } + case "claim-consume": + if !output.Found || decodeContent(t, output) != "original-claim" { + t.Fatalf("pinned claim consume = %#v", output) + } + if _, err := os.Lstat(filepath.Join(original, "oidc", filename)); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("pinned claim source remains: %v", err) + } + if contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(replacementRoot, "oidc", filename), maximumOIDCStateBytes); err != nil || string(contents) != "replacement-claim" { + t.Fatalf("replacement claim contents = %q error = %v", contents, err) + } + case "read-claim": + if !output.Found || decodeContent(t, output) != "original-pair" { + t.Fatalf("pinned claim read = %#v", output) + } + case "remove-claim": + for _, name := range []string{filename, asClaimFilename(filename)} { + if _, err := os.Lstat(filepath.Join(original, "oidc", name)); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("pinned claim remove left original %s: %v", name, err) + } + if _, err := os.Lstat(filepath.Join(replacementRoot, "oidc", name)); err != nil { + t.Fatalf("replacement claim pair was mutated: %v", err) + } + } + } + }) + } +} + +func TestProtocolPinsTheOriginalPrivateRootForLayoutCreation(t *testing.T) { + parent := privateTestRoot(t) + root := filepath.Join(parent, "auth") + replacement := filepath.Join(parent, "replacement") + moved := filepath.Join(parent, "auth-original") + for _, directory := range []string{root, replacement} { + if err := safeio.EnsurePrivateDirectory(directory); err != nil { + t.Fatal(err) + } + } + swapped := false + blocked := false + restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) { + if stage != "after-auth-root-open" || swapped || blocked { + return + } + if runtime.GOOS == "windows" { + if err := os.Rename(root, moved); err == nil { + t.Fatal("retained Windows root handle permitted rename") + } + blocked = true + return + } + if err := os.Rename(root, moved); err != nil { + t.Fatal(err) + } + if err := os.Rename(replacement, root); err != nil { + t.Fatal(err) + } + swapped = true + }) + t.Cleanup(restore) + + runRequest(t, request{Version: 1, Operation: "ensure-layout", Root: root}) + if !swapped && !blocked { + t.Fatal("layout creation did not expose the retained-root test hook") + } + original := root + replacementRoot := replacement + if swapped { + original = moved + replacementRoot = root + } + for _, child := range []string{"sessions", "oidc"} { + if err := safeio.ValidatePrivateDirectory(filepath.Join(original, child)); err != nil { + t.Fatalf("pinned %s child: %v", child, err) + } + if _, err := os.Lstat(filepath.Join(replacementRoot, child)); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("replacement root received layout mutation: %v", err) + } + } +} + +func TestProtocolValidateRootRechecksChildrenObservedMissingUnderTheRetainedRoot(t *testing.T) { + root := filepath.Join(privateTestRoot(t), "auth") + if err := safeio.EnsurePrivateDirectory(root); err != nil { + t.Fatal(err) + } + installedUnsafeChild := false + restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) { + if stage != "after-auth-layout-first-pass" || installedUnsafeChild { + return + } + if err := os.Mkdir(filepath.Join(root, "sessions"), 0o755); err != nil { + t.Fatal(err) + } + installedUnsafeChild = true + }) + t.Cleanup(restore) + + runRejected(t, request{Version: 1, Operation: "validate-root", Root: root}) + if !installedUnsafeChild { + t.Fatal("layout validation did not expose the missing-child recheck hook") + } +} + +func TestProtocolValidateLayoutPinsTheOriginalPOSIXRootAcrossMissingAndUnsafeChildren(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("native Windows retained-handle coverage uses DACLs in storage_windows_test.go") + } + for _, scenario := range []struct { + name string + originalUnsafe bool + replacementUnsafe bool + wantAccepted bool + }{ + {name: "rejects an unsafe retained child after a safe lexical replacement", originalUnsafe: true, wantAccepted: false}, + {name: "accepts missing retained children despite an unsafe lexical replacement", replacementUnsafe: true, wantAccepted: true}, + } { + t.Run(scenario.name, func(t *testing.T) { + parent := privateTestRoot(t) + root := filepath.Join(parent, "auth") + replacement := filepath.Join(parent, "replacement") + moved := filepath.Join(parent, "auth-original") + for _, directory := range []string{root, replacement} { + if err := safeio.EnsurePrivateDirectory(directory); err != nil { + t.Fatal(err) + } + } + if scenario.originalUnsafe { + if err := os.Mkdir(filepath.Join(root, "sessions"), 0o755); err != nil { + t.Fatal(err) + } + } + if scenario.replacementUnsafe { + if err := os.Mkdir(filepath.Join(replacement, "sessions"), 0o755); err != nil { + t.Fatal(err) + } + } + swapped := false + restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) { + if stage != "after-auth-root-open" || swapped { + return + } + if err := os.Rename(root, moved); err != nil { + t.Fatal(err) + } + if err := os.Rename(replacement, root); err != nil { + t.Fatal(err) + } + swapped = true + }) + t.Cleanup(restore) + input := request{Version: 1, Operation: "validate-root", Root: root} + if scenario.wantAccepted { + runRequest(t, input) + } else { + runRejected(t, input) + } + if !swapped { + t.Fatal("layout validation did not retain and expose the original root") + } + for _, child := range []string{"sessions", "oidc"} { + if _, err := os.Lstat(filepath.Join(moved, child)); !errors.Is(err, os.ErrNotExist) && !(scenario.originalUnsafe && child == "sessions") { + t.Fatalf("validate-root mutated retained %s child: %v", child, err) + } + } + }) + } +} + func TestProtocolReadsOnlyBoundedPrivateAuthConfig(t *testing.T) { root := privateTestRoot(t) filename := "auth.yaml" @@ -199,6 +562,36 @@ func TestProtocolReadsOnlyBoundedPrivateAuthConfig(t *testing.T) { runRejected(t, request{Version: 1, Operation: "read-auth-config", Root: root, Directory: "sessions", Filename: filename}) } +func TestProtocolReadsLocalUsersOnlyAsABoundedPrivateRegularFile(t *testing.T) { + root := privateTestRoot(t) + filename := "users.yaml" + path := filepath.Join(root, filename) + contents := []byte("version: 1\nusers: []\n") + if err := safeio.WriteCanonicalNewPrivateFile(path, contents, 0o600); err != nil { + t.Fatal(err) + } + read := runRequest(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: filename}) + if !read.Found || decodeContent(t, read) != string(contents) { + t.Fatalf("read-local-users = %#v", read) + } + + hardLink := filepath.Join(root, "users-copy.yaml") + if err := os.Link(path, hardLink); err != nil { + t.Fatal(err) + } + runRejected(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: filename}) + if err := os.Remove(hardLink); err != nil { + t.Fatal(err) + } + if err := os.Remove(path); err != nil { + t.Fatal(err) + } + if err := safeio.WriteCanonicalNewPrivateFile(path, bytes.Repeat([]byte("x"), maximumAuthConfigBytes+1), 0o600); err != nil { + t.Fatal(err) + } + runRejected(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: filename}) +} + func TestProtocolPermitsBoundedReservationSlotsOnlyForOIDCRecords(t *testing.T) { root := filepath.Join(privateTestRoot(t), "auth") slot := "slot-00.json" diff --git a/tools/tht/internal/authstorage/storage_windows_test.go b/tools/tht/internal/authstorage/storage_windows_test.go index fbd76bb9..264913e7 100644 --- a/tools/tht/internal/authstorage/storage_windows_test.go +++ b/tools/tht/internal/authstorage/storage_windows_test.go @@ -71,6 +71,96 @@ func TestProtocolValidatesAndCreatesTheCompleteWindowsSessionLayout(t *testing.T runRejected(t, request{Version: 1, Operation: "validate-root", Root: root}) } +func TestProtocolReadsWindowsLocalUsersWithOwnerOnlyDACLAndNoReparseFallback(t *testing.T) { + root := filepath.Join(t.TempDir(), "auth") + if err := safeio.EnsurePrivateDirectory(root); err != nil { + t.Fatal(err) + } + path := filepath.Join(root, "users.yaml") + contents := []byte("version: 1\nusers: []\n") + if err := safeio.WriteCanonicalNewPrivateFile(path, contents, 0o600); err != nil { + t.Fatal(err) + } + read := runRequest(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: "users.yaml"}) + if !read.Found || decodeContent(t, read) != string(contents) { + t.Fatalf("read-local-users = %#v", read) + } + if err := setPermissiveDACL(path); err != nil { + t.Fatal(err) + } + runRejected(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: "users.yaml"}) + + if err := os.Remove(path); err != nil { + t.Fatal(err) + } + if err := os.Symlink(filepath.Join(root, "missing-target.yaml"), path); err != nil { + t.Skipf("Windows host does not permit test symlink creation: %v", err) + } + runRejected(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: "users.yaml"}) +} + +func TestProtocolValidateLayoutPinsWindowsRootAcrossMissingAndUnsafeChildren(t *testing.T) { + for _, scenario := range []struct { + name string + originalUnsafe bool + replacementUnsafe bool + wantAccepted bool + }{ + {name: "rejects unsafe retained child", originalUnsafe: true, wantAccepted: false}, + {name: "accepts missing retained child while replacement is unsafe", replacementUnsafe: true, wantAccepted: true}, + } { + t.Run(scenario.name, func(t *testing.T) { + parent := t.TempDir() + root := filepath.Join(parent, "auth") + replacement := filepath.Join(parent, "replacement") + moved := filepath.Join(parent, "auth-original") + for _, directory := range []string{root, replacement} { + if err := safeio.EnsurePrivateDirectory(directory); err != nil { + t.Fatal(err) + } + } + if scenario.originalUnsafe { + path := filepath.Join(root, "sessions") + if err := safeio.EnsurePrivateDirectory(path); err != nil { + t.Fatal(err) + } + if err := setPermissiveDACL(path); err != nil { + t.Fatal(err) + } + } + if scenario.replacementUnsafe { + path := filepath.Join(replacement, "sessions") + if err := safeio.EnsurePrivateDirectory(path); err != nil { + t.Fatal(err) + } + if err := setPermissiveDACL(path); err != nil { + t.Fatal(err) + } + } + blocked := false + restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) { + if stage != "after-auth-root-open" || blocked { + return + } + if err := os.Rename(root, moved); err == nil { + t.Fatal("retained Windows root handle permitted rename") + } + blocked = true + }) + t.Cleanup(restore) + input := request{Version: 1, Operation: "validate-root", Root: root} + if scenario.wantAccepted { + runRequest(t, input) + } else { + runRejected(t, input) + } + if !blocked { + t.Fatal("layout validation did not retain the Windows root handle") + } + }) + } +} + func setPermissiveDACL(path string) error { world, err := windows.StringToSid("S-1-1-0") if err != nil { diff --git a/tools/tht/internal/safeio/private_root.go b/tools/tht/internal/safeio/private_root.go new file mode 100644 index 00000000..c1e9459b --- /dev/null +++ b/tools/tht/internal/safeio/private_root.go @@ -0,0 +1,67 @@ +package safeio + +import ( + "strings" + "sync" +) + +// PrivateDirectoryHandle pins one private directory for the duration of a storage operation. +// Implementations use descriptor-relative operations on POSIX and retained no-delete handles on +// Windows. Callers must close every returned child before releasing its parent. +type PrivateDirectoryHandle interface { + Close() error + Validate() error + OpenChild(name string, ensure bool) (PrivateDirectoryHandle, bool, error) + CreateRegular(name string, contents []byte) (bool, error) + ReadRegular(name string, maximum int64) ([]byte, bool, error) + ReplaceRegular(name string, contents []byte) error + RemoveRegular(name string) (bool, error) + ListPage(maximumEntries int, afterName string, validName func(string) bool, validLinks func(string, uint64) bool) (PrivateDirectoryPage, error) + ClaimRegular(source, claim string) (bool, error) + ReadClaim(source, claim string, maximum int64) ([]byte, bool, error) + RemoveClaim(source, claim string) (bool, error) +} + +// OpenPrivateDirectory validates or creates the final private directory while retaining the +// opened canonical directory handle. With ensure=false, found=false means the final component is +// absent but its already-opened parent proves the same operation could safely create it. +func OpenPrivateDirectory(path string, ensure bool) (PrivateDirectoryHandle, bool, error) { + if err := ValidateCanonicalPath(path); err != nil { + return nil, false, ErrUnsafeFile + } + return openPrivateDirectory(path, ensure) +} + +func validPrivateLeafName(name string) bool { + return name != "" && name != "." && name != ".." && !strings.ContainsAny(name, "\\/:\x00") +} + +var privateDirectoryTestHook struct { + sync.RWMutex + hook func(string) +} + +// SetPrivateDirectoryTestHookForTest installs a deterministic race hook for internal package +// tests. It has no production effect unless a test explicitly installs one. +func SetPrivateDirectoryTestHookForTest(hook func(string)) func() { + privateDirectoryTestHook.Lock() + previous := privateDirectoryTestHook.hook + privateDirectoryTestHook.hook = hook + privateDirectoryTestHook.Unlock() + return func() { + privateDirectoryTestHook.Lock() + privateDirectoryTestHook.hook = previous + privateDirectoryTestHook.Unlock() + } +} + +// NotifyPrivateDirectoryTestHookForTest marks an internal retained-root boundary. It is called +// only by storage code and lets tests install deterministic directory replacement races. +func NotifyPrivateDirectoryTestHookForTest(stage string) { + privateDirectoryTestHook.RLock() + hook := privateDirectoryTestHook.hook + privateDirectoryTestHook.RUnlock() + if hook != nil { + hook(stage) + } +} diff --git a/tools/tht/internal/safeio/private_root_unix.go b/tools/tht/internal/safeio/private_root_unix.go new file mode 100644 index 00000000..532c7c39 --- /dev/null +++ b/tools/tht/internal/safeio/private_root_unix.go @@ -0,0 +1,420 @@ +//go:build !windows + +package safeio + +import ( + "errors" + "io" + "os" + "sort" + "time" + + "golang.org/x/sys/unix" +) + +type unixPrivateDirectory struct { + descriptor int +} + +func openPrivateDirectory(path string, ensure bool) (PrivateDirectoryHandle, bool, error) { + parents, err := openCanonicalUnixParent(path) + if err != nil { + return nil, false, ErrUnsafeFile + } + defer parents.Close() + return openPrivateUnixDirectoryAt(parents.parent, parents.target, ensure) +} + +func openPrivateUnixDirectoryAt(parent int, name string, ensure bool) (PrivateDirectoryHandle, bool, error) { + if parent < 0 || !validPrivateLeafName(name) { + return nil, false, ErrUnsafeFile + } + for attempt := 0; attempt < 2; attempt++ { + descriptor, err := unix.Openat(parent, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0) + if err == nil { + value := &unixPrivateDirectory{descriptor: descriptor} + if value.Validate() != nil { + _ = value.Close() + return nil, false, ErrUnsafeFile + } + return value, true, nil + } + if !errors.Is(err, unix.ENOENT) { + return nil, false, ErrUnsafeFile + } + if !ensure { + if unix.Faccessat(parent, ".", unix.W_OK|unix.X_OK, unix.AT_EACCESS) != nil { + return nil, false, ErrUnsafeFile + } + return nil, false, nil + } + if err := unix.Mkdirat(parent, name, 0o700); err != nil && !errors.Is(err, unix.EEXIST) { + return nil, false, ErrUnsafeFile + } + } + return nil, false, ErrUnsafeFile +} + +func (directory *unixPrivateDirectory) Close() error { + if directory == nil || directory.descriptor < 0 { + return nil + } + err := unix.Close(directory.descriptor) + directory.descriptor = -1 + if err != nil { + return ErrUnsafeFile + } + return nil +} + +func (directory *unixPrivateDirectory) Validate() error { + if directory == nil || directory.descriptor < 0 { + return ErrUnsafeFile + } + var stat unix.Stat_t + if err := unix.Fstat(directory.descriptor, &stat); err != nil || !privateUnixDirectoryStat(&stat) { + return ErrUnsafeFile + } + return nil +} + +func (directory *unixPrivateDirectory) OpenChild(name string, ensure bool) (PrivateDirectoryHandle, bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(name) { + return nil, false, ErrUnsafeFile + } + child, found, err := openPrivateUnixDirectoryAt(directory.descriptor, name, ensure) + if err != nil || directory.Validate() != nil { + if child != nil { + _ = child.Close() + } + return nil, false, ErrUnsafeFile + } + return child, found, nil +} + +func privateUnixRootRegular(stat *unix.Stat_t, links uint64) bool { + return stat != nil && stat.Mode&unix.S_IFMT == unix.S_IFREG && uint64(stat.Nlink) == links && + stat.Uid == uint32(os.Geteuid()) && stat.Mode&0o7777 == 0o600 +} + +func sameUnixRootRegular(left, right unix.Stat_t) bool { + return left.Dev == right.Dev && left.Ino == right.Ino && left.Size == right.Size && + left.Mtim == right.Mtim && left.Ctim == right.Ctim && left.Mode == right.Mode && left.Nlink == right.Nlink +} + +func requirePrivateUnixRootRegularAt(directory int, name string, links uint64) (unix.Stat_t, error) { + var stat unix.Stat_t + if err := unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW); err != nil || !privateUnixRootRegular(&stat, links) { + return unix.Stat_t{}, ErrUnsafeFile + } + return stat, nil +} + +func privateUnixRootRegularAtAllowedLinks(directory int, name string, allowed ...uint64) error { + var stat unix.Stat_t + if err := unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW); err != nil { + return ErrUnsafeFile + } + for _, links := range allowed { + if privateUnixRootRegular(&stat, links) { + return nil + } + } + return ErrUnsafeFile +} + +func requireSamePrivateUnixRootPairAt(directory int, source, claim string) error { + left, err := requirePrivateUnixRootRegularAt(directory, source, 2) + if err != nil { + return ErrUnsafeFile + } + right, err := requirePrivateUnixRootRegularAt(directory, claim, 2) + if err != nil || !sameUnixPrivateFile(left, right) { + return ErrUnsafeFile + } + return nil +} + +func isPrivateUnixRootClaimAbsentOrOrphanAt(directory int, source, claim string) bool { + var sourceStat, claimStat unix.Stat_t + sourceErr := unix.Fstatat(directory, source, &sourceStat, unix.AT_SYMLINK_NOFOLLOW) + claimErr := unix.Fstatat(directory, claim, &claimStat, unix.AT_SYMLINK_NOFOLLOW) + if errors.Is(sourceErr, unix.ENOENT) && errors.Is(claimErr, unix.ENOENT) { + return true + } + return errors.Is(sourceErr, unix.ENOENT) && claimErr == nil && privateUnixRootRegular(&claimStat, 1) +} + +func writeAllPrivateRoot(file *os.File, contents []byte) error { + for written := 0; written < len(contents); { + count, err := file.Write(contents[written:]) + written += count + if err != nil { + return err + } + if count == 0 { + return io.ErrShortWrite + } + } + return file.Sync() +} + +func (directory *unixPrivateDirectory) CreateRegular(name string, contents []byte) (bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 { + return false, ErrUnsafeFile + } + descriptor, err := unix.Openat(directory.descriptor, name, + unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0o600) + if errors.Is(err, unix.EEXIST) { + if _, existingErr := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); existingErr != nil || directory.Validate() != nil { + return false, ErrUnsafeFile + } + return false, nil + } + if err != nil { + return false, ErrUnsafeFile + } + file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-create") + if file == nil { + _ = unix.Close(descriptor) + _ = unix.Unlinkat(directory.descriptor, name, 0) + return false, ErrUnsafeFile + } + failed := true + defer func() { + if failed { + _ = unix.Unlinkat(directory.descriptor, name, 0) + } + }() + if unix.Fchmod(descriptor, 0o600) != nil { + _ = file.Close() + return false, ErrUnsafeFile + } + var stat unix.Stat_t + if unix.Fstat(descriptor, &stat) != nil || !privateUnixRootRegular(&stat, 1) || writeAllPrivateRoot(file, contents) != nil || file.Close() != nil { + return false, ErrUnsafeFile + } + if directory.Validate() != nil || unix.Fsync(directory.descriptor) != nil { + return false, ErrUnsafeFile + } + failed = false + return true, nil +} + +func (directory *unixPrivateDirectory) ReadRegular(name string, maximum int64) ([]byte, bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(name) || maximum < 0 || maximum == int64(^uint64(0)>>1) { + return nil, false, ErrUnsafeFile + } + before, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1) + if err != nil { + var stat unix.Stat_t + if statErr := unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW); errors.Is(statErr, unix.ENOENT) { + return nil, false, nil + } + return nil, false, ErrUnsafeFile + } + descriptor, err := unix.Openat(directory.descriptor, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0) + if err != nil { + return nil, false, ErrUnsafeFile + } + file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-read") + if file == nil { + _ = unix.Close(descriptor) + return nil, false, ErrUnsafeFile + } + contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1)) + var opened unix.Stat_t + statErr := unix.Fstat(descriptor, &opened) + closeErr := file.Close() + var current unix.Stat_t + currentErr := unix.Fstatat(directory.descriptor, name, ¤t, unix.AT_SYMLINK_NOFOLLOW) + if readErr != nil || statErr != nil || closeErr != nil || currentErr != nil || int64(len(contents)) > maximum || + !privateUnixRootRegular(&opened, 1) || !privateUnixRootRegular(¤t, 1) || + !sameUnixRootRegular(before, opened) || !sameUnixRootRegular(opened, current) || directory.Validate() != nil { + return nil, false, ErrUnsafeFile + } + return contents, true, nil +} + +func (directory *unixPrivateDirectory) ReplaceRegular(name string, contents []byte) error { + if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 { + return ErrUnsafeFile + } + if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil { + return ErrUnsafeFile + } + temporary, err := writePrivateTemporaryAt(directory.descriptor, contents) + if err != nil { + return ErrUnsafeFile + } + defer func() { _ = unix.Unlinkat(directory.descriptor, temporary, 0) }() + if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil || directory.Validate() != nil { + return ErrUnsafeFile + } + if err := unix.Renameat(directory.descriptor, temporary, directory.descriptor, name); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil { + return ErrUnsafeFile + } + return nil +} + +func (directory *unixPrivateDirectory) RemoveRegular(name string) (bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(name) { + return false, ErrUnsafeFile + } + if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil { + var stat unix.Stat_t + if statErr := unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW); errors.Is(statErr, unix.ENOENT) { + return false, nil + } + return false, ErrUnsafeFile + } + if err := unix.Unlinkat(directory.descriptor, name, 0); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil { + return false, ErrUnsafeFile + } + return true, nil +} + +func (directory *unixPrivateDirectory) ListPage( + maximumEntries int, + afterName string, + validName func(string) bool, + validLinks func(string, uint64) bool, +) (PrivateDirectoryPage, error) { + if directory.Validate() != nil || maximumEntries < 1 || maximumEntries > 4096 || validName == nil || validLinks == nil || (afterName != "" && !validName(afterName)) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + var before unix.Stat_t + if unix.Fstat(directory.descriptor, &before) != nil || !privateUnixDirectoryStat(&before) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + duplicate, err := unix.Dup(directory.descriptor) + if err != nil { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + file := os.NewFile(uintptr(duplicate), "tht-safeio-private-root-list") + if file == nil { + _ = unix.Close(duplicate) + return PrivateDirectoryPage{}, ErrUnsafeFile + } + defer file.Close() + seen := make(map[string]struct{}, maximumEntries+1) + selected := make([]PrivateDirectoryEntry, 0, maximumEntries+1) + scanned := 0 + for { + entries, readErr := file.ReadDir(1) + if readErr != nil && !errors.Is(readErr, io.EOF) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + if len(entries) == 0 { + break + } + if len(entries) != 1 { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + scanned++ + if scanned > maximumPrivateDirectoryPageScanEntries { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + name := entries[0].Name() + if !validPrivateLeafName(name) || !validName(name) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + if _, duplicate := seen[name]; duplicate { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + seen[name] = struct{}{} + var stat unix.Stat_t + if unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW) != nil || + !privateUnixRootRegular(&stat, uint64(stat.Nlink)) || !validLinks(name, uint64(stat.Nlink)) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + if name > afterName { + selected = appendBoundedPrivateDirectoryEntry(selected, PrivateDirectoryEntry{ + Name: name, ModifiedUnixMs: time.Unix(stat.Mtim.Sec, stat.Mtim.Nsec).UnixMilli(), + }, maximumEntries+1) + } + if errors.Is(readErr, io.EOF) { + break + } + } + var after unix.Stat_t + if unix.Fstat(directory.descriptor, &after) != nil || !privateUnixDirectoryStat(&after) || + before.Dev != after.Dev || before.Ino != after.Ino || before.Mode != after.Mode || before.Mtim != after.Mtim || before.Ctim != after.Ctim || directory.Validate() != nil { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + sort.Slice(selected, func(left, right int) bool { return selected[left].Name < selected[right].Name }) + more := len(selected) > maximumEntries + if more { + selected = selected[:maximumEntries] + } + return PrivateDirectoryPage{Entries: selected, More: more}, nil +} + +func (directory *unixPrivateDirectory) ClaimRegular(source, claim string) (bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) { + return false, ErrUnsafeFile + } + if err := privateUnixRootRegularAtAllowedLinks(directory.descriptor, source, 1); err != nil { + if requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim) == nil || isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) { + return false, nil + } + return false, ErrUnsafeFile + } + if err := unix.Linkat(directory.descriptor, source, directory.descriptor, claim, 0); err != nil { + if errors.Is(err, unix.EEXIST) && privateUnixRootRegularAtAllowedLinks(directory.descriptor, claim, 1, 2) == nil { + return false, nil + } + return false, ErrUnsafeFile + } + if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil { + return false, ErrUnsafeFile + } + return true, nil +} + +func (directory *unixPrivateDirectory) ReadClaim(source, claim string, maximum int64) ([]byte, bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) || maximum < 0 || maximum == int64(^uint64(0)>>1) { + return nil, false, ErrUnsafeFile + } + if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil { + if isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) { + return nil, false, nil + } + return nil, false, ErrUnsafeFile + } + descriptor, err := unix.Openat(directory.descriptor, source, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0) + if err != nil { + return nil, false, ErrUnsafeFile + } + file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-claim") + if file == nil { + _ = unix.Close(descriptor) + return nil, false, ErrUnsafeFile + } + contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1)) + var after unix.Stat_t + statErr := unix.Fstat(descriptor, &after) + closeErr := file.Close() + if readErr != nil || statErr != nil || closeErr != nil || int64(len(contents)) > maximum || !privateUnixRootRegular(&after, 2) || + requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim) != nil || directory.Validate() != nil { + return nil, false, ErrUnsafeFile + } + return contents, true, nil +} + +func (directory *unixPrivateDirectory) RemoveClaim(source, claim string) (bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) { + return false, ErrUnsafeFile + } + if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil { + if isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) { + return false, nil + } + return false, ErrUnsafeFile + } + if err := unix.Unlinkat(directory.descriptor, source, 0); err != nil || privateUnixRootRegularAtAllowedLinks(directory.descriptor, claim, 1) != nil || + unix.Unlinkat(directory.descriptor, claim, 0) != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil { + return false, ErrUnsafeFile + } + return true, nil +} diff --git a/tools/tht/internal/safeio/private_root_windows.go b/tools/tht/internal/safeio/private_root_windows.go new file mode 100644 index 00000000..c86697e3 --- /dev/null +++ b/tools/tht/internal/safeio/private_root_windows.go @@ -0,0 +1,921 @@ +//go:build windows + +package safeio + +import ( + "errors" + "io" + "os" + "runtime" + "sort" + "time" + "unsafe" + + "golang.org/x/sys/windows" +) + +// windowsPrivateDirectory keeps the root's canonical component handles alive, then uses NT +// RootDirectory-relative opens for every descendant. Unlike a lexical child path, an NT relative +// object name is resolved by the already-open directory handle and cannot be redirected by a +// rename, replacement, or reparse point at the original root path. +type windowsPrivateDirectory struct { + anchors *windowsParentHandles + parent windows.Handle + handle windows.Handle + info windows.ByHandleFileInformation +} + +type windowsPrivateRegularAt struct { + handle windows.Handle + info windows.ByHandleFileInformation +} + +func openPrivateDirectory(path string, ensure bool) (PrivateDirectoryHandle, bool, error) { + anchors, target, err := openCanonicalWindowsParent(path) + if err != nil || anchors == nil || len(anchors.handles) == 0 { + if anchors != nil { + anchors.Close() + } + return nil, false, ErrUnsafeFile + } + parent := anchors.handles[len(anchors.handles)-1] + handle, found, err := openWindowsPrivateDirectoryAt(parent, target, false) + if err != nil { + anchors.Close() + return nil, false, ErrUnsafeFile + } + if !found { + // The final root is absent. The retained canonical parent must prove that the same + // ensure operation could create it; validation itself must remain side-effect free. + // FILE_APPEND_DATA is the Win32 spelling of directory FILE_ADD_SUBDIRECTORY. + probe, probeErr := openWindowsComponentWithAccess(anchors.directory, true, windows.FILE_APPEND_DATA) + if probeErr != nil { + anchors.Close() + return nil, false, ErrUnsafeFile + } + _ = windows.CloseHandle(probe) + if !ensure { + anchors.Close() + return nil, false, nil + } + // The canonical parent chain remains pinned by anchors; this extra handle supplies + // FILE_ADD_SUBDIRECTORY for the one initial root creation without reopening a child + // beneath the private root lexically. + writableParent, writableErr := openWindowsComponentWithAccess(anchors.directory, true, windows.FILE_APPEND_DATA) + if writableErr != nil { + anchors.Close() + return nil, false, ErrUnsafeFile + } + handle, found, err = openWindowsPrivateDirectoryAt(writableParent, target, true) + _ = windows.CloseHandle(writableParent) + if err != nil || !found { + anchors.Close() + return nil, false, ErrUnsafeFile + } + } + value := &windowsPrivateDirectory{anchors: anchors, handle: handle} + if value.captureAndValidate() != nil { + _ = value.Close() + return nil, false, ErrUnsafeFile + } + return value, true, nil +} + +func openWindowsPrivateDirectoryAt(parent windows.Handle, name string, ensure bool) (windows.Handle, bool, error) { + if parent == 0 || !validPrivateLeafName(name) { + return 0, false, ErrUnsafeFile + } + for attempt := 0; attempt < 2; attempt++ { + handle, err := openWindowsRelativeDirectory(parent, name) + if err == nil { + return handle, true, nil + } + if !isWindowsRelativeNotFound(err) { + return 0, false, ErrUnsafeFile + } + if !ensure { + return 0, false, nil + } + handle, err = createWindowsRelativePrivateDirectory(parent, name) + if err == nil { + return handle, true, nil + } + } + return 0, false, ErrUnsafeFile +} + +func openWindowsRelativeDirectory(parent windows.Handle, name string) (windows.Handle, error) { + handle, err := openWindowsRelativeObject( + parent, + name, + // The retained directory handle is also the RootDirectory for create, rename, + // hard-link, and delete operations below, so it needs the owner's full private + // directory capability rather than a read-only probe handle. + windows.FILE_GENERIC_READ|windows.FILE_GENERIC_WRITE|windows.DELETE, + windows.FILE_OPEN, + windows.FILE_DIRECTORY_FILE|windows.FILE_SYNCHRONOUS_IO_NONALERT|windows.FILE_OPEN_REPARSE_POINT, + nil, + ) + if err != nil { + return 0, err + } + if _, err := privateWindowsDirectoryInfo(handle); err != nil { + _ = windows.CloseHandle(handle) + return 0, ErrUnsafeFile + } + return handle, nil +} + +func createWindowsRelativePrivateDirectory(parent windows.Handle, name string) (windows.Handle, error) { + security, err := newOwnerOnlySecurityDescriptor() + if err != nil { + return 0, ErrUnsafeFile + } + defer security.Close() + handle, err := openWindowsRelativeObject( + parent, + name, + windows.FILE_GENERIC_READ|windows.FILE_GENERIC_WRITE|windows.DELETE, + windows.FILE_CREATE, + windows.FILE_DIRECTORY_FILE|windows.FILE_SYNCHRONOUS_IO_NONALERT|windows.FILE_OPEN_REPARSE_POINT, + security, + ) + if err != nil { + return 0, err + } + if _, err := privateWindowsDirectoryInfo(handle); err != nil { + _ = markWindowsHandleForDelete(handle) + _ = windows.CloseHandle(handle) + return 0, ErrUnsafeFile + } + return handle, nil +} + +func openWindowsRelativeObject( + parent windows.Handle, + name string, + access uint32, + disposition uint32, + options uint32, + security *ownerOnlySecurityDescriptor, +) (windows.Handle, error) { + if parent == 0 || !validPrivateLeafName(name) { + return 0, ErrUnsafeFile + } + objectName, err := windows.NewNTUnicodeString(name) + if err != nil { + return 0, ErrUnsafeFile + } + attributes := &windows.OBJECT_ATTRIBUTES{ + Length: uint32(unsafe.Sizeof(windows.OBJECT_ATTRIBUTES{})), + RootDirectory: parent, + ObjectName: objectName, + Attributes: windows.OBJ_CASE_INSENSITIVE, + SecurityDescriptor: nil, + } + if security != nil { + attributes.SecurityDescriptor = security.descriptor + } + var ( + handle windows.Handle + status windows.IO_STATUS_BLOCK + allocationSize int64 + ) + err = windows.NtCreateFile( + &handle, + access, + attributes, + &status, + &allocationSize, + windows.FILE_ATTRIBUTE_NORMAL, + windowsRetainedHandleShareMode, + disposition, + options, + 0, + 0, + ) + runtime.KeepAlive(objectName) + runtime.KeepAlive(security) + if err != nil { + return 0, err + } + return handle, nil +} + +func privateWindowsDirectoryInfo(handle windows.Handle) (windows.ByHandleFileInformation, error) { + var info windows.ByHandleFileInformation + if handle == 0 || windows.GetFileInformationByHandle(handle, &info) != nil || + info.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 || + info.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY == 0 || + validateOwnerOnlyDACL(handle) != nil { + return windows.ByHandleFileInformation{}, ErrUnsafeFile + } + return info, nil +} + +func privateWindowsRegularInfo(handle windows.Handle, allowedLinks ...uint32) (windows.ByHandleFileInformation, error) { + var info windows.ByHandleFileInformation + if handle == 0 || windows.GetFileInformationByHandle(handle, &info) != nil || + info.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 || + info.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY != 0 || + validateOwnerOnlyDACL(handle) != nil { + return windows.ByHandleFileInformation{}, ErrUnsafeFile + } + for _, links := range allowedLinks { + if info.NumberOfLinks == links { + return info, nil + } + } + return windows.ByHandleFileInformation{}, ErrUnsafeFile +} + +func isWindowsRelativeNotFound(err error) bool { + return errors.Is(err, windows.ERROR_FILE_NOT_FOUND) || + errors.Is(err, windows.ERROR_PATH_NOT_FOUND) || + errors.Is(err, windows.STATUS_NO_SUCH_FILE) || + errors.Is(err, windows.STATUS_OBJECT_NAME_NOT_FOUND) || + errors.Is(err, windows.STATUS_OBJECT_PATH_NOT_FOUND) +} + +func isWindowsRelativeCollision(err error) bool { + return errors.Is(err, windows.ERROR_FILE_EXISTS) || + errors.Is(err, windows.ERROR_ALREADY_EXISTS) || + errors.Is(err, windows.STATUS_OBJECT_NAME_COLLISION) +} + +func (directory *windowsPrivateDirectory) captureAndValidate() error { + if directory == nil || directory.handle == 0 { + return ErrUnsafeFile + } + info, err := privateWindowsDirectoryInfo(directory.handle) + if err != nil { + return ErrUnsafeFile + } + directory.info = info + return nil +} + +func (directory *windowsPrivateDirectory) Close() error { + if directory == nil { + return nil + } + var result error + if directory.handle != 0 { + if err := windows.CloseHandle(directory.handle); err != nil { + result = ErrUnsafeFile + } + directory.handle = 0 + } + if directory.parent != 0 { + if err := windows.CloseHandle(directory.parent); err != nil { + result = ErrUnsafeFile + } + directory.parent = 0 + } + if directory.anchors != nil { + directory.anchors.Close() + directory.anchors = nil + } + return result +} + +func sameWindowsPrivateDirectoryIdentity(left, right windows.ByHandleFileInformation) bool { + return left.VolumeSerialNumber == right.VolumeSerialNumber && left.FileIndexHigh == right.FileIndexHigh && + left.FileIndexLow == right.FileIndexLow && left.FileAttributes == right.FileAttributes +} + +func sameWindowsPrivateDirectorySnapshot(left, right windows.ByHandleFileInformation) bool { + return sameWindowsPrivateDirectoryIdentity(left, right) && left.LastWriteTime == right.LastWriteTime +} + +func (directory *windowsPrivateDirectory) Validate() error { + if directory == nil || directory.handle == 0 || (directory.anchors == nil && directory.parent == 0) { + return ErrUnsafeFile + } + if directory.anchors != nil && len(directory.anchors.handles) == 0 { + return ErrUnsafeFile + } + if directory.parent != 0 { + if _, err := privateWindowsDirectoryInfo(directory.parent); err != nil { + return ErrUnsafeFile + } + } + current, err := privateWindowsDirectoryInfo(directory.handle) + if err != nil || !sameWindowsPrivateDirectoryIdentity(directory.info, current) { + return ErrUnsafeFile + } + return nil +} + +func duplicateWindowsRetainedHandle(handle windows.Handle) (windows.Handle, error) { + if handle == 0 { + return 0, ErrUnsafeFile + } + var duplicate windows.Handle + process := windows.CurrentProcess() + if err := windows.DuplicateHandle(process, handle, process, &duplicate, 0, false, windows.DUPLICATE_SAME_ACCESS); err != nil { + return 0, ErrUnsafeFile + } + return duplicate, nil +} + +func (directory *windowsPrivateDirectory) OpenChild(name string, ensure bool) (PrivateDirectoryHandle, bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(name) { + return nil, false, ErrUnsafeFile + } + parent, err := duplicateWindowsRetainedHandle(directory.handle) + if err != nil { + return nil, false, ErrUnsafeFile + } + handle, found, err := openWindowsPrivateDirectoryAt(parent, name, ensure) + if err != nil || !found { + _ = windows.CloseHandle(parent) + if err != nil { + return nil, false, ErrUnsafeFile + } + return nil, false, nil + } + child := &windowsPrivateDirectory{parent: parent, handle: handle} + if child.captureAndValidate() != nil || directory.Validate() != nil { + _ = child.Close() + return nil, false, ErrUnsafeFile + } + return child, true, nil +} + +func openWindowsPrivateRegularAt( + parent windows.Handle, + name string, + access uint32, + allowedLinks ...uint32, +) (*windowsPrivateRegularAt, error) { + handle, err := openWindowsRelativeObject( + parent, + name, + access, + windows.FILE_OPEN, + windows.FILE_NON_DIRECTORY_FILE|windows.FILE_SYNCHRONOUS_IO_NONALERT|windows.FILE_OPEN_REPARSE_POINT, + nil, + ) + if err != nil { + return nil, err + } + info, err := privateWindowsRegularInfo(handle, allowedLinks...) + if err != nil { + _ = windows.CloseHandle(handle) + return nil, ErrUnsafeFile + } + return &windowsPrivateRegularAt{handle: handle, info: info}, nil +} + +func openWindowsPrivateRegularAtAllowedLinks( + parent windows.Handle, + name string, + access uint32, + allowedLinks ...uint32, +) (*windowsPrivateRegularAt, error) { + var ( + lastError error + unsafeFound bool + ) + for _, links := range allowedLinks { + value, err := openWindowsPrivateRegularAt(parent, name, access, links) + if err == nil { + return value, nil + } + lastError = err + if !isWindowsRelativeNotFound(err) { + unsafeFound = true + } + } + if unsafeFound { + return nil, ErrUnsafeFile + } + return nil, lastError +} + +func createWindowsPrivateRegularAt(parent windows.Handle, name string) (*windowsPrivateRegularAt, error) { + security, err := newOwnerOnlySecurityDescriptor() + if err != nil { + return nil, ErrUnsafeFile + } + defer security.Close() + handle, err := openWindowsRelativeObject( + parent, + name, + windows.FILE_GENERIC_READ|windows.FILE_GENERIC_WRITE|windows.DELETE, + windows.FILE_CREATE, + windows.FILE_NON_DIRECTORY_FILE|windows.FILE_SYNCHRONOUS_IO_NONALERT|windows.FILE_OPEN_REPARSE_POINT, + security, + ) + if err != nil { + return nil, err + } + info, err := privateWindowsRegularInfo(handle, 1) + if err != nil { + _ = markWindowsHandleForDelete(handle) + _ = windows.CloseHandle(handle) + return nil, ErrUnsafeFile + } + return &windowsPrivateRegularAt{handle: handle, info: info}, nil +} + +func (value *windowsPrivateRegularAt) Close() error { + if value == nil || value.handle == 0 { + return nil + } + handle := value.handle + value.handle = 0 + if err := windows.CloseHandle(handle); err != nil { + return ErrUnsafeFile + } + return nil +} + +func writeWindowsPrivateRegular(value *windowsPrivateRegularAt, contents []byte) error { + if value == nil || value.handle == 0 || len(contents) == 0 { + return ErrUnsafeFile + } + for remaining := contents; len(remaining) > 0; { + var written uint32 + if err := windows.WriteFile(value.handle, remaining, &written, nil); err != nil || written == 0 || int(written) > len(remaining) { + return ErrUnsafeFile + } + remaining = remaining[written:] + } + if err := windows.FlushFileBuffers(value.handle); err != nil { + return ErrUnsafeFile + } + info, err := privateWindowsRegularInfo(value.handle, 1) + if err != nil { + return ErrUnsafeFile + } + value.info = info + return nil +} + +func readWindowsPrivateRegular(value *windowsPrivateRegularAt, maximum int64, links uint32) ([]byte, error) { + if value == nil || value.handle == 0 || maximum < 0 || maximum == int64(^uint64(0)>>1) { + return nil, ErrUnsafeFile + } + contents := make([]byte, 0, 4096) + buffer := make([]byte, 4096) + for { + var read uint32 + err := windows.ReadFile(value.handle, buffer, &read, nil) + if read > 0 { + if int64(len(contents))+int64(read) > maximum { + return nil, ErrUnsafeFile + } + contents = append(contents, buffer[:read]...) + } + if err != nil { + if errors.Is(err, windows.ERROR_HANDLE_EOF) { + break + } + return nil, ErrUnsafeFile + } + if read == 0 { + break + } + } + after, err := privateWindowsRegularInfo(value.handle, links) + if err != nil || !sameWindowsPrivateFile(value.info, after) { + return nil, ErrUnsafeFile + } + value.info = after + return contents, nil +} + +func markWindowsHandleForDelete(handle windows.Handle) error { + if handle == 0 { + return ErrUnsafeFile + } + buffer := [1]byte{1} + var status windows.IO_STATUS_BLOCK + if err := windows.NtSetInformationFile(handle, &status, &buffer[0], uint32(len(buffer)), windows.FileDispositionInformation); err != nil { + return ErrUnsafeFile + } + return nil +} + +func closeAndDeleteWindowsPrivateRegular(value *windowsPrivateRegularAt) error { + if value == nil || value.handle == 0 || markWindowsHandleForDelete(value.handle) != nil || value.Close() != nil { + return ErrUnsafeFile + } + return nil +} + +func (directory *windowsPrivateDirectory) CreateRegular(name string, contents []byte) (bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 { + return false, ErrUnsafeFile + } + value, err := createWindowsPrivateRegularAt(directory.handle, name) + if err != nil { + existing, existingErr := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1) + if existingErr == nil { + _ = existing.Close() + return false, nil + } + return false, ErrUnsafeFile + } + published := false + defer func() { + if !published { + _ = closeAndDeleteWindowsPrivateRegular(value) + } + }() + if writeWindowsPrivateRegular(value, contents) != nil || directory.Validate() != nil { + return false, ErrUnsafeFile + } + if value.Close() != nil { + return false, ErrUnsafeFile + } + published = true + return true, nil +} + +func (directory *windowsPrivateDirectory) ReadRegular(name string, maximum int64) ([]byte, bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(name) || maximum < 0 || maximum == int64(^uint64(0)>>1) { + return nil, false, ErrUnsafeFile + } + value, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1) + if isWindowsRelativeNotFound(err) { + return nil, false, nil + } + if err != nil { + return nil, false, ErrUnsafeFile + } + defer value.Close() + contents, err := readWindowsPrivateRegular(value, maximum, 1) + if err != nil || directory.Validate() != nil { + return nil, false, ErrUnsafeFile + } + return contents, true, nil +} + +type windowsRelativeNameInformation struct { + Flags uint32 + RootDirectory windows.Handle + FileNameLength uint32 + FileName [1]uint16 +} + +func setWindowsRelativeNameInformation( + handle windows.Handle, + parent windows.Handle, + name string, + class uint32, + flags uint32, +) error { + if handle == 0 || parent == 0 || !validPrivateLeafName(name) { + return ErrUnsafeFile + } + encoded, err := windows.UTF16FromString(name) + if err != nil || len(encoded) < 2 { + return ErrUnsafeFile + } + nameBytes := (len(encoded) - 1) * 2 + var header windowsRelativeNameInformation + size := int(unsafe.Offsetof(header.FileName)) + nameBytes + buffer := make([]byte, size) + value := (*windowsRelativeNameInformation)(unsafe.Pointer(&buffer[0])) + value.Flags = flags + value.RootDirectory = parent + value.FileNameLength = uint32(nameBytes) + copy(unsafe.Slice(&value.FileName[0], len(encoded)-1), encoded[:len(encoded)-1]) + var status windows.IO_STATUS_BLOCK + if err := windows.NtSetInformationFile(handle, &status, &buffer[0], uint32(len(buffer)), class); err != nil { + return ErrUnsafeFile + } + runtime.KeepAlive(encoded) + runtime.KeepAlive(buffer) + return nil +} + +func renameWindowsPrivateRegularAt(handle windows.Handle, parent windows.Handle, name string) error { + return setWindowsRelativeNameInformation(handle, parent, name, windows.FileRenameInformation, windows.FILE_RENAME_REPLACE_IF_EXISTS) +} + +func linkWindowsPrivateRegularAt(handle windows.Handle, parent windows.Handle, name string) error { + return setWindowsRelativeNameInformation(handle, parent, name, windows.FileLinkInformation, 0) +} + +func createWindowsPrivateTemporaryAt(parent windows.Handle, contents []byte) (*windowsPrivateRegularAt, error) { + for attempt := 0; attempt < 16; attempt++ { + name, err := randomTemporaryName() + if err != nil { + return nil, ErrUnsafeFile + } + value, err := createWindowsPrivateRegularAt(parent, name) + if err != nil { + if isWindowsRelativeCollision(err) { + continue + } + return nil, ErrUnsafeFile + } + if writeWindowsPrivateRegular(value, contents) == nil { + return value, nil + } + _ = closeAndDeleteWindowsPrivateRegular(value) + return nil, ErrUnsafeFile + } + return nil, ErrUnsafeFile +} + +func (directory *windowsPrivateDirectory) ReplaceRegular(name string, contents []byte) error { + if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 { + return ErrUnsafeFile + } + existing, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1) + if err != nil { + return ErrUnsafeFile + } + if existing.Close() != nil { + return ErrUnsafeFile + } + temporary, err := createWindowsPrivateTemporaryAt(directory.handle, contents) + if err != nil { + return ErrUnsafeFile + } + renamed := false + defer func() { + if !renamed { + _ = closeAndDeleteWindowsPrivateRegular(temporary) + } + }() + current, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1) + if err != nil || current.Close() != nil || directory.Validate() != nil { + return ErrUnsafeFile + } + if renameWindowsPrivateRegularAt(temporary.handle, directory.handle, name) != nil || temporary.Close() != nil { + return ErrUnsafeFile + } + renamed = true + replaced, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1) + if err != nil || replaced.Close() != nil || directory.Validate() != nil { + return ErrUnsafeFile + } + return nil +} + +func (directory *windowsPrivateDirectory) RemoveRegular(name string) (bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(name) { + return false, ErrUnsafeFile + } + value, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ|windows.DELETE, 1) + if isWindowsRelativeNotFound(err) { + return false, nil + } + if err != nil { + return false, ErrUnsafeFile + } + if closeAndDeleteWindowsPrivateRegular(value) != nil || directory.Validate() != nil { + return false, ErrUnsafeFile + } + return true, nil +} + +func (directory *windowsPrivateDirectory) ListPage( + maximumEntries int, + afterName string, + validName func(string) bool, + validLinks func(string, uint64) bool, +) (PrivateDirectoryPage, error) { + if directory.Validate() != nil || maximumEntries < 1 || maximumEntries > 4096 || validName == nil || validLinks == nil || + (afterName != "" && !validName(afterName)) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + var before windows.ByHandleFileInformation + if err := windows.GetFileInformationByHandle(directory.handle, &before); err != nil || + !sameWindowsPrivateDirectoryIdentity(directory.info, before) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + duplicate, err := duplicateWindowsRetainedHandle(directory.handle) + if err != nil { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + file := os.NewFile(uintptr(duplicate), "tht-safeio-private-root-list") + if file == nil { + _ = windows.CloseHandle(duplicate) + return PrivateDirectoryPage{}, ErrUnsafeFile + } + defer file.Close() + seen := make(map[string]struct{}, maximumEntries+1) + selected := make([]PrivateDirectoryEntry, 0, maximumEntries+1) + scanned := 0 + for { + entries, readErr := file.ReadDir(1) + if readErr != nil && !errors.Is(readErr, io.EOF) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + if len(entries) == 0 { + break + } + if len(entries) != 1 { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + scanned++ + if scanned > maximumPrivateDirectoryPageScanEntries { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + name := entries[0].Name() + if !validPrivateLeafName(name) || !validName(name) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + if _, duplicate := seen[name]; duplicate { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + seen[name] = struct{}{} + value, valueErr := openWindowsPrivateRegularAtAllowedLinks(directory.handle, name, windows.FILE_GENERIC_READ, 1, 2) + if valueErr != nil { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + info := value.info + if value.Close() != nil { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + if !validLinks(name, uint64(info.NumberOfLinks)) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + if name > afterName { + selected = appendBoundedPrivateDirectoryEntry(selected, PrivateDirectoryEntry{ + Name: name, ModifiedUnixMs: time.Unix(0, info.LastWriteTime.Nanoseconds()).UnixMilli(), + }, maximumEntries+1) + } + if errors.Is(readErr, io.EOF) { + break + } + } + var after windows.ByHandleFileInformation + if windows.GetFileInformationByHandle(directory.handle, &after) != nil || directory.Validate() != nil || + !sameWindowsPrivateDirectorySnapshot(before, after) { + return PrivateDirectoryPage{}, ErrUnsafeFile + } + sort.Slice(selected, func(left, right int) bool { return selected[left].Name < selected[right].Name }) + more := len(selected) > maximumEntries + if more { + selected = selected[:maximumEntries] + } + return PrivateDirectoryPage{Entries: selected, More: more}, nil +} + +func sameWindowsRelativeClaim(source, claim *windowsPrivateRegularAt) bool { + return source != nil && claim != nil && sameWindowsPrivateFile(source.info, claim.info) +} + +func windowsRelativeClaimPairExists(directory *windowsPrivateDirectory, source, claim string) (bool, error) { + left, err := openWindowsPrivateRegularAt(directory.handle, source, windows.FILE_GENERIC_READ, 2) + if isWindowsRelativeNotFound(err) { + return false, nil + } + if err != nil { + return false, ErrUnsafeFile + } + defer left.Close() + right, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 2) + if isWindowsRelativeNotFound(err) { + return false, nil + } + if err != nil { + return false, ErrUnsafeFile + } + defer right.Close() + return sameWindowsRelativeClaim(left, right), nil +} + +func windowsRelativeClaimAbsentOrOrphan(directory *windowsPrivateDirectory, source, claim string) (bool, error) { + current, err := openWindowsPrivateRegularAtAllowedLinks(directory.handle, source, windows.FILE_GENERIC_READ, 1, 2) + if err == nil { + _ = current.Close() + return false, nil + } + if !isWindowsRelativeNotFound(err) { + return false, ErrUnsafeFile + } + orphan, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 1) + if err == nil { + _ = orphan.Close() + return true, nil + } + if isWindowsRelativeNotFound(err) { + return true, nil + } + return false, ErrUnsafeFile +} + +func (directory *windowsPrivateDirectory) ClaimRegular(source, claim string) (bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) { + return false, ErrUnsafeFile + } + value, err := openWindowsPrivateRegularAt( + directory.handle, + source, + windows.FILE_GENERIC_READ|windows.FILE_GENERIC_WRITE|windows.DELETE, + 1, + ) + if err != nil { + pair, pairErr := windowsRelativeClaimPairExists(directory, source, claim) + if pairErr == nil && pair { + return false, nil + } + orphan, orphanErr := windowsRelativeClaimAbsentOrOrphan(directory, source, claim) + if orphanErr == nil && orphan { + return false, nil + } + return false, ErrUnsafeFile + } + defer value.Close() + if linkWindowsPrivateRegularAt(value.handle, directory.handle, claim) != nil { + existing, existingErr := openWindowsPrivateRegularAtAllowedLinks(directory.handle, claim, windows.FILE_GENERIC_READ, 1, 2) + if existingErr == nil { + _ = existing.Close() + return false, nil + } + return false, ErrUnsafeFile + } + after, err := privateWindowsRegularInfo(value.handle, 2) + if err != nil { + return false, ErrUnsafeFile + } + value.info = after + claimed, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 2) + if err != nil { + return false, ErrUnsafeFile + } + defer claimed.Close() + if !sameWindowsRelativeClaim(value, claimed) || directory.Validate() != nil { + return false, ErrUnsafeFile + } + return true, nil +} + +func (directory *windowsPrivateDirectory) ReadClaim(source, claim string, maximum int64) ([]byte, bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) || + maximum < 0 || maximum == int64(^uint64(0)>>1) { + return nil, false, ErrUnsafeFile + } + value, err := openWindowsPrivateRegularAt(directory.handle, source, windows.FILE_GENERIC_READ, 2) + if err != nil { + orphan, orphanErr := windowsRelativeClaimAbsentOrOrphan(directory, source, claim) + if orphanErr == nil && orphan { + return nil, false, nil + } + return nil, false, ErrUnsafeFile + } + defer value.Close() + claimed, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 2) + if isWindowsRelativeNotFound(err) { + return nil, false, nil + } + if err != nil { + return nil, false, ErrUnsafeFile + } + defer claimed.Close() + if !sameWindowsRelativeClaim(value, claimed) { + return nil, false, ErrUnsafeFile + } + contents, err := readWindowsPrivateRegular(value, maximum, 2) + if err != nil { + return nil, false, ErrUnsafeFile + } + afterClaim, err := privateWindowsRegularInfo(claimed.handle, 2) + if err != nil || !sameWindowsPrivateFile(value.info, afterClaim) || directory.Validate() != nil { + return nil, false, ErrUnsafeFile + } + return contents, true, nil +} + +func (directory *windowsPrivateDirectory) RemoveClaim(source, claim string) (bool, error) { + if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) { + return false, ErrUnsafeFile + } + value, err := openWindowsPrivateRegularAt(directory.handle, source, windows.FILE_GENERIC_READ|windows.DELETE, 2) + if err != nil { + orphan, orphanErr := windowsRelativeClaimAbsentOrOrphan(directory, source, claim) + if orphanErr == nil && orphan { + return false, nil + } + return false, ErrUnsafeFile + } + claimed, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 2) + if isWindowsRelativeNotFound(err) { + _ = value.Close() + return false, nil + } + if err != nil || !sameWindowsRelativeClaim(value, claimed) { + _ = value.Close() + if claimed != nil { + _ = claimed.Close() + } + return false, ErrUnsafeFile + } + if claimed.Close() != nil || closeAndDeleteWindowsPrivateRegular(value) != nil { + return false, ErrUnsafeFile + } + remaining, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ|windows.DELETE, 1) + if err != nil || closeAndDeleteWindowsPrivateRegular(remaining) != nil || directory.Validate() != nil { + return false, ErrUnsafeFile + } + return true, nil +}