From 0a2c6672315111feca5e6a51d0d76801de316be8 Mon Sep 17 00:00:00 2001 From: mptyl Date: Mon, 17 Aug 2026 19:32:03 +0200 Subject: [PATCH] fix(auth): close cancellation and workspace races --- frontend/src/shell/WorkspaceManager.test.tsx | 403 ++++++++++++++++++ frontend/src/shell/WorkspaceManager.tsx | 201 +++++---- .../internal/compose/process_termination.go | 57 ++- .../compose/process_termination_test.go | 82 ++-- tools/tht/internal/compose/process_unix.go | 19 + tools/tht/internal/compose/process_windows.go | 194 +++++++-- .../compose/process_windows_static_test.go | 17 +- .../internal/compose/process_windows_test.go | 50 ++- tools/tht/internal/compose/runner.go | 16 +- tools/tht/internal/compose/runner_test.go | 4 +- 10 files changed, 822 insertions(+), 221 deletions(-) diff --git a/frontend/src/shell/WorkspaceManager.test.tsx b/frontend/src/shell/WorkspaceManager.test.tsx index 0587ae51..c318f4e3 100644 --- a/frontend/src/shell/WorkspaceManager.test.tsx +++ b/frontend/src/shell/WorkspaceManager.test.tsx @@ -62,6 +62,38 @@ function renderDeniedManager() { ); } +type ManagerSettings = { + open: boolean; + canManageWorkspace: boolean; + canManageSecrets: boolean; +}; + +function renderManagerWithSettings(initial: Partial = {}) { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + const onClose = vi.fn(); + let settings: ManagerSettings = { + open: true, + canManageWorkspace: true, + canManageSecrets: true, + ...initial, + }; + const tree = () => ( + + + + ); + const view = render(tree()); + return { + ...view, + client, + onClose, + rerender(next: Partial) { + settings = { ...settings, ...next }; + view.rerender(tree()); + }, + }; +} + beforeEach(() => { localStorage.clear(); setAuthState(authenticatedWorkspaceUser); @@ -610,6 +642,377 @@ test("a delayed runtime-secret save from user A cannot repopulate user B's cache expect(screen.queryByText("Runtime secrets saved. Stored values remain hidden.")).not.toBeInTheDocument(); }); +test("a deferred secret save cannot update a reopened dialog", async () => { + const user = userEvent.setup(); + let release!: () => void; + let started!: () => void; + let settled!: () => void; + const held = new Promise((resolve) => { release = resolve; }); + const requestStarted = new Promise((resolve) => { started = resolve; }); + const requestSettled = new Promise((resolve) => { settled = resolve; }); + server.use(http.put("/api/workspaces/psd-clinical/secrets", async () => { + started(); + try { + await held; + return HttpResponse.json(runtimeConfiguration(true)); + } finally { + settled(); + } + })); + const manager = renderManagerWithSettings(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + await user.type(await screen.findByLabelText("Data warehouse password"), "one-time-password"); + await user.click(screen.getByRole("button", { name: "Save entered secrets" })); + await requestStarted; + + await user.click(screen.getByRole("button", { name: "Close workspace management" })); + manager.rerender({ open: false }); + manager.rerender({ open: true }); + await user.type(await screen.findByLabelText("Data warehouse password"), "new-secret"); + expect(await screen.findByRole("button", { name: "Save entered secrets" })).not.toBeDisabled(); + const cacheBeforeRelease = manager.client.getQueryData(["workspace-runtime-configuration", "psd-clinical"]); + + act(() => release()); + await act(async () => { + await requestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + + expect(manager.client.getQueryData(["workspace-runtime-configuration", "psd-clinical"])).toEqual(cacheBeforeRelease); + expect(screen.queryByText("Runtime secrets saved. Stored values remain hidden.")).not.toBeInTheDocument(); +}); + +test("a deferred secret save cannot update a workspace selected again after a cross-workspace transition", async () => { + const user = userEvent.setup(); + const alternateId = "other-clinical"; + const alternateWorkspace = canonicalWorkspaceFixture(alternateId); + const alternateRevision = workspaceRevisionFixture(alternateId); + let release!: () => void; + let started!: () => void; + let settled!: () => void; + const held = new Promise((resolve) => { release = resolve; }); + const requestStarted = new Promise((resolve) => { started = resolve; }); + const requestSettled = new Promise((resolve) => { settled = resolve; }); + server.use( + http.get("/api/workspaces", () => HttpResponse.json([ + workspaceSummaryFixture("psd-clinical", { + displayName: "PSD Clinical", description: "Clinical data", configurationState: "configuration_required", revision, + }), + workspaceSummaryFixture(alternateId, { + displayName: "Other Clinical", description: "Other data", configurationState: "configuration_required", revision: alternateRevision, + }), + ])), + http.get(`/api/workspaces/${alternateId}`, () => HttpResponse.json({ workspace: alternateWorkspace, revision: alternateRevision })), + http.get(`/api/workspaces/${alternateId}/runtime-configuration`, () => HttpResponse.json({ + workspaceId: alternateId, revision: alternateRevision, configurationState: "configuration_required", requirements: [{ ...requirement }], + })), + http.put("/api/workspaces/psd-clinical/secrets", async () => { + started(); + try { + await held; + return HttpResponse.json(runtimeConfiguration(true)); + } finally { + settled(); + } + }), + ); + const manager = renderManagerWithSettings(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + await user.type(await screen.findByLabelText("Data warehouse password"), "one-time-password"); + await user.click(screen.getByRole("button", { name: "Save entered secrets" })); + await requestStarted; + + await user.click(screen.getByRole("button", { name: "Other Clinical" })); + expect(await screen.findByRole("heading", { name: "Other Clinical" })).toBeVisible(); + await user.click(screen.getByRole("button", { name: "PSD Clinical" })); + await user.type(await screen.findByLabelText("Data warehouse password"), "new-secret"); + expect(await screen.findByRole("button", { name: "Save entered secrets" })).not.toBeDisabled(); + const cacheBeforeRelease = manager.client.getQueryData(["workspace-runtime-configuration", "psd-clinical"]); + + act(() => release()); + await act(async () => { + await requestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + + expect(manager.client.getQueryData(["workspace-runtime-configuration", "psd-clinical"])).toEqual(cacheBeforeRelease); + expect(screen.queryByText("Runtime secrets saved. Stored values remain hidden.")).not.toBeInTheDocument(); +}); + +test("a deferred secret save cannot update a context after its secret permission changes", async () => { + const user = userEvent.setup(); + let release!: () => void; + let started!: () => void; + let settled!: () => void; + const held = new Promise((resolve) => { release = resolve; }); + const requestStarted = new Promise((resolve) => { started = resolve; }); + const requestSettled = new Promise((resolve) => { settled = resolve; }); + server.use(http.put("/api/workspaces/psd-clinical/secrets", async () => { + started(); + try { + await held; + return HttpResponse.json(runtimeConfiguration(true)); + } finally { + settled(); + } + })); + const manager = renderManagerWithSettings(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + await user.type(await screen.findByLabelText("Data warehouse password"), "one-time-password"); + await user.click(screen.getByRole("button", { name: "Save entered secrets" })); + await requestStarted; + + manager.rerender({ canManageSecrets: false }); + manager.rerender({ canManageSecrets: true }); + await user.type(await screen.findByLabelText("Data warehouse password"), "new-secret"); + expect(await screen.findByRole("button", { name: "Save entered secrets" })).not.toBeDisabled(); + const cacheBeforeRelease = manager.client.getQueryData(["workspace-runtime-configuration", "psd-clinical"]); + + act(() => release()); + await act(async () => { + await requestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + + expect(manager.client.getQueryData(["workspace-runtime-configuration", "psd-clinical"])).toEqual(cacheBeforeRelease); + expect(screen.queryByText("Runtime secrets saved. Stored values remain hidden.")).not.toBeInTheDocument(); +}); + +test("a deferred repository update cannot update a reopened dialog or refetch its cache", async () => { + const user = userEvent.setup(); + let release!: () => void; + let started!: () => void; + let settled!: () => void; + let statusRequests = 0; + let workspaceRequests = 0; + const held = new Promise((resolve) => { release = resolve; }); + const requestStarted = new Promise((resolve) => { started = resolve; }); + const requestSettled = new Promise((resolve) => { settled = resolve; }); + server.use( + http.get("/api/workspace-registry/status", () => { + statusRequests++; + return HttpResponse.json({ branch: "main", head: "a".repeat(40), ahead: 0, behind: 0, degraded: false }); + }), + http.get("/api/workspaces", () => { + workspaceRequests++; + return HttpResponse.json([workspaceSummaryFixture("psd-clinical", { + displayName: "PSD Clinical", description: "Clinical data", configurationState: "configuration_required", revision, + })]); + }), + http.post("/api/workspace-registry/pull", async () => { + started(); + try { + await held; + return HttpResponse.json({ branch: "main", head: "b".repeat(40), ahead: 0, behind: 0, degraded: false }); + } finally { + settled(); + } + }), + ); + const manager = renderManagerWithSettings(); + await user.click(await screen.findByRole("button", { name: "Update workspace repository" })); + await requestStarted; + + await user.click(screen.getByRole("button", { name: "Close workspace management" })); + manager.rerender({ open: false }); + manager.rerender({ open: true }); + expect(await screen.findByRole("button", { name: "Update workspace repository" })).not.toBeDisabled(); + const requestsBeforeRelease = { statusRequests, workspaceRequests }; + + act(() => release()); + await act(async () => { + await requestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + + expect({ statusRequests, workspaceRequests }).toEqual(requestsBeforeRelease); + expect(screen.queryByText("Workspace repository updated and validated.")).not.toBeInTheDocument(); +}); + +test("a deferred repository update cannot update a context after a cross-workspace transition", async () => { + const user = userEvent.setup(); + const alternateId = "other-clinical"; + const alternateWorkspace = canonicalWorkspaceFixture(alternateId); + const alternateRevision = workspaceRevisionFixture(alternateId); + let release!: () => void; + let started!: () => void; + let settled!: () => void; + let statusRequests = 0; + let workspaceRequests = 0; + const held = new Promise((resolve) => { release = resolve; }); + const requestStarted = new Promise((resolve) => { started = resolve; }); + const requestSettled = new Promise((resolve) => { settled = resolve; }); + server.use( + http.get("/api/workspace-registry/status", () => { + statusRequests++; + return HttpResponse.json({ branch: "main", head: "a".repeat(40), ahead: 0, behind: 0, degraded: false }); + }), + http.get("/api/workspaces", () => { + workspaceRequests++; + return HttpResponse.json([ + workspaceSummaryFixture("psd-clinical", { + displayName: "PSD Clinical", description: "Clinical data", configurationState: "configuration_required", revision, + }), + workspaceSummaryFixture(alternateId, { + displayName: "Other Clinical", description: "Other data", configurationState: "configuration_required", revision: alternateRevision, + }), + ]); + }), + http.get(`/api/workspaces/${alternateId}`, () => HttpResponse.json({ workspace: alternateWorkspace, revision: alternateRevision })), + http.get(`/api/workspaces/${alternateId}/runtime-configuration`, () => HttpResponse.json({ + workspaceId: alternateId, revision: alternateRevision, configurationState: "configuration_required", requirements: [{ ...requirement }], + })), + http.post("/api/workspace-registry/pull", async () => { + started(); + try { + await held; + return HttpResponse.json({ branch: "main", head: "b".repeat(40), ahead: 0, behind: 0, degraded: false }); + } finally { + settled(); + } + }), + ); + const manager = renderManagerWithSettings(); + await user.click(await screen.findByRole("button", { name: "Update workspace repository" })); + await requestStarted; + + await user.click(screen.getByRole("button", { name: "Other Clinical" })); + expect(await screen.findByRole("heading", { name: "Other Clinical" })).toBeVisible(); + await user.click(within(screen.getByRole("navigation", { name: "Workspaces" })) + .getByRole("button", { name: "Back to Level 1" })); + expect(await screen.findByRole("button", { name: "Update workspace repository" })).not.toBeDisabled(); + const requestsBeforeRelease = { statusRequests, workspaceRequests }; + + act(() => release()); + await act(async () => { + await requestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + + expect({ statusRequests, workspaceRequests }).toEqual(requestsBeforeRelease); + expect(screen.queryByText("Workspace repository updated and validated.")).not.toBeInTheDocument(); +}); + +test("a deferred repository update cannot update a context after workspace permission changes", async () => { + const user = userEvent.setup(); + let release!: () => void; + let started!: () => void; + let settled!: () => void; + let statusRequests = 0; + let workspaceRequests = 0; + const held = new Promise((resolve) => { release = resolve; }); + const requestStarted = new Promise((resolve) => { started = resolve; }); + const requestSettled = new Promise((resolve) => { settled = resolve; }); + server.use( + http.get("/api/workspace-registry/status", () => { + statusRequests++; + return HttpResponse.json({ branch: "main", head: "a".repeat(40), ahead: 0, behind: 0, degraded: false }); + }), + http.get("/api/workspaces", () => { + workspaceRequests++; + return HttpResponse.json([workspaceSummaryFixture("psd-clinical", { + displayName: "PSD Clinical", description: "Clinical data", configurationState: "configuration_required", revision, + })]); + }), + http.post("/api/workspace-registry/pull", async () => { + started(); + try { + await held; + return HttpResponse.json({ branch: "main", head: "b".repeat(40), ahead: 0, behind: 0, degraded: false }); + } finally { + settled(); + } + }), + ); + const manager = renderManagerWithSettings(); + await user.click(await screen.findByRole("button", { name: "Update workspace repository" })); + await requestStarted; + + manager.rerender({ canManageWorkspace: false }); + manager.rerender({ canManageWorkspace: true }); + expect(await screen.findByRole("button", { name: "Update workspace repository" })).not.toBeDisabled(); + const requestsBeforeRelease = { statusRequests, workspaceRequests }; + + act(() => release()); + await act(async () => { + await requestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + + expect({ statusRequests, workspaceRequests }).toEqual(requestsBeforeRelease); + expect(screen.queryByText("Workspace repository updated and validated.")).not.toBeInTheDocument(); +}); + +test("a deferred validation cannot update a context after workspace permission changes", async () => { + const user = userEvent.setup(); + let release!: () => void; + let started!: () => void; + let settled!: () => void; + const held = new Promise((resolve) => { release = resolve; }); + const requestStarted = new Promise((resolve) => { started = resolve; }); + const requestSettled = new Promise((resolve) => { settled = resolve; }); + server.use(http.post("/api/workspaces/validate", async () => { + started(); + try { + await held; + return HttpResponse.json({ workspace, contract: {}, activatable: true, diagnostics: [], authentication: readyAuthentication }); + } finally { + settled(); + } + })); + const manager = renderManagerWithSettings(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + await user.click(screen.getByRole("button", { name: "Validate workspace source" })); + await requestStarted; + + manager.rerender({ canManageWorkspace: false }); + manager.rerender({ canManageWorkspace: true }); + expect(await screen.findByRole("button", { name: "Validate workspace source" })).not.toBeDisabled(); + + act(() => release()); + await act(async () => { + await requestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + + expect(screen.queryByTestId("workspace-authentication")).not.toBeInTheDocument(); +}); + +test("a deferred connection test cannot update a context after workspace permission changes", async () => { + const user = userEvent.setup(); + let release!: () => void; + let started!: () => void; + let settled!: () => void; + const held = new Promise((resolve) => { release = resolve; }); + const requestStarted = new Promise((resolve) => { started = resolve; }); + const requestSettled = new Promise((resolve) => { settled = resolve; }); + server.use(http.post("/api/workspaces/psd-clinical/test", async () => { + started(); + try { + await held; + return HttpResponse.json({ activatable: true, diagnostics: [], authentication: readyAuthentication }); + } finally { + settled(); + } + })); + const manager = renderManagerWithSettings(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + await user.click(screen.getByRole("button", { name: "Test workspace connections" })); + await requestStarted; + + manager.rerender({ canManageWorkspace: false }); + manager.rerender({ canManageWorkspace: true }); + expect(await screen.findByRole("button", { name: "Test workspace connections" })).not.toBeDisabled(); + + act(() => release()); + await act(async () => { + await requestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + + expect(screen.queryByTestId("workspace-authentication")).not.toBeInTheDocument(); +}); + test("closing clears unsaved secret fields", async () => { const user = userEvent.setup(); const onClose = vi.fn(); diff --git a/frontend/src/shell/WorkspaceManager.tsx b/frontend/src/shell/WorkspaceManager.tsx index 3a718330..c0678f42 100644 --- a/frontend/src/shell/WorkspaceManager.tsx +++ b/frontend/src/shell/WorkspaceManager.tsx @@ -1,4 +1,4 @@ -import { useEffect, useMemo, useRef, useState } from "react"; +import { useEffect, useLayoutEffect, useMemo, useRef, useState } from "react"; import { useQuery, useQueryClient } from "@tanstack/react-query"; import { AlertCircle, @@ -26,7 +26,12 @@ import { type AuthDiagnostics, type WorkspaceRuntimeConfiguration, } from "../api/workspaces"; -import { captureAuthOperation, isAuthOperationCurrent, StaleAuthOperationError } from "../auth/authOperation"; +import { + captureAuthOperation, + isAuthOperationCurrent, + StaleAuthOperationError, + type AuthOperationGuard, +} from "../auth/authOperation"; import { Button } from "../components/ui/button"; import { Dialog, @@ -62,6 +67,17 @@ function stateLabel(state: "ready" | "configuration_required"): string { const workspaceAuthoringGuideUrl = "https://github.com/mptyl/ThothII/blob/main/docs/install/local-workspace-registry.md#prepare-and-publish-a-workspace-source"; +type WorkspaceOperationContext = Readonly<{ + open: boolean; + canManageWorkspace: boolean; + canManageSecrets: boolean; +}>; + +type WorkspaceOperationGuard = Readonly<{ + operation: AuthOperationGuard; + context: WorkspaceOperationContext; +}>; + export function WorkspaceManager({ open, onClose, @@ -87,50 +103,97 @@ export function WorkspaceManager({ const operationEpochRef = useRef(0); const diagnosticEpochRef = useRef(0); const selectedIdRef = useRef(selectedId); - const wasOpenRef = useRef(open); + const contextRef = useRef({ open, canManageWorkspace, canManageSecrets }); + const committedContextRef = useRef(contextRef.current); selectedIdRef.current = selectedId; + contextRef.current = { open, canManageWorkspace, canManageSecrets }; - const clearGlobalMessages = () => { + const clearWorkspacePresentation = () => { setNotice(undefined); setDiagnostics([]); - }; - - const clearDiagnosticMessages = () => { - diagnosticEpochRef.current += 1; setValidationNotice(undefined); setValidationDiagnostics([]); setConnectionNotice(undefined); setConnectionDiagnostics([]); setAuthentication(undefined); - setBusyAction((current) => current === "validate" || current === "test" ? undefined : current); }; - const clearMessages = () => { - clearGlobalMessages(); - clearDiagnosticMessages(); + const cancelWorkspaceQueries = () => { + void queryClient.cancelQueries({ queryKey: ["workspace-repository-status"] }); + void queryClient.cancelQueries({ queryKey: ["workspaces"] }); + void queryClient.cancelQueries({ queryKey: ["workspace"] }); + void queryClient.cancelQueries({ queryKey: ["workspace-runtime-configuration"] }); }; - useEffect(() => () => { operationEpochRef.current += 1; }, []); - useEffect(() => { - if (wasOpenRef.current && !open) { + const clearWorkspaceCaches = () => { + queryClient.removeQueries({ queryKey: ["workspace-repository-status"] }); + queryClient.removeQueries({ queryKey: ["workspaces"] }); + queryClient.removeQueries({ queryKey: ["workspace"] }); + queryClient.removeQueries({ queryKey: ["workspace-runtime-configuration"] }); + }; + + // This is the sole invalidation boundary for deferred workspace work. Call it before capturing + // an operation guard; doing it after capture would immediately invalidate the new operation. + const invalidateWorkspaceContext = ({ + clearSecrets = true, + clearCaches = false, + }: { clearSecrets?: boolean; clearCaches?: boolean } = {}) => { + operationEpochRef.current += 1; + diagnosticEpochRef.current += 1; + cancelWorkspaceQueries(); + if (clearCaches) clearWorkspaceCaches(); + setBusyAction(undefined); + if (clearSecrets) { setSecretValues({}); - clearMessages(); } - wasOpenRef.current = open; - }, [open]); + clearWorkspacePresentation(); + }; - async function guardedQuery(request: () => Promise, targetId?: string): Promise { - const guard = captureAuthOperation({ + useEffect(() => () => { + operationEpochRef.current += 1; + diagnosticEpochRef.current += 1; + }, []); + + useLayoutEffect(() => { + const previous = committedContextRef.current; + const current = contextRef.current; + const openChanged = previous.open !== current.open; + const permissionsChanged = previous.canManageWorkspace !== current.canManageWorkspace + || previous.canManageSecrets !== current.canManageSecrets; + if (openChanged || permissionsChanged) { + invalidateWorkspaceContext({ clearCaches: permissionsChanged }); + } + committedContextRef.current = current; + }, [open, canManageWorkspace, canManageSecrets]); + + function captureWorkspaceOperation(targetId?: string): WorkspaceOperationGuard | null { + const context = contextRef.current; + if (!context.open) return null; + const operation = captureAuthOperation({ sessionId: targetId ?? null, disposalEpoch: operationEpochRef.current, }); + return operation ? { operation, context } : null; + } + + function isWorkspaceOperationCurrent(guard: WorkspaceOperationGuard | null, targetId?: string): boolean { + if (!guard) return false; + const context = contextRef.current; + return context.open === guard.context.open + && context.canManageWorkspace === guard.context.canManageWorkspace + && context.canManageSecrets === guard.context.canManageSecrets + && isAuthOperationCurrent(guard.operation, { + sessionId: targetId ?? null, + disposalEpoch: operationEpochRef.current, + }); + } + + async function guardedQuery(request: () => Promise, targetId?: string): Promise { + const guard = captureWorkspaceOperation(targetId); if (!guard) throw new StaleAuthOperationError(); const result = await request(); - const currentSessionId = targetId === undefined ? null : selectedIdRef.current; - if (!isAuthOperationCurrent(guard, { - sessionId: currentSessionId, - disposalEpoch: operationEpochRef.current, - })) throw new StaleAuthOperationError(); + const currentTargetId = targetId === undefined ? undefined : selectedIdRef.current; + if (!isWorkspaceOperationCurrent(guard, currentTargetId)) throw new StaleAuthOperationError(); return result; } @@ -161,90 +224,84 @@ export function WorkspaceManager({ }); const close = () => { - setSecretValues({}); - clearMessages(); + invalidateWorkspaceContext(); onClose(); }; const selectWorkspace = (id: string) => { + invalidateWorkspaceContext(); setSelectedId(id); - setSecretValues({}); - clearMessages(); }; const showLevelOne = () => { + invalidateWorkspaceContext(); setSelectedId(undefined); - setSecretValues({}); - clearMessages(); }; async function updateRepository() { - const guard = captureAuthOperation({ disposalEpoch: operationEpochRef.current }); + invalidateWorkspaceContext({ clearCaches: true }); + const guard = captureWorkspaceOperation(); if (!guard) return; - clearMessages(); + const targetId = selectedIdRef.current; setBusyAction("repository"); try { await pullWorkspaceRegistry(); - if (!isAuthOperationCurrent(guard, { disposalEpoch: operationEpochRef.current })) return; + if (!isWorkspaceOperationCurrent(guard)) return; await Promise.all([ statusQuery.refetch(), workspacesQuery.refetch(), - selectedId ? detailQuery.refetch() : Promise.resolve(), - selectedId ? runtimeQuery.refetch() : Promise.resolve(), + targetId ? detailQuery.refetch() : Promise.resolve(), + targetId ? runtimeQuery.refetch() : Promise.resolve(), ]); - if (!isAuthOperationCurrent(guard, { disposalEpoch: operationEpochRef.current })) return; + if (!isWorkspaceOperationCurrent(guard)) return; setNotice("Workspace repository updated and validated."); } catch (error) { - if (isAuthOperationCurrent(guard, { disposalEpoch: operationEpochRef.current })) { + if (isWorkspaceOperationCurrent(guard)) { setDiagnostics([publicError(error, "git_unavailable: Workspace repository could not be updated")]); } } finally { - if (isAuthOperationCurrent(guard, { disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); + if (isWorkspaceOperationCurrent(guard)) setBusyAction(undefined); } } async function validateSource() { if (!detailQuery.data) return; - const guard = captureAuthOperation({ sessionId: selectedId, disposalEpoch: operationEpochRef.current }); + const targetId = selectedId; + const source = detailQuery.data.workspace; + invalidateWorkspaceContext({ clearSecrets: false }); + const guard = captureWorkspaceOperation(targetId); if (!guard) return; - const diagnosticEpoch = ++diagnosticEpochRef.current; + const diagnosticEpoch = diagnosticEpochRef.current; setBusyAction("validate"); - clearGlobalMessages(); - setAuthentication(undefined); - setValidationNotice(undefined); - setValidationDiagnostics([]); try { - const result = await validateWorkspace(detailQuery.data.workspace); + const result = await validateWorkspace(source); if (diagnosticEpoch !== diagnosticEpochRef.current || - !isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; + !isWorkspaceOperationCurrent(guard, selectedIdRef.current)) return; setAuthentication(result.authentication); setValidationNotice(result.activatable ? "Workspace source and authentication are valid." : "Workspace source is valid."); } catch (error) { if (diagnosticEpoch === diagnosticEpochRef.current && - isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) { + isWorkspaceOperationCurrent(guard, selectedIdRef.current)) { setValidationDiagnostics([publicError(error, "workspace_invalid: Workspace validation could not be completed")]); } } finally { if (diagnosticEpoch === diagnosticEpochRef.current && - isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); + isWorkspaceOperationCurrent(guard, selectedIdRef.current)) setBusyAction(undefined); } } async function testConnections() { if (!selectedId) return; const targetId = selectedId; - const guard = captureAuthOperation({ sessionId: targetId, disposalEpoch: operationEpochRef.current }); + invalidateWorkspaceContext({ clearSecrets: false }); + const guard = captureWorkspaceOperation(targetId); if (!guard) return; - const diagnosticEpoch = ++diagnosticEpochRef.current; + const diagnosticEpoch = diagnosticEpochRef.current; setBusyAction("test"); - clearGlobalMessages(); - setAuthentication(undefined); - setConnectionNotice(undefined); - setConnectionDiagnostics([]); try { - const result = await testWorkspace(selectedId); + const result = await testWorkspace(targetId); if (diagnosticEpoch !== diagnosticEpochRef.current || - !isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; + !isWorkspaceOperationCurrent(guard, selectedIdRef.current)) return; setAuthentication(result.authentication); const issues = result.diagnostics.filter(({ level }) => level !== "info"); const informational = result.diagnostics.find(({ level }) => level === "info"); @@ -258,70 +315,70 @@ export function WorkspaceManager({ } } catch (error) { if (diagnosticEpoch === diagnosticEpochRef.current && - isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) { + isWorkspaceOperationCurrent(guard, selectedIdRef.current)) { setConnectionDiagnostics([publicError(error, "connector_unavailable: Workspace connections could not be tested")]); } } finally { if (diagnosticEpoch === diagnosticEpochRef.current && - isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); + isWorkspaceOperationCurrent(guard, selectedIdRef.current)) setBusyAction(undefined); } } async function saveSecrets() { if (!selectedId) return; const targetId = selectedId; - const guard = captureAuthOperation({ sessionId: targetId, disposalEpoch: operationEpochRef.current }); - if (!guard) return; const values = Object.fromEntries( Object.entries(secretValues).filter(([, value]) => value.length > 0), ); if (Object.keys(values).length === 0) return; + invalidateWorkspaceContext({ clearSecrets: false }); + const guard = captureWorkspaceOperation(targetId); + if (!guard) return; setBusyAction("save-secrets"); - clearMessages(); try { const configuration = await saveWorkspaceSecrets(targetId, values); - if (!isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; + if (!isWorkspaceOperationCurrent(guard, selectedIdRef.current)) return; queryClient.setQueryData( ["workspace-runtime-configuration", targetId], configuration, ); setSecretValues({}); await workspacesQuery.refetch(); - if (!isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; + if (!isWorkspaceOperationCurrent(guard, selectedIdRef.current)) return; setNotice("Runtime secrets saved. Stored values remain hidden."); } catch (error) { - if (isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) { + if (isWorkspaceOperationCurrent(guard, selectedIdRef.current)) { setDiagnostics([publicError(error, "workspace_invalid: Runtime secrets could not be saved")]); } } finally { - if (isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); + if (isWorkspaceOperationCurrent(guard, selectedIdRef.current)) setBusyAction(undefined); } } async function forgetSecret(requirementId: string) { if (!selectedId) return; const targetId = selectedId; - const guard = captureAuthOperation({ sessionId: targetId, disposalEpoch: operationEpochRef.current }); + invalidateWorkspaceContext({ clearSecrets: false }); + const guard = captureWorkspaceOperation(targetId); if (!guard) return; setBusyAction(`forget:${requirementId}`); - clearMessages(); try { const configuration = await forgetWorkspaceSecret(targetId, requirementId); - if (!isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; + if (!isWorkspaceOperationCurrent(guard, selectedIdRef.current)) return; queryClient.setQueryData( ["workspace-runtime-configuration", targetId], configuration, ); setSecretValues((current) => ({ ...current, [requirementId]: "" })); await workspacesQuery.refetch(); - if (!isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; + if (!isWorkspaceOperationCurrent(guard, selectedIdRef.current)) return; setNotice("Stored secret forgotten."); } catch (error) { - if (isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) { + if (isWorkspaceOperationCurrent(guard, selectedIdRef.current)) { setDiagnostics([publicError(error, "workspace_invalid: Stored secret could not be forgotten")]); } } finally { - if (isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); + if (isWorkspaceOperationCurrent(guard, selectedIdRef.current)) setBusyAction(undefined); } } diff --git a/tools/tht/internal/compose/process_termination.go b/tools/tht/internal/compose/process_termination.go index 87b731d0..201ac785 100644 --- a/tools/tht/internal/compose/process_termination.go +++ b/tools/tht/internal/compose/process_termination.go @@ -1,26 +1,50 @@ package compose -import "time" +import ( + "errors" + "os" + "time" +) -// terminateWindowsProcess reports a reap failure only when the process has not exited within the -// final bound. taskkill and Process.Kill can both race with a child that has already been reaped. -func terminateWindowsProcess( +// processLifetime owns a started Docker CLI process until it has been reaped and every process +// it may have started has been released. Platform implementations must make close idempotent. +type processLifetime interface { + terminate(done <-chan error) error + close() error +} + +var startProcessForRunner = startProcess + +var terminateProcessForRunner = func(process processLifetime, done <-chan error) error { + return process.terminate(done) +} + +// terminateProcessTree terminates the owned tree before consulting the direct child's reaping +// result. A direct child can be reaped while descendants survive, so a ready done channel is never +// evidence that tree termination can be skipped. The direct kill is only a fallback when tree +// ownership could not terminate the tree; an already-reaped direct child is benign in that case. +func terminateProcessTree( done <-chan error, terminateTree func() error, terminateDirect func() error, bound time.Duration, ) error { - if processReaped(done) { - return nil + treeErr := terminateTree() + var directErr error + if treeErr != nil { + directErr = terminateDirect() + if errors.Is(directErr, os.ErrProcessDone) { + directErr = nil + } } - _ = terminateTree() - if processReaped(done) { - return nil - } - _ = terminateDirect() - if processReaped(done) { + reapErr := waitForProcessReap(done, bound) + if treeErr == nil && directErr == nil && reapErr == nil { return nil } + return errors.Join(ErrProcessReap, treeErr, directErr, reapErr) +} + +func waitForProcessReap(done <-chan error, bound time.Duration) error { timer := time.NewTimer(bound) defer timer.Stop() select { @@ -30,12 +54,3 @@ func terminateWindowsProcess( return ErrProcessReap } } - -func processReaped(done <-chan error) bool { - select { - case <-done: - return true - default: - return false - } -} diff --git a/tools/tht/internal/compose/process_termination_test.go b/tools/tht/internal/compose/process_termination_test.go index 7ef6eedf..6352bbc5 100644 --- a/tools/tht/internal/compose/process_termination_test.go +++ b/tools/tht/internal/compose/process_termination_test.go @@ -2,88 +2,80 @@ package compose import ( "errors" + "os" "testing" + "time" ) -func TestWindowsCancellationTreatsProcessReapedBeforeTreeTerminationAsSuccess(t *testing.T) { +func TestTerminationKillsTheTreeEvenWhenTheDirectChildIsAlreadyReaped(t *testing.T) { done := make(chan error, 1) done <- nil + treeFailure := errors.New("tree termination failed") treeCalls := 0 directCalls := 0 - err := terminateWindowsProcess(done, func() error { + err := terminateProcessTree(done, func() error { treeCalls++ - return errors.New("tree termination should not run") + return treeFailure }, func() error { directCalls++ - return errors.New("direct termination should not run") - }, 0) + return os.ErrProcessDone + }, time.Second) - if err != nil { - t.Fatalf("terminateWindowsProcess() error = %v, want nil", err) + if treeCalls != 1 { + t.Fatalf("tree termination calls = %d, want 1", treeCalls) } - if treeCalls != 0 || directCalls != 0 { - t.Fatalf("termination calls = tree:%d direct:%d, want none", treeCalls, directCalls) + if directCalls != 1 { + t.Fatalf("direct termination calls = %d, want 1 fallback after tree failure", directCalls) + } + if !errors.Is(err, treeFailure) || !errors.Is(err, ErrProcessReap) { + t.Fatalf("terminateProcessTree() error = %v, want tree failure and ErrProcessReap", err) + } + if errors.Is(err, os.ErrProcessDone) { + t.Fatalf("terminateProcessTree() error = %v, must treat an already-reaped direct child as benign", err) } } -func TestWindowsCancellationTreatsProcessReapedAfterTreeTerminationAsSuccess(t *testing.T) { +func TestTerminationWaitsForTheOriginalProcessAfterTreeTermination(t *testing.T) { done := make(chan error, 1) + done <- nil directCalls := 0 - err := terminateWindowsProcess(done, func() error { - done <- nil - return ErrProcessReap - }, func() error { + err := terminateProcessTree(done, func() error { return nil }, func() error { directCalls++ - return errors.New("direct termination should not run") - }, 0) + return errors.New("direct fallback should not run after a successful tree termination") + }, time.Second) if err != nil { - t.Fatalf("terminateWindowsProcess() error = %v, want nil", err) + t.Fatalf("terminateProcessTree() error = %v, want nil", err) } if directCalls != 0 { t.Fatalf("direct termination calls = %d, want 0", directCalls) } } -func TestWindowsCancellationTreatsProcessReapedAfterDirectTerminationAsSuccess(t *testing.T) { - done := make(chan error, 1) - directCalls := 0 +func TestTerminationJoinsTreeDirectAndReapFailures(t *testing.T) { + done := make(chan error) + treeFailure := errors.New("tree termination failed") + directFailure := errors.New("direct termination failed") - err := terminateWindowsProcess(done, func() error { - return ErrProcessReap - }, func() error { - directCalls++ - done <- nil - return ErrProcessReap - }, 0) + err := terminateProcessTree(done, func() error { return treeFailure }, func() error { return directFailure }, 0) - if err != nil { - t.Fatalf("terminateWindowsProcess() error = %v, want nil", err) - } - if directCalls != 1 { - t.Fatalf("direct termination calls = %d, want 1", directCalls) + for _, want := range []error{treeFailure, directFailure, ErrProcessReap} { + if !errors.Is(err, want) { + t.Fatalf("terminateProcessTree() error = %v, want %v", err, want) + } } } -func TestWindowsCancellationReportsReapFailureWhenProcessSurvives(t *testing.T) { +func TestTerminationReportsReapFailureWhenTheTreeTerminatesButTheProcessNeverReaps(t *testing.T) { done := make(chan error) - treeCalls := 0 - directCalls := 0 - err := terminateWindowsProcess(done, func() error { - treeCalls++ - return ErrProcessReap - }, func() error { - directCalls++ - return ErrProcessReap + err := terminateProcessTree(done, func() error { return nil }, func() error { + return errors.New("direct fallback should not run") }, 0) if !errors.Is(err, ErrProcessReap) { - t.Fatalf("terminateWindowsProcess() error = %v, want ErrProcessReap", err) - } - if treeCalls != 1 || directCalls != 1 { - t.Fatalf("termination calls = tree:%d direct:%d, want one each", treeCalls, directCalls) + t.Fatalf("terminateProcessTree() error = %v, want ErrProcessReap", err) } } diff --git a/tools/tht/internal/compose/process_unix.go b/tools/tht/internal/compose/process_unix.go index 0bc62f9a..5b597adf 100644 --- a/tools/tht/internal/compose/process_unix.go +++ b/tools/tht/internal/compose/process_unix.go @@ -15,6 +15,25 @@ func configureProcess(command *exec.Cmd) { command.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} } +type unixProcessLifetime struct { + command *exec.Cmd +} + +func startProcess(command *exec.Cmd) (processLifetime, error) { + if err := command.Start(); err != nil { + return nil, err + } + return unixProcessLifetime{command: command}, nil +} + +func (process unixProcessLifetime) terminate(done <-chan error) error { + return terminateProcess(process.command, done) +} + +func (unixProcessLifetime) close() error { + return nil +} + func terminateProcess(command *exec.Cmd, done <-chan error) error { if command.Process == nil { return nil diff --git a/tools/tht/internal/compose/process_windows.go b/tools/tht/internal/compose/process_windows.go index 8170e7ea..d9e92c03 100644 --- a/tools/tht/internal/compose/process_windows.go +++ b/tools/tht/internal/compose/process_windows.go @@ -3,75 +3,181 @@ package compose import ( - "context" - "io" + "errors" + "os" "os/exec" - "path/filepath" - "strconv" "syscall" "time" + "unsafe" "golang.org/x/sys/windows" ) const finalTerminationBound = 2 * time.Second -const processTreeTerminationBound = 2 * time.Second const createNewProcessGroup = 0x00000200 - -var windowsSystemDirectory = windows.GetSystemDirectory +const createSuspendedProcess = 0x00000004 +const jobTerminationExitCode = 1 func configureProcess(command *exec.Cmd) { - command.SysProcAttr = &syscall.SysProcAttr{CreationFlags: createNewProcessGroup} + command.SysProcAttr = &syscall.SysProcAttr{ + CreationFlags: createNewProcessGroup | createSuspendedProcess, + } } -func terminateProcess(command *exec.Cmd, done <-chan error) error { - if command.Process == nil { +type windowsProcessJob struct { + handle windows.Handle +} + +func newWindowsProcessJob() (*windowsProcessJob, error) { + handle, err := windows.CreateJobObject(nil, nil) + if err != nil { + return nil, err + } + limits := windows.JOBOBJECT_EXTENDED_LIMIT_INFORMATION{} + limits.BasicLimitInformation.LimitFlags = windows.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE + if _, err := windows.SetInformationJobObject( + handle, + windows.JobObjectExtendedLimitInformation, + uintptr(unsafe.Pointer(&limits)), + uint32(unsafe.Sizeof(limits)), + ); err != nil { + return nil, errors.Join(err, windows.CloseHandle(handle)) + } + return &windowsProcessJob{handle: handle}, nil +} + +func (job *windowsProcessJob) close() error { + if job == nil || job.handle == 0 { return nil } - return terminateWindowsProcess( + handle := job.handle + job.handle = 0 + return windows.CloseHandle(handle) +} + +type windowsProcessLifetime struct { + command *exec.Cmd + job *windowsProcessJob +} + +// startProcess creates the Docker process suspended, assigns it to a kill-on-close Job Object, +// then resumes its one primary thread. The suspension closes the startup window in which a direct +// child could otherwise create descendants before the Job Object begins to own them. +func startProcess(command *exec.Cmd) (processLifetime, error) { + job, err := newWindowsProcessJob() + if err != nil { + return nil, err + } + if err := command.Start(); err != nil { + return nil, errors.Join(err, job.close()) + } + if command.Process == nil { + return nil, errors.Join(ErrProcessReap, abortWindowsStart(command, job)) + } + if err := assignWindowsProcessToJob(job, command.Process.Pid); err != nil { + return nil, errors.Join(err, abortWindowsStart(command, job)) + } + if err := resumeWindowsProcess(command.Process.Pid); err != nil { + return nil, errors.Join(err, abortWindowsStart(command, job)) + } + return &windowsProcessLifetime{command: command, job: job}, nil +} + +func (process *windowsProcessLifetime) terminate(done <-chan error) error { + if process == nil || process.command == nil || process.command.Process == nil || process.job == nil || process.job.handle == 0 { + return ErrProcessReap + } + return terminateProcessTree( done, - func() error { return terminateWindowsProcessTree(command.Process.Pid) }, - command.Process.Kill, + func() error { return windows.TerminateJobObject(process.job.handle, jobTerminationExitCode) }, + process.command.Process.Kill, finalTerminationBound, ) } -func terminateWindowsProcessTree(pid int) error { - ctx, cancel := context.WithTimeout(context.Background(), processTreeTerminationBound) - defer cancel() - taskkillPath, err := systemTaskkillPath() - if err != nil { - return ErrProcessReap +func (process *windowsProcessLifetime) close() error { + if process == nil { + return nil } - command := windowsTreeKillCommand(ctx, taskkillPath, pid) - command.Stdout = io.Discard - command.Stderr = io.Discard - if err := command.Run(); err != nil { - return ErrProcessReap - } - if ctx.Err() != nil { - return ErrProcessReap - } - return nil + return process.job.close() } -// systemTaskkillPath resolves taskkill through Windows' protected system-directory API, never PATH. -func systemTaskkillPath() (string, error) { - systemDirectory, err := windowsSystemDirectory() +func assignWindowsProcessToJob(job *windowsProcessJob, pid int) error { + process, err := windows.OpenProcess(windows.PROCESS_SET_QUOTA|windows.PROCESS_TERMINATE, false, uint32(pid)) if err != nil { - return "", ErrProcessReap + return err } - systemDirectory = filepath.Clean(systemDirectory) - if systemDirectory == "." || !filepath.IsAbs(systemDirectory) { - return "", ErrProcessReap - } - taskkillPath := filepath.Join(systemDirectory, "taskkill.exe") - if !filepath.IsAbs(taskkillPath) || filepath.Dir(taskkillPath) != systemDirectory { - return "", ErrProcessReap - } - return taskkillPath, nil + return errors.Join( + windows.AssignProcessToJobObject(job.handle, process), + windows.CloseHandle(process), + ) } -func windowsTreeKillCommand(ctx context.Context, taskkillPath string, pid int) *exec.Cmd { - return exec.CommandContext(ctx, taskkillPath, "/PID", strconv.Itoa(pid), "/T", "/F") +func resumeWindowsProcess(pid int) error { + snapshot, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPTHREAD, 0) + if err != nil { + return err + } + threadID, threadErr := suspendedPrimaryThreadID(snapshot, uint32(pid)) + closeSnapshotErr := windows.CloseHandle(snapshot) + if threadErr != nil || closeSnapshotErr != nil { + return errors.Join(threadErr, closeSnapshotErr) + } + thread, err := windows.OpenThread(windows.THREAD_SUSPEND_RESUME, false, threadID) + if err != nil { + return err + } + _, resumeErr := windows.ResumeThread(thread) + closeThreadErr := windows.CloseHandle(thread) + return errors.Join(resumeErr, closeThreadErr) +} + +func suspendedPrimaryThreadID(snapshot windows.Handle, pid uint32) (uint32, error) { + entry := windows.ThreadEntry32{Size: uint32(unsafe.Sizeof(windows.ThreadEntry32{}))} + if err := windows.Thread32First(snapshot, &entry); err != nil { + return 0, err + } + threadIDs := make([]uint32, 0, 1) + for { + if entry.OwnerProcessID == pid { + threadIDs = append(threadIDs, entry.ThreadID) + } + entry.Size = uint32(unsafe.Sizeof(windows.ThreadEntry32{})) + err := windows.Thread32Next(snapshot, &entry) + if errors.Is(err, windows.ERROR_NO_MORE_FILES) { + break + } + if err != nil { + return 0, err + } + } + if len(threadIDs) != 1 { + return 0, ErrProcessReap + } + return threadIDs[0], nil +} + +// abortWindowsStart runs only while the direct process is still suspended or immediately after +// it was resumed unsuccessfully. It waits exactly once, so the normal runner path never races a +// second Wait or leaves a waiter goroutine behind. +func abortWindowsStart(command *exec.Cmd, job *windowsProcessJob) error { + var cleanupErr error + if job != nil && job.handle != 0 { + cleanupErr = errors.Join(cleanupErr, windows.TerminateJobObject(job.handle, jobTerminationExitCode)) + } + if command.Process != nil { + if err := command.Process.Kill(); err != nil && !errors.Is(err, os.ErrProcessDone) { + cleanupErr = errors.Join(cleanupErr, err) + } + if err := command.Wait(); err != nil { + var exitError *exec.ExitError + if !errors.As(err, &exitError) { + cleanupErr = errors.Join(cleanupErr, err) + } + } + } + if job != nil { + cleanupErr = errors.Join(cleanupErr, job.close()) + } + return cleanupErr } diff --git a/tools/tht/internal/compose/process_windows_static_test.go b/tools/tht/internal/compose/process_windows_static_test.go index 480330e4..01edbdab 100644 --- a/tools/tht/internal/compose/process_windows_static_test.go +++ b/tools/tht/internal/compose/process_windows_static_test.go @@ -6,28 +6,29 @@ import ( "testing" ) -func TestWindowsTerminationUsesBoundedExactPIDTreeKillWithoutAShell(t *testing.T) { +func TestWindowsTerminationOwnsTheProcessTreeWithAKillOnCloseJobObjectWithoutAShell(t *testing.T) { source, err := os.ReadFile("process_windows.go") if err != nil { t.Fatal(err) } text := string(source) for _, required := range []string{ - "context.WithTimeout", "windows.GetSystemDirectory", "filepath.Clean", "filepath.IsAbs", - `filepath.Join(systemDirectory, "taskkill.exe")`, "systemTaskkillPath", - `exec.CommandContext(ctx, taskkillPath, "/PID", strconv.Itoa(pid), "/T", "/F")`, - "io.Discard", "ErrProcessReap", + "windows.CreateJobObject", "windows.JOBOBJECT_EXTENDED_LIMIT_INFORMATION", + "windows.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE", "windows.SetInformationJobObject", + "windows.AssignProcessToJobObject", "windows.TerminateJobObject", "windows.OpenProcess", + "windows.CreateToolhelp32Snapshot", "windows.ResumeThread", "windows.CloseHandle", + "createSuspendedProcess", } { if !strings.Contains(text, required) { t.Errorf("process_windows.go is missing %q", required) } } for _, forbidden := range []string{ - "cmd.exe", "powershell", `exec.Command("taskkill.exe"`, - `exec.CommandContext(ctx, "taskkill.exe"`, `exec.LookPath("taskkill.exe")`, + "taskkill", "cmd.exe", "powershell", "exec.Command(", "exec.CommandContext(", + "exec.LookPath(", "windows.GetSystemDirectory", } { if strings.Contains(strings.ToLower(text), strings.ToLower(forbidden)) { - t.Errorf("process_windows.go contains unsafe command form %q", forbidden) + t.Errorf("process_windows.go contains unsafe or unowned process-tree mechanism %q", forbidden) } } } diff --git a/tools/tht/internal/compose/process_windows_test.go b/tools/tht/internal/compose/process_windows_test.go index c0f62867..daef0fce 100644 --- a/tools/tht/internal/compose/process_windows_test.go +++ b/tools/tht/internal/compose/process_windows_test.go @@ -3,38 +3,44 @@ package compose import ( - "context" - "errors" - "path/filepath" - "reflect" + "os/exec" "testing" + "unsafe" + + "golang.org/x/sys/windows" ) -func TestWindowsTreeKillCommandUsesTrustedAbsoluteSystemPathAndExactPIDArgumentArray(t *testing.T) { - original := windowsSystemDirectory - windowsSystemDirectory = func() (string, error) { return `C:\\Windows\\System32`, nil } - t.Cleanup(func() { windowsSystemDirectory = original }) - - taskkillPath, err := systemTaskkillPath() +func TestWindowsProcessJobUsesKillOnClose(t *testing.T) { + job, err := newWindowsProcessJob() if err != nil { t.Fatal(err) } - if !filepath.IsAbs(taskkillPath) { - t.Fatalf("taskkill path = %q, want absolute trusted system path", taskkillPath) + t.Cleanup(func() { + if closeErr := job.close(); closeErr != nil { + t.Error(closeErr) + } + }) + + var limits windows.JOBOBJECT_EXTENDED_LIMIT_INFORMATION + if err := windows.QueryInformationJobObject( + job.handle, + windows.JobObjectExtendedLimitInformation, + uintptr(unsafe.Pointer(&limits)), + uint32(unsafe.Sizeof(limits)), + nil, + ); err != nil { + t.Fatal(err) } - command := windowsTreeKillCommand(context.Background(), taskkillPath, 4242) - want := []string{taskkillPath, "/PID", "4242", "/T", "/F"} - if !reflect.DeepEqual(command.Args, want) { - t.Fatalf("taskkill args = %#v, want %#v", command.Args, want) + if limits.BasicLimitInformation.LimitFlags&windows.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE == 0 { + t.Fatalf("job limits = %#x, want JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE", limits.BasicLimitInformation.LimitFlags) } } -func TestWindowsSystemTaskkillPathRejectsRelativeSystemDirectory(t *testing.T) { - original := windowsSystemDirectory - windowsSystemDirectory = func() (string, error) { return `System32`, nil } - t.Cleanup(func() { windowsSystemDirectory = original }) +func TestWindowsProcessConfigurationStartsSuspendedBeforeJobAssignment(t *testing.T) { + command := exec.Command("docker") + configureProcess(command) - if _, err := systemTaskkillPath(); !errors.Is(err, ErrProcessReap) { - t.Fatalf("systemTaskkillPath() error = %v, want ErrProcessReap", err) + if command.SysProcAttr == nil || command.SysProcAttr.CreationFlags&createSuspendedProcess == 0 { + t.Fatalf("configureProcess() flags = %#x, want CREATE_SUSPENDED", command.SysProcAttr.CreationFlags) } } diff --git a/tools/tht/internal/compose/runner.go b/tools/tht/internal/compose/runner.go index 16dbd64f..a3a10b76 100644 --- a/tools/tht/internal/compose/runner.go +++ b/tools/tht/internal/compose/runner.go @@ -31,8 +31,6 @@ var ErrProcessReap = errors.New("Docker command process could not be reaped") // ErrContainerCleanup reports that Docker could not confirm removal of a named one-shot container. var ErrContainerCleanup = errors.New("Docker one-shot container cleanup failed") -var terminateProcessForRunner = terminateProcess - // CaptureLimits bounds each captured stream while the child is running. type CaptureLimits struct { StdoutBytes int @@ -115,7 +113,8 @@ func (r execRunner) runBounded(ctx context.Context, args []string, stdin io.Read stderr := newCappedBuffer(limits.StderrBytes, overflow, stderrObserver) command.Stdout = stdout command.Stderr = stderr - if err := command.Start(); err != nil { + process, err := startProcessForRunner(command) + if err != nil { return startFailure(err) } done := make(chan error, 1) @@ -128,16 +127,19 @@ func (r execRunner) runBounded(ctx context.Context, args []string, stdin io.Read case <-ctx.Done(): interrupted = true processErr = ctx.Err() - if err := terminateProcessForRunner(command, done); err != nil { - lifecycleErr = errors.Join(lifecycleErr, ErrProcessReap) + if err := terminateProcessForRunner(process, done); err != nil { + lifecycleErr = errors.Join(lifecycleErr, err) } case <-overflow: interrupted = true processErr = ErrOutputLimit - if err := terminateProcessForRunner(command, done); err != nil { - lifecycleErr = errors.Join(lifecycleErr, ErrProcessReap) + if err := terminateProcessForRunner(process, done); err != nil { + lifecycleErr = errors.Join(lifecycleErr, err) } } + if err := process.close(); err != nil { + lifecycleErr = errors.Join(lifecycleErr, ErrProcessReap, err) + } if containerName != "" { if err := r.cleanupOneShotContainer(containerName); err != nil { lifecycleErr = errors.Join(lifecycleErr, ErrContainerCleanup) diff --git a/tools/tht/internal/compose/runner_test.go b/tools/tht/internal/compose/runner_test.go index 0e6a8d8a..afb30a2c 100644 --- a/tools/tht/internal/compose/runner_test.go +++ b/tools/tht/internal/compose/runner_test.go @@ -245,8 +245,8 @@ func TestRunnerSurfacesCleanupFailureWithoutContainerName(t *testing.T) { func TestRunnerPropagatesAReapingFailure(t *testing.T) { original := terminateProcessForRunner - terminateProcessForRunner = func(command *exec.Cmd, done <-chan error) error { - _ = terminateProcess(command, done) + terminateProcessForRunner = func(process processLifetime, done <-chan error) error { + _ = process.terminate(done) return ErrProcessReap } t.Cleanup(func() { terminateProcessForRunner = original })