diff --git a/frontend/src/shell/WorkspaceManager.test.tsx b/frontend/src/shell/WorkspaceManager.test.tsx index 075be8a5..ced56ec3 100644 --- a/frontend/src/shell/WorkspaceManager.test.tsx +++ b/frontend/src/shell/WorkspaceManager.test.tsx @@ -273,8 +273,68 @@ test("clears a previous authentication result as soon as validation is retried", await waitFor(() => expect(calls).toBe(2)); expect(screen.queryByTestId("workspace-authentication")).not.toBeInTheDocument(); - releaseRetry(); - expect(await screen.findByTestId("workspace-authentication")).toBeVisible(); + releaseRetry(); + expect(await screen.findByTestId("workspace-authentication")).toBeVisible(); +}); + +test("ignores an older validation response that completes after a newer connection test", async () => { + const user = userEvent.setup(); + let releaseValidation!: () => void; + let releaseTest!: () => void; + let validationStarted!: () => void; + let testStarted!: () => void; + let validationSettled!: () => void; + const heldValidation = new Promise((resolve) => { releaseValidation = resolve; }); + const heldTest = new Promise((resolve) => { releaseTest = resolve; }); + const validationRequestStarted = new Promise((resolve) => { validationStarted = resolve; }); + const testRequestStarted = new Promise((resolve) => { testStarted = resolve; }); + const validationRequestSettled = new Promise((resolve) => { validationSettled = resolve; }); + server.use( + http.post("/api/workspaces/validate", async () => { + validationStarted(); + try { + await heldValidation; + return HttpResponse.json({ + workspace, contract: {}, activatable: true, diagnostics: [], authentication: readyAuthentication, + }); + } finally { + validationSettled(); + } + }), + http.post("/api/workspaces/psd-clinical/test", async () => { + testStarted(); + await heldTest; + return HttpResponse.json({ + activatable: false, + diagnostics: [], + authentication: { + ready: false, + mode: "oidc", + checks: [{ level: "error", code: "oidc_secret_missing", message: "Newer connection result." }], + }, + }); + }), + ); + renderManager(); + await user.click(await screen.findByRole("button", { name: "PSD Clinical" })); + + await user.click(screen.getByRole("button", { name: "Validate workspace source" })); + await validationRequestStarted; + await user.click(screen.getByRole("button", { name: "Test workspace connections" })); + await testRequestStarted; + expect(screen.queryByTestId("workspace-authentication")).not.toBeInTheDocument(); + + act(() => releaseTest()); + const authentication = await screen.findByTestId("workspace-authentication"); + expect(within(authentication).getByText(/Newer connection result/)).toBeVisible(); + + act(() => releaseValidation()); + await act(async () => { + await validationRequestSettled; + await new Promise((resolve) => { setTimeout(resolve, 50); }); + }); + expect(within(authentication).getByText(/Newer connection result/)).toBeVisible(); + expect(within(authentication).queryByText("Passed")).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 aee9013a..a70f0c00 100644 --- a/frontend/src/shell/WorkspaceManager.tsx +++ b/frontend/src/shell/WorkspaceManager.tsx @@ -85,6 +85,7 @@ export function WorkspaceManager({ const [authentication, setAuthentication] = useState(); const [busyAction, setBusyAction] = useState(); const operationEpochRef = useRef(0); + const diagnosticEpochRef = useRef(0); const selectedIdRef = useRef(selectedId); selectedIdRef.current = selectedId; useEffect(() => () => { operationEpochRef.current += 1; }, []); @@ -192,6 +193,7 @@ export function WorkspaceManager({ if (!detailQuery.data) return; const guard = captureAuthOperation({ sessionId: selectedId, disposalEpoch: operationEpochRef.current }); if (!guard) return; + const diagnosticEpoch = ++diagnosticEpochRef.current; setBusyAction("validate"); clearGlobalMessages(); setAuthentication(undefined); @@ -199,15 +201,18 @@ export function WorkspaceManager({ setValidationDiagnostics([]); try { const result = await validateWorkspace(detailQuery.data.workspace); - if (!isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; + if (diagnosticEpoch !== diagnosticEpochRef.current || + !isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; setAuthentication(result.authentication); setValidationNotice(result.activatable ? "Workspace source and authentication are valid." : "Workspace source is valid."); } catch (error) { - if (isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) { + if (diagnosticEpoch === diagnosticEpochRef.current && + isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) { setValidationDiagnostics([publicError(error, "workspace_invalid: Workspace validation could not be completed")]); } } finally { - if (isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); + if (diagnosticEpoch === diagnosticEpochRef.current && + isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); } } @@ -216,6 +221,7 @@ export function WorkspaceManager({ const targetId = selectedId; const guard = captureAuthOperation({ sessionId: targetId, disposalEpoch: operationEpochRef.current }); if (!guard) return; + const diagnosticEpoch = ++diagnosticEpochRef.current; setBusyAction("test"); clearGlobalMessages(); setAuthentication(undefined); @@ -223,7 +229,8 @@ export function WorkspaceManager({ setConnectionDiagnostics([]); try { const result = await testWorkspace(selectedId); - if (!isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; + if (diagnosticEpoch !== diagnosticEpochRef.current || + !isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) return; setAuthentication(result.authentication); const issues = result.diagnostics.filter(({ level }) => level !== "info"); const informational = result.diagnostics.find(({ level }) => level === "info"); @@ -236,11 +243,13 @@ export function WorkspaceManager({ : "Workspace connection test completed."); } } catch (error) { - if (isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) { + if (diagnosticEpoch === diagnosticEpochRef.current && + isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) { setConnectionDiagnostics([publicError(error, "connector_unavailable: Workspace connections could not be tested")]); } } finally { - if (isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); + if (diagnosticEpoch === diagnosticEpochRef.current && + isAuthOperationCurrent(guard, { sessionId: selectedIdRef.current, disposalEpoch: operationEpochRef.current })) setBusyAction(undefined); } } diff --git a/tools/tht/internal/authconfig/commands.go b/tools/tht/internal/authconfig/commands.go index eefa8d4a..7d6e9d48 100644 --- a/tools/tht/internal/authconfig/commands.go +++ b/tools/tht/internal/authconfig/commands.go @@ -16,6 +16,7 @@ import ( "os/exec" "path/filepath" "strings" + "sync" "time" "unicode" "unicode/utf16" @@ -37,6 +38,10 @@ const interactiveAuthCheckTimeout = 12 * time.Minute const maxAuthDiagnosticOutputBytes = 64 * 1024 +const maxDevicePromptLineBytes = 4 * 1024 + +const maxDeviceVerificationURIBytes = 2 * 1024 + var errCommandRefused = errors.New("authentication command refused") var writeNewAuthFile = safeio.WriteCanonicalNewFile @@ -131,14 +136,11 @@ func checkCommand(ctx context.Context, installation config.Installation, args [] if err != nil { return authFailure(stderr, authMessage(err)) } - report, prompt, err := runCheck(ctx, installation, runner, interactive, false) + report, err := runCheck(ctx, installation, runner, interactive, false, stderr) if err != nil { fmt.Fprintln(stderr, "tht: authentication diagnostics could not be completed") return 1 } - if interactive && prompt != "" { - fmt.Fprintln(stderr, prompt) - } if jsonMode { if err := json.NewEncoder(stdout).Encode(report); err != nil { fmt.Fprintln(stderr, "tht: authentication diagnostic report could not be written") @@ -292,28 +294,102 @@ func sanitizeAuthDiagnostics(report AuthDiagnostics, secrets []string) AuthDiagn return report } -func devicePrompt(stderr string, secrets []string) string { - if len(stderr) > maxAuthDiagnosticOutputBytes { +type devicePromptStream struct { + mu sync.Mutex + pending []byte + discarding bool + emitted bool + secrets []string + output io.Writer + writeErr error +} + +func newDevicePromptStream(secrets []string, outputWriter io.Writer) *devicePromptStream { + return &devicePromptStream{ + pending: make([]byte, 0, 256), secrets: append([]string(nil), secrets...), output: outputWriter, + } +} + +func (stream *devicePromptStream) Observe(chunk []byte) { + stream.mu.Lock() + defer stream.mu.Unlock() + for _, value := range chunk { + if value == '\n' { + if !stream.discarding { + stream.emitLineLocked(string(stream.pending)) + } + stream.pending = stream.pending[:0] + stream.discarding = false + continue + } + if stream.discarding { + continue + } + if len(stream.pending) >= maxDevicePromptLineBytes { + stream.pending = stream.pending[:0] + stream.discarding = true + continue + } + stream.pending = append(stream.pending, value) + } +} + +func (stream *devicePromptStream) Finish() error { + stream.mu.Lock() + defer stream.mu.Unlock() + if !stream.discarding && len(stream.pending) > 0 { + stream.emitLineLocked(string(stream.pending)) + } + stream.pending = nil + stream.discarding = false + return stream.writeErr +} + +func (stream *devicePromptStream) emitLineLocked(line string) { + if stream.emitted || stream.output == nil { + return + } + prompt := parseDevicePromptLine(output.Sanitize(line, stream.secrets)) + if prompt == "" { + return + } + stream.emitted = true + if _, err := fmt.Fprintln(stream.output, prompt); err != nil { + stream.writeErr = errors.New("authentication device prompt could not be written") + } +} + +func parseDevicePromptLine(line string) string { + const prefix = "Open " + const separator = " and enter code " + if len(line) == 0 || len(line) > maxDevicePromptLineBytes || !utf8.ValidString(line) || + strings.TrimSpace(line) != line || !strings.HasPrefix(line, prefix) { return "" } - for _, line := range strings.Split(output.Sanitize(stderr, secrets), "\n") { - const prefix = "Open " - const separator = " and enter code " - if !strings.HasPrefix(line, prefix) { - continue + for _, character := range line { + if unicode.IsControl(character) { + return "" } - parts := strings.Split(strings.TrimPrefix(line, prefix), separator) - if len(parts) != 2 || !validDeviceVerificationURI(parts[0]) || !validDeviceUserCode(parts[1]) { - continue - } - return prefix + parts[0] + separator + parts[1] } - return "" + payload := strings.TrimPrefix(line, prefix) + if strings.Count(payload, separator) != 1 { + return "" + } + verificationURI, userCode, found := strings.Cut(payload, separator) + if !found || !validDeviceVerificationURI(verificationURI) || !validDeviceUserCode(userCode) { + return "" + } + return prefix + verificationURI + separator + userCode } func validDeviceVerificationURI(value string) bool { + if len(value) == 0 || len(value) > maxDeviceVerificationURIBytes || !utf8.ValidString(value) { + return false + } parsed, err := url.Parse(value) - return err == nil && parsed.Scheme == "https" && parsed.Host != "" && parsed.User == nil && parsed.Fragment == "" + return err == nil && parsed.IsAbs() && parsed.Scheme == "https" && parsed.Host != "" && + parsed.Hostname() != "" && parsed.User == nil && parsed.Opaque == "" && parsed.Fragment == "" && + !strings.ContainsAny(value, " \t\r\n") } func validDeviceUserCode(value string) bool { @@ -334,17 +410,16 @@ func validDeviceUserCode(value string) bool { // Check invokes the exact backend diagnostic command. Normal host checks always use one-shot // Compose execution; aggregate doctor passes useRunningCore=true only after confirming core runs. func Check(ctx context.Context, installation config.Installation, runner compose.Runner, interactive, useRunningCore bool) (AuthDiagnostics, error) { - report, _, err := runCheck(ctx, installation, runner, interactive, useRunningCore) - return report, err + return runCheck(ctx, installation, runner, interactive, useRunningCore, nil) } -func runCheck(ctx context.Context, installation config.Installation, runner compose.Runner, interactive, useRunningCore bool) (AuthDiagnostics, string, error) { +func runCheck(ctx context.Context, installation config.Installation, runner compose.Runner, interactive, useRunningCore bool, promptOutput io.Writer) (AuthDiagnostics, error) { if runner == nil { - return AuthDiagnostics{}, "", errors.New("authentication diagnostic runner is unavailable") + return AuthDiagnostics{}, errors.New("authentication diagnostic runner is unavailable") } secrets, err := authenticationSecretValues(installation) if err != nil { - return AuthDiagnostics{}, "", err + return AuthDiagnostics{}, err } timeout := authCheckTimeout if interactive { @@ -356,35 +431,48 @@ func runCheck(ctx context.Context, installation config.Installation, runner comp if !useRunningCore { containerName, err := compose.NewOneShotContainerName("thothii-auth-check") if err != nil { - return AuthDiagnostics{}, "", errors.New("authentication diagnostic command failed") + return AuthDiagnostics{}, errors.New("authentication diagnostic command failed") } command = []string{"run", "--rm", "--no-deps", "--no-TTY", "--name", containerName, "core", "node", "dist/auth/diagnostic-command.js", "--json"} } if interactive { command = append(command, "--interactive") } - result, err := compose.RunBounded(runner, bounded, installation.ComposeArgs(command...), nil, compose.CaptureLimits{ + var promptStream *devicePromptStream + var stderrObserver func([]byte) + if interactive && promptOutput != nil { + promptStream = newDevicePromptStream(secrets, promptOutput) + stderrObserver = promptStream.Observe + } + result, err := compose.RunBoundedStreaming(runner, bounded, installation.ComposeArgs(command...), nil, compose.CaptureLimits{ StdoutBytes: maxAuthDiagnosticOutputBytes, StderrBytes: maxAuthDiagnosticOutputBytes, - }) + }, stderrObserver) + if promptStream != nil { + if promptErr := promptStream.Finish(); promptErr != nil { + return AuthDiagnostics{}, promptErr + } + } var exitError *exec.ExitError - validProcessOutcome := result.ExitCode == 0 && err == nil || - result.ExitCode == 1 && (err == nil || errors.As(err, &exitError)) + unsafeLifecycle := errors.Is(err, compose.ErrOutputLimit) || errors.Is(err, compose.ErrProcessReap) || + errors.Is(err, compose.ErrContainerCleanup) || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) + validProcessOutcome := !unsafeLifecycle && (result.ExitCode == 0 && err == nil || + result.ExitCode == 1 && (err == nil || errors.As(err, &exitError))) if !validProcessOutcome { - return AuthDiagnostics{}, "", errors.New("authentication diagnostic command failed") + return AuthDiagnostics{}, errors.New("authentication diagnostic command failed") } report, decodeErr := decodeAuthDiagnostics(result.Stdout) if decodeErr != nil { - return AuthDiagnostics{}, "", errors.New("authentication diagnostic report is invalid") + return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid") } if report.Ready != (result.ExitCode == 0) { - return AuthDiagnostics{}, "", errors.New("authentication diagnostic command failed") + return AuthDiagnostics{}, errors.New("authentication diagnostic command failed") } safe := sanitizeAuthDiagnostics(report, secrets) if !validAuthDiagnostics(safe) { - return AuthDiagnostics{}, "", errors.New("authentication diagnostic report is invalid") + return AuthDiagnostics{}, errors.New("authentication diagnostic report is invalid") } - return safe, devicePrompt(result.Stderr, secrets), nil + return safe, nil } func authFailure(stderr io.Writer, message string) int { diff --git a/tools/tht/internal/authconfig/commands_test.go b/tools/tht/internal/authconfig/commands_test.go index 0d10f329..21eb4958 100644 --- a/tools/tht/internal/authconfig/commands_test.go +++ b/tools/tht/internal/authconfig/commands_test.go @@ -9,7 +9,9 @@ import ( "os" "path/filepath" "strings" + "sync" "testing" + "time" "github.com/aritmolab/thothii/tools/tht/internal/compose" "github.com/aritmolab/thothii/tools/tht/internal/config" @@ -49,6 +51,7 @@ func TestAuthCheckRunsOneShotCoreDiagnosticWithPristineJSON(t *testing.T) { func TestAuthCheckEmitsValidFailedReportFromRealExitError(t *testing.T) { installation := authInstallation(newAuthDirectory(t)) runner := compose.NewRunner(writeAuthExecutable(t, `#!/bin/sh +if [ "$1" = "container" ]; then exit 0; fi printf '%s\n' '{"ready":false,"mode":"oidc","checks":[{"level":"error","code":"oidc_secret_missing","message":"A required OIDC or group catalog secret is unavailable."}]}' exit 1 `)) @@ -187,6 +190,90 @@ func TestAuthCheckInteractiveForwardsOnlyTheValidatedDevicePrompt(t *testing.T) assertAuthOneShotCommand(t, installation, calls[0], true) } +func TestAuthCheckInteractiveStreamsOnlyASanitizedPromptBeforePollingCompletes(t *testing.T) { + installation := authInstallation(newAuthDirectory(t)) + root := filepath.Dir(installation.EnvFile) + secretFile := filepath.Join(root, "prompt-secret") + if err := os.WriteFile(secretFile, []byte("prompt-secret-sentinel"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(installation.EnvFile, []byte("PROMPT_TOKEN_FILE="+secretFile+"\n"), 0o600); err != nil { + t.Fatal(err) + } + started := filepath.Join(root, "started") + release := filepath.Join(root, "release") + t.Setenv("THT_AUTH_STREAM_STARTED", started) + t.Setenv("THT_AUTH_STREAM_RELEASE", release) + runner := compose.NewRunner(writeAuthExecutable(t, `#!/bin/sh +if [ "$1" = "container" ]; then exit 0; fi +printf '%s\n' 'untrusted prompt-secret-sentinel child output' >&2 +printf '%s' 'Open https://issuer.example.test/device?token=' >&2 +printf '%s\n' 'prompt-secret-sentinel and enter code ABCD-EFGH' >&2 +: > "$THT_AUTH_STREAM_STARTED" +while [ ! -f "$THT_AUTH_STREAM_RELEASE" ]; do sleep 0.01; done +printf '%s\n' '{"ready":true,"mode":"oidc","checks":[{"level":"info","code":"auth_ready","message":"Authentication is ready."}]}' +`)) + var stdout bytes.Buffer + stderr := &synchronizedBuffer{} + done := make(chan int, 1) + go func() { + done <- RunWithRunner(context.Background(), installation, []string{"check", "--json", "--interactive"}, strings.NewReader(""), &stdout, stderr, runner) + }() + + if !waitForFile(started, 5*time.Second) { + t.Fatal("interactive diagnostic did not start") + } + deadline := time.Now().Add(2 * time.Second) + streamed := false + for time.Now().Before(deadline) { + if stderr.String() == "Open https://issuer.example.test/device?token=[REDACTED] and enter code ABCD-EFGH\n" { + streamed = true + break + } + time.Sleep(10 * time.Millisecond) + } + if err := os.WriteFile(release, []byte("release"), 0o600); err != nil { + t.Fatal(err) + } + select { + case code := <-done: + if code != 0 { + t.Fatalf("interactive auth check = %d, stdout=%q stderr=%q", code, stdout.String(), stderr.String()) + } + case <-time.After(5 * time.Second): + t.Fatal("interactive diagnostic did not finish") + } + if !streamed { + t.Fatalf("device prompt was not streamed before polling completed: stderr=%q", stderr.String()) + } + if strings.Contains(stderr.String(), "prompt-secret-sentinel") || strings.Contains(stderr.String(), "untrusted") { + t.Fatalf("interactive auth leaked raw child stderr: %q", stderr.String()) + } + if !json.Valid(stdout.Bytes()) { + t.Fatalf("interactive JSON stdout is not pristine: %q", stdout.String()) + } +} + +func TestDevicePromptGrammarRejectsUnboundedAmbiguousOrUnsafeLines(t *testing.T) { + valid := "Open https://issuer.example.test/device?tenant=one and enter code ABCD-EFGH" + if got := parseDevicePromptLine(valid); got != valid { + t.Fatalf("valid prompt = %q, want %q", got, valid) + } + for _, line := range []string{ + "Open http://issuer.example.test/device and enter code ABCD-EFGH", + "Open https://user@issuer.example.test/device and enter code ABCD-EFGH", + "Open https://issuer.example.test/device#fragment and enter code ABCD-EFGH", + "Open https://issuer.example.test/device and enter code ABCD-EFGH and enter code IJKL-MNOP", + "Open https://issuer.example.test/device and enter code ABCD EFGH", + "Open https://issuer.example.test/device\r and enter code ABCD-EFGH", + "Open https://issuer.example.test/" + strings.Repeat("a", maxDevicePromptLineBytes) + " and enter code ABCD-EFGH", + } { + if prompt := parseDevicePromptLine(line); prompt != "" { + t.Fatalf("unsafe prompt accepted: %q", prompt) + } + } +} + func TestAuthCheckFailsBeforeExecutionWhenDeclaredSecretCorpusIsIncomplete(t *testing.T) { installation := authInstallation(newAuthDirectory(t)) installation.EnvFile = filepath.Join(t.TempDir(), "operator.env") @@ -201,13 +288,29 @@ func TestAuthCheckFailsBeforeExecutionWhenDeclaredSecretCorpusIsIncomplete(t *te }) var stdout, stderr bytes.Buffer - code := RunWithRunner(context.Background(), installation, []string{"check", "--json"}, strings.NewReader(""), &stdout, &stderr, runner) + code := RunWithRunner(context.Background(), installation, []string{"check", "--json", "--interactive"}, strings.NewReader(""), &stdout, &stderr, runner) if code != 1 || calls != 0 || stdout.Len() != 0 || stderr.String() != "tht: authentication diagnostics could not be completed\n" { t.Fatalf("incomplete corpus was not refused before execution: code=%d calls=%d stdout=%q stderr=%q", code, calls, stdout.String(), stderr.String()) } } +func TestAuthCheckRejectsAValidReportWhenOneShotCleanupFails(t *testing.T) { + runner := runnerFunc(func(_ context.Context, _ []string, _ io.Reader) (compose.Result, error) { + return compose.Result{ + Stdout: `{"ready":true,"mode":"oidc","checks":[{"level":"info","code":"auth_ready","message":"Authentication is ready."}]}`, + ExitCode: 0, + }, compose.ErrContainerCleanup + }) + var stdout, stderr bytes.Buffer + + code := RunWithRunner(context.Background(), authInstallation(newAuthDirectory(t)), []string{"check", "--json"}, strings.NewReader(""), &stdout, &stderr, runner) + + if code != 1 || stdout.Len() != 0 || stderr.String() != "tht: authentication diagnostics could not be completed\n" { + t.Fatalf("cleanup failure accepted: code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String()) + } +} + func TestAuthCheckRedactsFailedCoreOutputAndRejectsMalformedReports(t *testing.T) { installation := authInstallation(newAuthDirectory(t)) secretPath := filepath.Join(t.TempDir(), "auth-check-secret") @@ -249,6 +352,34 @@ func (run runnerFunc) Run(ctx context.Context, args []string, input io.Reader) ( return run(ctx, args, input) } +type synchronizedBuffer struct { + mu sync.Mutex + buffer bytes.Buffer +} + +func (b *synchronizedBuffer) Write(value []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buffer.Write(value) +} + +func (b *synchronizedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buffer.String() +} + +func waitForFile(path string, timeout time.Duration) bool { + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + if _, err := os.Stat(path); err == nil { + return true + } + time.Sleep(10 * time.Millisecond) + } + return false +} + func assertAuthOneShotCommand(t *testing.T, installation config.Installation, got []string, interactive bool) { t.Helper() prefix := installation.ComposeArgs("run", "--rm", "--no-deps", "--no-TTY", "--name") diff --git a/tools/tht/internal/compose/process_windows.go b/tools/tht/internal/compose/process_windows.go index b95d6607..2fdf5d56 100644 --- a/tools/tht/internal/compose/process_windows.go +++ b/tools/tht/internal/compose/process_windows.go @@ -3,12 +3,16 @@ package compose import ( + "context" + "io" "os/exec" + "strconv" "syscall" "time" ) const finalTerminationBound = 2 * time.Second +const processTreeTerminationBound = 2 * time.Second const createNewProcessGroup = 0x00000200 func configureProcess(command *exec.Cmd) { @@ -19,11 +23,34 @@ 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 } } + +func terminateWindowsProcessTree(pid int) error { + ctx, cancel := context.WithTimeout(context.Background(), processTreeTerminationBound) + defer cancel() + command := windowsTreeKillCommand(ctx, 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 +} + +func windowsTreeKillCommand(ctx context.Context, pid int) *exec.Cmd { + return exec.CommandContext(ctx, "taskkill.exe", "/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 new file mode 100644 index 00000000..c9428832 --- /dev/null +++ b/tools/tht/internal/compose/process_windows_static_test.go @@ -0,0 +1,28 @@ +package compose + +import ( + "os" + "strings" + "testing" +) + +func TestWindowsTerminationUsesBoundedExactPIDTreeKillWithoutAShell(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", "exec.CommandContext", `"taskkill.exe"`, `"/PID"`, + "strconv.Itoa", `"/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"`} { + 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 new file mode 100644 index 00000000..84b90c6c --- /dev/null +++ b/tools/tht/internal/compose/process_windows_test.go @@ -0,0 +1,17 @@ +//go:build windows + +package compose + +import ( + "context" + "reflect" + "testing" +) + +func TestWindowsTreeKillCommandUsesExactPIDArgumentArray(t *testing.T) { + command := windowsTreeKillCommand(context.Background(), 4242) + want := []string{"taskkill.exe", "/PID", "4242", "/T", "/F"} + if !reflect.DeepEqual(command.Args, want) { + t.Fatalf("taskkill args = %#v, want %#v", command.Args, want) + } +} diff --git a/tools/tht/internal/compose/runner.go b/tools/tht/internal/compose/runner.go index 9113f74c..16dbd64f 100644 --- a/tools/tht/internal/compose/runner.go +++ b/tools/tht/internal/compose/runner.go @@ -55,6 +55,10 @@ type boundedRunner interface { RunBounded(context.Context, []string, io.Reader, CaptureLimits) (Result, error) } +type boundedStreamingRunner interface { + RunBoundedStreaming(context.Context, []string, io.Reader, CaptureLimits, func([]byte)) (Result, error) +} + // execRunner executes the Docker CLI. It never invokes a shell. type execRunner struct { binary string @@ -78,10 +82,16 @@ func (r execRunner) Run(ctx context.Context, args []string, stdin io.Reader) (Re // RunBounded invokes Docker while enforcing both stream limits during capture. func (r execRunner) RunBounded(ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits) (Result, error) { - return r.runBounded(ctx, args, stdin, limits, true) + return r.runBounded(ctx, args, stdin, limits, true, nil) } -func (r execRunner) runBounded(ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits, manageOneShot bool) (Result, error) { +// RunBoundedStreaming retains bounded stderr capture while synchronously observing only the +// retained prefix. It is used for the one validated interactive device-flow prompt. +func (r execRunner) RunBoundedStreaming(ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits, stderrObserver func([]byte)) (Result, error) { + return r.runBounded(ctx, args, stdin, limits, true, stderrObserver) +} + +func (r execRunner) runBounded(ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits, manageOneShot bool, stderrObserver func([]byte)) (Result, error) { if err := validCaptureLimits(limits); err != nil { return Result{}, err } @@ -101,8 +111,8 @@ func (r execRunner) runBounded(ctx context.Context, args []string, stdin io.Read configureProcess(command) command.Stdin = stdin overflow := make(chan struct{}, 1) - stdout := newCappedBuffer(limits.StdoutBytes, overflow) - stderr := newCappedBuffer(limits.StderrBytes, overflow) + stdout := newCappedBuffer(limits.StdoutBytes, overflow, nil) + stderr := newCappedBuffer(limits.StderrBytes, overflow, stderrObserver) command.Stdout = stdout command.Stderr = stderr if err := command.Start(); err != nil { @@ -128,7 +138,7 @@ func (r execRunner) runBounded(ctx context.Context, args []string, stdin io.Read lifecycleErr = errors.Join(lifecycleErr, ErrProcessReap) } } - if interrupted && containerName != "" { + if containerName != "" { if err := r.cleanupOneShotContainer(containerName); err != nil { lifecycleErr = errors.Join(lifecycleErr, ErrContainerCleanup) } @@ -138,20 +148,27 @@ func (r execRunner) runBounded(ctx context.Context, args []string, stdin io.Read result.ExitCode = command.ProcessState.ExitCode() } if stdout.Overflowed() || stderr.Overflowed() { - return result, errors.Join(ErrOutputLimit, lifecycleErr) + return result, joinLifecycleError(ErrOutputLimit, lifecycleErr) } if interrupted { - return result, errors.Join(processErr, lifecycleErr) + return result, joinLifecycleError(processErr, lifecycleErr) } if processErr == nil { - return result, nil + return result, lifecycleErr } var exitError *exec.ExitError if errors.As(processErr, &exitError) { result.ExitCode = exitError.ExitCode() - return result, processErr + return result, joinLifecycleError(processErr, lifecycleErr) } - return result, processErr + return result, joinLifecycleError(processErr, lifecycleErr) +} + +func joinLifecycleError(processErr, lifecycleErr error) error { + if lifecycleErr == nil { + return processErr + } + return errors.Join(processErr, lifecycleErr) } // NewOneShotContainerName returns a Docker-safe, cross-process unique name with a bounded prefix. @@ -253,7 +270,7 @@ func (r execRunner) cleanupOneShotContainer(name string) error { return nil } limits := CaptureLimits{StdoutBytes: cleanupCaptureBytes, StderrBytes: cleanupCaptureBytes} - if _, err := r.runBounded(ctx, []string{"container", "rm", "-f", name}, nil, limits, false); err != nil { + if _, err := r.runBounded(ctx, []string{"container", "rm", "-f", name}, nil, limits, false, nil); err != nil { return ErrContainerCleanup } exists, err = r.oneShotContainerExists(ctx, name) @@ -267,7 +284,7 @@ func (r execRunner) oneShotContainerExists(ctx context.Context, name string) (bo limits := CaptureLimits{StdoutBytes: cleanupCaptureBytes, StderrBytes: cleanupCaptureBytes} result, err := r.runBounded(ctx, []string{ "container", "ls", "--all", "--quiet", "--filter", "name=^/" + name + "$", - }, nil, limits, false) + }, nil, limits, false, nil) if err != nil { return false, ErrContainerCleanup } @@ -294,6 +311,22 @@ func RunBounded(runner Runner, ctx context.Context, args []string, stdin io.Read return result, err } +// RunBoundedStreaming uses the production runner's during-capture observer. Compatibility +// runners are observed only after their already-bounded result returns. +func RunBoundedStreaming(runner Runner, ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits, stderrObserver func([]byte)) (Result, error) { + if err := validCaptureLimits(limits); err != nil { + return Result{}, err + } + if streaming, ok := runner.(boundedStreamingRunner); ok { + return streaming.RunBoundedStreaming(ctx, args, stdin, limits, stderrObserver) + } + result, err := RunBounded(runner, ctx, args, stdin, limits) + if stderrObserver != nil && result.Stderr != "" { + stderrObserver([]byte(result.Stderr)) + } + return result, err +} + func validCaptureLimits(limits CaptureLimits) error { if limits.StdoutBytes < 1 || limits.StderrBytes < 1 || limits.StdoutBytes > maximumCaptureBytes || limits.StderrBytes > maximumCaptureBytes { @@ -323,23 +356,24 @@ type cappedBuffer struct { contents []byte maximum int overflow chan<- struct{} + observer func([]byte) exceeded bool } -func newCappedBuffer(maximum int, overflow chan<- struct{}) *cappedBuffer { +func newCappedBuffer(maximum int, overflow chan<- struct{}, observer func([]byte)) *cappedBuffer { capacity := maximum if capacity > 4096 { capacity = 4096 } - return &cappedBuffer{contents: make([]byte, 0, capacity), maximum: maximum, overflow: overflow} + return &cappedBuffer{contents: make([]byte, 0, capacity), maximum: maximum, overflow: overflow, observer: observer} } func (b *cappedBuffer) Write(value []byte) (int, error) { b.mu.Lock() - defer b.mu.Unlock() remaining := b.maximum - len(b.contents) + kept := 0 if remaining > 0 { - kept := len(value) + kept = len(value) if kept > remaining { kept = remaining } @@ -352,6 +386,14 @@ func (b *cappedBuffer) Write(value []byte) (int, error) { default: } } + var observed []byte + if kept > 0 && b.observer != nil { + observed = append([]byte(nil), value[:kept]...) + } + b.mu.Unlock() + if len(observed) > 0 { + b.observer(observed) + } return len(value), nil } @@ -397,3 +439,11 @@ func (r InstallationRunner) RunBounded(ctx context.Context, args []string, stdin } return RunBounded(r.Runner, ctx, args, stdin, limits) } + +// RunBoundedStreaming preserves the stderr observer across installation argument injection. +func (r InstallationRunner) RunBoundedStreaming(ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits, stderrObserver func([]byte)) (Result, error) { + if len(args) > 0 && args[0] == "compose" { + return RunBoundedStreaming(r.Runner, ctx, r.Installation.ComposeArgs(args[1:]...), stdin, limits, stderrObserver) + } + return RunBoundedStreaming(r.Runner, ctx, args, stdin, limits, stderrObserver) +} diff --git a/tools/tht/internal/compose/runner_test.go b/tools/tht/internal/compose/runner_test.go index 4e358bcc..0e6a8d8a 100644 --- a/tools/tht/internal/compose/runner_test.go +++ b/tools/tht/internal/compose/runner_test.go @@ -116,6 +116,107 @@ func TestRunnerCleansUpNamedComposeContainerAfterCancellation(t *testing.T) { assertContainerCleanup(t, logFile, marker, name) } +func TestRunnerVerifiesAndRemovesNamedComposeContainerAfterNormalExit(t *testing.T) { + logFile := filepath.Join(t.TempDir(), "calls.log") + marker := filepath.Join(t.TempDir(), "container-present") + if err := os.WriteFile(marker, []byte("present"), 0o600); err != nil { + t.Fatal(err) + } + t.Setenv("THT_RUNNER_TEST_LOG", logFile) + t.Setenv("THT_RUNNER_TEST_MARKER", marker) + t.Setenv("THT_RUNNER_TEST_MODE", "normal") + runner := NewRunner(writeExecutable(t, cleanupAwareDockerScript)) + name := "thothii-cleanup-normal-sentinel" + + result, err := RunBounded(runner, context.Background(), []string{ + "compose", "run", "--rm", "--name", name, "core", + }, nil, CaptureLimits{StdoutBytes: 1024, StderrBytes: 1024}) + + if err != nil || result.ExitCode != 0 { + t.Fatalf("RunBounded() result=%#v error=%v, want clean exit", result, err) + } + assertContainerCleanup(t, logFile, marker, name) +} + +func TestRunnerVerifiesAndRemovesNamedComposeContainerAfterExpectedExitOne(t *testing.T) { + logFile := filepath.Join(t.TempDir(), "calls.log") + marker := filepath.Join(t.TempDir(), "container-present") + if err := os.WriteFile(marker, []byte("present"), 0o600); err != nil { + t.Fatal(err) + } + t.Setenv("THT_RUNNER_TEST_LOG", logFile) + t.Setenv("THT_RUNNER_TEST_MARKER", marker) + t.Setenv("THT_RUNNER_TEST_MODE", "exit-one") + runner := NewRunner(writeExecutable(t, cleanupAwareDockerScript)) + name := "thothii-cleanup-exit-one-sentinel" + + result, err := RunBounded(runner, context.Background(), []string{ + "compose", "run", "--rm", "--name", name, "core", + }, nil, CaptureLimits{StdoutBytes: 1024, StderrBytes: 1024}) + + var exitError *exec.ExitError + if !errors.As(err, &exitError) || errors.Is(err, ErrContainerCleanup) || result.ExitCode != 1 { + t.Fatalf("RunBounded() result=%#v error=%v, want original exit 1", result, err) + } + assertContainerCleanup(t, logFile, marker, name) +} + +func TestRunnerAcceptsANamedOneShotAlreadyRemovedByRm(t *testing.T) { + logFile := filepath.Join(t.TempDir(), "calls.log") + marker := filepath.Join(t.TempDir(), "container-absent") + t.Setenv("THT_RUNNER_TEST_LOG", logFile) + t.Setenv("THT_RUNNER_TEST_MARKER", marker) + t.Setenv("THT_RUNNER_TEST_MODE", "normal") + runner := NewRunner(writeExecutable(t, cleanupAwareDockerScript)) + name := "thothii-cleanup-already-removed-sentinel" + + result, err := RunBounded(runner, context.Background(), []string{ + "compose", "run", "--rm", "--name", name, "core", + }, nil, CaptureLimits{StdoutBytes: 1024, StderrBytes: 1024}) + + if err != nil || result.ExitCode != 0 { + t.Fatalf("RunBounded() result=%#v error=%v, want clean exit", result, err) + } + calls, readErr := os.ReadFile(logFile) + if readErr != nil { + t.Fatal(readErr) + } + if !strings.Contains(string(calls), "container ls --all --quiet --filter name=^/"+name+"$") || + strings.Contains(string(calls), "container rm -f "+name) { + t.Fatalf("Docker calls = %q, want absence verification without forced removal", calls) + } + if _, statErr := os.Stat(marker); !os.IsNotExist(statErr) { + t.Fatalf("named one-shot unexpectedly left residue: %v", statErr) + } +} + +func TestRunnerJoinsCleanupFailureAfterNormalExit(t *testing.T) { + logFile := filepath.Join(t.TempDir(), "calls.log") + marker := filepath.Join(t.TempDir(), "container-present") + if err := os.WriteFile(marker, []byte("present"), 0o600); err != nil { + t.Fatal(err) + } + t.Setenv("THT_RUNNER_TEST_LOG", logFile) + t.Setenv("THT_RUNNER_TEST_MARKER", marker) + t.Setenv("THT_RUNNER_TEST_MODE", "normal-cleanup-fails") + runner := NewRunner(writeExecutable(t, cleanupAwareDockerScript)) + name := "thothii-cleanup-failure-sentinel" + + result, err := RunBounded(runner, context.Background(), []string{ + "compose", "run", "--rm", "--name", name, "core", + }, nil, CaptureLimits{StdoutBytes: 1024, StderrBytes: 1024}) + + if result.ExitCode != 0 || !errors.Is(err, ErrContainerCleanup) { + t.Fatalf("RunBounded() result=%#v error=%v, want cleanup failure joined to exit 0", result, err) + } + if strings.Contains(err.Error(), name) { + t.Fatalf("RunBounded() exposed the container name: %v", err) + } + if _, statErr := os.Stat(marker); statErr != nil { + t.Fatalf("cleanup failure test unexpectedly removed marker: %v", statErr) + } +} + func TestRunnerSurfacesCleanupFailureWithoutContainerName(t *testing.T) { logFile := filepath.Join(t.TempDir(), "calls.log") marker := filepath.Join(t.TempDir(), "container-present") @@ -270,13 +371,19 @@ if [ "$1" = "container" ] && [ "$2" = "ls" ]; then exit 0 fi if [ "$1" = "container" ] && [ "$2" = "rm" ]; then - if [ "$THT_RUNNER_TEST_MODE" = "cleanup-fails" ]; then exit 9; fi + case "$THT_RUNNER_TEST_MODE" in *cleanup-fails) exit 9 ;; esac rm -f "$THT_RUNNER_TEST_MARKER" exit 0 fi if [ "$THT_RUNNER_TEST_MODE" = "flood" ]; then while :; do printf '0123456789abcdef'; printf 'fedcba9876543210' >&2; done fi +if [ "$THT_RUNNER_TEST_MODE" = "normal" ] || [ "$THT_RUNNER_TEST_MODE" = "normal-cleanup-fails" ]; then + exit 0 +fi +if [ "$THT_RUNNER_TEST_MODE" = "exit-one" ]; then + exit 1 +fi trap '' TERM INT while :; do sleep 1; done `