diff --git a/backend/src/auth/windows-auth-storage.ts b/backend/src/auth/windows-auth-storage.ts index 3c9c62fc..c6b21738 100644 --- a/backend/src/auth/windows-auth-storage.ts +++ b/backend/src/auth/windows-auth-storage.ts @@ -1,5 +1,6 @@ import { spawn } from "node:child_process"; import { win32 } from "node:path"; +import type { Readable, Writable } from "node:stream"; import { z } from "zod"; const PROTOCOL_VERSION = 1; @@ -46,11 +47,29 @@ export interface WindowsAuthStorageInvocationResult { stderr: Buffer; } +interface WindowsAuthStorageChild { + readonly stdin: Writable | null; + readonly stdout: Readable | null; + readonly stderr: Readable | null; + kill(signal?: NodeJS.Signals | number): boolean; + on(event: "error", listener: (error: Error) => void): this; + once(event: "error", listener: (error: Error) => void): this; + once(event: "close", listener: (code: number | null, signal: NodeJS.Signals | null) => void): this; +} + +type WindowsAuthStorageSpawn = ( + executable: string, + args: readonly string[], + options: { shell: false; windowsHide: true; stdio: ["pipe", "pipe", "pipe"]; env: NodeJS.ProcessEnv }, +) => WindowsAuthStorageChild; + 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; + /** Test-only child-launch seam; production uses the fixed no-shell Node child-process launcher. */ + spawnChild?: WindowsAuthStorageSpawn; } const responseSchema = z.strictObject({ @@ -141,65 +160,111 @@ function environmentForBridge(): NodeJS.ProcessEnv { }; } -async function invokeTht(invocation: WindowsAuthStorageInvocation): Promise { +const spawnTht: WindowsAuthStorageSpawn = (executable, args, options) => spawn(executable, [...args], options); + +async function invokeTht( + invocation: WindowsAuthStorageInvocation, + spawnChild: WindowsAuthStorageSpawn = spawnTht, +): Promise { return new Promise((resolve, reject) => { let settled = false; + let closeSeen = false; + let aborted = false; + let closeCode: number | null = null; + let closeSignal: NodeJS.Signals | null = null; let timeout: NodeJS.Timeout | undefined; const stdout: Buffer[] = []; const stderr: Buffer[] = []; let stdoutBytes = 0; let stderrBytes = 0; + let child: WindowsAuthStorageChild | undefined; const settle = (callback: () => void): void => { if (settled) return; settled = true; if (timeout !== undefined) clearTimeout(timeout); callback(); }; - let child: ReturnType; + const stopStream = (stream: Writable | Readable | null | undefined): void => { + try { stream?.destroy(); } catch { /* abort is already fail-closed */ } + }; + const finishAfterClose = (): void => { + if (!closeSeen || settled) return; + if (aborted || typeof closeCode !== "number" || !Number.isInteger(closeCode) || closeSignal !== null) { + settle(() => reject(invalid())); + return; + } + const code = closeCode; + settle(() => resolve({ + code, + stdout: Buffer.concat(stdout), + stderr: Buffer.concat(stderr), + })); + }; + const abort = (): void => { + if (aborted || settled) return; + aborted = true; + if (timeout !== undefined) clearTimeout(timeout); + if (child !== undefined) { + stopStream(child.stdin); + stopStream(child.stdout); + stopStream(child.stderr); + try { child.kill(); } catch { /* the close handler still owns settlement */ } + } + finishAfterClose(); + }; try { - child = spawn(invocation.executable, [...invocation.args], { + child = spawnChild(invocation.executable, invocation.args, { shell: false, windowsHide: true, stdio: ["pipe", "pipe", "pipe"], env: environmentForBridge(), }); } catch { - reject(invalid()); + settle(() => reject(invalid())); return; } + child.once("close", (code, signal) => { + closeSeen = true; + closeCode = code; + closeSignal = signal; + if (code === null || signal !== null) aborted = true; + finishAfterClose(); + }); + child.on("error", abort); if (!child.stdin || !child.stdout || !child.stderr) { - try { child.kill(); } catch { /* unavailable child streams fail closed */ } - reject(invalid()); + abort(); 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())); + abort(); }, invocation.timeoutMs); - child.once("error", () => settle(() => reject(invalid()))); stdoutStream.on("data", (chunk: Buffer) => { + if (aborted) return; stdoutBytes += chunk.length; if (stdoutBytes > MAX_RESPONSE_BYTES) { - try { child.kill(); } catch { /* child failure is converted below */ } - settle(() => reject(invalid())); + abort(); return; } stdout.push(Buffer.from(chunk)); }); stderrStream.on("data", (chunk: Buffer) => { + if (aborted) return; stderrBytes += chunk.length; - if (stderrBytes <= MAX_RESPONSE_BYTES) stderr.push(Buffer.from(chunk)); + if (stderrBytes > MAX_RESPONSE_BYTES) { + abort(); + return; + } + 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); + stdin.once("error", abort); + try { + stdin.end(invocation.input); + } catch { + abort(); + } }); } @@ -214,7 +279,7 @@ function contentFrom(response: BridgeResponse, maximum: number): Buffer | undefi export function createWindowsAuthStorageBridge(options: WindowsAuthStorageBridgeOptions = {}): WindowsAuthStorageBridge { const executable = safeThtExecutable(options.thtExecutable); - const invoke = options.invoke ?? invokeTht; + const invoke = options.invoke ?? ((invocation: WindowsAuthStorageInvocation) => invokeTht(invocation, options.spawnChild)); const request = async (value: BridgeRequest): Promise => { try { const response = await invoke({ diff --git a/backend/test/auth-session-store.test.ts b/backend/test/auth-session-store.test.ts index 0ffcc142..b8754943 100644 --- a/backend/test/auth-session-store.test.ts +++ b/backend/test/auth-session-store.test.ts @@ -59,6 +59,7 @@ import { type AuthSessionStore, type SessionCreateInput, } from "../src/auth/session-store.js"; +import { createWindowsAuthStorageBridge } from "../src/auth/windows-auth-storage.js"; const roots: string[] = []; const base = new Date("2030-01-02T03:04:05.000Z"); @@ -579,6 +580,65 @@ describe("file-backed auth session store", () => { } }); + test("prunes nonempty Windows directories through the Go lower-camel list DTO", async () => { + const records = new Map(); + const key = (directory: string, entry: string) => `${directory}/${entry}`; + const response = (value: Record) => ({ + code: 0, + stdout: Buffer.from(`${JSON.stringify({ version: 1, ok: true, ...value })}\n`), + stderr: Buffer.alloc(0), + }); + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\tht.exe", + invoke: async ({ input }) => { + const request = JSON.parse(input.toString("utf8")) as { + operation: string; + directory: "sessions" | "oidc"; + filename?: string; + contentBase64?: string; + }; + const entry = request.filename === undefined ? undefined : key(request.directory, request.filename); + switch (request.operation) { + case "create": + if (entry === undefined || request.contentBase64 === undefined || records.has(entry)) return response({ created: false }); + records.set(entry, Buffer.from(request.contentBase64, "base64")); + return response({ created: true }); + case "read": + return entry === undefined || !records.has(entry) + ? response({}) + : response({ found: true, contentBase64: records.get(entry)?.toString("base64") }); + case "remove": + return response({ removed: entry !== undefined && records.delete(entry) }); + case "list": + return response({ + // Raw lower-camel entry objects, exactly as authstorage's Go response emits them. + entries: [...records.keys()] + .filter((value) => value.startsWith(`${request.directory}/`)) + .map((value) => ({ name: value.slice(request.directory.length + 1), modifiedUnixMs: base.getTime() })), + }); + default: + throw new Error("unexpected bridge operation"); + } + }, + }); + 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 }); + await create(store, { idleTtlMs: 60_000, absoluteTtlMs: 60_000 }); + await store.createOidcState({ nonce: "n".repeat(43), codeVerifier: "v".repeat(43), returnTo: "/" }, base); + + await expect(store.prune(new Date(base.getTime() + 11 * 60_000))).resolves.toBe(2); + 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 index 1f044c3c..b0cbb0a5 100644 --- a/backend/test/windows-auth-storage.test.ts +++ b/backend/test/windows-auth-storage.test.ts @@ -1,9 +1,31 @@ -import { describe, expect, test } from "vitest"; +import { EventEmitter } from "node:events"; +import { PassThrough } from "node:stream"; +import { describe, expect, test, vi } from "vitest"; import { createWindowsAuthStorageBridge } from "../src/auth/windows-auth-storage.js"; const root = "C:\\ProgramData\\ThothII\\auth"; const filename = "a".repeat(64) + ".json"; +class FakeBridgeChild extends EventEmitter { + readonly stdin = new PassThrough(); + readonly stdout = new PassThrough(); + readonly stderr = new PassThrough(); + readonly kill = vi.fn(() => true); + + close(code = 0, signal: NodeJS.Signals | null = null): void { + this.emit("close", code, signal); + } +} + +function bridgeForChild(child: FakeBridgeChild) { + const spawnChild = vi.fn(() => child); + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\tht.exe", + spawnChild, + } as never); + return { bridge, spawnChild }; +} + 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 }> = []; @@ -71,4 +93,100 @@ describe("Windows auth-storage bridge", () => { expect(() => createWindowsAuthStorageBridge({ thtExecutable: "tht.exe && unexpected" })) .toThrow("auth_session_store_invalid"); }); + + test("parses the lower-camel list DTO emitted by the Go helper for a nonempty directory", async () => { + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\tht.exe", + invoke: async () => ({ + code: 0, + // This is the raw JSON object emitted by authstorage.response after Go's DTO encoding. + stdout: Buffer.from(`{"version":1,"ok":true,"entries":[{"name":"${filename}","modifiedUnixMs":1893456245000}]}\n`), + stderr: Buffer.alloc(0), + }), + }); + + await expect(bridge.list(root, "oidc")).resolves.toEqual([ + { name: filename, modifiedUnixMs: 1_893_456_245_000 }, + ]); + }); + + test("aborts a stdin-closed looping helper on timeout and waits for close", async () => { + vi.useFakeTimers(); + const child = new FakeBridgeChild(); + const { bridge, spawnChild } = bridgeForChild(child); + const pending = bridge.list(root, "sessions"); + const outcome = pending.then(() => "resolved", () => "rejected"); + try { + await vi.advanceTimersByTimeAsync(5_000); + expect(spawnChild).toHaveBeenCalledOnce(); + expect(child.kill).toHaveBeenCalledOnce(); + expect(child.stdin.destroyed).toBe(true); + await expect(Promise.race([outcome, Promise.resolve("pending")])).resolves.toBe("pending"); + + child.close(); + await expect(outcome).resolves.toBe("rejected"); + } finally { + child.close(); + vi.useRealTimers(); + } + }); + + test.each(["stdout", "stderr"] as const)("aborts a %s-flooding helper and waits for close", async (stream) => { + const child = new FakeBridgeChild(); + const { bridge, spawnChild } = bridgeForChild(child); + const pending = bridge.list(root, "sessions"); + const outcome = pending.then(() => "resolved", () => "rejected"); + await Promise.resolve(); + expect(spawnChild).toHaveBeenCalledOnce(); + + child[stream].write(Buffer.alloc(64 * 1024 + 1)); + await Promise.resolve(); + expect(child.kill).toHaveBeenCalledOnce(); + expect(child.stdin.destroyed).toBe(true); + await expect(Promise.race([outcome, Promise.resolve("pending")])).resolves.toBe("pending"); + + child.close(); + await expect(outcome).resolves.toBe("rejected"); + }); + + test("aborts a helper when its stdin errors and waits for termination", async () => { + const child = new FakeBridgeChild(); + const { bridge, spawnChild } = bridgeForChild(child); + const stdinErrored = new Promise((resolve) => child.stdin.once("error", () => resolve())); + child.stdin.once("finish", () => child.stdin.destroy(new Error("stdin failure"))); + const pending = bridge.list(root, "sessions"); + const outcome = pending.then(() => "resolved", () => "rejected"); + await Promise.resolve(); + expect(spawnChild).toHaveBeenCalledOnce(); + + await stdinErrored; + expect(child.kill).toHaveBeenCalledOnce(); + child.close(); + await expect(outcome).resolves.toBe("rejected"); + }); + + test("aborts an errored child exactly once and waits for its close event", async () => { + const child = new FakeBridgeChild(); + const { bridge, spawnChild } = bridgeForChild(child); + const pending = bridge.list(root, "sessions"); + const outcome = pending.then(() => "resolved", () => "rejected"); + await Promise.resolve(); + expect(spawnChild).toHaveBeenCalledOnce(); + + child.emit("error", new Error("helper error")); + child.emit("error", new Error("duplicate helper error")); + expect(child.kill).toHaveBeenCalledOnce(); + await expect(Promise.race([outcome, Promise.resolve("pending")])).resolves.toBe("pending"); + child.close(); + await expect(outcome).resolves.toBe("rejected"); + }); + + test("fails closed when launching the helper throws before a child exists", async () => { + const bridge = createWindowsAuthStorageBridge({ + thtExecutable: "C:\\tht.exe", + spawnChild: () => { throw new Error("launch detail must not escape"); }, + }); + + await expect(bridge.list(root, "sessions")).rejects.toThrow("auth_session_store_invalid"); + }); }); diff --git a/tools/tht/internal/authstorage/storage.go b/tools/tht/internal/authstorage/storage.go index e3e749d8..7dfd74f1 100644 --- a/tools/tht/internal/authstorage/storage.go +++ b/tools/tht/internal/authstorage/storage.go @@ -3,6 +3,7 @@ package authstorage import ( + "bytes" "context" "encoding/base64" "encoding/json" @@ -60,7 +61,11 @@ func Run(ctx context.Context, args []string, stdin io.Reader, stdout, stderr io. if err := ctx.Err(); err != nil { return fail(stderr) } - decoder := json.NewDecoder(io.LimitReader(stdin, maximumProtocolBytes+1)) + payload, err := io.ReadAll(io.LimitReader(stdin, maximumProtocolBytes+1)) + if err != nil || len(payload) > maximumProtocolBytes { + return fail(stderr) + } + decoder := json.NewDecoder(bytes.NewReader(payload)) decoder.DisallowUnknownFields() var input request if err := decoder.Decode(&input); err != nil { @@ -229,7 +234,7 @@ func createPrivate(directory, filename string, contents []byte) (bool, error) { } else if !errors.Is(err, os.ErrNotExist) { return false, errInvalid } - if err := safeio.WriteCanonicalNewFile(path, contents, 0o600); err == nil { + if err := safeio.WriteCanonicalNewPrivateFile(path, contents, 0o600); err == nil { return true, nil } if safeio.ValidatePrivateRegular(path) == nil { diff --git a/tools/tht/internal/authstorage/storage_test.go b/tools/tht/internal/authstorage/storage_test.go index 2ffb98fd..ea310344 100644 --- a/tools/tht/internal/authstorage/storage_test.go +++ b/tools/tht/internal/authstorage/storage_test.go @@ -7,6 +7,7 @@ import ( "encoding/json" "os" "path/filepath" + "strings" "sync" "testing" @@ -43,6 +44,44 @@ func TestProtocolCreatesReadsReplacesListsAndRemovesPrivateRecord(t *testing.T) } } +func TestProtocolListSerializesLowerCamelBridgeDTO(t *testing.T) { + root := filepath.Join(privateTestRoot(t), "auth") + filename := "cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc.json" + runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))}) + input, err := json.Marshal(request{Version: 1, Operation: "list", Root: root, Directory: "sessions"}) + 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()) + } + if !strings.Contains(stdout.String(), `"name":"`+filename+`"`) || strings.Contains(stdout.String(), `"Name":`) { + t.Fatalf("list bridge JSON = %q, want lower-camel entry fields", stdout.String()) + } +} + +func TestRunRejectsProtocolOverflowAndTrailingJSONValues(t *testing.T) { + root := filepath.Join(privateTestRoot(t), "auth") + valid, err := json.Marshal(request{Version: 1, Operation: "list", Root: root, Directory: "sessions"}) + if err != nil { + t.Fatal(err) + } + if len(valid) >= maximumProtocolBytes { + t.Fatal("test request unexpectedly consumes protocol bound") + } + + overflowWhitespace := append(append([]byte(nil), valid...), bytes.Repeat([]byte(" "), maximumProtocolBytes+1-len(valid))...) + runRawRejected(t, overflowWhitespace) + + withinBound := append(append([]byte(nil), valid...), []byte("{}")...) + runRawRejected(t, withinBound) + + crossingBound := append(append([]byte(nil), valid...), bytes.Repeat([]byte(" "), maximumProtocolBytes-len(valid)-1)...) + crossingBound = append(crossingBound, []byte("{}")...) + runRawRejected(t, crossingBound) +} + func TestProtocolClaimConsumeIsAtomicAcrossConcurrentRequests(t *testing.T) { root := filepath.Join(privateTestRoot(t), "auth") filename := "dddddddddddddddddddddddddddddddddddddddddddddddddddddddddddddddd.json" @@ -140,6 +179,14 @@ func runRejected(t *testing.T, value request) { } } +func runRawRejected(t *testing.T, input []byte) { + t.Helper() + 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) diff --git a/tools/tht/internal/authstorage/storage_windows_test.go b/tools/tht/internal/authstorage/storage_windows_test.go index c296ffed..dd0510a2 100644 --- a/tools/tht/internal/authstorage/storage_windows_test.go +++ b/tools/tht/internal/authstorage/storage_windows_test.go @@ -45,6 +45,16 @@ func TestProtocolRejectsPermissiveDACLAndReparseRoot(t *testing.T) { runRejected(t, request{Version: 1, Operation: "create", Root: linkedRoot, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("private"))}) } +func TestProtocolRejectsRecordCreationAfterSessionsDirectoryDACLBecomesPermissive(t *testing.T) { + root := filepath.Join(t.TempDir(), "auth") + filename := "abababababababababababababababababababababababababababababababab.json" + runRequest(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions"}) + if err := setPermissiveDACL(filepath.Join(root, "sessions")); err != nil { + t.Fatal(err) + } + runRejected(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))}) +} + func setPermissiveDACL(path string) error { world, err := windows.StringToSid("S-1-1-0") if err != nil { diff --git a/tools/tht/internal/safeio/files.go b/tools/tht/internal/safeio/files.go index 4a5ca745..e61a3357 100644 --- a/tools/tht/internal/safeio/files.go +++ b/tools/tht/internal/safeio/files.go @@ -8,6 +8,7 @@ import ( "io" "os" "path/filepath" + "sort" "strings" "unicode/utf8" ) @@ -81,8 +82,8 @@ func ReadCanonicalPrivateRegular(path string, maximum int64) ([]byte, error) { // 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 + Name string `json:"name"` + ModifiedUnixMs int64 `json:"modifiedUnixMs"` } // ListCanonicalPrivateDirectory lists regular, non-symlinked direct children from an owner-only @@ -94,8 +95,16 @@ func ListCanonicalPrivateDirectory(path string, maximumEntries int) ([]PrivateDi if err := ValidatePrivateDirectory(path); err != nil { return nil, ErrUnsafeFile } - entries, err := os.ReadDir(path) - if err != nil || len(entries) > maximumEntries { + directory, err := os.Open(path) + if err != nil { + return nil, ErrUnsafeFile + } + defer directory.Close() + entries, err := directory.ReadDir(maximumEntries + 1) + if err != nil && !errors.Is(err, io.EOF) { + return nil, ErrUnsafeFile + } + if len(entries) > maximumEntries { return nil, ErrUnsafeFile } result := make([]PrivateDirectoryEntry, 0, len(entries)) @@ -109,6 +118,7 @@ func ListCanonicalPrivateDirectory(path string, maximumEntries int) ([]PrivateDi } result = append(result, PrivateDirectoryEntry{Name: entry.Name(), ModifiedUnixMs: info.ModTime().UnixMilli()}) } + sort.Slice(result, func(left, right int) bool { return result[left].Name < result[right].Name }) if err := ValidatePrivateDirectory(path); err != nil { return nil, ErrUnsafeFile } @@ -116,6 +126,17 @@ func ListCanonicalPrivateDirectory(path string, maximumEntries int) ([]PrivateDi } func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error { + return writeCanonicalNewFile(path, contents, mode, false) +} + +// WriteCanonicalNewPrivateFile is the authentication-storage variant of exclusive file +// creation. It requires the final parent directory to already have platform-specific owner-only +// protection and preserves that check while the platform primitive opens the parent. +func WriteCanonicalNewPrivateFile(path string, contents []byte, mode os.FileMode) error { + return writeCanonicalNewFile(path, contents, mode, true) +} + +func writeCanonicalNewFile(path string, contents []byte, mode os.FileMode, requirePrivateParent bool) error { if err := ValidateCanonicalPath(path); err != nil { return err } @@ -123,6 +144,9 @@ func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error if err := requireCanonicalDirectory(parent); err != nil { return err } + if requirePrivateParent && ValidatePrivateDirectory(parent) != nil { + return ErrUnsafeFile + } if info, err := os.Lstat(path); err == nil { if !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || info.Mode()&os.ModeType != 0 { return ErrUnsafeFile @@ -131,7 +155,15 @@ func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error } else if !errors.Is(err, os.ErrNotExist) { return ErrUnsafeFile } - file, err := createCanonicalNewPrivateFile(path, mode) + var ( + file *os.File + err error + ) + if requirePrivateParent { + file, err = createCanonicalNewPrivateParentFile(path, mode) + } else { + file, err = createCanonicalNewPrivateFile(path, mode) + } if err != nil { return ErrUnsafeFile } diff --git a/tools/tht/internal/safeio/files_test.go b/tools/tht/internal/safeio/files_test.go index e8a2e81b..6f32a7f3 100644 --- a/tools/tht/internal/safeio/files_test.go +++ b/tools/tht/internal/safeio/files_test.go @@ -2,6 +2,7 @@ package safeio import ( "errors" + "fmt" "os" "path/filepath" "testing" @@ -93,3 +94,34 @@ func TestWriteCanonicalNewFileRejectsExistingTargets(t *testing.T) { t.Fatalf("WriteCanonicalNewFile(existing) error = %v, want ErrUnsafeFile", err) } } + +func TestListCanonicalPrivateDirectoryBoundsAndSortsValidatedEntries(t *testing.T) { + temporaryRoot, err := filepath.EvalSymlinks(os.TempDir()) + if err != nil { + t.Fatal(err) + } + root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-list-") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.RemoveAll(root) }) + if err := ProtectPrivateDirectory(root); err != nil { + t.Fatal(err) + } + for index := 255; index >= 0; index-- { + path := filepath.Join(root, fmt.Sprintf("%064x.json", index)) + if err := WriteCanonicalNewFile(path, []byte("record"), 0o600); err != nil { + t.Fatalf("WriteCanonicalNewFile(%d) error = %v", index, err) + } + } + entries, err := ListCanonicalPrivateDirectory(root, 256) + if err != nil || len(entries) != 256 || entries[0].Name != fmt.Sprintf("%064x.json", 0) || entries[255].Name != fmt.Sprintf("%064x.json", 255) { + t.Fatalf("bounded ordered entries = %#v error = %v", entries, err) + } + if err := WriteCanonicalNewFile(filepath.Join(root, fmt.Sprintf("%064x.json", 256)), []byte("record"), 0o600); err != nil { + t.Fatal(err) + } + if _, err := ListCanonicalPrivateDirectory(root, 256); !errors.Is(err, ErrUnsafeFile) { + t.Fatalf("257-entry listing error = %v, want ErrUnsafeFile", err) + } +} diff --git a/tools/tht/internal/safeio/private_unix.go b/tools/tht/internal/safeio/private_unix.go index fdc8f1d8..e5c5a4ff 100644 --- a/tools/tht/internal/safeio/private_unix.go +++ b/tools/tht/internal/safeio/private_unix.go @@ -65,6 +65,13 @@ func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, err return file, nil } +func createCanonicalNewPrivateParentFile(path string, mode os.FileMode) (*os.File, error) { + if err := ValidatePrivateDirectory(filepath.Dir(path)); err != nil { + return nil, ErrUnsafeFile + } + return createCanonicalNewPrivateFile(path, mode) +} + // ValidatePrivateRegular requires a canonical, single-link private regular file. func ValidatePrivateRegular(path string) error { if err := ValidateCanonicalPath(path); err != nil { diff --git a/tools/tht/internal/safeio/private_windows.go b/tools/tht/internal/safeio/private_windows.go index 03864260..dc916e83 100644 --- a/tools/tht/internal/safeio/private_windows.go +++ b/tools/tht/internal/safeio/private_windows.go @@ -96,8 +96,16 @@ 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) { + return createCanonicalNewFile(path, mode, false) +} + +func createCanonicalNewPrivateParentFile(path string, mode os.FileMode) (*os.File, error) { + return createCanonicalNewFile(path, mode, true) +} + +func createCanonicalNewFile(path string, mode os.FileMode, requirePrivateParent bool) (*os.File, error) { parents, target, err := openCanonicalWindowsParent(path) - if err != nil || len(parents.handles) == 0 || validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil { + if err != nil || len(parents.handles) == 0 || (requirePrivateParent && validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil) { if parents != nil { parents.Close() } diff --git a/tools/tht/internal/safeio/private_windows_test.go b/tools/tht/internal/safeio/private_windows_test.go index 26985f3f..6d70b590 100644 --- a/tools/tht/internal/safeio/private_windows_test.go +++ b/tools/tht/internal/safeio/private_windows_test.go @@ -73,6 +73,22 @@ func TestCreateCanonicalNewPrivateFileInstallsOwnerOnlyDACLAtCreation(t *testing } } +func TestGenericNewFileAllowsInheritedOperatorParentButAuthNewFileRequiresPrivateParent(t *testing.T) { + directory := filepath.Join(t.TempDir(), "operator-output") + if err := os.Mkdir(directory, 0o700); err != nil { + t.Fatal(err) + } + if err := setPermissiveDACL(directory); err != nil { + t.Fatal(err) + } + if err := WriteCanonicalNewFile(filepath.Join(directory, "candidates.yaml"), []byte("reviewed: []\n"), 0o600); err != nil { + t.Fatalf("generic operator output error = %v", err) + } + if err := WriteCanonicalNewPrivateFile(filepath.Join(directory, "auth.json"), []byte("record"), 0o600); !errors.Is(err, ErrUnsafeFile) { + t.Fatalf("auth record in inherited directory error = %v, want ErrUnsafeFile", err) + } +} + func TestWithWindowsSecurityDescriptorKeepsOwnedDescriptorValidDuringInspection(t *testing.T) { directory := filepath.Join(t.TempDir(), "auth") if err := os.Mkdir(directory, 0o700); err != nil { diff --git a/tools/tht/internal/workspaceops/operations_windows_test.go b/tools/tht/internal/workspaceops/operations_windows_test.go new file mode 100644 index 00000000..6c20ea49 --- /dev/null +++ b/tools/tht/internal/workspaceops/operations_windows_test.go @@ -0,0 +1,44 @@ +//go:build windows + +package workspaceops + +import ( + "context" + "os" + "path/filepath" + "reflect" + "strings" + "testing" + + "github.com/aritmolab/thothii/tools/tht/internal/compose" +) + +func TestExecuteSuggestFksWritesOutputInInheritedOperatorDirectory(t *testing.T) { + installation := testInstallation(t) + operatorDirectory := filepath.Join(filepath.Dir(installation.Path), "operator-output") + if err := os.Mkdir(operatorDirectory, 0o700); err != nil { + t.Fatal(err) + } + output := filepath.Join(operatorDirectory, "candidates.yaml") + runner := &fakeRunner{run: func(args []string, _ string) (compose.Result, error) { + switch { + case contains(args, "config", "--format", "json"): + return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local"}}}`}, nil + case reflect.DeepEqual(args, []string{"image", "inspect", "--format", "{{.Id}}", "thothii-core:local"}): + return compose.Result{Stdout: "sha256:" + strings.Repeat("d", 64)}, nil + case contains(args, "workspace-maintenance", "schema-suggest-fks"): + return compose.Result{Stdout: `{"schemaVersion":1,"status":"blocked","code":"manual_review_required","workspaceId":"abc","workspaceRevision":"1234567890abcdef1234567890abcdef12345678","descriptorBlob":"sha256:` + strings.Repeat("e", 64) + `","operation":"schema-suggest-fks","completedStages":[],"suggestedFksYaml":"reviewed: []\n"}`}, nil + default: + t.Fatalf("unexpected Docker invocation: %#v", args) + return compose.Result{}, nil + } + }} + + if _, err := Execute(context.Background(), installation, runner, SuggestFksRequest{baseRequest: baseRequest{Workspace: "abc"}, Output: output}); err != nil { + t.Fatalf("workspace schema suggest-fks --output error = %v", err) + } + contents, err := os.ReadFile(output) + if err != nil || string(contents) != "reviewed: []\n" { + t.Fatalf("workspace schema suggest-fks output = %q error = %v", contents, err) + } +}