From c9b02fc57ec27a7309676df27bbaf6ce32495f45 Mon Sep 17 00:00:00 2001 From: mptyl Date: Sun, 16 Aug 2026 21:43:03 +0200 Subject: [PATCH] fix(auth): add Windows session storage bridge --- backend/src/auth/session-store.ts | 145 ++++++++- backend/src/auth/windows-auth-storage.ts | 278 +++++++++++++++++ backend/test/auth-session-store.test.ts | 81 ++++- backend/test/windows-auth-storage.test.ts | 74 +++++ tools/tht/cmd/tht/main.go | 4 + tools/tht/internal/authstorage/storage.go | 291 ++++++++++++++++++ .../tht/internal/authstorage/storage_test.go | 164 ++++++++++ .../authstorage/storage_windows_test.go | 71 +++++ tools/tht/internal/safeio/claim_unix.go | 169 ++++++++++ tools/tht/internal/safeio/claim_windows.go | 210 +++++++++++++ tools/tht/internal/safeio/files.go | 80 +++++ .../tht/internal/safeio/private_read_unix.go | 17 + .../internal/safeio/private_read_windows.go | 44 +++ tools/tht/internal/safeio/private_windows.go | 10 +- 14 files changed, 1630 insertions(+), 8 deletions(-) create mode 100644 backend/src/auth/windows-auth-storage.ts create mode 100644 backend/test/windows-auth-storage.test.ts create mode 100644 tools/tht/internal/authstorage/storage.go create mode 100644 tools/tht/internal/authstorage/storage_test.go create mode 100644 tools/tht/internal/authstorage/storage_windows_test.go create mode 100644 tools/tht/internal/safeio/claim_unix.go create mode 100644 tools/tht/internal/safeio/claim_windows.go create mode 100644 tools/tht/internal/safeio/private_read_unix.go create mode 100644 tools/tht/internal/safeio/private_read_windows.go diff --git a/backend/src/auth/session-store.ts b/backend/src/auth/session-store.ts index 76339052..ad8216c1 100644 --- a/backend/src/auth/session-store.ts +++ b/backend/src/auth/session-store.ts @@ -22,6 +22,10 @@ import { dirname, isAbsolute, join, normalize } from "node:path"; import { z } from "zod"; import type { PrincipalContext } from "./principal.js"; import type { AuthSessionRecord, OidcStateRecord, Permission, Role } from "./types.js"; +import { + createWindowsAuthStorageBridge, + type WindowsAuthStorageBridge, +} from "./windows-auth-storage.js"; const TOKEN_BYTES = 32; const TOKEN_PATTERN = /^[A-Za-z0-9_-]{43}$/; @@ -96,6 +100,11 @@ export interface AuthSessionStore { consumeOidcState(state: string, now?: Date): Promise; } +/** Narrow test seam for the native Windows tht-backed storage adaptor. */ +export interface FileAuthSessionStoreOptions { + windowsStorageBridge?: WindowsAuthStorageBridge; +} + interface FileIdentity { dev: number; ino: number; @@ -300,10 +309,8 @@ function privateDirectory(path: string): void { } function storageDirectories(root: string): StorageDirectories { - // Node's chmod is not a Windows DACL boundary. The existing Go operator store has a - // CreateFile security-descriptor path, but no equivalent safe Node primitive is available. - // Refuse before probing or creating the configured root rather than publishing browser state - // with inherited ACLs. + // 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. if (process.platform === "win32") throw invalid(); if (typeof root !== "string" || root.length === 0 || root.includes("\0") || !isAbsolute(root) || normalize(root) !== root) throw invalid(); @@ -526,6 +533,15 @@ function parseOidcStateRecord(source: string): OidcStateRecord { } } +function parseWindowsRecord(contents: Buffer, maximumBytes: number, parse: (source: string) => T): T { + if (!Buffer.isBuffer(contents) || contents.length === 0 || contents.length > maximumBytes) throw invalid(); + try { + return parse(new TextDecoder("utf-8", { fatal: true }).decode(contents)); + } catch { + throw invalid(); + } +} + interface OidcStateClaim { state: TrustedFile; claimIdentity: FileIdentity; @@ -701,7 +717,20 @@ export function deriveCsrfToken(sessionToken: string): string { } } -export function createFileAuthSessionStore(root: string, validity?: AuthSessionValidity): AuthSessionStore { +export function createFileAuthSessionStore( + root: string, + validity?: AuthSessionValidity, + options: FileAuthSessionStoreOptions = {}, +): AuthSessionStore { + const windowsStorage = process.platform === "win32" + ? options.windowsStorageBridge ?? createWindowsAuthStorageBridge() + : undefined; + + function requiredWindowsStorage(): WindowsAuthStorageBridge { + if (windowsStorage === undefined) throw invalid(); + return windowsStorage; + } + async function createSession(input: SessionCreateInput, now = new Date()): Promise { const nowMs = dateMilliseconds(now); let validated: z.infer; @@ -730,6 +759,16 @@ export function createFileAuthSessionStore(root: string, validity?: AuthSessionV 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 = storageDirectories(root); for (let attempt = 0; attempt < 8; attempt += 1) { const token = randomBytes(TOKEN_BYTES).toString("base64url"); @@ -746,6 +785,24 @@ export function createFileAuthSessionStore(root: string, validity?: AuthSessionV 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, validity)) return record; + } catch { + await bridge.remove(root, "sessions", filename); + throw invalid(); + } + await bridge.remove(root, "sessions", filename); + return undefined; + } const directories = storageDirectories(root); const trusted = readTrusted(directories.sessions, filename, MAX_SESSION_RECORD_BYTES, parseSessionRecord); if (!trusted) return undefined; @@ -769,6 +826,27 @@ export function createFileAuthSessionStore(root: string, validity?: AuthSessionV 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 directories = storageDirectories(root); const trusted = readTrusted(directories.sessions, filename, MAX_SESSION_RECORD_BYTES, parseSessionRecord); if (!trusted) return; @@ -798,6 +876,10 @@ export function createFileAuthSessionStore(root: string, validity?: AuthSessionV 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 = storageDirectories(root); removeTrusted(directories.sessions, filename); }); @@ -822,6 +904,14 @@ export function createFileAuthSessionStore(root: string, validity?: AuthSessionV expiresAt: isoAt(expiresMs), }; const contents = serialize(record, MAX_OIDC_STATE_RECORD_BYTES); + if (process.platform === "win32") { + const bridge = requiredWindowsStorage(); + for (let attempt = 0; attempt < 8; attempt += 1) { + const state = randomBytes(TOKEN_BYTES).toString("base64url"); + if (await bridge.create(root, "oidc", digestFilename(state), contents)) return { state, record }; + } + throw invalid(); + } const directories = storageDirectories(root); for (let attempt = 0; attempt < 8; attempt += 1) { const state = randomBytes(TOKEN_BYTES).toString("base64url"); @@ -835,6 +925,12 @@ export function createFileAuthSessionStore(root: string, validity?: AuthSessionV 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); + return oidcStateExpired(record, nowMs) ? undefined : record; + } const directories = storageDirectories(root); const claim = claimOidcState(directories.oidc, filename); // An installed claim belongs to another process/store instance. Only the process which @@ -851,6 +947,45 @@ export function createFileAuthSessionStore(root: string, validity?: AuthSessionV async function prune(now = new Date()): Promise { const nowMs = dateMilliseconds(now); + if (process.platform === "win32") { + const bridge = requiredWindowsStorage(); + const sessionEntries = await bridge.list(root, "sessions"); + const oidcEntries = await bridge.list(root, "oidc"); + let removed = 0; + for (const entry of sessionEntries) { + if (!DIGEST_FILENAME_PATTERN.test(entry.name)) throw invalid(); + const contents = await bridge.read(root, "sessions", entry.name); + if (!contents) continue; + const record = parseWindowsRecord(contents, MAX_SESSION_RECORD_BYTES, parseSessionRecord); + if (sessionExpired(record, nowMs) && await bridge.remove(root, "sessions", entry.name)) removed += 1; + } + 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])); + if (stateNames.size + claimEntries.size !== oidcEntries.length) throw invalid(); + 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; + } + return removed; + } const directories = storageDirectories(root); let removed = 0; const pruneDirectory = async ( diff --git a/backend/src/auth/windows-auth-storage.ts b/backend/src/auth/windows-auth-storage.ts new file mode 100644 index 00000000..3c9c62fc --- /dev/null +++ b/backend/src/auth/windows-auth-storage.ts @@ -0,0 +1,278 @@ +import { spawn } from "node:child_process"; +import { win32 } from "node:path"; +import { z } from "zod"; + +const PROTOCOL_VERSION = 1; +const MAX_PROTOCOL_BYTES = 64 * 1024; +const MAX_RESPONSE_BYTES = 64 * 1024; +const MAX_SESSION_BYTES = 16 * 1024; +const MAX_OIDC_BYTES = 8 * 1024; +const MAX_ENTRIES = 256; +const TIMEOUT_MS = 5_000; +const DIGEST_FILENAME = /^[a-f0-9]{64}\.json$/; +const CLAIM_FILENAME = /^[a-f0-9]{64}\.claim$/; + +const invalid = (): Error => new Error("auth_session_store_invalid"); + +export type WindowsAuthStorageDirectory = "sessions" | "oidc"; + +export interface WindowsAuthStorageEntry { + name: string; + modifiedUnixMs: number; +} + +/** Internal adapter boundary for the file-session store's native Windows path. */ +export interface WindowsAuthStorageBridge { + 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; + remove(root: string, directory: WindowsAuthStorageDirectory, filename: string): Promise; + list(root: string, directory: WindowsAuthStorageDirectory): Promise; + claimConsume(root: string, filename: string): Promise; + readClaim(root: string, filename: string): Promise; + removeClaim(root: string, filename: string): Promise; +} + +export interface WindowsAuthStorageInvocation { + executable: string; + args: readonly string[]; + input: Buffer; + timeoutMs: number; +} + +export interface WindowsAuthStorageInvocationResult { + code: number; + stdout: Buffer; + stderr: Buffer; +} + +export interface WindowsAuthStorageBridgeOptions { + /** Test-only transport seam. Production always uses the no-shell child-process invocation. */ + invoke?: (invocation: WindowsAuthStorageInvocation) => Promise; + /** Optional configured tht path. Defaults to THT_BIN, then the safe bare command `tht`. */ + thtExecutable?: string; +} + +const responseSchema = z.strictObject({ + version: z.literal(PROTOCOL_VERSION), + ok: z.literal(true), + created: z.boolean().optional(), + replaced: z.boolean().optional(), + removed: z.boolean().optional(), + found: z.boolean().optional(), + contentBase64: z.string().max(MAX_PROTOCOL_BYTES).optional(), + entries: z.array(z.strictObject({ + name: z.string().max(128), + modifiedUnixMs: z.number().int().safe().nonnegative(), + })).max(MAX_ENTRIES).optional(), +}); + +type BridgeResponse = z.infer; + +interface BridgeRequest { + version: typeof PROTOCOL_VERSION; + operation: "create" | "read" | "replace" | "remove" | "list" | "claim-consume" | "read-claim" | "remove-claim"; + root: string; + directory: WindowsAuthStorageDirectory; + filename?: string; + contentBase64?: string; +} + +function directoryMaximum(directory: WindowsAuthStorageDirectory): number { + return directory === "sessions" ? MAX_SESSION_BYTES : MAX_OIDC_BYTES; +} + +function canonicalBase64(value: string, maximum: number): Buffer { + if (typeof value !== "string" || value.length > Math.ceil(maximum / 3) * 4) throw invalid(); + try { + const decoded = Buffer.from(value, "base64"); + if (decoded.length === 0 || decoded.length > maximum || decoded.toString("base64") !== value) throw invalid(); + return decoded; + } catch { + throw invalid(); + } +} + +function validateRoot(root: string): void { + if (typeof root !== "string" || root.length === 0 || /[\u0000-\u001f\u007f]/.test(root) + || !win32.isAbsolute(root) || win32.normalize(root) !== root) throw invalid(); +} + +function validateFilename(filename: string, claim = false): void { + if (typeof filename !== "string" || !(claim ? CLAIM_FILENAME : DIGEST_FILENAME).test(filename)) throw invalid(); +} + +function safeThtExecutable(value: string | undefined): string { + const executable = value ?? process.env.THT_BIN ?? "tht"; + if (typeof executable !== "string" || executable.length === 0 || /[\u0000-\u001f\u007f]/.test(executable)) throw invalid(); + if (executable === "tht" || executable === "tht.exe") return executable; + if (win32.isAbsolute(executable) && win32.normalize(executable) === executable && /\.exe$/i.test(executable)) return executable; + throw invalid(); +} + +function parseResponse(result: WindowsAuthStorageInvocationResult): BridgeResponse { + if (!Number.isInteger(result.code) || result.code !== 0 || !Buffer.isBuffer(result.stdout) + || !Buffer.isBuffer(result.stderr) || result.stdout.length === 0 || result.stdout.length > MAX_RESPONSE_BYTES) { + throw invalid(); + } + try { + const source = new TextDecoder("utf-8", { fatal: true }).decode(result.stdout); + return responseSchema.parse(JSON.parse(source)); + } catch { + throw invalid(); + } +} + +function encodedRequest(request: BridgeRequest): Buffer { + validateRoot(request.root); + if (request.filename !== undefined) validateFilename(request.filename); + if (request.contentBase64 !== undefined) canonicalBase64(request.contentBase64, directoryMaximum(request.directory)); + const encoded = Buffer.from(JSON.stringify(request), "utf8"); + if (encoded.length === 0 || encoded.length > MAX_PROTOCOL_BYTES) throw invalid(); + return encoded; +} + +function environmentForBridge(): NodeJS.ProcessEnv { + const path = process.env.PATH; + const systemRoot = process.env.SystemRoot ?? process.env.SYSTEMROOT; + return { + ...(path === undefined ? {} : { PATH: path }), + ...(systemRoot === undefined ? {} : { SystemRoot: systemRoot }), + }; +} + +async function invokeTht(invocation: WindowsAuthStorageInvocation): Promise { + return new Promise((resolve, reject) => { + let settled = false; + let timeout: NodeJS.Timeout | undefined; + const stdout: Buffer[] = []; + const stderr: Buffer[] = []; + let stdoutBytes = 0; + let stderrBytes = 0; + const settle = (callback: () => void): void => { + if (settled) return; + settled = true; + if (timeout !== undefined) clearTimeout(timeout); + callback(); + }; + let child: ReturnType; + try { + child = spawn(invocation.executable, [...invocation.args], { + shell: false, + windowsHide: true, + stdio: ["pipe", "pipe", "pipe"], + env: environmentForBridge(), + }); + } catch { + reject(invalid()); + return; + } + if (!child.stdin || !child.stdout || !child.stderr) { + try { child.kill(); } catch { /* unavailable child streams fail closed */ } + reject(invalid()); + return; + } + const stdin = child.stdin; + const stdoutStream = child.stdout; + const stderrStream = child.stderr; + timeout = setTimeout(() => { + try { child.kill(); } catch { /* child failure is converted below */ } + settle(() => reject(invalid())); + }, invocation.timeoutMs); + child.once("error", () => settle(() => reject(invalid()))); + stdoutStream.on("data", (chunk: Buffer) => { + stdoutBytes += chunk.length; + if (stdoutBytes > MAX_RESPONSE_BYTES) { + try { child.kill(); } catch { /* child failure is converted below */ } + settle(() => reject(invalid())); + return; + } + stdout.push(Buffer.from(chunk)); + }); + stderrStream.on("data", (chunk: Buffer) => { + stderrBytes += chunk.length; + if (stderrBytes <= MAX_RESPONSE_BYTES) stderr.push(Buffer.from(chunk)); + }); + child.once("close", (code) => settle(() => resolve({ + code: code ?? -1, + stdout: Buffer.concat(stdout), + stderr: Buffer.concat(stderr), + }))); + stdin.once("error", () => settle(() => reject(invalid()))); + stdin.end(invocation.input); + }); +} + +function contentFrom(response: BridgeResponse, maximum: number): Buffer | undefined { + if (response.found !== true) { + if (response.contentBase64 !== undefined) throw invalid(); + return undefined; + } + if (response.contentBase64 === undefined) throw invalid(); + return canonicalBase64(response.contentBase64, maximum); +} + +export function createWindowsAuthStorageBridge(options: WindowsAuthStorageBridgeOptions = {}): WindowsAuthStorageBridge { + const executable = safeThtExecutable(options.thtExecutable); + const invoke = options.invoke ?? invokeTht; + const request = async (value: BridgeRequest): Promise => { + try { + const response = await invoke({ + executable, + args: ["_auth-storage"], + input: encodedRequest(value), + timeoutMs: TIMEOUT_MS, + }); + return parseResponse(response); + } catch { + throw invalid(); + } + }; + const recordRequest = (operation: BridgeRequest["operation"], root: string, directory: WindowsAuthStorageDirectory, filename: string, contents?: Buffer): BridgeRequest => ({ + version: PROTOCOL_VERSION, + operation, + root, + directory, + filename, + ...(contents === undefined ? {} : { contentBase64: contents.toString("base64") }), + }); + + return { + 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)); + if (response.created === undefined) throw invalid(); + return response.created; + }, + async read(root, directory, filename) { + return contentFrom(await request(recordRequest("read", root, directory, filename)), directoryMaximum(directory)); + }, + async replace(root, directory, filename, contents) { + if (!Buffer.isBuffer(contents) || contents.length === 0 || contents.length > directoryMaximum(directory)) throw invalid(); + const response = await request(recordRequest("replace", root, directory, filename, contents)); + if (response.replaced !== true) throw invalid(); + }, + async remove(root, directory, filename) { + const response = await request(recordRequest("remove", root, directory, filename)); + return response.removed === true; + }, + async list(root, directory) { + const response = await request({ version: PROTOCOL_VERSION, operation: "list", root, directory }); + if (response.entries === undefined) throw invalid(); + for (const entry of response.entries) { + if (!DIGEST_FILENAME.test(entry.name) && !CLAIM_FILENAME.test(entry.name)) throw invalid(); + } + return response.entries.map((entry) => ({ name: entry.name, modifiedUnixMs: entry.modifiedUnixMs })); + }, + async claimConsume(root, filename) { + return contentFrom(await request(recordRequest("claim-consume", root, "oidc", filename)), MAX_OIDC_BYTES); + }, + async readClaim(root, filename) { + return contentFrom(await request(recordRequest("read-claim", root, "oidc", filename)), MAX_OIDC_BYTES); + }, + async removeClaim(root, filename) { + const response = await request(recordRequest("remove-claim", root, "oidc", filename)); + return response.removed === true; + }, + }; +} diff --git a/backend/test/auth-session-store.test.ts b/backend/test/auth-session-store.test.ts index fa659b56..0ffcc142 100644 --- a/backend/test/auth-session-store.test.ts +++ b/backend/test/auth-session-store.test.ts @@ -488,20 +488,97 @@ describe("file-backed auth session store", () => { await expectStoreInvalid(store.revoke(created.token)); }); - test("rejects native Windows storage before any state write", async () => { + test("routes native Windows session creation through an injected storage bridge", async () => { const storageRoot = root(); const store = validStore(join(storageRoot, "windows-auth-state")); + const createRecord = vi.fn(async () => true); const originalPlatform = Object.getOwnPropertyDescriptor(process, "platform"); if (!originalPlatform) throw new Error("platform descriptor unavailable"); Object.defineProperty(process, "platform", { configurable: true, value: "win32" }); try { - await expectStoreInvalid(create(store)); + const created = await create( + createFileAuthSessionStore(join(storageRoot, "windows-auth-state"), { + currentAuthConfigRevision: () => revision, + findLocalUser: async () => validLocalUser, + }, { windowsStorageBridge: { create: createRecord } } as never), + ); + expect(created.token).toMatch(/^[A-Za-z0-9_-]{43}$/); + expect(createRecord).toHaveBeenCalledTimes(1); expect(existsSync(join(storageRoot, "windows-auth-state"))).toBe(false); } finally { Object.defineProperty(process, "platform", originalPlatform); } }); + test("routes native Windows lifecycle and OIDC operations only through the storage bridge", async () => { + const records = new Map(); + const calls: string[] = []; + const key = (directory: string, filename: string) => `${directory}/${filename}`; + const bridge = { + create: async (_root: string, directory: string, filename: string, contents: Buffer) => { + calls.push("create"); + const entry = key(directory, filename); + if (records.has(entry)) return false; + records.set(entry, Buffer.from(contents)); + return true; + }, + read: async (_root: string, directory: string, filename: string) => { + calls.push("read"); + const value = records.get(key(directory, filename)); + return value === undefined ? undefined : Buffer.from(value); + }, + replace: async (_root: string, directory: string, filename: string, contents: Buffer) => { + calls.push("replace"); + records.set(key(directory, filename), Buffer.from(contents)); + }, + remove: async (_root: string, directory: string, filename: string) => { + calls.push("remove"); + return records.delete(key(directory, filename)); + }, + list: async (_root: string, directory: string) => { + calls.push("list"); + return [...records.keys()] + .filter((entry) => entry.startsWith(`${directory}/`)) + .map((entry) => ({ name: entry.slice(directory.length + 1), modifiedUnixMs: base.getTime() })); + }, + claimConsume: async (_root: string, filename: string) => { + calls.push("claim-consume"); + const entry = key("oidc", filename); + const value = records.get(entry); + records.delete(entry); + return value === undefined ? undefined : Buffer.from(value); + }, + readClaim: async () => undefined, + removeClaim: async () => false, + }; + const originalPlatform = Object.getOwnPropertyDescriptor(process, "platform"); + if (!originalPlatform) throw new Error("platform descriptor unavailable"); + Object.defineProperty(process, "platform", { configurable: true, value: "win32" }); + try { + const store = createFileAuthSessionStore("C:\\ProgramData\\ThothII\\auth", { + currentAuthConfigRevision: () => revision, + findLocalUser: async () => validLocalUser, + }, { windowsStorageBridge: bridge } as never); + const session = await create(store); + await expect(store.resolve(session.token)).resolves.toMatchObject({ subject: session.record.subject }); + await store.touch(session.token, new Date(base.getTime() + 5 * 60_000)); + await store.revoke(session.token); + await expect(store.resolve(session.token)).resolves.toBeUndefined(); + + const oidc = await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base); + await expect(store.consumeOidcState(oidc.state, new Date(base.getTime() + 9 * 60_000))) + .resolves.toMatchObject({ nonce: "n".repeat(43) }); + await expect(store.consumeOidcState(oidc.state)).resolves.toBeUndefined(); + await create(store, { idleTtlMs: 60_000, absoluteTtlMs: 60_000 }); + await store.createOidcState({ nonce: "x".repeat(43), codeVerifier: "y".repeat(43), returnTo: "/" }, base); + await expect(store.prune(new Date(base.getTime() + 11 * 60_000))).resolves.toBe(2); + expect(calls).toEqual(expect.arrayContaining(["create", "read", "replace", "remove", "claim-consume", "list"])); + expect(records).toHaveLength(0); + } finally { + Object.defineProperty(process, "platform", originalPlatform); + } + }); + test("revokes on config, local-user, revision, enabled, or role mismatch before returning", async () => { const storageRoot = root(); let currentRevision = revision; diff --git a/backend/test/windows-auth-storage.test.ts b/backend/test/windows-auth-storage.test.ts new file mode 100644 index 00000000..1f044c3c --- /dev/null +++ b/backend/test/windows-auth-storage.test.ts @@ -0,0 +1,74 @@ +import { describe, expect, test } from "vitest"; +import { createWindowsAuthStorageBridge } from "../src/auth/windows-auth-storage.js"; + +const root = "C:\\ProgramData\\ThothII\\auth"; +const filename = "a".repeat(64) + ".json"; + +describe("Windows auth-storage bridge", () => { + test("uses hidden tht argv and sends record bytes only over bounded stdin", async () => { + const calls: Array<{ executable: string; args: readonly string[]; input: Buffer; timeoutMs: number }> = []; + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\Program Files\\ThothII\\tht.exe", + invoke: async (call) => { + calls.push(call); + return { code: 0, stdout: Buffer.from('{"version":1,"ok":true,"created":true}\n'), stderr: Buffer.alloc(0) }; + }, + }); + + await expect(bridge.create(root, "sessions", filename, Buffer.from('{"subject":"record-data"}'))).resolves.toBe(true); + expect(calls).toHaveLength(1); + expect(calls[0]).toMatchObject({ + executable: "C:\\Program Files\\ThothII\\tht.exe", + args: ["_auth-storage"], + }); + expect(JSON.stringify(calls[0].args)).not.toContain("record-data"); + expect(JSON.parse(calls[0].input.toString("utf8"))).toMatchObject({ + version: 1, + operation: "create", + root, + directory: "sessions", + filename, + contentBase64: Buffer.from('{"subject":"record-data"}').toString("base64"), + }); + expect(calls[0].timeoutMs).toBeGreaterThan(0); + }); + + 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) } }, + { label: "unexpected stdout", result: { code: 0, stdout: Buffer.from('{"version":1,"ok":true}\nextra'), stderr: Buffer.alloc(0) } }, + { label: "unexpected JSON field", result: { code: 0, stdout: Buffer.from('{"version":1,"ok":true,"created":true,"detail":"secret"}\n'), stderr: Buffer.alloc(0) } }, + ])("fails closed on $label bridge output", async ({ result }) => { + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\tht.exe", + invoke: async () => result, + }); + + await expect(bridge.create(root, "sessions", filename, Buffer.from("record"))) + .rejects.toThrow("auth_session_store_invalid"); + }); + + test("fails closed on a bridge timeout without disclosing request content", async () => { + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\tht.exe", + invoke: async () => { throw new Error("timeout secret-record"); }, + }); + + await expect(bridge.create(root, "sessions", filename, Buffer.from("secret-record"))) + .rejects.toThrow("auth_session_store_invalid"); + }); + + test("rejects a claimed-read response without bounded record bytes", async () => { + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\tht.exe", + invoke: async () => ({ code: 0, stdout: Buffer.from('{"version":1,"ok":true,"found":true}\n'), stderr: Buffer.alloc(0) }), + }); + + await expect(bridge.readClaim(root, filename)).rejects.toThrow("auth_session_store_invalid"); + }); + + test("rejects an executable value that would require shell parsing", () => { + expect(() => createWindowsAuthStorageBridge({ thtExecutable: "tht.exe && unexpected" })) + .toThrow("auth_session_store_invalid"); + }); +}); diff --git a/tools/tht/cmd/tht/main.go b/tools/tht/cmd/tht/main.go index a0041124..4f03019e 100644 --- a/tools/tht/cmd/tht/main.go +++ b/tools/tht/cmd/tht/main.go @@ -15,6 +15,7 @@ import ( "strings" "github.com/aritmolab/thothii/tools/tht/internal/authconfig" + "github.com/aritmolab/thothii/tools/tht/internal/authstorage" "github.com/aritmolab/thothii/tools/tht/internal/backup" "github.com/aritmolab/thothii/tools/tht/internal/compose" "github.com/aritmolab/thothii/tools/tht/internal/config" @@ -99,6 +100,9 @@ func main() { } func run(ctx context.Context, args []string, stdout, stderr io.Writer) int { + if len(args) > 0 && args[0] == "_auth-storage" { + return authstorage.Run(ctx, args[1:], os.Stdin, stdout, stderr) + } if len(args) == 1 && (args[0] == "--help" || args[0] == "-h") { fmt.Fprint(stdout, usage) return 0 diff --git a/tools/tht/internal/authstorage/storage.go b/tools/tht/internal/authstorage/storage.go new file mode 100644 index 00000000..e3e749d8 --- /dev/null +++ b/tools/tht/internal/authstorage/storage.go @@ -0,0 +1,291 @@ +// Package authstorage implements tht's hidden, stdin/stdout-only bridge for protected browser +// session files. It deliberately has no operator-facing commands or installation configuration. +package authstorage + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "io" + "os" + "path/filepath" + "regexp" + "strings" + + "github.com/aritmolab/thothii/tools/tht/internal/safeio" +) + +const ( + protocolVersion = 1 + maximumProtocolBytes = 64 * 1024 + maximumSessionBytes = 16 * 1024 + maximumOIDCStateBytes = 8 * 1024 + maximumEntries = 256 +) + +var ( + digestFilename = regexp.MustCompile(`^[a-f0-9]{64}\.json$`) + claimFilename = regexp.MustCompile(`^[a-f0-9]{64}\.claim$`) + errInvalid = errors.New("auth storage request invalid") +) + +type request struct { + Version int `json:"version"` + Operation string `json:"operation"` + Root string `json:"root"` + Directory string `json:"directory"` + Filename string `json:"filename,omitempty"` + ContentBase64 string `json:"contentBase64,omitempty"` +} + +type response struct { + Version int `json:"version"` + OK bool `json:"ok"` + Created bool `json:"created,omitempty"` + Replaced bool `json:"replaced,omitempty"` + Removed bool `json:"removed,omitempty"` + Found bool `json:"found,omitempty"` + Claimed bool `json:"claimed,omitempty"` + ContentBase64 string `json:"contentBase64,omitempty"` + Entries []safeio.PrivateDirectoryEntry `json:"entries,omitempty"` +} + +// Run accepts exactly one strict JSON request on stdin and emits exactly one JSON response on +// stdout. All diagnostic text is fixed and goes only to stderr. +func Run(ctx context.Context, args []string, stdin io.Reader, stdout, stderr io.Writer) int { + if len(args) != 0 || ctx == nil || stdin == nil || stdout == nil || stderr == nil { + return fail(stderr) + } + if err := ctx.Err(); err != nil { + return fail(stderr) + } + decoder := json.NewDecoder(io.LimitReader(stdin, maximumProtocolBytes+1)) + decoder.DisallowUnknownFields() + var input request + if err := decoder.Decode(&input); err != nil { + return fail(stderr) + } + var extra any + if err := decoder.Decode(&extra); err != io.EOF { + return fail(stderr) + } + if err := ctx.Err(); err != nil { + return fail(stderr) + } + result, err := execute(input) + if err != nil || ctx.Err() != nil { + return fail(stderr) + } + encoder := json.NewEncoder(stdout) + encoder.SetEscapeHTML(false) + if err := encoder.Encode(result); err != nil { + return fail(stderr) + } + return 0 +} + +func fail(stderr io.Writer) int { + _, _ = io.WriteString(stderr, "tht: auth storage request failed\n") + return 1 +} + +func execute(input request) (response, error) { + if input.Version != protocolVersion || !validDirectory(input.Directory) { + return response{}, errInvalid + } + directory, err := storageDirectory(input.Root, input.Directory) + if err != nil { + return response{}, errInvalid + } + switch input.Operation { + case "create": + contents, err := decodeContents(input) + if err != nil || !digestFilename.MatchString(input.Filename) { + return response{}, errInvalid + } + created, err := createPrivate(directory, input.Filename, contents) + if err != nil { + return response{}, errInvalid + } + return response{Version: protocolVersion, OK: true, Created: created}, nil + case "read": + if !digestFilename.MatchString(input.Filename) { + return response{}, errInvalid + } + contents, found, err := readPrivate(directory, input.Filename, recordMaximum(input.Directory)) + if err != nil { + return response{}, errInvalid + } + return contentResponse(found, contents), nil + case "replace": + contents, err := decodeContents(input) + if err != nil || !digestFilename.MatchString(input.Filename) { + return response{}, errInvalid + } + if err := safeio.ReplaceCanonicalRegular(filepath.Join(directory, input.Filename), contents, 0o600); err != nil { + return response{}, errInvalid + } + return response{Version: protocolVersion, OK: true, Replaced: true}, nil + case "remove": + if !digestFilename.MatchString(input.Filename) && !claimFilename.MatchString(input.Filename) { + return response{}, errInvalid + } + removed, err := removePrivate(directory, input.Filename) + if err != nil { + return response{}, errInvalid + } + return response{Version: protocolVersion, OK: true, Removed: removed}, nil + case "list": + entries, err := safeio.ListCanonicalPrivateDirectory(directory, maximumEntries) + if err != nil { + return response{}, errInvalid + } + for _, entry := range entries { + if !digestFilename.MatchString(entry.Name) && !claimFilename.MatchString(entry.Name) { + return response{}, errInvalid + } + } + return response{Version: protocolVersion, OK: true, Entries: entries}, nil + case "claim-consume": + if input.Directory != "oidc" || !digestFilename.MatchString(input.Filename) { + return response{}, errInvalid + } + return claimConsume(directory, input.Filename) + case "read-claim": + if input.Directory != "oidc" || !digestFilename.MatchString(input.Filename) { + return response{}, errInvalid + } + contents, found, err := safeio.ReadCanonicalPrivateClaim( + filepath.Join(directory, input.Filename), + filepath.Join(directory, asClaimFilename(input.Filename)), + maximumOIDCStateBytes, + ) + if err != nil { + return response{}, errInvalid + } + return contentResponse(found, contents), nil + case "remove-claim": + if input.Directory != "oidc" || !digestFilename.MatchString(input.Filename) { + return response{}, errInvalid + } + removed, err := safeio.RemoveCanonicalPrivateClaim( + filepath.Join(directory, input.Filename), + filepath.Join(directory, asClaimFilename(input.Filename)), + ) + if err != nil { + return response{}, errInvalid + } + return response{Version: protocolVersion, OK: true, Removed: removed}, nil + default: + return response{}, errInvalid + } +} + +func contentResponse(found bool, contents []byte) response { + if !found { + return response{Version: protocolVersion, OK: true} + } + return response{Version: protocolVersion, OK: true, Found: true, ContentBase64: base64.StdEncoding.EncodeToString(contents)} +} + +func storageDirectory(root, directory string) (string, error) { + if !filepath.IsAbs(root) || filepath.Clean(root) != root || strings.ContainsRune(root, '\x00') || 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" +} + +func recordMaximum(directory string) int64 { + if directory == "oidc" { + return maximumOIDCStateBytes + } + return maximumSessionBytes +} + +func decodeContents(input request) ([]byte, error) { + contents, err := base64.StdEncoding.DecodeString(input.ContentBase64) + if err != nil || len(contents) == 0 || int64(len(contents)) > recordMaximum(input.Directory) { + return nil, errInvalid + } + 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) { + return false, errInvalid + } + if err := safeio.WriteCanonicalNewFile(path, contents, 0o600); err == nil { + return true, nil + } + if safeio.ValidatePrivateRegular(path) == nil { + return false, nil + } + return false, errInvalid +} + +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) + if err != nil { + return nil, false, errInvalid + } + return contents, true, 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 { + return false, errInvalid + } + if err := safeio.RemoveCanonicalPrivateRegular(path); err != nil { + return false, errInvalid + } + return true, 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) + if err != nil { + return response{}, errInvalid + } + if !claimed { + return response{Version: protocolVersion, OK: true}, nil + } + contents, found, err := safeio.ReadCanonicalPrivateClaim(source, claim, maximumOIDCStateBytes) + if err != nil || !found { + return response{}, errInvalid + } + removed, err := safeio.RemoveCanonicalPrivateClaim(source, claim) + if err != nil || !removed { + return response{}, errInvalid + } + return contentResponse(true, contents), nil +} + +func asClaimFilename(filename string) string { + return strings.TrimSuffix(filename, ".json") + ".claim" +} diff --git a/tools/tht/internal/authstorage/storage_test.go b/tools/tht/internal/authstorage/storage_test.go new file mode 100644 index 00000000..2ffb98fd --- /dev/null +++ b/tools/tht/internal/authstorage/storage_test.go @@ -0,0 +1,164 @@ +package authstorage + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "os" + "path/filepath" + "sync" + "testing" + + "github.com/aritmolab/thothii/tools/tht/internal/safeio" + "github.com/aritmolab/thothii/tools/tht/internal/testsupport" +) + +func TestProtocolCreatesReadsReplacesListsAndRemovesPrivateRecord(t *testing.T) { + root := filepath.Join(privateTestRoot(t), "auth") + filename := "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa.json" + + created := runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("first"))}) + if !created.Created { + t.Fatal("create did not report a new record") + } + if duplicate := runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("other"))}); duplicate.Created { + t.Fatal("duplicate exclusive create reported success") + } + read := runRequest(t, request{Version: 1, Operation: "read", Root: root, Directory: "sessions", Filename: filename}) + if !read.Found || decodeContent(t, read) != "first" { + t.Fatalf("read = %#v, want private first record", read) + } + updated := runRequest(t, request{Version: 1, Operation: "replace", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("second"))}) + if !updated.Replaced { + t.Fatal("replace did not report success") + } + listed := runRequest(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions"}) + if len(listed.Entries) != 1 || listed.Entries[0].Name != filename { + t.Fatalf("list = %#v, want exactly %q", listed.Entries, filename) + } + removed := runRequest(t, request{Version: 1, Operation: "remove", Root: root, Directory: "sessions", Filename: filename}) + if !removed.Removed { + t.Fatal("remove did not report success") + } +} + +func TestProtocolClaimConsumeIsAtomicAcrossConcurrentRequests(t *testing.T) { + root := filepath.Join(privateTestRoot(t), "auth") + filename := "dddddddddddddddddddddddddddddddddddddddddddddddddddddddddddddddd.json" + runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "oidc", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("oidc-record"))}) + + request := request{Version: 1, Operation: "claim-consume", Root: root, Directory: "oidc", Filename: filename} + results := make(chan response, 2) + errors := make(chan error, 2) + var group sync.WaitGroup + for range 2 { + group.Add(1) + go func() { + defer group.Done() + result, err := execute(request) + if err != nil { + errors <- err + return + } + results <- result + }() + } + group.Wait() + close(results) + close(errors) + for err := range errors { + t.Fatalf("concurrent claim error = %v", err) + } + found := 0 + for result := range results { + if result.Found { + found++ + } + } + if found != 1 { + t.Fatalf("winning claim count = %d, want 1", found) + } +} + +func TestProtocolRejectsBoundsReparseAndUnexpectedStorageNames(t *testing.T) { + root := filepath.Join(privateTestRoot(t), "auth") + filename := "eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee.json" + runRejected(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString(make([]byte, maximumSessionBytes+1))}) + + valid := request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))} + runRequest(t, valid) + if err := safeio.WriteCanonicalNewFile(filepath.Join(root, "sessions", "unexpected.txt"), []byte("junk"), 0o600); err != nil { + t.Fatal(err) + } + runRejected(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions"}) + + outer := privateTestRoot(t) + linkedRoot := filepath.Join(outer, "linked-auth") + testsupport.SymlinkOrSkip(t, filepath.Join(outer, "missing-real-auth"), linkedRoot) + runRejected(t, request{Version: 1, Operation: "create", Root: linkedRoot, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))}) +} + +func TestProtocolClaimConsumeHasOneWinner(t *testing.T) { + root := filepath.Join(privateTestRoot(t), "auth") + filename := "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb.json" + runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "oidc", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("oidc-record"))}) + + first := runRequest(t, request{Version: 1, Operation: "claim-consume", Root: root, Directory: "oidc", Filename: filename}) + second := runRequest(t, request{Version: 1, Operation: "claim-consume", Root: root, Directory: "oidc", Filename: filename}) + if !first.Found || decodeContent(t, first) != "oidc-record" || second.Found { + t.Fatalf("claim results = %#v, %#v", first, second) + } +} + +func runRequest(t *testing.T, value request) response { + t.Helper() + input, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + var stdout, stderr bytes.Buffer + if code := Run(context.Background(), nil, bytes.NewReader(input), &stdout, &stderr); code != 0 { + t.Fatalf("Run() code = %d stderr = %q", code, stderr.String()) + } + var output response + if err := json.Unmarshal(stdout.Bytes(), &output); err != nil || !output.OK || output.Version != 1 { + t.Fatalf("stdout = %q response = %#v error = %v", stdout.String(), output, err) + } + return output +} + +func runRejected(t *testing.T, value request) { + t.Helper() + input, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + var stdout, stderr bytes.Buffer + if code := Run(context.Background(), nil, bytes.NewReader(input), &stdout, &stderr); code == 0 || stdout.Len() != 0 || stderr.String() != "tht: auth storage request failed\n" { + t.Fatalf("Run() rejection code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String()) + } +} + +func decodeContent(t *testing.T, value response) string { + t.Helper() + contents, err := base64.StdEncoding.DecodeString(value.ContentBase64) + if err != nil { + t.Fatal(err) + } + return string(contents) +} + +func privateTestRoot(t *testing.T) string { + t.Helper() + temporary, err := filepath.EvalSymlinks(os.TempDir()) + if err != nil { + t.Fatal(err) + } + root, err := os.MkdirTemp(temporary, "tht-authstorage-") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.RemoveAll(root) }) + return root +} diff --git a/tools/tht/internal/authstorage/storage_windows_test.go b/tools/tht/internal/authstorage/storage_windows_test.go new file mode 100644 index 00000000..c296ffed --- /dev/null +++ b/tools/tht/internal/authstorage/storage_windows_test.go @@ -0,0 +1,71 @@ +//go:build windows + +package authstorage + +import ( + "encoding/base64" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/aritmolab/thothii/tools/tht/internal/safeio" + "golang.org/x/sys/windows" +) + +func TestProtocolCreatesRecordsWithOwnerOnlyDACLAndRejectsReparseRoot(t *testing.T) { + root := filepath.Join(t.TempDir(), "auth") + filename := "cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc.json" + runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("private"))}) + if err := safeio.ValidatePrivateDirectory(root); err != nil { + t.Fatalf("root DACL = %v", err) + } + if err := safeio.ValidatePrivateDirectory(filepath.Join(root, "sessions")); err != nil { + t.Fatalf("sessions DACL = %v", err) + } + if err := safeio.ValidatePrivateRegular(filepath.Join(root, "sessions", filename)); err != nil { + t.Fatalf("record DACL = %v", err) + } +} + +func TestProtocolRejectsPermissiveDACLAndReparseRoot(t *testing.T) { + root := filepath.Join(t.TempDir(), "auth") + filename := "ffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff.json" + runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("private"))}) + path := filepath.Join(root, "sessions", filename) + if err := setPermissiveDACL(path); err != nil { + t.Fatal(err) + } + runRejected(t, request{Version: 1, Operation: "read", Root: root, Directory: "sessions", Filename: filename}) + + linkedRoot := filepath.Join(t.TempDir(), "reparse-auth") + if err := os.Symlink(root, linkedRoot); err != nil { + t.Skipf("Windows host does not permit test symlink creation: %v", err) + } + runRejected(t, request{Version: 1, Operation: "create", Root: linkedRoot, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("private"))}) +} + +func setPermissiveDACL(path string) error { + world, err := windows.StringToSid("S-1-1-0") + if err != nil { + return err + } + var pinner runtime.Pinner + pinner.Pin(world) + defer pinner.Unpin() + acl, err := windows.ACLFromEntries([]windows.EXPLICIT_ACCESS{{ + AccessPermissions: windows.GENERIC_READ | windows.GENERIC_WRITE, + AccessMode: windows.GRANT_ACCESS, + Trustee: windows.TRUSTEE{ + TrusteeForm: windows.TRUSTEE_IS_SID, + TrusteeType: windows.TRUSTEE_IS_GROUP, + TrusteeValue: windows.TrusteeValueFromSID(world), + }, + }}, nil) + if err != nil { + return err + } + return windows.SetNamedSecurityInfo(path, windows.SE_FILE_OBJECT, + windows.DACL_SECURITY_INFORMATION|windows.PROTECTED_DACL_SECURITY_INFORMATION, + nil, nil, acl, nil) +} diff --git a/tools/tht/internal/safeio/claim_unix.go b/tools/tht/internal/safeio/claim_unix.go new file mode 100644 index 00000000..b6832bf8 --- /dev/null +++ b/tools/tht/internal/safeio/claim_unix.go @@ -0,0 +1,169 @@ +//go:build !windows + +package safeio + +import ( + "errors" + "io" + "os" + + "golang.org/x/sys/unix" +) + +func claimCanonicalPrivateRegular(source, claim string) (bool, error) { + directory, sourceName, err := openCanonicalParentDirectory(source) + if err != nil { + return false, ErrUnsafeFile + } + defer unix.Close(directory) + claimName := claim[len(claim)-len(claimBaseName(claim)):] + if err := requirePrivateRegularAtLinks(directory, sourceName, 1); err != nil { + if requireSamePrivatePairAt(directory, sourceName, claimName) == nil { + return false, nil + } + if isPrivateClaimAbsentOrOrphanAt(directory, sourceName, claimName) { + return false, nil + } + return false, ErrUnsafeFile + } + if err := unix.Linkat(directory, sourceName, directory, claimName, 0); err != nil { + if errors.Is(err, unix.EEXIST) && privateRegularAtAllowedLinks(directory, claimName, 1, 2) == nil { + return false, nil + } + return false, ErrUnsafeFile + } + if err := requireSamePrivatePairAt(directory, sourceName, claimName); err != nil { + return false, ErrUnsafeFile + } + if err := unix.Fsync(directory); err != nil { + return false, ErrUnsafeFile + } + return true, nil +} + +func readCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) { + directory, sourceName, err := openCanonicalParentDirectory(source) + if err != nil { + return nil, false, ErrUnsafeFile + } + defer unix.Close(directory) + claimName := claim[len(claim)-len(claimBaseName(claim)):] + if err := requireSamePrivatePairAt(directory, sourceName, claimName); err != nil { + if isPrivateClaimAbsentOrOrphanAt(directory, sourceName, claimName) { + return nil, false, nil + } + return nil, false, ErrUnsafeFile + } + descriptor, err := unix.Openat(directory, sourceName, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0) + if err != nil { + return nil, false, ErrUnsafeFile + } + file := os.NewFile(uintptr(descriptor), "tht-safeio-oidc-claim") + if file == nil { + unix.Close(descriptor) + return nil, false, ErrUnsafeFile + } + defer file.Close() + var before unix.Stat_t + if err := unix.Fstat(descriptor, &before); err != nil || !privateRegularStatWithLinks(before, 2) { + return nil, false, ErrUnsafeFile + } + contents, err := io.ReadAll(io.LimitReader(file, maximum+1)) + if err != nil || int64(len(contents)) > maximum { + return nil, false, ErrUnsafeFile + } + var after unix.Stat_t + if err := unix.Fstat(descriptor, &after); err != nil || !sameUnixPrivateFile(before, after) || !privateRegularStatWithLinks(after, 2) { + return nil, false, ErrUnsafeFile + } + if err := requireSamePrivatePairAt(directory, sourceName, claimName); err != nil { + return nil, false, ErrUnsafeFile + } + return contents, true, nil +} + +func removeCanonicalPrivateClaim(source, claim string) (bool, error) { + directory, sourceName, err := openCanonicalParentDirectory(source) + if err != nil { + return false, ErrUnsafeFile + } + defer unix.Close(directory) + claimName := claim[len(claim)-len(claimBaseName(claim)):] + if err := requireSamePrivatePairAt(directory, sourceName, claimName); err != nil { + if isPrivateClaimAbsentOrOrphanAt(directory, sourceName, claimName) { + return false, nil + } + return false, ErrUnsafeFile + } + if err := unix.Unlinkat(directory, sourceName, 0); err != nil { + return false, ErrUnsafeFile + } + if err := requirePrivateRegularAtLinks(directory, claimName, 1); err != nil { + return false, ErrUnsafeFile + } + if err := unix.Unlinkat(directory, claimName, 0); err != nil { + return false, ErrUnsafeFile + } + if err := unix.Fsync(directory); err != nil { + return false, ErrUnsafeFile + } + return true, nil +} + +func claimBaseName(path string) string { + for index := len(path) - 1; index >= 0; index-- { + if path[index] == byte(os.PathSeparator) { + return path[index+1:] + } + } + return path +} + +func privateRegularAtAllowedLinks(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 privateRegularStatWithLinks(stat, links) { + return nil + } + } + return ErrUnsafeFile +} + +func requirePrivateRegularAtLinks(directory int, name string, links uint64) error { + return privateRegularAtAllowedLinks(directory, name, links) +} + +func requireSamePrivatePairAt(directory int, source, claim string) error { + var left, right unix.Stat_t + if err := unix.Fstatat(directory, source, &left, unix.AT_SYMLINK_NOFOLLOW); err != nil || !privateRegularStatWithLinks(left, 2) { + return ErrUnsafeFile + } + if err := unix.Fstatat(directory, claim, &right, unix.AT_SYMLINK_NOFOLLOW); err != nil || !privateRegularStatWithLinks(right, 2) { + return ErrUnsafeFile + } + if !sameUnixPrivateFile(left, right) { + return ErrUnsafeFile + } + return nil +} + +func isPrivateClaimAbsentOrOrphanAt(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 && privateRegularStatWithLinks(claimStat, 1) +} + +func privateRegularStatWithLinks(stat unix.Stat_t, links uint64) bool { + return stat.Mode&unix.S_IFMT == unix.S_IFREG && uint64(stat.Nlink) == links && stat.Mode&0o7777 == 0o600 +} + +func sameUnixPrivateFile(left, right unix.Stat_t) bool { + return left.Dev == right.Dev && left.Ino == right.Ino && left.Size == right.Size && left.Mtim == right.Mtim +} diff --git a/tools/tht/internal/safeio/claim_windows.go b/tools/tht/internal/safeio/claim_windows.go new file mode 100644 index 00000000..e9f1be46 --- /dev/null +++ b/tools/tht/internal/safeio/claim_windows.go @@ -0,0 +1,210 @@ +//go:build windows + +package safeio + +import ( + "errors" + "io" + "os" + "path/filepath" + + "golang.org/x/sys/windows" +) + +type windowsPrivateRegular struct { + parents *windowsParentHandles + handle windows.Handle + path string + info windows.ByHandleFileInformation +} + +func (value *windowsPrivateRegular) Close() { + if value.handle != 0 { + _ = windows.CloseHandle(value.handle) + value.handle = 0 + } + if value.parents != nil { + value.parents.Close() + value.parents = nil + } +} + +func openWindowsPrivateRegular(path string, links uint32) (*windowsPrivateRegular, error) { + parents, target, err := openCanonicalWindowsParent(path) + if err != nil || len(parents.handles) == 0 || validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil { + if parents != nil { + parents.Close() + } + return nil, ErrUnsafeFile + } + fullPath := filepath.Join(parents.directory, target) + handle, err := windows.CreateFile( + windows.StringToUTF16Ptr(fullPath), + windows.GENERIC_READ, + windowsRetainedHandleShareMode, + nil, + windows.OPEN_EXISTING, + windows.FILE_FLAG_OPEN_REPARSE_POINT|windows.FILE_ATTRIBUTE_NORMAL, + 0, + ) + if err != nil { + parents.Close() + return nil, err + } + var info windows.ByHandleFileInformation + if err := windows.GetFileInformationByHandle(handle, &info); err != nil || info.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 || info.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY != 0 || info.NumberOfLinks != links || validateOwnerOnlyDACL(handle) != nil { + _ = windows.CloseHandle(handle) + parents.Close() + return nil, ErrUnsafeFile + } + return &windowsPrivateRegular{parents: parents, handle: handle, path: fullPath, info: info}, nil +} + +func claimCanonicalPrivateRegular(source, claim string) (bool, error) { + sourceFile, err := openWindowsPrivateRegular(source, 1) + if err != nil { + if windowsPrivateClaimPairExists(source, claim) { + return false, nil + } + if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) { + return false, nil + } + return false, ErrUnsafeFile + } + defer sourceFile.Close() + if err := windows.CreateHardLink(windows.StringToUTF16Ptr(claim), windows.StringToUTF16Ptr(sourceFile.path), 0); err != nil { + if errors.Is(err, windows.ERROR_FILE_EXISTS) || errors.Is(err, windows.ERROR_ALREADY_EXISTS) { + if existing, existingErr := openWindowsPrivateRegularWithAllowedLinks(claim, 1, 2); existingErr == nil { + existing.Close() + return false, nil + } + } + return false, ErrUnsafeFile + } + if err := windows.GetFileInformationByHandle(sourceFile.handle, &sourceFile.info); err != nil || sourceFile.info.NumberOfLinks != 2 { + return false, ErrUnsafeFile + } + claimFile, err := openWindowsPrivateRegular(claim, 2) + if err != nil { + return false, ErrUnsafeFile + } + defer claimFile.Close() + if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { + return false, ErrUnsafeFile + } + return true, nil +} + +func readCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) { + sourceFile, err := openWindowsPrivateRegular(source, 2) + if err != nil { + if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) { + return nil, false, nil + } + return nil, false, ErrUnsafeFile + } + defer sourceFile.Close() + claimFile, err := openWindowsPrivateRegular(claim, 2) + if err != nil { + if isWindowsNotFound(err) { + return nil, false, nil + } + return nil, false, ErrUnsafeFile + } + defer claimFile.Close() + if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { + return nil, false, ErrUnsafeFile + } + file := os.NewFile(uintptr(sourceFile.handle), "tht-safeio-oidc-claim") + if file == nil { + return nil, false, ErrUnsafeFile + } + contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1)) + file.Close() + sourceFile.handle = 0 + if readErr != nil || int64(len(contents)) > maximum { + return nil, false, ErrUnsafeFile + } + if err := windows.GetFileInformationByHandle(claimFile.handle, &claimFile.info); err != nil || claimFile.info.NumberOfLinks != 2 || !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { + return nil, false, ErrUnsafeFile + } + return contents, true, nil +} + +func removeCanonicalPrivateClaim(source, claim string) (bool, error) { + sourceFile, err := openWindowsPrivateRegular(source, 2) + if err != nil { + if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) { + return false, nil + } + return false, ErrUnsafeFile + } + claimFile, err := openWindowsPrivateRegular(claim, 2) + if err != nil { + sourceFile.Close() + if isWindowsNotFound(err) { + return false, nil + } + return false, ErrUnsafeFile + } + if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) { + sourceFile.Close() + claimFile.Close() + return false, ErrUnsafeFile + } + sourceFile.Close() + claimFile.Close() + if err := windows.DeleteFile(windows.StringToUTF16Ptr(source)); err != nil { + return false, ErrUnsafeFile + } + remaining, err := openWindowsPrivateRegular(claim, 1) + if err != nil { + return false, ErrUnsafeFile + } + remaining.Close() + if err := windows.DeleteFile(windows.StringToUTF16Ptr(claim)); err != nil { + return false, ErrUnsafeFile + } + return true, nil +} + +func openWindowsPrivateRegularWithAllowedLinks(path string, allowed ...uint32) (*windowsPrivateRegular, error) { + for _, links := range allowed { + value, err := openWindowsPrivateRegular(path, links) + if err == nil { + return value, nil + } + } + return nil, ErrUnsafeFile +} + +func windowsPrivateClaimAbsentOrOrphan(claim string) bool { + claimFile, err := openWindowsPrivateRegular(claim, 1) + if err != nil { + return isWindowsNotFound(err) + } + claimFile.Close() + return true +} + +func windowsPrivateClaimPairExists(source, claim string) bool { + sourceFile, sourceErr := openWindowsPrivateRegular(source, 2) + if sourceErr != nil { + return false + } + defer sourceFile.Close() + claimFile, claimErr := openWindowsPrivateRegular(claim, 2) + if claimErr != nil { + return false + } + defer claimFile.Close() + return sameWindowsPrivateFile(sourceFile.info, claimFile.info) +} + +func isWindowsNotFound(err error) bool { + return errors.Is(err, windows.ERROR_FILE_NOT_FOUND) || errors.Is(err, windows.ERROR_PATH_NOT_FOUND) +} + +func sameWindowsPrivateFile(left, right windows.ByHandleFileInformation) bool { + return left.VolumeSerialNumber == right.VolumeSerialNumber && left.FileIndexHigh == right.FileIndexHigh && left.FileIndexLow == right.FileIndexLow && left.FileSizeHigh == right.FileSizeHigh && left.FileSizeLow == right.FileSizeLow && left.LastWriteTime == right.LastWriteTime +} diff --git a/tools/tht/internal/safeio/files.go b/tools/tht/internal/safeio/files.go index 4800c24a..4a5ca745 100644 --- a/tools/tht/internal/safeio/files.go +++ b/tools/tht/internal/safeio/files.go @@ -71,6 +71,50 @@ func ReadCanonicalUTF8(path string, maximum int64) (string, error) { return string(contents), nil } +// ReadCanonicalPrivateRegular reads one owner-only private record and revalidates its metadata +// after the bounded read. It intentionally rejects ordinary hard links; OIDC's explicitly named +// atomic claim pair uses ReadCanonicalPrivateClaim instead. +func ReadCanonicalPrivateRegular(path string, maximum int64) ([]byte, error) { + return readCanonicalPrivateRegular(path, maximum) +} + +// PrivateDirectoryEntry is a bounded, untrusted directory listing item. Callers must still +// validate each filename and record before using it. +type PrivateDirectoryEntry struct { + Name string + ModifiedUnixMs int64 +} + +// ListCanonicalPrivateDirectory lists regular, non-symlinked direct children from an owner-only +// directory. It returns no content and bounds the number of entries before allocating output. +func ListCanonicalPrivateDirectory(path string, maximumEntries int) ([]PrivateDirectoryEntry, error) { + if maximumEntries < 1 || maximumEntries > 4096 { + return nil, ErrUnsafeFile + } + if err := ValidatePrivateDirectory(path); err != nil { + return nil, ErrUnsafeFile + } + entries, err := os.ReadDir(path) + if err != nil || len(entries) > maximumEntries { + return nil, ErrUnsafeFile + } + result := make([]PrivateDirectoryEntry, 0, len(entries)) + for _, entry := range entries { + if entry.Name() == "" || strings.Contains(entry.Name(), string(filepath.Separator)) { + return nil, ErrUnsafeFile + } + info, err := os.Lstat(filepath.Join(path, entry.Name())) + if err != nil || !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 { + return nil, ErrUnsafeFile + } + result = append(result, PrivateDirectoryEntry{Name: entry.Name(), ModifiedUnixMs: info.ModTime().UnixMilli()}) + } + if err := ValidatePrivateDirectory(path); err != nil { + return nil, ErrUnsafeFile + } + return result, nil +} + func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error { if err := ValidateCanonicalPath(path); err != nil { return err @@ -127,6 +171,42 @@ func RemoveCanonicalPrivateRegular(path string) error { return removeCanonicalPrivateRegular(path) } +// ClaimCanonicalPrivateRegular atomically creates a second, explicit private hard link to one +// existing record. It is used only for digest-named, single-use OIDC state claims. +func ClaimCanonicalPrivateRegular(source, claim string) (bool, error) { + if err := validateClaimPaths(source, claim); err != nil { + return false, ErrUnsafeFile + } + return claimCanonicalPrivateRegular(source, claim) +} + +// ReadCanonicalPrivateClaim reads a verified two-link source/claim pair. found=false means the +// state has already been consumed or a winning process is between its two removal steps. +func ReadCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) { + if maximum < 0 || maximum == int64(^uint64(0)>>1) || validateClaimPaths(source, claim) != nil { + return nil, false, ErrUnsafeFile + } + return readCanonicalPrivateClaim(source, claim, maximum) +} + +// RemoveCanonicalPrivateClaim removes exactly a verified two-link source/claim pair. +func RemoveCanonicalPrivateClaim(source, claim string) (bool, error) { + if err := validateClaimPaths(source, claim); err != nil { + return false, ErrUnsafeFile + } + return removeCanonicalPrivateClaim(source, claim) +} + +func validateClaimPaths(source, claim string) error { + if ValidateCanonicalPath(source) != nil || ValidateCanonicalPath(claim) != nil || filepath.Dir(source) != filepath.Dir(claim) { + return ErrUnsafeFile + } + if err := ValidatePrivateDirectory(filepath.Dir(source)); err != nil { + return ErrUnsafeFile + } + return nil +} + func randomTemporaryName() (string, error) { bytes := make([]byte, 16) if _, err := rand.Read(bytes); err != nil { diff --git a/tools/tht/internal/safeio/private_read_unix.go b/tools/tht/internal/safeio/private_read_unix.go new file mode 100644 index 00000000..b6ff17ba --- /dev/null +++ b/tools/tht/internal/safeio/private_read_unix.go @@ -0,0 +1,17 @@ +//go:build !windows + +package safeio + +func readCanonicalPrivateRegular(path string, maximum int64) ([]byte, error) { + if err := ValidatePrivateRegular(path); err != nil { + return nil, ErrUnsafeFile + } + contents, err := ReadCanonicalRegular(path, maximum) + if err != nil { + return nil, ErrUnsafeFile + } + if err := ValidatePrivateRegular(path); err != nil { + return nil, ErrUnsafeFile + } + return contents, nil +} diff --git a/tools/tht/internal/safeio/private_read_windows.go b/tools/tht/internal/safeio/private_read_windows.go new file mode 100644 index 00000000..35267712 --- /dev/null +++ b/tools/tht/internal/safeio/private_read_windows.go @@ -0,0 +1,44 @@ +//go:build windows + +package safeio + +import ( + "io" + "os" + + "golang.org/x/sys/windows" +) + +func readCanonicalPrivateRegular(path string, maximum int64) ([]byte, error) { + if maximum < 0 || maximum == int64(^uint64(0)>>1) { + return nil, ErrUnsafeFile + } + value, err := openWindowsPrivateRegular(path, 1) + if err != nil { + return nil, ErrUnsafeFile + } + defer value.Close() + file := os.NewFile(uintptr(value.handle), "tht-safeio-private-read") + if file == nil { + return nil, ErrUnsafeFile + } + before := value.info + contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1)) + var after = value.info + identityErr := windows.GetFileInformationByHandle(value.handle, &after) + daclErr := validateOwnerOnlyDACL(value.handle) + closeErr := file.Close() + value.handle = 0 + if readErr != nil || identityErr != nil || daclErr != nil || !sameWindowsPrivateFile(before, after) || after.NumberOfLinks != 1 || closeErr != nil || int64(len(contents)) > maximum { + return nil, ErrUnsafeFile + } + current, err := openWindowsPrivateRegular(path, 1) + if err != nil { + return nil, ErrUnsafeFile + } + current.Close() + if !sameWindowsPrivateFile(before, current.info) { + return nil, ErrUnsafeFile + } + return contents, nil +} diff --git a/tools/tht/internal/safeio/private_windows.go b/tools/tht/internal/safeio/private_windows.go index fb2ad335..03864260 100644 --- a/tools/tht/internal/safeio/private_windows.go +++ b/tools/tht/internal/safeio/private_windows.go @@ -96,6 +96,14 @@ func ProtectPrivateRegular(path string) error { // createCanonicalNewPrivateFile installs the owner-only protected DACL in the CreateFile call, so // another mutation can never observe a newly-created lock with an inherited/default DACL. func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, error) { + parents, target, err := openCanonicalWindowsParent(path) + if err != nil || len(parents.handles) == 0 || validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil { + if parents != nil { + parents.Close() + } + return nil, ErrUnsafeFile + } + defer parents.Close() security, err := newOwnerOnlySecurityDescriptor() if err != nil { return nil, ErrUnsafeFile @@ -106,7 +114,7 @@ func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, err SecurityDescriptor: security.descriptor, } handle, err := windows.CreateFile( - windows.StringToUTF16Ptr(path), + windows.StringToUTF16Ptr(filepath.Join(parents.directory, target)), windows.GENERIC_WRITE, windowsRetainedHandleShareMode, attributes,