fix(auth): add Windows session storage bridge
This commit is contained in:
@@ -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<OidcStateRecord | undefined>;
|
||||
}
|
||||
|
||||
/** 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<T>(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<OidcStateRecord>;
|
||||
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<CreatedAuthSession> {
|
||||
const nowMs = dateMilliseconds(now);
|
||||
let validated: z.infer<typeof sessionInputSchema>;
|
||||
@@ -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<number> {
|
||||
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 (
|
||||
|
||||
@@ -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<boolean>;
|
||||
read(root: string, directory: WindowsAuthStorageDirectory, filename: string): Promise<Buffer | undefined>;
|
||||
replace(root: string, directory: WindowsAuthStorageDirectory, filename: string, contents: Buffer): Promise<void>;
|
||||
remove(root: string, directory: WindowsAuthStorageDirectory, filename: string): Promise<boolean>;
|
||||
list(root: string, directory: WindowsAuthStorageDirectory): Promise<WindowsAuthStorageEntry[]>;
|
||||
claimConsume(root: string, filename: string): Promise<Buffer | undefined>;
|
||||
readClaim(root: string, filename: string): Promise<Buffer | undefined>;
|
||||
removeClaim(root: string, filename: string): Promise<boolean>;
|
||||
}
|
||||
|
||||
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<WindowsAuthStorageInvocationResult>;
|
||||
/** 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<typeof responseSchema>;
|
||||
|
||||
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<WindowsAuthStorageInvocationResult> {
|
||||
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<typeof spawn>;
|
||||
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<BridgeResponse> => {
|
||||
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;
|
||||
},
|
||||
};
|
||||
}
|
||||
@@ -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<string, Buffer>();
|
||||
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;
|
||||
|
||||
@@ -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");
|
||||
});
|
||||
});
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user