From d12b629420036d8491c782cfa287e411c862275b Mon Sep 17 00:00:00 2001 From: mptyl Date: Mon, 17 Aug 2026 18:41:40 +0200 Subject: [PATCH] fix(auth): harden diagnostic cleanup races --- frontend/src/shell/WorkspaceManager.test.tsx | 158 ++++++++++++++++++ frontend/src/shell/WorkspaceManager.tsx | 46 +++-- .../internal/compose/process_termination.go | 41 +++++ .../compose/process_termination_test.go | 89 ++++++++++ tools/tht/internal/compose/process_windows.go | 49 ++++-- .../compose/process_windows_static_test.go | 11 +- .../internal/compose/process_windows_test.go | 29 +++- 7 files changed, 387 insertions(+), 36 deletions(-) create mode 100644 tools/tht/internal/compose/process_termination.go create mode 100644 tools/tht/internal/compose/process_termination_test.go diff --git a/frontend/src/shell/WorkspaceManager.test.tsx b/frontend/src/shell/WorkspaceManager.test.tsx index ced56ec3..0587ae51 100644 --- a/frontend/src/shell/WorkspaceManager.test.tsx +++ b/frontend/src/shell/WorkspaceManager.test.tsx @@ -327,6 +327,7 @@ test("ignores an older validation response that completes after a newer connecti act(() => releaseTest()); const authentication = await screen.findByTestId("workspace-authentication"); expect(within(authentication).getByText(/Newer connection result/)).toBeVisible(); + expect(screen.getByRole("button", { name: "Validate workspace source" })).not.toBeDisabled(); act(() => releaseValidation()); await act(async () => { @@ -335,6 +336,163 @@ test("ignores an older validation response that completes after a newer connecti }); expect(within(authentication).getByText(/Newer connection result/)).toBeVisible(); expect(within(authentication).queryByText("Passed")).not.toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Validate workspace source" })).not.toBeDisabled(); +}); + +test("does not apply a diagnostic that completes after its dialog is closed and reopened", 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 client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + const onClose = vi.fn(); + const view = render( + + + , + ); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + await user.click(screen.getByRole("button", { name: "Validate workspace source" })); + await requestStarted; + + await user.click(screen.getByRole("button", { name: "Close workspace management" })); + expect(onClose).toHaveBeenCalledTimes(1); + view.rerender( + + + , + ); + view.rerender( + + + , + ); + 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("does not apply a diagnostic that completes after returning through Level 1", 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(); + } + })); + renderManager(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + await user.click(screen.getByRole("button", { name: "Validate workspace source" })); + await requestStarted; + + await user.click(within(screen.getByRole("navigation", { name: "Workspaces" })) + .getByRole("button", { name: "Back to Level 1" })); + expect(screen.getByTestId("workspace-overview")).toBeVisible(); + await user.click(screen.getByRole("button", { name: "PSD Clinical" })); + 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("does not apply a diagnostic after changing workspaces and returning", 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.post("/api/workspaces/validate", async () => { + started(); + try { + await held; + return HttpResponse.json({ + workspace, contract: {}, activatable: true, diagnostics: [], authentication: readyAuthentication, + }); + } finally { + settled(); + } + }), + ); + renderManager(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + await user.click(screen.getByRole("button", { name: "Validate workspace source" })); + 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" })); + 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("never renders a hostile authentication field rejected by the API decoder", async () => { diff --git a/frontend/src/shell/WorkspaceManager.tsx b/frontend/src/shell/WorkspaceManager.tsx index a70f0c00..3a718330 100644 --- a/frontend/src/shell/WorkspaceManager.tsx +++ b/frontend/src/shell/WorkspaceManager.tsx @@ -87,8 +87,37 @@ export function WorkspaceManager({ const operationEpochRef = useRef(0); const diagnosticEpochRef = useRef(0); const selectedIdRef = useRef(selectedId); + const wasOpenRef = useRef(open); selectedIdRef.current = selectedId; + + const clearGlobalMessages = () => { + 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(); + }; + useEffect(() => () => { operationEpochRef.current += 1; }, []); + useEffect(() => { + if (wasOpenRef.current && !open) { + setSecretValues({}); + clearMessages(); + } + wasOpenRef.current = open; + }, [open]); async function guardedQuery(request: () => Promise, targetId?: string): Promise { const guard = captureAuthOperation({ @@ -131,21 +160,6 @@ export function WorkspaceManager({ enabled: Boolean(open && selectedId), }); - const clearMessages = () => { - setNotice(undefined); - setDiagnostics([]); - setValidationNotice(undefined); - setValidationDiagnostics([]); - setConnectionNotice(undefined); - setConnectionDiagnostics([]); - setAuthentication(undefined); - }; - - const clearGlobalMessages = () => { - setNotice(undefined); - setDiagnostics([]); - }; - const close = () => { setSecretValues({}); clearMessages(); @@ -167,8 +181,8 @@ export function WorkspaceManager({ async function updateRepository() { const guard = captureAuthOperation({ disposalEpoch: operationEpochRef.current }); if (!guard) return; - setBusyAction("repository"); clearMessages(); + setBusyAction("repository"); try { await pullWorkspaceRegistry(); if (!isAuthOperationCurrent(guard, { disposalEpoch: operationEpochRef.current })) return; diff --git a/tools/tht/internal/compose/process_termination.go b/tools/tht/internal/compose/process_termination.go new file mode 100644 index 00000000..87b731d0 --- /dev/null +++ b/tools/tht/internal/compose/process_termination.go @@ -0,0 +1,41 @@ +package compose + +import "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( + done <-chan error, + terminateTree func() error, + terminateDirect func() error, + bound time.Duration, +) error { + if processReaped(done) { + return nil + } + _ = terminateTree() + if processReaped(done) { + return nil + } + _ = terminateDirect() + if processReaped(done) { + return nil + } + timer := time.NewTimer(bound) + defer timer.Stop() + select { + case <-done: + return nil + case <-timer.C: + 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 new file mode 100644 index 00000000..7ef6eedf --- /dev/null +++ b/tools/tht/internal/compose/process_termination_test.go @@ -0,0 +1,89 @@ +package compose + +import ( + "errors" + "testing" +) + +func TestWindowsCancellationTreatsProcessReapedBeforeTreeTerminationAsSuccess(t *testing.T) { + done := make(chan error, 1) + done <- nil + treeCalls := 0 + directCalls := 0 + + err := terminateWindowsProcess(done, func() error { + treeCalls++ + return errors.New("tree termination should not run") + }, func() error { + directCalls++ + return errors.New("direct termination should not run") + }, 0) + + if err != nil { + t.Fatalf("terminateWindowsProcess() error = %v, want nil", err) + } + if treeCalls != 0 || directCalls != 0 { + t.Fatalf("termination calls = tree:%d direct:%d, want none", treeCalls, directCalls) + } +} + +func TestWindowsCancellationTreatsProcessReapedAfterTreeTerminationAsSuccess(t *testing.T) { + done := make(chan error, 1) + directCalls := 0 + + err := terminateWindowsProcess(done, func() error { + done <- nil + return ErrProcessReap + }, func() error { + directCalls++ + return errors.New("direct termination should not run") + }, 0) + + if err != nil { + t.Fatalf("terminateWindowsProcess() 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 + + err := terminateWindowsProcess(done, func() error { + return ErrProcessReap + }, func() error { + directCalls++ + done <- nil + return ErrProcessReap + }, 0) + + if err != nil { + t.Fatalf("terminateWindowsProcess() error = %v, want nil", err) + } + if directCalls != 1 { + t.Fatalf("direct termination calls = %d, want 1", directCalls) + } +} + +func TestWindowsCancellationReportsReapFailureWhenProcessSurvives(t *testing.T) { + done := make(chan error) + treeCalls := 0 + directCalls := 0 + + err := terminateWindowsProcess(done, func() error { + treeCalls++ + return ErrProcessReap + }, func() error { + directCalls++ + return ErrProcessReap + }, 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) + } +} diff --git a/tools/tht/internal/compose/process_windows.go b/tools/tht/internal/compose/process_windows.go index 2fdf5d56..8170e7ea 100644 --- a/tools/tht/internal/compose/process_windows.go +++ b/tools/tht/internal/compose/process_windows.go @@ -6,15 +6,20 @@ import ( "context" "io" "os/exec" + "path/filepath" "strconv" "syscall" "time" + + "golang.org/x/sys/windows" ) const finalTerminationBound = 2 * time.Second const processTreeTerminationBound = 2 * time.Second const createNewProcessGroup = 0x00000200 +var windowsSystemDirectory = windows.GetSystemDirectory + func configureProcess(command *exec.Cmd) { command.SysProcAttr = &syscall.SysProcAttr{CreationFlags: createNewProcessGroup} } @@ -23,23 +28,22 @@ func terminateProcess(command *exec.Cmd, done <-chan error) error { if command.Process == nil { return nil } - treeErr := terminateWindowsProcessTree(command.Process.Pid) - _ = command.Process.Kill() - select { - case <-done: - if treeErr != nil { - return ErrProcessReap - } - return nil - case <-time.After(finalTerminationBound): - return ErrProcessReap - } + return terminateWindowsProcess( + done, + func() error { return terminateWindowsProcessTree(command.Process.Pid) }, + command.Process.Kill, + finalTerminationBound, + ) } func terminateWindowsProcessTree(pid int) error { ctx, cancel := context.WithTimeout(context.Background(), processTreeTerminationBound) defer cancel() - command := windowsTreeKillCommand(ctx, pid) + taskkillPath, err := systemTaskkillPath() + if err != nil { + return ErrProcessReap + } + command := windowsTreeKillCommand(ctx, taskkillPath, pid) command.Stdout = io.Discard command.Stderr = io.Discard if err := command.Run(); err != nil { @@ -51,6 +55,23 @@ func terminateWindowsProcessTree(pid int) error { return nil } -func windowsTreeKillCommand(ctx context.Context, pid int) *exec.Cmd { - return exec.CommandContext(ctx, "taskkill.exe", "/PID", strconv.Itoa(pid), "/T", "/F") +// systemTaskkillPath resolves taskkill through Windows' protected system-directory API, never PATH. +func systemTaskkillPath() (string, error) { + systemDirectory, err := windowsSystemDirectory() + if err != nil { + return "", ErrProcessReap + } + 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 +} + +func windowsTreeKillCommand(ctx context.Context, taskkillPath string, pid int) *exec.Cmd { + return exec.CommandContext(ctx, taskkillPath, "/PID", strconv.Itoa(pid), "/T", "/F") } diff --git a/tools/tht/internal/compose/process_windows_static_test.go b/tools/tht/internal/compose/process_windows_static_test.go index c9428832..480330e4 100644 --- a/tools/tht/internal/compose/process_windows_static_test.go +++ b/tools/tht/internal/compose/process_windows_static_test.go @@ -13,14 +13,19 @@ func TestWindowsTerminationUsesBoundedExactPIDTreeKillWithoutAShell(t *testing.T } text := string(source) for _, required := range []string{ - "context.WithTimeout", "exec.CommandContext", `"taskkill.exe"`, `"/PID"`, - "strconv.Itoa", `"/T"`, `"/F"`, "io.Discard", "ErrProcessReap", + "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", } { 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"`} { + for _, forbidden := range []string{ + "cmd.exe", "powershell", `exec.Command("taskkill.exe"`, + `exec.CommandContext(ctx, "taskkill.exe"`, `exec.LookPath("taskkill.exe")`, + } { if strings.Contains(strings.ToLower(text), strings.ToLower(forbidden)) { t.Errorf("process_windows.go contains unsafe command form %q", forbidden) } diff --git a/tools/tht/internal/compose/process_windows_test.go b/tools/tht/internal/compose/process_windows_test.go index 84b90c6c..c0f62867 100644 --- a/tools/tht/internal/compose/process_windows_test.go +++ b/tools/tht/internal/compose/process_windows_test.go @@ -4,14 +4,37 @@ package compose import ( "context" + "errors" + "path/filepath" "reflect" "testing" ) -func TestWindowsTreeKillCommandUsesExactPIDArgumentArray(t *testing.T) { - command := windowsTreeKillCommand(context.Background(), 4242) - want := []string{"taskkill.exe", "/PID", "4242", "/T", "/F"} +func TestWindowsTreeKillCommandUsesTrustedAbsoluteSystemPathAndExactPIDArgumentArray(t *testing.T) { + original := windowsSystemDirectory + windowsSystemDirectory = func() (string, error) { return `C:\\Windows\\System32`, nil } + t.Cleanup(func() { windowsSystemDirectory = original }) + + taskkillPath, err := systemTaskkillPath() + if err != nil { + t.Fatal(err) + } + if !filepath.IsAbs(taskkillPath) { + t.Fatalf("taskkill path = %q, want absolute trusted system path", taskkillPath) + } + 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) } } + +func TestWindowsSystemTaskkillPathRejectsRelativeSystemDirectory(t *testing.T) { + original := windowsSystemDirectory + windowsSystemDirectory = func() (string, error) { return `System32`, nil } + t.Cleanup(func() { windowsSystemDirectory = original }) + + if _, err := systemTaskkillPath(); !errors.Is(err, ErrProcessReap) { + t.Fatalf("systemTaskkillPath() error = %v, want ErrProcessReap", err) + } +}