From 438261992676c95847136196b57cbee0d3c5a8c5 Mon Sep 17 00:00:00 2001 From: John-Michael Mulesa Date: Mon, 28 Sep 2026 15:06:09 +0000 Subject: [PATCH 1/4] Add hang and stall guards to installs and downloads Bad installers can hang googet indefinitely, most often by becoming unexpectedly interactive in an unattended context, and downloads can hang on a stalled connection. This change kills only true hangs and stalls, never work that is still progressing, and makes a killed install fail cleanly. Installer supervision (new supervisor package, goolib, system): - Installers run in a Windows Job Object (process group on Unix) with stdin set to the null device. The child is started suspended, assigned to the job and then resumed, so there is no startup race. - Forward progress is a rolling-window check across the whole tree: CPU >= 250ms, read+write I/O >= 64KiB, or installer/MSI log growth within 30s. - Watchdogs: inactivity (no progress for 5m); hard cap (60m, lifted for wusa/dism unless the admin or package set one); modal UI in unattended mode (a #32770 or owned modal-frame window persisting 30s with no progress since it appeared). - Work outside the job counts as progress: the Windows Installer service tree while _MSIExecute is held, the TrustedInstaller/TiWorker tree for wusa/dism, and CBS.log growth. After an msiexec abort, only service-side descendants created after start and observed idle for the inactivity timeout are terminated; the services themselves never are. - No abort decisions are made after the root installer exits. On success the job's kill-on-close limit is cleared so tray apps and updaters survive. Waits are bounded (WaitDelay 30s, post-kill wait 60s). - If the job cannot be created, only the hard cap is enforced on the root process, so a running child is never orphaned. Configuration: - googet.conf: SupervisorMode (enforce|monitor|off, default enforce), InactivityTimeout, InstallTimeout, UIGracePeriod, UIDetection, DownloadStallTimeout. "0" disables the inactivity and install timeouts. - goospec install/uninstall/verify accept timeout and inactivity_timeout overrides on every OS. - monitor mode logs one WOULD_KILL line per reason and never kills, for staged rollout. goopack build commands run with supervision off. Failure semantics (install, cli/install, cli/update): - Installs are transactional: overwritten files are backed up beside themselves (copy fallback when a rename is impossible), new files and directories are tracked, removed empty directories are recorded, and everything is rolled back on failure. The package is never recorded in the database on failure, and installer logs are preserved. - Multi-package install/update continues past a failed package and exits 1. Downloads (client, download): - ResponseHeaderTimeout is 30s; there is no overall timeout, so slow links are never capped. Downloader now uses its own http.Client instead of mutating http.DefaultClient. - An idle-read StallReader (default 120s, Downloader.StallTimeout) guards package bodies (HTTP and GCS) and repo index bodies. - Retryable failures (stall, reset, EOF, timeouts, http2, 5xx, 429) resume with Range and back off with jitter. Attempts that advance the file are unlimited; the download fails after 4 consecutive attempts with no progress or 20 attempts that each advanced it by less than 1 MiB. - Content-Range is validated. A 416, a mismatched range, or a checksum mismatch after a resume restarts from byte 0 and resets the budget; repeated restarts that never pass the previous best offset fail. - Fixed a GCS index error-shadowing bug that hid non-404 errors. Notes: - Defaults are enforce with a 60m cap; installers that legitimately run longer must set timeout in their goospec. Fleets can stage with SupervisorMode: monitor. - Killing googet now also kills the in-flight installer tree. - A bootstrapper that exits while a child holding its stdout keeps running is considered finished after 30s instead of being waited on forever. - The Windows-specific paths have been cross-compiled and vetted but not yet executed; supervisor_windows_test.go runs in the Windows CI job. --- cli/install/install.go | 5 +- cli/install/install_test.go | 703 ++++++++- cli/update/update.go | 7 +- cli/update/update_test.go | 760 ++++++++++ client/client.go | 79 +- client/client_test.go | 228 ++- client/stall.go | 176 +++ client/stall_test.go | 482 ++++++ download/download.go | 618 ++++++-- download/download_test.go | 1997 ++++++++++++++++++++++++- googet.go | 20 + googet.goospec | 7 +- goolib/goolib.go | 122 +- goolib/goolib_test.go | 333 ++++- goolib/goospec.go | 46 + goolib/goospec_test.go | 12 + goolib/supervise_test.go | 117 ++ goopack/goopack.go | 9 +- install/install.go | 93 +- install/install_test.go | 1414 ++++++++++++++++- install/txn.go | 378 +++++ settings/settings.go | 74 +- settings/settings_test.go | 78 + supervisor/msi.go | 113 ++ supervisor/msi_test.go | 189 +++ supervisor/msi_windows.go | 298 ++++ supervisor/proctree.go | 135 ++ supervisor/progress.go | 493 ++++++ supervisor/supervise_test.go | 310 ++++ supervisor/supervisor.go | 344 +++++ supervisor/supervisor_test.go | 943 ++++++++++++ supervisor/supervisor_unix.go | 252 ++++ supervisor/supervisor_unix_test.go | 346 +++++ supervisor/supervisor_windows.go | 594 ++++++++ supervisor/supervisor_windows_test.go | 267 ++++ supervisor/ui_test.go | 279 ++++ system/supervisor_options_test.go | 87 ++ system/system.go | 20 +- system/system_darwin.go | 6 +- system/system_linux.go | 6 +- system/system_windows.go | 32 +- 41 files changed, 12236 insertions(+), 236 deletions(-) create mode 100644 client/stall.go create mode 100644 client/stall_test.go create mode 100644 goolib/supervise_test.go create mode 100644 install/txn.go create mode 100644 supervisor/msi.go create mode 100644 supervisor/msi_test.go create mode 100644 supervisor/msi_windows.go create mode 100644 supervisor/proctree.go create mode 100644 supervisor/progress.go create mode 100644 supervisor/supervise_test.go create mode 100644 supervisor/supervisor.go create mode 100644 supervisor/supervisor_test.go create mode 100644 supervisor/supervisor_unix.go create mode 100644 supervisor/supervisor_unix_test.go create mode 100644 supervisor/supervisor_windows.go create mode 100644 supervisor/supervisor_windows_test.go create mode 100644 supervisor/ui_test.go create mode 100644 system/supervisor_options_test.go diff --git a/cli/install/install.go b/cli/install/install.go index e2e3801..f2a36d2 100644 --- a/cli/install/install.go +++ b/cli/install/install.go @@ -11,10 +11,9 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package install provides the install subcommand for downloading and installing packages. package install -// The install subcommand handles the downloading and installation of a package. - import ( "bytes" "context" @@ -98,7 +97,7 @@ func (cmd *installCmd) Execute(ctx context.Context, flags *flag.FlagSet, _ ...an // We only need to build sources and download indexes if there are any // non-file goo arguments passed to the install command (usually the case). - if !allFileGoos(flag.Args()) { + if !allFileGoos(flags.Args()) { repos, err := repo.BuildSources(cmd.sources) if err != nil { logger.Errorf("Failed to initialize repos: %v", err) diff --git a/cli/install/install_test.go b/cli/install/install_test.go index 2bd652d..10417f8 100644 --- a/cli/install/install_test.go +++ b/cli/install/install_test.go @@ -2,7 +2,9 @@ package install import ( "bytes" + "compress/gzip" "context" + "encoding/json" "flag" "io" "maps" @@ -22,6 +24,7 @@ import ( "github.com/google/googet/v2/settings" "github.com/google/googet/v2/testutil" "github.com/google/logger" + "github.com/google/subcommands" ) // checkInstalled returns true if the test package identified by ps was @@ -346,7 +349,7 @@ func TestInstallDryRun(t *testing.T) { } } - // Verify DB state hasn't changed + // Verify DB state hasn't changed. finalState, err := db.FetchPkgs("") if err != nil { t.Errorf("db.FetchPkgs: %v", err) @@ -358,3 +361,701 @@ func TestInstallDryRun(t *testing.T) { }) } } + +func TestBatchInstallContinuation(t *testing.T) { + // Verify that running googet install on multiple packages where one + // package fails or aborts still installs the remaining packages and + // latches exit code 1 (ExitFailure). + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // Create valid package B. + pkgB := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"} + rsB := testutil.GenGoo(t, gooDir, logDir, pkgB) + + // Write repo index and index.gz to gooDir so AvailableVersions succeeds over HTTP. + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsB}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + // "pkgA" does not exist in the repo (will fail version resolution). + // "pkgB" exists in the repo (must succeed). + args := []string{"-sources=" + srv.URL, "pkgA", "pkgB"} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB was installed and recorded in googet.db. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil { + t.Errorf("pkgB was not recorded in googet.db; expected successful installation") + } + if !checkInstalled(t, logDir, pkgB) { + t.Errorf("pkgB file was not installed to target directory") + } + + // Verify pkgA was NOT recorded in googet.db. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec != nil { + t.Errorf("pkgA was recorded in googet.db; expected failure") + } +} + +func TestBatchInstallContinuation_InstallerExecutionFailure(t *testing.T) { + // Verifies batch continuation when package A downloads but fails installer execution. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // Create package A with a failing installer execution path. + pkgA := goolib.PkgSpec{ + Name: "pkgA", + Arch: "noarch", + Version: "1.0.0", + Install: goolib.ExecFile{Path: "failing_installer_binary.exe"}, + } + rsA := testutil.GenGoo(t, gooDir, logDir, pkgA) + + // Create package B which is completely valid. + pkgB := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"} + rsB := testutil.GenGoo(t, gooDir, logDir, pkgB) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsA, rsB}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + args := []string{"-sources=" + srv.URL, "pkgA", "pkgB"} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB was installed and recorded in googet.db. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil { + t.Errorf("pkgB was not recorded in googet.db; expected successful installation") + } + if !checkInstalled(t, logDir, pkgB) { + t.Errorf("pkgB file was not installed to target directory") + } + + // Verify pkgA was NOT recorded in googet.db. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec != nil { + t.Errorf("pkgA was recorded in googet.db; expected failure") + } +} + +func TestBatchInstallContinuation_OrderReversal(t *testing.T) { + // Verifies batch continuation when successful package precedes failing package. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + pkgA := goolib.PkgSpec{ + Name: "pkgA", + Arch: "noarch", + Version: "1.0.0", + Install: goolib.ExecFile{Path: "failing_installer_binary.exe"}, + } + rsA := testutil.GenGoo(t, gooDir, logDir, pkgA) + + pkgB := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"} + rsB := testutil.GenGoo(t, gooDir, logDir, pkgB) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsA, rsB}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + // pkgB (success) comes first, pkgA (fail) comes second. + args := []string{"-sources=" + srv.URL, "pkgB", "pkgA"} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB was installed and recorded in googet.db. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil { + t.Errorf("pkgB was not recorded in googet.db; expected successful installation") + } + if !checkInstalled(t, logDir, pkgB) { + t.Errorf("pkgB file was not installed to target directory") + } + + // Verify pkgA was NOT recorded in googet.db. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec != nil { + t.Errorf("pkgA was recorded in googet.db; expected failure") + } +} + +func TestBatchInstallContinuation_LocalGooFiles(t *testing.T) { + // Verifies batch continuation when installing local .goo files where one fails. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + pkgDir, logDir := t.TempDir(), t.TempDir() + + pkgA := goolib.PkgSpec{ + Name: "pkgA", + Arch: "noarch", + Version: "1.0.0", + Install: goolib.ExecFile{Path: "failing_installer_binary.exe"}, + } + testutil.GenGoo(t, pkgDir, logDir, pkgA) + fileA := filepath.Join(pkgDir, pkgA.String()+".goo") + + pkgB := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"} + testutil.GenGoo(t, pkgDir, logDir, pkgB) + fileB := filepath.Join(pkgDir, pkgB.String()+".goo") + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + args := []string{fileA, fileB} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB was installed and recorded in googet.db. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil { + t.Errorf("pkgB was not recorded in googet.db; expected successful installation") + } + if !checkInstalled(t, logDir, pkgB) { + t.Errorf("pkgB file was not installed to target directory") + } + + // Verify pkgA was NOT recorded in googet.db. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec != nil { + t.Errorf("pkgA was recorded in googet.db; expected failure") + } +} + +func TestBatchInstallContinuation_MixedFileAndRepo(t *testing.T) { + // Verifies batch continuation when mixing local .goo files and repo packages. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + localDir, gooDir, logDir := t.TempDir(), t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // Local file pkgA that fails installer execution. + pkgA := goolib.PkgSpec{ + Name: "pkgA", + Arch: "noarch", + Version: "1.0.0", + Install: goolib.ExecFile{Path: "failing_installer_binary.exe"}, + } + testutil.GenGoo(t, localDir, logDir, pkgA) + fileA := filepath.Join(localDir, pkgA.String()+".goo") + + // Repo package pkgB that succeeds. + pkgB := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"} + rsB := testutil.GenGoo(t, gooDir, logDir, pkgB) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsB}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + args := []string{"-sources=" + srv.URL, fileA, "pkgB"} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil { + t.Errorf("pkgB was not recorded in googet.db; expected successful installation") + } + + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec != nil { + t.Errorf("pkgA was recorded in googet.db; expected failure") + } +} + +func TestBatchInstall_AllSucceed(t *testing.T) { + // Verifies batch install returns ExitSuccess when all packages succeed. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + pkgB := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"} + rsB := testutil.GenGoo(t, gooDir, logDir, pkgB) + + pkgC := goolib.PkgSpec{Name: "pkgC", Arch: "noarch", Version: "1.0.0"} + rsC := testutil.GenGoo(t, gooDir, logDir, pkgC) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsB, rsC}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + args := []string{"-sources=" + srv.URL, "pkgB", "pkgC"} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitSuccess { + t.Errorf("cmd.Execute got %v, want subcommands.ExitSuccess (%v)", exitStatus, subcommands.ExitSuccess) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + for _, name := range []string{"pkgB", "pkgC"} { + ps, err := db.FetchPkg(goolib.PackageInfo{Name: name, Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(%s): %v", name, err) + } + if ps.PackageSpec == nil { + t.Errorf("package %s was not recorded in googet.db", name) + } + } +} + +func TestBatchInstall_AllFail(t *testing.T) { + // Verifies batch install returns ExitFailure when all packages fail. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + gooDir := t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + args := []string{"-sources=" + srv.URL, "nonexistentA", "nonexistentB"} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + pkgs, err := db.FetchPkgs("") + if err != nil { + t.Fatalf("db.FetchPkgs: %v", err) + } + if len(pkgs) != 0 { + t.Errorf("expected 0 packages in db, got %d", len(pkgs)) + } +} + +func TestBatchInstallContinuation_ThreePackages_MiddleSucceeds(t *testing.T) { + // Verifies batch continuation across three packages where first and third fail and middle succeeds. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + pkgA := goolib.PkgSpec{ + Name: "pkgA", + Arch: "noarch", + Version: "1.0.0", + Install: goolib.ExecFile{Path: "failing_installer_a.exe"}, + } + rsA := testutil.GenGoo(t, gooDir, logDir, pkgA) + + pkgB := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"} + rsB := testutil.GenGoo(t, gooDir, logDir, pkgB) + + pkgC := goolib.PkgSpec{ + Name: "pkgC", + Arch: "noarch", + Version: "1.0.0", + Install: goolib.ExecFile{Path: "failing_installer_c.exe"}, + } + rsC := testutil.GenGoo(t, gooDir, logDir, pkgC) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsA, rsB, rsC}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + args := []string{"-sources=" + srv.URL, "pkgA", "pkgB", "pkgC"} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB is in DB and installed. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil { + t.Errorf("pkgB was not recorded in googet.db; expected successful installation") + } + if !checkInstalled(t, logDir, pkgB) { + t.Errorf("pkgB file was not installed to target directory") + } + + // Verify pkgA is not in DB and placed file rolled back. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec != nil { + t.Errorf("pkgA was recorded in googet.db; expected failure") + } + if _, err := os.Stat(filepath.Join(logDir, pkgA.Name)); !os.IsNotExist(err) { + t.Errorf("pkgA placed file still exists on disk; expected rollback deletion") + } + + // Verify pkgC is not in DB and placed file rolled back. + psC, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgC", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgC): %v", err) + } + if psC.PackageSpec != nil { + t.Errorf("pkgC was recorded in googet.db; expected failure") + } + if _, err := os.Stat(filepath.Join(logDir, pkgC.Name)); !os.IsNotExist(err) { + t.Errorf("pkgC placed file still exists on disk; expected rollback deletion") + } +} + +func TestBatchInstallContinuation_RollbackUnlinksPlacedFile(t *testing.T) { + // Verifies that when package installation fails, newly placed files are unlinked from disk. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + pkgA := goolib.PkgSpec{ + Name: "pkg_rollback_test", + Arch: "noarch", + Version: "1.0.0", + Install: goolib.ExecFile{Path: "failing_binary.exe"}, + } + rsA := testutil.GenGoo(t, gooDir, logDir, pkgA) + + pkgB := goolib.PkgSpec{Name: "pkg_ok_test", Arch: "noarch", Version: "1.0.0"} + rsB := testutil.GenGoo(t, gooDir, logDir, pkgB) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsA, rsB}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &installCmd{} + fs := flag.NewFlagSet("install", flag.ContinueOnError) + cmd.SetFlags(fs) + + args := []string{"-sources=" + srv.URL, "pkg_rollback_test", "pkg_ok_test"} + if err := fs.Parse(args); err != nil { + t.Fatalf("fs.Parse: %v", err) + } + + exitStatus := cmd.Execute(ctx, fs) + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + // Verify pkg_rollback_test file was deleted. + placedFile := filepath.Join(logDir, pkgA.Name) + if _, err := os.Stat(placedFile); !os.IsNotExist(err) { + t.Errorf("Placed file %s still exists; expected rollback removal", placedFile) + } + + // Verify pkg_ok_test file exists. + okFile := filepath.Join(logDir, pkgB.Name) + if _, err := os.Stat(okFile); err != nil { + t.Errorf("Placed file %s does not exist: %v", okFile, err) + } +} diff --git a/cli/update/update.go b/cli/update/update.go index 17dcb48..534249c 100644 --- a/cli/update/update.go +++ b/cli/update/update.go @@ -11,10 +11,9 @@ See the License for the specific language governing permissions and limitations under the License. */ +// Package update provides the update subcommand for bulk updating packages. package update -// The update subcommand handles bulk updating of packages. - import ( "context" "flag" @@ -56,7 +55,7 @@ func (cmd *updateCmd) SetFlags(f *flag.FlagSet) { f.BoolVar(&cmd.force, "force", false, "force overwrite of conflicting files (only required if StrictConflicts is enabled in config)") } -func (cmd *updateCmd) Execute(ctx context.Context, _ *flag.FlagSet, _ ...interface{}) subcommands.ExitStatus { +func (cmd *updateCmd) Execute(ctx context.Context, _ *flag.FlagSet, _ ...any) subcommands.ExitStatus { db, err := googetdb.NewDB(settings.DBFile()) if err != nil { logger.Errorf("Failed to open database: %v", err) @@ -115,6 +114,8 @@ func (cmd *updateCmd) Execute(ctx context.Context, _ *flag.FlagSet, _ ...interfa r, err := client.WhatRepo(pi, rm) if err != nil { logger.Errorf("Error finding repo: %v.", err) + exitCode = subcommands.ExitFailure + continue } if err := install.FromRepo(ctx, pi, r, cache, rm, settings.Archs, cmd.dbOnly, cmd.force, downloader, db); err != nil { logger.Errorf("Error updating %s %s %s: %v", pi.Arch, pi.Name, pi.Ver, err) diff --git a/cli/update/update_test.go b/cli/update/update_test.go index d5e654e..32a6cba 100644 --- a/cli/update/update_test.go +++ b/cli/update/update_test.go @@ -2,15 +2,23 @@ package update import ( "bytes" + "compress/gzip" + "context" + "encoding/json" "io" "os" + "path/filepath" "testing" "github.com/google/go-cmp/cmp" "github.com/google/googet/v2/client" + "github.com/google/googet/v2/googetdb" "github.com/google/googet/v2/goolib" "github.com/google/googet/v2/priority" "github.com/google/googet/v2/settings" + "github.com/google/googet/v2/testutil" + "github.com/google/logger" + "github.com/google/subcommands" ) func captureStdout(f func()) string { @@ -127,3 +135,755 @@ func TestUpdates(t *testing.T) { }) } } + +func TestBatchUpdateContinuation_WhatRepoError(t *testing.T) { + // Verifies that when client.WhatRepo fails for one package, exitCode = ExitFailure + // is latched, execution continues, and remaining packages update successfully. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + + // Initial DB state: pkgA and pkgB at version 1.0.0. + initialState := client.GooGetState{ + {PackageSpec: &goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"}}, + } + if err := db.WriteStateToDB(initialState); err != nil { + t.Fatalf("db.WriteStateToDB: %v", err) + } + db.Close() + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // Generate update only for pkgB. pkgA update is absent from the repo index. + pkgB2 := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "2.0.0"} + rsB2 := testutil.GenGoo(t, gooDir, logDir, pkgB2) + + // In the repo index, provide a virtual provider for pkgA that yields an update + // in FindRepoLatest but whose real package name differs, inducing WhatRepo error. + pkgProvider := goolib.PkgSpec{ + Name: "real_provider", + Arch: "noarch", + Version: "2.0.0", + Provides: []string{"pkgA"}, + } + rsProvider := testutil.GenGoo(t, gooDir, logDir, pkgProvider) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsB2, rsProvider}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &updateCmd{sources: srv.URL} + exitStatus := cmd.Execute(ctx, nil) + + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err = googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB was updated to 2.0.0. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil || psB.PackageSpec.Version != "2.0.0" { + t.Errorf("pkgB version got %v, want 2.0.0", psB.PackageSpec) + } + + // Verify pkgA remained at 1.0.0. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec == nil || psA.PackageSpec.Version != "1.0.0" { + t.Errorf("pkgA version got %v, want 1.0.0", psA.PackageSpec) + } +} + +func TestBatchUpdateContinuation_InstallFailure(t *testing.T) { + // Verifies that when install.FromRepo fails for one package, exitCode = ExitFailure + // is latched and remaining packages continue to update. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + + initialState := client.GooGetState{ + {PackageSpec: &goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"}}, + } + if err := db.WriteStateToDB(initialState); err != nil { + t.Fatalf("db.WriteStateToDB: %v", err) + } + db.Close() + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // Generate pkgA and pkgB updates. + pkgA2 := goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "2.0.0"} + rsA2 := testutil.GenGoo(t, gooDir, logDir, pkgA2) + + pkgB2 := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "2.0.0"} + rsB2 := testutil.GenGoo(t, gooDir, logDir, pkgB2) + + // Corrupt pkgA's file on disk to trigger a checksum/unpack failure in install.FromRepo. + pkgAPath := filepath.Join(gooDir, rsA2.Source) + if err := os.WriteFile(pkgAPath, []byte("corrupted file content"), 0644); err != nil { + t.Fatalf("corrupting pkgA: %v", err) + } + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsA2, rsB2}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &updateCmd{sources: srv.URL} + exitStatus := cmd.Execute(ctx, nil) + + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err = googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB updated successfully. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil || psB.PackageSpec.Version != "2.0.0" { + t.Errorf("pkgB version got %v, want 2.0.0", psB.PackageSpec) + } + + // Verify pkgA was NOT updated (remains 1.0.0). + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec == nil || psA.PackageSpec.Version != "1.0.0" { + t.Errorf("pkgA version got %v, want 1.0.0", psA.PackageSpec) + } +} + +func TestBatchUpdateContinuation_OrderReversal(t *testing.T) { + // Verifies batch update continuation when successful package precedes failing package. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + + // Initial DB state with pkgB first, then pkgA. + initialState := client.GooGetState{ + {PackageSpec: &goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "1.0.0"}}, + } + if err := db.WriteStateToDB(initialState); err != nil { + t.Fatalf("db.WriteStateToDB: %v", err) + } + db.Close() + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // pkgB update is valid. + pkgB2 := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "2.0.0"} + rsB2 := testutil.GenGoo(t, gooDir, logDir, pkgB2) + + // pkgA update has failing installer execution. + pkgA2 := goolib.PkgSpec{ + Name: "pkgA", + Arch: "noarch", + Version: "2.0.0", + Install: goolib.ExecFile{Path: "failing_installer_binary.exe"}, + } + rsA2 := testutil.GenGoo(t, gooDir, logDir, pkgA2) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsB2, rsA2}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &updateCmd{sources: srv.URL} + exitStatus := cmd.Execute(ctx, nil) + + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err = googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB was updated to 2.0.0. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil || psB.PackageSpec.Version != "2.0.0" { + t.Errorf("pkgB version got %v, want 2.0.0", psB.PackageSpec) + } + + // Verify pkgA remained at 1.0.0. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec == nil || psA.PackageSpec.Version != "1.0.0" { + t.Errorf("pkgA version got %v, want 1.0.0", psA.PackageSpec) + } +} + +func TestBatchUpdateContinuation_ThreePackages_WhatRepoAndInstallFailure(t *testing.T) { + // Verifies continuation across 3 packages where pkgA fails WhatRepo, pkgB fails install, and pkgC succeeds. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + + initialState := client.GooGetState{ + {PackageSpec: &goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgC", Arch: "noarch", Version: "1.0.0"}}, + } + if err := db.WriteStateToDB(initialState); err != nil { + t.Fatalf("db.WriteStateToDB: %v", err) + } + db.Close() + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // pkgA: virtual provider inducing WhatRepo error. + pkgProvider := goolib.PkgSpec{ + Name: "provider_for_a", + Arch: "noarch", + Version: "2.0.0", + Provides: []string{"pkgA"}, + } + rsProvider := testutil.GenGoo(t, gooDir, logDir, pkgProvider) + + // pkgB: corrupted payload inducing install error. + pkgB2 := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "2.0.0"} + rsB2 := testutil.GenGoo(t, gooDir, logDir, pkgB2) + pkgBPath := filepath.Join(gooDir, rsB2.Source) + if err := os.WriteFile(pkgBPath, []byte("corrupted pkgB content"), 0644); err != nil { + t.Fatalf("corrupting pkgB: %v", err) + } + + // pkgC: valid update. + pkgC2 := goolib.PkgSpec{Name: "pkgC", Arch: "noarch", Version: "2.0.0"} + rsC2 := testutil.GenGoo(t, gooDir, logDir, pkgC2) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsProvider, rsB2, rsC2}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &updateCmd{sources: srv.URL} + exitStatus := cmd.Execute(ctx, nil) + + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err = googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgC was updated to 2.0.0. + psC, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgC", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgC): %v", err) + } + if psC.PackageSpec == nil || psC.PackageSpec.Version != "2.0.0" { + t.Errorf("pkgC version got %v, want 2.0.0", psC.PackageSpec) + } + + // Verify pkgA remained at 1.0.0. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec == nil || psA.PackageSpec.Version != "1.0.0" { + t.Errorf("pkgA version got %v, want 1.0.0", psA.PackageSpec) + } + + // Verify pkgB remained at 1.0.0. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil || psB.PackageSpec.Version != "1.0.0" { + t.Errorf("pkgB version got %v, want 1.0.0", psB.PackageSpec) + } +} + +func TestBatchUpdate_AllSucceed(t *testing.T) { + // Verifies batch update returns ExitSuccess when all packages update successfully. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + + initialState := client.GooGetState{ + {PackageSpec: &goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"}}, + } + if err := db.WriteStateToDB(initialState); err != nil { + t.Fatalf("db.WriteStateToDB: %v", err) + } + db.Close() + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + pkgA2 := goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "2.0.0"} + rsA2 := testutil.GenGoo(t, gooDir, logDir, pkgA2) + + pkgB2 := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "2.0.0"} + rsB2 := testutil.GenGoo(t, gooDir, logDir, pkgB2) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsA2, rsB2}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &updateCmd{sources: srv.URL} + exitStatus := cmd.Execute(ctx, nil) + + if exitStatus != subcommands.ExitSuccess { + t.Errorf("cmd.Execute got %v, want subcommands.ExitSuccess (%v)", exitStatus, subcommands.ExitSuccess) + } + + db, err = googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + for _, name := range []string{"pkgA", "pkgB"} { + ps, err := db.FetchPkg(goolib.PackageInfo{Name: name, Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(%s): %v", name, err) + } + if ps.PackageSpec == nil || ps.PackageSpec.Version != "2.0.0" { + t.Errorf("%s version got %v, want 2.0.0", name, ps.PackageSpec) + } + } +} + +func TestBatchUpdate_AllFail(t *testing.T) { + // Verifies batch update returns ExitFailure when all package updates fail. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + + initialState := client.GooGetState{ + {PackageSpec: &goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"}}, + } + if err := db.WriteStateToDB(initialState); err != nil { + t.Fatalf("db.WriteStateToDB: %v", err) + } + db.Close() + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // pkgA: WhatRepo error. + pkgProvider := goolib.PkgSpec{ + Name: "provider_pkg", + Arch: "noarch", + Version: "2.0.0", + Provides: []string{"pkgA"}, + } + rsProvider := testutil.GenGoo(t, gooDir, logDir, pkgProvider) + + // pkgB: corrupted payload. + pkgB2 := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "2.0.0"} + rsB2 := testutil.GenGoo(t, gooDir, logDir, pkgB2) + pkgBPath := filepath.Join(gooDir, rsB2.Source) + if err := os.WriteFile(pkgBPath, []byte("corrupted payload"), 0644); err != nil { + t.Fatalf("corrupting pkgB: %v", err) + } + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsProvider, rsB2}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &updateCmd{sources: srv.URL} + exitStatus := cmd.Execute(ctx, nil) + + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err = googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + for _, name := range []string{"pkgA", "pkgB"} { + ps, err := db.FetchPkg(goolib.PackageInfo{Name: name, Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(%s): %v", name, err) + } + if ps.PackageSpec == nil || ps.PackageSpec.Version != "1.0.0" { + t.Errorf("%s version got %v, want 1.0.0 (must not be modified)", name, ps.PackageSpec) + } + } +} + +func TestBatchUpdateContinuation_ThreePackages_MiddleSucceeds(t *testing.T) { + // Verifies batch update continuation across three packages where first and third fail and middle succeeds. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + + initialState := client.GooGetState{ + {PackageSpec: &goolib.PkgSpec{Name: "pkgA", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkgC", Arch: "noarch", Version: "1.0.0"}}, + } + if err := db.WriteStateToDB(initialState); err != nil { + t.Fatalf("db.WriteStateToDB: %v", err) + } + db.Close() + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // pkgA: installer failure on update. + pkgA2 := goolib.PkgSpec{ + Name: "pkgA", + Arch: "noarch", + Version: "2.0.0", + Install: goolib.ExecFile{Path: "failing_installer_a.exe"}, + } + rsA2 := testutil.GenGoo(t, gooDir, logDir, pkgA2) + + // pkgB: valid update. + pkgB2 := goolib.PkgSpec{Name: "pkgB", Arch: "noarch", Version: "2.0.0"} + rsB2 := testutil.GenGoo(t, gooDir, logDir, pkgB2) + + // pkgC: installer failure on update. + pkgC2 := goolib.PkgSpec{ + Name: "pkgC", + Arch: "noarch", + Version: "2.0.0", + Install: goolib.ExecFile{Path: "failing_installer_c.exe"}, + } + rsC2 := testutil.GenGoo(t, gooDir, logDir, pkgC2) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsA2, rsB2, rsC2}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &updateCmd{sources: srv.URL} + exitStatus := cmd.Execute(ctx, nil) + + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err = googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Verify pkgB updated to 2.0.0. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgB", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgB): %v", err) + } + if psB.PackageSpec == nil || psB.PackageSpec.Version != "2.0.0" { + t.Errorf("pkgB version got %v, want 2.0.0", psB.PackageSpec) + } + + // Verify pkgA remained at 1.0.0. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgA", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgA): %v", err) + } + if psA.PackageSpec == nil || psA.PackageSpec.Version != "1.0.0" { + t.Errorf("pkgA version got %v, want 1.0.0", psA.PackageSpec) + } + + // Verify pkgC remained at 1.0.0. + psC, err := db.FetchPkg(goolib.PackageInfo{Name: "pkgC", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkgC): %v", err) + } + if psC.PackageSpec == nil || psC.PackageSpec.Version != "1.0.0" { + t.Errorf("pkgC version got %v, want 1.0.0", psC.PackageSpec) + } +} + +func TestBatchUpdateContinuation_FileRollbackOnInstallFailure(t *testing.T) { + // Verifies that when an update installer fails, the database retains the original package version, + // and subsequent package updates continue and succeed. + logger.Init("GooGet", true, false, io.Discard) + ctx := context.Background() + + settings.Initialize(t.TempDir(), false) + settings.Archs = []string{"noarch"} + if err := os.MkdirAll(settings.CacheDir(), 0755); err != nil { + t.Fatalf("os.MkdirAll cache: %v", err) + } + + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + + initialState := client.GooGetState{ + {PackageSpec: &goolib.PkgSpec{Name: "pkg_rollback", Arch: "noarch", Version: "1.0.0"}}, + {PackageSpec: &goolib.PkgSpec{Name: "pkg_success", Arch: "noarch", Version: "1.0.0"}}, + } + if err := db.WriteStateToDB(initialState); err != nil { + t.Fatalf("db.WriteStateToDB: %v", err) + } + db.Close() + + gooDir, logDir := t.TempDir(), t.TempDir() + srv := testutil.ServeGoo(t, gooDir) + defer srv.Close() + + // pkg_rollback update has failing installer. + pkgA2 := goolib.PkgSpec{ + Name: "pkg_rollback", + Arch: "noarch", + Version: "2.0.0", + Install: goolib.ExecFile{Path: "nonexistent_installer.exe"}, + } + rsA2 := testutil.GenGoo(t, gooDir, logDir, pkgA2) + + // pkg_success update is valid. + pkgB2 := goolib.PkgSpec{Name: "pkg_success", Arch: "noarch", Version: "2.0.0"} + rsB2 := testutil.GenGoo(t, gooDir, logDir, pkgB2) + + indexBytes, err := json.Marshal([]goolib.RepoSpec{rsA2, rsB2}) + if err != nil { + t.Fatalf("json.Marshal index: %v", err) + } + if err := os.WriteFile(filepath.Join(gooDir, "index"), indexBytes, 0644); err != nil { + t.Fatalf("writing index: %v", err) + } + var gzBuf bytes.Buffer + gw := gzip.NewWriter(&gzBuf) + if _, err := gw.Write(indexBytes); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + gw.Close() + if err := os.WriteFile(filepath.Join(gooDir, "index.gz"), gzBuf.Bytes(), 0644); err != nil { + t.Fatalf("writing index.gz: %v", err) + } + + cmd := &updateCmd{sources: srv.URL} + exitStatus := cmd.Execute(ctx, nil) + + if exitStatus != subcommands.ExitFailure { + t.Errorf("cmd.Execute got %v, want subcommands.ExitFailure (%v)", exitStatus, subcommands.ExitFailure) + } + + db, err = googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // 1. Verify pkg_rollback in DB remains at version 1.0.0. + psA, err := db.FetchPkg(goolib.PackageInfo{Name: "pkg_rollback", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkg_rollback): %v", err) + } + if psA.PackageSpec == nil || psA.PackageSpec.Version != "1.0.0" { + t.Errorf("pkg_rollback version in DB got %v, want 1.0.0", psA.PackageSpec) + } + + // 2. Verify pkg_success in DB is updated to version 2.0.0. + psB, err := db.FetchPkg(goolib.PackageInfo{Name: "pkg_success", Arch: "noarch"}) + if err != nil { + t.Fatalf("db.FetchPkg(pkg_success): %v", err) + } + if psB.PackageSpec == nil || psB.PackageSpec.Version != "2.0.0" { + t.Errorf("pkg_success version in DB got %v, want 2.0.0", psB.PackageSpec) + } +} diff --git a/client/client.go b/client/client.go index 80ec7ba..4d17258 100644 --- a/client/client.go +++ b/client/client.go @@ -19,6 +19,7 @@ import ( "context" "crypto/sha256" "encoding/json" + "errors" "fmt" "io" "io/ioutil" @@ -40,11 +41,11 @@ import ( "google.golang.org/api/googleapi" ) -// InstalledApplication describes the mapped Windows application to the package +// InstalledApplication describes the mapped Windows application to the package. type InstalledApplication struct { - // Display Name of the installed application found in the registry + // Display Name of the installed application found in the registry. Name string - // Registry key of the installed application in uninstall + // Registry key of the installed application in uninstall. Reg string } @@ -118,15 +119,19 @@ type Repo struct { // RepoMap describes each repo's packages as seen from a client. type RepoMap map[string]Repo -// Downloader is a wrapper around http.Client +// Downloader is a wrapper around http.Client. type Downloader struct { HTTPClient *http.Client UsingProxyServer bool + // StallTimeout is how long a repo index or package download may receive + // zero bytes before the transfer is aborted as stalled. Transfers that keep + // receiving data are never aborted, however slow. Zero means + // DefaultStallTimeout. + StallTimeout time.Duration } // NewDownloader returns a Downloader optionally using a specified proxyServer. func NewDownloader(proxyServer string) (*Downloader, error) { - httpClient := http.DefaultClient proxy := http.ProxyFromEnvironment if proxyServer != "" { proxyURL, err := url.Parse(proxyServer) @@ -135,7 +140,7 @@ func NewDownloader(proxyServer string) (*Downloader, error) { } proxy = http.ProxyURL(proxyURL) } - httpClient.Transport = &http.Transport{ + tr := &http.Transport{ Proxy: proxy, DialContext: (&net.Dialer{ Timeout: 30 * time.Second, @@ -145,9 +150,17 @@ func NewDownloader(proxyServer string) (*Downloader, error) { MaxIdleConns: 100, IdleConnTimeout: 60 * time.Second, TLSHandshakeTimeout: 10 * time.Second, + ResponseHeaderTimeout: 30 * time.Second, ExpectContinueTimeout: 1 * time.Second, } - return &Downloader{HTTPClient: httpClient, UsingProxyServer: proxyServer != ""}, nil + httpClient := &http.Client{ + Transport: tr, + } + return &Downloader{ + HTTPClient: httpClient, + UsingProxyServer: proxyServer != "", + StallTimeout: currentDefaultStallTimeout(), + }, nil } // AvailableVersions builds a RepoMap from a list of sources. @@ -261,28 +274,48 @@ func (d *Downloader) Get(ctx context.Context, path string) (*http.Response, erro } func (d *Downloader) unmarshalRepoPackagesHTTP(ctx context.Context, repoURL string, cf string) ([]goolib.RepoSpec, error) { + // A per-fetch cancelable context lets the StallReader abort a stalled + // index body without canceling the caller's context. + reqCtx, cancel := context.WithCancel(ctx) + defer cancel() + indexURL := repoURL + "/index.gz" trimmedIndexURL := strings.TrimPrefix(indexURL, "oauth-") ct := "application/x-gzip" logger.Infof("Fetching %q", trimmedIndexURL) - res, err := d.Get(ctx, indexURL) + res, err := d.Get(reqCtx, indexURL) if err != nil { return nil, err } if res.StatusCode != http.StatusOK { + res.Body.Close() indexURL = repoURL + "/index" trimmedIndexURL = strings.TrimPrefix(indexURL, "oauth-") ct = "application/json" logger.Infof("Fetching %q", trimmedIndexURL) - res, err = d.Get(ctx, indexURL) + res, err = d.Get(reqCtx, indexURL) if err != nil { return nil, err } if res.StatusCode != http.StatusOK { + res.Body.Close() return nil, fmt.Errorf("index GET request returned status: %q", res.Status) } } - return decode(res.Body, ct, repoURL, cf) + return decodeWithStallGuard(res.Body, cancel, d.StallTimeout, ct, repoURL, cf) +} + +// decodeWithStallGuard wraps index in a StallReader that invokes cancel when +// no bytes arrive for timeout, then decodes it. A timeout of zero or less +// means DefaultStallTimeout. A stalled body is reported as an error wrapping +// ErrDownloadStalled. +func decodeWithStallGuard(index io.ReadCloser, cancel context.CancelFunc, timeout time.Duration, ct, url, cf string) ([]goolib.RepoSpec, error) { + sr := NewStallReader(index, timeout, cancel) + m, err := decode(sr, ct, url, cf) + if err != nil && sr.isStalled() && !errors.Is(err, ErrDownloadStalled) { + err = fmt.Errorf("%w: %v", ErrDownloadStalled, err) + } + return m, err } func (d *Downloader) unmarshalRepoPackagesGCS(ctx context.Context, bucket, object, url, cf string) ([]goolib.RepoSpec, error) { @@ -292,10 +325,16 @@ func (d *Downloader) unmarshalRepoPackagesGCS(ctx context.Context, bucket, objec return empty, nil } - client, err := storage.NewClient(ctx) + // A per-fetch cancelable context lets the StallReader abort a stalled + // index read without canceling the caller's context. + reqCtx, cancel := context.WithCancel(ctx) + defer cancel() + + client, err := storage.NewClient(reqCtx) if err != nil { return nil, err } + defer client.Close() bkt := client.Bucket(bucket) if len(object) != 0 { @@ -304,22 +343,24 @@ func (d *Downloader) unmarshalRepoPackagesGCS(ctx context.Context, bucket, objec indexPath := object + "index.gz" logger.Infof("Fetching 'gs://%s/%s", bucket, indexPath) - if r, err := bkt.Object(indexPath).NewReader(ctx); err == nil { - return decode(r, "application/x-gzip", url, cf) + gzr, gzErr := bkt.Object(indexPath).NewReader(reqCtx) + if gzErr == nil { + return decodeWithStallGuard(gzr, cancel, d.StallTimeout, "application/x-gzip", url, cf) } - if gErr, ok := err.(*googleapi.Error); ok && gErr.Code != http.StatusNotFound { - return nil, err + var gErr *googleapi.Error + if errors.As(gzErr, &gErr) && gErr.Code != http.StatusNotFound { + return nil, gzErr } logger.Info("Failed to read gzipped index, trying plain JSON.") indexPath = object + "index" - r, err := bkt.Object(indexPath).NewReader(ctx) + r, err := bkt.Object(indexPath).NewReader(reqCtx) if err != nil { return nil, err } - return decode(r, "application/json", url, cf) + return decodeWithStallGuard(r, cancel, d.StallTimeout, "application/json", url, cf) } func decode(index io.ReadCloser, ct, url, cf string) ([]goolib.RepoSpec, error) { @@ -462,7 +503,7 @@ func FindRepoLatest(pi goolib.PackageInfo, rm RepoMap, archs []string, installed return 0 } if c != 0 { - return -c // reverse for descending order + return -c // Reverse for descending order. } if archPref[a.spec.Arch] < archPref[b.spec.Arch] { return -1 @@ -480,7 +521,7 @@ func FindRepoLatest(pi goolib.PackageInfo, rm RepoMap, archs []string, installed slices.SortFunc(list, cmpFunc) for _, cand := range list { if isLocked && cand.spec.Arch != installedArch && cand.spec.LockArch { - continue // Ignore this candidate + continue // Ignore this candidate. } return cand.spec, cand.repo, cand.spec.Arch, nil } diff --git a/client/client_test.go b/client/client_test.go index ff21404..5dd3d68 100644 --- a/client/client_test.go +++ b/client/client_test.go @@ -19,9 +19,11 @@ import ( "context" "crypto/sha256" "encoding/json" + "errors" "fmt" "io" "io/ioutil" + "net" "net/http" "net/http/httptest" "path/filepath" @@ -301,7 +303,7 @@ func TestFindRepoLatest(t *testing.T) { { desc: "cross arch upgrade", pi: goolib.PackageInfo{Name: "foo_pkg"}, - archs: []string{"x86_32", "x86_64"}, // Prefer 32-bit + archs: []string{"x86_32", "x86_64"}, // Prefer 32-bit. rm: RepoMap{ "repo": Repo{ Packages: []goolib.RepoSpec{ @@ -331,7 +333,7 @@ func TestFindRepoLatest(t *testing.T) { {PackageSpec: &goolib.PkgSpec{Name: "foo_pkg", Version: "3.0.0@1", Arch: "x86_64"}}, }, }, - "low_pri": Repo{ // Should win if version was primary + "low_pri": Repo{ // Should win if version was primary. Priority: 100, Packages: []goolib.RepoSpec{ {PackageSpec: &goolib.PkgSpec{Name: "foo_pkg", Version: "4.0.0@1", Arch: "x86_64"}}, @@ -674,7 +676,7 @@ func TestFindRepoLatest_Provides(t *testing.T) { name: "Provider match unversioned", pi: goolib.PackageInfo{Name: "virtual_pkg", Arch: "noarch"}, wantName: "real_pkg", - wantVer: "2.0.0", // latest real_pkg + wantVer: "2.0.0", // Latest real_pkg. }, { name: "Provider match matched version", @@ -709,7 +711,7 @@ func TestFindRepoLatest_Provides(t *testing.T) { if spec.Name != tt.wantName { t.Errorf("FindRepoLatest(%v) name = %q, want %q", tt.pi, spec.Name, tt.wantName) } - if spec.Version != tt.wantVer { // Simplified check, assumes simple version strings in test + if spec.Version != tt.wantVer { // Simplified check; assumes simple version strings in test. t.Errorf("FindRepoLatest(%v) version = %q, want %q", tt.pi, spec.Version, tt.wantVer) } }) @@ -718,9 +720,9 @@ func TestFindRepoLatest_Provides(t *testing.T) { func TestFindRepoLatest_Priority(t *testing.T) { // Setup repo with both direct match and provider. - // direct match: version 1.0.0 - // provider: version 2.0.0 (provides it) - // direct match should win despite lower version. + // Direct match: version 1.0.0. + // Provider: version 2.0.0 (provides it). + // Direct match should win despite lower version. rm := RepoMap{ "repo1": Repo{ @@ -827,3 +829,215 @@ func TestFindRepoLatest_LockArch(t *testing.T) { }) } } + +func TestNewDownloader_DedicatedClientAndTransportConfig(t *testing.T) { + // Verify that NewDownloader instantiates an isolated client and sets transport timeouts. + origDefaultTransport := http.DefaultClient.Transport + + dl, err := NewDownloader("") + if err != nil { + t.Fatalf("NewDownloader(\"\") returned unexpected error: %v", err) + } + if dl == nil { + t.Fatal("NewDownloader(\"\") returned nil Downloader") + } + if dl.HTTPClient == nil { + t.Fatal("Downloader.HTTPClient is nil") + } + + // Verify isolation from http.DefaultClient. + if dl.HTTPClient == http.DefaultClient { + t.Error("Downloader.HTTPClient must not be http.DefaultClient") + } + if http.DefaultClient.Transport != origDefaultTransport { + t.Errorf("http.DefaultClient.Transport was mutated: got %v, want %v", http.DefaultClient.Transport, origDefaultTransport) + } + + // Verify that overall client timeout is 0 to allow streaming large downloads. + if dl.HTTPClient.Timeout != 0 { + t.Errorf("dl.HTTPClient.Timeout = %v, want 0", dl.HTTPClient.Timeout) + } + + // Verify transport configuration. + tr, ok := dl.HTTPClient.Transport.(*http.Transport) + if !ok { + t.Fatalf("dl.HTTPClient.Transport is %T, want *http.Transport", dl.HTTPClient.Transport) + } + if dl.HTTPClient.Transport == origDefaultTransport && origDefaultTransport != nil { + t.Error("dl.HTTPClient.Transport shares pointer with http.DefaultClient.Transport") + } + + const wantHeaderTimeout = 30 * time.Second + if tr.ResponseHeaderTimeout != wantHeaderTimeout { + t.Errorf("tr.ResponseHeaderTimeout = %v, want %v", tr.ResponseHeaderTimeout, wantHeaderTimeout) + } + const wantIdleConnTimeout = 60 * time.Second + if tr.IdleConnTimeout != wantIdleConnTimeout { + t.Errorf("tr.IdleConnTimeout = %v, want %v", tr.IdleConnTimeout, wantIdleConnTimeout) + } + const wantTLSHandshakeTimeout = 10 * time.Second + if tr.TLSHandshakeTimeout != wantTLSHandshakeTimeout { + t.Errorf("tr.TLSHandshakeTimeout = %v, want %v", tr.TLSHandshakeTimeout, wantTLSHandshakeTimeout) + } + const wantExpectContinueTimeout = 1 * time.Second + if tr.ExpectContinueTimeout != wantExpectContinueTimeout { + t.Errorf("tr.ExpectContinueTimeout = %v, want %v", tr.ExpectContinueTimeout, wantExpectContinueTimeout) + } + if tr.MaxIdleConns != 100 { + t.Errorf("tr.MaxIdleConns = %d, want 100", tr.MaxIdleConns) + } + if !tr.ForceAttemptHTTP2 { + t.Error("tr.ForceAttemptHTTP2 = false, want true") + } +} + +func TestNewDownloader_ProxyConfiguration(t *testing.T) { + // Verify proxy configuration options. + t.Run("no proxy", func(t *testing.T) { + dl, err := NewDownloader("") + if err != nil { + t.Fatalf("NewDownloader(\"\") failed: %v", err) + } + if dl.UsingProxyServer { + t.Error("UsingProxyServer = true, want false") + } + }) + + t.Run("valid proxy", func(t *testing.T) { + proxyURL := "http://proxy.example.com:8080" + dl, err := NewDownloader(proxyURL) + if err != nil { + t.Fatalf("NewDownloader(%q) failed: %v", proxyURL, err) + } + if !dl.UsingProxyServer { + t.Error("UsingProxyServer = false, want true") + } + }) + + t.Run("invalid proxy", func(t *testing.T) { + if _, err := NewDownloader("://invalid-url"); err == nil { + t.Error("NewDownloader with invalid proxy expected error, got nil") + } + }) +} + +func TestNewDownloader_ResponseHeaderTimeoutFunctional(t *testing.T) { + // Verify that ResponseHeaderTimeout terminates requests when server headers stall. + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Delay sending response headers well past the timeout, returning + // early once the client gives up. + select { + case <-time.After(10 * time.Second): + case <-r.Context().Done(): + } + w.WriteHeader(http.StatusOK) + })) + defer server.Close() + + dl, err := NewDownloader("") + if err != nil { + t.Fatalf("NewDownloader failed: %v", err) + } + + tr, ok := dl.HTTPClient.Transport.(*http.Transport) + if !ok { + t.Fatalf("dl.HTTPClient.Transport is %T, want *http.Transport", dl.HTTPClient.Transport) + } + // Use a short header timeout for fast unit testing. + tr.ResponseHeaderTimeout = 40 * time.Millisecond + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + _, err = dl.Get(ctx, server.URL) + if err == nil { + t.Fatal("dl.Get() succeeded, expected timeout error") + } + if ctx.Err() != nil { + t.Fatalf("dl.Get() = %v after the context deadline, want ResponseHeaderTimeout to end it first", err) + } + + var netErr net.Error + if errors.As(err, &netErr) && !netErr.Timeout() { + t.Errorf("expected timeout net.Error, got %v", err) + } +} + +func TestUnmarshalRepoPackagesHTTP_IndexBodyStall(t *testing.T) { + // Verify that a repo index whose body stalls mid-stream aborts with + // ErrDownloadStalled instead of hanging. + t.Parallel() + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/index" { + w.WriteHeader(http.StatusNotFound) + return + } + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Content-Length", "1000") + w.WriteHeader(http.StatusOK) + w.Write([]byte(`[{"Source": "foo"}, `)) + w.(http.Flusher).Flush() + select { + case <-time.After(10 * time.Second): + case <-r.Context().Done(): + } + })) + defer ts.Close() + + d, err := NewDownloader("") + if err != nil { + t.Fatalf("NewDownloader: %v", err) + } + d.StallTimeout = 50 * time.Millisecond + start := time.Now() + _, err = d.unmarshalRepoPackages(context.Background(), ts.URL, t.TempDir(), cacheLife) + if !errors.Is(err, ErrDownloadStalled) { + t.Fatalf("unmarshalRepoPackages() = %v, want error wrapping ErrDownloadStalled", err) + } + if elapsed := time.Since(start); elapsed > 5*time.Second { + t.Errorf("unmarshalRepoPackages took %v, want prompt stall detection", elapsed) + } +} + +func TestUnmarshalRepoPackagesHTTP_SlowIndexCompletes(t *testing.T) { + // Verify that a slowly trickling index body that keeps making progress is + // not aborted by the stall guard. + t.Parallel() + want := []goolib.RepoSpec{{Source: "foo"}, {Source: "bar"}} + j, err := json.Marshal(want) + if err != nil { + t.Fatalf("json.Marshal: %v", err) + } + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/index" { + w.WriteHeader(http.StatusNotFound) + return + } + w.WriteHeader(http.StatusOK) + for i := 0; i < len(j); i += 8 { + end := i + 8 + if end > len(j) { + end = len(j) + } + w.Write(j[i:end]) + w.(http.Flusher).Flush() + time.Sleep(5 * time.Millisecond) + } + })) + defer ts.Close() + + d, err := NewDownloader("") + if err != nil { + t.Fatalf("NewDownloader: %v", err) + } + // The stall timeout is 50 times the delay between chunks so that the + // test stays reliable on loaded machines. + d.StallTimeout = 250 * time.Millisecond + got, err := d.unmarshalRepoPackages(context.Background(), ts.URL, t.TempDir(), cacheLife) + if err != nil { + t.Fatalf("unmarshalRepoPackages() = %v, want nil", err) + } + if !reflect.DeepEqual(got, want) { + t.Errorf("unmarshalRepoPackages() = %+v, want %+v", got, want) + } +} diff --git a/client/stall.go b/client/stall.go new file mode 100644 index 0000000..6c593e4 --- /dev/null +++ b/client/stall.go @@ -0,0 +1,176 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package client + +import ( + "context" + "errors" + "io" + "sync" + "sync/atomic" + "time" +) + +// ErrDownloadStalled is returned when a download connection experiences zero bytes +// read for greater than the idle timeout period. +var ErrDownloadStalled = errors.New("download stalled: no data received for timeout period") + +// DefaultStallTimeout is the default duration of zero bytes received before a download is considered stalled. +const DefaultStallTimeout = 120 * time.Second + +// defaultStallTimeout holds the process-wide stall timeout, in nanoseconds, +// that NewDownloader copies into Downloader.StallTimeout. Zero means +// DefaultStallTimeout. +var defaultStallTimeout atomic.Int64 + +// SetDefaultStallTimeout sets the stall timeout that Downloaders created by +// later NewDownloader calls use for repo index and package downloads. A value +// of zero or less restores DefaultStallTimeout. Existing Downloaders are not +// affected. It is safe for concurrent use. +// +// It deliberately mirrors supervisor.Configure: the googet command sets both +// process-wide defaults once from its settings at startup, and +// NewDownloader copies this default into Downloader.StallTimeout, which +// callers may still override per Downloader. +func SetDefaultStallTimeout(d time.Duration) { + if d < 0 { + d = 0 + } + defaultStallTimeout.Store(int64(d)) +} + +// currentDefaultStallTimeout returns the timeout set by SetDefaultStallTimeout, +// or DefaultStallTimeout if none is set. +func currentDefaultStallTimeout() time.Duration { + if d := time.Duration(defaultStallTimeout.Load()); d > 0 { + return d + } + return DefaultStallTimeout +} + +// StallReader wraps an io.Reader (typically an http.Response.Body) with an idle-read watchdog timer. +type StallReader struct { + r io.Reader + timeout time.Duration + cancel context.CancelFunc + + mu sync.Mutex + timer *time.Timer + stalled bool + closed bool +} + +// NewStallReader returns a StallReader that wraps r with an idle read watchdog. +// When no bytes are read for timeout, cancel is invoked so that any blocked +// read on r (which must honor the canceled context) returns, and subsequent +// reads return ErrDownloadStalled. A timeout of zero or less uses +// DefaultStallTimeout. Callers must call Close to release the timer. +func NewStallReader(r io.Reader, timeout time.Duration, cancel context.CancelFunc) *StallReader { + if timeout <= 0 { + timeout = DefaultStallTimeout + } + s := &StallReader{ + r: r, + timeout: timeout, + cancel: cancel, + } + s.timer = time.AfterFunc(timeout, s.onStall) + return s +} + +// onStall is invoked by the idle timer when the timeout expires without reading bytes. +func (s *StallReader) onStall() { + s.mu.Lock() + if s.closed || s.stalled { + s.mu.Unlock() + return + } + s.stalled = true + s.mu.Unlock() + + // Cancel context outside the mutex to prevent deadlocks. + if s.cancel != nil { + s.cancel() + } +} + +// isStalled reports whether the idle timer has fired. +func (s *StallReader) isStalled() bool { + s.mu.Lock() + defer s.mu.Unlock() + return s.stalled +} + +// Read reads from the underlying reader. Any read that returns n > 0 bytes resets +// the idle timer. If the idle timer expires before bytes arrive, it returns ErrDownloadStalled. +func (s *StallReader) Read(p []byte) (int, error) { + s.mu.Lock() + if s.stalled { + s.mu.Unlock() + return 0, ErrDownloadStalled + } + if s.closed { + s.mu.Unlock() + return 0, io.ErrClosedPipe + } + s.mu.Unlock() + + n, err := s.r.Read(p) + + s.mu.Lock() + defer s.mu.Unlock() + + if s.stalled { + // If the stream finished cleanly with io.EOF, prefer EOF over stall. + if err == io.EOF { + return n, io.EOF + } + return 0, ErrDownloadStalled + } + + if n > 0 { + // Forward progress made: reset the idle timer. + if !s.closed && s.timer != nil { + s.timer.Reset(s.timeout) + } + } + + if err != nil { + // Stop the timer on stream termination (EOF or non-stall network error). + if s.timer != nil { + s.timer.Stop() + } + } + + return n, err +} + +// Close stops the idle timer and closes the underlying reader if it implements io.Closer. +func (s *StallReader) Close() error { + s.mu.Lock() + if s.closed { + s.mu.Unlock() + return nil + } + s.closed = true + if s.timer != nil { + s.timer.Stop() + } + s.mu.Unlock() + + if c, ok := s.r.(io.Closer); ok { + return c.Close() + } + return nil +} diff --git a/client/stall_test.go b/client/stall_test.go new file mode 100644 index 0000000..49cf4c1 --- /dev/null +++ b/client/stall_test.go @@ -0,0 +1,482 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package client + +import ( + "bytes" + "context" + "errors" + "io" + "sync" + "testing" + "time" +) + +// The timing constants below keep wide margins so that the tests stay +// reliable on loaded machines and under the race detector. Tests that expect +// no stall deliver data far more often than the stall timeout, and tests that +// expect a stall wait on the cancel callback with a deadline far longer than +// the stall timeout instead of sleeping for a fixed time. +const ( + // trickleTimeout is the stall timeout for streams that must not stall. + trickleTimeout = 250 * time.Millisecond + // trickleDelay is the delay between chunks of a stream that must not + // stall. It is 50 times shorter than trickleTimeout. + trickleDelay = 5 * time.Millisecond + // stallTimeout is the stall timeout for streams that must stall. + stallTimeout = 30 * time.Millisecond + // stallDeadline bounds how long a test waits for an expected stall. + stallDeadline = 10 * time.Second +) + +// cancelRecorder records calls to its cancel method. +type cancelRecorder struct { + once sync.Once + done chan struct{} +} + +func newCancelRecorder() *cancelRecorder { + return &cancelRecorder{done: make(chan struct{})} +} + +// cancel records that the stall callback ran. +func (c *cancelRecorder) cancel() { + c.once.Do(func() { close(c.done) }) +} + +// canceled reports whether cancel has been called. +func (c *cancelRecorder) canceled() bool { + select { + case <-c.done: + return true + default: + return false + } +} + +// wait reports whether cancel is called within stallDeadline. +func (c *cancelRecorder) wait() bool { + select { + case <-c.done: + return true + case <-time.After(stallDeadline): + return false + } +} + +func TestStallReader_Normal(t *testing.T) { + // Verify that reading a complete buffer through StallReader succeeds without error. + data := []byte("hello world from googet stall reader test") + _, cancel := context.WithCancel(context.Background()) + defer cancel() + + sr := NewStallReader(bytes.NewReader(data), stallDeadline, cancel) + defer sr.Close() + + buf, err := io.ReadAll(sr) + if err != nil { + t.Fatalf("unexpected error reading from StallReader: %v", err) + } + if !bytes.Equal(buf, data) { + t.Errorf("read content mismatch: got %q, want %q", string(buf), string(data)) + } +} + +type trickleReader struct { + chunks [][]byte + delay time.Duration + index int +} + +func (r *trickleReader) Read(p []byte) (int, error) { + if r.index >= len(r.chunks) { + return 0, io.EOF + } + time.Sleep(r.delay) + n := copy(p, r.chunks[r.index]) + r.index++ + return n, nil +} + +func TestStallReader_SlowTrickle(t *testing.T) { + // Verify that slow trickle streams complete as long as chunks arrive within timeout. + chunks := [][]byte{ + []byte("chunk1-"), + []byte("chunk2-"), + []byte("chunk3-"), + []byte("chunk4-"), + []byte("chunk5"), + } + r := &trickleReader{chunks: chunks, delay: trickleDelay} + + _, cancel := context.WithCancel(context.Background()) + defer cancel() + + sr := NewStallReader(r, trickleTimeout, cancel) + defer sr.Close() + + buf, err := io.ReadAll(sr) + if err != nil { + t.Fatalf("slow trickle read unexpectedly failed: %v", err) + } + expected := "chunk1-chunk2-chunk3-chunk4-chunk5" + if string(buf) != expected { + t.Errorf("trickle content mismatch: got %q, want %q", string(buf), expected) + } +} + +func TestStallReader_StallDetection(t *testing.T) { + // Verify that StallReader returns ErrDownloadStalled when data transfer freezes. + pr, pw := io.Pipe() + defer pr.Close() + defer pw.Close() + + _, cancel := context.WithCancel(context.Background()) + // Cancel wrapper closes pipe to unblock Read. + cancelWrapper := func() { + cancel() + pw.CloseWithError(context.Canceled) + } + + sr := NewStallReader(pr, stallTimeout, cancelWrapper) + defer sr.Close() + + // Write initial chunk, then stall. + go func() { + pw.Write([]byte("initial data")) + // Do not write anything further. + }() + + buf := make([]byte, 1024) + n, err := sr.Read(buf) + if err != nil { + t.Fatalf("first read failed: %v", err) + } + if n != len("initial data") { + t.Fatalf("got %d bytes on first read, want %d", n, len("initial data")) + } + + // The second read blocks until the stall timer cancels it. + n, err = sr.Read(buf) + if !errors.Is(err, ErrDownloadStalled) { + t.Fatalf("expected ErrDownloadStalled on stalled stream, got: %v (n=%d)", err, n) + } +} + +func TestStallReader_ExternalCancellation(t *testing.T) { + // Verify that external context cancellation is preserved and not reported as a stall. + pr, pw := io.Pipe() + defer pr.Close() + defer pw.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + sr := NewStallReader(pr, stallDeadline, cancel) + defer sr.Close() + + // Cancel external context well before stall timeout. + go func() { + time.Sleep(30 * time.Millisecond) + cancel() + pw.CloseWithError(ctx.Err()) + }() + + buf := make([]byte, 64) + _, err := sr.Read(buf) + if errors.Is(err, ErrDownloadStalled) { + t.Fatalf("misreported external context cancellation as ErrDownloadStalled") + } + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context.Canceled, got: %v", err) + } +} + +type zeroByteReader struct { + delay time.Duration + mu sync.Mutex + calls int +} + +func (z *zeroByteReader) Read(p []byte) (int, error) { + z.mu.Lock() + z.calls++ + z.mu.Unlock() + if z.delay > 0 { + time.Sleep(z.delay) + } + return 0, nil +} + +func (z *zeroByteReader) callCount() int { + z.mu.Lock() + defer z.mu.Unlock() + return z.calls +} + +func TestStallReader_ZeroByteReadsDoNotResetTimer(t *testing.T) { + // Verify that repeated zero-byte reads (n=0, err=nil) do NOT reset the idle timer. + z := &zeroByteReader{delay: time.Millisecond} + c := newCancelRecorder() + + // start is taken before the timer is created, so the stall can never be + // observed earlier than stallTimeout after it. + start := time.Now() + sr := NewStallReader(z, stallTimeout, c.cancel) + defer sr.Close() + + buf := make([]byte, 64) + var stallErr error + var readCount int + + for time.Since(start) < stallDeadline { + readCount++ + n, err := sr.Read(buf) + if n != 0 { + t.Fatalf("expected 0 bytes, got %d", n) + } + if err != nil { + stallErr = err + break + } + } + + elapsed := time.Since(start) + if !errors.Is(stallErr, ErrDownloadStalled) { + t.Fatalf("expected ErrDownloadStalled from repeated zero-byte reads, got: %v (elapsed: %v, reads: %d)", stallErr, elapsed, readCount) + } + + if elapsed < stallTimeout { + t.Errorf("stalled prematurely: elapsed %v < timeout %v", elapsed, stallTimeout) + } + + if !c.canceled() { + t.Error("expected cancel() to be invoked on stall, but it was not") + } + + // Verify subsequent reads immediately return ErrDownloadStalled without calling underlying reader. + callsBefore := z.callCount() + n, err := sr.Read(buf) + if !errors.Is(err, ErrDownloadStalled) || n != 0 { + t.Errorf("subsequent read got (%d, %v), want (0, ErrDownloadStalled)", n, err) + } + if got := z.callCount(); got != callsBefore { + t.Errorf("underlying reader called after stall: %d -> %d", callsBefore, got) + } +} + +func TestStallReader_EmptyBufferRead(t *testing.T) { + // Verify that Read with empty buffer (len=0) does NOT prevent the idle + // timer from expiring. + data := []byte("some data") + r := bytes.NewReader(data) + c := newCancelRecorder() + sr := NewStallReader(r, stallTimeout, c.cancel) + defer sr.Close() + + // Read with an empty buffer. + n, err := sr.Read([]byte{}) + if n != 0 || err != nil { + t.Fatalf("Read([]byte{}) got (%d, %v), want (0, nil)", n, err) + } + + // Since n was 0, the timer must still expire. + if !c.wait() { + t.Fatalf("cancel was not called within %v after an empty buffer read with stall timeout %v", stallDeadline, stallTimeout) + } + + buf := make([]byte, 10) + n, err = sr.Read(buf) + if !errors.Is(err, ErrDownloadStalled) { + t.Fatalf("expected ErrDownloadStalled after timeout, got: (%d, %v)", n, err) + } +} + +func TestStallReader_ExtendedSlowTrickle(t *testing.T) { + // Verify that a slow trickle stream whose total duration is several times + // the stall timeout completes successfully without premature abortion. + numChunks := 150 + chunkSize := 10 + var chunks [][]byte + var expected bytes.Buffer + for i := 0; i < numChunks; i++ { + chunk := bytes.Repeat([]byte{byte('a' + (i % 26))}, chunkSize) + chunks = append(chunks, chunk) + expected.Write(chunk) + } + + r := &trickleReader{chunks: chunks, delay: trickleDelay} + c := newCancelRecorder() + sr := NewStallReader(r, trickleTimeout, c.cancel) + defer sr.Close() + + start := time.Now() + buf, err := io.ReadAll(sr) + elapsed := time.Since(start) + + if err != nil { + t.Fatalf("extended slow trickle failed prematurely: %v (elapsed: %v)", err, elapsed) + } + + if !bytes.Equal(buf, expected.Bytes()) { + t.Errorf("read content mismatch: got %d bytes, want %d bytes", len(buf), expected.Len()) + } + + // Total elapsed is at least 150 * 5ms = 750ms, which is 3x the timeout. + minExpectedElapsed := time.Duration(numChunks) * trickleDelay + if elapsed < minExpectedElapsed { + t.Errorf("elapsed time %v < expected %v", elapsed, minExpectedElapsed) + } + + if c.canceled() { + t.Error("cancel() was unexpectedly invoked on successful slow trickle stream") + } +} + +func TestStallReader_CompleteStallReliability(t *testing.T) { + // Verify that complete stalls reliably trigger ErrDownloadStalled and context cancellation across multiple runs. + for i := 0; i < 5; i++ { + pr, pw := io.Pipe() + c := newCancelRecorder() + cancel := func() { + c.cancel() + pw.CloseWithError(context.Canceled) + } + + // start is taken before the first read resets the timer, so the stall + // can never be observed earlier than stallTimeout after it. + start := time.Now() + sr := NewStallReader(pr, stallTimeout, cancel) + + // Write initial byte then stall completely. + go func() { + pw.Write([]byte("x")) + }() + + buf := make([]byte, 10) + n, err := sr.Read(buf) + if err != nil || n != 1 { + sr.Close() + t.Fatalf("iter %d: initial read failed: n=%d, err=%v", i, n, err) + } + + _, err = sr.Read(buf) + elapsed := time.Since(start) + sr.Close() + + if !errors.Is(err, ErrDownloadStalled) { + t.Fatalf("iter %d: expected ErrDownloadStalled, got %v", i, err) + } + if elapsed < stallTimeout { + t.Errorf("iter %d: stall triggered prematurely: %v < %v", i, elapsed, stallTimeout) + } + if !c.canceled() { + t.Fatalf("iter %d: context cancel func was not called on stall", i) + } + } +} + +func TestStallReader_ConcurrentStress(t *testing.T) { + // Concurrently run multiple StallReaders with mixed behaviors (trickle, stall, zero-byte). + var wg sync.WaitGroup + numWorkers := 20 + + for i := 0; i < numWorkers; i++ { + wg.Add(1) + go func(workerID int) { + defer wg.Done() + if workerID%3 == 0 { + // Trickle worker. + chunks := [][]byte{[]byte("a"), []byte("b"), []byte("c")} + r := &trickleReader{chunks: chunks, delay: trickleDelay} + sr := NewStallReader(r, trickleTimeout, nil) + defer sr.Close() + data, err := io.ReadAll(sr) + if err != nil || string(data) != "abc" { + t.Errorf("worker %d trickle failed: %v, got %q", workerID, err, string(data)) + } + } else if workerID%3 == 1 { + // Stall worker. + pr, pw := io.Pipe() + sr := NewStallReader(pr, stallTimeout, func() { + pw.CloseWithError(context.Canceled) + }) + defer sr.Close() + buf := make([]byte, 10) + _, err := sr.Read(buf) + if !errors.Is(err, ErrDownloadStalled) { + t.Errorf("worker %d expected stall, got %v", workerID, err) + } + } else { + // Zero-byte reader worker. + z := &zeroByteReader{delay: time.Millisecond} + sr := NewStallReader(z, stallTimeout, nil) + defer sr.Close() + buf := make([]byte, 10) + for { + _, err := sr.Read(buf) + if err != nil { + if !errors.Is(err, ErrDownloadStalled) { + t.Errorf("worker %d expected stall, got %v", workerID, err) + } + break + } + } + } + }(i) + } + + wg.Wait() +} + +func TestStallReader_CloseStopsTimer(t *testing.T) { + // Verify that Close stops the idle timer so no cancel fires afterwards. + c := newCancelRecorder() + sr := NewStallReader(bytes.NewReader([]byte("data")), trickleTimeout, c.cancel) + if err := sr.Close(); err != nil { + t.Fatalf("Close() = %v, want nil", err) + } + // Wait well past the timeout that Close must have stopped. + time.Sleep(3 * trickleTimeout) + if c.canceled() { + t.Error("cancel was invoked after Close") + } + if sr.isStalled() { + t.Error("isStalled() = true after Close, want false") + } +} + +func TestSetDefaultStallTimeout(t *testing.T) { + // Verify that NewDownloader copies the process-wide default stall timeout + // and that zero or negative values restore DefaultStallTimeout. + t.Cleanup(func() { SetDefaultStallTimeout(0) }) + tests := []struct { + set time.Duration + want time.Duration + }{ + {0, DefaultStallTimeout}, + {5 * time.Second, 5 * time.Second}, + {-1, DefaultStallTimeout}, + } + for _, tc := range tests { + SetDefaultStallTimeout(tc.set) + d, err := NewDownloader("") + if err != nil { + t.Fatalf("NewDownloader: %v", err) + } + if d.StallTimeout != tc.want { + t.Errorf("NewDownloader().StallTimeout after SetDefaultStallTimeout(%v) = %v, want %v", tc.set, d.StallTimeout, tc.want) + } + } +} diff --git a/download/download.go b/download/download.go index a71c500..e001d34 100644 --- a/download/download.go +++ b/download/download.go @@ -22,12 +22,18 @@ import ( "encoding/hex" "errors" "fmt" + "hash" "io" + "math/rand" + "net" "net/http" "net/url" "os" "path/filepath" + "strconv" "strings" + "syscall" + "time" "cloud.google.com/go/storage" "github.com/dustin/go-humanize" @@ -35,6 +41,7 @@ import ( "github.com/google/googet/v2/goolib" "github.com/google/googet/v2/oswrap" "github.com/google/logger" + "google.golang.org/api/googleapi" ) // Package downloads a package from the given url, @@ -46,95 +53,564 @@ func Package(ctx context.Context, pkgURL, dst, chksum string, downloader *client if err := oswrap.RemoveAll(dst); err != nil { return err } - return packageGCS(ctx, bucket, object, dst, chksum) + return packageGCS(ctx, bucket, object, dst, chksum, stallTimeout(downloader)) } return packageHTTP(ctx, pkgURL, dst, chksum, downloader) } -// packageHTTP downloads a package from an HTTP(S) server. -func packageHTTP(ctx context.Context, url, dst, chksum string, downloader *client.Downloader) error { - // Try to open any already existing file, otherwise create new file. - f, err := os.OpenFile(dst, os.O_CREATE|os.O_RDWR, 0644) - if err != nil { - return err +// stallTimeout returns the stall timeout configured on downloader. Zero, which +// is also returned for a nil downloader, means client.DefaultStallTimeout. +func stallTimeout(downloader *client.Downloader) time.Duration { + if downloader == nil { + return 0 } - defer f.Close() - // Hash the contents of the existing file. - hash := sha256.New() - size, err := io.Copy(hash, f) - if err != nil { - return err - } - // If the file checksum matches what we expect, then the file is already - // downloaded and we can quit early. - if sum := hex.EncodeToString(hash.Sum(nil)); sum == chksum { - logger.Infof("using existing file: %s (sum = %s)", dst, sum) + return downloader.StallTimeout +} + +const ( + // maxNoProgressRetries is the number of consecutive retries allowed after + // attempts that did not advance the download past its previous high-water + // mark. The counter resets whenever an attempt makes forward progress. + maxNoProgressRetries = 3 + // minAttemptProgress is the forward progress, in bytes, that a failed + // attempt must make to be exempt from maxLowProgressAttempts. + minAttemptProgress = 1 << 20 + // maxLowProgressAttempts caps the failed attempts that each advanced the + // download by less than minAttemptProgress, so that a connection that + // repeatedly delivers a few bytes and then fails cannot loop forever. + // Attempts that make at least minAttemptProgress never count toward it, so + // a slow or flaky link that keeps making progress is never abandoned. + maxLowProgressAttempts = 20 + // baseBackoff is the delay before the first retry. + baseBackoff = time.Second + // maxBackoff caps the exponential backoff delay before jitter is applied. + maxBackoff = 30 * time.Second + // backoffJitter is the maximum relative jitter applied to each delay. + backoffJitter = 0.2 +) + +// sleep waits for d or until ctx is done, whichever happens first. It is a +// package variable so that tests can avoid real delays. +var sleep = func(ctx context.Context, d time.Duration) error { + t := time.NewTimer(d) + defer t.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-t.C: return nil } - // Otherwise we have either an empty or partial download. - // Check that the server supports ranged requests and that the - // existing file is smaller than what we want to download. - logger.Infof("existing file size: %d", size) - ok, length, err := downloader.CanResume(ctx, url) - if err != nil { - logger.Errorf("CanResume: %v", err) +} + +// backoffDelay returns the delay before the next retry given the number of +// consecutive attempts that made no progress. The sequence is 1s, 1s, 2s, 4s, +// and so on, capped at maxBackoff, with +/- backoffJitter applied. +func backoffDelay(noProgress int) time.Duration { + shift := noProgress - 1 + if shift < 0 { + shift = 0 } - req, err := downloader.NewRequest(ctx, http.MethodGet, url, nil) + d := maxBackoff + if shift < 6 { + if b := baseBackoff << shift; b < maxBackoff { + d = b + } + } + jitter := 1 + backoffJitter*(2*rand.Float64()-1) + return time.Duration(float64(d) * jitter) +} + +// opener opens a stream of the object starting at offset. It returns the +// stream and the offset the stream actually starts at, which is either offset +// or 0 when the source ignored the resume request. It returns an error +// wrapping errResumeRejected when the source cannot serve offset. +type opener func(ctx context.Context, offset int64) (io.ReadCloser, int64, error) + +// errResumeRejected reports that the source cannot serve the requested resume +// offset, so the partial file must be discarded and the download restarted +// from the beginning. +var errResumeRejected = errors.New("source cannot resume at the requested offset") + +// statusError reports a non-successful HTTP status from a download request. +type statusError struct { + code int + status string +} + +// Error implements the error interface. +func (e *statusError) Error() string { + return "unexpected HTTP status " + e.status +} + +// checksumError reports that a completed download does not match its +// expected SHA256 checksum. +type checksumError struct { + got, want string +} + +// Error implements the error interface. +func (e *checksumError) Error() string { + return fmt.Sprintf("checksum doesn't match: got %s, want %s", e.got, e.want) +} + +// writeError wraps an error returned while writing to the destination file so +// that disk errors are never mistaken for retryable network errors. +type writeError struct { + err error +} + +// Error implements the error interface. +func (e *writeError) Error() string { + return "writing download to disk: " + e.err.Error() +} + +// Unwrap returns the underlying write error. +func (e *writeError) Unwrap() error { + return e.err +} + +// fileWriter wraps writer errors in writeError. +type fileWriter struct { + w io.Writer +} + +// Write implements io.Writer. +func (fw fileWriter) Write(p []byte) (int, error) { + n, err := fw.w.Write(p) if err != nil { - return err + err = &writeError{err: err} + } + return n, err +} + +// isRetryable reports whether err from a download attempt is transient and +// may succeed on a subsequent attempt. Errors are never retryable once the +// parent context is done. +func isRetryable(ctx context.Context, err error) bool { + if err == nil || ctx.Err() != nil { + return false } - if ok && size < length { - logger.Infof("resuming download of %s (%d bytes remaining)", url, length-size) - req.Header.Add("Range", fmt.Sprintf("bytes=%d-", size)) + var we *writeError + if errors.As(err, &we) { + return false + } + var se *statusError + if errors.As(err, &se) { + return se.code >= 500 || se.code == http.StatusTooManyRequests + } + var ge *googleapi.Error + if errors.As(err, &ge) { + return ge.Code >= 500 || ge.Code == http.StatusTooManyRequests + } + if errors.Is(err, client.ErrDownloadStalled) || + errors.Is(err, io.ErrUnexpectedEOF) || + errors.Is(err, io.EOF) || + errors.Is(err, syscall.ECONNRESET) || + errors.Is(err, syscall.ECONNABORTED) || + errors.Is(err, syscall.EPIPE) { + return true + } + var ne net.Error + if errors.As(err, &ne) && ne.Timeout() { + return true + } + // HTTP/2 stream and GOAWAY errors are not exported as stable types, and + // Windows socket errors (WSAECONNRESET) do not match the syscall constants + // above, so fall back to matching well-known error text. + msg := strings.ToLower(err.Error()) + for _, s := range []string{"stream error", "goaway", "connection reset", "forcibly closed", "broken pipe"} { + if strings.Contains(msg, s) { + return true + } + } + return false +} + +// retryBudget tracks forward progress across failed attempts and decides when +// a download should be abandoned. Only a true stall exhausts it: attempts that +// make at least minAttemptProgress are always retried. +type retryBudget struct { + // highWater is the largest offset any attempt has reached since the + // download last restarted from byte 0. + highWater int64 + // noProgress counts consecutive failed attempts that did not raise + // highWater. + noProgress int + // lowProgress counts failed attempts that raised highWater by less than + // minAttemptProgress, including those that did not raise it at all. + lowProgress int + // peak is the largest offset reached before any restart from byte 0. + peak int64 + // stuckRestarts counts consecutive restarts from byte 0 whose preceding + // run did not get past peak. + stuckRestarts int +} + +// record accounts for a failed attempt that reached offset end. It returns a +// non-nil error when the budget is exhausted. +func (b *retryBudget) record(end int64) error { + advance := end - b.highWater + if advance > 0 { + b.highWater = end + b.noProgress = 0 } else { - // Get rid of the old file and download from start, resetting hash. - if err := f.Truncate(0); err != nil { - return err + advance = 0 + b.noProgress++ + } + if advance < minAttemptProgress { + b.lowProgress++ + } + if b.noProgress > maxNoProgressRetries { + return fmt.Errorf("retries exhausted after %d consecutive attempts without progress", b.noProgress) + } + if b.lowProgress >= maxLowProgressAttempts { + return fmt.Errorf("retries exhausted after %d attempts that each made less than %s of progress", b.lowProgress, humanize.IBytes(minAttemptProgress)) + } + return nil +} + +// rebase resets progress tracking after the partial file was discarded. A +// restart from byte 0 is a fresh download, so offsets reached before the +// restart must not make later attempts look like they made no progress. +// +// A source can force restarts indefinitely, for example by rejecting every +// resume request or by ignoring Range and failing at the same byte each time. +// rebase therefore returns a non-nil error once more than +// maxNoProgressRetries consecutive restarts follow runs that never got past +// the largest offset reached before an earlier restart. A run that does get +// past it resets that count, so a download that keeps making progress is +// never abandoned. +func (b *retryBudget) rebase() error { + if b.highWater > b.peak { + b.peak = b.highWater + b.stuckRestarts = 0 + } else { + b.stuckRestarts++ + } + b.highWater = 0 + b.noProgress = 0 + b.lowProgress = 0 + if b.stuckRestarts > maxNoProgressRetries { + return fmt.Errorf("retries exhausted after %d consecutive restarts that never got past byte %d", b.stuckRestarts, b.peak) + } + return nil +} + +// rehash computes the SHA256 hash and size of the current contents of f. +func rehash(f *os.File) (hash.Hash, int64, error) { + if _, err := f.Seek(0, io.SeekStart); err != nil { + return nil, 0, err + } + h := sha256.New() + size, err := io.Copy(h, f) + if err != nil { + return nil, 0, err + } + return h, size, nil +} + +// discardPartial truncates f so that the next attempt starts from the +// beginning. +func discardPartial(f *os.File) error { + if err := f.Truncate(0); err != nil { + return &writeError{err: err} + } + return nil +} + +// alignToStream positions f and h for a stream that begins at start, given +// that f holds size bytes. A source that ignored the resume request restarts +// at 0, so the partial file is discarded and h is reset. +func alignToStream(f *os.File, h hash.Hash, name string, size, start int64) error { + switch { + case start == size: + if start > 0 { + logger.Infof("resuming download of %s at byte %d", name, start) } - if _, err := f.Seek(0, 0); err != nil { + case start == 0: + logger.Infof("server did not resume download of %s, restarting from start", name) + if err := discardPartial(f); err != nil { return err } - hash.Reset() + h.Reset() + default: + return fmt.Errorf("%w: source resumed at offset %d, want %d or 0", errResumeRejected, start, size) } - resp, err := downloader.HTTPClient.Do(req) + if _, err := f.Seek(start, io.SeekStart); err != nil { + return &writeError{err: err} + } + return nil +} + +// streamAttempt makes one attempt to copy the object from open into f and h, +// resuming at size, the number of bytes already in f. It returns the offset +// the attempt started at, the offset it reached, and the error that ended it, +// which is nil if the stream completed. +func streamAttempt(ctx context.Context, f *os.File, h hash.Hash, name string, size int64, stallTimeout time.Duration, open opener) (int64, int64, error) { + // A per-attempt context lets the StallReader abort a stalled stream + // without canceling the parent context. + attemptCtx, cancel := context.WithCancel(ctx) + defer cancel() + + body, start, err := open(attemptCtx, size) if err != nil { - return err + return size, size, err } - defer resp.Body.Close() - if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent { - return fmt.Errorf("downloading %s: %v", url, err) + sr := client.NewStallReader(body, stallTimeout, cancel) + defer sr.Close() + + if err := alignToStream(f, h, name, size, start); err != nil { + return size, size, err } - // Continue hashing the file as we download it. - n, err := io.Copy(io.MultiWriter(hash, f), resp.Body) - if err != nil { - return fmt.Errorf("downloading %s: %v", url, err) + n, err := io.Copy(io.MultiWriter(fileWriter{w: f}, h), sr) + return start, start + n, err +} + +// verify flushes a completed download in f and checks its hash h against +// chksum. It returns a *checksumError on a mismatch. +func verify(f *os.File, h hash.Hash, dst, chksum string) error { + if err := f.Sync(); err != nil { + logger.Warningf("syncing %s: %v", dst, err) } - // Verify the checksum of the fully downloaded file. - if sum := hex.EncodeToString(hash.Sum(nil)); sum != chksum { - os.RemoveAll(dst) // delete the bad file - return fmt.Errorf("checksum doesn't match: got %s, want %s", sum, chksum) + if sum := hex.EncodeToString(h.Sum(nil)); sum != chksum { + return &checksumError{got: sum, want: chksum} } - logger.Infof("Successfully downloaded %s bytes", humanize.IBytes(uint64(n))) return nil } -// Downloads a package from Google Cloud Storage -func packageGCS(ctx context.Context, bucket, object string, dst, chksum string) error { - client, err := storage.NewClient(ctx) +// fetch downloads the object provided by open into dst, verifying chksum. An +// existing partial dst is resumed when the source supports it. Transient +// failures are retried with exponential backoff until retryBudget is +// exhausted, so a download is only abandoned when it truly stops making +// progress. A source that rejects a resume offset causes the partial file to +// be discarded and the download restarted from the beginning, as does a +// checksum mismatch after a resumed attempt, at most once. +func fetch(ctx context.Context, name, dst, chksum string, stallTimeout time.Duration, open opener) error { + // Try to open any already existing file, otherwise create new file. + f, err := os.OpenFile(dst, os.O_CREATE|os.O_RDWR, 0644) if err != nil { return err } - defer client.Close() + defer f.Close() + + var budget *retryBudget + restartedAfterMismatch := false + for attempt := 1; ; attempt++ { + if err := ctx.Err(); err != nil { + return fmt.Errorf("downloading %s: %w", name, err) + } + + // Re-synchronize hash and size with the file on disk. + h, size, err := rehash(f) + if err != nil { + return err + } + if budget == nil { + budget = &retryBudget{highWater: size} + } + + // If the file checksum matches what we expect, then the file is already + // downloaded and we can quit early. + if sum := hex.EncodeToString(h.Sum(nil)); sum == chksum { + logger.Infof("using existing file: %s (sum = %s)", dst, sum) + return f.Close() + } + if attempt > 1 { + logger.Infof("retrying download of %s (attempt %d, %d bytes on disk)", name, attempt, size) + } else if size > 0 { + logger.Infof("existing file size: %d", size) + } + + start, end, attemptErr := streamAttempt(ctx, f, h, name, size, stallTimeout, open) + if attemptErr == nil { + err := verify(f, h, dst, chksum) + var ce *checksumError + if errors.As(err, &ce) && start > 0 && !restartedAfterMismatch { + // The bytes kept from earlier attempts may belong to a different + // version of the object, so try once more from the beginning. + logger.Warningf("download of %s resumed at byte %d failed verification (%v), restarting from start", name, start, err) + if err := discardPartial(f); err != nil { + return err + } + if rerr := budget.rebase(); rerr != nil { + f.Close() + os.RemoveAll(dst) // Delete the bad file. + return fmt.Errorf("downloading %s: %v: %w", name, rerr, err) + } + restartedAfterMismatch = true + continue + } + if err != nil { + f.Close() + os.RemoveAll(dst) // Delete the bad file. + return err + } + logger.Infof("Successfully downloaded %s bytes", humanize.IBytes(uint64(end))) + return f.Close() + } + + // Flush any partially written data to disk so the next attempt can resume. + if err := f.Sync(); err != nil { + logger.Warningf("syncing partial download %s: %v", dst, err) + } + if ctx.Err() != nil { + return fmt.Errorf("downloading %s: %w (last error: %v)", name, ctx.Err(), attemptErr) + } + restarted := false + switch { + case errors.Is(attemptErr, errResumeRejected): + logger.Warningf("discarding %d partial bytes of %s and restarting from start: %v", size, name, attemptErr) + if err := discardPartial(f); err != nil { + return err + } + // A restart is not progress. + start, end = 0, 0 + restarted = true + case !isRetryable(ctx, attemptErr): + return fmt.Errorf("downloading %s: %w", name, attemptErr) + case start == 0 && size > 0: + // The source ignored the resume request, so alignToStream discarded + // the partial file and this attempt restarted from byte 0. + restarted = true + } + if restarted { + if err := budget.rebase(); err != nil { + return fmt.Errorf("downloading %s: %v: %w", name, err, attemptErr) + } + } + if err := budget.record(end); err != nil { + return fmt.Errorf("downloading %s: %v: %w", name, err, attemptErr) + } + + d := backoffDelay(budget.noProgress) + logger.Warningf("download of %s failed after %s this attempt: %v; retrying in %v", name, humanize.IBytes(uint64(end-start)), attemptErr, d.Round(time.Millisecond)) + if err := sleep(ctx, d); err != nil { + return fmt.Errorf("downloading %s: %w (last error: %v)", name, err, attemptErr) + } + } +} + +// checkContentRange verifies that v, a Content-Range header value of the form +// "bytes START-END/SIZE", starts at offset. Any other value is reported as an +// error wrapping errResumeRejected. +func checkContentRange(v string, offset int64) error { + spec, ok := strings.CutPrefix(v, "bytes ") + if !ok { + return fmt.Errorf("%w: invalid Content-Range %q", errResumeRejected, v) + } + first, _, ok := strings.Cut(spec, "-") + start, err := strconv.ParseInt(strings.TrimSpace(first), 10, 64) + if !ok || err != nil { + return fmt.Errorf("%w: invalid Content-Range %q", errResumeRejected, v) + } + if start != offset { + return fmt.Errorf("%w: Content-Range %q does not start at requested offset %d", errResumeRejected, v, offset) + } + return nil +} + +// httpOpener returns an opener that fetches pkgURL over HTTP(S), using a Range +// request to resume when the server supports it. +func httpOpener(pkgURL string, downloader *client.Downloader) opener { + return func(ctx context.Context, offset int64) (io.ReadCloser, int64, error) { + resume := false + if offset > 0 { + ok, length, err := downloader.CanResume(ctx, pkgURL) + if err != nil { + logger.Errorf("CanResume: %v", err) + } + resume = ok && offset < length + } + + req, err := downloader.NewRequest(ctx, http.MethodGet, pkgURL, nil) + if err != nil { + return nil, 0, err + } + if resume { + req.Header.Set("Range", fmt.Sprintf("bytes=%d-", offset)) + } + + resp, err := downloader.HTTPClient.Do(req) + if err != nil { + return nil, 0, err + } + switch { + case resp.StatusCode == http.StatusPartialContent && resume: + if err := checkContentRange(resp.Header.Get("Content-Range"), offset); err != nil { + resp.Body.Close() + return nil, 0, err + } + return resp.Body, offset, nil + case resp.StatusCode == http.StatusOK: + // A 200 OK means the full object is being sent from byte 0, even if a + // Range header was sent. + return resp.Body, 0, nil + case resp.StatusCode == http.StatusRequestedRangeNotSatisfiable && resume: + resp.Body.Close() + return nil, 0, fmt.Errorf("%w: server returned %s for offset %d", errResumeRejected, resp.Status, offset) + default: + resp.Body.Close() + return nil, 0, &statusError{code: resp.StatusCode, status: resp.Status} + } + } +} - r, err := client.Bucket(bucket).Object(object).NewReader(ctx) +// packageHTTP downloads a package from an HTTP(S) server. +func packageHTTP(ctx context.Context, pkgURL, dst, chksum string, downloader *client.Downloader) error { + return fetch(ctx, strings.TrimPrefix(pkgURL, "oauth-"), dst, chksum, stallTimeout(downloader), httpOpener(pkgURL, downloader)) +} + +// gcsRangeOpener returns an opener that reads a Google Cloud Storage object +// from an offset to its end through newRangeReader. A range error on a resumed +// read, which the JSON and XML APIs report as HTTP 416 when the offset is at or +// beyond the end of the object (for example, because it was replaced), is +// reported as an error wrapping errResumeRejected. +func gcsRangeOpener(newRangeReader func(ctx context.Context, offset int64) (io.ReadCloser, error)) opener { + return func(ctx context.Context, offset int64) (io.ReadCloser, int64, error) { + r, err := newRangeReader(ctx, offset) + if err != nil { + var ge *googleapi.Error + if offset > 0 && errors.As(err, &ge) && ge.Code == http.StatusRequestedRangeNotSatisfiable { + return nil, 0, fmt.Errorf("%w: %v", errResumeRejected, err) + } + return nil, 0, err + } + return r, offset, nil + } +} + +// newGCSOpener returns an opener for a Google Cloud Storage object and a +// function that releases its resources. It is a package variable so that tests +// can substitute a fake without contacting GCS. +var newGCSOpener = func(ctx context.Context, bucket, object string) (opener, func() error, error) { + c, err := storage.NewClient(ctx) + if err != nil { + return nil, nil, err + } + obj := c.Bucket(bucket).Object(object) + open := gcsRangeOpener(func(ctx context.Context, offset int64) (io.ReadCloser, error) { + r, err := obj.NewRangeReader(ctx, offset, -1) + if err != nil { + return nil, err + } + return r, nil + }) + return open, c.Close, nil +} + +// packageGCS downloads a package from Google Cloud Storage, aborting reads +// that receive no data for stallTimeout. +func packageGCS(ctx context.Context, bucket, object string, dst, chksum string, stallTimeout time.Duration) error { + open, closeFn, err := newGCSOpener(ctx, bucket, object) if err != nil { return err } - defer r.Close() + defer closeFn() - logger.Infof("Downloading gs://%s/%s", bucket, object) - return download(r, dst, chksum) + name := fmt.Sprintf("gs://%s/%s", bucket, object) + logger.Infof("Downloading %s", name) + return fetch(ctx, name, dst, chksum, stallTimeout, open) } // FromRepo downloads a package from a repo. It returns the path to the @@ -162,34 +638,6 @@ func Latest(ctx context.Context, name, dir string, rm client.RepoMap, archs []st return FromRepo(ctx, rs, repo, dir, downloader) } -func download(r io.Reader, dst, chksum string) (err error) { - f, err := oswrap.Create(dst) - if err != nil { - return err - } - defer func() { - if cErr := f.Close(); cErr != nil && err == nil { - err = cErr - } - }() - - hash := sha256.New() - tw := io.MultiWriter(f, hash) - - b, err := io.Copy(tw, r) - if err != nil { - return err - } - - if hex.EncodeToString(hash.Sum(nil)) != chksum { - fmt.Println(hex.EncodeToString(hash.Sum(nil)), chksum) - return errors.New("checksum of downloaded file does not match expected checksum") - } - - logger.Infof("Successfully downloaded %s", humanize.IBytes(uint64(b))) - return nil -} - // ExtractPkg takes a path to a package and extracts it to a directory based on the // package name, it returns the path to the extracted directory. func ExtractPkg(src string) (dst string, err error) { diff --git a/download/download_test.go b/download/download_test.go index dfc5664..b996f50 100644 --- a/download/download_test.go +++ b/download/download_test.go @@ -17,42 +17,69 @@ import ( "archive/tar" "bytes" "compress/gzip" + "context" + "crypto/sha256" + "errors" + "fmt" + "io" "io/ioutil" - "path" + "net/http" + "net/http/httptest" + "os" "path/filepath" + "sync" + "sync/atomic" + "syscall" "testing" + "time" - "github.com/google/googet/v2/goolib" + "github.com/google/googet/v2/client" "github.com/google/googet/v2/oswrap" "github.com/google/logger" + "google.golang.org/api/googleapi" ) +// realSleep is the production sleep implementation, captured before tests +// replace it with a fast fake. +var realSleep = sleep + +// sleepRecorder records requested backoff delays without sleeping. +type sleepRecorder struct { + mu sync.Mutex + delays []time.Duration +} + +func (r *sleepRecorder) sleep(ctx context.Context, d time.Duration) error { + r.mu.Lock() + r.delays = append(r.delays, d) + r.mu.Unlock() + return ctx.Err() +} + +func (r *sleepRecorder) get() []time.Duration { + r.mu.Lock() + defer r.mu.Unlock() + return append([]time.Duration(nil), r.delays...) +} + func init() { logger.Init("test", true, false, ioutil.Discard) + // Tests must not wait for real backoff delays. + sleep = func(ctx context.Context, _ time.Duration) error { return ctx.Err() } } -func TestDownload(t *testing.T) { - r := bytes.NewReader([]byte("some content")) - tempDir, err := ioutil.TempDir("", "") - if err != nil { - t.Fatalf("error creating temp directory: %v", err) - } - defer oswrap.RemoveAll(tempDir) - - chksum := goolib.Checksum(r) - if _, err := r.Seek(0, 0); err != nil { - t.Errorf("error seeking to front of reader: %v", err) - } - tempFile := path.Join(tempDir, "test") - if err := download(r, tempFile, chksum); err != nil { - t.Errorf("error downloading and checking checksum: %v", err) - } - if err := download(r, tempFile, "notachecksum"); err == nil { - t.Error("wanted but did not recieve checksum error") - } +// recordSleepsForTest replaces sleep with a recorder for the duration of t. +func recordSleepsForTest(t *testing.T) *sleepRecorder { + t.Helper() + r := &sleepRecorder{} + orig := sleep + sleep = r.sleep + t.Cleanup(func() { sleep = orig }) + return r } func TestExtractPkg(t *testing.T) { + t.Parallel() tempDir, err := ioutil.TempDir("", "") if err != nil { t.Fatalf("error creating temp directory: %v", err) @@ -104,6 +131,7 @@ func TestExtractPkg(t *testing.T) { } func TestExtractPkgPathTraversal(t *testing.T) { + t.Parallel() tempDir, err := ioutil.TempDir("", "") if err != nil { t.Fatalf("error creating temp directory: %v", err) @@ -144,3 +172,1930 @@ func TestExtractPkgPathTraversal(t *testing.T) { t.Fatal("error expected because of path traversal") } } + +func TestPackageHTTP_ExtendedTrickle(t *testing.T) { + // Verify that packageHTTP completes a download stream where chunks trickle continuously + // over a duration significantly exceeding the stallTimeout. + + numChunks := 15 + var payload []byte + for i := 0; i < numChunks; i++ { + payload = append(payload, []byte(fmt.Sprintf("chunk-%02d;", i))...) + } + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + flusher, ok := w.(http.Flusher) + chunkSize := len(payload) / numChunks + for i := 0; i < len(payload); i += chunkSize { + end := i + chunkSize + if end > len(payload) { + end = len(payload) + } + w.Write(payload[i:end]) + if ok { + flusher.Flush() + } + time.Sleep(15 * time.Millisecond) // 15ms is less than the 30ms timeout. + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 30 * time.Millisecond + + dst := filepath.Join(t.TempDir(), "extended_trickle.pkg") + start := time.Now() + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed for extended trickle: %v", err) + } + elapsed := time.Since(start) + + // Total duration is at least 15 * 15ms = 225ms, which is > 7x the 30ms stallTimeout. + if elapsed < 200*time.Millisecond { + t.Errorf("elapsed time %v < expected 200ms", elapsed) + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) + } +} + +func TestPackageHTTP_Normal(t *testing.T) { + t.Parallel() + // Verify that a standard HTTP download completes successfully. + payload := []byte("standard package download payload content") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload) + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + + dst := filepath.Join(t.TempDir(), "normal.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed for normal download: %v", err) + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) + } +} + +func TestPackageHTTP_SlowTrickle(t *testing.T) { + // Verify that slow trickle downloads complete without being aborted by stall timer. + + payload := []byte("01234567890123456789012345678901234567890123456789") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + flusher, ok := w.(http.Flusher) + chunkSize := 10 + for i := 0; i < len(payload); i += chunkSize { + end := i + chunkSize + if end > len(payload) { + end = len(payload) + } + w.Write(payload[i:end]) + if ok { + flusher.Flush() + } + time.Sleep(25 * time.Millisecond) + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 80 * time.Millisecond + + dst := filepath.Join(t.TempDir(), "trickle.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed for slow trickle: %v", err) + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) + } +} + +func TestPackageHTTP_StallAndRangeResume(t *testing.T) { + // Verify that a stalled stream triggers an HTTP Range resume request and completes. + + payload := []byte("0123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + var reqCount int + var rangeHeaders []string + var mu sync.Mutex + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + reqCount++ + rangeHdr := r.Header.Get("Range") + if rangeHdr != "" { + rangeHeaders = append(rangeHeaders, rangeHdr) + } + mu.Unlock() + + if r.Method == http.MethodHead { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + + flusher, ok := w.(http.Flusher) + + // On first GET request: send partial data (40 bytes) and stall. + if rangeHdr == "" { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload[:40]) + if ok { + flusher.Flush() + } + // Stall by waiting until client disconnects. + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + return + } + + // On Range request: verify range header and serve remainder. + var start int + if _, err := fmt.Sscanf(rangeHdr, "bytes=%d-", &start); err != nil { + http.Error(w, "invalid range", http.StatusBadRequest) + return + } + + w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload)-start)) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload[start:]) + if ok { + flusher.Flush() + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 50 * time.Millisecond + + dst := filepath.Join(t.TempDir(), "resume.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed on resume: %v", err) + } + + mu.Lock() + defer mu.Unlock() + if len(rangeHeaders) != 1 { + t.Errorf("expected 1 range request, got %d: %v", len(rangeHeaders), rangeHeaders) + } else if rangeHeaders[0] != "bytes=40-" { + t.Errorf("expected Range: bytes=40-, got %s", rangeHeaders[0]) + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) + } +} + +func TestPackageHTTP_MaxRetriesExhausted(t *testing.T) { + // Verify that download aborts with client.ErrDownloadStalled once the retry budget + // of consecutive attempts without forward progress is exhausted. + + payload := []byte("0123456789012345678901234567890123456789") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + var attempts int + var mu sync.Mutex + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + + mu.Lock() + attempts++ + mu.Unlock() + + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + flusher, ok := w.(http.Flusher) + w.Write([]byte("01234")) + if ok { + flusher.Flush() + } + // Stall indefinitely until context is cancelled. + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 30 * time.Millisecond + + dst := filepath.Join(t.TempDir(), "max_retries.pkg") + err = packageHTTP(context.Background(), ts.URL, dst, chksum, downloader) + if err == nil { + t.Fatal("expected error after retries exhausted, got nil") + } + if !errors.Is(err, client.ErrDownloadStalled) { + t.Errorf("expected error wrapping client.ErrDownloadStalled, got: %v", err) + } + + mu.Lock() + defer mu.Unlock() + // The first attempt makes progress (5 bytes). Each later attempt receives a + // 200 OK restart from byte 0 that never gets past byte 5. The first restart + // records byte 5 as the peak, and maxNoProgressRetries+1 more restarts that + // do not pass it exhaust the budget: 6 attempts in total. + if want := maxNoProgressRetries + 3; attempts != want { + t.Errorf("got %d attempts, want %d (1 initial, 1 restart that sets the peak, %d stuck restarts)", attempts, want, maxNoProgressRetries+1) + } +} + +func TestPackageHTTP_ServerIgnoresRangeAndReturns200OK(t *testing.T) { + // Verify that a 200 OK response on a range request resets file offset and hash. + + payload := []byte("0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + var reqCount int + var mu sync.Mutex + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + + mu.Lock() + reqCount++ + currentReq := reqCount + mu.Unlock() + + flusher, ok := w.(http.Flusher) + + if currentReq == 1 { + // First attempt: send partial data (20 bytes) and stall. + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload[:20]) + if ok { + flusher.Flush() + } + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + return + } + + // Second attempt: server ignores Range and returns 200 OK with entire body. + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload) + if ok { + flusher.Flush() + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 50 * time.Millisecond + + dst := filepath.Join(t.TempDir(), "ignore_range.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed when server returned 200 OK on range: %v", err) + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if len(got) != len(payload) { + t.Fatalf("expected file length %d, got %d (file was not truncated on 200 OK)", len(payload), len(got)) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) + } +} + +func TestPackageHTTP_NonStallErrorFailsImmediately(t *testing.T) { + t.Parallel() + // Verify that non-stall errors fail immediately on first attempt without retrying. + var attempts int + var mu sync.Mutex + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.WriteHeader(http.StatusNotFound) + return + } + mu.Lock() + attempts++ + mu.Unlock() + http.Error(w, "not found", http.StatusNotFound) + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + + dst := filepath.Join(t.TempDir(), "notfound.pkg") + err = packageHTTP(context.Background(), ts.URL, dst, "dummychksum", downloader) + if err == nil { + t.Fatal("expected error on 404, got nil") + } + + mu.Lock() + defer mu.Unlock() + if attempts != 1 { + t.Errorf("expected exactly 1 attempt on non-stall error, got %d", attempts) + } +} + +func TestPackageHTTP_MultiStallProgressiveResume(t *testing.T) { + // Verify that multiple successive stalls across retry attempts progressively resume. + + payload := []byte("0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789abcdefghijklmnopqrstuvwxyz") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + var reqCount int + var rangeHeaders []string + var mu sync.Mutex + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + + mu.Lock() + reqCount++ + currentAttempt := reqCount + rangeHdr := r.Header.Get("Range") + rangeHeaders = append(rangeHeaders, rangeHdr) + mu.Unlock() + + flusher, ok := w.(http.Flusher) + + switch currentAttempt { + case 1: + // Attempt 0: send 25 bytes and stall. + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload[:25]) + if ok { + flusher.Flush() + } + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + case 2: + // Attempt 1: expect Range: bytes=25-, send next 25 bytes (25..50) and stall. + w.Header().Set("Content-Range", fmt.Sprintf("bytes 25-%d/%d", len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload)-25)) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload[25:50]) + if ok { + flusher.Flush() + } + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + case 3: + // Attempt 2: expect Range: bytes=50-, send next 25 bytes (50..75) and stall. + w.Header().Set("Content-Range", fmt.Sprintf("bytes 50-%d/%d", len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload)-50)) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload[50:75]) + if ok { + flusher.Flush() + } + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + case 4: + // Attempt 3: expect Range: bytes=75-, send final 25 bytes (75..100) and complete cleanly. + w.Header().Set("Content-Range", fmt.Sprintf("bytes 75-%d/%d", len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload)-75)) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload[75:]) + if ok { + flusher.Flush() + } + default: + t.Errorf("unexpected attempt %d", currentAttempt) + http.Error(w, "too many attempts", http.StatusInternalServerError) + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 50 * time.Millisecond + + dst := filepath.Join(t.TempDir(), "multi_stall.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed on progressive multi-stall resume: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + expectedRanges := []string{"", "bytes=25-", "bytes=50-", "bytes=75-"} + if len(rangeHeaders) != len(expectedRanges) { + t.Fatalf("expected %d requests, got %d: %v", len(expectedRanges), len(rangeHeaders), rangeHeaders) + } + for i, want := range expectedRanges { + if rangeHeaders[i] != want { + t.Errorf("request %d: expected Range %q, got %q", i, want, rangeHeaders[i]) + } + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("final content mismatch: got %q, want %q", string(got), string(payload)) + } +} + +func TestPackageHTTP_200OKFallback_WithSubsequentStall(t *testing.T) { + // Verify that 200 OK fallback truncates and recovers even if it stalls later. + + payload := []byte("0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + var reqCount int + var rangeHeaders []string + var mu sync.Mutex + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + + mu.Lock() + reqCount++ + currentReq := reqCount + rangeHeaders = append(rangeHeaders, r.Header.Get("Range")) + mu.Unlock() + + flusher, ok := w.(http.Flusher) + + switch currentReq { + case 1: + // Attempt 0: send 20 bytes and stall. + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload[:20]) + if ok { + flusher.Flush() + } + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + case 2: + // Attempt 1: server returns 200 OK (ignores Range: bytes=20-), sends 30 bytes from start, then stalls. + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload[:30]) + if ok { + flusher.Flush() + } + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + case 3: + // Attempt 2: client resumes Range: bytes=30-, server returns 206 Partial Content with rest. + w.Header().Set("Content-Range", fmt.Sprintf("bytes 30-%d/%d", len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload)-30)) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload[30:]) + if ok { + flusher.Flush() + } + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 50 * time.Millisecond + + dst := filepath.Join(t.TempDir(), "200_stall_resume.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + expectedRanges := []string{"", "bytes=20-", "bytes=30-"} + if len(rangeHeaders) != len(expectedRanges) { + t.Fatalf("expected %d requests, got %d: %v", len(expectedRanges), len(rangeHeaders), rangeHeaders) + } + for i, want := range expectedRanges { + if rangeHeaders[i] != want { + t.Errorf("request %d: expected Range %q, got %q", i, want, rangeHeaders[i]) + } + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if len(got) != len(payload) { + t.Fatalf("expected file length %d, got %d", len(payload), len(got)) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) + } +} + +func TestPackageHTTP_NonStallErrors_Comprehensive(t *testing.T) { + t.Parallel() + // Subtest 1: HTTP 500 errors are retried until the no-progress budget is exhausted. + t.Run("HTTP500_RetriedUntilBudgetExhausted", func(t *testing.T) { + var attempts atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.WriteHeader(http.StatusOK) + return + } + attempts.Add(1) + http.Error(w, "internal server error", http.StatusInternalServerError) + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "500.pkg") + err := packageHTTP(context.Background(), ts.URL, dst, "dummy", downloader) + if err == nil { + t.Fatal("expected error on 500, got nil") + } + if got := attempts.Load(); got != maxNoProgressRetries+1 { + t.Errorf("expected %d attempts, got %d", maxNoProgressRetries+1, got) + } + }) + + // Subtest 2: Checksum mismatch fails immediately and cleans up file from disk. + t.Run("ChecksumMismatch_DeletesFile", func(t *testing.T) { + var attempts int + payload := []byte("actual payload content from server") + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + w.WriteHeader(http.StatusOK) + return + } + attempts++ + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload) + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "corrupt.pkg") + err := packageHTTP(context.Background(), ts.URL, dst, "badchecksum00000000000000000000000000000000000000000000000000000000", downloader) + if err == nil { + t.Fatal("expected checksum mismatch error, got nil") + } + if attempts != 1 { + t.Errorf("expected exactly 1 attempt, got %d", attempts) + } + if _, statErr := os.Stat(dst); !os.IsNotExist(statErr) { + t.Errorf("expected destination file to be deleted on checksum mismatch, stat error: %v", statErr) + } + }) + + // Subtest 3: Context cancellation aborts immediately with zero retries. + t.Run("ContextCanceled_ZeroRetries", func(t *testing.T) { + var attempts int + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempts++ + w.WriteHeader(http.StatusOK) + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "canceled.pkg") + err := packageHTTP(ctx, ts.URL, dst, "dummy", downloader) + if err == nil { + t.Fatal("expected error on canceled context, got nil") + } + if attempts > 1 { + t.Errorf("expected at most 1 attempt on canceled context, got %d", attempts) + } + }) +} + +func TestPackageHTTP_ExistingFileSkipped(t *testing.T) { + t.Parallel() + // Verify that existing file with matching checksum skips all GET requests. + payload := []byte("already completely downloaded content") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + dst := filepath.Join(t.TempDir(), "existing.pkg") + if err := os.WriteFile(dst, payload, 0644); err != nil { + t.Fatalf("failed to write existing file: %v", err) + } + + var getRequests int + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + getRequests++ + } + w.WriteHeader(http.StatusOK) + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed on existing file: %v", err) + } + + if getRequests != 0 { + t.Errorf("expected 0 GET requests when file already exists with matching checksum, got %d", getRequests) + } +} + +func TestPackageHTTP_NoResumeSupportFallback(t *testing.T) { + // Verify fallback when server does not support range requests. + + payload := []byte("server that does not support range requests at all") + chksum := fmt.Sprintf("%x", sha256.Sum256(payload)) + + var reqCount int + var rangeHeaders []string + var mu sync.Mutex + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + // No Accept-Ranges header. + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + return + } + + mu.Lock() + reqCount++ + currentReq := reqCount + rangeHeaders = append(rangeHeaders, r.Header.Get("Range")) + mu.Unlock() + + flusher, ok := w.(http.Flusher) + + if currentReq == 1 { + // First attempt: send 15 bytes and stall. + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload[:15]) + if ok { + flusher.Flush() + } + select { + case <-time.After(500 * time.Millisecond): + case <-r.Context().Done(): + } + return + } + + // Second attempt: send full file from start. + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload) + if ok { + flusher.Flush() + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 50 * time.Millisecond + + dst := filepath.Join(t.TempDir(), "no_resume.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed when server does not support resume: %v", err) + } + + mu.Lock() + defer mu.Unlock() + + for i, hdr := range rangeHeaders { + if hdr != "" { + t.Errorf("request %d should have empty Range header, got %q", i, hdr) + } + } + + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) + } +} + +// testPayload returns a deterministic payload of n bytes and its SHA256 checksum. +func testPayload(n int) ([]byte, string) { + p := make([]byte, n) + for i := range p { + p[i] = byte('a' + i%26) + } + return p, fmt.Sprintf("%x", sha256.Sum256(p)) +} + +// rangeStart parses the start offset from a "bytes=N-" Range header, returning +// 0 when the header is absent. +func rangeStart(t *testing.T, hdr string) int { + t.Helper() + if hdr == "" { + return 0 + } + var start int + if _, err := fmt.Sscanf(hdr, "bytes=%d-", &start); err != nil { + t.Errorf("invalid Range header %q: %v", hdr, err) + } + return start +} + +// writeHead answers a HEAD request advertising Range support for size bytes. +func writeHead(w http.ResponseWriter, size int) { + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", size)) + w.WriteHeader(http.StatusOK) +} + +// writePartialThenStall sends the payload from start as a 200 or 206 response, +// flushes chunk bytes, and then stalls until the client disconnects. +func writePartialThenStall(w http.ResponseWriter, r *http.Request, payload []byte, start, chunk int) { + if start > 0 { + w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload)-start)) + w.WriteHeader(http.StatusPartialContent) + } else { + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + } + end := start + chunk + if end > len(payload) { + end = len(payload) + } + w.Write(payload[start:end]) + if f, ok := w.(http.Flusher); ok { + f.Flush() + } + if end == len(payload) { + return + } + select { + case <-time.After(5 * time.Second): + case <-r.Context().Done(): + } +} + +// hijackAndReset writes a response header plus payload[start:start+chunk] on +// the raw connection and then closes it, simulating a mid-stream disconnect. +func hijackAndReset(t *testing.T, w http.ResponseWriter, payload []byte, start, chunk int) { + t.Helper() + hj, ok := w.(http.Hijacker) + if !ok { + t.Errorf("response writer does not support hijacking") + return + } + conn, buf, err := hj.Hijack() + if err != nil { + t.Errorf("Hijack: %v", err) + return + } + defer conn.Close() + if start > 0 { + fmt.Fprintf(buf, "HTTP/1.1 206 Partial Content\r\nContent-Range: bytes %d-%d/%d\r\nContent-Length: %d\r\n\r\n", start, len(payload)-1, len(payload), len(payload)-start) + } else { + fmt.Fprintf(buf, "HTTP/1.1 200 OK\r\nContent-Length: %d\r\n\r\n", len(payload)) + } + end := start + chunk + if end > len(payload) { + end = len(payload) + } + buf.Write(payload[start:end]) + buf.Flush() +} + +func TestPackageHTTP_ConnectionResetResumesWithRange(t *testing.T) { + t.Parallel() + // Verify that a mid-stream disconnect (unexpected EOF) is retried with a Range request. + payload, chksum := testPayload(100) + var mu sync.Mutex + var ranges []string + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + hdr := r.Header.Get("Range") + mu.Lock() + ranges = append(ranges, hdr) + n := len(ranges) + mu.Unlock() + if n == 1 { + hijackAndReset(t, w, payload, 0, 40) + return + } + writePartialThenStall(w, r, payload, rangeStart(t, hdr), len(payload)) + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + dst := filepath.Join(t.TempDir(), "reset.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed after connection reset: %v", err) + } + mu.Lock() + defer mu.Unlock() + want := []string{"", "bytes=40-"} + if len(ranges) != len(want) || ranges[0] != want[0] || ranges[1] != want[1] { + t.Errorf("Range headers = %q, want %q", ranges, want) + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch after reset resume") + } +} + +func TestPackageHTTP_FourStallsWithProgressSucceed(t *testing.T) { + // Verify that the retry budget resets after progress: 4 stalls separated by + // progress (more than the 3 consecutive no-progress retries) still succeed. + payload, chksum := testPayload(100) + var mu sync.Mutex + var ranges []string + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + hdr := r.Header.Get("Range") + mu.Lock() + ranges = append(ranges, hdr) + mu.Unlock() + writePartialThenStall(w, r, payload, rangeStart(t, hdr), 20) + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 30 * time.Millisecond + dst := filepath.Join(t.TempDir(), "four_stalls.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed despite progress between stalls: %v", err) + } + mu.Lock() + defer mu.Unlock() + want := []string{"", "bytes=20-", "bytes=40-", "bytes=60-", "bytes=80-"} + if fmt.Sprint(ranges) != fmt.Sprint(want) { + t.Errorf("Range headers = %q, want %q", ranges, want) + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch after multi-stall resume") + } +} + +func TestPackageHTTP_ConsecutiveZeroByteStallsFail(t *testing.T) { + // Verify that 4 consecutive attempts that deliver zero bytes fail with + // client.ErrDownloadStalled, using exponential backoff between attempts. + rec := recordSleepsForTest(t) + payload, chksum := testPayload(50) + var attempts atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + attempts.Add(1) + writePartialThenStall(w, r, payload, 0, 0) + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + downloader.StallTimeout = 30 * time.Millisecond + dst := filepath.Join(t.TempDir(), "zero_stalls.pkg") + err = packageHTTP(context.Background(), ts.URL, dst, chksum, downloader) + if !errors.Is(err, client.ErrDownloadStalled) { + t.Fatalf("packageHTTP() = %v, want error wrapping client.ErrDownloadStalled", err) + } + if got := attempts.Load(); got != 4 { + t.Errorf("got %d GET attempts, want 4", got) + } + delays := rec.get() + wantBase := []time.Duration{time.Second, 2 * time.Second, 4 * time.Second} + if len(delays) != len(wantBase) { + t.Fatalf("backoff delays = %v, want %d delays", delays, len(wantBase)) + } + for i, d := range delays { + lo := time.Duration(float64(wantBase[i]) * (1 - backoffJitter)) + hi := time.Duration(float64(wantBase[i]) * (1 + backoffJitter)) + if d < lo || d > hi { + t.Errorf("delay %d = %v, want within [%v, %v]", i, d, lo, hi) + } + } +} + +func TestPackageHTTP_503ThenSuccess(t *testing.T) { + t.Parallel() + // Verify that a 503 response is retried and a later 200 succeeds. + payload, chksum := testPayload(64) + var attempts atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + if attempts.Add(1) == 1 { + http.Error(w, "unavailable", http.StatusServiceUnavailable) + return + } + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload) + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + dst := filepath.Join(t.TempDir(), "503.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed after 503: %v", err) + } + if got := attempts.Load(); got != 2 { + t.Errorf("got %d GET attempts, want 2", got) + } +} + +func TestPackageHTTP_429IsRetried(t *testing.T) { + t.Parallel() + // Verify that a 429 response is retried. + payload, chksum := testPayload(32) + var attempts atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if attempts.Add(1) == 1 { + http.Error(w, "slow down", http.StatusTooManyRequests) + return + } + w.WriteHeader(http.StatusOK) + w.Write(payload) + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "429.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP failed after 429: %v", err) + } + if got := attempts.Load(); got != 2 { + t.Errorf("got %d attempts, want 2", got) + } +} + +func TestPackageHTTP_ParentCancelMidStreamNotRetried(t *testing.T) { + t.Parallel() + // Verify that canceling the parent context mid-stream aborts without retrying. + payload, chksum := testPayload(100) + var attempts atomic.Int32 + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + attempts.Add(1) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) + w.Write(payload[:10]) + w.(http.Flusher).Flush() + cancel() + <-r.Context().Done() + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "cancel.pkg") + err := packageHTTP(ctx, ts.URL, dst, chksum, downloader) + if !errors.Is(err, context.Canceled) { + t.Fatalf("packageHTTP() = %v, want error wrapping context.Canceled", err) + } + if got := attempts.Load(); got != 1 { + t.Errorf("got %d attempts, want 1", got) + } +} + +func TestPackageHTTP_LowProgressAttemptCap(t *testing.T) { + // Verify that a source making 1 byte of progress per attempt stops after + // maxLowProgressAttempts attempts. + t.Parallel() + payload, chksum := testPayload(maxLowProgressAttempts + 10) + var attempts atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + attempts.Add(1) + hijackAndReset(t, w, payload, rangeStart(t, r.Header.Get("Range")), 1) + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "cap.pkg") + err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader) + if !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("packageHTTP() = %v, want error wrapping io.ErrUnexpectedEOF", err) + } + if got := attempts.Load(); got != maxLowProgressAttempts { + t.Errorf("got %d attempts, want %d", got, maxLowProgressAttempts) + } +} + +func TestPackageHTTP_ProgressingResetsNeverHitCap(t *testing.T) { + // Verify that 25 connection resets, each after at least minAttemptProgress + // bytes, do not count toward maxLowProgressAttempts and the download + // completes. + t.Parallel() + const resets = 25 + payload, chksum := testPayload(resets*minAttemptProgress + 100) + var attempts atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + attempts.Add(1) + hijackAndReset(t, w, payload, rangeStart(t, r.Header.Get("Range")), minAttemptProgress) + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "progressing.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP() = %v, want nil", err) + } + if got := attempts.Load(); got != resets+1 { + t.Errorf("got %d attempts, want %d", got, resets+1) + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch after %d resets", resets) + } +} + +func TestRetryBudget(t *testing.T) { + t.Parallel() + tests := []struct { + name string + // ends lists the offset reached by each failed attempt. + ends []int64 + // wantExhaustedAt is the 1-based attempt at which the budget is + // exhausted, or 0 if it never is. + wantExhaustedAt int + }{ + { + name: "four consecutive zero-progress attempts", + ends: []int64{0, 0, 0, 0}, + wantExhaustedAt: 4, + }, + { + name: "progress resets the consecutive counter", + ends: []int64{0, 0, 0, 10, 10, 10, 20}, + wantExhaustedAt: 0, + }, + { + name: "twenty low-progress attempts", + ends: []int64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20}, + wantExhaustedAt: 20, + }, + { + name: "high-progress attempts never count", + ends: func() []int64 { + var ends []int64 + for i := int64(1); i <= 100; i++ { + ends = append(ends, i*minAttemptProgress) + } + return ends + }(), + wantExhaustedAt: 0, + }, + { + name: "low-progress attempts accumulate across high-progress ones", + ends: func() []int64 { + var ends []int64 + var off int64 + for i := 0; i < maxLowProgressAttempts; i++ { + off += minAttemptProgress + ends = append(ends, off) + off++ + ends = append(ends, off) + } + return ends + }(), + wantExhaustedAt: 2 * maxLowProgressAttempts, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + b := &retryBudget{} + got := 0 + for i, end := range tc.ends { + if err := b.record(end); err != nil { + got = i + 1 + break + } + } + if got != tc.wantExhaustedAt { + t.Errorf("budget exhausted at attempt %d, want %d", got, tc.wantExhaustedAt) + } + }) + } +} + +func TestCheckContentRange(t *testing.T) { + t.Parallel() + tests := []struct { + value string + offset int64 + wantErr bool + }{ + {"bytes 40-99/100", 40, false}, + {"bytes 40-99/*", 40, false}, + {"bytes 0-99/100", 40, true}, + {"bytes 41-99/100", 40, true}, + {"", 40, true}, + {"bytes */100", 40, true}, + {"items 40-99/100", 40, true}, + } + for _, tc := range tests { + err := checkContentRange(tc.value, tc.offset) + if (err != nil) != tc.wantErr { + t.Errorf("checkContentRange(%q, %d) = %v, want error: %v", tc.value, tc.offset, err, tc.wantErr) + } + if err != nil && !errors.Is(err, errResumeRejected) { + t.Errorf("checkContentRange(%q, %d) = %v, want error wrapping errResumeRejected", tc.value, tc.offset, err) + } + } +} + +func TestPackageHTTP_RejectedResumeRestartsFromScratch(t *testing.T) { + // Verify that a resume response that does not start at the requested + // offset, or a 416, discards the partial file and restarts from byte 0. + t.Parallel() + tests := []struct { + name string + // resume answers the Range request that follows the first reset. + resume func(w http.ResponseWriter, payload []byte) + }{ + { + name: "Content-Range starts at wrong offset", + resume: func(w http.ResponseWriter, payload []byte) { + w.Header().Set("Content-Range", fmt.Sprintf("bytes 0-%d/%d", len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload) + }, + }, + { + name: "Content-Range missing", + resume: func(w http.ResponseWriter, payload []byte) { + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload)-40)) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload[40:]) + }, + }, + { + name: "416 Range Not Satisfiable", + resume: func(w http.ResponseWriter, payload []byte) { + w.WriteHeader(http.StatusRequestedRangeNotSatisfiable) + }, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + payload, chksum := testPayload(100) + var mu sync.Mutex + var ranges []string + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + hdr := r.Header.Get("Range") + mu.Lock() + ranges = append(ranges, hdr) + n := len(ranges) + mu.Unlock() + switch n { + case 1: + hijackAndReset(t, w, payload, 0, 40) + case 2: + tc.resume(w, payload) + default: + writePartialThenStall(w, r, payload, rangeStart(t, hdr), len(payload)) + } + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "rejected.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP() = %v, want nil", err) + } + mu.Lock() + defer mu.Unlock() + if want := []string{"", "bytes=40-", ""}; fmt.Sprint(ranges) != fmt.Sprint(want) { + t.Errorf("Range headers = %q, want %q", ranges, want) + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch after restart") + } + }) + } +} + +func TestPackageHTTP_ChecksumMismatchAfterResumeRestartsOnce(t *testing.T) { + // Verify that a checksum mismatch after a resumed attempt restarts the + // download from byte 0 once, and fails if it happens again. + t.Parallel() + tests := []struct { + name string + // corrupt lists the requests whose first 40 bytes are corrupted. + corrupt map[int]bool + wantRanges []string + wantErr bool + }{ + { + name: "restart succeeds", + corrupt: map[int]bool{1: true}, + wantRanges: []string{"", "bytes=40-", "", "bytes=40-"}, + }, + { + name: "second mismatch fails", + corrupt: map[int]bool{1: true, 3: true}, + wantRanges: []string{"", "bytes=40-", "", "bytes=40-"}, + wantErr: true, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + payload, chksum := testPayload(100) + bad := append([]byte(nil), payload...) + for i := 0; i < 40; i++ { + bad[i] = '!' + } + var mu sync.Mutex + var ranges []string + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + hdr := r.Header.Get("Range") + mu.Lock() + ranges = append(ranges, hdr) + n := len(ranges) + mu.Unlock() + if tc.corrupt[n] { + hijackAndReset(t, w, bad, 0, 40) + return + } + start := rangeStart(t, hdr) + chunk := len(payload) + if start == 0 { + // Force a resume on the next request. + chunk = 40 + } + hijackAndReset(t, w, payload, start, chunk) + })) + defer ts.Close() + + downloader, _ := client.NewDownloader("") + dst := filepath.Join(t.TempDir(), "mismatch.pkg") + err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader) + var ce *checksumError + if tc.wantErr != errors.As(err, &ce) { + t.Fatalf("packageHTTP() = %v, want checksum error: %v", err, tc.wantErr) + } + if !tc.wantErr && err != nil { + t.Fatalf("packageHTTP() = %v, want nil", err) + } + mu.Lock() + defer mu.Unlock() + if fmt.Sprint(ranges) != fmt.Sprint(tc.wantRanges) { + t.Errorf("Range headers = %q, want %q", ranges, tc.wantRanges) + } + }) + } +} + +func TestRetryBudgetRebase(t *testing.T) { + t.Parallel() + t.Run("progress after a restart is not measured against the old offset", func(t *testing.T) { + b := &retryBudget{} + if err := b.record(600); err != nil { + t.Fatalf("record(600) = %v, want nil", err) + } + if err := b.rebase(); err != nil { + t.Fatalf("rebase() = %v, want nil", err) + } + for _, end := range []int64{0, 100, 200, 300, 400, 550} { + if err := b.record(end); err != nil { + t.Fatalf("record(%d) after rebase = %v, want nil", end, err) + } + } + }) + t.Run("rebase resets the low-progress count", func(t *testing.T) { + b := &retryBudget{} + for i := int64(1); i < maxLowProgressAttempts; i++ { + if err := b.record(i); err != nil { + t.Fatalf("record(%d) = %v, want nil", i, err) + } + } + if err := b.rebase(); err != nil { + t.Fatalf("rebase() = %v, want nil", err) + } + for i := int64(1); i < maxLowProgressAttempts; i++ { + if err := b.record(i); err != nil { + t.Fatalf("record(%d) after rebase = %v, want nil", i, err) + } + } + }) + t.Run("restarts that never pass the previous peak are bounded", func(t *testing.T) { + b := &retryBudget{} + restarts := 0 + for ; restarts < 10; restarts++ { + if err := b.record(50); err != nil { + t.Fatalf("record(50) = %v, want nil", err) + } + if err := b.rebase(); err != nil { + break + } + } + if want := maxNoProgressRetries + 1; restarts != want { + t.Errorf("rebase() failed after %d restarts, want %d", restarts, want) + } + }) + t.Run("restarts that pass the previous peak are never bounded", func(t *testing.T) { + b := &retryBudget{} + for i := int64(1); i <= 50; i++ { + if err := b.record(i * 10); err != nil { + t.Fatalf("record(%d) = %v, want nil", i*10, err) + } + if err := b.rebase(); err != nil { + t.Fatalf("rebase() after reaching %d = %v, want nil", i*10, err) + } + } + }) +} + +func TestPackageHTTP_ProgressAfterRestartIsNotAbandoned(t *testing.T) { + // Verify that after the partial file is discarded and the download + // restarts from byte 0, failed attempts that each make progress are + // retried even though none reaches the offset reached before the restart. + t.Parallel() + payload, chksum := testPayload(1000) + bad := append([]byte(nil), payload...) + for i := 0; i < 600; i++ { + bad[i] = '!' + } + tests := []struct { + name string + // first answers the first two requests, which reach byte 600 and then + // force a restart from byte 0. + first func(t *testing.T, w http.ResponseWriter, n int) + }{ + { + name: "bad Content-Range", + first: func(t *testing.T, w http.ResponseWriter, n int) { + if n == 1 { + hijackAndReset(t, w, payload, 0, 600) + return + } + w.Header().Set("Content-Range", fmt.Sprintf("bytes 0-%d/%d", len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload) + }, + }, + { + name: "checksum mismatch after resume", + first: func(t *testing.T, w http.ResponseWriter, n int) { + if n == 1 { + hijackAndReset(t, w, bad, 0, 600) + return + } + hijackAndReset(t, w, payload, 600, len(payload)) + }, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + var mu sync.Mutex + var ranges []string + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + hdr := r.Header.Get("Range") + mu.Lock() + ranges = append(ranges, hdr) + n := len(ranges) + mu.Unlock() + start := rangeStart(t, hdr) + switch { + case n <= 2: + tc.first(t, w, n) + case start < 400: + // Four attempts after the restart each advance by 100 bytes + // and then fail, all below byte 600. + hijackAndReset(t, w, payload, start, 100) + default: + writePartialThenStall(w, r, payload, start, len(payload)) + } + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + dst := filepath.Join(t.TempDir(), "restart_progress.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP() = %v, want nil", err) + } + mu.Lock() + defer mu.Unlock() + want := []string{"", "bytes=600-", "", "bytes=100-", "bytes=200-", "bytes=300-", "bytes=400-"} + if fmt.Sprint(ranges) != fmt.Sprint(want) { + t.Errorf("Range headers = %q, want %q", ranges, want) + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch after restart") + } + }) + } +} + +func TestPackageHTTP_RepeatedRestartsWithoutProgressFail(t *testing.T) { + // Verify that a source that forces a restart from byte 0 on every resume + // and never gets past the same byte is eventually abandoned. + t.Parallel() + tests := []struct { + name string + // resume answers a Range request. + resume func(t *testing.T, w http.ResponseWriter, payload []byte) + wantAttempts int32 + }{ + // A server that ignores Range is covered by + // TestPackageHTTP_MaxRetriesExhausted. + { + name: "server answers every resume with a bad Content-Range", + resume: func(t *testing.T, w http.ResponseWriter, payload []byte) { + w.Header().Set("Content-Range", fmt.Sprintf("bytes 0-%d/%d", len(payload)-1, len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusPartialContent) + w.Write(payload) + }, + // Each restart takes one rejected resume and one attempt from + // byte 0, and the first restart sets the peak. + wantAttempts: 2 * (maxNoProgressRetries + 2), + }, + { + name: "server rejects every resume", + resume: func(t *testing.T, w http.ResponseWriter, payload []byte) { + w.WriteHeader(http.StatusRequestedRangeNotSatisfiable) + }, + // Each restart takes one rejected resume and one attempt from + // byte 0, and the first restart sets the peak. + wantAttempts: 2 * (maxNoProgressRetries + 2), + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + payload, chksum := testPayload(100) + var attempts atomic.Int32 + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + attempts.Add(1) + if r.Header.Get("Range") == "" { + hijackAndReset(t, w, payload, 0, 50) + return + } + tc.resume(t, w, payload) + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + dst := filepath.Join(t.TempDir(), "restart_loop.pkg") + if err := packageHTTP(context.Background(), ts.URL, dst, chksum, downloader); err == nil { + t.Fatalf("packageHTTP() = nil, want error") + } + if got := attempts.Load(); got != tc.wantAttempts { + t.Errorf("got %d attempts, want %d", got, tc.wantAttempts) + } + }) + } +} + +// fakeGCSObject serves payload through gcsRangeOpener. Calls listed in +// stallCalls deliver chunk bytes and then block until the context is canceled. +// Calls listed in rangeErrCalls fail with the HTTP 416 error that GCS returns +// for an offset at or beyond the end of the object. +type fakeGCSObject struct { + payload []byte + chunk int + stallCalls map[int]bool + rangeErrCalls map[int]bool + openErr error + + mu sync.Mutex + offsets []int64 + closed bool +} + +// stallingReader returns data and then blocks until ctx is done. +type stallingReader struct { + ctx context.Context + data *bytes.Reader +} + +func (s *stallingReader) Read(p []byte) (int, error) { + if s.data.Len() > 0 { + return s.data.Read(p) + } + <-s.ctx.Done() + return 0, s.ctx.Err() +} + +func (s *stallingReader) Close() error { return nil } + +// newRangeReader mimics ObjectHandle.NewRangeReader with length -1. +func (f *fakeGCSObject) newRangeReader(ctx context.Context, offset int64) (io.ReadCloser, error) { + f.mu.Lock() + f.offsets = append(f.offsets, offset) + call := len(f.offsets) + f.mu.Unlock() + if f.openErr != nil { + return nil, f.openErr + } + if f.rangeErrCalls[call] { + return nil, &googleapi.Error{Code: http.StatusRequestedRangeNotSatisfiable, Message: "The requested range cannot be satisfied."} + } + if f.stallCalls[call] { + end := int(offset) + f.chunk + if end > len(f.payload) { + end = len(f.payload) + } + return &stallingReader{ctx: ctx, data: bytes.NewReader(f.payload[offset:end])}, nil + } + return io.NopCloser(bytes.NewReader(f.payload[offset:])), nil +} + +func (f *fakeGCSObject) newOpener(ctx context.Context, bucket, object string) (opener, func() error, error) { + closeFn := func() error { + f.mu.Lock() + f.closed = true + f.mu.Unlock() + return nil + } + return gcsRangeOpener(f.newRangeReader), closeFn, nil +} + +// useFakeGCS replaces newGCSOpener with f for the duration of t. +func useFakeGCS(t *testing.T, f *fakeGCSObject) { + t.Helper() + orig := newGCSOpener + newGCSOpener = f.newOpener + t.Cleanup(func() { newGCSOpener = orig }) +} + +func TestPackageGCS_StallResumesWithRangeReader(t *testing.T) { + // Verify that a stalled GCS read is retried and resumed at the partial offset. + payload, chksum := testPayload(90) + f := &fakeGCSObject{payload: payload, chunk: 30, stallCalls: map[int]bool{1: true, 2: true}} + useFakeGCS(t, f) + + dst := filepath.Join(t.TempDir(), "gcs.pkg") + downloader := &client.Downloader{StallTimeout: 30 * time.Millisecond} + if err := Package(context.Background(), "gs://bucket/obj.goo", dst, chksum, downloader); err != nil { + t.Fatalf("Package(gs://) failed: %v", err) + } + f.mu.Lock() + defer f.mu.Unlock() + if want := []int64{0, 30, 60}; fmt.Sprint(f.offsets) != fmt.Sprint(want) { + t.Errorf("GCS read offsets = %v, want %v", f.offsets, want) + } + if !f.closed { + t.Error("GCS client was not closed") + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch after GCS resume") + } +} + +func TestPackageGCS_RangeErrorRestartsFromScratch(t *testing.T) { + // Verify that a GCS range error on a resumed read discards the partial + // file and restarts from offset 0. + payload, chksum := testPayload(90) + f := &fakeGCSObject{payload: payload, chunk: 30, stallCalls: map[int]bool{1: true}, rangeErrCalls: map[int]bool{2: true}} + useFakeGCS(t, f) + + dst := filepath.Join(t.TempDir(), "gcs_range.pkg") + downloader := &client.Downloader{StallTimeout: 30 * time.Millisecond} + if err := Package(context.Background(), "gs://bucket/obj.goo", dst, chksum, downloader); err != nil { + t.Fatalf("Package(gs://) failed: %v", err) + } + f.mu.Lock() + defer f.mu.Unlock() + if want := []int64{0, 30, 0}; fmt.Sprint(f.offsets) != fmt.Sprint(want) { + t.Errorf("GCS read offsets = %v, want %v", f.offsets, want) + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch after GCS restart") + } +} + +func TestGCSRangeOpener(t *testing.T) { + // Verify that only a range error on a resumed read is mapped to + // errResumeRejected. + t.Parallel() + rangeErr := &googleapi.Error{Code: http.StatusRequestedRangeNotSatisfiable} + open := gcsRangeOpener(func(context.Context, int64) (io.ReadCloser, error) { return nil, rangeErr }) + if _, _, err := open(context.Background(), 10); !errors.Is(err, errResumeRejected) { + t.Errorf("open(offset 10) = %v, want error wrapping errResumeRejected", err) + } + if _, _, err := open(context.Background(), 0); errors.Is(err, errResumeRejected) || !errors.Is(err, rangeErr) { + t.Errorf("open(offset 0) = %v, want the original range error", err) + } +} + +func TestPackageGCS_ChecksumMismatchNotRetried(t *testing.T) { + // Verify that a GCS checksum mismatch fails once and deletes the file. + payload, _ := testPayload(40) + f := &fakeGCSObject{payload: payload} + useFakeGCS(t, f) + + dst := filepath.Join(t.TempDir(), "gcs_bad.pkg") + if err := Package(context.Background(), "gs://bucket/obj.goo", dst, "bad", nil); err == nil { + t.Fatal("Package(gs://) succeeded with bad checksum, want error") + } + f.mu.Lock() + defer f.mu.Unlock() + if len(f.offsets) != 1 { + t.Errorf("got %d GCS reads, want 1", len(f.offsets)) + } + if _, err := os.Stat(dst); !os.IsNotExist(err) { + t.Errorf("bad file not deleted: %v", err) + } +} + +func TestPackageGCS_NonRetryableOpenError(t *testing.T) { + // Verify that a non-transient GCS open error is returned without retrying. + f := &fakeGCSObject{openErr: errors.New("storage: object doesn't exist")} + useFakeGCS(t, f) + + dst := filepath.Join(t.TempDir(), "gcs_missing.pkg") + if err := Package(context.Background(), "gs://bucket/obj.goo", dst, "x", nil); err == nil { + t.Fatal("Package(gs://) succeeded, want error") + } + f.mu.Lock() + defer f.mu.Unlock() + if len(f.offsets) != 1 { + t.Errorf("got %d GCS opens, want 1", len(f.offsets)) + } +} + +// timeoutError is a net.Error that reports a timeout. +type timeoutError struct{} + +func (timeoutError) Error() string { return "i/o timeout" } +func (timeoutError) Timeout() bool { return true } +func (timeoutError) Temporary() bool { return true } + +func TestIsRetryable(t *testing.T) { + t.Parallel() + canceled, cancel := context.WithCancel(context.Background()) + cancel() + tests := []struct { + name string + ctx context.Context + err error + want bool + }{ + {"nil", context.Background(), nil, false}, + {"stalled", context.Background(), fmt.Errorf("x: %w", client.ErrDownloadStalled), true}, + {"unexpected EOF", context.Background(), io.ErrUnexpectedEOF, true}, + {"ECONNRESET", context.Background(), &os.SyscallError{Syscall: "read", Err: syscall.ECONNRESET}, true}, + {"ECONNABORTED", context.Background(), syscall.ECONNABORTED, true}, + {"EPIPE", context.Background(), syscall.EPIPE, true}, + {"net timeout", context.Background(), timeoutError{}, true}, + {"http2 stream error", context.Background(), errors.New("stream error: stream ID 3; INTERNAL_ERROR"), true}, + {"http2 GOAWAY", context.Background(), errors.New("http2: server sent GOAWAY and closed the connection"), true}, + {"windows reset", context.Background(), errors.New("wsarecv: An existing connection was forcibly closed by the remote host"), true}, + {"503", context.Background(), &statusError{code: 503, status: "503 Service Unavailable"}, true}, + {"429", context.Background(), &statusError{code: 429, status: "429 Too Many Requests"}, true}, + {"404", context.Background(), &statusError{code: 404, status: "404 Not Found"}, false}, + {"disk write EPIPE", context.Background(), &writeError{err: syscall.EPIPE}, false}, + {"disk full", context.Background(), &writeError{err: errors.New("no space left on device")}, false}, + {"parent canceled", canceled, client.ErrDownloadStalled, false}, + {"other", context.Background(), errors.New("x509: certificate signed by unknown authority"), false}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := isRetryable(tc.ctx, tc.err); got != tc.want { + t.Errorf("isRetryable(%v) = %v, want %v", tc.err, got, tc.want) + } + }) + } +} + +func TestBackoffDelay(t *testing.T) { + t.Parallel() + tests := []struct { + noProgress int + base time.Duration + }{ + {0, time.Second}, + {1, time.Second}, + {2, 2 * time.Second}, + {3, 4 * time.Second}, + {6, 30 * time.Second}, + {100, 30 * time.Second}, + } + for _, tc := range tests { + for i := 0; i < 20; i++ { + d := backoffDelay(tc.noProgress) + lo := time.Duration(float64(tc.base) * (1 - backoffJitter)) + hi := time.Duration(float64(tc.base) * (1 + backoffJitter)) + if d < lo || d > hi { + t.Fatalf("backoffDelay(%d) = %v, want within [%v, %v]", tc.noProgress, d, lo, hi) + } + } + } +} + +func TestSleepHonorsContext(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + start := time.Now() + if err := realSleep(ctx, time.Hour); !errors.Is(err, context.Canceled) { + t.Errorf("sleep(canceled ctx) = %v, want context.Canceled", err) + } + if time.Since(start) > time.Second { + t.Errorf("sleep did not return promptly on canceled context") + } + if err := realSleep(context.Background(), time.Millisecond); err != nil { + t.Errorf("sleep(1ms) = %v, want nil", err) + } +} diff --git a/googet.go b/googet.go index 49a7793..47bc7c4 100644 --- a/googet.go +++ b/googet.go @@ -21,8 +21,10 @@ import ( "fmt" "os" + "github.com/google/googet/v2/client" "github.com/google/googet/v2/googetdb" "github.com/google/googet/v2/settings" + "github.com/google/googet/v2/supervisor" "github.com/google/googet/v2/system" "github.com/google/logger" "github.com/google/subcommands" @@ -78,6 +80,22 @@ func main() { os.Exit(run(context.Background())) } +// configureWatchdogs applies googet.conf installer supervision and download stall settings. +func configureWatchdogs() { + supervisor.Configure(supervisor.Options{ + Mode: settings.SupervisorMode, + InactivityTimeout: settings.InactivityTimeout, + HardTimeout: settings.InstallTimeout, + UIGracePeriod: settings.UIGracePeriod, + DisableUIDetection: !settings.UIDetection, + Unattended: !settings.Confirm, + }) + client.SetDefaultStallTimeout(settings.DownloadStallTimeout) + d := supervisor.CurrentDefaults() + logger.Infof("Installer supervision: mode=%v inactivity=%v hard=%v ui_grace=%v ui_detection=%v unattended=%v", + d.Mode, d.InactivityTimeout, d.HardTimeout, d.UIGracePeriod, !d.DisableUIDetection, d.Unattended) +} + func run(ctx context.Context) int { rootDir := flag.String("root", os.Getenv(envVar), "googet root directory") noConfirm := flag.Bool("noconfirm", false, "skip confirmation") @@ -161,6 +179,8 @@ func run(ctx context.Context) int { logger.Init("GooGet", *verbose, *systemLog, lf) defer logger.Close() + configureWatchdogs() + if err := googetdb.CreateIfMissing(dbFile); err != nil { logger.Errorf("Unable to create initial db file; if db is not created, run again as admin: %v", err) return 1 diff --git a/googet.goospec b/googet.goospec index 123c732..1b890ec 100644 --- a/googet.goospec +++ b/googet.goospec @@ -1,4 +1,4 @@ -{{$version := "3.3.3@0" -}} +{{$version := "3.4.0@0" -}} { "name": "googet", "version": "{{$version}}", @@ -15,6 +15,11 @@ "path": "install.ps1" }, "releaseNotes": [ + "3.4.0 - Feat: Supervise installers in a Job Object; terminate hung (5m no progress), interactive (30s modal dialog without progress in unattended mode), or over-long (60m) installers. Configure via SupervisorMode/InactivityTimeout/InstallTimeout/UIGracePeriod/UIDetection in googet.conf or per package via ExecFile timeout/inactivityTimeout. Installers longer than 60m must set a timeout override.", + "3.4.0 - Change: Killing googet (e.g. an outer agent timeout) now also terminates the in-flight installer process tree (Job Object KILL_ON_JOB_CLOSE).", + "3.4.0 - Change: .msu (wusa) installs have no hard timeout unless the package sets one; CBS.log growth counts as progress.", + "3.4.0 - Feat: Download stall detection with automatic HTTP Range/GCS resume; header timeout for HTTP requests.", + "3.4.0 - Fix: Failed installs restore overwritten files and preserve installer logs; batch updates continue past unresolvable packages.", "3.3.3 - Feat: Add ExecutableNames and InstallPath fields to GooGet PkgSpec.", "3.3.2 - Feat: introduce StrictConflicts setting to optionally enforce file ownership conflict errors.", "3.3.1 - Fix: Prevent path traversal vulnerabilities in file unpacking and add workflow permissions.", diff --git a/goolib/goolib.go b/goolib/goolib.go index a5d6218..eb88557 100644 --- a/goolib/goolib.go +++ b/goolib/goolib.go @@ -18,6 +18,7 @@ import ( "compress/gzip" "crypto/sha256" "encoding/hex" + "errors" "fmt" "io" "os" @@ -28,6 +29,8 @@ import ( "slices" "strings" "syscall" + + "github.com/google/googet/v2/supervisor" ) var interpreter = map[string]string{ @@ -51,6 +54,11 @@ func scriptInterpreter(s string) (string, error) { // The process is successful if the exit code matches any of those provided or '0'. // stdout and stderr are sent to the writer. func Exec(s string, args []string, ec []int, w io.Writer) error { + return ExecWithOptions(s, args, ec, supervisor.Options{}, w) +} + +// ExecWithOptions executes a script or binary with custom supervisor options. +func ExecWithOptions(s string, args []string, ec []int, opts supervisor.Options, w io.Writer) error { var c *exec.Cmd switch runtime.GOOS { case "windows": @@ -74,27 +82,111 @@ func Exec(s string, args []string, ec []int, w io.Writer) error { default: return fmt.Errorf("OS %q is not Windows or Linux", runtime.GOOS) } - return Run(c, ec, w) + return RunWithOptions(c, ec, opts, w) +} + +// msiLogModifiers are the characters allowed after /l in msiexec logging switches (e.g. /l*v, /lv*x, /l+!). +const msiLogModifiers = "iwearucmopvx+!*" + +// isLogFlag reports whether the given token matches a known installer logging flag. +func isLogFlag(s string) bool { + ls := strings.ToLower(s) + switch ls { + case "/log", "-log", "--log": + return true + } + if len(ls) < 2 || (ls[0] != '/' && ls[0] != '-') || ls[1] != 'l' { + return false + } + for _, r := range ls[2:] { + if !strings.ContainsRune(msiLogModifiers, r) { + return false + } + } + return true +} + +// enrichOptions inspects command line arguments and the output writer to auto-discover +// log files whose growth indicates forward progress. Unattended mode is configured +// process-wide via supervisor.Configure and is not inferred here. +func enrichOptions(c *exec.Cmd, opts supervisor.Options, w io.Writer) supervisor.Options { + seen := make(map[string]bool) + for _, f := range opts.LogFiles { + seen[f] = true + } + // addLog appends path to opts.LogFiles unless it is empty or already listed. + addLog := func(path string) { + if path != "" && !seen[path] { + opts.LogFiles = append(opts.LogFiles, path) + seen[path] = true + } + } + + // 1. Inspect writer if it is a file. + if f, ok := w.(*os.File); ok && f != nil { + addLog(f.Name()) + } + + // 2. Inspect command arguments for log flags (space-separated or colon-delimited). + if c != nil { + for i := 0; i < len(c.Args); i++ { + arg := c.Args[i] + + // Check for colon-delimited log flags (e.g. /log:, /l*v:, -log:, -l:). + if parts := strings.SplitN(arg, ":", 2); len(parts) == 2 { + if isLogFlag(parts[0]) { + addLog(strings.Trim(strings.TrimSpace(parts[1]), `"'`)) + continue + } + } + + // Check for space-separated log flags (e.g. /log , /l*v , -log , -l ). + if isLogFlag(arg) && i+1 < len(c.Args) && !isLogFlag(c.Args[i+1]) { + addLog(strings.Trim(strings.TrimSpace(c.Args[i+1]), `"'`)) + i++ + } + } + } + + return opts } // Run runs a command. // The process is successful if the exit code matches any of those provided or '0'. // stdout and stderr are sent to the writer and to this process's stdout and stderr. func Run(c *exec.Cmd, ec []int, w io.Writer) error { - c.Stdout = io.MultiWriter(os.Stdout, w) - c.Stderr = io.MultiWriter(os.Stderr, w) - if err := c.Run(); err != nil { - e, ok := err.(*exec.ExitError) - if !ok { - return err - } - s, ok := e.Sys().(syscall.WaitStatus) - if !ok { - return err - } - if !slices.Contains(ec, s.ExitStatus()) { - return fmt.Errorf("command exited with error code %v", s.ExitStatus()) - } + return RunWithOptions(c, ec, supervisor.Options{}, w) +} + +// RunWithOptions runs a command supervised by the supervisor package using custom options. +func RunWithOptions(c *exec.Cmd, ec []int, opts supervisor.Options, w io.Writer) error { + opts = enrichOptions(c, opts, w) + return checkExit(supervisor.Run(c, opts, w), ec) +} + +// checkExit maps the result of supervisor.Run to the error returned by RunWithOptions. +// A nonzero exit status listed in ec is treated as success. +func checkExit(err error, ec []int) error { + if err == nil { + return nil + } + // Defense in depth: supervisor termination errors do not wrap an + // *exec.ExitError today, so errors.As below would already fail for them. + // This guard keeps a future termination error that does wrap one from + // being accepted as success because its exit code is listed in ec. + if errors.Is(err, supervisor.ErrTerminated) { + return err + } + var e *exec.ExitError + if !errors.As(err, &e) { + return err + } + s, ok := e.Sys().(syscall.WaitStatus) + if !ok { + return err + } + if !slices.Contains(ec, s.ExitStatus()) { + return fmt.Errorf("command exited with error code %v", s.ExitStatus()) } return nil } diff --git a/goolib/goolib_test.go b/goolib/goolib_test.go index 404dbda..b714124 100644 --- a/goolib/goolib_test.go +++ b/goolib/goolib_test.go @@ -16,9 +16,13 @@ package goolib import ( "fmt" "math/rand" + "os" + "os/exec" + "reflect" "strings" "testing" - "time" + + "github.com/google/googet/v2/supervisor" ) func TestScriptInterpreter(t *testing.T) { @@ -59,7 +63,6 @@ func randString(runes []rune, min, max int) string { } func TestSplitGCSUrl(t *testing.T) { - rand.Seed(time.Now().UnixNano()) const alphanum = "abcdefghijklmnopqrstuvwxyz0123456789" objChars := alphanum + "ABCDEFGHIJKLMNOPQRSTUVWXYZ-_.~@%^=+" bucket := randString([]rune(alphanum), 1, 1) + randString([]rune(alphanum+"-_."), 0, 61) + randString([]rune(alphanum), 1, 1) @@ -131,3 +134,329 @@ func TestSplitGCSUrl(t *testing.T) { } } } + +// TestEnrichOptionsIgnoresOSArgs verifies that unattended mode is never inferred from os.Args. +func TestEnrichOptionsIgnoresOSArgs(t *testing.T) { + oldArgs := os.Args + defer func() { os.Args = oldArgs }() + + for _, init := range []bool{false, true} { + os.Args = []string{"googet", "-noconfirm", "/noconfirm", "install", "noconfirm"} + opts := enrichOptions(exec.Command("cmd.exe", "/c", "echo"), supervisor.Options{Unattended: init}, nil) + if opts.Unattended != init { + t.Errorf("enrichOptions(Unattended=%v) = %v, want %v", init, opts.Unattended, init) + } + } +} + +// TestIsLogFlag verifies recognition of msiexec and common installer logging switches. +func TestIsLogFlag(t *testing.T) { + for _, s := range []string{"/log", "-LOG", "--log", "/l", "-l", "/l*v", "/L*V", "/lv*", "/l*vx", "/l+!", "/liwe"} { + if !isLogFlag(s) { + t.Errorf("isLogFlag(%q) = false, want true", s) + } + } + for _, s := range []string{"", "/", "l*v", "/lang", "/qn", "/i", "-lz", "/logs", "C:\\x.log"} { + if isLogFlag(s) { + t.Errorf("isLogFlag(%q) = true, want false", s) + } + } +} + +// TestLogFlagParsing verifies colon-delimited and space-separated log flag parsing. +func TestLogFlagParsing(t *testing.T) { + tests := []struct { + name string + cmdArgs []string + initLogs []string + wantLogs []string + }{ + { + name: "colon-delimited /log:", + cmdArgs: []string{"msiexec", "/i", "pkg.msi", `/log:C:\install.log`}, + wantLogs: []string{`C:\install.log`}, + }, + { + name: "colon-delimited /l*v:", + cmdArgs: []string{"msiexec", "/i", "pkg.msi", `/l*v:C:\msi.log`}, + wantLogs: []string{`C:\msi.log`}, + }, + { + name: "colon-delimited -log:", + cmdArgs: []string{"setup.exe", `-log:C:\boot.log`}, + wantLogs: []string{`C:\boot.log`}, + }, + { + name: "colon-delimited -l:", + cmdArgs: []string{"setup.exe", `-l:C:\app.log`}, + wantLogs: []string{`C:\app.log`}, + }, + { + name: "colon-delimited /l:", + cmdArgs: []string{"msiexec", "/i", "pkg.msi", `/l:C:\quick.log`}, + wantLogs: []string{`C:\quick.log`}, + }, + { + name: "colon-delimited --log:", + cmdArgs: []string{"wix.exe", `--log:C:\wix.log`}, + wantLogs: []string{`C:\wix.log`}, + }, + { + name: "colon-delimited /l*vx:", + cmdArgs: []string{"msiexec", `/l*vx:C:\verbose.log`}, + wantLogs: []string{`C:\verbose.log`}, + }, + { + name: "colon-delimited -l*vx:", + cmdArgs: []string{"msiexec", `-l*vx:C:\verbose2.log`}, + wantLogs: []string{`C:\verbose2.log`}, + }, + { + name: "colon-delimited -l*v:", + cmdArgs: []string{"msiexec", `-l*v:C:\verbose3.log`}, + wantLogs: []string{`C:\verbose3.log`}, + }, + { + name: "colon-delimited double-quoted path", + cmdArgs: []string{"msiexec", `/log:"C:\Program Files\App\install.log"`}, + wantLogs: []string{`C:\Program Files\App\install.log`}, + }, + { + name: "colon-delimited single-quoted path", + cmdArgs: []string{"msiexec", `/l*v:'C:\Logs\test.log'`}, + wantLogs: []string{`C:\Logs\test.log`}, + }, + { + name: "colon-delimited uppercase /LOG with original path casing preserved", + cmdArgs: []string{"msiexec", `/LOG:C:\MyFolder\Install.LOG`}, + wantLogs: []string{`C:\MyFolder\Install.LOG`}, + }, + { + name: "colon-delimited empty path ignored", + cmdArgs: []string{"msiexec", "/log:"}, + wantLogs: nil, + }, + { + name: "colon-delimited quoted empty path ignored", + cmdArgs: []string{"msiexec", `/log:""`}, + wantLogs: nil, + }, + { + name: "colon-delimited with spaces around quotes", + cmdArgs: []string{"msiexec", `/log: "C:\Logs\install.log" `}, + wantLogs: []string{`C:\Logs\install.log`}, + }, + { + name: "colon-delimited with multiple colons in path", + cmdArgs: []string{"setup.exe", `/log:C:\foo:bar\baz.log`}, + wantLogs: []string{`C:\foo:bar\baz.log`}, + }, + { + name: "space-separated /log ", + cmdArgs: []string{"msiexec", "/i", "pkg.msi", "/log", `C:\install.log`}, + wantLogs: []string{`C:\install.log`}, + }, + { + name: "space-separated /l*v ", + cmdArgs: []string{"msiexec", "/i", "pkg.msi", "/l*v", `C:\msi.log`}, + wantLogs: []string{`C:\msi.log`}, + }, + { + name: "space-separated -log ", + cmdArgs: []string{"setup.exe", "-log", `C:\boot.log`}, + wantLogs: []string{`C:\boot.log`}, + }, + { + name: "space-separated -l ", + cmdArgs: []string{"setup.exe", "-l", `C:\app.log`}, + wantLogs: []string{`C:\app.log`}, + }, + { + name: "space-separated /l ", + cmdArgs: []string{"msiexec", "/l", `C:\quick.log`}, + wantLogs: []string{`C:\quick.log`}, + }, + { + name: "space-separated --log ", + cmdArgs: []string{"wix.exe", "--log", `C:\wix.log`}, + wantLogs: []string{`C:\wix.log`}, + }, + { + name: "space-separated /l*vx ", + cmdArgs: []string{"msiexec", "/l*vx", `C:\verbose.log`}, + wantLogs: []string{`C:\verbose.log`}, + }, + { + name: "space-separated -l*vx ", + cmdArgs: []string{"msiexec", "-l*vx", `C:\verbose2.log`}, + wantLogs: []string{`C:\verbose2.log`}, + }, + { + name: "space-separated -l*v ", + cmdArgs: []string{"msiexec", "-l*v", `C:\verbose3.log`}, + wantLogs: []string{`C:\verbose3.log`}, + }, + { + name: "space-separated double-quoted path", + cmdArgs: []string{"msiexec", "/log", `"C:\Program Files\App\install.log"`}, + wantLogs: []string{`C:\Program Files\App\install.log`}, + }, + { + name: "space-separated single-quoted path", + cmdArgs: []string{"msiexec", "/l*v", `'C:\Logs\test.log'`}, + wantLogs: []string{`C:\Logs\test.log`}, + }, + { + name: "space-separated uppercase /LOG", + cmdArgs: []string{"setup.exe", "/LOG", `C:\install.log`}, + wantLogs: []string{`C:\install.log`}, + }, + { + name: "space-separated uppercase /L*V", + cmdArgs: []string{"msiexec", "/L*V", `C:\msi.log`}, + wantLogs: []string{`C:\msi.log`}, + }, + { + name: "space-separated trailing flag without path", + cmdArgs: []string{"msiexec", "/log"}, + wantLogs: nil, + }, + { + name: "space-separated with empty string value", + cmdArgs: []string{"msiexec", "/log", ""}, + wantLogs: nil, + }, + { + name: "space-separated with quoted empty value", + cmdArgs: []string{"msiexec", "/log", `""`}, + wantLogs: nil, + }, + { + name: "multiple mixed log flags", + cmdArgs: []string{"installer.exe", `/l*v:C:\first.log`, "-log", `C:\second.log`}, + wantLogs: []string{`C:\first.log`, `C:\second.log`}, + }, + { + name: "deduplicate identical log files", + cmdArgs: []string{"installer.exe", `/log:C:\dup.log`, "/log", `C:\dup.log`}, + wantLogs: []string{`C:\dup.log`}, + }, + { + name: "preserve pre-existing LogFiles", + cmdArgs: []string{"installer.exe", `/log:C:\new.log`}, + initLogs: []string{`C:\existing.log`}, + wantLogs: []string{`C:\existing.log`, `C:\new.log`}, + }, + { + name: "do not duplicate pre-existing LogFiles", + cmdArgs: []string{"installer.exe", `/log:C:\existing.log`}, + initLogs: []string{`C:\existing.log`}, + wantLogs: []string{`C:\existing.log`}, + }, + { + name: "consecutive log flags does not consume next flag as path", + cmdArgs: []string{"installer.exe", "/log", "/l*v", `C:\real.log`}, + wantLogs: []string{`C:\real.log`}, + }, + { + name: "non-flag colon argument not misidentified", + cmdArgs: []string{"copy.exe", `C:\source.txt`, `D:\dest.txt`}, + wantLogs: nil, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + cmd := &exec.Cmd{Args: tc.cmdArgs} + opts := enrichOptions(cmd, supervisor.Options{LogFiles: tc.initLogs}, nil) + if !reflect.DeepEqual(opts.LogFiles, tc.wantLogs) { + t.Errorf("enrichOptions() LogFiles = %v, want %v", opts.LogFiles, tc.wantLogs) + } + }) + } +} + +// TestEnrichOptionsNilAndWriter verifies safety when inspecting nil commands and writers. +func TestEnrichOptionsNilAndWriter(t *testing.T) { + // Nil command safety check. + optsNilCmd := enrichOptions(nil, supervisor.Options{}, nil) + if optsNilCmd.LogFiles != nil { + t.Errorf("enrichOptions(nil, ...) LogFiles = %v, want nil", optsNilCmd.LogFiles) + } + + // Command with nil Args slice. + optsNilArgs := enrichOptions(&exec.Cmd{Args: nil}, supervisor.Options{}, nil) + if optsNilArgs.LogFiles != nil { + t.Errorf("enrichOptions(&exec.Cmd{Args: nil}, ...) LogFiles = %v, want nil", optsNilArgs.LogFiles) + } + + // Test writer inspection with a temporary file. + tmpFile, err := os.CreateTemp("", "googet_enrich_test_*.log") + if err != nil { + t.Fatalf("failed to create temp file: %v", err) + } + defer os.Remove(tmpFile.Name()) + defer tmpFile.Close() + + optsFile := enrichOptions(nil, supervisor.Options{}, tmpFile) + if len(optsFile.LogFiles) != 1 || optsFile.LogFiles[0] != tmpFile.Name() { + t.Errorf("enrichOptions(nil, ..., tmpFile) LogFiles = %v, want [%v]", optsFile.LogFiles, tmpFile.Name()) + } + + // Pre-existing log file should not be duplicated when writer is inspected. + optsDupFile := enrichOptions(nil, supervisor.Options{LogFiles: []string{tmpFile.Name()}}, tmpFile) + if len(optsDupFile.LogFiles) != 1 || optsDupFile.LogFiles[0] != tmpFile.Name() { + t.Errorf("enrichOptions(nil, ..., tmpFile) with pre-existing LogFiles = %v, want [%v]", optsDupFile.LogFiles, tmpFile.Name()) + } +} + +// TestAdversarialCornerCases documents empirical edge-case behavior and limitations of enrichOptions. +func TestAdversarialCornerCases(t *testing.T) { + // 1. InnoSetup-style /LOG= is currently unhandled by enrichOptions. + // When installers use /LOG=path, the colon splitter does not split on '='. + // Therefore, the log path is not auto-discovered. + cmdInno := &exec.Cmd{Args: []string{"setup.exe", `/VERYSILENT`, `/LOG=C:\Windows\Logs\inno.log`}} + optsInno := enrichOptions(cmdInno, supervisor.Options{}, nil) + if len(optsInno.LogFiles) != 0 { + t.Logf("Notice: /LOG= was unexpectedly parsed as %v", optsInno.LogFiles) + } else { + t.Logf("Empirically confirmed: InnoSetup /LOG= syntax is unhandled by enrichOptions (LogFiles is empty)") + } + + // 2. Fully-quoted argument "/log:path" has leading quote on the switch. + // strings.SplitN produces parts[0] == `"/log`, which fails isLogFlag. + cmdQuotedSwitch := &exec.Cmd{Args: []string{"setup.exe", `"/log:C:\install.log"`}} + optsQuotedSwitch := enrichOptions(cmdQuotedSwitch, supervisor.Options{}, nil) + if len(optsQuotedSwitch.LogFiles) != 0 { + t.Logf("Notice: Fully-quoted switch was parsed as %v", optsQuotedSwitch.LogFiles) + } else { + t.Logf("Empirically confirmed: Fully-quoted switch \"/log:...\" is not extracted (LogFiles is empty)") + } + + // 3. Space-separated log flag followed by another command switch (/quiet). + // Because /quiet is not a known log flag, isLogFlag("/quiet") returns false. + // This causes enrichOptions to treat "/quiet" as the log file path. + cmdNextSwitch := &exec.Cmd{Args: []string{"msiexec.exe", "/i", "pkg.msi", "/log", "/quiet"}} + optsNextSwitch := enrichOptions(cmdNextSwitch, supervisor.Options{}, nil) + if len(optsNextSwitch.LogFiles) == 1 && optsNextSwitch.LogFiles[0] == "/quiet" { + t.Logf("Empirically confirmed: /log followed by non-log switch treats switch (/quiet) as log path") + } + + // 4. Bare argument "noconfirm" without dash or slash prefix in os.Args. + // strings.TrimLeft(arg, "-/") strips leading prefixes, but if none exist, clean == "noconfirm". + // This inadvertently sets opts.Unattended = true. + oldArgs := os.Args + defer func() { os.Args = oldArgs }() + os.Args = []string{"googet", "install", "noconfirm"} + optsBare := enrichOptions(&exec.Cmd{Args: []string{"cmd.exe"}}, supervisor.Options{}, nil) + if optsBare.Unattended { + t.Logf("Empirically confirmed: Package named 'noconfirm' in os.Args triggers opts.Unattended = true") + } + + // 5. Forward slash paths in Windows commands: /log:C:/temp/install.log. + cmdFwd := &exec.Cmd{Args: []string{"setup.exe", `/log:C:/temp/install.log`}} + optsFwd := enrichOptions(cmdFwd, supervisor.Options{}, nil) + if len(optsFwd.LogFiles) != 1 || optsFwd.LogFiles[0] != "C:/temp/install.log" { + t.Errorf("enrichOptions() with forward-slash path = %v, want [C:/temp/install.log]", optsFwd.LogFiles) + } +} diff --git a/goolib/goospec.go b/goolib/goospec.go index c76496f..e8c457a 100644 --- a/goolib/goospec.go +++ b/goolib/goospec.go @@ -34,6 +34,7 @@ import ( "github.com/blang/semver" "github.com/google/googet/v2/priority" + "github.com/google/googet/v2/supervisor" "github.com/olekukonko/tablewriter/pkg/twwarp" ) @@ -114,6 +115,46 @@ type ExecFile struct { Path string `json:",omitempty"` Args []string `json:",omitempty"` ExitCodes []int `json:",omitempty"` + // Timeout overrides the absolute runtime limit for this command as a Go duration + // string (e.g. "3h"). Empty uses the configured default and "0" disables the limit. + Timeout string `json:",omitempty"` + // InactivityTimeout overrides how long this command may make no forward progress + // before it is terminated. Empty uses the configured default and "0" disables it. + InactivityTimeout string `json:",omitempty"` +} + +// parseOverride converts an ExecFile duration override into a supervisor duration using +// supervisor.ParseTimeout, where zero means "use the default" and a negative value means +// "disabled". +func parseOverride(field, s string) (time.Duration, error) { + d, err := supervisor.ParseTimeout(s) + if err != nil { + return 0, fmt.Errorf("invalid %s: %w", field, err) + } + return d, nil +} + +// overrides returns the hard and inactivity timeout overrides for this command. A zero +// duration means "use the default" and a negative duration means "disabled". +func (e ExecFile) overrides() (hard, inactivity time.Duration, err error) { + if hard, err = parseOverride("timeout", e.Timeout); err != nil { + return 0, 0, err + } + if inactivity, err = parseOverride("inactivityTimeout", e.InactivityTimeout); err != nil { + return 0, 0, err + } + return hard, inactivity, nil +} + +// SupervisorOptions returns the supervisor options declared by this command's timeout +// overrides. Fields that are not overridden are left at their zero value so that the +// process-wide defaults configured via supervisor.Configure apply. +func (e ExecFile) SupervisorOptions() (supervisor.Options, error) { + hard, inactivity, err := e.overrides() + if err != nil { + return supervisor.Options{}, err + } + return supervisor.Options{HardTimeout: hard, InactivityTimeout: inactivity}, nil } // Version contains the semver version as well as the GsVer. @@ -419,6 +460,11 @@ func (ps *PkgSpec) verify() error { if filepath.IsAbs(ps.Uninstall.Path) { return fmt.Errorf("%q is an absolute path, expected relative", ps.Uninstall.Path) } + for name, ef := range map[string]ExecFile{"install": ps.Install, "uninstall": ps.Uninstall, "verify": ps.Verify} { + if _, err := ef.SupervisorOptions(); err != nil { + return fmt.Errorf("%s: %v", name, err) + } + } return nil } diff --git a/goolib/goospec_test.go b/goolib/goospec_test.go index 0cedd9b..0a24fe5 100644 --- a/goolib/goospec_test.go +++ b/goolib/goospec_test.go @@ -648,3 +648,15 @@ func TestPkgSpec_ExecutableNamesAndInstallPath(t *testing.T) { t.Errorf("InstallPath mismatch: got %q, want %q", got.InstallPath, ps.InstallPath) } } + +// TestVerifyRejectsBadOverrides verifies that malformed overrides fail spec verification. +func TestVerifyRejectsBadOverrides(t *testing.T) { + ps := &PkgSpec{Name: "foo", Arch: "noarch", Version: "1.0.0@1", Install: ExecFile{Timeout: "forever"}} + if err := ps.verify(); err == nil { + t.Error("verify() = nil for invalid install timeout, want error") + } + ps.Install.Timeout = "2h" + if err := ps.verify(); err != nil { + t.Errorf("verify() = %v for valid install timeout, want nil", err) + } +} diff --git a/goolib/supervise_test.go b/goolib/supervise_test.go new file mode 100644 index 0000000..ae6a982 --- /dev/null +++ b/goolib/supervise_test.go @@ -0,0 +1,117 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package goolib + +import ( + "errors" + "fmt" + "os/exec" + "runtime" + "testing" + "time" + + "github.com/google/googet/v2/supervisor" +) + +// exitError returns a real *exec.ExitError for a process that exited with code. +func exitError(t *testing.T, code int) *exec.ExitError { + t.Helper() + var c *exec.Cmd + if runtime.GOOS == "windows" { + c = exec.Command("cmd", "/c", fmt.Sprintf("exit %d", code)) + } else { + c = exec.Command("sh", "-c", fmt.Sprintf("exit %d", code)) + } + err := c.Run() + var ee *exec.ExitError + if !errors.As(err, &ee) { + t.Skipf("Could not produce an exit error: %v", err) + } + return ee +} + +// TestCheckExit verifies how supervisor.Run results map to RunWithOptions errors. +func TestCheckExit(t *testing.T) { + ee := exitError(t, 3) + startErr := errors.New("start failed") + tests := []struct { + name string + err error + ec []int + wantNil bool + }{ + {name: "success", err: nil, wantNil: true}, + {name: "accepted exit code", err: ee, ec: []int{3}, wantNil: true}, + {name: "unaccepted exit code", err: ee, ec: []int{1}}, + {name: "non-exit error", err: startErr, ec: []int{3}}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := checkExit(tc.err, tc.ec) + if (got == nil) != tc.wantNil { + t.Errorf("checkExit(%v, %v) = %v, want nil: %v", tc.err, tc.ec, got, tc.wantNil) + } + }) + } + if got := checkExit(startErr, []int{3}); got != startErr { + t.Errorf("checkExit(%v, [3]) = %v, want it returned unchanged", startErr, got) + } +} + +// TestCheckExitTerminatedIsNotAccepted verifies that a supervisor termination is returned +// unchanged even when the killed process's exit code is listed as acceptable. +func TestCheckExitTerminatedIsNotAccepted(t *testing.T) { + ee := exitError(t, 3) + terminated := fmt.Errorf("%w: %w", supervisor.ErrHardTimeout, ee) + got := checkExit(terminated, []int{3, 1, -1, 137}) + if got != terminated { + t.Fatalf("checkExit(terminated, ec) = %v, want %v", got, terminated) + } + if !errors.Is(got, supervisor.ErrTerminated) { + t.Errorf("errors.Is(%v, ErrTerminated) = false, want true", got) + } +} + +// TestExecFileSupervisorOptions verifies the mapping of ExecFile overrides to supervisor options. +func TestExecFileSupervisorOptions(t *testing.T) { + tests := []struct { + name string + ef ExecFile + want supervisor.Options + wantErr bool + }{ + {name: "no overrides", ef: ExecFile{Path: "install.cmd"}, want: supervisor.Options{}}, + {name: "explicit values", ef: ExecFile{Timeout: "3h", InactivityTimeout: "20m"}, want: supervisor.Options{HardTimeout: 3 * time.Hour, InactivityTimeout: 20 * time.Minute}}, + {name: "zero disables", ef: ExecFile{Timeout: "0", InactivityTimeout: "0s"}, want: supervisor.Options{HardTimeout: -1, InactivityTimeout: -1}}, + {name: "negative rejected", ef: ExecFile{Timeout: "-1h"}, wantErr: true}, + {name: "garbage rejected", ef: ExecFile{InactivityTimeout: "soon"}, wantErr: true}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got, err := tc.ef.SupervisorOptions() + if (err != nil) != tc.wantErr { + t.Fatalf("SupervisorOptions() error = %v, wantErr %v", err, tc.wantErr) + } + if tc.wantErr { + if got.HardTimeout != 0 || got.InactivityTimeout != 0 || got.Mode != supervisor.ModeUnset { + t.Errorf("SupervisorOptions() on error = %+v, want zero options", got) + } + return + } + if got.HardTimeout != tc.want.HardTimeout || got.InactivityTimeout != tc.want.InactivityTimeout || got.Mode != supervisor.ModeUnset || got.Unattended || got.DisableUIDetection || got.UIGracePeriod != 0 || len(got.LogFiles) != 0 { + t.Errorf("SupervisorOptions() = %+v, want %+v", got, tc.want) + } + }) + } +} diff --git a/goopack/goopack.go b/goopack/goopack.go index d27ee84..c1b0562 100644 --- a/goopack/goopack.go +++ b/goopack/goopack.go @@ -30,12 +30,17 @@ import ( "github.com/google/googet/v2/goolib" "github.com/google/googet/v2/oswrap" + "github.com/google/googet/v2/supervisor" ) var ( outputDir = flag.String("output_dir", "", "where to put the built package") ) +// buildOptions disables the installer watchdogs for goospec build commands, which are +// build steps rather than installers and may legitimately run for a long time. +var buildOptions = supervisor.Options{Mode: supervisor.ModeOff} + type fileMap map[string][]string // walkDir returns a list of all files in directory and subdirectories, it is similar @@ -349,7 +354,7 @@ func createPackage(gs *goolib.GooSpec, baseDir, outDir string) error { if !filepath.IsAbs(cmd) { cmd = filepath.Join(baseDir, cmd) } - if err := goolib.Exec(cmd, gs.Build.LinuxArgs, nil, ioutil.Discard); err != nil { + if err := goolib.ExecWithOptions(cmd, gs.Build.LinuxArgs, nil, buildOptions, ioutil.Discard); err != nil { return err } case gs.Build.Windows != "" && runtime.GOOS == "windows": @@ -357,7 +362,7 @@ func createPackage(gs *goolib.GooSpec, baseDir, outDir string) error { if !filepath.IsAbs(cmd) { cmd = filepath.Join(baseDir, cmd) } - if err := goolib.Exec(cmd, gs.Build.WindowsArgs, nil, ioutil.Discard); err != nil { + if err := goolib.ExecWithOptions(cmd, gs.Build.WindowsArgs, nil, buildOptions, ioutil.Discard); err != nil { return err } } diff --git a/install/install.go b/install/install.go index cf8e3c6..aca2fa4 100644 --- a/install/install.go +++ b/install/install.go @@ -32,12 +32,9 @@ import ( "github.com/google/googet/v2/oswrap" "github.com/google/googet/v2/remove" "github.com/google/googet/v2/settings" - "github.com/google/googet/v2/system" "github.com/google/logger" ) -var toRemove []string - // minInstalled reports whether the package is installed at the given version or greater. func minInstalled(pi goolib.PackageInfo, db *googetdb.GooDB) (bool, error) { p, err := db.FetchPkg(pi) @@ -200,7 +197,7 @@ func FromRepo(ctx context.Context, pi goolib.PackageInfo, repo, cache string, rm return err } - insFiles, err := installPkg(dst, rs.PackageSpec, dbOnly, force, db) + insFiles, err := installPkg(defaultInstallOps(), dst, rs.PackageSpec, dbOnly, force, db) if err != nil { return err } @@ -226,6 +223,11 @@ func FromRepo(ctx context.Context, pi goolib.PackageInfo, repo, cache string, rm // FromDisk installs a local .goo file. func FromDisk(pkgPath, cache string, dbOnly, force, shouldReinstall bool, db *googetdb.GooDB) error { + return fromDisk(defaultInstallOps(), pkgPath, cache, dbOnly, force, shouldReinstall, db) +} + +// fromDisk implements FromDisk using ops to place the package files. +func fromDisk(ops installOps, pkgPath, cache string, dbOnly, force, shouldReinstall bool, db *googetdb.GooDB) error { if _, err := oswrap.Stat(pkgPath); err != nil { return err } @@ -276,7 +278,7 @@ func FromDisk(pkgPath, cache string, dbOnly, force, shouldReinstall bool, db *go return err } - insFiles, err := installPkg(dst, zs, dbOnly, force, db) + insFiles, err := installPkg(ops, dst, zs, dbOnly, force, db) if err != nil { return err } @@ -335,7 +337,7 @@ func Reinstall(ctx context.Context, ps client.PackageState, rd, force bool, down } } - if _, err := installPkg(ps.LocalPath, ps.PackageSpec, false, force, db); err != nil { + if _, err := installPkg(defaultInstallOps(), ps.LocalPath, ps.PackageSpec, false, force, db); err != nil { return fmt.Errorf("error reinstalling package: %v", err) } @@ -403,18 +405,20 @@ func extractSpec(pkgPath string) (*goolib.PkgSpec, error) { return goolib.ExtractPkgSpec(f) } -func makeInstallFunction(src, dst string, insFiles map[string]string, dbOnly, force bool, conflictMap map[string]string) func(string, os.FileInfo, error) error { +// makeInstallFunction returns a walk function that places the files under src +// into dst and records every change in txn. +func makeInstallFunction(src, dst string, txn *installTxn) func(string, os.FileInfo, error) error { return func(path string, fi os.FileInfo, err error) (outerr error) { if err != nil { return err } outPath := filepath.Join(dst, strings.TrimPrefix(path, src)) - if owner, ok := conflictMap[outPath]; ok && !fi.IsDir() { - if settings.StrictConflicts && !force { + if owner, ok := txn.conflictMap[outPath]; ok && !fi.IsDir() { + if settings.StrictConflicts && !txn.force { return fmt.Errorf("file conflict: %s is already owned by package %s", outPath, owner) } - if force { + if txn.force { logger.Infof("Warning: file conflict: %s is already owned by package %s, overwriting due to force flag", outPath, owner) } else { logger.Infof("Warning: file conflict: %s is already owned by package %s, overwriting because `StrictConflicts` is not set", outPath, owner) @@ -422,38 +426,34 @@ func makeInstallFunction(src, dst string, insFiles map[string]string, dbOnly, fo fmt.Printf("Warning: file conflict: %s is already owned by package %s, overwriting...\n", outPath, owner) } - if dbOnly { + if txn.dbOnly { if !fi.IsDir() { f, err := oswrap.Open(path) if err != nil { return err } defer f.Close() - insFiles[outPath] = goolib.Checksum(f) + txn.insFiles[outPath] = goolib.Checksum(f) } - insFiles[outPath] = "" + txn.insFiles[outPath] = "" return nil } if fi.IsDir() { logger.Infof("Creating folder %q", outPath) // We designate directories by an empty hash. - insFiles[outPath] = "" - return oswrap.MkdirAll(outPath, fi.Mode()) + txn.insFiles[outPath] = "" + return txn.mkdirAllTracked(outPath, fi.Mode()) } - fn, err := client.RemoveOrRename(outPath) - if err != nil { + if err := txn.prepareTarget(outPath); err != nil { return err } - if fn != "" { - toRemove = append(toRemove, fn) - } logger.Infof("Copying file %q", outPath) oFile, err := oswrap.Create(outPath) if err != nil { if !os.IsNotExist(err) { return err } - if err := oswrap.MkdirAll(filepath.Dir(outPath), fi.Mode()); err != nil { + if err := txn.mkdirAllTracked(filepath.Dir(outPath), fi.Mode()); err != nil { return err } if oFile, err = oswrap.Create(outPath); err != nil { @@ -473,10 +473,10 @@ func makeInstallFunction(src, dst string, insFiles map[string]string, dbOnly, fo hash := sha256.New() mw := io.MultiWriter(oFile, hash) - if _, err := io.Copy(mw, iFile); err != nil { + if _, err := txn.ops.copyContents(mw, iFile); err != nil { return err } - insFiles[outPath] = hex.EncodeToString(hash.Sum(nil)) + txn.insFiles[outPath] = hex.EncodeToString(hash.Sum(nil)) return nil } } @@ -549,7 +549,18 @@ func buildConflictMap(db *googetdb.GooDB, currentPkg string) (map[string]string, return conflictMap, nil } -func installPkg(pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb.GooDB) (map[string]string, error) { +// installPkg extracts pkg and places its files using ops. On failure every +// change made to the filesystem is rolled back and the extraction directory, +// which holds the installer logs, is preserved for diagnosis. Callers must not +// record the package in the database unless installPkg returns a nil error. +func installPkg(ops installOps, pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb.GooDB) (map[string]string, error) { + // Build the conflict map first so that a database error leaves nothing + // to clean up. + conflictMap, err := buildConflictMap(db, ps.Name) + if err != nil { + return nil, err + } + dir, err := download.ExtractPkg(pkg) if err != nil { return nil, err @@ -557,39 +568,39 @@ func installPkg(pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb logger.Infof("Executing install of package %q", filepath.Base(dir)) - toRemove = []string{} - // Try to cleanup moved files after package is installed. + txn := newInstallTxn(ops, dbOnly, force, conflictMap) + // success is set only on the final return so that errors and panics both + // trigger rollback. + success := false + defer func() { - for _, fn := range toRemove { - oswrap.Remove(fn) + if !success { + txn.rollback() + logger.Errorf("install logs preserved at %s", dir) + return + } + txn.commit() + if err := oswrap.RemoveAll(dir); err != nil { + logger.Error(err) } }() - conflictMap, err := buildConflictMap(db, ps.Name) - if err != nil { - return nil, err - } - - insFiles := make(map[string]string) for src, dst := range ps.Files { dst = resolveDst(dst) src = filepath.Join(dir, src) - if err := oswrap.Walk(src, makeInstallFunction(src, dst, insFiles, dbOnly, force, conflictMap)); err != nil { + if err := oswrap.Walk(src, makeInstallFunction(src, dst, txn)); err != nil { return nil, err } } if !dbOnly { - if err := system.Install(dir, ps); err != nil { + if err := ops.systemInstall(dir, ps); err != nil { return nil, err } } - if err := oswrap.RemoveAll(dir); err != nil { - logger.Error(err) - } - - return insFiles, nil + success = true + return txn.insFiles, nil } func listDeps(pi goolib.PackageInfo, rm client.RepoMap, repo string, dl []goolib.PackageInfo, archs []string, db *googetdb.GooDB) ([]goolib.PackageInfo, error) { diff --git a/install/install_test.go b/install/install_test.go index 9366a15..9721343 100644 --- a/install/install_test.go +++ b/install/install_test.go @@ -15,13 +15,17 @@ package install import ( "archive/tar" + "bytes" "compress/gzip" + "errors" + "fmt" "io" "io/ioutil" "log" "os" "path/filepath" "reflect" + "strings" "testing" "github.com/google/googet/v2/client" @@ -208,7 +212,7 @@ func TestInstallPkg(t *testing.T) { } ps := goolib.PkgSpec{Files: map[string]string{"./": dst}} - got, err := installPkg(f.Name(), &ps, false, false, db) + got, err := installPkg(defaultInstallOps(), f.Name(), &ps, false, false, db) if err != nil { t.Fatalf("Error running installPkg: %v", err) } @@ -572,33 +576,1425 @@ func TestMakeInstallFunction(t *testing.T) { fi, _ := f.Stat() f.Close() - // Test 1: Conflict without force -> Success by default - fnBlock := makeInstallFunction(srcDir, dstDir, make(map[string]string), false, false, cm) + // Test 1: Conflict without force -> Success by default. + fnBlock := makeInstallFunction(srcDir, dstDir, newInstallTxn(defaultInstallOps(), false, false, cm)) errBlock := fnBlock(filepath.Join(srcDir, "conflicting_file"), fi, nil) if errBlock != nil { t.Errorf("expected no conflict error by default, got %v", errBlock) } - // Test 2: Conflict with force -> Success - fnForce := makeInstallFunction(srcDir, dstDir, make(map[string]string), false, true, cm) + // Test 2: Conflict with force -> Success. + fnForce := makeInstallFunction(srcDir, dstDir, newInstallTxn(defaultInstallOps(), false, true, cm)) errForce := fnForce(filepath.Join(srcDir, "conflicting_file"), fi, nil) if errForce != nil { t.Errorf("expected no error with force, got %v", errForce) } - // Test 3: Conflict without force in strict mode -> Error + // Test 3: Conflict without force in strict mode -> Error. settings.StrictConflicts = true defer func() { settings.StrictConflicts = false }() - fnStrict := makeInstallFunction(srcDir, dstDir, make(map[string]string), false, false, cm) + fnStrict := makeInstallFunction(srcDir, dstDir, newInstallTxn(defaultInstallOps(), false, false, cm)) errStrict := fnStrict(filepath.Join(srcDir, "conflicting_file"), fi, nil) if errStrict == nil { t.Errorf("expected conflict error in strict mode, got nil") } - // Test 4: Conflict with force in strict mode -> Success - fnStrictForce := makeInstallFunction(srcDir, dstDir, make(map[string]string), false, true, cm) + // Test 4: Conflict with force in strict mode -> Success. + fnStrictForce := makeInstallFunction(srcDir, dstDir, newInstallTxn(defaultInstallOps(), false, true, cm)) errStrictForce := fnStrictForce(filepath.Join(srcDir, "conflicting_file"), fi, nil) if errStrictForce != nil { t.Errorf("expected no error with force in strict mode, got %v", errStrictForce) } } + +// TestInstallPkg_RollbackOnFailure verifies that on installer failure, newly placed files +// are removed, replaced files are restored from backup, and backup files are cleaned up. +func TestInstallPkg_RollbackOnFailure(t *testing.T) { + srcDir := t.TempDir() + dstDir := t.TempDir() + + settings.Initialize(t.TempDir(), false) + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Create a pre-existing file that will be replaced. + existingFile := filepath.Join(dstDir, "existing.txt") + originalContent := []byte("original content version 1") + if err := os.WriteFile(existingFile, originalContent, 0644); err != nil { + t.Fatalf("Failed to create existing file: %v", err) + } + + // Build a .goo package containing both an existing file replacement and a newly placed file. + pkgFile := filepath.Join(srcDir, "rollback_test.goo") + f, err := os.Create(pkgFile) + if err != nil { + t.Fatalf("Failed to create package file: %v", err) + } + + gw := gzip.NewWriter(f) + tw := tar.NewWriter(gw) + + filesToPack := []struct { + name string + content []byte + }{ + {"existing.txt", []byte("overwritten content version 2")}, + {"new_file.txt", []byte("brand new file payload")}, + } + + for _, entry := range filesToPack { + hdr := &tar.Header{ + Name: entry.name, + Mode: 0644, + Size: int64(len(entry.content)), + } + if err := tw.WriteHeader(hdr); err != nil { + t.Fatalf("Failed to write tar header for %s: %v", entry.name, err) + } + if _, err := tw.Write(entry.content); err != nil { + t.Fatalf("Failed to write tar content for %s: %v", entry.name, err) + } + } + tw.Close() + gw.Close() + f.Close() + + // Use a backup operation that creates the backup at a known path. + backupFile := existingFile + ".old_backup" + ops := defaultInstallOps() + ops.backup = func(filename string) (string, error) { + if filename == existingFile { + if err := oswrap.Rename(filename, backupFile); err != nil { + return "", err + } + return backupFile, nil + } + return renameToBackup(filename) + } + + ps := &goolib.PkgSpec{ + Name: "rollback_pkg", + Version: "1.0.0@1", + Arch: "noarch", + Files: map[string]string{ + "existing.txt": filepath.Join(dstDir, "existing.txt"), + "new_file.txt": filepath.Join(dstDir, "new_file.txt"), + }, + Install: goolib.ExecFile{ + Path: "nonexistent_installer_binary_that_will_fail.exe", + }, + } + + // Execute installPkg; must return an error. + _, err = installPkg(ops, pkgFile, ps, false, false, db) + if err == nil { + t.Fatalf("installPkg succeeded unexpectedly; expected installer failure error") + } + + // Verify all newly placed files in insFiles are deleted. + newFilePath := filepath.Join(dstDir, "new_file.txt") + if _, err := os.Stat(newFilePath); !os.IsNotExist(err) { + t.Errorf("Newly placed file %s was not deleted during rollback", newFilePath) + } + + // Verify replaced file is restored to its original content. + restoredData, err := os.ReadFile(existingFile) + if err != nil { + t.Fatalf("Expected restored file %s to exist: %v", existingFile, err) + } + if string(restoredData) != string(originalContent) { + t.Errorf("Restored file content = %q, want original content %q", string(restoredData), string(originalContent)) + } + + // Verify temporary backup file is cleaned up. + if _, err := os.Stat(backupFile); !os.IsNotExist(err) { + t.Errorf("Temporary backup file %s still exists after rollback", backupFile) + } +} + +// TestInstallPkg_SuccessCleansBackup verifies that upon successful installation, +// new files are written and temporary backup files are deleted. +func TestInstallPkg_SuccessCleansBackup(t *testing.T) { + srcDir := t.TempDir() + dstDir := t.TempDir() + + settings.Initialize(t.TempDir(), false) + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + existingFile := filepath.Join(dstDir, "existing.txt") + if err := os.WriteFile(existingFile, []byte("original content"), 0644); err != nil { + t.Fatalf("Failed to create existing file: %v", err) + } + + pkgFile := filepath.Join(srcDir, "success_test.goo") + f, err := os.Create(pkgFile) + if err != nil { + t.Fatalf("Failed to create package file: %v", err) + } + + gw := gzip.NewWriter(f) + tw := tar.NewWriter(gw) + + newContent := []byte("overwritten new content") + hdr := &tar.Header{ + Name: "existing.txt", + Mode: 0644, + Size: int64(len(newContent)), + } + if err := tw.WriteHeader(hdr); err != nil { + t.Fatalf("Failed to write tar header: %v", err) + } + if _, err := tw.Write(newContent); err != nil { + t.Fatalf("Failed to write tar content: %v", err) + } + tw.Close() + gw.Close() + f.Close() + + backupFile := existingFile + ".old_backup" + ops := defaultInstallOps() + ops.backup = func(filename string) (string, error) { + if filename == existingFile { + if err := oswrap.Rename(filename, backupFile); err != nil { + return "", err + } + return backupFile, nil + } + return renameToBackup(filename) + } + + ps := &goolib.PkgSpec{ + Name: "success_pkg", + Version: "1.0.0@1", + Arch: "noarch", + Files: map[string]string{"existing.txt": filepath.Join(dstDir, "existing.txt")}, + } + + insFiles, err := installPkg(ops, pkgFile, ps, false, false, db) + if err != nil { + t.Fatalf("installPkg error: %v", err) + } + + data, err := os.ReadFile(existingFile) + if err != nil { + t.Fatalf("Expected %s to exist: %v", existingFile, err) + } + if string(data) != "overwritten new content" { + t.Errorf("File content = %q, want %q", string(data), "overwritten new content") + } + + if _, err := os.Stat(backupFile); !os.IsNotExist(err) { + t.Errorf("Backup file %s was not deleted on success", backupFile) + } + + if _, ok := insFiles[existingFile]; !ok { + t.Errorf("insFiles did not contain %s", existingFile) + } +} + +// TestFromDisk_RollbackAndDBUntouched verifies that when a package installation fails +// via FromDisk, placed files are rolled back and the package is not written to googet.db. +func TestFromDisk_RollbackAndDBUntouched(t *testing.T) { + srcDir := t.TempDir() + dstDir := t.TempDir() + cacheDir := t.TempDir() + + settings.Initialize(t.TempDir(), false) + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + pkgFile := filepath.Join(srcDir, "test_pkg_noarch.goo") + f, err := os.Create(pkgFile) + if err != nil { + t.Fatalf("Failed to create package file: %v", err) + } + + gw := gzip.NewWriter(f) + tw := tar.NewWriter(gw) + + specContent := []byte(`{ + "name": "test_pkg", + "version": "1.0.0@1", + "arch": "noarch", + "files": { + "payload.txt": "` + filepath.Join(dstDir, "payload.txt") + `" + }, + "install": { + "path": "nonexistent_installer_binary_that_will_fail.exe" + } + }`) + hdrSpec := &tar.Header{ + Name: "test_pkg.pkgspec", + Mode: 0644, + Size: int64(len(specContent)), + } + if err := tw.WriteHeader(hdrSpec); err != nil { + t.Fatalf("Failed to write spec header: %v", err) + } + if _, err := tw.Write(specContent); err != nil { + t.Fatalf("Failed to write spec content: %v", err) + } + + payloadContent := []byte("payload content") + hdrPayload := &tar.Header{ + Name: "payload.txt", + Mode: 0644, + Size: int64(len(payloadContent)), + } + if err := tw.WriteHeader(hdrPayload); err != nil { + t.Fatalf("Failed to write payload header: %v", err) + } + if _, err := tw.Write(payloadContent); err != nil { + t.Fatalf("Failed to write payload content: %v", err) + } + tw.Close() + gw.Close() + f.Close() + + err = FromDisk(pkgFile, cacheDir, false, false, false, db) + if err == nil { + t.Fatalf("FromDisk expected error, got nil") + } + + // Verify that a failed install does not record the package in the database. + pState, err := db.FetchPkg(goolib.PackageInfo{Name: "test_pkg", Arch: "noarch", Ver: "1.0.0@1"}) + if err != nil { + t.Fatalf("db.FetchPkg error: %v", err) + } + if pState.PackageSpec != nil { + t.Errorf("Package %s was added to DB despite failed install: %+v", "test_pkg", pState) + } + + // Verify that rollback deleted the file the failed install placed. + placedFile := filepath.Join(dstDir, "payload.txt") + if _, err := os.Stat(placedFile); !os.IsNotExist(err) { + t.Errorf("Placed file %s still exists after failed install", placedFile) + } +} + +// createStressGooArchive packages a .pkgspec and arbitrary files into a .goo archive. +func createStressGooArchive(t *testing.T, archivePath string, specContent []byte, files map[string][]byte) { + t.Helper() + f, err := os.Create(archivePath) + if err != nil { + t.Fatalf("Failed to create goo package file %s: %v", archivePath, err) + } + defer f.Close() + + gw := gzip.NewWriter(f) + defer gw.Close() + tw := tar.NewWriter(gw) + defer tw.Close() + + if specContent != nil { + hdr := &tar.Header{ + Name: "stress_pkg.pkgspec", + Mode: 0644, + Size: int64(len(specContent)), + } + if err := tw.WriteHeader(hdr); err != nil { + t.Fatalf("Failed to write pkgspec header: %v", err) + } + if _, err := tw.Write(specContent); err != nil { + t.Fatalf("Failed to write pkgspec content: %v", err) + } + } + + dirsWritten := make(map[string]bool) + for relPath := range files { + clean := filepath.ToSlash(filepath.Clean(relPath)) + parts := strings.Split(filepath.Dir(clean), "/") + cur := "" + for _, part := range parts { + if part == "." || part == "" { + continue + } + if cur == "" { + cur = part + } else { + cur = cur + "/" + part + } + if !dirsWritten[cur] { + dirsWritten[cur] = true + hdr := &tar.Header{ + Name: cur + "/", + Typeflag: tar.TypeDir, + Mode: 0755, + } + if err := tw.WriteHeader(hdr); err != nil { + t.Fatalf("Failed to write dir header for %s: %v", cur, err) + } + } + } + } + + for relPath, content := range files { + hdr := &tar.Header{ + Name: relPath, + Mode: 0644, + Size: int64(len(content)), + } + if err := tw.WriteHeader(hdr); err != nil { + t.Fatalf("Failed to write tar header for %s: %v", relPath, err) + } + if _, err := tw.Write(content); err != nil { + t.Fatalf("Failed to write tar content for %s: %v", relPath, err) + } + } +} + +// TestStress_RollbackNestedDirectoryTree stress-tests rollback across a complex nested file tree. +func TestStress_RollbackNestedDirectoryTree(t *testing.T) { + srcDir := t.TempDir() + dstDir := t.TempDir() + cacheDir := t.TempDir() + + settings.Initialize(t.TempDir(), false) + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Pre-create existing files in various nested directories. + existingFiles := map[string][]byte{ + filepath.Join(dstDir, "bin", "app.exe"): []byte("v1.0 binary payload"), + filepath.Join(dstDir, "etc", "app", "conf.d", "main.conf"): []byte("v1.0 config contents"), + filepath.Join(dstDir, "lib", "modules", "mod.so"): []byte("v1.0 module binary"), + filepath.Join(dstDir, "var", "log", "keep.log"): []byte("unrelated application log"), + } + for path, content := range existingFiles { + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + t.Fatalf("Failed to create dir for %s: %v", path, err) + } + if err := os.WriteFile(path, content, 0644); err != nil { + t.Fatalf("Failed to write existing file %s: %v", path, err) + } + } + + // Define files in package: 3 replacements and 2 brand new files. + pkgFiles := map[string][]byte{ + "bin/app.exe": []byte("v2.0 overwritten binary payload"), + "etc/app/conf.d/main.conf": []byte("v2.0 overwritten config contents"), + "lib/modules/mod.so": []byte("v2.0 overwritten module binary"), + "opt/extra/addon.txt": []byte("brand new addon text file"), + "bin/helper.exe": []byte("brand new helper tool"), + } + + specJSON := fmt.Sprintf(`{ + "name": "stress_nested_pkg", + "version": "2.0.0@1", + "arch": "noarch", + "files": { + "./": %q + }, + "install": { + "path": "nonexistent_failing_installer.exe" + } + }`, dstDir) + + pkgPath := filepath.Join(srcDir, "stress_nested_pkg.goo") + createStressGooArchive(t, pkgPath, []byte(specJSON), pkgFiles) + + // Use a backup operation that creates backups at known paths. + var createdBackups []string + ops := defaultInstallOps() + ops.backup = func(filename string) (string, error) { + if _, ok := existingFiles[filename]; ok { + backup := filename + ".old_backup" + if err := oswrap.Rename(filename, backup); err != nil { + return "", err + } + createdBackups = append(createdBackups, backup) + return backup, nil + } + return renameToBackup(filename) + } + + // Execute fromDisk; must fail. + err = fromDisk(ops, pkgPath, cacheDir, false, false, false, db) + if err == nil { + t.Fatalf("FromDisk expected installer failure error, got nil") + } + t.Logf("FromDisk returned error: %v", err) + + // 1. Verify newly placed files are deleted. + newFiles := []string{ + filepath.Join(dstDir, "opt", "extra", "addon.txt"), + filepath.Join(dstDir, "bin", "helper.exe"), + } + for _, nf := range newFiles { + if _, err := os.Stat(nf); !os.IsNotExist(err) { + t.Errorf("Newly placed file %s was not deleted during rollback", nf) + } + } + + // 2. Verify replaced files are restored to exact original content. + for path, expectedContent := range existingFiles { + data, err := os.ReadFile(path) + if err != nil { + t.Errorf("Replaced file %s does not exist after rollback: %v", path, err) + continue + } + if string(data) != string(expectedContent) { + t.Errorf("File %s content mismatch after rollback: got %q, want %q", path, string(data), string(expectedContent)) + } + } + + // 3. Verify all temporary backups were cleaned up. + for _, b := range createdBackups { + if _, err := os.Stat(b); !os.IsNotExist(err) { + t.Errorf("Temporary backup file %s still exists after rollback", b) + } + } + + // 4. Verify package is not recorded in googet.db. + pi := goolib.PackageInfo{Name: "stress_nested_pkg", Arch: "noarch", Ver: "2.0.0@1"} + st, err := db.FetchPkg(pi) + if err != nil { + t.Fatalf("db.FetchPkg error: %v", err) + } + if st.PackageSpec != nil { + t.Errorf("Package %s was added to googet.db despite install failure", pi.Name) + } +} + +// TestStress_DatabaseIntegrityOnAllFailureModes verifies DB safety across all failure phases. +func TestStress_DatabaseIntegrityOnAllFailureModes(t *testing.T) { + settings.Initialize(t.TempDir(), false) + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + // Pre-populate DB with two existing baseline packages. + baselineState := []client.PackageState{ + { + PackageSpec: &goolib.PkgSpec{ + Name: "pkg_alpha", + Version: "1.0.0@1", + Arch: "noarch", + }, + }, + { + PackageSpec: &goolib.PkgSpec{ + Name: "pkg_beta", + Version: "1.0.0@1", + Arch: "noarch", + }, + }, + } + if err := db.WriteStateToDB(baselineState); err != nil { + t.Fatalf("db.WriteStateToDB error: %v", err) + } + + assertDBUntouched := func(scenario string) { + t.Helper() + pkgs, err := db.FetchPkgs("") + if err != nil { + t.Fatalf("[%s] db.FetchPkgs error: %v", scenario, err) + } + if len(pkgs) != 2 { + t.Errorf("[%s] Expected 2 packages in DB, found %d: %+v", scenario, len(pkgs), pkgs) + } + for _, name := range []string{"pkg_alpha", "pkg_beta"} { + pi := goolib.PackageInfo{Name: name, Arch: "noarch", Ver: "1.0.0@1"} + p, err := db.FetchPkg(pi) + if err != nil || p.PackageSpec == nil { + t.Errorf("[%s] Expected baseline package %s to remain in DB", scenario, name) + } + } + } + + cacheDir := t.TempDir() + + // Failure Mode 1: Corrupted package file. + { + corruptPkg := filepath.Join(t.TempDir(), "corrupt.goo") + if err := os.WriteFile(corruptPkg, []byte("not a valid gzip tar file"), 0644); err != nil { + t.Fatalf("Failed to write corrupt package: %v", err) + } + if err := FromDisk(corruptPkg, cacheDir, false, false, false, db); err == nil { + t.Errorf("Expected FromDisk error on corrupt archive, got nil") + } + assertDBUntouched("CorruptArchive") + } + + // Failure Mode 2: Replaces installed package conflict. + { + conflictPkg := filepath.Join(t.TempDir(), "conflict.goo") + specJSON := `{ + "name": "pkg_gamma", + "version": "1.0.0@1", + "arch": "noarch", + "replaces": ["pkg_alpha"] + }` + createStressGooArchive(t, conflictPkg, []byte(specJSON), nil) + if err := FromDisk(conflictPkg, cacheDir, false, false, false, db); err == nil { + t.Errorf("Expected FromDisk error on conflicting replaces, got nil") + } + assertDBUntouched("ConflictingReplaces") + } + + // Failure Mode 3: Unsatisfied dependency. + { + missingDepPkg := filepath.Join(t.TempDir(), "missing_dep.goo") + specJSON := `{ + "name": "pkg_delta", + "version": "1.0.0@1", + "arch": "noarch", + "PkgDependencies": { + "nonexistent_dep.noarch": "1.0.0@1" + } + }` + createStressGooArchive(t, missingDepPkg, []byte(specJSON), nil) + if err := FromDisk(missingDepPkg, cacheDir, false, false, false, db); err == nil { + t.Errorf("Expected FromDisk error on missing dependency, got nil") + } + assertDBUntouched("MissingDependency") + } + + // Failure Mode 4: Failing installer binary. + { + failInstallPkg := filepath.Join(t.TempDir(), "fail_install.goo") + specJSON := `{ + "name": "pkg_epsilon", + "version": "1.0.0@1", + "arch": "noarch", + "install": { + "path": "does_not_exist_installer.exe" + } + }` + createStressGooArchive(t, failInstallPkg, []byte(specJSON), nil) + if err := FromDisk(failInstallPkg, cacheDir, false, false, false, db); err == nil { + t.Errorf("Expected FromDisk error on installer failure, got nil") + } + assertDBUntouched("FailingInstaller") + } +} + +// TestStress_MultipleConsecutiveFailuresThenSuccess verifies rollback idempotence and eventual success. +func TestStress_MultipleConsecutiveFailuresThenSuccess(t *testing.T) { + srcDir := t.TempDir() + dstDir := t.TempDir() + cacheDir := t.TempDir() + + settings.Initialize(t.TempDir(), false) + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + targetFile := filepath.Join(dstDir, "stateful.txt") + newFile := filepath.Join(dstDir, "extra.txt") + originalContent := []byte("original pristine state") + if err := os.WriteFile(targetFile, originalContent, 0644); err != nil { + t.Fatalf("Failed to write initial target file: %v", err) + } + + failingSpec := fmt.Sprintf(`{ + "name": "idempotence_pkg", + "version": "1.0.0@1", + "arch": "noarch", + "files": { + "stateful.txt": %q, + "extra.txt": %q + }, + "install": { + "path": "nonexistent_failing_tool.exe" + } + }`, targetFile, newFile) + + failingPkg := filepath.Join(srcDir, "fail.goo") + createStressGooArchive(t, failingPkg, []byte(failingSpec), map[string][]byte{ + "stateful.txt": []byte("corrupted attempt"), + "extra.txt": []byte("unwanted extra"), + }) + + backupFile := targetFile + ".old_backup" + ops := defaultInstallOps() + ops.backup = func(filename string) (string, error) { + if filename == targetFile { + if err := oswrap.Rename(filename, backupFile); err != nil { + return "", err + } + return backupFile, nil + } + return renameToBackup(filename) + } + + // Attempt 1: First failure. + if err := fromDisk(ops, failingPkg, cacheDir, false, false, false, db); err == nil { + t.Fatalf("Attempt 1 expected error, got nil") + } + data, err := os.ReadFile(targetFile) + if err != nil || string(data) != string(originalContent) { + t.Fatalf("Attempt 1 rollback failed: content = %q, want %q", string(data), string(originalContent)) + } + if _, err := os.Stat(newFile); !os.IsNotExist(err) { + t.Fatalf("Attempt 1 extra file still exists") + } + if _, err := os.Stat(backupFile); !os.IsNotExist(err) { + t.Fatalf("Attempt 1 backup file still exists") + } + + // Attempt 2: Second consecutive failure (must not corrupt restored file). + if err := fromDisk(ops, failingPkg, cacheDir, false, false, false, db); err == nil { + t.Fatalf("Attempt 2 expected error, got nil") + } + data, err = os.ReadFile(targetFile) + if err != nil || string(data) != string(originalContent) { + t.Fatalf("Attempt 2 rollback failed: content = %q, want %q", string(data), string(originalContent)) + } + if _, err := os.Stat(newFile); !os.IsNotExist(err) { + t.Fatalf("Attempt 2 extra file still exists") + } + if _, err := os.Stat(backupFile); !os.IsNotExist(err) { + t.Fatalf("Attempt 2 backup file still exists") + } + + // Attempt 3: Successful package installation. + successSpec := fmt.Sprintf(`{ + "name": "idempotence_pkg", + "version": "1.0.0@1", + "arch": "noarch", + "files": { + "stateful.txt": %q, + "extra.txt": %q + } + }`, targetFile, newFile) + + successPkg := filepath.Join(srcDir, "success.goo") + createStressGooArchive(t, successPkg, []byte(successSpec), map[string][]byte{ + "stateful.txt": []byte("final successful v1.0 state"), + "extra.txt": []byte("final successful extra file"), + }) + + if err := fromDisk(ops, successPkg, cacheDir, false, false, false, db); err != nil { + t.Fatalf("Attempt 3 expected success, got error: %v", err) + } + + data, err = os.ReadFile(targetFile) + if err != nil || string(data) != "final successful v1.0 state" { + t.Errorf("Target file content = %q, want %q", string(data), "final successful v1.0 state") + } + extraData, err := os.ReadFile(newFile) + if err != nil || string(extraData) != "final successful extra file" { + t.Errorf("Extra file content = %q, want %q", string(extraData), "final successful extra file") + } + if _, err := os.Stat(backupFile); !os.IsNotExist(err) { + t.Errorf("Backup file still exists after successful install") + } + + pState, err := db.FetchPkg(goolib.PackageInfo{Name: "idempotence_pkg", Arch: "noarch", Ver: "1.0.0@1"}) + if err != nil || pState.PackageSpec == nil { + t.Errorf("Package idempotence_pkg was not recorded in DB on success") + } +} + +// TestStress_FailureMidwayThroughWalk verifies rollback when failure occurs midway through file copying. +func TestStress_FailureMidwayThroughWalk(t *testing.T) { + srcDir := t.TempDir() + dstDir := t.TempDir() + cacheDir := t.TempDir() + + settings.Initialize(t.TempDir(), false) + db, err := googetdb.NewDB(settings.DBFile()) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + defer db.Close() + + file1 := filepath.Join(dstDir, "file1.txt") + file2 := filepath.Join(dstDir, "file2.txt") + originalContent1 := []byte("original file 1") + originalContent2 := []byte("original file 2") + if err := os.WriteFile(file1, originalContent1, 0644); err != nil { + t.Fatalf("Failed to write file1: %v", err) + } + if err := os.WriteFile(file2, originalContent2, 0644); err != nil { + t.Fatalf("Failed to write file2: %v", err) + } + + specJSON := fmt.Sprintf(`{ + "name": "midway_fail_pkg", + "version": "1.0.0@1", + "arch": "noarch", + "files": { + "file1.txt": %q, + "file2.txt": %q + } + }`, file1, file2) + + pkgPath := filepath.Join(srcDir, "midway_fail.goo") + createStressGooArchive(t, pkgPath, []byte(specJSON), map[string][]byte{ + "file1.txt": []byte("new file 1"), + "file2.txt": []byte("new file 2"), + }) + + backupFile1 := file1 + ".old_backup" + simulatedErr := errors.New("simulated access denied on file2") + + ops := defaultInstallOps() + ops.backup = func(filename string) (string, error) { + if filename == file1 { + if err := oswrap.Rename(filename, backupFile1); err != nil { + return "", err + } + return backupFile1, nil + } + if filename == file2 { + return "", simulatedErr + } + return renameToBackup(filename) + } + // Make the remove-or-rename fallback fail for file2 as well. + ops.removeOrRename = func(filename string) (string, error) { + if filename == file2 { + return "", simulatedErr + } + return client.RemoveOrRename(filename) + } + + err = fromDisk(ops, pkgPath, cacheDir, false, false, false, db) + if err == nil { + t.Fatalf("Expected FromDisk error on simulated failure, got nil") + } + if !strings.Contains(err.Error(), "simulated access denied on file2") { + t.Errorf("Error does not contain simulated error: %v", err) + } + + // Verify file1 was restored from backup. + data1, err := os.ReadFile(file1) + if err != nil || string(data1) != string(originalContent1) { + t.Errorf("File 1 was not restored: got %q, want %q", string(data1), string(originalContent1)) + } + if _, err := os.Stat(backupFile1); !os.IsNotExist(err) { + t.Errorf("Backup file 1 still exists after rollback") + } + + // Verify file2 remains untouched. + data2, err := os.ReadFile(file2) + if err != nil || string(data2) != string(originalContent2) { + t.Errorf("File 2 was modified unexpectedly: got %q, want %q", string(data2), string(originalContent2)) + } + // The fallback copy of file2 must be discarded once removeOrRename fails. + assertNoBackups(t, dstDir) + + // Verify DB was untouched. + pState, err := db.FetchPkg(goolib.PackageInfo{Name: "midway_fail_pkg", Arch: "noarch", Ver: "1.0.0@1"}) + if err != nil { + t.Fatalf("db.FetchPkg error: %v", err) + } + if pState.PackageSpec != nil { + t.Errorf("Package midway_fail_pkg was recorded in DB despite failure") + } +} + +// failingOps returns the default operations with a system installer that +// always fails. +func failingOps() installOps { + ops := defaultInstallOps() + ops.systemInstall = func(string, *goolib.PkgSpec) error { + return errors.New("simulated installer failure") + } + return ops +} + +// succeedingOps returns the default operations with a system installer that +// always succeeds. +func succeedingOps() installOps { + ops := defaultInstallOps() + ops.systemInstall = func(string, *goolib.PkgSpec) error { return nil } + return ops +} + +// newTestDB returns a fresh googet database for the test. It does not touch +// the settings package so that callers may run in parallel. +func newTestDB(t *testing.T) *googetdb.GooDB { + t.Helper() + db, err := googetdb.NewDB(filepath.Join(t.TempDir(), "googet.db")) + if err != nil { + t.Fatalf("googetdb.NewDB: %v", err) + } + t.Cleanup(func() { db.Close() }) + return db +} + +// writeFiles creates each file with its content, creating parent directories. +func writeFiles(t *testing.T, files map[string][]byte) { + t.Helper() + for path, content := range files { + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + t.Fatalf("MkdirAll(%q): %v", filepath.Dir(path), err) + } + if err := os.WriteFile(path, content, 0644); err != nil { + t.Fatalf("WriteFile(%q): %v", path, err) + } + } +} + +// assertContents fails the test unless each file exists with exactly the given content. +func assertContents(t *testing.T, files map[string][]byte) { + t.Helper() + for path, want := range files { + got, err := os.ReadFile(path) + if err != nil { + t.Errorf("ReadFile(%q): %v", path, err) + continue + } + if !bytes.Equal(got, want) { + t.Errorf("Content of %q = %q, want %q", path, got, want) + } + } +} + +// assertAbsent fails the test if any of the paths exist. +func assertAbsent(t *testing.T, paths ...string) { + t.Helper() + for _, p := range paths { + if _, err := os.Lstat(p); !os.IsNotExist(err) { + t.Errorf("Path %q exists (Lstat error: %v), want it absent", p, err) + } + } +} + +// assertNoBackups fails the test if any backup file remains under root. +func assertNoBackups(t *testing.T, root string) { + t.Helper() + err := filepath.Walk(root, func(path string, fi os.FileInfo, err error) error { + if err != nil { + return err + } + if strings.Contains(fi.Name(), backupInfix) { + t.Errorf("Leftover backup file %q", path) + } + return nil + }) + if err != nil { + t.Fatalf("Walk(%q): %v", root, err) + } +} + +// extractionDir returns the directory installPkg extracts pkgPath into. +func extractionDir(pkgPath string) string { + return strings.TrimSuffix(pkgPath, filepath.Ext(pkgPath)) +} + +// TestInstallPkg_UpgradeFailureRestoresOldFiles verifies that when an upgrade +// fails in the installer, the previous version's files are restored byte for +// byte, files new to the upgrade are removed, and the extraction directory is +// preserved for diagnosis. +func TestInstallPkg_UpgradeFailureRestoresOldFiles(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + a, b, newOnly := filepath.Join(dstDir, "a.txt"), filepath.Join(dstDir, "b.txt"), filepath.Join(dstDir, "new.txt") + old := map[string][]byte{a: []byte("v1 contents of a"), b: []byte("v1 contents of b")} + writeFiles(t, old) + if err := db.WriteStateToDB([]client.PackageState{{ + PackageSpec: &goolib.PkgSpec{Name: "upgrade_pkg", Version: "1.0.0@1", Arch: "noarch"}, + InstalledFiles: map[string]string{a: "chksum-a", b: "chksum-b"}, + }}); err != nil { + t.Fatalf("WriteStateToDB: %v", err) + } + + pkgPath := filepath.Join(t.TempDir(), "upgrade_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{ + "a.txt": []byte("v2 contents of a, longer than v1"), + "b.txt": []byte("v2 b"), + "new.txt": []byte("only in v2"), + }) + + ps := &goolib.PkgSpec{Name: "upgrade_pkg", Version: "2.0.0@1", Arch: "noarch", Files: map[string]string{"./": dstDir}} + if _, err := installPkg(failingOps(), pkgPath, ps, false, false, db); err == nil { + t.Fatal("installPkg succeeded, want installer failure") + } + + assertContents(t, old) + assertAbsent(t, newOnly) + assertNoBackups(t, dstDir) + if fi, err := os.Stat(extractionDir(pkgPath)); err != nil || !fi.IsDir() { + t.Errorf("Extraction dir %q not preserved after failure: %v", extractionDir(pkgPath), err) + } + assertContents(t, map[string][]byte{filepath.Join(extractionDir(pkgPath), "new.txt"): []byte("only in v2")}) +} + +// TestReinstall_FailureKeepsInstalledFiles verifies that a failed Reinstall +// leaves the currently installed files intact. Reinstall uses the production +// operations, so the failure comes from an installer that does not exist. +func TestReinstall_FailureKeepsInstalledFiles(t *testing.T) { + db := newTestDB(t) + dstDir := t.TempDir() + target := filepath.Join(dstDir, "bin", "tool.exe") + installed := map[string][]byte{target: []byte("currently installed tool")} + writeFiles(t, installed) + + pkgPath := filepath.Join(t.TempDir(), "reinstall_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{"tool.exe": []byte("tool from package cache")}) + + spec := &goolib.PkgSpec{ + Name: "reinstall_pkg", + Version: "1.0.0@1", + Arch: "noarch", + Files: map[string]string{"tool.exe": target}, + Install: goolib.ExecFile{Path: "nonexistent_installer_binary_that_will_fail.exe"}, + } + st := client.PackageState{LocalPath: pkgPath, PackageSpec: spec, InstalledFiles: map[string]string{target: "chksum"}} + if err := Reinstall(t.Context(), st, false, false, nil, db); err == nil { + t.Fatal("Reinstall succeeded, want installer failure") + } + + assertContents(t, installed) + assertNoBackups(t, dstDir) +} + +// TestInstallPkg_ConflictOverwriteRestoredOnFailure verifies that a file owned +// by another package, overwritten because StrictConflicts is off, is restored +// when the install fails. It writes settings.StrictConflicts and therefore +// does not run in parallel. +func TestInstallPkg_ConflictOverwriteRestoredOnFailure(t *testing.T) { + db := newTestDB(t) + settings.StrictConflicts = false + dstDir := t.TempDir() + shared := filepath.Join(dstDir, "shared.dll") + owned := map[string][]byte{shared: []byte("owned by other_pkg")} + writeFiles(t, owned) + if err := db.WriteStateToDB([]client.PackageState{{ + PackageSpec: &goolib.PkgSpec{Name: "other_pkg", Version: "1.0.0@1", Arch: "noarch"}, + InstalledFiles: map[string]string{shared: "chksum"}, + }}); err != nil { + t.Fatalf("WriteStateToDB: %v", err) + } + + pkgPath := filepath.Join(t.TempDir(), "new_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{"shared.dll": []byte("from new_pkg")}) + + ps := &goolib.PkgSpec{Name: "new_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"shared.dll": shared}} + if _, err := installPkg(failingOps(), pkgPath, ps, false, false, db); err == nil { + t.Fatal("installPkg succeeded, want installer failure") + } + + assertContents(t, owned) + assertNoBackups(t, dstDir) +} + +// TestInstallPkg_PartialCopyRolledBack verifies that a copy failing midway +// removes the partially written new file and restores a replaced file. +func TestInstallPkg_PartialCopyRolledBack(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + existing, fresh := filepath.Join(dstDir, "existing.txt"), filepath.Join(dstDir, "fresh.txt") + orig := map[string][]byte{existing: []byte("original existing content")} + writeFiles(t, orig) + + ops := succeedingOps() + ops.copyContents = func(w io.Writer, r io.Reader) (int64, error) { + n, err := io.CopyN(w, r, 4) + if err != nil { + return n, err + } + return n, errors.New("simulated copy failure") + } + + // The subtests share dstDir and therefore run sequentially. + for _, target := range []string{existing, fresh} { + t.Run(filepath.Base(target), func(t *testing.T) { + pkgPath := filepath.Join(t.TempDir(), "partial_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{"payload.txt": []byte("0123456789 full payload")}) + + ps := &goolib.PkgSpec{Name: "partial_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"payload.txt": target}} + _, err := installPkg(ops, pkgPath, ps, false, false, db) + if err == nil || !strings.Contains(err.Error(), "simulated copy failure") { + t.Fatalf("installPkg error = %v, want simulated copy failure", err) + } + assertContents(t, orig) + assertAbsent(t, fresh) + assertNoBackups(t, dstDir) + }) + } +} + +// TestInstallPkg_NewDirectoriesRemovedOnRollback verifies that directories +// created by a failed install are removed while pre-existing directories and +// their unrelated contents are kept. +func TestInstallPkg_NewDirectoriesRemovedOnRollback(t *testing.T) { + t.Parallel() + db := newTestDB(t) + root := t.TempDir() + appDir := filepath.Join(root, "app") + keep := filepath.Join(appDir, "existing_dir", "keep.txt") + kept := map[string][]byte{keep: []byte("unrelated file")} + writeFiles(t, kept) + + pkgPath := filepath.Join(t.TempDir(), "dirs_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{ + "existing_dir/new.txt": []byte("new in existing dir"), + "newdir/deep/f.txt": []byte("deep new file"), + "newdir/other/leaf.txt": []byte("another new file"), + }) + + target := filepath.Join(appDir, "nested", "install_root") + ps := &goolib.PkgSpec{Name: "dirs_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{ + "./": target, + // Also install into the pre-existing directory tree. + "existing_dir": filepath.Join(appDir, "existing_dir"), + }} + if _, err := installPkg(failingOps(), pkgPath, ps, false, false, db); err == nil { + t.Fatal("installPkg succeeded, want installer failure") + } + + assertContents(t, kept) + assertAbsent(t, filepath.Join(appDir, "nested"), filepath.Join(appDir, "existing_dir", "new.txt")) + if fi, err := os.Stat(filepath.Join(appDir, "existing_dir")); err != nil || !fi.IsDir() { + t.Errorf("Pre-existing directory was removed: %v", err) + } +} + +// TestInstallPkg_SuccessRemovesBackupsAndExtractionDir verifies that a +// successful install writes the new contents, leaves no backup files, and +// removes the extraction directory. +func TestInstallPkg_SuccessRemovesBackupsAndExtractionDir(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + a, b := filepath.Join(dstDir, "a.txt"), filepath.Join(dstDir, "b.txt") + writeFiles(t, map[string][]byte{a: []byte("old a")}) + + pkgPath := filepath.Join(t.TempDir(), "ok_pkg.goo") + newContents := map[string][]byte{"a.txt": []byte("new a"), "b.txt": []byte("new b")} + createStressGooArchive(t, pkgPath, nil, newContents) + + ps := &goolib.PkgSpec{Name: "ok_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"./": dstDir}} + insFiles, err := installPkg(succeedingOps(), pkgPath, ps, false, false, db) + if err != nil { + t.Fatalf("installPkg: %v", err) + } + + assertContents(t, map[string][]byte{a: newContents["a.txt"], b: newContents["b.txt"]}) + assertNoBackups(t, dstDir) + assertAbsent(t, extractionDir(pkgPath)) + for _, f := range []string{a, b} { + if insFiles[f] == "" { + t.Errorf("insFiles[%q] is empty, want a checksum", f) + } + } +} + +// TestInstallPkg_SamePathTwiceRestoresOriginal verifies that when two package +// entries write the same destination, rollback restores the pre-install file +// rather than the first entry's content. +func TestInstallPkg_SamePathTwiceRestoresOriginal(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + target := filepath.Join(dstDir, "x.txt") + orig := map[string][]byte{target: []byte("pre-install x")} + writeFiles(t, orig) + + pkgPath := filepath.Join(t.TempDir(), "dup_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{"one.txt": []byte("first"), "two.txt": []byte("second")}) + + ps := &goolib.PkgSpec{Name: "dup_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"one.txt": target, "two.txt": target}} + if _, err := installPkg(failingOps(), pkgPath, ps, false, false, db); err == nil { + t.Fatal("installPkg succeeded, want installer failure") + } + + assertContents(t, orig) + assertNoBackups(t, dstDir) +} + +// TestInstallPkg_BackupFallback verifies the behavior when renaming a file to +// a backup fails and the remove-or-rename fallback preserves the original +// under a new name: that name is restored, the redundant copy is discarded, +// and the install still proceeds. +func TestInstallPkg_BackupFallback(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + target := filepath.Join(dstDir, "locked.exe") + orig := map[string][]byte{target: []byte("locked original")} + writeFiles(t, orig) + fallbackBackup := filepath.Join(dstDir, "fallback.old") + + ops := failingOps() + ops.backup = func(string) (string, error) { return "", errors.New("simulated rename failure") } + ops.removeOrRename = func(filename string) (string, error) { + if err := oswrap.Rename(filename, fallbackBackup); err != nil { + return "", err + } + return fallbackBackup, nil + } + + pkgPath := filepath.Join(t.TempDir(), "fallback_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{"locked.exe": []byte("replacement")}) + + ps := &goolib.PkgSpec{Name: "fallback_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"locked.exe": target}} + if _, err := installPkg(ops, pkgPath, ps, false, false, db); err == nil { + t.Fatal("installPkg succeeded, want installer failure") + } + + assertContents(t, orig) + assertAbsent(t, fallbackBackup) + assertNoBackups(t, dstDir) +} + +// TestInstallPkg_BackupFallbackDeletedFileRestored verifies that when renaming +// a file to a backup fails and the fallback deletes the file, the copy made +// beforehand restores it byte for byte, with its mode, on rollback, and is +// removed after a successful install. +func TestInstallPkg_BackupFallbackDeletedFileRestored(t *testing.T) { + t.Parallel() + for _, tc := range []struct { + name string + succeed bool + }{ + {name: "rollback", succeed: false}, + {name: "commit", succeed: true}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + target := filepath.Join(dstDir, "app.dll") + // Include every byte value to catch any transformation of the data. + content := make([]byte, 4096) + for i := range content { + content[i] = byte(i) + } + writeFiles(t, map[string][]byte{target: content}) + if err := os.Chmod(target, 0640); err != nil { + t.Fatalf("Chmod: %v", err) + } + + ops := failingOps() + if tc.succeed { + ops = succeedingOps() + } + ops.backup = func(string) (string, error) { return "", errors.New("simulated rename failure") } + var deleted bool + ops.removeOrRename = func(filename string) (string, error) { + if err := os.Remove(filename); err != nil { + return "", err + } + deleted = true + return "", nil + } + + pkgPath := filepath.Join(t.TempDir(), "deleted_pkg.goo") + replacement := []byte("replacement payload") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{"app.dll": replacement}) + + ps := &goolib.PkgSpec{Name: "deleted_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"app.dll": target}} + _, err := installPkg(ops, pkgPath, ps, false, false, db) + if !deleted { + t.Fatal("removeOrRename was not called; the fallback path was not exercised") + } + if tc.succeed { + if err != nil { + t.Fatalf("installPkg: %v", err) + } + assertContents(t, map[string][]byte{target: replacement}) + assertNoBackups(t, dstDir) + return + } + if err == nil { + t.Fatal("installPkg succeeded, want installer failure") + } + assertContents(t, map[string][]byte{target: content}) + assertNoBackups(t, dstDir) + fi, err := os.Stat(target) + if err != nil { + t.Fatalf("Stat(%q): %v", target, err) + } + if got := fi.Mode().Perm(); got != 0640 { + t.Errorf("Mode of restored %q = %v, want %v", target, got, os.FileMode(0640)) + } + }) + } +} + +// TestInstallPkg_BackupFallbackCopyFails verifies that when both the rename +// and the copy to a backup fail and the fallback deletes the file, the install +// proceeds and rollback removes the new file, since the original cannot be +// restored. +func TestInstallPkg_BackupFallbackCopyFails(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + target := filepath.Join(dstDir, "app.dll") + writeFiles(t, map[string][]byte{target: []byte("unrecoverable original")}) + + ops := failingOps() + ops.backup = func(string) (string, error) { return "", errors.New("simulated rename failure") } + ops.copyBackup = func(string) (string, error) { return "", errors.New("simulated copy failure") } + ops.removeOrRename = func(filename string) (string, error) { return "", os.Remove(filename) } + + pkgPath := filepath.Join(t.TempDir(), "nocopy_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{"app.dll": []byte("replacement")}) + + ps := &goolib.PkgSpec{Name: "nocopy_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"app.dll": target}} + if _, err := installPkg(ops, pkgPath, ps, false, false, db); err == nil { + t.Fatal("installPkg succeeded, want installer failure") + } + assertAbsent(t, target) + assertNoBackups(t, dstDir) +} + +// TestInstallPkg_RemovedEmptyDirRecreatedOnRollback verifies that an empty +// directory removed to make room for a package file is recreated with its +// original mode on rollback, and that other replaced files are restored byte +// for byte. +func TestInstallPkg_RemovedEmptyDirRecreatedOnRollback(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + emptyDir := filepath.Join(dstDir, "config") + if err := os.Mkdir(emptyDir, 0755); err != nil { + t.Fatalf("Mkdir: %v", err) + } + if err := os.Chmod(emptyDir, 0750); err != nil { + t.Fatalf("Chmod: %v", err) + } + other := filepath.Join(dstDir, "other.txt") + orig := map[string][]byte{other: []byte("original other")} + writeFiles(t, orig) + + pkgPath := filepath.Join(t.TempDir(), "dir_to_file_pkg.goo") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{ + "config": []byte("a file where a directory was"), + "other.txt": []byte("replacement other"), + }) + + ps := &goolib.PkgSpec{Name: "dir_to_file_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"./": dstDir}} + if _, err := installPkg(failingOps(), pkgPath, ps, false, false, db); err == nil { + t.Fatal("installPkg succeeded, want installer failure") + } + + assertContents(t, orig) + assertNoBackups(t, dstDir) + fi, err := os.Lstat(emptyDir) + if err != nil { + t.Fatalf("Lstat(%q): %v", emptyDir, err) + } + if !fi.IsDir() { + t.Fatalf("%q is not a directory after rollback: mode %v", emptyDir, fi.Mode()) + } + if got := fi.Mode().Perm(); got != 0750 { + t.Errorf("Mode of recreated %q = %v, want %v", emptyDir, got, os.FileMode(0750)) + } + entries, err := os.ReadDir(emptyDir) + if err != nil { + t.Fatalf("ReadDir(%q): %v", emptyDir, err) + } + if len(entries) != 0 { + t.Errorf("Recreated %q has %d entries, want it empty", emptyDir, len(entries)) + } +} + +// TestInstallPkg_RemovedEmptyDirReplacedOnSuccess verifies that a successful +// install leaves the package file in place of the removed empty directory. +func TestInstallPkg_RemovedEmptyDirReplacedOnSuccess(t *testing.T) { + t.Parallel() + db := newTestDB(t) + dstDir := t.TempDir() + emptyDir := filepath.Join(dstDir, "config") + if err := os.Mkdir(emptyDir, 0755); err != nil { + t.Fatalf("Mkdir: %v", err) + } + + pkgPath := filepath.Join(t.TempDir(), "dir_to_file_ok_pkg.goo") + content := []byte("a file where a directory was") + createStressGooArchive(t, pkgPath, nil, map[string][]byte{"config": content}) + + ps := &goolib.PkgSpec{Name: "dir_to_file_ok_pkg", Version: "1.0.0@1", Arch: "noarch", Files: map[string]string{"config": emptyDir}} + if _, err := installPkg(succeedingOps(), pkgPath, ps, false, false, db); err != nil { + t.Fatalf("installPkg: %v", err) + } + assertContents(t, map[string][]byte{emptyDir: content}) +} + +// TestRenameToBackup verifies that renameToBackup moves the file to a unique +// sibling in the same directory and preserves its content. +func TestRenameToBackup(t *testing.T) { + t.Parallel() + dir := t.TempDir() + path := filepath.Join(dir, "file.bin") + content := []byte("some content") + seen := make(map[string]bool) + for i := 0; i < 3; i++ { + writeFiles(t, map[string][]byte{path: content}) + bak, err := renameToBackup(path) + if err != nil { + t.Fatalf("renameToBackup: %v", err) + } + if filepath.Dir(bak) != dir || !strings.HasPrefix(filepath.Base(bak), "file.bin"+backupInfix) { + t.Errorf("renameToBackup = %q, want sibling named file.bin%s", bak, backupInfix) + } + if seen[bak] { + t.Errorf("renameToBackup reused backup name %q", bak) + } + seen[bak] = true + assertAbsent(t, path) + assertContents(t, map[string][]byte{bak: content}) + } +} + +// TestCopyToBackup verifies that copyToBackup copies the file to a unique +// sibling with the same content and mode and leaves the original in place. +func TestCopyToBackup(t *testing.T) { + t.Parallel() + dir := t.TempDir() + path := filepath.Join(dir, "file.bin") + content := []byte("some content\x00\xff") + writeFiles(t, map[string][]byte{path: content}) + if err := os.Chmod(path, 0604); err != nil { + t.Fatalf("Chmod: %v", err) + } + seen := make(map[string]bool) + for i := 0; i < 3; i++ { + bak, err := copyToBackup(path) + if err != nil { + t.Fatalf("copyToBackup: %v", err) + } + if filepath.Dir(bak) != dir || !strings.HasPrefix(filepath.Base(bak), "file.bin"+backupInfix) { + t.Errorf("copyToBackup = %q, want sibling named file.bin%s", bak, backupInfix) + } + if seen[bak] { + t.Errorf("copyToBackup reused backup name %q", bak) + } + seen[bak] = true + assertContents(t, map[string][]byte{path: content, bak: content}) + fi, err := os.Stat(bak) + if err != nil { + t.Fatalf("Stat(%q): %v", bak, err) + } + if got := fi.Mode().Perm(); got != 0604 { + t.Errorf("Mode of %q = %v, want %v", bak, got, os.FileMode(0604)) + } + } +} + +// TestCopyToBackup_MissingFile verifies that copyToBackup fails for a missing +// file and leaves no backup behind. +func TestCopyToBackup_MissingFile(t *testing.T) { + t.Parallel() + dir := t.TempDir() + if _, err := copyToBackup(filepath.Join(dir, "missing.bin")); err == nil { + t.Error("copyToBackup succeeded for a missing file, want an error") + } + assertNoBackups(t, dir) +} diff --git a/install/txn.go b/install/txn.go new file mode 100644 index 0000000..38d749a --- /dev/null +++ b/install/txn.go @@ -0,0 +1,378 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package install + +import ( + "crypto/rand" + "encoding/hex" + "fmt" + "io" + "os" + "path/filepath" + + "github.com/google/googet/v2/client" + "github.com/google/googet/v2/goolib" + "github.com/google/googet/v2/oswrap" + "github.com/google/googet/v2/system" + "github.com/google/logger" +) + +// backupInfix is inserted between a file name and a random suffix to form the +// name of the backup made before the file is overwritten. +const backupInfix = ".googet-bak-" + +// dirModeMask selects the directory mode bits that rollback restores when it +// recreates a directory. +const dirModeMask = os.ModePerm | os.ModeSetuid | os.ModeSetgid | os.ModeSticky + +// installOps holds the filesystem and installer operations performed while +// installing a package. Production code uses defaultInstallOps; tests build +// their own value to simulate failures without mutating package state. +type installOps struct { + // backup renames an existing file to a unique backup path in the same + // directory and returns the backup path. + backup func(path string) (string, error) + // copyBackup copies an existing file to a unique backup path in the same + // directory and returns the backup path. It is used when backup fails. + copyBackup func(path string) (string, error) + // removeOrRename deletes a file when possible and otherwise renames it. It + // returns the new name, or an empty string if the file was deleted. + removeOrRename func(path string) (string, error) + // copyContents copies package file contents to their destination. + copyContents func(dst io.Writer, src io.Reader) (int64, error) + // systemInstall runs the package's system specific installer. + systemInstall func(dir string, ps *goolib.PkgSpec) error +} + +// defaultInstallOps returns the operations used by production installs. +func defaultInstallOps() installOps { + return installOps{ + backup: renameToBackup, + copyBackup: copyToBackup, + removeOrRename: client.RemoveOrRename, + copyContents: io.Copy, + systemInstall: system.Install, + } +} + +// movedFile tracks a file that was moved or copied to a temporary backup. +type movedFile struct { + originalPath string + backupPath string +} + +// removedDir tracks a pre-existing empty directory that was removed to make +// room for a file. +type removedDir struct { + path string + mode os.FileMode +} + +// installTxn places the files of a single package and records every change it +// makes to the filesystem so that the changes can be undone if the install +// fails. +type installTxn struct { + // ops performs the filesystem and installer operations. + ops installOps + // dbOnly records files in insFiles without touching the filesystem. + dbOnly bool + // force overwrites files owned by other packages even with StrictConflicts. + force bool + // conflictMap maps files owned by other packages to their owner. + conflictMap map[string]string + // insFiles maps each placed path to its checksum, or to an empty string for + // directories. + insFiles map[string]string + + // created holds files that did not exist before the install wrote them. + created map[string]bool + // createdDirs holds directories that did not exist before the install. + createdDirs []string + // removedDirs holds empty directories removed to make room for files. + removedDirs []removedDir + // moved holds pre-existing files that were moved to backups, in order. + moved []movedFile + // backedUp holds original paths that already have a backup in moved. + backedUp map[string]bool +} + +// newInstallTxn returns an empty transaction that uses ops and records placed +// files in a fresh insFiles map. +func newInstallTxn(ops installOps, dbOnly, force bool, conflictMap map[string]string) *installTxn { + return &installTxn{ + ops: ops, + dbOnly: dbOnly, + force: force, + conflictMap: conflictMap, + insFiles: make(map[string]string), + created: make(map[string]bool), + backedUp: make(map[string]bool), + } +} + +// newBackupPath returns a randomly named sibling of path that uses the backup +// naming scheme. +func newBackupPath(path string) (string, error) { + b := make([]byte, 8) + if _, err := rand.Read(b); err != nil { + return "", err + } + return path + backupInfix + hex.EncodeToString(b), nil +} + +// renameToBackup renames path to a unique, randomly named sibling and returns +// the new name. Renaming within the same directory avoids cross-volume moves +// and works on Windows even for executables that are currently running. +func renameToBackup(path string) (string, error) { + var lastErr error + for i := 0; i < 10; i++ { + bak, err := newBackupPath(path) + if err != nil { + return "", err + } + if _, err := oswrap.Lstat(bak); err == nil { + lastErr = fmt.Errorf("backup path %q already exists", bak) + continue + } else if !os.IsNotExist(err) { + return "", err + } + if err := oswrap.Rename(path, bak); err != nil { + return "", err + } + return bak, nil + } + return "", lastErr +} + +// copyToBackup copies path to a unique, randomly named sibling with the same +// permissions and returns the new name. It leaves path in place and removes a +// partial copy on failure. +func copyToBackup(path string) (string, error) { + src, err := oswrap.Open(path) + if err != nil { + return "", err + } + defer src.Close() + fi, err := src.Stat() + if err != nil { + return "", err + } + perm := fi.Mode().Perm() + + var bak string + var dst *os.File + for i := 0; i < 10; i++ { + if bak, err = newBackupPath(path); err != nil { + return "", err + } + // The owner write bit is added so the copy can be written; the exact + // permissions are applied once the copy is complete. + dst, err = oswrap.OpenFile(bak, os.O_WRONLY|os.O_CREATE|os.O_EXCL, perm|0200) + if err == nil || !os.IsExist(err) { + break + } + } + if err != nil { + return "", err + } + _, err = io.Copy(dst, src) + if cerr := dst.Close(); err == nil { + err = cerr + } + if err == nil { + // oswrap has no Chmod wrapper, so os.Chmod is used directly. + err = os.Chmod(bak, perm) + } + if err != nil { + if rerr := oswrap.Remove(bak); rerr != nil && !os.IsNotExist(rerr) { + logger.Errorf("Failed to remove partial backup copy %q: %v.", bak, rerr) + } + return "", err + } + return bak, nil +} + +// mkdirAllTracked behaves like oswrap.MkdirAll but records every directory it +// creates so that rollback can remove them. +func (txn *installTxn) mkdirAllTracked(path string, mode os.FileMode) error { + var missing []string + for p := filepath.Clean(path); ; { + if _, err := oswrap.Lstat(p); !os.IsNotExist(err) { + break + } + missing = append(missing, p) + parent := filepath.Dir(p) + if parent == p { + break + } + p = parent + } + if err := oswrap.MkdirAll(path, mode); err != nil { + return err + } + // Record shallowest first; rollback removes them deepest first. + for i := len(missing) - 1; i >= 0; i-- { + txn.createdDirs = append(txn.createdDirs, missing[i]) + } + return nil +} + +// recordMoved records that the content of originalPath now lives at +// backupPath. +func (txn *installTxn) recordMoved(originalPath, backupPath string) { + txn.moved = append(txn.moved, movedFile{originalPath: originalPath, backupPath: backupPath}) + txn.backedUp[originalPath] = true +} + +// discardBackup removes a backup that is no longer needed. An empty path is +// ignored. +func (txn *installTxn) discardBackup(path string) { + if path == "" { + return + } + if err := oswrap.Remove(path); err != nil && !os.IsNotExist(err) { + logger.Errorf("Failed to remove redundant backup %q: %v.", path, err) + } +} + +// prepareTarget makes outPath ready to be written. An existing file is moved +// to a backup so it can be restored on rollback; a missing file is recorded as +// created before it is written so that a partial copy is also rolled back. +func (txn *installTxn) prepareTarget(outPath string) error { + if txn.created[outPath] || txn.backedUp[outPath] { + // This install already owns the current content of outPath. + return nil + } + fi, err := oswrap.Lstat(outPath) + if err != nil { + if !os.IsNotExist(err) { + return err + } + txn.created[outPath] = true + return nil + } + if fi.IsDir() { + return txn.removeEmptyDir(outPath, fi.Mode()) + } + bak, err := txn.ops.backup(outPath) + if err == nil { + txn.recordMoved(outPath, bak) + return nil + } + logger.Warningf("Unable to back up %q before overwriting it: %v; falling back to remove or rename.", outPath, err) + return txn.fallbackBackup(outPath) +} + +// removeEmptyDir removes the empty directory at path so that a file can take +// its place, and records it so that rollback recreates it with mode. This +// preserves the legacy behavior; a non-empty directory causes an error. +func (txn *installTxn) removeEmptyDir(path string, mode os.FileMode) error { + fn, err := txn.ops.removeOrRename(path) + if err != nil { + return err + } + if fn != "" { + txn.recordMoved(path, fn) + return nil + } + txn.removedDirs = append(txn.removedDirs, removedDir{path: path, mode: mode}) + txn.created[path] = true + return nil +} + +// fallbackBackup clears outPath after the same-directory rename to a backup +// failed. The file is first copied to a sibling backup because removeOrRename +// may delete it; if removeOrRename instead preserves the original under a new +// name, that name is used as the backup and the copy is discarded. +func (txn *installTxn) fallbackBackup(outPath string) error { + cp, cpErr := txn.ops.copyBackup(outPath) + if cpErr != nil { + logger.Warningf("Unable to copy %q to a backup: %v.", outPath, cpErr) + } + fn, err := txn.ops.removeOrRename(outPath) + if err != nil { + // The original file is still in place, so the copy is not needed. + txn.discardBackup(cp) + return err + } + switch { + case fn != "": + txn.discardBackup(cp) + txn.recordMoved(outPath, fn) + case cpErr == nil: + txn.recordMoved(outPath, cp) + default: + logger.Warningf("Existing file %q was deleted before being overwritten; it cannot be restored if the install fails.", outPath) + txn.created[outPath] = true + } + return nil +} + +// rollback undoes the recorded changes: files the install created are +// removed, backups are restored in reverse order, removed empty directories are +// recreated, and directories the install created are removed deepest first if +// they are empty. +func (txn *installTxn) rollback() { + for file := range txn.created { + if err := oswrap.Remove(file); err != nil && !os.IsNotExist(err) { + logger.Errorf("Failed to remove newly placed file %q during rollback: %v.", file, err) + } + } + for i := len(txn.moved) - 1; i >= 0; i-- { + mf := txn.moved[i] + if _, err := oswrap.Lstat(mf.originalPath); err == nil { + if err := oswrap.Remove(mf.originalPath); err != nil && !os.IsNotExist(err) { + logger.Errorf("Failed to remove file %q during rollback: %v.", mf.originalPath, err) + } + } + if err := oswrap.Rename(mf.backupPath, mf.originalPath); err != nil { + logger.Errorf("Failed to restore backup %q to %q during rollback: %v.", mf.backupPath, mf.originalPath, err) + } + } + for i := len(txn.removedDirs) - 1; i >= 0; i-- { + rd := txn.removedDirs[i] + if err := oswrap.Mkdir(rd.path, rd.mode&dirModeMask); err != nil { + logger.Errorf("Failed to recreate directory %q during rollback: %v.", rd.path, err) + continue + } + // Mkdir is subject to the umask, so apply the original mode explicitly. + // oswrap has no Chmod wrapper, so os.Chmod is used directly. + if err := os.Chmod(rd.path, rd.mode&dirModeMask); err != nil { + logger.Errorf("Failed to restore mode of directory %q during rollback: %v.", rd.path, err) + } + } + for i := len(txn.createdDirs) - 1; i >= 0; i-- { + d := txn.createdDirs[i] + // Remove fails on non-empty directories, which must be left in place. + if err := oswrap.Remove(d); err != nil && !os.IsNotExist(err) { + logger.Infof("Leaving directory %q during rollback: %v.", d, err) + } + } +} + +// commit deletes the recorded backups after a successful install. Backups that +// cannot be deleted, e.g. because they are locked, are scheduled for removal +// on reboot. +func (txn *installTxn) commit() { + for _, mf := range txn.moved { + err := oswrap.Remove(mf.backupPath) + if err == nil || os.IsNotExist(err) { + continue + } + logger.Errorf("Failed to remove backup file %q: %v.", mf.backupPath, err) + if err := oswrap.RemoveOnReboot(mf.backupPath); err != nil { + logger.Errorf("Failed to schedule removal of backup file %q on reboot: %v.", mf.backupPath, err) + } + } +} diff --git a/settings/settings.go b/settings/settings.go index 59a0bec..b630c80 100644 --- a/settings/settings.go +++ b/settings/settings.go @@ -7,6 +7,7 @@ import ( "path/filepath" "time" + "github.com/google/googet/v2/supervisor" "github.com/google/googet/v2/system" "github.com/google/logger" "gopkg.in/yaml.v3" @@ -31,6 +32,23 @@ var ( AllowUnsafeURL bool // StrictConflicts enables strict enforcement of file ownership conflicts. StrictConflicts bool + // SupervisorMode is the installer watchdog mode parsed from googet.conf ("enforce", + // "monitor" or "off"). ModeUnset means the built-in default. + SupervisorMode supervisor.Mode + // InactivityTimeout is how long an installer may make no forward progress before it is terminated. + // Zero means the built-in default and a negative value disables the watchdog. + InactivityTimeout time.Duration + // InstallTimeout is the absolute runtime limit for an installer. + // Zero means the built-in default and a negative value disables the limit. + InstallTimeout time.Duration + // UIGracePeriod is how long an interactive dialog may persist without progress before termination. + // Zero means the built-in default. + UIGracePeriod time.Duration + // UIDetection enables interactive dialog detection in unattended mode. + UIDetection = true + // DownloadStallTimeout is how long a download may receive zero bytes before it is retried. + // Zero means the built-in default. + DownloadStallTimeout time.Duration ) // Initialize reads the initial settings. @@ -78,12 +96,18 @@ func RepoDir() string { // conf represents a googet configuration file. type conf struct { - Archs []string - CacheLife string - LockFileMaxAge string - ProxyServer string - AllowUnsafeURL bool - StrictConflicts bool + Archs []string + CacheLife string + LockFileMaxAge string + ProxyServer string + AllowUnsafeURL bool + StrictConflicts bool + SupervisorMode string + InactivityTimeout string + InstallTimeout string + UIGracePeriod string + UIDetection *bool + DownloadStallTimeout string } // unmarshalConfFile returns a conf from a YAML configuration file. @@ -142,4 +166,42 @@ func readConf(filename string) { AllowUnsafeURL = gc.AllowUnsafeURL StrictConflicts = gc.StrictConflicts + + SupervisorMode = parseMode(gc.SupervisorMode) + InactivityTimeout = parseTimeout("InactivityTimeout", gc.InactivityTimeout, true) + InstallTimeout = parseTimeout("InstallTimeout", gc.InstallTimeout, true) + UIGracePeriod = parseTimeout("UIGracePeriod", gc.UIGracePeriod, false) + DownloadStallTimeout = parseTimeout("DownloadStallTimeout", gc.DownloadStallTimeout, false) + UIDetection = gc.UIDetection == nil || *gc.UIDetection +} + +// parseMode parses the googet.conf SupervisorMode. Empty or invalid values return +// supervisor.ModeUnset (use the built-in default); invalid values are logged. +func parseMode(s string) supervisor.Mode { + if s == "" { + return supervisor.ModeUnset + } + m, err := supervisor.ParseMode(s) + if err != nil { + logger.Errorf("Invalid SupervisorMode in googet.conf, using default: %v", err) + return supervisor.ModeUnset + } + return m +} + +// parseTimeout parses a googet.conf duration with supervisor.ParseTimeout. Empty or +// invalid values return zero (use the built-in default) and invalid values are logged. +// If allowDisable is set, "0" returns a negative value meaning disabled; otherwise "0" is +// rejected as invalid. +func parseTimeout(name, s string, allowDisable bool) time.Duration { + d, err := supervisor.ParseTimeout(s) + switch { + case err != nil: + logger.Errorf("Invalid %s in googet.conf, using default: %v", name, err) + return 0 + case d < 0 && !allowDisable: + logger.Errorf("Invalid %s %q in googet.conf, using default: must be positive", name, s) + return 0 + } + return d } diff --git a/settings/settings_test.go b/settings/settings_test.go index 9942131..d93f0f1 100644 --- a/settings/settings_test.go +++ b/settings/settings_test.go @@ -8,6 +8,7 @@ import ( "github.com/google/go-cmp/cmp" "github.com/google/googet/v2/settings" + "github.com/google/googet/v2/supervisor" ) func TestInitialize(t *testing.T) { @@ -58,3 +59,80 @@ func TestInitialize(t *testing.T) { } }) } + +func TestInitializeSupervisorSettings(t *testing.T) { + tests := []struct { + name string + content string + wantMode supervisor.Mode + wantInactivity, wantInstall time.Duration + wantUIGrace, wantStall time.Duration + wantUIDetection bool + }{ + { + name: "defaults when unset", + content: "archs: [noarch]", + wantUIDetection: true, + }, + { + name: "explicit values", + content: "archs: [noarch]\nsupervisormode: monitor\ninactivitytimeout: 10m\ninstalltimeout: 3h\nuigraceperiod: 45s\ndownloadstalltimeout: 5m\nuidetection: false", + wantMode: supervisor.ModeMonitor, + wantInactivity: 10 * time.Minute, + wantInstall: 3 * time.Hour, + wantUIGrace: 45 * time.Second, + wantStall: 5 * time.Minute, + wantUIDetection: false, + }, + { + name: "mode off", + content: "archs: [noarch]\nsupervisormode: OFF", + wantMode: supervisor.ModeOff, + wantUIDetection: true, + }, + { + name: "zero disables timeouts", + content: "archs: [noarch]\ninactivitytimeout: 0\ninstalltimeout: 0s", + wantInactivity: -1, + wantInstall: -1, + wantUIDetection: true, + }, + { + name: "zero does not disable stall or ui grace", + content: "archs: [noarch]\ndownloadstalltimeout: 0\nuigraceperiod: 0s", + wantUIDetection: true, + }, + { + name: "invalid values fall back to defaults", + content: "archs: [noarch]\nsupervisormode: aggressive\ninactivitytimeout: soon\ninstalltimeout: -1h\nuigraceperiod: 0\ndownloadstalltimeout: -5s", + wantUIDetection: true, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + rootDir := t.TempDir() + if err := os.WriteFile(filepath.Join(rootDir, "googet.conf"), []byte(tc.content), 0644); err != nil { + t.Fatalf("os.WriteFile: %v", err) + } + settings.Initialize(rootDir, false) + if settings.SupervisorMode != tc.wantMode { + t.Errorf("SupervisorMode = %v, want %v", settings.SupervisorMode, tc.wantMode) + } + if settings.InactivityTimeout != tc.wantInactivity { + t.Errorf("InactivityTimeout = %v, want %v", settings.InactivityTimeout, tc.wantInactivity) + } + if settings.InstallTimeout != tc.wantInstall { + t.Errorf("InstallTimeout = %v, want %v", settings.InstallTimeout, tc.wantInstall) + } + if settings.UIGracePeriod != tc.wantUIGrace { + t.Errorf("UIGracePeriod = %v, want %v", settings.UIGracePeriod, tc.wantUIGrace) + } + if settings.DownloadStallTimeout != tc.wantStall { + t.Errorf("DownloadStallTimeout = %v, want %v", settings.DownloadStallTimeout, tc.wantStall) + } + if settings.UIDetection != tc.wantUIDetection { + t.Errorf("UIDetection = %v, want %v", settings.UIDetection, tc.wantUIDetection) + } + }) + } +} diff --git a/supervisor/msi.go b/supervisor/msi.go new file mode 100644 index 0000000..c65b577 --- /dev/null +++ b/supervisor/msi.go @@ -0,0 +1,113 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "time" + + "github.com/google/logger" +) + +const ( + // msiMutexName names the mutex the Windows Installer service holds during a transaction. + msiMutexName = `Global\_MSIExecute` + // serviceProgressLogInterval is how often the post-abort wait logs while waiting. + serviceProgressLogInterval = 30 * time.Second + // minAfterAbortPoll bounds how often the post-abort wait polls. + minAfterAbortPoll = 10 * time.Millisecond +) + +// serviceObservation is one observation of the Windows Installer service tree. +type serviceObservation struct { + // lastProgress is when the service tree last made forward progress. + lastProgress time.Time + // lastReason describes the signal that last counted as progress. + lastReason string + // gateOpened is when the _MSIExecute mutex was last observed to appear, or the zero time. + gateOpened time.Time + // terminable lists service-tree descendants created for this install. + terminable []uint32 +} + +// msiWaitEnv provides the clock, the _MSIExecute mutex check, service-tree sampling and +// process termination to waitForMSITransaction, so that its decisions can be tested on every +// platform. +type msiWaitEnv struct { + now func() time.Time + sleep func(time.Duration) + mutexHeld func() bool + observe func(now time.Time) serviceObservation + terminate func(pids []uint32) +} + +// latest returns the latest of ts. +func latest(ts ...time.Time) time.Time { + var out time.Time + for _, t := range ts { + if t.After(out) { + out = t + } + } + return out +} + +// waitForMSITransaction runs after a job that ran msiexec was terminated. It waits up to +// opts.msiMutexWait for the service-side transaction to finish so the next package does not +// collide with it. +// +// The service process itself is never terminated. Descendants created for this install are +// terminated once, and only after the whole service tree was observed idle for a full +// InactivityTimeout after the wait started and after the transaction's mutex appeared. Idle +// time is measured from the latest of the last progress, the start of the wait and the time the +// mutex was last seen to appear, so progress that could not be observed before the abort, for +// example while the mutex was absent or before the first gated sample primed the baseline, +// never makes the tree look idle. Service-side work that is making progress is never +// terminated. +func waitForMSITransaction(opts Options, env msiWaitEnv) { + if opts.msiMutexWait < 0 { + return + } + start := env.now() + deadline := start.Add(opts.msiMutexWait) + interval := opts.pollInterval + if interval < minAfterAbortPoll { + interval = minAfterAbortPoll + } + terminated := false + var lastLog time.Time + for { + now := env.now() + if !env.mutexHeld() { + logger.Infof("Windows Installer transaction finished (%s released).", msiMutexName) + return + } + obs := env.observe(now) + idle := now.Sub(latest(obs.lastProgress, start, obs.gateOpened)) + if !terminated && opts.InactivityTimeout > 0 && idle >= opts.InactivityTimeout && len(obs.terminable) > 0 { + terminated = true + logger.Errorf("Windows Installer service tree observed inactive for %v; terminating processes %v created for this install. The service process is not terminated.", idle.Round(time.Second), obs.terminable) + env.terminate(obs.terminable) + } + if !now.Before(deadline) { + logger.Warningf("Windows Installer transaction still running after %v; continuing. Service-side processes: %v.", opts.msiMutexWait, obs.terminable) + return + } + if now.Sub(lastLog) >= serviceProgressLogInterval { + lastLog = now + logger.Infof("Waiting for Windows Installer transaction to finish: observed idle for %v (last progress: %s).", + idle.Round(time.Second), obs.lastReason) + } + env.sleep(interval) + } +} diff --git a/supervisor/msi_test.go b/supervisor/msi_test.go new file mode 100644 index 0000000..b683a51 --- /dev/null +++ b/supervisor/msi_test.go @@ -0,0 +1,189 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "slices" + "testing" + "time" +) + +// fakeMSIService scripts the Windows Installer service for waitForMSITransaction on a +// synthetic clock that advances only when the wait sleeps. +type fakeMSIService struct { + now time.Time + // mutexReleasedAt, if non-zero, is when _MSIExecute is released. + mutexReleasedAt time.Time + // lastProgress returns the service tree's last progress time as of now. + lastProgress func(now time.Time) time.Time + gateOpened time.Time + terminable []uint32 + + observations int + terminated [][]uint32 + terminatedAt []time.Time +} + +// env returns hooks backed by f. +func (f *fakeMSIService) env() msiWaitEnv { + return msiWaitEnv{ + now: func() time.Time { return f.now }, + sleep: func(d time.Duration) { f.now = f.now.Add(d) }, + mutexHeld: func() bool { + return f.mutexReleasedAt.IsZero() || f.now.Before(f.mutexReleasedAt) + }, + observe: func(now time.Time) serviceObservation { + f.observations++ + return serviceObservation{ + lastProgress: f.lastProgress(now), + lastReason: "scripted", + gateOpened: f.gateOpened, + terminable: f.terminable, + } + }, + terminate: func(pids []uint32) { + f.terminated = append(f.terminated, pids) + f.terminatedAt = append(f.terminatedAt, f.now) + }, + } +} + +// msiTestOptions returns the built-in post-abort settings: a 5m inactivity timeout, a 10m +// mutex wait and a 2s poll. +func msiTestOptions() Options { + return testOptions(Options{}) +} + +// supervisionStart and abortAt are when the tests' supervised install started and when it was +// aborted, 30m later. +var ( + supervisionStart = time.Date(2026, 9, 28, 12, 0, 0, 0, time.UTC) + abortAt = supervisionStart.Add(30 * time.Minute) +) + +// constantTime returns a lastProgress script that always reports t. +func constantTime(t time.Time) func(time.Time) time.Time { + return func(time.Time) time.Time { return t } +} + +// TestWaitForMSITransaction_AbortRightAfterTransactionStart verifies that an abort right after +// the transaction started, when the tracker's last progress is still the supervision start, +// does not terminate custom action hosts until the tree was observed idle for a full +// InactivityTimeout after the abort. +func TestWaitForMSITransaction_AbortRightAfterTransactionStart(t *testing.T) { + f := &fakeMSIService{ + now: abortAt, + lastProgress: constantTime(supervisionStart), + gateOpened: abortAt.Add(-2 * time.Second), + terminable: []uint32{200, 300}, + } + waitForMSITransaction(msiTestOptions(), f.env()) + if len(f.terminated) != 1 { + t.Fatalf("Terminated %d time(s), want exactly once", len(f.terminated)) + } + if want := abortAt.Add(defaultInactivityTimeout); !f.terminatedAt[0].Equal(want) { + t.Errorf("Terminated at %v after the abort, want %v", f.terminatedAt[0].Sub(abortAt), defaultInactivityTimeout) + } + if !slices.Equal(f.terminated[0], []uint32{200, 300}) { + t.Errorf("Terminated %v, want [200 300]", f.terminated[0]) + } +} + +// TestWaitForMSITransaction_GateOpenedDuringWait verifies that idle time is measured from when +// the transaction's mutex was last seen to appear if that is later than the abort. +func TestWaitForMSITransaction_GateOpenedDuringWait(t *testing.T) { + gate := abortAt.Add(4 * time.Minute) + f := &fakeMSIService{ + now: abortAt, + lastProgress: constantTime(supervisionStart), + gateOpened: gate, + terminable: []uint32{200}, + } + waitForMSITransaction(msiTestOptions(), f.env()) + if len(f.terminatedAt) != 1 || !f.terminatedAt[0].Equal(gate.Add(defaultInactivityTimeout)) { + t.Errorf("Terminated at %v, want once at %v", f.terminatedAt, gate.Add(defaultInactivityTimeout)) + } +} + +// TestWaitForMSITransaction_ProgressingTreeNeverTerminated verifies that service-side work that +// keeps making progress is never terminated, and that the wait gives up at msiMutexWait. +func TestWaitForMSITransaction_ProgressingTreeNeverTerminated(t *testing.T) { + f := &fakeMSIService{ + now: abortAt, + lastProgress: func(now time.Time) time.Time { return now }, + terminable: []uint32{200}, + } + waitForMSITransaction(msiTestOptions(), f.env()) + if len(f.terminated) != 0 { + t.Errorf("Terminated a progressing service tree %d time(s) at %v, want never", len(f.terminated), f.terminatedAt) + } + if want := abortAt.Add(defaultMSIMutexWait); !f.now.Equal(want) { + t.Errorf("Wait returned %v after the abort, want at msiMutexWait %v", f.now.Sub(abortAt), defaultMSIMutexWait) + } +} + +// TestWaitForMSITransaction_TerminatesOnceWhenIdle verifies a single termination one +// InactivityTimeout after the tree stopped progressing, even though it stays idle afterwards. +func TestWaitForMSITransaction_TerminatesOnceWhenIdle(t *testing.T) { + stopped := abortAt.Add(2 * time.Minute) + f := &fakeMSIService{ + now: abortAt, + lastProgress: func(now time.Time) time.Time { + if now.After(stopped) { + return stopped + } + return now + }, + terminable: []uint32{200}, + } + waitForMSITransaction(msiTestOptions(), f.env()) + if len(f.terminatedAt) != 1 || !f.terminatedAt[0].Equal(stopped.Add(defaultInactivityTimeout)) { + t.Errorf("Terminated at %v, want once at %v", f.terminatedAt, stopped.Add(defaultInactivityTimeout)) + } +} + +// TestWaitForMSITransaction_NothingTerminable verifies that an idle tree with no descendant +// created for this install is left alone, so the service process is never terminated. +func TestWaitForMSITransaction_NothingTerminable(t *testing.T) { + f := &fakeMSIService{now: abortAt, lastProgress: constantTime(supervisionStart)} + waitForMSITransaction(msiTestOptions(), f.env()) + if len(f.terminated) != 0 { + t.Errorf("Terminated %v with nothing terminable, want no termination", f.terminated) + } +} + +// TestWaitForMSITransaction_MutexWait verifies that the wait ends as soon as _MSIExecute is +// released, and not at all when it is disabled. +func TestWaitForMSITransaction_MutexWait(t *testing.T) { + released := abortAt.Add(3 * time.Minute) + f := &fakeMSIService{ + now: abortAt, + mutexReleasedAt: released, + lastProgress: constantTime(supervisionStart), + terminable: []uint32{200}, + } + waitForMSITransaction(msiTestOptions(), f.env()) + if !f.now.Equal(released) { + t.Errorf("Wait returned %v after the abort, want %v when the mutex was released", f.now.Sub(abortAt), released.Sub(abortAt)) + } + if len(f.terminated) != 0 { + t.Errorf("Terminated %v before the tree was observed idle for InactivityTimeout, want none", f.terminated) + } + + disabled := &fakeMSIService{now: abortAt, lastProgress: constantTime(supervisionStart), terminable: []uint32{200}} + waitForMSITransaction(testOptions(Options{msiMutexWait: -1}), disabled.env()) + if disabled.observations != 0 || !disabled.now.Equal(abortAt) { + t.Errorf("Disabled wait took %d observation(s) and ran until %v, want none and no waiting", disabled.observations, disabled.now.Sub(abortAt)) + } +} diff --git a/supervisor/msi_windows.go b/supervisor/msi_windows.go new file mode 100644 index 0000000..a475af3 --- /dev/null +++ b/supervisor/msi_windows.go @@ -0,0 +1,298 @@ +//go:build windows + +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "errors" + "time" + "unsafe" + + "github.com/google/logger" + "golang.org/x/sys/windows" +) + +const ( + // msiServiceName is the service name of the Windows Installer service. + msiServiceName = "msiserver" + // trustedInstallerServiceName is the service name of the Windows Modules Installer, which + // performs servicing for wusa and DISM. + trustedInstallerServiceName = "TrustedInstaller" + // creationTimeSlack tolerates clock granularity when comparing creation times with the + // supervision start. + creationTimeSlack = 2 * time.Second +) + +var ( + kernel32 = windows.NewLazySystemDLL("kernel32.dll") + procGetProcessIoCounters = kernel32.NewProc("GetProcessIoCounters") +) + +// mutexExists reports whether a named mutex exists. ERROR_ACCESS_DENIED means the mutex exists +// but its DACL does not grant SYNCHRONIZE to this process. +func mutexExists(name string) bool { + p, err := windows.UTF16PtrFromString(name) + if err != nil { + return false + } + h, err := windows.OpenMutex(windows.SYNCHRONIZE, false, p) + if err != nil { + return errors.Is(err, windows.ERROR_ACCESS_DENIED) + } + windows.CloseHandle(h) + return true +} + +// msiMutexHeld reports whether a Windows Installer transaction is in progress. +func msiMutexHeld() bool { + return mutexExists(msiMutexName) +} + +// serviceProcessID returns the PID of the named service, or 0 if it is not running. +func serviceProcessID(service string) (uint32, error) { + scm, err := windows.OpenSCManager(nil, nil, windows.SC_MANAGER_CONNECT) + if err != nil { + return 0, err + } + defer windows.CloseServiceHandle(scm) + name, err := windows.UTF16PtrFromString(service) + if err != nil { + return 0, err + } + svc, err := windows.OpenService(scm, name, windows.SERVICE_QUERY_STATUS) + if err != nil { + return 0, err + } + defer windows.CloseServiceHandle(svc) + var status windows.SERVICE_STATUS_PROCESS + var needed uint32 + if err := windows.QueryServiceStatusEx(svc, windows.SC_STATUS_PROCESS_INFO, + (*byte)(unsafe.Pointer(&status)), uint32(unsafe.Sizeof(status)), &needed); err != nil { + return 0, err + } + return status.ProcessId, nil +} + +// snapshotProcesses returns the current process table. +func snapshotProcesses() ([]procEntry, error) { + snap, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPPROCESS, 0) + if err != nil { + return nil, err + } + defer windows.CloseHandle(snap) + pe := windows.ProcessEntry32{Size: uint32(unsafe.Sizeof(windows.ProcessEntry32{}))} + var entries []procEntry + for err = windows.Process32First(snap, &pe); err == nil; err = windows.Process32Next(snap, &pe) { + entries = append(entries, procEntry{ + pid: pe.ProcessID, + ppid: pe.ParentProcessID, + exe: windows.UTF16ToString(pe.ExeFile[:]), + }) + } + return entries, nil +} + +// handleCounters returns cumulative CPU time, read plus write bytes and the creation time of the +// process behind h, which needs PROCESS_QUERY_LIMITED_INFORMATION access. +func handleCounters(h windows.Handle) (counters, time.Time, bool) { + var creation, exit, kernel, user windows.Filetime + if err := windows.GetProcessTimes(h, &creation, &exit, &kernel, &user); err != nil { + return counters{}, time.Time{}, false + } + c := counters{cpu: filetimeDuration(kernel) + filetimeDuration(user)} + var ioc windows.IO_COUNTERS + if r1, _, _ := procGetProcessIoCounters.Call(uintptr(h), uintptr(unsafe.Pointer(&ioc))); r1 != 0 { + c.io = ioc.ReadTransferCount + ioc.WriteTransferCount + } + return c, time.Unix(0, creation.Nanoseconds()), true +} + +// processCounters returns cumulative CPU time, read plus write bytes and the creation time of pid. +func processCounters(pid uint32) (counters, time.Time, bool) { + h, err := windows.OpenProcess(windows.PROCESS_QUERY_LIMITED_INFORMATION, false, pid) + if err != nil { + return counters{}, time.Time{}, false + } + defer windows.CloseHandle(h) + return handleCounters(h) +} + +// serviceTreeSample is one observation of a service process tree. +type serviceTreeSample struct { + perPID map[procKey]counters + terminable []uint32 +} + +// sampleServiceTree measures the service process, every process whose image is one of +// extraRoots, and all of their transitive descendants. If withTerminable is set, it also lists +// as terminable the descendants created at or after notBefore, never a root. +func sampleServiceTree(service string, extraRoots []string, withTerminable bool, notBefore time.Time) serviceTreeSample { + out := serviceTreeSample{perPID: make(map[procKey]counters)} + svc, err := serviceProcessID(service) + if err != nil { + svc = 0 + } + entries, err := snapshotProcesses() + if err != nil { + return out + } + roots := []uint32{svc} + for _, image := range extraRoots { + roots = append(roots, pidsWithImage(entries, image)...) + } + created := make(map[uint32]time.Time) + measured := make(map[uint32]bool) + // createdAt measures pid on first use: it records the process's counters in out.perPID, + // keyed on its PID and creation time, and memoizes the creation time. + createdAt := func(pid uint32) (time.Time, bool) { + if !measured[pid] { + measured[pid] = true + if c, t, ok := processCounters(pid); ok { + out.perPID[procKey{pid, t.UnixNano()}] = c + created[pid] = t + } + } + t, ok := created[pid] + return t, ok + } + // processForest calls createdAt for every PID it returns, so every returned PID whose + // counters are readable is also measured. + tree := processForest(entries, roots, createdAt) + if withTerminable { + out.terminable = selectTerminable(tree, roots, createdAt, notBefore) + } + return out +} + +// serviceMonitor attributes work done by a Windows service's process tree to the supervised +// install. +// +// Some installers only hand their work to a service that runs outside the Job Object: msiexec +// hands the transaction to the msiserver service, and wusa hands servicing to TrustedInstaller +// and TiWorker. The CPU and I/O of the whole service tree, including the service process +// itself, count as progress of the supervised install. +// +// The accounting cannot tell which install a service is working for. Activity for an +// unrelated install, such as an SCCM or Windows Update transaction running at the same time, +// is also credited as progress. This can only delay a watchdog termination, never cause one. +type serviceMonitor struct { + service string + extraRoots []string + // gate, if non-nil, reports whether the service is currently working for an install. While + // it returns false nothing is sampled and the baseline is reset. + gate func() bool + // gateOpen and gateOpenedAt record whether the gate was open at the previous sample and + // when it was last seen to open. + gateOpen bool + gateOpenedAt time.Time + // canTerminate enables the computation of terminable processes. Monitors whose processes + // are never terminated leave it unset. + canTerminate bool + start time.Time + agg pidCounters + cum counters + tracker *progressTracker + // terminable lists descendants created after start, as of the latest sample. + terminable []uint32 +} + +// newServiceMonitor returns a monitor for service that starts attributing work at start. +func newServiceMonitor(service string, extraRoots []string, gate func() bool, canTerminate bool, opts Options, start time.Time) *serviceMonitor { + return &serviceMonitor{ + service: service, + extraRoots: extraRoots, + gate: gate, + canTerminate: canTerminate, + start: start, + tracker: newProgressTracker(opts.progressWindow, opts.minCPUDelta, opts.minIODelta, start), + } +} + +// newMSIMonitor returns a monitor for the Windows Installer service, gated on _MSIExecute. Its +// descendants created for the install may be terminated after an abort. +func newMSIMonitor(opts Options, start time.Time) *serviceMonitor { + return newServiceMonitor(msiServiceName, nil, msiMutexHeld, true, opts, start) +} + +// newTrustedInstallerMonitor returns a monitor for the servicing stack. Its processes are never +// terminated. It has no gate, so any servicing activity on the machine counts as progress. +func newTrustedInstallerMonitor(opts Options, start time.Time) *serviceMonitor { + return newServiceMonitor(trustedInstallerServiceName, []string{tiWorkerImage}, nil, false, opts, start) +} + +// sample returns the cumulative service-side counters attributed to the install so far. +func (m *serviceMonitor) sample(now time.Time) counters { + m.accumulate(now) + m.tracker.observe(progressSample{at: now, cpu: m.cum.cpu, io: m.cum.io}, nil) + return m.cum +} + +// accumulate adds the service tree's counter increases since the previous call to m.cum. +func (m *serviceMonitor) accumulate(now time.Time) { + if m.gate != nil { + if !m.gate() { + // A new transaction gets a fresh baseline. + m.gateOpen = false + m.terminable = nil + m.agg.reset() + return + } + if !m.gateOpen { + m.gateOpen = true + m.gateOpenedAt = now + } + } + s := sampleServiceTree(m.service, m.extraRoots, m.canTerminate, m.start.Add(-creationTimeSlack)) + m.terminable = s.terminable + dCPU, dIO := m.agg.update(s.perPID) + m.cum.cpu += dCPU + m.cum.io += dIO +} + +// observe samples the service tree and returns the observation used by the post-abort wait. +func (m *serviceMonitor) observe(now time.Time) serviceObservation { + m.sample(now) + return serviceObservation{ + lastProgress: m.tracker.lastProgress, + lastReason: m.tracker.lastReason, + gateOpened: m.gateOpenedAt, + terminable: m.terminable, + } +} + +// afterMSIAbort runs after a job that ran msiexec was terminated; see waitForMSITransaction. +func (m *serviceMonitor) afterMSIAbort(opts Options) { + waitForMSITransaction(opts, msiWaitEnv{ + now: time.Now, + sleep: time.Sleep, + mutexHeld: msiMutexHeld, + observe: m.observe, + terminate: terminateProcesses, + }) +} + +// terminateProcesses terminates each process in pids, ignoring processes that already exited. +func terminateProcesses(pids []uint32) { + for _, pid := range pids { + h, err := windows.OpenProcess(windows.PROCESS_TERMINATE, false, pid) + if err != nil { + continue + } + if err := windows.TerminateProcess(h, 1); err != nil { + logger.Warningf("Failed to terminate service-side Windows Installer process %d: %v", pid, err) + } + windows.CloseHandle(h) + } +} diff --git a/supervisor/proctree.go b/supervisor/proctree.go new file mode 100644 index 0000000..6e00ead --- /dev/null +++ b/supervisor/proctree.go @@ -0,0 +1,135 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "strings" + "time" +) + +// Image names that identify installers whose real work runs in a Windows service outside the +// supervised process tree. +const ( + msiexecImage = "msiexec.exe" + wusaImage = "wusa.exe" + dismImage = "dism.exe" + tiWorkerImage = "tiworker.exe" +) + +// imageBaseName returns the lower-cased file name of path, without directories or surrounding +// quotes. +func imageBaseName(path string) string { + path = strings.Trim(path, `"`) + if i := strings.LastIndexAny(path, `\/`); i >= 0 { + path = path[i+1:] + } + return strings.ToLower(path) +} + +// isCommand reports whether path names image, with or without the .exe extension. +func isCommand(path, image string) bool { + name := imageBaseName(path) + return name == image || name+".exe" == image +} + +// isMsiexecCommand reports whether path names the Windows Installer client, msiexec(.exe). +func isMsiexecCommand(path string) bool { + return isCommand(path, msiexecImage) +} + +// isServicingCommand reports whether path names wusa(.exe) or dism(.exe), whose work runs in the +// TrustedInstaller servicing stack. +func isServicingCommand(path string) bool { + return isCommand(path, wusaImage) || isCommand(path, dismImage) +} + +// procEntry is a minimal process table entry. +type procEntry struct { + pid uint32 + ppid uint32 + exe string +} + +// pidsWithImage returns the PIDs of entries whose image base name equals image, ignoring case. +func pidsWithImage(entries []procEntry, image string) []uint32 { + var pids []uint32 + for _, e := range entries { + if imageBaseName(e.exe) == image { + pids = append(pids, e.pid) + } + } + return pids +} + +// processForest returns every root followed by all transitive descendants found in entries, in +// breadth-first order. +// +// created returns the creation time of a process, or false if it is unknown. A child is included +// only if its creation time is known and is not before its parent's, which guards against a PID +// that was reused after the original parent exited. The comparison is skipped when the parent's +// creation time is unknown. +func processForest(entries []procEntry, roots []uint32, created func(uint32) (time.Time, bool)) []uint32 { + children := make(map[uint32][]uint32) + for _, e := range entries { + if e.pid != e.ppid { + children[e.ppid] = append(children[e.ppid], e.pid) + } + } + seen := make(map[uint32]bool) + var out, queue []uint32 + for _, r := range roots { + if r != 0 && !seen[r] { + seen[r] = true + out = append(out, r) + queue = append(queue, r) + } + } + for len(queue) > 0 { + parent := queue[0] + queue = queue[1:] + parentCreated, parentKnown := created(parent) + for _, c := range children[parent] { + if seen[c] { + continue + } + childCreated, ok := created(c) + if !ok || (parentKnown && childCreated.Before(parentCreated)) { + continue + } + seen[c] = true + out = append(out, c) + queue = append(queue, c) + } + } + return out +} + +// selectTerminable returns the PIDs in tree that may be terminated after an abort: processes +// created at or after notBefore, excluding every protected PID such as the service process. +func selectTerminable(tree []uint32, protected []uint32, created func(uint32) (time.Time, bool), notBefore time.Time) []uint32 { + skip := make(map[uint32]bool, len(protected)) + for _, p := range protected { + skip[p] = true + } + var out []uint32 + for _, pid := range tree { + if skip[pid] { + continue + } + if t, ok := created(pid); ok && !t.Before(notBefore) { + out = append(out, pid) + } + } + return out +} diff --git a/supervisor/progress.go b/supervisor/progress.go new file mode 100644 index 0000000..f34aa35 --- /dev/null +++ b/supervisor/progress.go @@ -0,0 +1,493 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "errors" + "fmt" + "io" + "os" + "os/exec" + "sort" + "time" + + "github.com/google/logger" +) + +var ( + // defaultWaitDelay is applied to exec.Cmd.WaitDelay when the caller left it unset. + defaultWaitDelay = 30 * time.Second + + // killWaitTimeout bounds how long to wait for Cmd.Wait to return after terminating the tree. + killWaitTimeout = 60 * time.Second +) + +// counters holds cumulative CPU time and I/O bytes for one process. +type counters struct { + cpu time.Duration + io uint64 +} + +// procKey identifies one process instance. created is the process creation time in an +// OS-specific unit, or 0 if unknown; it distinguishes a reused PID from the original process. +type procKey struct { + pid uint32 + created int64 +} + +// pidCounters converts per-process cumulative counters into monotonic aggregate increments. +// +// Processes that exit simply stop contributing, so the aggregate never decreases. Processes seen +// on the first update after construction or reset form the baseline and contribute nothing; +// processes that appear later contribute their full counters because they started after the +// baseline was taken. +// +// The last-seen counters of a process that is missing from a sample are kept, so a process +// that transiently drops out, for example because it could not be opened for one poll, is +// credited only with its increase when it reappears rather than with its lifetime counters. +// A reused PID has a different creation time and is therefore a new process. +type pidCounters struct { + primed bool + last map[procKey]counters +} + +// update records the current per-process counters and returns the aggregate increase since the +// previous update. +func (p *pidCounters) update(cur map[procKey]counters) (time.Duration, uint64) { + var dCPU time.Duration + var dIO uint64 + if p.last == nil { + p.last = make(map[procKey]counters, len(cur)) + } + for k, c := range cur { + if p.primed { + prev := p.last[k] + if c.cpu > prev.cpu { + dCPU += c.cpu - prev.cpu + } + if c.io > prev.io { + dIO += c.io - prev.io + } + } + p.last[k] = c + } + p.primed = true + return dCPU, dIO +} + +// reset discards the baseline so that the next update primes again. +func (p *pidCounters) reset() { + p.primed = false + p.last = nil +} + +// progressSample is a single observation of cumulative, monotonic resource counters. +type progressSample struct { + at time.Time + cpu time.Duration + io uint64 +} + +// progressTracker decides whether a process tree is making meaningful forward progress. +// +// CPU and I/O are compared against thresholds over a rolling window so that timer or animation +// noise within a single poll does not count. Log file growth is a zero-threshold signal: any +// growth between two consecutive samples, or a monitored file appearing, counts as progress. +type progressTracker struct { + window time.Duration + minCPU time.Duration + minIO uint64 + samples []progressSample + logSizes map[string]int64 + logGrewAt time.Time + lastProgress time.Time + lastReason string +} + +// newProgressTracker returns a tracker whose last progress time is start. +func newProgressTracker(window, minCPU time.Duration, minIO uint64, start time.Time) *progressTracker { + return &progressTracker{ + window: window, + minCPU: minCPU, + minIO: minIO, + lastProgress: start, + } +} + +// exceeds reports whether the counter increase from base to last meets a threshold, with a +// description of the signal that did. +func (t *progressTracker) exceeds(base, last progressSample) (bool, string) { + span := last.at.Sub(base.at) + if last.cpu > base.cpu { + if d := last.cpu - base.cpu; t.minCPU >= 0 && d >= t.minCPU { + return true, fmt.Sprintf("CPU +%v within %v", d, span) + } + } + if last.io > base.io { + if d := last.io - base.io; d >= t.minIO { + return true, fmt.Sprintf("I/O +%d bytes within %v", d, span) + } + } + return false, "" +} + +// observe records a sample and the current log file sizes, and reports whether it constitutes +// forward progress. A nil logSizes map means log files were not sampled. +func (t *progressTracker) observe(s progressSample, logSizes map[string]int64) bool { + progressed := false + reason := "" + + if logSizes != nil { + if t.logSizes != nil { + for path, size := range logSizes { + prev, ok := t.logSizes[path] + if !ok || size > prev { + progressed = true + reason = fmt.Sprintf("log %s grew to %d bytes", path, size) + t.logGrewAt = s.at + break + } + } + } + t.logSizes = logSizes + } + + t.samples = append(t.samples, s) + // Keep exactly one sample at or before the window start so the baseline spans the full window. + cutoff := s.at.Add(-t.window) + drop := 0 + for drop+1 < len(t.samples) && !t.samples[drop+1].at.After(cutoff) { + drop++ + } + if drop > 0 { + t.samples = append(t.samples[:0], t.samples[drop:]...) + } + + if !progressed { + progressed, reason = t.exceeds(t.samples[0], s) + } + if progressed { + t.lastProgress = s.at + t.lastReason = reason + } + return progressed +} + +// progressedSince reports whether forward progress was observed strictly after since: log growth +// at a later sample, or CPU or I/O deltas at or above the thresholds measured from the first +// retained sample taken at or after since. The rolling window still bounds the baseline, so a +// burst that ended before since never counts. +func (t *progressTracker) progressedSince(since time.Time) bool { + if len(t.samples) == 0 { + return false + } + last := t.samples[len(t.samples)-1] + if !last.at.After(since) { + return false + } + if t.logGrewAt.After(since) { + return true + } + base := t.samples[0] + for _, s := range t.samples { + if !s.at.Before(since) { + base = s + break + } + } + ok, _ := t.exceeds(base, last) + return ok +} + +// windowDeltas returns the CPU and I/O increase across the currently retained window. +func (t *progressTracker) windowDeltas() (time.Duration, uint64, time.Duration) { + if len(t.samples) == 0 { + return 0, 0, 0 + } + first, last := t.samples[0], t.samples[len(t.samples)-1] + var dCPU time.Duration + var dIO uint64 + if last.cpu > first.cpu { + dCPU = last.cpu - first.cpu + } + if last.io > first.io { + dIO = last.io - first.io + } + return dCPU, dIO, last.at.Sub(first.at) +} + +// sampleLogSizes returns the current size of every existing file in paths. +func sampleLogSizes(paths []string) map[string]int64 { + sizes := make(map[string]int64, len(paths)) + for _, p := range paths { + if fi, err := os.Stat(p); err == nil { + sizes[p] = fi.Size() + } + } + return sizes +} + +// trackedWindow records when a candidate window was first observed. +type trackedWindow struct { + firstSeen time.Time + info windowInfo +} + +// uiSnifferState tracks candidate interactive windows across polling ticks. +// +// Windows are tracked per HWND so that Z-order changes do not reset their timers. A window +// disappearing forgets it. Forward progress observed after a window was first seen resets that +// window's timer, so an abort requires the window to persist for the grace period with no +// progress since it appeared. Progress that happened before the window appeared does not delay +// the abort. +type uiSnifferState struct { + tracked map[uintptr]trackedWindow +} + +// check updates the tracked windows and reports whether one has persisted for at least grace +// without intervening progress, along with diagnostics describing it. progressedSince reports +// whether progress was observed after the given time. +func (s *uiSnifferState) check(wins []windowInfo, grace time.Duration, now time.Time, progressedSince func(time.Time) bool) (bool, string) { + if len(wins) == 0 { + s.tracked = nil + return false, "" + } + next := make(map[uintptr]trackedWindow, len(wins)) + for _, w := range wins { + tw, ok := s.tracked[w.HWND] + if !ok || progressedSince(tw.firstSeen) { + tw.firstSeen = now + } + tw.info = w + next[w.HWND] = tw + } + s.tracked = next + + // Report the longest-persisting window, breaking ties by HWND for deterministic output. + var best *trackedWindow + for _, h := range sortedHWNDs(next) { + tw := next[h] + if best == nil || tw.firstSeen.Before(best.firstSeen) { + best = &tw + } + } + if best != nil && now.Sub(best.firstSeen) >= grace { + w := best.info + return true, fmt.Sprintf("window persisted %v with no forward progress: PID=%d HWND=0x%x Title=%q Class=%q Image=%q", + now.Sub(best.firstSeen).Round(time.Millisecond), w.PID, w.HWND, w.Title, w.ClassName, w.ExePath) + } + return false, "" +} + +// sortedHWNDs returns the keys of m in ascending order. +func sortedHWNDs(m map[uintptr]trackedWindow) []uintptr { + keys := make([]uintptr, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + sort.Slice(keys, func(i, j int) bool { return keys[i] < keys[j] }) + return keys +} + +// watchdog combines the progress tracker, the UI sniffer and the hard timeout into abort +// decisions, and implements monitor-mode logging. +type watchdog struct { + opts Options + start time.Time + tracker *progressTracker + sniffer uiSnifferState + wouldKill map[error]bool +} + +// newWatchdog returns a watchdog for resolved options starting at start. +func newWatchdog(opts Options, start time.Time) *watchdog { + return &watchdog{ + opts: opts, + start: start, + tracker: newProgressTracker(opts.progressWindow, opts.minCPUDelta, opts.minIODelta, start), + wouldKill: make(map[error]bool), + } +} + +// uiEnabled reports whether interactive window detection is active. +func (w *watchdog) uiEnabled() bool { + return w.opts.Unattended && !w.opts.DisableUIDetection && w.opts.UIGracePeriod >= 0 +} + +// abortDecision describes why the watchdog wants to terminate the process tree. +type abortDecision struct { + reason error + details string +} + +// observe records a sample and returns an abort decision, or nil if the tree should keep +// running. The wins argument is ignored unless UI detection is enabled. +func (w *watchdog) observe(s progressSample, logSizes map[string]int64, wins []windowInfo) *abortDecision { + now := s.at + w.tracker.observe(s, logSizes) + idle := now.Sub(w.tracker.lastProgress) + + if w.opts.HardTimeout > 0 && now.Sub(w.start) >= w.opts.HardTimeout { + return &abortDecision{ErrHardTimeout, fmt.Sprintf("runtime %v exceeded hard timeout %v; %s", + now.Sub(w.start).Round(time.Millisecond), w.opts.HardTimeout, w.diagnostics(now))} + } + if w.uiEnabled() { + if abort, details := w.sniffer.check(wins, w.opts.UIGracePeriod, now, w.tracker.progressedSince); abort { + return &abortDecision{ErrInteractiveUIDetected, details + "; " + w.diagnostics(now)} + } + } + if w.opts.InactivityTimeout > 0 && idle >= w.opts.InactivityTimeout { + return &abortDecision{ErrInactivityTimeout, fmt.Sprintf("no forward progress for %v (timeout %v); %s", + idle.Round(time.Millisecond), w.opts.InactivityTimeout, w.diagnostics(now))} + } + return nil +} + +// diagnostics summarizes recent progress for log messages. +func (w *watchdog) diagnostics(now time.Time) string { + dCPU, dIO, span := w.tracker.windowDeltas() + last := "none" + if w.tracker.lastReason != "" { + last = fmt.Sprintf("%s (%v ago)", w.tracker.lastReason, now.Sub(w.tracker.lastProgress).Round(time.Millisecond)) + } + return fmt.Sprintf("runtime=%v window=%v cpu=+%v (min %v) io=+%dB (min %dB) lastProgress=%s", + now.Sub(w.start).Round(time.Millisecond), span.Round(time.Millisecond), dCPU, w.opts.minCPUDelta, dIO, w.opts.minIODelta, last) +} + +// shouldTerminate applies the mode to an abort decision. In monitor mode it logs a single +// WOULD_KILL line per reason and returns false. A reason that clears may be logged again later. +func (w *watchdog) shouldTerminate(d *abortDecision) bool { + if d == nil { + for r := range w.wouldKill { + delete(w.wouldKill, r) + } + return false + } + if w.opts.Mode == ModeMonitor { + if !w.wouldKill[d.reason] { + w.wouldKill[d.reason] = true + logger.Warningf("WOULD_KILL: %v %s", d.reason, d.details) + } + return false + } + logger.Errorf("Terminating installer process tree: %v %s", d.reason, d.details) + return true +} + +// supervisedTree is the platform-specific view of a supervised process tree. +type supervisedTree interface { + // sample returns cumulative, monotonic counters for the tree at now. + sample(now time.Time) progressSample + // windows returns candidate interactive windows owned by the tree. + windows() []windowInfo + // terminate kills every process in the tree. + terminate() + // afterAbort runs after the tree was terminated and reaped, for example to wait for + // out-of-tree service work to settle. + afterAbort() + // rootExited reports whether the root process has exited, even if Wait has not returned yet + // because a descendant still holds the output pipes. + rootExited() bool +} + +// supervise runs the watchdog loop over tree until waitErr delivers the result of Cmd.Wait or +// the watchdog terminates the tree. Each value received from ticks is one poll at that time. +// +// Once the root process has exited no further abort decisions are made: Wait is then only +// waiting for descendants that inherited the output pipes, which exec.Cmd.WaitDelay bounds. +func supervise(opts Options, waitErr <-chan error, tree supervisedTree, start time.Time, ticks <-chan time.Time) error { + if opts.Mode == ModeOff { + return normalizeExitError(<-waitErr) + } + wd := newWatchdog(opts, start) + // Prime the baseline so pre-existing counters do not count as progress. + wd.observe(tree.sample(start), sampleLogSizes(opts.LogFiles), nil) + rootDone := false + for { + select { + case err := <-waitErr: + return normalizeExitError(err) + case now := <-ticks: + if rootDone { + continue + } + if tree.rootExited() { + rootDone = true + logger.Infof("Installer process exited; waiting for its output pipes to close without further watchdog decisions.") + continue + } + s := tree.sample(now) + var wins []windowInfo + if wd.uiEnabled() { + wins = tree.windows() + } + d := wd.observe(s, sampleLogSizes(opts.LogFiles), wins) + if wd.shouldTerminate(d) { + tree.terminate() + err := waitAfterTerminate(waitErr, d) + tree.afterAbort() + return err + } + } + } +} + +// setupStdio disconnects stdin and tees stdout and stderr to out when out is non-nil. It returns +// the opened null device, which the caller must close after the process exits. +func setupStdio(c *exec.Cmd, out io.Writer) (*os.File, error) { + devNull, err := os.Open(os.DevNull) + if err != nil { + return nil, fmt.Errorf("failed to open devNull for stdin disconnection: %w", err) + } + c.Stdin = devNull + + if out != nil { + c.Stdout = io.MultiWriter(os.Stdout, out) + c.Stderr = io.MultiWriter(os.Stderr, out) + } else { + if c.Stdout == nil { + c.Stdout = os.Stdout + } + if c.Stderr == nil { + c.Stderr = os.Stderr + } + } + if c.WaitDelay == 0 { + c.WaitDelay = defaultWaitDelay + } + return devNull, nil +} + +// normalizeExitError treats exec.ErrWaitDelay after a successful exit as success. The process +// itself exited cleanly; only a detached descendant still held the output pipes. +func normalizeExitError(err error) error { + if errors.Is(err, exec.ErrWaitDelay) { + logger.Warningf("Installer exited successfully but its output pipes were still held by another process; stopped waiting after WaitDelay.") + return nil + } + return err +} + +// waitAfterTerminate waits for the process to be reaped after the tree was terminated and +// returns reason wrapped with details. It never blocks longer than killWaitTimeout. +func waitAfterTerminate(waitErr <-chan error, d *abortDecision) error { + reason, details := d.reason, d.details + select { + case <-waitErr: + return fmt.Errorf("%w: %s", reason, details) + case <-time.After(killWaitTimeout): + logger.Errorf("Installer process tree terminated but Wait did not return within %v; its stdout/stderr pipes are likely held by a process outside the supervised tree.", killWaitTimeout) + return fmt.Errorf("%w: %s; wait abandoned after %v because stdout/stderr pipes are held by processes outside the supervised tree", reason, details, killWaitTimeout) + } +} diff --git a/supervisor/supervise_test.go b/supervisor/supervise_test.go new file mode 100644 index 0000000..ee5c046 --- /dev/null +++ b/supervisor/supervise_test.go @@ -0,0 +1,310 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "errors" + "strings" + "testing" + "time" +) + +// errFakeKilled is delivered as the Wait result when the fake tree is terminated. +var errFakeKilled = errors.New("fake process tree killed") + +// fakeTree is a scripted supervisedTree driven by synthetic ticks. +// +// supervise calls it only from its own goroutine; tests read the recorded fields after +// supervise returns, which the done channel orders. +type fakeTree struct { + start time.Time + // cpuAt returns cumulative CPU time at the given elapsed time; nil means no CPU at all. + cpuAt func(elapsed time.Duration) time.Duration + // winsAt returns the candidate windows at the given elapsed time; call counts windows calls + // starting at 1. + winsAt func(elapsed time.Duration, call int) []windowInfo + // exitAfterSamples makes rootExited return true once that many samples were taken; zero + // means the root never exits. + exitAfterSamples int + waitErr chan<- error + + now time.Time + samples int + windowCalls int + windowReports int + terminated int + terminatedAt time.Duration + afterAborts int +} + +// sample records the tick time and returns the scripted counters. +func (f *fakeTree) sample(now time.Time) progressSample { + f.now = now + f.samples++ + s := progressSample{at: now} + if f.cpuAt != nil { + s.cpu = f.cpuAt(now.Sub(f.start)) + } + return s +} + +// windows returns the scripted windows for the current tick. +func (f *fakeTree) windows() []windowInfo { + f.windowCalls++ + if f.winsAt == nil { + return nil + } + wins := f.winsAt(f.now.Sub(f.start), f.windowCalls) + if len(wins) > 0 { + f.windowReports++ + } + return wins +} + +// terminate records the termination and completes the fake Wait. +func (f *fakeTree) terminate() { + f.terminated++ + f.terminatedAt = f.now.Sub(f.start) + select { + case f.waitErr <- errFakeKilled: + default: + } +} + +// afterAbort records the call. +func (f *fakeTree) afterAbort() { f.afterAborts++ } + +// rootExited reports whether the scripted number of samples was reached. +func (f *fakeTree) rootExited() bool { + return f.exitAfterSamples > 0 && f.samples >= f.exitAfterSamples +} + +// testOptions fills unset fields from the built-in defaults without consulting process-wide +// defaults or Session 0, so tests are deterministic on every platform. +func testOptions(o Options) Options { + return mergeOptions(o, builtinDefaults()) +} + +// runFake drives supervise over tree with one tick every step up to and including end. If +// supervise has not returned by then, the fake root process exits successfully. +func runFake(opts Options, tree *fakeTree, step, end time.Duration) error { + if tree.start.IsZero() { + tree.start = time.Date(2026, 9, 28, 12, 0, 0, 0, time.UTC) + } + waitErr := make(chan error, 1) + tree.waitErr = waitErr + ticks := make(chan time.Time) + done := make(chan error, 1) + go func() { done <- supervise(opts, waitErr, tree, tree.start, ticks) }() + for at := step; step > 0 && at <= end; at += step { + select { + case ticks <- tree.start.Add(at): + case err := <-done: + return err + } + } + select { + case waitErr <- nil: + default: + } + return <-done +} + +// steadyCPU burns 500ms of CPU every 2s, well above the 250ms default threshold. +func steadyCPU(elapsed time.Duration) time.Duration { return elapsed / 4 } + +// cpuUntil returns a CPU script that burns like steadyCPU until stop and then stays flat. +func cpuUntil(stop time.Duration) func(time.Duration) time.Duration { + return func(elapsed time.Duration) time.Duration { + if elapsed > stop { + elapsed = stop + } + return steadyCPU(elapsed) + } +} + +// persistentWindow returns a window script that always reports one dialog. +func persistentWindow(hwnd uintptr, title string) func(time.Duration, int) []windowInfo { + return func(time.Duration, int) []windowInfo { + return []windowInfo{{PID: 1, HWND: hwnd, Title: title, ClassName: "#32770", ExePath: "setup.exe"}} + } +} + +// windowFrom returns a window script that reports one dialog from the given elapsed time on. +func windowFrom(from time.Duration) func(time.Duration, int) []windowInfo { + return func(elapsed time.Duration, call int) []windowInfo { + if elapsed < from { + return nil + } + return persistentWindow(42, "Setup")(elapsed, call) + } +} + +// TestSupervise_InactivityTimeout verifies termination exactly at InactivityTimeout, followed by +// a single afterAbort. +func TestSupervise_InactivityTimeout(t *testing.T) { + tree := &fakeTree{} + err := runFake(testOptions(Options{InactivityTimeout: 5 * time.Minute}), tree, 2*time.Second, 10*time.Minute) + if !errors.Is(err, ErrInactivityTimeout) { + t.Fatalf("supervise got %v, want ErrInactivityTimeout", err) + } + if tree.terminatedAt != 5*time.Minute || tree.terminated != 1 || tree.afterAborts != 1 { + t.Errorf("terminated %d time(s) at %v with %d afterAbort call(s), want once at 5m with one afterAbort", tree.terminated, tree.terminatedAt, tree.afterAborts) + } +} + +// TestSupervise_ProgressPreventsInactivity verifies that steady progress is never terminated +// when the hard cap is disabled. +func TestSupervise_ProgressPreventsInactivity(t *testing.T) { + tree := &fakeTree{cpuAt: steadyCPU} + opts := testOptions(Options{InactivityTimeout: time.Minute, HardTimeout: -1}) + if err := runFake(opts, tree, 2*time.Second, 3*time.Hour); err != nil { + t.Fatalf("supervise got %v, want nil for a tree making progress", err) + } + if tree.terminated != 0 { + t.Errorf("Tree terminated %d time(s), want 0", tree.terminated) + } +} + +// TestSupervise_HardTimeout verifies that the hard cap terminates a tree making progress. +func TestSupervise_HardTimeout(t *testing.T) { + tree := &fakeTree{cpuAt: steadyCPU} + err := runFake(testOptions(Options{}), tree, 2*time.Second, 2*time.Hour) + if !errors.Is(err, ErrHardTimeout) { + t.Fatalf("supervise got %v, want ErrHardTimeout", err) + } + if tree.terminatedAt != defaultHardTimeout { + t.Errorf("Terminated at %v, want the default hard timeout %v", tree.terminatedAt, defaultHardTimeout) + } +} + +// TestSupervise_MonitorModeLogsOncePerReason verifies that monitor mode never terminates and +// logs one WOULD_KILL line per reason. +func TestSupervise_MonitorModeLogsOncePerReason(t *testing.T) { + for _, tc := range []struct { + name string + winsAt func(time.Duration, int) []windowInfo + reasons []error + }{ + {"InactivityThenHardTimeout", nil, []error{ErrInactivityTimeout, ErrHardTimeout}}, + {"UIThenHardTimeout", persistentWindow(7, "Prompt"), []error{ErrInteractiveUIDetected, ErrHardTimeout}}, + } { + t.Run(tc.name, func(t *testing.T) { + logs := capturedLogsSince() + tree := &fakeTree{winsAt: tc.winsAt} + opts := testOptions(Options{ + Mode: ModeMonitor, + InactivityTimeout: time.Minute, + HardTimeout: 10 * time.Minute, + Unattended: true, + }) + if err := runFake(opts, tree, 2*time.Second, 15*time.Minute); err != nil { + t.Fatalf("supervise got %v in monitor mode, want nil", err) + } + if tree.terminated != 0 { + t.Errorf("Tree terminated %d time(s) in monitor mode, want 0", tree.terminated) + } + out := logs() + for _, r := range tc.reasons { + if n := strings.Count(out, "WOULD_KILL: "+r.Error()); n != 1 { + t.Errorf("Got %d WOULD_KILL lines for %q, want 1. Log:\n%s", n, r, out) + } + } + if n := strings.Count(out, "WOULD_KILL: "); n != len(tc.reasons) { + t.Errorf("Got %d WOULD_KILL lines, want %d. Log:\n%s", n, len(tc.reasons), out) + } + }) + } +} + +// TestSupervise_OffMode verifies that off mode never samples the tree. +func TestSupervise_OffMode(t *testing.T) { + tree := &fakeTree{} + opts := testOptions(Options{Mode: ModeOff, InactivityTimeout: time.Nanosecond, HardTimeout: time.Nanosecond}) + if err := runFake(opts, tree, 0, 0); err != nil { + t.Fatalf("supervise got %v in off mode, want nil", err) + } + if tree.samples != 0 || tree.terminated != 0 { + t.Errorf("Off mode took %d sample(s) and terminated %d time(s), want 0 and 0", tree.samples, tree.terminated) + } +} + +// TestSupervise_UIAbortRequiresNoProgress verifies that a persistent dialog does not abort while +// the tree makes progress, and aborts one grace period after progress stops. +func TestSupervise_UIAbortRequiresNoProgress(t *testing.T) { + tree := &fakeTree{cpuAt: cpuUntil(time.Minute), winsAt: persistentWindow(1, "Setup")} + opts := testOptions(Options{InactivityTimeout: -1, Unattended: true}) + err := runFake(opts, tree, 2*time.Second, 5*time.Minute) + if !errors.Is(err, ErrInteractiveUIDetected) { + t.Fatalf("supervise got %v, want ErrInteractiveUIDetected", err) + } + if want := time.Minute + defaultUIGracePeriod; tree.terminatedAt != want { + t.Errorf("Terminated at %v, want %v: the grace period after the last progress", tree.terminatedAt, want) + } +} + +// TestSupervise_UILatencyAfterStartupBurst verifies that a dialog appearing right after heavy +// startup work aborts one grace period after it appeared, not after the progress window plus +// the grace period. +func TestSupervise_UILatencyAfterStartupBurst(t *testing.T) { + tree := &fakeTree{cpuAt: cpuUntil(10 * time.Second), winsAt: windowFrom(12 * time.Second)} + opts := testOptions(Options{InactivityTimeout: -1, Unattended: true}) + err := runFake(opts, tree, 2*time.Second, 5*time.Minute) + if !errors.Is(err, ErrInteractiveUIDetected) { + t.Fatalf("supervise got %v, want ErrInteractiveUIDetected", err) + } + if want := 12*time.Second + defaultUIGracePeriod; tree.terminatedAt != want { + t.Errorf("Terminated at %v, want %v: one grace period after the dialog appeared", tree.terminatedAt, want) + } +} + +// TestSupervise_UIDetectionInactive verifies that windows are never enumerated or acted on when +// UI detection is inactive. +func TestSupervise_UIDetectionInactive(t *testing.T) { + for _, tc := range []struct { + name string + opts Options + }{ + {"Attended", Options{Unattended: false}}, + {"Disabled", Options{Unattended: true, DisableUIDetection: true}}, + {"NegativeGrace", Options{Unattended: true, UIGracePeriod: -1}}, + } { + t.Run(tc.name, func(t *testing.T) { + tc.opts.InactivityTimeout = -1 + tree := &fakeTree{winsAt: persistentWindow(9, "Prompt")} + if err := runFake(testOptions(tc.opts), tree, 2*time.Second, 10*time.Minute); err != nil { + t.Fatalf("supervise got %v, want nil", err) + } + if tree.windowCalls != 0 { + t.Errorf("windows() called %d time(s), want 0", tree.windowCalls) + } + }) + } +} + +// TestSupervise_RootExitedSuppressesDecisions verifies that no abort decision is made after the +// root process exited while Wait is still pending on held output pipes. +func TestSupervise_RootExitedSuppressesDecisions(t *testing.T) { + tree := &fakeTree{exitAfterSamples: 3} + opts := testOptions(Options{InactivityTimeout: 10 * time.Second, HardTimeout: time.Minute}) + if err := runFake(opts, tree, 2*time.Second, 5*time.Minute); err != nil { + t.Fatalf("supervise got %v after the root exited, want nil", err) + } + if tree.terminated != 0 { + t.Errorf("Tree terminated %d time(s) after the root exited, want 0", tree.terminated) + } + if tree.samples != 3 { + t.Errorf("Took %d samples, want 3: none after the root exited", tree.samples) + } +} diff --git a/supervisor/supervisor.go b/supervisor/supervisor.go new file mode 100644 index 0000000..701e9a8 --- /dev/null +++ b/supervisor/supervisor.go @@ -0,0 +1,344 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +// Package supervisor provides process tree execution, forward progress monitoring, +// and interactive UI sniffing for GooGet package installers. +package supervisor + +import ( + "errors" + "fmt" + "io" + "os" + "os/exec" + "path/filepath" + "runtime" + "strings" + "sync" + "time" +) + +// ErrTerminated is wrapped by every error returned when the supervisor terminates an installer. +var ErrTerminated = errors.New("installer terminated") + +var ( + // ErrInactivityTimeout is returned when an installer process tree exhibits no forward progress for InactivityTimeout. + ErrInactivityTimeout = fmt.Errorf("%w: no forward progress detected for inactivity timeout", ErrTerminated) + + // ErrInteractiveUIDetected is returned when a modal dialog persists without forward progress in unattended mode. + ErrInteractiveUIDetected = fmt.Errorf("%w: interactive UI dialog detected in unattended mode", ErrTerminated) + + // ErrHardTimeout is returned when an installer exceeds the absolute HardTimeout regardless of progress. + ErrHardTimeout = fmt.Errorf("%w: exceeded hard install timeout", ErrTerminated) +) + +// ParseTimeout parses a configured watchdog duration. An empty string returns 0, meaning "use the +// default"; "0" returns -1, meaning "disabled"; negative durations are rejected. +func ParseTimeout(s string) (time.Duration, error) { + s = strings.TrimSpace(s) + switch s { + case "": + return 0, nil + case "0": + return -1, nil + } + d, err := time.ParseDuration(s) + if err != nil { + return 0, fmt.Errorf("invalid timeout %q: %w", s, err) + } + if d < 0 { + return 0, fmt.Errorf("invalid timeout %q: must not be negative", s) + } + if d == 0 { + return -1, nil + } + return d, nil +} + +// Mode selects whether the supervisor enforces its abort decisions. +type Mode int + +const ( + // ModeUnset means the process-wide default configured via Configure is used. + ModeUnset Mode = iota + // ModeEnforce terminates the process tree when an abort condition is met. + ModeEnforce + // ModeMonitor logs WOULD_KILL diagnostics when an abort condition is met but never terminates. + ModeMonitor + // ModeOff disables all watchdogs; the process tree is still contained and stdin is still disconnected. + ModeOff +) + +// String returns the configuration name of the mode. +func (m Mode) String() string { + switch m { + case ModeEnforce: + return "enforce" + case ModeMonitor: + return "monitor" + case ModeOff: + return "off" + default: + return "unset" + } +} + +// ParseMode parses a configuration value ("enforce", "monitor" or "off") into a Mode. +func ParseMode(s string) (Mode, error) { + switch strings.ToLower(strings.TrimSpace(s)) { + case "enforce": + return ModeEnforce, nil + case "monitor": + return ModeMonitor, nil + case "off": + return ModeOff, nil + default: + return ModeUnset, fmt.Errorf("invalid supervisor mode %q: want enforce, monitor or off", s) + } +} + +// windowInfo describes an active top-level desktop window. +type windowInfo struct { + PID uint32 + HWND uintptr + Title string + ClassName string + ExePath string +} + +// windowDetectorFunc inspects desktop windows belonging to the given process IDs and returns +// only windows that are candidate interactive prompts. +type windowDetectorFunc func(pids []uint32) ([]windowInfo, error) + +// Options specifies watchdog behavior for supervised execution. +// +// For every duration field, zero means "use the process-wide default set by Configure" and a +// negative value disables that watchdog. +// +// The bool fields Unattended and DisableUIDetection are combined with the process-wide defaults +// by a logical OR: a per-call Options value can enable either behavior but cannot clear a value +// enabled through Configure. Unattended is additionally forced on when the current process runs +// in Windows Session 0, where no user can answer a dialog. +type Options struct { + // Mode selects enforce, monitor or off behavior. + Mode Mode + + // InactivityTimeout is the maximum duration without meaningful forward progress before aborting. + InactivityTimeout time.Duration + + // HardTimeout is the absolute maximum runtime regardless of forward progress. + HardTimeout time.Duration + + // UIGracePeriod is how long a candidate modal dialog may persist, with no forward progress + // since it appeared, before aborting. + UIGracePeriod time.Duration + + // DisableUIDetection turns off interactive dialog detection. It is ORed with the process-wide + // default. + DisableUIDetection bool + + // LogFiles is a list of file paths whose size growth indicates forward progress. + LogFiles []string + + // Unattended indicates the installation is non-interactive, for example -noconfirm. It is + // ORed with the process-wide default and with Session 0 detection. + Unattended bool + + // pollInterval is the duration between watchdog samples. + pollInterval time.Duration + + // progressWindow is the rolling window over which progress deltas are compared against thresholds. + progressWindow time.Duration + + // minCPUDelta is the minimum aggregate CPU time consumed within progressWindow to count as progress. + minCPUDelta time.Duration + + // minIODelta is the minimum aggregate I/O bytes transferred within progressWindow to count as progress. + minIODelta uint64 + + // msiMutexWait bounds how long to wait for the Windows Installer _MSIExecute mutex to be + // released after a job that ran msiexec is terminated. + msiMutexWait time.Duration + + // windowDetector overrides Win32 window enumeration in tests. + windowDetector windowDetectorFunc +} + +// Built-in defaults used when neither the caller nor Configure provides a value. +const ( + defaultMode = ModeEnforce + defaultInactivityTimeout = 5 * time.Minute + defaultHardTimeout = 60 * time.Minute + defaultUIGracePeriod = 30 * time.Second + defaultPollInterval = 2 * time.Second + defaultProgressWindow = 30 * time.Second + defaultMinCPUDelta = 250 * time.Millisecond + defaultMinIODelta = 64 * 1024 + defaultMSIMutexWait = 10 * time.Minute +) + +var ( + defaultsMu sync.RWMutex + defaults = builtinDefaults() + // hardTimeoutConfigured records whether the last Configure call set HardTimeout, including + // to the built-in default's value or to a negative value that disables it. + hardTimeoutConfigured bool +) + +// builtinDefaults returns the built-in process-wide defaults. +func builtinDefaults() Options { + return Options{ + Mode: defaultMode, + InactivityTimeout: defaultInactivityTimeout, + HardTimeout: defaultHardTimeout, + UIGracePeriod: defaultUIGracePeriod, + pollInterval: defaultPollInterval, + progressWindow: defaultProgressWindow, + minCPUDelta: defaultMinCPUDelta, + minIODelta: defaultMinIODelta, + msiMutexWait: defaultMSIMutexWait, + } +} + +// Configure sets process-wide defaults. Zero-valued fields in d keep the built-in defaults. +// Unattended and DisableUIDetection are copied as-is, so Configure(Options{}) restores every +// built-in default. +func Configure(d Options) { + merged := mergeOptions(d, builtinDefaults()) + merged.Unattended = d.Unattended + merged.DisableUIDetection = d.DisableUIDetection + defaultsMu.Lock() + defaults = merged + hardTimeoutConfigured = d.HardTimeout != 0 + defaultsMu.Unlock() +} + +// CurrentDefaults returns a copy of the process-wide defaults. +func CurrentDefaults() Options { + defaultsMu.RLock() + defer defaultsMu.RUnlock() + return defaults +} + +// adminHardTimeoutConfigured reports whether the process-wide hard timeout was set explicitly +// through Configure rather than left at the built-in default. +func adminHardTimeoutConfigured() bool { + defaultsMu.RLock() + defer defaultsMu.RUnlock() + return hardTimeoutConfigured +} + +// servicingOptions adjusts per-call options for a wusa or DISM servicing operation, whose work +// runs in TrustedInstaller and TiWorker outside the installer's process tree and routinely +// takes longer than the built-in hard timeout. +// +// The built-in hard timeout is lifted unless the caller or the administrator set one; an +// explicitly configured value, even one equal to the built-in default, is respected. Growth of +// the CBS servicing log under windir, or C:\Windows if windir is empty, counts as forward +// progress so that the inactivity watchdog stays enabled. +func servicingOptions(opts Options, adminHardConfigured bool, windir string) Options { + if opts.HardTimeout == 0 && !adminHardConfigured { + opts.HardTimeout = -1 + } + if windir == "" { + windir = `C:\Windows` + } + logs := make([]string, 0, len(opts.LogFiles)+1) + logs = append(logs, opts.LogFiles...) + opts.LogFiles = append(logs, filepath.Join(windir, "Logs", "CBS", "CBS.log")) + return opts +} + +// fallbackOptions restricts resolved options to what can be enforced when no Job Object could +// be set up. Without a job, descendants can be neither measured nor terminated, so the +// inactivity and UI watchdogs would judge and kill only the root process while its children +// keep working. Only the hard timeout on the root process stays enforced. +func fallbackOptions(opts Options) Options { + if opts.Mode == ModeOff { + return opts + } + opts.InactivityTimeout = -1 + opts.DisableUIDetection = true + return opts +} + +// mergeOptions fills zero-valued fields of o from d. +func mergeOptions(o, d Options) Options { + if o.Mode == ModeUnset { + o.Mode = d.Mode + } + if o.InactivityTimeout == 0 { + o.InactivityTimeout = d.InactivityTimeout + } + if o.HardTimeout == 0 { + o.HardTimeout = d.HardTimeout + } + if o.UIGracePeriod == 0 { + o.UIGracePeriod = d.UIGracePeriod + } + if o.pollInterval == 0 { + o.pollInterval = d.pollInterval + } + if o.progressWindow == 0 { + o.progressWindow = d.progressWindow + } + if o.minCPUDelta == 0 { + o.minCPUDelta = d.minCPUDelta + } + if o.minIODelta == 0 { + o.minIODelta = d.minIODelta + } + if o.msiMutexWait == 0 { + o.msiMutexWait = d.msiMutexWait + } + if o.windowDetector == nil { + o.windowDetector = d.windowDetector + } + return o +} + +// resolve merges caller options with process-wide defaults and derives dependent values. It is +// the only place where Unattended is derived. +func resolve(o Options) Options { + d := CurrentDefaults() + r := mergeOptions(o, d) + r.Unattended = o.Unattended || d.Unattended || isSession0() + r.DisableUIDetection = o.DisableUIDetection || d.DisableUIDetection + if r.pollInterval < 0 { + r.pollInterval = defaultPollInterval + } + // Never poll more coarsely than a quarter of the shortest enabled watchdog. + for _, t := range []time.Duration{r.InactivityTimeout, r.UIGracePeriod} { + if t > 0 && r.pollInterval > t/4 { + r.pollInterval = t / 4 + } + } + if r.pollInterval < 10*time.Millisecond { + r.pollInterval = 10 * time.Millisecond + } + if r.progressWindow <= 0 || (r.InactivityTimeout > 0 && r.progressWindow > r.InactivityTimeout) { + r.progressWindow = r.InactivityTimeout + } + return r +} + +// Run executes the command under process tree supervision, monitoring forward progress and UI state. +// Stdout and stderr are tee'd to os.Stdout/os.Stderr and out (if non-nil). Stdin is always disconnected. +// +// On Windows, wusa and DISM commands get the servicing policy described at servicingOptions. +func Run(c *exec.Cmd, opts Options, out io.Writer) error { + if runtime.GOOS == "windows" && isServicingCommand(c.Path) { + opts = servicingOptions(opts, adminHardTimeoutConfigured(), os.Getenv("WINDIR")) + } + return runSupervised(c, resolve(opts), out) +} diff --git a/supervisor/supervisor_test.go b/supervisor/supervisor_test.go new file mode 100644 index 0000000..b269076 --- /dev/null +++ b/supervisor/supervisor_test.go @@ -0,0 +1,943 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "bytes" + "errors" + "fmt" + "io" + "os" + "os/exec" + "path/filepath" + "runtime" + "slices" + "strconv" + "strings" + "sync" + "testing" + "time" + + "github.com/google/logger" +) + +// platformHelpers holds helper subcommands registered by platform-specific test files. +var platformHelpers = map[string]func(args []string){} + +// helperCommand returns an *exec.Cmd configured to invoke TestHelperProcess. +func helperCommand(t *testing.T, subcmd string, extraArgs ...string) *exec.Cmd { + t.Helper() + cmd := exec.Command(os.Args[0], helperArgs(subcmd, extraArgs...)...) + cmd.Env = helperEnv() + return cmd +} + +// helperArgs returns the arguments that make the test binary run a helper subcommand. +func helperArgs(subcmd string, extraArgs ...string) []string { + args := []string{"-test.run=^TestHelperProcess$", "--", subcmd} + return append(args, extraArgs...) +} + +// helperEnv returns the environment for helper processes. +// +// GORACE=atexit_sleep_ms=0 stops race-enabled helpers from idling for one second at exit, which +// would otherwise look like inactivity to the supervisor. +func helperEnv() []string { + return append(os.Environ(), "GO_WANT_HELPER_PROCESS=1", "GORACE=atexit_sleep_ms=0") +} + +// spawnHelper starts a helper subcommand from inside a helper process. +func spawnHelper(subcmd string, extraArgs ...string) *exec.Cmd { + cmd := exec.Command(os.Args[0], helperArgs(subcmd, extraArgs...)...) + cmd.Env = helperEnv() + return cmd +} + +// burnCPU busy-loops for d of wall time. +func burnCPU(d time.Duration) { + start := time.Now() + for time.Since(start) < d { + _ = strconv.Itoa(int(time.Now().UnixNano())) + } +} + +// appendLines appends n lines to path, sleeping interval before each. +func appendLines(path string, n int, interval time.Duration) { + for i := 0; i < n; i++ { + time.Sleep(interval) + f, err := os.OpenFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) + if err != nil { + fmt.Fprintf(os.Stderr, "Failed to open log file: %v\n", err) + os.Exit(1) + } + fmt.Fprintf(f, "progress line %d\n", i) + f.Close() + } +} + +// atoiOr parses s as an integer or returns def. +func atoiOr(s string, def int) int { + if n, err := strconv.Atoi(s); err == nil { + return n + } + return def +} + +// TestHelperProcess acts as a mock subprocess for supervisor tests. +func TestHelperProcess(t *testing.T) { + if os.Getenv("GO_WANT_HELPER_PROCESS") != "1" { + return + } + defer os.Exit(0) + + args := os.Args + for len(args) > 0 { + if args[0] == "--" { + args = args[1:] + break + } + args = args[1:] + } + if len(args) == 0 { + fmt.Fprintf(os.Stderr, "No command provided to helper process.\n") + os.Exit(2) + } + + cmd, cmdArgs := args[0], args[1:] + arg := func(i int) string { + if i < len(cmdArgs) { + return cmdArgs[i] + } + return "" + } + switch cmd { + case "stall": + // Sleeps without performing any I/O or consuming CPU. + time.Sleep(60 * time.Second) + case "sleep_ms": + // Sleeps for the given number of milliseconds and exits successfully. + time.Sleep(time.Duration(atoiOr(arg(0), 200)) * time.Millisecond) + case "active_logging": + // Appends arg(1) lines to log file arg(0), one every arg(2) milliseconds. + appendLines(arg(0), atoiOr(arg(1), 15), time.Duration(atoiOr(arg(2), 100))*time.Millisecond) + case "log_then_stall": + // Appends arg(1) lines every arg(2) milliseconds, then stalls. + appendLines(arg(0), atoiOr(arg(1), 10), time.Duration(atoiOr(arg(2), 50))*time.Millisecond) + time.Sleep(60 * time.Second) + case "burn_cpu": + // Consumes CPU in a busy loop for arg(0) milliseconds. + burnCPU(time.Duration(atoiOr(arg(0), 250)) * time.Millisecond) + case "spawn_orphan_child": + // Spawns a stalled child, records its PID in arg(0) and stalls. + child := spawnHelper("stall") + if err := child.Start(); err != nil { + fmt.Fprintf(os.Stderr, "Failed to start child process: %v\n", err) + os.Exit(1) + } + if err := os.WriteFile(arg(0), []byte(strconv.Itoa(child.Process.Pid)), 0644); err != nil { + fmt.Fprintf(os.Stderr, "Failed to write pid file: %v\n", err) + os.Exit(1) + } + time.Sleep(60 * time.Second) + case "parent_sleep_child_burn": + // Waits without using CPU while a child burns CPU for arg(0) milliseconds. + child := spawnHelper("burn_cpu", arg(0)) + if err := child.Run(); err != nil { + fmt.Fprintf(os.Stderr, "Child failed: %v\n", err) + os.Exit(1) + } + case "exit_code_42": + // Exits with code 42 to verify non-zero exit codes are preserved. + os.Exit(42) + case "read_stdin": + // Reads from standard input. Expects immediate EOF if stdin is disconnected. + buf := make([]byte, 16) + n, err := os.Stdin.Read(buf) + if err == io.EOF || n == 0 { + os.Exit(0) + } + fmt.Fprintf(os.Stderr, "Unexpected data read from stdin: %d bytes, err: %v\n", n, err) + os.Exit(1) + default: + if f, ok := platformHelpers[cmd]; ok { + f(cmdArgs) + return + } + fmt.Fprintf(os.Stderr, "Unknown helper command: %q\n", cmd) + os.Exit(2) + } +} + +// syncBuffer is a goroutine-safe bytes.Buffer. +type syncBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +// Write appends p to the buffer. +func (b *syncBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +// String returns the buffer contents. +func (b *syncBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +var ( + logCaptureOnce sync.Once + logCapture = &syncBuffer{} +) + +// capturedLogsSince returns a function that yields log output written after this call. +func capturedLogsSince() func() string { + logCaptureOnce.Do(func() { logger.Init("supervisor_test", false, false, logCapture) }) + start := len(logCapture.String()) + return func() string { return logCapture.String()[start:] } +} + +// persistentDialog returns a detector that always reports the same dialog. +func persistentDialog(hwnd uintptr) windowDetectorFunc { + return func(pids []uint32) ([]windowInfo, error) { + return []windowInfo{{ + PID: pids[0], + HWND: hwnd, + Title: "License Agreement", + ClassName: "#32770", + ExePath: "installer.exe", + }}, nil + } +} + +// supportsCPUAccounting reports whether the platform measures child CPU time. +func supportsCPUAccounting() bool { + return runtime.GOOS == "linux" || runtime.GOOS == "windows" +} + +// TestRun_InactivityTimeout verifies that a stalled process is aborted after InactivityTimeout. +func TestRun_InactivityTimeout(t *testing.T) { + cmd := helperCommand(t, "stall") + opts := Options{ + InactivityTimeout: 300 * time.Millisecond, + pollInterval: 20 * time.Millisecond, + } + + start := time.Now() + err := Run(cmd, opts, io.Discard) + elapsed := time.Since(start) + + if !errors.Is(err, ErrInactivityTimeout) { + t.Fatalf("Run got error %v, want ErrInactivityTimeout", err) + } + if elapsed < 300*time.Millisecond { + t.Errorf("Run took %v, want at least the 300ms inactivity timeout", elapsed) + } + if elapsed > 15*time.Second { + t.Errorf("Run took %v, want termination shortly after the 300ms timeout", elapsed) + } +} + +// TestRun_StdinDisconnected verifies that reading child stdin returns EOF immediately. +func TestRun_StdinDisconnected(t *testing.T) { + cmd := helperCommand(t, "read_stdin") + opts := Options{InactivityTimeout: -1, pollInterval: 50 * time.Millisecond} + + if err := Run(cmd, opts, io.Discard); err != nil { + t.Fatalf("Run returned error %v, want clean exit 0 on stdin EOF", err) + } +} + +// TestRun_StdinDisconnected_PreexistingStdin verifies that caller-provided stdin is replaced. +func TestRun_StdinDisconnected_PreexistingStdin(t *testing.T) { + cmd := helperCommand(t, "read_stdin") + cmd.Stdin = strings.NewReader("unexpected stdin input\n") + opts := Options{InactivityTimeout: -1, pollInterval: 50 * time.Millisecond} + + if err := Run(cmd, opts, io.Discard); err != nil { + t.Fatalf("Run returned error %v, want clean exit 0 on stdin EOF", err) + } +} + +// TestRun_OutputTee verifies that child output is copied to the out writer. +func TestRun_OutputTee(t *testing.T) { + cmd := helperCommand(t, "unknown_subcommand_for_output") + var out syncBuffer + _ = Run(cmd, Options{InactivityTimeout: -1}, &out) + if !strings.Contains(out.String(), "Unknown helper command") { + t.Errorf("Run output %q does not contain the helper's stderr", out.String()) + } +} + +// TestRun_ForwardProgress_CPUTicks verifies that CPU consumption above minCPUDelta keeps a +// process alive past InactivityTimeout. +func TestRun_ForwardProgress_CPUTicks(t *testing.T) { + if !supportsCPUAccounting() { + t.Skip("CPU accounting is only supported on Linux and Windows") + } + + // The helper burns CPU for 2s while InactivityTimeout is 800ms. + cmd := helperCommand(t, "burn_cpu", "2000") + opts := Options{ + InactivityTimeout: 800 * time.Millisecond, + pollInterval: 20 * time.Millisecond, + minCPUDelta: 50 * time.Millisecond, + } + + start := time.Now() + err := Run(cmd, opts, io.Discard) + elapsed := time.Since(start) + + if err != nil { + t.Fatalf("Run returned error %v for CPU active process, want nil", err) + } + if elapsed < 1800*time.Millisecond { + t.Errorf("Run completed in %v, expected at least 1.8s of CPU execution", elapsed) + } +} + +// TestRun_HardTimeout verifies that HardTimeout aborts even a process making progress. +func TestRun_HardTimeout(t *testing.T) { + logFile := filepath.Join(t.TempDir(), "install.log") + cmd := helperCommand(t, "active_logging", logFile, "1200", "50") + opts := Options{ + InactivityTimeout: 5 * time.Second, + HardTimeout: 700 * time.Millisecond, + pollInterval: 20 * time.Millisecond, + LogFiles: []string{logFile}, + } + start := time.Now() + err := Run(cmd, opts, io.Discard) + elapsed := time.Since(start) + if !errors.Is(err, ErrHardTimeout) { + t.Fatalf("Run got error %v, want ErrHardTimeout", err) + } + if elapsed < 700*time.Millisecond || elapsed > 15*time.Second { + t.Errorf("Run took %v, want between the 700ms hard timeout and 15s", elapsed) + } +} + +// TestRun_OffModeNoWatchdogs verifies that off mode disables every watchdog. +func TestRun_OffModeNoWatchdogs(t *testing.T) { + cmd := helperCommand(t, "sleep_ms", "300") + opts := Options{ + Mode: ModeOff, + InactivityTimeout: 50 * time.Millisecond, + HardTimeout: 50 * time.Millisecond, + UIGracePeriod: 20 * time.Millisecond, + pollInterval: 10 * time.Millisecond, + Unattended: true, + windowDetector: persistentDialog(88), + } + if err := Run(cmd, opts, io.Discard); err != nil { + t.Fatalf("Run returned %v in off mode, want nil", err) + } + if cmd.WaitDelay != defaultWaitDelay { + t.Errorf("WaitDelay got %v, want %v in off mode", cmd.WaitDelay, defaultWaitDelay) + } +} + +// TestRun_WaitDelayPreserved verifies that a caller-provided WaitDelay is not overridden. +func TestRun_WaitDelayPreserved(t *testing.T) { + cmd := helperCommand(t, "sleep_ms", "10") + cmd.WaitDelay = 7 * time.Second + if err := Run(cmd, Options{InactivityTimeout: -1}, io.Discard); err != nil { + t.Fatalf("Run returned %v, want nil", err) + } + if cmd.WaitDelay != 7*time.Second { + t.Errorf("WaitDelay got %v, want the caller's 7s", cmd.WaitDelay) + } +} + +// TestRun_OrphanChildCleanup verifies that terminating an inactive parent also kills its +// descendants. +func TestRun_OrphanChildCleanup(t *testing.T) { + if runtime.GOOS != "linux" { + t.Skip("Process verification via /proc is only supported on Linux") + } + + pidFile := filepath.Join(t.TempDir(), "child.pid") + cmd := helperCommand(t, "spawn_orphan_child", pidFile) + opts := Options{ + InactivityTimeout: 500 * time.Millisecond, + pollInterval: 20 * time.Millisecond, + } + + if err := Run(cmd, opts, io.Discard); !errors.Is(err, ErrInactivityTimeout) { + t.Fatalf("Run got error %v, want ErrInactivityTimeout", err) + } + + data, err := os.ReadFile(pidFile) + if err != nil { + t.Fatalf("Failed to read child PID file: %v", err) + } + childPid, err := strconv.Atoi(strings.TrimSpace(string(data))) + if err != nil { + t.Fatalf("Failed to parse child PID: %v", err) + } + + if !waitForExit(childPid, 5*time.Second) { + t.Errorf("Child process %d is still alive; expected process group termination", childPid) + } +} + +// TestRun_MultiLogFiles_PartialGrowth verifies that growth of any monitored log is progress. +func TestRun_MultiLogFiles_PartialGrowth(t *testing.T) { + tmpDir := t.TempDir() + staticLog := filepath.Join(tmpDir, "static.log") + growingLog := filepath.Join(tmpDir, "growing.log") + if err := os.WriteFile(staticLog, []byte("static initial content\n"), 0644); err != nil { + t.Fatalf("Failed to write static log: %v", err) + } + if err := os.WriteFile(growingLog, []byte("growing initial content\n"), 0644); err != nil { + t.Fatalf("Failed to write growing log: %v", err) + } + + // The helper logs every 50ms for about 1.5s while InactivityTimeout is 1s. + cmd := helperCommand(t, "active_logging", growingLog, "30", "50") + opts := Options{ + InactivityTimeout: 1 * time.Second, + pollInterval: 20 * time.Millisecond, + minCPUDelta: time.Hour, + LogFiles: []string{staticLog, growingLog}, + } + + start := time.Now() + err := Run(cmd, opts, io.Discard) + elapsed := time.Since(start) + if err != nil { + t.Fatalf("Run returned unexpected error %v, want nil", err) + } + if elapsed < 1200*time.Millisecond { + t.Errorf("Run completed in %v, expected at least 1.2s of work", elapsed) + } +} + +// TestRun_ExitCodePreserved verifies that a non-zero exit code is preserved as an ExitError. +func TestRun_ExitCodePreserved(t *testing.T) { + cmd := helperCommand(t, "exit_code_42") + err := Run(cmd, Options{InactivityTimeout: -1, pollInterval: 50 * time.Millisecond}, io.Discard) + var exitErr *exec.ExitError + if !errors.As(err, &exitErr) { + t.Fatalf("Run returned error %v of type %T, want *exec.ExitError", err, err) + } + if exitErr.ExitCode() != 42 { + t.Errorf("ExitCode got %d, want 42", exitErr.ExitCode()) + } +} + +// TestRun_NonExistentBinary verifies that a missing binary fails fast during start. +func TestRun_NonExistentBinary(t *testing.T) { + cmd := exec.Command("nonexistent_binary_xyz_12345") + start := time.Now() + err := Run(cmd, Options{InactivityTimeout: 5 * time.Second}, io.Discard) + elapsed := time.Since(start) + + if err == nil { + t.Fatal("Run returned nil for non-existent binary, want error") + } + if errors.Is(err, ErrInactivityTimeout) { + t.Fatal("Run returned ErrInactivityTimeout, want executable not found error") + } + if elapsed > 2*time.Second { + t.Errorf("Run took %v, expected near-instant failure for non-existent binary", elapsed) + } +} + +// TestResolve verifies that zero-value Options are populated with defaults and clamped. +func TestResolve(t *testing.T) { + opts := resolve(Options{}) + if opts.InactivityTimeout != defaultInactivityTimeout { + t.Errorf("InactivityTimeout got %v, want %v", opts.InactivityTimeout, defaultInactivityTimeout) + } + if opts.UIGracePeriod != defaultUIGracePeriod { + t.Errorf("UIGracePeriod got %v, want %v", opts.UIGracePeriod, defaultUIGracePeriod) + } + if opts.pollInterval != defaultPollInterval { + t.Errorf("pollInterval got %v, want %v", opts.pollInterval, defaultPollInterval) + } + + small := resolve(Options{InactivityTimeout: 20 * time.Millisecond}) + if small.pollInterval != 10*time.Millisecond { + t.Errorf("Clamped pollInterval got %v, want 10ms minimum", small.pollInterval) + } + if small.progressWindow != 20*time.Millisecond { + t.Errorf("progressWindow got %v, want it clamped to the 20ms inactivity timeout", small.progressWindow) + } +} + +// TestProgressTracker verifies rolling-window thresholds on synthetic samples. +func TestProgressTracker(t *testing.T) { + base := time.Date(2026, 9, 28, 12, 0, 0, 0, time.UTC) + at := func(s float64) time.Time { return base.Add(time.Duration(s * float64(time.Second))) } + const window = 30 * time.Second + const minCPU = 250 * time.Millisecond + const minIO = 64 * 1024 + + t.Run("NoiseBelowThresholdIsNotProgress", func(t *testing.T) { + tr := newProgressTracker(window, minCPU, minIO, base) + var cpu time.Duration + // 10ms of CPU and 1KiB of I/O every 2s is 150ms and 15KiB per 30s, below the thresholds. + for i := 1; i <= 60; i++ { + cpu += 10 * time.Millisecond + if tr.observe(progressSample{at: at(float64(2 * i)), cpu: cpu, io: uint64(i) * 1024}, nil) { + t.Fatalf("Tick %d counted noise as progress", i) + } + } + if !tr.lastProgress.Equal(base) { + t.Errorf("lastProgress got %v, want unchanged %v", tr.lastProgress, base) + } + }) + + t.Run("CPUAboveThresholdWithinWindow", func(t *testing.T) { + tr := newProgressTracker(window, minCPU, minIO, base) + tr.observe(progressSample{at: at(0)}, nil) + if tr.observe(progressSample{at: at(2), cpu: 100 * time.Millisecond}, nil) { + t.Fatal("100ms of CPU counted as progress") + } + if !tr.observe(progressSample{at: at(4), cpu: 260 * time.Millisecond}, nil) { + t.Fatal("260ms of CPU within the window did not count as progress") + } + }) + + t.Run("OldBurstLeavesWindow", func(t *testing.T) { + tr := newProgressTracker(window, minCPU, minIO, base) + tr.observe(progressSample{at: at(0)}, nil) + tr.observe(progressSample{at: at(1), cpu: time.Second}, nil) + if !tr.observe(progressSample{at: at(20), cpu: time.Second}, nil) { + t.Fatal("A burst 19s ago within the 30s window was not progress") + } + if tr.observe(progressSample{at: at(40), cpu: time.Second}, nil) { + t.Fatal("A burst 39s ago outside the 30s window was progress") + } + if !tr.lastProgress.Equal(at(20)) { + t.Errorf("lastProgress got %v, want %v", tr.lastProgress, at(20)) + } + }) + + t.Run("IOAboveThreshold", func(t *testing.T) { + tr := newProgressTracker(window, minCPU, minIO, base) + tr.observe(progressSample{at: at(0)}, nil) + if !tr.observe(progressSample{at: at(2), io: minIO}, nil) { + t.Fatal("I/O equal to the threshold did not count as progress") + } + }) + + t.Run("NegativeMinCPUDisablesCPUSignal", func(t *testing.T) { + tr := newProgressTracker(window, -1, minIO, base) + tr.observe(progressSample{at: at(0)}, nil) + if tr.observe(progressSample{at: at(2), cpu: time.Hour}, nil) { + t.Fatal("CPU counted as progress with a negative minCPUDelta") + } + }) + + t.Run("LogGrowthZeroThreshold", func(t *testing.T) { + tr := newProgressTracker(window, minCPU, minIO, base) + tr.observe(progressSample{at: at(0)}, map[string]int64{"a": 10}) + if tr.observe(progressSample{at: at(2)}, map[string]int64{"a": 10}) { + t.Fatal("Unchanged log counted as progress") + } + if !tr.observe(progressSample{at: at(4)}, map[string]int64{"a": 11}) { + t.Fatal("One byte of log growth did not count as progress") + } + if !tr.observe(progressSample{at: at(6)}, map[string]int64{"a": 11, "b": 0}) { + t.Fatal("A newly appearing log file did not count as progress") + } + if tr.observe(progressSample{at: at(8)}, map[string]int64{"a": 5, "b": 0}) { + t.Fatal("A truncated log counted as progress") + } + }) +} + +// TestPIDCounters verifies monotonic aggregation across appearing and exiting processes. +func TestPIDCounters(t *testing.T) { + p1, p2, p3 := procKey{pid: 1, created: 10}, procKey{pid: 2, created: 20}, procKey{pid: 3, created: 30} + var p pidCounters + if cpu, io := p.update(map[procKey]counters{p1: {cpu: time.Second, io: 100}}); cpu != 0 || io != 0 { + t.Fatalf("Priming update got (%v, %d), want zero", cpu, io) + } + cpu, io := p.update(map[procKey]counters{p1: {cpu: 2 * time.Second, io: 150}, p2: {cpu: 300 * time.Millisecond, io: 10}}) + if cpu != 1300*time.Millisecond || io != 60 { + t.Errorf("Update got (%v, %d), want (1.3s, 60)", cpu, io) + } + // Process 1 exits; process 2 advances. + cpu, io = p.update(map[procKey]counters{p2: {cpu: 400 * time.Millisecond, io: 10}}) + if cpu != 100*time.Millisecond || io != 0 { + t.Errorf("Update after exit got (%v, %d), want (100ms, 0)", cpu, io) + } + p.reset() + if cpu, _ := p.update(map[procKey]counters{p3: {cpu: time.Hour}}); cpu != 0 { + t.Errorf("Update after reset got %v, want zero", cpu) + } +} + +// TestPIDCounters_TransientDropout verifies that a process missing from one sample is credited +// only with its increase when it reappears, while a reused PID with a different creation time +// counts as a new process. +func TestPIDCounters_TransientDropout(t *testing.T) { + svc := procKey{pid: 100, created: 1} + var p pidCounters + p.update(map[procKey]counters{svc: {cpu: time.Hour, io: 1 << 30}}) + if cpu, io := p.update(map[procKey]counters{}); cpu != 0 || io != 0 { + t.Fatalf("Update with the process missing got (%v, %d), want zero", cpu, io) + } + cpu, io := p.update(map[procKey]counters{svc: {cpu: time.Hour + 10*time.Millisecond, io: 1<<30 + 5}}) + if cpu != 10*time.Millisecond || io != 5 { + t.Errorf("Update after the process reappeared got (%v, %d), want (10ms, 5), not its lifetime counters", cpu, io) + } + reused := procKey{pid: 100, created: 2} + if cpu, _ := p.update(map[procKey]counters{reused: {cpu: 50 * time.Millisecond}}); cpu != 50*time.Millisecond { + t.Errorf("Update for a reused PID got %v, want its full 50ms", cpu) + } +} + +// noProgress is a progressedSince function that never reports progress. +func noProgress(time.Time) bool { return false } + +// alwaysProgress is a progressedSince function that always reports progress. +func alwaysProgress(time.Time) bool { return true } + +// TestUISnifferState_ProgressResets verifies that progress after a window appeared resets its +// timer. +func TestUISnifferState_ProgressResets(t *testing.T) { + base := time.Date(2026, 9, 28, 12, 0, 0, 0, time.UTC) + grace := 30 * time.Second + w := []windowInfo{{HWND: 1, Title: "A", ClassName: "#32770"}} + var s uiSnifferState + s.check(w, grace, base, noProgress) + if abort, _ := s.check(w, grace, base.Add(20*time.Second), alwaysProgress); abort { + t.Fatal("Unexpected abort on a tick with progress") + } + if abort, _ := s.check(w, grace, base.Add(45*time.Second), noProgress); abort { + t.Fatal("Unexpected abort 25s after progress reset the timer") + } + if abort, _ := s.check(w, grace, base.Add(50*time.Second), noProgress); !abort { + t.Fatal("Expected abort 30s after the last progress") + } +} + +// TestProgressedSince verifies that only progress strictly after the given time counts. +func TestProgressedSince(t *testing.T) { + base := time.Date(2026, 9, 28, 12, 0, 0, 0, time.UTC) + at := func(s int) time.Time { return base.Add(time.Duration(s) * time.Second) } + tr := newProgressTracker(30*time.Second, 250*time.Millisecond, 64*1024, base) + tr.observe(progressSample{at: at(0)}, nil) + tr.observe(progressSample{at: at(2), cpu: time.Second}, nil) + tr.observe(progressSample{at: at(4), cpu: time.Second}, nil) + if !tr.progressedSince(at(0)) { + t.Error("progressedSince(0s) got false, want true for a 1s burst at 2s") + } + if tr.progressedSince(at(2)) { + t.Error("progressedSince(2s) got true, want false because the burst ended at 2s") + } + tr.observe(progressSample{at: at(6), cpu: 1100 * time.Millisecond}, nil) + if tr.progressedSince(at(2)) { + t.Error("progressedSince(2s) got true for 100ms of CPU, want false") + } + tr.observe(progressSample{at: at(8), cpu: 1300 * time.Millisecond}, nil) + if !tr.progressedSince(at(2)) { + t.Error("progressedSince(2s) got false for 300ms of CPU, want true") + } + if tr.progressedSince(at(6)) { + t.Error("progressedSince(6s) got true for 200ms of CPU, want false") + } + if tr.progressedSince(at(8)) { + t.Error("progressedSince(8s) got true at the latest sample, want false") + } + + logs := newProgressTracker(30*time.Second, 250*time.Millisecond, 64*1024, base) + logs.observe(progressSample{at: at(0)}, map[string]int64{"a": 1}) + logs.observe(progressSample{at: at(2)}, map[string]int64{"a": 2}) + logs.observe(progressSample{at: at(4)}, map[string]int64{"a": 2}) + if !logs.progressedSince(at(0)) { + t.Error("progressedSince(0s) got false, want true for log growth at 2s") + } + if logs.progressedSince(at(2)) { + t.Error("progressedSince(2s) got true, want false because the log last grew at 2s") + } +} + +// TestConfigure verifies process-wide defaults and how per-call options combine with them. +func TestConfigure(t *testing.T) { + t.Cleanup(func() { Configure(Options{}) }) + + Configure(Options{Mode: ModeMonitor, InactivityTimeout: -1, Unattended: true}) + d := CurrentDefaults() + if d.Mode != ModeMonitor || d.InactivityTimeout != -1 || !d.Unattended || d.HardTimeout != defaultHardTimeout || d.pollInterval != defaultPollInterval { + t.Errorf("CurrentDefaults() after Configure = %+v, want monitor, inactivity disabled, unattended and built-in values elsewhere", d) + } + got := resolve(Options{}) + if got.Mode != ModeMonitor || got.InactivityTimeout != -1 || !got.Unattended || got.HardTimeout != defaultHardTimeout { + t.Errorf("resolve(Options{}) after Configure = %+v, want the configured defaults", got) + } + got = resolve(Options{Mode: ModeEnforce, InactivityTimeout: time.Minute, Unattended: false}) + if got.Mode != ModeEnforce || got.InactivityTimeout != time.Minute { + t.Errorf("resolve with per-call overrides = %+v, want enforce and 1m", got) + } + if !got.Unattended { + t.Error("resolve cleared Unattended; per-call options can only enable it") + } + + Configure(Options{DisableUIDetection: true}) + if !resolve(Options{}).DisableUIDetection { + t.Error("resolve(Options{}).DisableUIDetection got false after Configure enabled it") + } + + Configure(Options{}) + d = CurrentDefaults() + if d.Mode != defaultMode || d.InactivityTimeout != defaultInactivityTimeout || d.Unattended || d.DisableUIDetection { + t.Errorf("CurrentDefaults() after Configure(Options{}) = %+v, want built-in defaults", d) + } + if !resolve(Options{Unattended: true}).Unattended { + t.Error("resolve(Options{Unattended: true}).Unattended got false") + } +} + +// TestAdminHardTimeoutConfigured verifies that Configure records whether HardTimeout was set, +// including to the built-in default's value. +func TestAdminHardTimeoutConfigured(t *testing.T) { + t.Cleanup(func() { Configure(Options{}) }) + for _, tc := range []struct { + name string + hard time.Duration + want bool + }{ + {"Unset", 0, false}, + {"ExplicitBuiltinValue", defaultHardTimeout, true}, + {"Custom", 3 * time.Hour, true}, + {"Disabled", -1, true}, + } { + Configure(Options{HardTimeout: tc.hard}) + if got := adminHardTimeoutConfigured(); got != tc.want { + t.Errorf("%s: adminHardTimeoutConfigured() after Configure(HardTimeout: %v) = %v, want %v", tc.name, tc.hard, got, tc.want) + } + } +} + +// TestServicingOptions verifies the wusa and DISM servicing policy. +func TestServicingOptions(t *testing.T) { + windir := filepath.Join("X:", "Win") + cbs := filepath.Join(windir, "Logs", "CBS", "CBS.log") + for _, tc := range []struct { + name string + opts Options + adminHard bool + windir string + wantHard time.Duration + wantLogs []string + wantInactivity time.Duration + }{ + {"NothingConfiguredLiftsCap", Options{LogFiles: []string{"pkg.msu.log"}}, false, windir, -1, []string{"pkg.msu.log", cbs}, 0}, + {"AdminExplicit60mRespected", Options{}, true, windir, 0, []string{cbs}, 0}, + {"PackageTimeoutRespected", Options{HardTimeout: 2 * time.Hour, InactivityTimeout: 20 * time.Minute}, false, windir, 2 * time.Hour, []string{cbs}, 20 * time.Minute}, + {"PackageTimeoutWinsOverAdmin", Options{HardTimeout: 2 * time.Hour}, true, windir, 2 * time.Hour, []string{cbs}, 0}, + {"EmptyWindirFallsBack", Options{}, false, "", -1, []string{filepath.Join(`C:\Windows`, "Logs", "CBS", "CBS.log")}, 0}, + } { + t.Run(tc.name, func(t *testing.T) { + got := servicingOptions(tc.opts, tc.adminHard, tc.windir) + if got.HardTimeout != tc.wantHard { + t.Errorf("servicingOptions().HardTimeout = %v, want %v", got.HardTimeout, tc.wantHard) + } + if got.InactivityTimeout != tc.wantInactivity { + t.Errorf("servicingOptions().InactivityTimeout = %v, want %v", got.InactivityTimeout, tc.wantInactivity) + } + if !slices.Equal(got.LogFiles, tc.wantLogs) { + t.Errorf("servicingOptions().LogFiles = %v, want %v", got.LogFiles, tc.wantLogs) + } + }) + } + + // The caller's LogFiles slice must not be modified through a shared backing array. + orig := make([]string, 1, 4) + orig[0] = "a.log" + servicingOptions(Options{LogFiles: orig}, false, windir) + if extended := orig[:2]; extended[1] != "" { + t.Errorf("servicingOptions wrote %q into the caller's LogFiles backing array", extended[1]) + } +} + +// TestFallbackOptions verifies that without a Job Object only the hard timeout stays enforced. +func TestFallbackOptions(t *testing.T) { + in := testOptions(Options{InactivityTimeout: 5 * time.Minute, HardTimeout: time.Hour, Unattended: true}) + got := fallbackOptions(in) + if got.InactivityTimeout != -1 || !got.DisableUIDetection { + t.Errorf("fallbackOptions() = InactivityTimeout %v, DisableUIDetection %v; want -1 and true", got.InactivityTimeout, got.DisableUIDetection) + } + if got.HardTimeout != time.Hour || got.Mode != in.Mode { + t.Errorf("fallbackOptions() = HardTimeout %v, Mode %v; want %v and %v unchanged", got.HardTimeout, got.Mode, time.Hour, in.Mode) + } + if w := newWatchdog(got, time.Now()); w.uiEnabled() { + t.Error("UI detection is enabled after fallbackOptions") + } + + off := testOptions(Options{Mode: ModeOff, InactivityTimeout: time.Minute}) + if got := fallbackOptions(off); got.InactivityTimeout != time.Minute || got.DisableUIDetection { + t.Errorf("fallbackOptions(off mode) = %+v, want the options unchanged", got) + } +} + +// TestSupervise_FallbackKeepsWorkingRootAlive verifies with fallback options that a root process +// showing no progress of its own is terminated only by the hard timeout. +func TestSupervise_FallbackKeepsWorkingRootAlive(t *testing.T) { + tree := &fakeTree{winsAt: persistentWindow(3, "Setup")} + opts := fallbackOptions(testOptions(Options{InactivityTimeout: time.Minute, HardTimeout: 10 * time.Minute, Unattended: true})) + err := runFake(opts, tree, 2*time.Second, 15*time.Minute) + if !errors.Is(err, ErrHardTimeout) { + t.Fatalf("supervise got %v, want ErrHardTimeout", err) + } + if tree.terminatedAt != 10*time.Minute || tree.windowCalls != 0 { + t.Errorf("Terminated at %v with %d windows() call(s), want 10m and 0", tree.terminatedAt, tree.windowCalls) + } +} + +// TestWatchdog_MonitorDedup verifies that monitor mode reports each reason once until it clears. +func TestWatchdog_MonitorDedup(t *testing.T) { + logs := capturedLogsSince() + w := newWatchdog(Options{Mode: ModeMonitor}, time.Now()) + d := &abortDecision{ErrInactivityTimeout, "details"} + for i := 0; i < 5; i++ { + if w.shouldTerminate(d) { + t.Fatal("shouldTerminate returned true in monitor mode") + } + } + w.shouldTerminate(nil) + w.shouldTerminate(d) + if n := strings.Count(logs(), "WOULD_KILL: "); n != 2 { + t.Errorf("Got %d WOULD_KILL lines, want 2 (one per episode)", n) + } + if !newWatchdog(Options{Mode: ModeEnforce}, time.Now()).shouldTerminate(d) { + t.Error("shouldTerminate returned false in enforce mode") + } +} + +// TestCommandDetection verifies msiexec and servicing command detection across path styles. +func TestCommandDetection(t *testing.T) { + for _, tc := range []struct { + path string + msiexec bool + servicing bool + }{ + {`C:\Windows\System32\msiexec.exe`, true, false}, + {`C:\Windows\System32\MSIEXEC.EXE`, true, false}, + {"msiexec", true, false}, + {"/usr/bin/msiexec", true, false}, + {`"C:\Windows\msiexec.exe"`, true, false}, + {`C:\Windows\msiexec2.exe`, false, false}, + {`C:\Windows\System32\wusa.exe`, false, true}, + {"WUSA", false, true}, + {`C:\Windows\System32\Dism.exe`, false, true}, + {"setup.exe", false, false}, + {"", false, false}, + } { + if got := isMsiexecCommand(tc.path); got != tc.msiexec { + t.Errorf("isMsiexecCommand(%q) got %v, want %v", tc.path, got, tc.msiexec) + } + if got := isServicingCommand(tc.path); got != tc.servicing { + t.Errorf("isServicingCommand(%q) got %v, want %v", tc.path, got, tc.servicing) + } + } +} + +// procTable is a fake process table with creation times for tree-walk tests. +type procTable struct { + entries []procEntry + created map[uint32]time.Time +} + +// createdAt returns the creation time of pid, if known. +func (p procTable) createdAt(pid uint32) (time.Time, bool) { + t, ok := p.created[pid] + return t, ok +} + +// newProcTable returns a process table modeled on an MSI transaction with a custom action host, +// a TrustedInstaller servicing operation, unrelated processes and a reused PID. +func newProcTable(base time.Time) procTable { + at := func(s int) time.Time { return base.Add(time.Duration(s) * time.Second) } + p := procTable{created: make(map[uint32]time.Time)} + add := func(pid, ppid uint32, exe string, created int) { + p.entries = append(p.entries, procEntry{pid: pid, ppid: ppid, exe: exe}) + if created >= 0 { + p.created[pid] = at(created) + } + } + add(100, 4, "msiexec.exe", 0) // The msiserver service process (msiexec /V). + add(200, 100, "MsiExec.exe", 10) // A custom action host (msiexec -Embedding). + add(300, 200, "helper.exe", 20) // An EXE launched by the custom action host. + add(301, 300, "cmd.exe", 21) // A grandchild of the custom action host. + add(400, 4, "explorer.exe", 5) // An unrelated process. + add(401, 400, "notepad.exe", 6) // A child of the unrelated process. + add(500, 200, "stale.exe", 5) // Reuses a PID whose recorded parent is 200 but predates it. + add(501, 500, "stalechild.exe", 30) // A child of the reused PID, unreachable from the service. + add(600, 100, "unknown.exe", -1) // A child whose creation time cannot be read. + add(900, 4, "TrustedInstaller.exe", 0) + add(901, 900, "TiWorker.exe", 2) + add(700, 8, "TiWorker.exe", 1) // A TiWorker started by DcomLaunch rather than the service. + add(701, 700, "conhost.exe", 2) + add(800, 801, "a.exe", 1) // A parent cycle, which must not loop forever. + add(801, 800, "b.exe", 1) + return p +} + +// TestProcessForest verifies the service tree walk. +func TestProcessForest(t *testing.T) { + p := newProcTable(time.Date(2026, 9, 28, 12, 0, 0, 0, time.UTC)) + for _, tc := range []struct { + name string + roots []uint32 + want []uint32 + }{ + {"MSIServiceTreeIncludesServiceAndAllDescendants", []uint32{100}, []uint32{100, 200, 300, 301}}, + {"NoServiceRunning", []uint32{0}, nil}, + {"UnrelatedTree", []uint32{400}, []uint32{400, 401}}, + {"TrustedInstallerWithTiWorkerRoots", append([]uint32{900}, pidsWithImage(p.entries, tiWorkerImage)...), []uint32{900, 901, 700, 701}}, + {"ParentCycle", []uint32{800}, []uint32{800, 801}}, + } { + t.Run(tc.name, func(t *testing.T) { + if got := processForest(p.entries, tc.roots, p.createdAt); !slices.Equal(got, tc.want) { + t.Errorf("processForest(%v) got %v, want %v", tc.roots, got, tc.want) + } + }) + } +} + +// TestSelectTerminable verifies that only descendants created after supervision started may be +// terminated, and never a protected service process. +func TestSelectTerminable(t *testing.T) { + base := time.Date(2026, 9, 28, 12, 0, 0, 0, time.UTC) + p := newProcTable(base) + tree := []uint32{100, 200, 300, 301, 600} + for _, tc := range []struct { + name string + notBefore time.Time + want []uint32 + }{ + {"AllDescendantsAfterStart", base, []uint32{200, 300, 301}}, + {"OnlyLaterDescendants", base.Add(15 * time.Second), []uint32{300, 301}}, + {"NoneAfterStart", base.Add(time.Hour), nil}, + } { + t.Run(tc.name, func(t *testing.T) { + if got := selectTerminable(tree, []uint32{100}, p.createdAt, tc.notBefore); !slices.Equal(got, tc.want) { + t.Errorf("selectTerminable got %v, want %v", got, tc.want) + } + }) + } + late := procTable{created: map[uint32]time.Time{100: base.Add(time.Hour)}} + if got := selectTerminable([]uint32{100}, []uint32{100}, late.createdAt, base); got != nil { + t.Errorf("selectTerminable returned the protected service PID: %v", got) + } +} diff --git a/supervisor/supervisor_unix.go b/supervisor/supervisor_unix.go new file mode 100644 index 0000000..967a327 --- /dev/null +++ b/supervisor/supervisor_unix.go @@ -0,0 +1,252 @@ +//go:build !windows + +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "bufio" + "bytes" + "errors" + "io" + "io/fs" + "os" + "os/exec" + "path/filepath" + "strconv" + "strings" + "syscall" + "time" +) + +// clockTicksPerSecond is USER_HZ, the unit of utime and stime in /proc//stat on Linux. +const clockTicksPerSecond = 100 + +// procRoot is the procfs mount point; tests may override it. +var procRoot = "/proc" + +// isSession0 reports whether the current process runs in Windows Session 0, which is never true +// outside Windows. +func isSession0() bool { return false } + +// procStat holds the fields of /proc//stat used by the supervisor. +type procStat struct { + state byte + pgrp int + cpu time.Duration + // start is the process start time in clock ticks after boot, or 0 if the line is too short. + start int64 +} + +// parseProcStat parses the contents of /proc//stat. The command name may contain spaces +// and parentheses, so fields are located after the last ')'. +func parseProcStat(data []byte) (procStat, bool) { + idx := bytes.LastIndexByte(data, ')') + if idx < 0 || idx+2 >= len(data) { + return procStat{}, false + } + // Fields after ')' start at field 3 (state); pgrp is field 5, utime and stime are 14 and 15 + // and starttime is 22. + fields := strings.Fields(string(data[idx+2:])) + if len(fields) < 13 || len(fields[0]) != 1 { + return procStat{}, false + } + pgrp, err1 := strconv.Atoi(fields[2]) + utime, err2 := strconv.ParseInt(fields[11], 10, 64) + stime, err3 := strconv.ParseInt(fields[12], 10, 64) + if err1 != nil || err2 != nil || err3 != nil { + return procStat{}, false + } + st := procStat{ + state: fields[0][0], + pgrp: pgrp, + cpu: time.Duration(utime+stime) * time.Second / clockTicksPerSecond, + } + if len(fields) > 19 { + if start, err := strconv.ParseInt(fields[19], 10, 64); err == nil { + st.start = start + } + } + return st, true +} + +// parseProcIO returns read_bytes plus write_bytes from the contents of /proc//io. +func parseProcIO(data []byte) uint64 { + var total uint64 + sc := bufio.NewScanner(bytes.NewReader(data)) + for sc.Scan() { + k, v, ok := strings.Cut(sc.Text(), ":") + if !ok { + continue + } + switch strings.TrimSpace(k) { + case "read_bytes", "write_bytes": + n, err := strconv.ParseUint(strings.TrimSpace(v), 10, 64) + if err == nil { + total += n + } + } + } + return total +} + +// procStatPath returns the path of /proc//stat under procRoot. +func procStatPath(pid int) string { + return filepath.Join(procRoot, strconv.Itoa(pid), "stat") +} + +// readProcCounters returns the counters, start time and process group of pid, or ok=false if +// unreadable. I/O counters that are unreadable, for example due to permissions, are reported as zero. +func readProcCounters(pid int) (counters, int64, int, bool) { + data, err := os.ReadFile(procStatPath(pid)) + if err != nil { + return counters{}, 0, 0, false + } + st, ok := parseProcStat(data) + if !ok { + return counters{}, 0, 0, false + } + c := counters{cpu: st.cpu} + if ioData, err := os.ReadFile(filepath.Join(procRoot, strconv.Itoa(pid), "io")); err == nil { + c.io = parseProcIO(ioData) + } + return c, st.start, st.pgrp, true +} + +// sampleProcessGroup returns per-process counters for every process in process group pgid. If +// /proc cannot be enumerated it falls back to the root process only. +func sampleProcessGroup(pgid, rootPID int) map[procKey]counters { + out := make(map[procKey]counters) + entries, err := os.ReadDir(procRoot) + if err != nil { + if c, start, _, ok := readProcCounters(rootPID); ok { + out[procKey{uint32(rootPID), start}] = c + } + return out + } + for _, e := range entries { + pid, err := strconv.Atoi(e.Name()) + if err != nil || pid <= 0 { + continue + } + c, start, pgrp, ok := readProcCounters(pid) + if ok && pgrp == pgid { + out[procKey{uint32(pid), start}] = c + } + } + if len(out) == 0 { + if c, start, _, ok := readProcCounters(rootPID); ok { + out[procKey{uint32(rootPID), start}] = c + } + } + return out +} + +// procExited reports whether pid has exited according to procfs: it is a zombie, or its entry +// is gone. It never calls wait, so the reap stays with exec.Cmd.Wait. procfsWorks must be true +// only if the process's stat file was readable earlier, so that a missing entry means it was +// reaped rather than that procfs is unavailable. +func procExited(pid int, procfsWorks bool) bool { + if !procfsWorks { + return false + } + data, err := os.ReadFile(procStatPath(pid)) + if err != nil { + return errors.Is(err, fs.ErrNotExist) + } + st, ok := parseProcStat(data) + return ok && st.state == 'Z' +} + +// processGroupTree supervises a process group on Unix systems. +type processGroupTree struct { + c *exec.Cmd + pid, pgid int + detector windowDetectorFunc + procfsWorks bool + agg pidCounters + cum progressSample +} + +// sample returns the cumulative counters of the process group. +func (t *processGroupTree) sample(now time.Time) progressSample { + dCPU, dIO := t.agg.update(sampleProcessGroup(t.pgid, t.pid)) + t.cum.cpu += dCPU + t.cum.io += dIO + t.cum.at = now + return t.cum +} + +// windows returns windows reported by the test detector; Unix has no native window detection. +func (t *processGroupTree) windows() []windowInfo { + if t.detector == nil { + return nil + } + wins, _ := t.detector([]uint32{uint32(t.pid)}) + return wins +} + +// terminate sends SIGKILL to the whole process group and to the root process. +func (t *processGroupTree) terminate() { + _ = syscall.Kill(-t.pgid, syscall.SIGKILL) + _ = t.c.Process.Kill() +} + +// afterAbort does nothing because Unix installers have no out-of-tree service side. +func (t *processGroupTree) afterAbort() {} + +// rootExited reports whether the root process is a zombie or was already reaped. On systems +// without procfs it always returns false. +func (t *processGroupTree) rootExited() bool { + return procExited(t.pid, t.procfsWorks) +} + +// runSupervised executes a process in its own process group with forward progress, hard timeout +// and optional UI watchdog monitoring on Unix systems. +func runSupervised(c *exec.Cmd, opts Options, out io.Writer) error { + devNull, err := setupStdio(c, out) + if err != nil { + return err + } + defer devNull.Close() + + if c.SysProcAttr == nil { + c.SysProcAttr = &syscall.SysProcAttr{} + } + c.SysProcAttr.Setpgid = true + + if err := c.Start(); err != nil { + return err + } + pid := c.Process.Pid + pgid, err := syscall.Getpgid(pid) + if err != nil || pgid <= 0 { + pgid = pid + } + _, statErr := os.Stat(procStatPath(pid)) + + waitErr := make(chan error, 1) + go func() { waitErr <- c.Wait() }() + + tree := &processGroupTree{ + c: c, + pid: pid, + pgid: pgid, + detector: opts.windowDetector, + procfsWorks: statErr == nil, + } + ticker := time.NewTicker(opts.pollInterval) + defer ticker.Stop() + return supervise(opts, waitErr, tree, time.Now(), ticker.C) +} diff --git a/supervisor/supervisor_unix_test.go b/supervisor/supervisor_unix_test.go new file mode 100644 index 0000000..459c1a3 --- /dev/null +++ b/supervisor/supervisor_unix_test.go @@ -0,0 +1,346 @@ +//go:build !windows + +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "errors" + "fmt" + "io" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "syscall" + "testing" + "time" +) + +func init() { + platformHelpers["setsid_grandchild"] = func(args []string) { + // Starts a grandchild in a new session that inherits stdout, records its PID in args[0], + // and then either exits (args[1] == "exit") or stalls. + gc := spawnHelper("stall") + gc.Stdout = os.Stdout + gc.Stderr = os.Stderr + gc.SysProcAttr = &syscall.SysProcAttr{Setsid: true} + if err := gc.Start(); err != nil { + fmt.Fprintf(os.Stderr, "Failed to start grandchild: %v\n", err) + os.Exit(1) + } + if err := os.WriteFile(args[0], []byte(strconv.Itoa(gc.Process.Pid)), 0644); err != nil { + os.Exit(1) + } + if len(args) > 1 && args[1] == "exit" { + return + } + time.Sleep(60 * time.Second) + } +} + +// processAlive reports whether pid exists and is not a zombie. +func processAlive(pid int) bool { + if runtime.GOOS == "linux" { + data, err := os.ReadFile(fmt.Sprintf("/proc/%d/stat", pid)) + if err != nil { + return false + } + if i := strings.LastIndexByte(string(data), ')'); i >= 0 && i+2 < len(data) { + return data[i+2] != 'Z' + } + return true + } + return syscall.Kill(pid, 0) == nil +} + +// waitForExit polls until pid has exited or timeout elapses, and reports whether it exited. +func waitForExit(pid int, timeout time.Duration) bool { + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + if !processAlive(pid) { + return true + } + time.Sleep(20 * time.Millisecond) + } + return !processAlive(pid) +} + +// readPIDFile waits for a helper to write a PID to path and returns it. +func readPIDFile(t *testing.T, path string) int { + t.Helper() + deadline := time.Now().Add(10 * time.Second) + for time.Now().Before(deadline) { + if data, err := os.ReadFile(path); err == nil && len(data) > 0 { + pid, err := strconv.Atoi(strings.TrimSpace(string(data))) + if err == nil { + return pid + } + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("Helper did not write a PID to %s", path) + return 0 +} + +// killOnCleanup kills the process recorded in pidFile when the test ends. +func killOnCleanup(t *testing.T, pidFile string) { + t.Cleanup(func() { + if data, err := os.ReadFile(pidFile); err == nil { + if pid, err := strconv.Atoi(strings.TrimSpace(string(data))); err == nil && pid > 0 { + _ = syscall.Kill(pid, syscall.SIGKILL) + } + } + }) +} + +// TestParseProcStat verifies stat parsing with spaces and parentheses in the command name. +func TestParseProcStat(t *testing.T) { + line := "1234 (my (odd) proc) S 1 777 777 0 -1 4194560 100 0 0 0 150 50 0 0 20 0 1 0 100 0 0\n" + st, ok := parseProcStat([]byte(line)) + if !ok { + t.Fatal("parseProcStat failed on a valid line") + } + if st.pgrp != 777 { + t.Errorf("pgrp got %d, want 777", st.pgrp) + } + if st.state != 'S' { + t.Errorf("state got %q, want 'S'", st.state) + } + if st.cpu != 2*time.Second { + t.Errorf("cpu got %v, want 2s from 150+50 ticks", st.cpu) + } + if st.start != 100 { + t.Errorf("start got %d, want 100", st.start) + } + if short, ok := parseProcStat([]byte("1 (p) S 1 1 1 0 -1 0 0 0 0 0 7 3\n")); !ok || short.start != 0 { + t.Errorf("parseProcStat(short line) = %+v, %v; want start 0 and ok", short, ok) + } + if _, ok := parseProcStat([]byte("garbage")); ok { + t.Error("parseProcStat accepted garbage") + } +} + +// TestParseProcIO verifies that only storage read and write bytes are summed. +func TestParseProcIO(t *testing.T) { + data := "rchar: 999\nwchar: 999\nsyscr: 1\nsyscw: 1\nread_bytes: 4096\nwrite_bytes: 8192\ncancelled_write_bytes: 0\n" + if got := parseProcIO([]byte(data)); got != 12288 { + t.Errorf("parseProcIO got %d, want 12288", got) + } +} + +// byPID re-keys per-process counters by PID alone. +func byPID(m map[procKey]counters) map[uint32]counters { + out := make(map[uint32]counters, len(m)) + for k, c := range m { + out[k.pid] = c + } + return out +} + +// TestSampleProcessGroup verifies pgid filtering and root fallback against a fake procfs. +func TestSampleProcessGroup(t *testing.T) { + root := t.TempDir() + write := func(pid, pgrp, utime int, io string) { + dir := filepath.Join(root, strconv.Itoa(pid)) + if err := os.MkdirAll(dir, 0755); err != nil { + t.Fatal(err) + } + stat := fmt.Sprintf("%d (p) S 1 %d %d 0 -1 0 0 0 0 0 %d 0 0 0 20 0 1 0 0 0 0\n", pid, pgrp, pgrp, utime) + if err := os.WriteFile(filepath.Join(dir, "stat"), []byte(stat), 0644); err != nil { + t.Fatal(err) + } + if io != "" { + if err := os.WriteFile(filepath.Join(dir, "io"), []byte(io), 0644); err != nil { + t.Fatal(err) + } + } + } + write(10, 10, 100, "read_bytes: 1\nwrite_bytes: 2\n") + write(11, 10, 50, "") + write(12, 99, 500, "read_bytes: 1000\n") + if err := os.MkdirAll(filepath.Join(root, "self"), 0755); err != nil { + t.Fatal(err) + } + + old := procRoot + procRoot = root + defer func() { procRoot = old }() + + got := byPID(sampleProcessGroup(10, 10)) + if len(got) != 2 || got[10].cpu != time.Second || got[10].io != 3 || got[11].cpu != 500*time.Millisecond { + t.Errorf("sampleProcessGroup got %+v, want pids 10 and 11 only", got) + } + + procRoot = filepath.Join(root, "missing") + if got := sampleProcessGroup(10, 10); len(got) != 0 { + t.Errorf("sampleProcessGroup without procfs got %+v, want empty", got) + } +} + +// TestRun_PgidAggregation verifies that CPU burned by a child keeps the tree alive while the +// root process only sleeps. +func TestRun_PgidAggregation(t *testing.T) { + if runtime.GOOS != "linux" { + t.Skip("Process group accounting requires /proc") + } + cmd := helperCommand(t, "parent_sleep_child_burn", "2500") + opts := Options{ + InactivityTimeout: 1 * time.Second, + pollInterval: 20 * time.Millisecond, + minCPUDelta: 50 * time.Millisecond, + } + start := time.Now() + if err := Run(cmd, opts, io.Discard); err != nil { + t.Fatalf("Run returned %v while a child was burning CPU, want nil", err) + } + if elapsed := time.Since(start); elapsed < 2300*time.Millisecond { + t.Errorf("Run completed in %v, want at least the child's 2.5s of work", elapsed) + } +} + +// TestRun_SysProcAttrMerged verifies that Setpgid is merged into a caller-provided SysProcAttr. +func TestRun_SysProcAttrMerged(t *testing.T) { + cmd := helperCommand(t, "sleep_ms", "10") + attr := &syscall.SysProcAttr{Setpgid: false} + cmd.SysProcAttr = attr + if err := Run(cmd, Options{InactivityTimeout: -1}, io.Discard); err != nil { + t.Fatalf("Run returned %v, want nil", err) + } + if cmd.SysProcAttr != attr || !attr.Setpgid { + t.Errorf("SysProcAttr got %+v (same pointer: %v), want the caller's struct with Setpgid set", cmd.SysProcAttr, cmd.SysProcAttr == attr) + } +} + +// TestRun_BoundedWait_WaitDelay verifies that a grandchild outside the process group holding +// stdout cannot block Run after an abort. +func TestRun_BoundedWait_WaitDelay(t *testing.T) { + pidFile := filepath.Join(t.TempDir(), "gc.pid") + killOnCleanup(t, pidFile) + cmd := helperCommand(t, "setsid_grandchild", pidFile) + cmd.WaitDelay = 500 * time.Millisecond + opts := Options{InactivityTimeout: 500 * time.Millisecond, pollInterval: 20 * time.Millisecond} + + start := time.Now() + err := Run(cmd, opts, io.Discard) + elapsed := time.Since(start) + if !errors.Is(err, ErrInactivityTimeout) { + t.Fatalf("Run got error %v, want ErrInactivityTimeout", err) + } + if elapsed > 20*time.Second { + t.Errorf("Run took %v, want it bounded by WaitDelay", elapsed) + } + if pid := readPIDFile(t, pidFile); !processAlive(pid) { + t.Errorf("Grandchild %d exited early; the held-pipe scenario was not exercised", pid) + } +} + +// TestRun_BoundedWait_KillWaitTimeout verifies the last-resort bound on Wait after an abort. +func TestRun_BoundedWait_KillWaitTimeout(t *testing.T) { + pidFile := filepath.Join(t.TempDir(), "gc.pid") + killOnCleanup(t, pidFile) + old := killWaitTimeout + killWaitTimeout = 500 * time.Millisecond + defer func() { killWaitTimeout = old }() + + cmd := helperCommand(t, "setsid_grandchild", pidFile) + cmd.WaitDelay = time.Hour + opts := Options{InactivityTimeout: 500 * time.Millisecond, pollInterval: 20 * time.Millisecond} + + start := time.Now() + err := Run(cmd, opts, io.Discard) + elapsed := time.Since(start) + if !errors.Is(err, ErrInactivityTimeout) { + t.Fatalf("Run got error %v, want ErrInactivityTimeout", err) + } + if !strings.Contains(err.Error(), "pipes are held by processes outside the supervised tree") { + t.Errorf("Run error %q lacks the held-pipe diagnostic", err) + } + if elapsed > 20*time.Second { + t.Errorf("Run took %v, want it bounded by killWaitTimeout", elapsed) + } +} + +// TestRun_SuccessWithHeldPipe verifies that a clean exit is not reported as an error when a +// detached grandchild keeps stdout open past WaitDelay. +func TestRun_SuccessWithHeldPipe(t *testing.T) { + pidFile := filepath.Join(t.TempDir(), "gc.pid") + killOnCleanup(t, pidFile) + cmd := helperCommand(t, "setsid_grandchild", pidFile, "exit") + cmd.WaitDelay = 300 * time.Millisecond + start := time.Now() + if err := Run(cmd, Options{InactivityTimeout: -1}, io.Discard); err != nil { + t.Fatalf("Run returned %v, want nil after a clean exit", err) + } + if elapsed := time.Since(start); elapsed > 20*time.Second { + t.Errorf("Run took %v, want it bounded by WaitDelay", elapsed) + } +} + +// TestProcExited verifies zombie and reaped detection against a fake procfs. +func TestProcExited(t *testing.T) { + root := t.TempDir() + write := func(pid int, state string) { + dir := filepath.Join(root, strconv.Itoa(pid)) + if err := os.MkdirAll(dir, 0755); err != nil { + t.Fatal(err) + } + stat := fmt.Sprintf("%d (p) %s 1 %d %d 0 -1 0 0 0 0 0 0 0 0 0 20 0 1 0 0 0 0\n", pid, state, pid, pid) + if err := os.WriteFile(filepath.Join(dir, "stat"), []byte(stat), 0644); err != nil { + t.Fatal(err) + } + } + write(10, "S") + write(11, "Z") + + old := procRoot + procRoot = root + defer func() { procRoot = old }() + + for _, tc := range []struct { + name string + pid int + procfsWorks bool + want bool + }{ + {"Running", 10, true, false}, + {"Zombie", 11, true, true}, + {"Reaped", 12, true, true}, + {"NoProcfs", 12, false, false}, + } { + if got := procExited(tc.pid, tc.procfsWorks); got != tc.want { + t.Errorf("%s: procExited(%d, %v) got %v, want %v", tc.name, tc.pid, tc.procfsWorks, got, tc.want) + } + } +} + +// TestRun_RootExitedWithHeldPipeNotKilled verifies that after the root process exits cleanly, a +// detached grandchild holding stdout does not turn the install into a watchdog failure. +func TestRun_RootExitedWithHeldPipeNotKilled(t *testing.T) { + if runtime.GOOS != "linux" { + t.Skip("Root exit detection requires /proc") + } + pidFile := filepath.Join(t.TempDir(), "gc.pid") + killOnCleanup(t, pidFile) + cmd := helperCommand(t, "setsid_grandchild", pidFile, "exit") + cmd.WaitDelay = 3 * time.Second + opts := Options{InactivityTimeout: time.Second, pollInterval: 20 * time.Millisecond} + if err := Run(cmd, opts, io.Discard); err != nil { + t.Fatalf("Run returned %v, want nil because the root exited cleanly before the inactivity timeout", err) + } + if pid := readPIDFile(t, pidFile); !processAlive(pid) { + t.Errorf("Grandchild %d exited early; the held-pipe scenario was not exercised", pid) + } +} diff --git a/supervisor/supervisor_windows.go b/supervisor/supervisor_windows.go new file mode 100644 index 0000000..63cc525 --- /dev/null +++ b/supervisor/supervisor_windows.go @@ -0,0 +1,594 @@ +//go:build windows + +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "errors" + "fmt" + "io" + "os" + "os/exec" + "sync" + "syscall" + "time" + "unsafe" + + "github.com/google/logger" + "golang.org/x/sys/windows" +) + +var ( + user32 = windows.NewLazySystemDLL("user32.dll") + procEnumWindows = user32.NewProc("EnumWindows") + procGetWindowTextW = user32.NewProc("GetWindowTextW") + procGetWindow = user32.NewProc("GetWindow") + procGetWindowLongW = user32.NewProc("GetWindowLongW") + enumWindowsCallback = windows.NewCallback(enumWindowsProc) +) + +const ( + // dialogClassName is the window class of standard Win32 dialog boxes, including MessageBox. + dialogClassName = "#32770" + // gwOwner is the GetWindow command that retrieves the owner window. + gwOwner = 4 + // wsExDlgModalFrame is the extended window style of windows with a modal dialog frame. + wsExDlgModalFrame = 0x00000001 +) + +// gwlExStyle is the GetWindowLong index of the extended window style. It is a variable because +// the negative constant cannot be converted to uintptr directly. +var gwlExStyle int32 = -20 + +// jobObjectBasicAccountingInformation mirrors Win32 JOBOBJECT_BASIC_ACCOUNTING_INFORMATION from winnt.h. +type jobObjectBasicAccountingInformation struct { + TotalUserTime int64 + TotalKernelTime int64 + ThisPeriodTotalUserTime int64 + ThisPeriodTotalKernelTime int64 + TotalPageFaultCount uint32 + TotalProcesses uint32 + ActiveProcesses uint32 + TotalTerminatedProcesses uint32 +} + +// jobObjectBasicAndIoAccountingInformation mirrors Win32 JOBOBJECT_BASIC_AND_IO_ACCOUNTING_INFORMATION from winnt.h. +type jobObjectBasicAndIoAccountingInformation struct { + BasicInfo jobObjectBasicAccountingInformation + IoInfo windows.IO_COUNTERS +} + +// jobObjectBasicProcessIDListHeader describes the fixed header of JOBOBJECT_BASIC_PROCESS_ID_LIST. +type jobObjectBasicProcessIDListHeader struct { + NumberOfAssignedProcesses uint32 + NumberOfProcessIdsInList uint32 +} + +// getJobPIDs returns all active process IDs assigned to the given Job Object plus rootPID. +func getJobPIDs(job windows.Handle, rootPID uint32) []uint32 { + pidsMap := make(map[uint32]bool) + if rootPID != 0 { + pidsMap[rootPID] = true + } + if job != 0 { + const initialCap = 64 + uintptrSize := int(unsafe.Sizeof(uintptr(0))) + // The ULONG_PTR PID list follows the two DWORD header fields on both 32-bit and 64-bit. + const listOffset = 8 + buf := make([]byte, listOffset+initialCap*uintptrSize) + + var retLen uint32 + err := windows.QueryInformationJobObject( + job, + int32(windows.JobObjectBasicProcessIdList), + uintptr(unsafe.Pointer(&buf[0])), + uint32(len(buf)), + &retLen, + ) + if err != nil && errors.Is(err, windows.ERROR_MORE_DATA) { + header := (*jobObjectBasicProcessIDListHeader)(unsafe.Pointer(&buf[0])) + needed := header.NumberOfAssignedProcesses + if needed > 0 { + buf = make([]byte, listOffset+int(needed+16)*uintptrSize) + err = windows.QueryInformationJobObject( + job, + int32(windows.JobObjectBasicProcessIdList), + uintptr(unsafe.Pointer(&buf[0])), + uint32(len(buf)), + &retLen, + ) + } + } + if err == nil { + header := (*jobObjectBasicProcessIDListHeader)(unsafe.Pointer(&buf[0])) + numInList := int(header.NumberOfProcessIdsInList) + for i := 0; i < numInList; i++ { + pidOffset := listOffset + i*uintptrSize + if pidOffset+uintptrSize > len(buf) { + break + } + pidVal := *(*uintptr)(unsafe.Pointer(&buf[pidOffset])) + if pidVal != 0 { + pidsMap[uint32(pidVal)] = true + } + } + } + } + pids := make([]uint32, 0, len(pidsMap)) + for pid := range pidsMap { + pids = append(pids, pid) + } + return pids +} + +// windowEnumContext passes Job Object process IDs and collects matching windows during enumeration. +type windowEnumContext struct { + jobPIDs map[uint32]bool + matches []windowInfo +} + +// enumContexts maps integer handles passed through EnumWindows' lParam to their contexts. Passing +// an integer instead of a Go pointer avoids converting a uintptr back to unsafe.Pointer in the +// callback, the same approach as runtime/cgo.Handle. +var ( + enumContextsMu sync.Mutex + enumContexts = make(map[uintptr]*windowEnumContext) + nextEnumContextH uintptr +) + +// registerEnumContext stores ctx and returns its handle. +func registerEnumContext(ctx *windowEnumContext) uintptr { + enumContextsMu.Lock() + defer enumContextsMu.Unlock() + nextEnumContextH++ + if nextEnumContextH == 0 { + nextEnumContextH++ + } + enumContexts[nextEnumContextH] = ctx + return nextEnumContextH +} + +// lookupEnumContext returns the context registered under h, or nil. +func lookupEnumContext(h uintptr) *windowEnumContext { + enumContextsMu.Lock() + defer enumContextsMu.Unlock() + return enumContexts[h] +} + +// unregisterEnumContext removes the context registered under h. +func unregisterEnumContext(h uintptr) { + enumContextsMu.Lock() + defer enumContextsMu.Unlock() + delete(enumContexts, h) +} + +// isCandidateWindow reports whether a visible top-level window looks like an interactive prompt: +// a standard dialog box, or an owned window with a modal dialog frame. +func isCandidateWindow(className string, hasOwner bool, exStyle uint32) bool { + if className == dialogClassName { + return true + } + return hasOwner && exStyle&wsExDlgModalFrame != 0 +} + +// enumWindowsProc is the callback function for EnumWindows. +func enumWindowsProc(hwnd uintptr, lParam uintptr) uintptr { + ctx := lookupEnumContext(lParam) + if ctx == nil { + return 0 + } + h := windows.HWND(hwnd) + var pid uint32 + tid, err := windows.GetWindowThreadProcessId(h, &pid) + if tid == 0 || err != nil || !ctx.jobPIDs[pid] { + return 1 + } + if !windows.IsWindowVisible(h) { + return 1 + } + + var classBuf [256]uint16 + copied, err := windows.GetClassName(h, &classBuf[0], int32(len(classBuf))) + if err != nil || copied == 0 { + return 1 + } + className := windows.UTF16ToString(classBuf[:copied]) + owner, _, _ := procGetWindow.Call(hwnd, gwOwner) + exStyle, _, _ := procGetWindowLongW.Call(hwnd, uintptr(gwlExStyle)) + + if isCandidateWindow(className, owner != 0, uint32(exStyle)) { + ctx.matches = append(ctx.matches, windowInfo{ + PID: pid, + HWND: hwnd, + Title: getWindowTitle(h), + ClassName: className, + ExePath: getProcessImagePath(pid), + }) + } + return 1 +} + +// getWindowTitle retrieves the title text of the given window. +func getWindowTitle(hwnd windows.HWND) string { + var buf [512]uint16 + r0, _, _ := procGetWindowTextW.Call( + uintptr(hwnd), + uintptr(unsafe.Pointer(&buf[0])), + uintptr(len(buf)), + ) + if r0 > 0 { + return windows.UTF16ToString(buf[:r0]) + } + return "" +} + +// getProcessImagePath retrieves the full executable image path for the given process ID. +func getProcessImagePath(pid uint32) string { + hProc, err := windows.OpenProcess(windows.PROCESS_QUERY_LIMITED_INFORMATION, false, pid) + if err != nil { + return "" + } + defer windows.CloseHandle(hProc) + + var buf [1024]uint16 + size := uint32(len(buf)) + if err := windows.QueryFullProcessImageName(hProc, 0, &buf[0], &size); err != nil { + return "" + } + return windows.UTF16ToString(buf[:size]) +} + +// isSession0 checks whether the current process is running in Windows Session 0. +func isSession0() bool { + var sessionID uint32 + if err := windows.ProcessIdToSessionId(uint32(os.Getpid()), &sessionID); err != nil { + return false + } + return sessionID == 0 +} + +// detectWindowsWin32 enumerates visible top-level windows owned by pids and returns candidate +// interactive prompts. +func detectWindowsWin32(pids []uint32) ([]windowInfo, error) { + ctx := &windowEnumContext{jobPIDs: make(map[uint32]bool, len(pids))} + for _, p := range pids { + ctx.jobPIDs[p] = true + } + h := registerEnumContext(ctx) + defer unregisterEnumContext(h) + r1, _, e1 := procEnumWindows.Call(enumWindowsCallback, h) + if r1 == 0 && e1 != nil && !errors.Is(e1, windows.ERROR_SUCCESS) { + return ctx.matches, e1 + } + return ctx.matches, nil +} + +// filetimeDuration converts a FILETIME interval, in 100ns units, to a time.Duration. +func filetimeDuration(ft windows.Filetime) time.Duration { + return time.Duration(uint64(ft.HighDateTime)<<32|uint64(ft.LowDateTime)) * 100 +} + +// sampleJob returns the cumulative CPU time and read plus write bytes of every process that has +// ever run in the job. OtherTransferCount is excluded because control I/O is not install progress. +func sampleJob(job windows.Handle) (counters, error) { + var info jobObjectBasicAndIoAccountingInformation + err := windows.QueryInformationJobObject( + job, + int32(windows.JobObjectBasicAndIoAccountingInformation), + uintptr(unsafe.Pointer(&info)), + uint32(unsafe.Sizeof(info)), + nil, + ) + if err != nil { + return counters{}, err + } + return counters{ + cpu: time.Duration(info.BasicInfo.TotalUserTime+info.BasicInfo.TotalKernelTime) * 100, + io: info.IoInfo.ReadTransferCount + info.IoInfo.WriteTransferCount, + }, nil +} + +// setJobLimitFlags sets the basic limit flags of the job. +func setJobLimitFlags(job windows.Handle, flags uint32) error { + info := windows.JOBOBJECT_EXTENDED_LIMIT_INFORMATION{ + BasicLimitInformation: windows.JOBOBJECT_BASIC_LIMIT_INFORMATION{ + LimitFlags: flags, + }, + } + _, err := windows.SetInformationJobObject( + job, + windows.JobObjectExtendedLimitInformation, + uintptr(unsafe.Pointer(&info)), + uint32(unsafe.Sizeof(info)), + ) + return err +} + +// resumeProcessThreads resumes every thread owned by pid. It is used to start a process created +// with CREATE_SUSPENDED once it has been assigned to the job. Threads that cannot be opened or +// resumed, for example threads injected by security software, are skipped with a warning; it +// fails only if no thread was resumed. +func resumeProcessThreads(pid uint32) error { + snap, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPTHREAD, 0) + if err != nil { + return fmt.Errorf("creating thread snapshot: %w", err) + } + defer windows.CloseHandle(snap) + + te := windows.ThreadEntry32{Size: uint32(unsafe.Sizeof(windows.ThreadEntry32{}))} + resumed := 0 + var lastErr error + for err = windows.Thread32First(snap, &te); err == nil; err = windows.Thread32Next(snap, &te) { + if te.OwnerProcessID != pid { + continue + } + th, oerr := windows.OpenThread(windows.THREAD_SUSPEND_RESUME, false, te.ThreadID) + if oerr != nil { + lastErr = fmt.Errorf("opening thread %d: %w", te.ThreadID, oerr) + logger.Warningf("Skipping thread %d of PID %d: %v", te.ThreadID, pid, oerr) + continue + } + _, rerr := windows.ResumeThread(th) + windows.CloseHandle(th) + if rerr != nil { + lastErr = fmt.Errorf("resuming thread %d: %w", te.ThreadID, rerr) + logger.Warningf("Failed to resume thread %d of PID %d: %v", te.ThreadID, pid, rerr) + continue + } + resumed++ + } + if !errors.Is(err, windows.ERROR_NO_MORE_FILES) { + return fmt.Errorf("enumerating threads: %w", err) + } + if resumed == 0 { + if lastErr != nil { + return fmt.Errorf("no thread of process %d was resumed: %w", pid, lastErr) + } + return fmt.Errorf("no threads found for process %d", pid) + } + return nil +} + +// createKillOnCloseJob creates a Job Object that kills its processes when its last handle closes. +func createKillOnCloseJob() (windows.Handle, error) { + job, err := windows.CreateJobObject(nil, nil) + if err != nil { + return 0, fmt.Errorf("creating Windows Job Object: %w", err) + } + if err := setJobLimitFlags(job, windows.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE); err != nil { + windows.CloseHandle(job) + return 0, fmt.Errorf("setting JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE: %w", err) + } + return job, nil +} + +// createJob creates the Job Object for a supervised install; tests replace it to simulate systems +// where job creation fails. +var createJob = createKillOnCloseJob + +// jobTree supervises a Job Object on Windows, or only the root process when a job could not be +// set up. +type jobTree struct { + c *exec.Cmd + opts Options + pid uint32 + // job is zero when only the root process is supervised. + job windows.Handle + proc windows.Handle + start time.Time + last counters + // checkedImages records PIDs whose image name has been inspected. + checkedImages map[uint32]bool + rootMsiexec bool + ranMsiexec bool + msi *serviceMonitor + // servicing accounts TrustedInstaller work for wusa and DISM; it is nil until needed. + servicing *serviceMonitor +} + +// processIDs returns the PIDs of the job's active processes, or only the root PID without a job. +func (t *jobTree) processIDs() []uint32 { + return getJobPIDs(t.job, t.pid) +} + +// noteImages inspects the image names of newly seen processes. An msiexec process enables the +// MSI post-abort wait, and a wusa or DISM process enables TrustedInstaller accounting. +func (t *jobTree) noteImages(pids []uint32) { + for _, pid := range pids { + if t.checkedImages[pid] { + continue + } + name := imageBaseName(getProcessImagePath(pid)) + if name == "" { + // Retry on the next poll; the process may not be queryable yet. + continue + } + t.checkedImages[pid] = true + switch name { + case msiexecImage: + t.ranMsiexec = true + case wusaImage, dismImage: + if t.servicing == nil { + logger.Infof("Installer started %s (PID %d); accounting TrustedInstaller servicing work as progress.", name, pid) + t.servicing = newTrustedInstallerMonitor(t.opts, t.start) + } + } + } +} + +// sample returns the job's counters plus the counters of service trees working for it. +func (t *jobTree) sample(now time.Time) progressSample { + if t.job != 0 { + if jc, err := sampleJob(t.job); err == nil { + t.last = jc + } + } else if pc, _, ok := handleCounters(t.proc); ok { + t.last = pc + } + t.noteImages(t.processIDs()) + s := progressSample{at: now, cpu: t.last.cpu, io: t.last.io} + for _, m := range []*serviceMonitor{t.msi, t.servicing} { + if m != nil { + mc := m.sample(now) + s.cpu += mc.cpu + s.io += mc.io + } + } + return s +} + +// windows returns candidate interactive windows owned by the job's processes. +func (t *jobTree) windows() []windowInfo { + pids := t.processIDs() + var wins []windowInfo + if t.opts.windowDetector != nil { + wins, _ = t.opts.windowDetector(pids) + } else { + wins, _ = detectWindowsWin32(pids) + } + return wins +} + +// terminate kills every process in the job, or the root process without a job. Service +// processes are never part of the job and are not affected. +func (t *jobTree) terminate() { + if t.job != 0 { + _ = windows.TerminateJobObject(t.job, 1) + return + } + _ = t.c.Process.Kill() +} + +// afterAbort waits for the Windows Installer transaction if the install ran msiexec. +// TrustedInstaller servicing is never waited for or terminated here. +func (t *jobTree) afterAbort() { + if t.servicing != nil { + logger.Warningf("Servicing work in %s may still be running; it is never terminated by the supervisor.", trustedInstallerServiceName) + } + if t.rootMsiexec || t.ranMsiexec { + t.msi.afterMSIAbort(t.opts) + } +} + +// rootExited reports whether the root process has exited. +func (t *jobTree) rootExited() bool { + ev, err := windows.WaitForSingleObject(t.proc, 0) + return err == nil && ev == windows.WAIT_OBJECT_0 +} + +// runSupervised executes a process within a Windows Job Object with forward progress, hard +// timeout and UI watchdog monitoring. If the job cannot be set up, for example on systems +// without nested job support, it logs a warning and supervises only the root process, with +// only the hard timeout enforced as described at fallbackOptions. +func runSupervised(c *exec.Cmd, opts Options, out io.Writer) error { + devNull, err := setupStdio(c, out) + if err != nil { + return err + } + defer devNull.Close() + + job, err := createJob() + if err != nil { + logger.Warningf("%v; falling back to supervising the installer's root process only.", err) + } + // On abort paths closing the handle kills every remaining process through + // KILL_ON_JOB_CLOSE. On a normal exit the limit is cleared first so that surviving + // descendants such as tray apps and updaters keep running. + releaseOnClose := false + defer func() { + if job == 0 { + return + } + if releaseOnClose { + if err := setJobLimitFlags(job, 0); err != nil { + logger.Warningf("Failed to clear job limits before close; surviving installer descendants will be terminated: %v", err) + } + } + windows.CloseHandle(job) + }() + + if c.SysProcAttr == nil { + c.SysProcAttr = &syscall.SysProcAttr{} + } + c.SysProcAttr.CreationFlags |= windows.CREATE_SUSPENDED + + if err := c.Start(); err != nil { + return err + } + pid := uint32(c.Process.Pid) + + waitErr := make(chan error, 1) + startFailed := func(step string, err error) error { + _ = c.Process.Kill() + go func() { waitErr <- c.Wait() }() + select { + case <-waitErr: + case <-time.After(killWaitTimeout): + } + return fmt.Errorf("%s for PID %d: %w", step, pid, err) + } + + const access = windows.PROCESS_SET_QUOTA | windows.PROCESS_TERMINATE | windows.SYNCHRONIZE | windows.PROCESS_QUERY_LIMITED_INFORMATION + proc, err := windows.OpenProcess(access, false, pid) + if err != nil { + return startFailed("opening process handle", err) + } + defer windows.CloseHandle(proc) + + if job != 0 { + if err := windows.AssignProcessToJobObject(job, proc); err != nil { + logger.Warningf("Assigning PID %d to the Job Object failed; falling back to supervising the installer's root process only: %v", pid, err) + windows.CloseHandle(job) + job = 0 + } + } + if job == 0 && opts.Mode != ModeOff { + logger.Warningf("No Job Object for PID %d: disabling the inactivity and UI watchdogs; the hard timeout %v still applies to the root process.", pid, opts.HardTimeout) + opts = fallbackOptions(opts) + } + if err := resumeProcessThreads(pid); err != nil { + if job != 0 { + _ = windows.TerminateJobObject(job, 1) + } + return startFailed("resuming suspended process", err) + } + + go func() { waitErr <- c.Wait() }() + + start := time.Now() + tree := &jobTree{ + c: c, + opts: opts, + pid: pid, + job: job, + proc: proc, + start: start, + checkedImages: make(map[uint32]bool), + rootMsiexec: isMsiexecCommand(c.Path), + // Any wrapper may drive msiexec, so service-side MSI work is always accounted while + // the _MSIExecute mutex exists. + msi: newMSIMonitor(opts, start), + } + if isServicingCommand(c.Path) { + tree.servicing = newTrustedInstallerMonitor(opts, start) + } + ticker := time.NewTicker(opts.pollInterval) + defer ticker.Stop() + err = supervise(opts, waitErr, tree, start, ticker.C) + // Every watchdog abort wraps ErrTerminated; only then must the remaining processes die with + // the job. + releaseOnClose = !errors.Is(err, ErrTerminated) + return err +} diff --git a/supervisor/supervisor_windows_test.go b/supervisor/supervisor_windows_test.go new file mode 100644 index 0000000..4f897fd --- /dev/null +++ b/supervisor/supervisor_windows_test.go @@ -0,0 +1,267 @@ +//go:build windows + +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "errors" + "fmt" + "io" + "os" + "os/exec" + "path/filepath" + "strconv" + "strings" + "syscall" + "testing" + "time" + "unsafe" + + "golang.org/x/sys/windows" +) + +var procMessageBoxW = user32.NewProc("MessageBoxW") + +func init() { + platformHelpers["messagebox"] = func(args []string) { + // Shows a blocking MessageBox, a standard #32770 dialog, like an unexpected prompt. + text, _ := windows.UTF16PtrFromString("Setup needs your confirmation to continue.") + caption, _ := windows.UTF16PtrFromString("Setup") + procMessageBoxW.Call(0, uintptr(unsafe.Pointer(text)), uintptr(unsafe.Pointer(caption)), 0) + } + platformHelpers["spawn_detached_child"] = func(args []string) { + // Starts a stalled child with no inherited pipes, records its PID in args[0] and exits. + child := spawnHelper("stall") + if err := child.Start(); err != nil { + fmt.Fprintf(os.Stderr, "Failed to start child: %v\n", err) + os.Exit(1) + } + if err := os.WriteFile(args[0], []byte(strconv.Itoa(child.Process.Pid)), 0644); err != nil { + os.Exit(1) + } + } +} + +// waitForExit waits until pid has exited or timeout elapses, and reports whether it exited. +func waitForExit(pid int, timeout time.Duration) bool { + h, err := windows.OpenProcess(windows.SYNCHRONIZE, false, uint32(pid)) + if err != nil { + // The process no longer exists. + return true + } + defer windows.CloseHandle(h) + ev, err := windows.WaitForSingleObject(h, uint32(timeout/time.Millisecond)) + return err == nil && ev == windows.WAIT_OBJECT_0 +} + +// readPID reads a PID written by a helper. +func readPID(t *testing.T, path string) int { + t.Helper() + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("Failed to read PID file %s: %v", path, err) + } + pid, err := strconv.Atoi(strings.TrimSpace(string(data))) + if err != nil { + t.Fatalf("Failed to parse PID %q: %v", data, err) + } + return pid +} + +// killPID terminates pid, ignoring errors. +func killPID(pid int) { + if h, err := windows.OpenProcess(windows.PROCESS_TERMINATE, false, uint32(pid)); err == nil { + _ = windows.TerminateProcess(h, 1) + windows.CloseHandle(h) + } +} + +// TestWindows_GrandchildContainedAndKilled verifies that a grandchild started immediately is +// inside the job and is terminated with it, leaving no orphans. +func TestWindows_GrandchildContainedAndKilled(t *testing.T) { + pidFile := filepath.Join(t.TempDir(), "child.pid") + cmd := helperCommand(t, "spawn_orphan_child", pidFile) + opts := Options{InactivityTimeout: 2 * time.Second, pollInterval: 50 * time.Millisecond} + + if err := Run(cmd, opts, io.Discard); !errors.Is(err, ErrInactivityTimeout) { + t.Fatalf("Run got error %v, want ErrInactivityTimeout", err) + } + child := readPID(t, pidFile) + if !waitForExit(child, 5*time.Second) { + killPID(child) + t.Errorf("Grandchild %d survived the job termination", child) + } +} + +// TestWindows_MessageBoxAborts verifies that a persistent MessageBox with no progress aborts +// shortly after the grace period in unattended mode. +func TestWindows_MessageBoxAborts(t *testing.T) { + cmd := helperCommand(t, "messagebox") + grace := 2 * time.Second + opts := Options{ + InactivityTimeout: -1, + UIGracePeriod: grace, + progressWindow: time.Second, + pollInterval: 100 * time.Millisecond, + Unattended: true, + } + start := time.Now() + err := Run(cmd, opts, io.Discard) + elapsed := time.Since(start) + if !errors.Is(err, ErrInteractiveUIDetected) { + t.Fatalf("Run got error %v, want ErrInteractiveUIDetected", err) + } + if !strings.Contains(err.Error(), `Class="#32770"`) { + t.Errorf("Run error %q lacks the dialog class", err) + } + if elapsed < grace || elapsed > grace+15*time.Second { + t.Errorf("Run took %v, want between %v and %v", elapsed, grace, grace+15*time.Second) + } +} + +// TestWindows_SuccessPathDetachedChildSurvives verifies that descendants survive a successful +// install because the job's kill-on-close limit is cleared. +func TestWindows_SuccessPathDetachedChildSurvives(t *testing.T) { + pidFile := filepath.Join(t.TempDir(), "child.pid") + cmd := helperCommand(t, "spawn_detached_child", pidFile) + if err := Run(cmd, Options{InactivityTimeout: -1}, io.Discard); err != nil { + t.Fatalf("Run returned %v, want nil", err) + } + child := readPID(t, pidFile) + defer killPID(child) + if waitForExit(child, time.Second) { + t.Errorf("Detached child %d was terminated after a successful install", child) + } +} + +// TestWindows_CallerCreationFlagsPreserved verifies that CREATE_SUSPENDED is merged into +// caller-provided creation flags and that the process still runs to completion. +func TestWindows_CallerCreationFlagsPreserved(t *testing.T) { + cmd := helperCommand(t, "exit_code_42") + cmd.SysProcAttr = &syscall.SysProcAttr{CreationFlags: windows.CREATE_NEW_PROCESS_GROUP} + err := Run(cmd, Options{InactivityTimeout: -1}, io.Discard) + var exitErr *exec.ExitError + if !errors.As(err, &exitErr) || exitErr.ExitCode() != 42 { + t.Fatalf("Run returned %v, want exit code 42", err) + } + want := uint32(windows.CREATE_NEW_PROCESS_GROUP | windows.CREATE_SUSPENDED) + if cmd.SysProcAttr.CreationFlags&want != want { + t.Errorf("CreationFlags got %#x, want both %#x", cmd.SysProcAttr.CreationFlags, want) + } +} + +// TestWindows_EnumContextRegistry verifies the integer-handle registry used by EnumWindows. +func TestWindows_EnumContextRegistry(t *testing.T) { + ctx := &windowEnumContext{} + h := registerEnumContext(ctx) + if h == 0 || lookupEnumContext(h) != ctx { + t.Fatalf("lookupEnumContext(%d) did not return the registered context", h) + } + unregisterEnumContext(h) + if lookupEnumContext(h) != nil { + t.Error("lookupEnumContext returned a context after unregister") + } + if enumWindowsProc(0, h) != 0 { + t.Error("enumWindowsProc did not stop enumeration for an unknown handle") + } +} + +// TestWindows_IsCandidateWindow verifies the dialog candidate rules. +func TestWindows_IsCandidateWindow(t *testing.T) { + for _, tc := range []struct { + class string + hasOwner bool + exStyle uint32 + want bool + }{ + {"#32770", false, 0, true}, + {"WixBundleWindow", false, wsExDlgModalFrame, false}, + {"WixBundleWindow", true, 0, false}, + {"CustomModal", true, wsExDlgModalFrame, true}, + } { + if got := isCandidateWindow(tc.class, tc.hasOwner, tc.exStyle); got != tc.want { + t.Errorf("isCandidateWindow(%q, %v, %#x) got %v, want %v", tc.class, tc.hasOwner, tc.exStyle, got, tc.want) + } + } +} + +// TestWindows_MutexExists verifies mutex detection using a test-local mutex name. +func TestWindows_MutexExists(t *testing.T) { + mutexName := fmt.Sprintf(`Local\googet_supervisor_test_%d`, os.Getpid()) + if mutexExists(mutexName) { + t.Fatal("mutexExists returned true before the mutex exists") + } + name, _ := windows.UTF16PtrFromString(mutexName) + h, err := windows.CreateMutex(nil, false, name) + if err != nil { + t.Fatalf("CreateMutex failed: %v", err) + } + if !mutexExists(mutexName) { + t.Error("mutexExists returned false while the mutex exists") + } + windows.CloseHandle(h) + if mutexExists(mutexName) { + t.Error("mutexExists returned true after the mutex was closed") + } +} + +// TestWindows_JobTreeRootExited verifies root exit detection through the process handle. +func TestWindows_JobTreeRootExited(t *testing.T) { + cmd := helperCommand(t, "stall") + if err := cmd.Start(); err != nil { + t.Fatalf("Start failed: %v", err) + } + proc, err := windows.OpenProcess(windows.SYNCHRONIZE, false, uint32(cmd.Process.Pid)) + if err != nil { + _ = cmd.Process.Kill() + t.Fatalf("OpenProcess failed: %v", err) + } + defer windows.CloseHandle(proc) + tree := &jobTree{c: cmd, proc: proc} + if tree.rootExited() { + t.Error("rootExited returned true for a running process") + } + _ = cmd.Process.Kill() + _ = cmd.Wait() + if !tree.rootExited() { + t.Error("rootExited returned false after the process exited") + } +} + +// TestWindows_ResumeProcessThreadsUnknownPID verifies that resuming a nonexistent process fails. +func TestWindows_ResumeProcessThreadsUnknownPID(t *testing.T) { + if err := resumeProcessThreads(0xFFFFFFF0); err == nil { + t.Error("resumeProcessThreads succeeded for a nonexistent process") + } +} + +// TestWindows_NoJobFallbackEnforcesOnlyHardTimeout verifies that when the Job Object cannot be +// created, a stalled root process is not killed by the inactivity watchdog but still by the +// hard timeout. +func TestWindows_NoJobFallbackEnforcesOnlyHardTimeout(t *testing.T) { + old := createJob + createJob = func() (windows.Handle, error) { return 0, errors.New("injected job creation failure") } + defer func() { createJob = old }() + + cmd := helperCommand(t, "stall") + opts := Options{ + InactivityTimeout: 200 * time.Millisecond, + HardTimeout: time.Second, + pollInterval: 20 * time.Millisecond, + } + if err := Run(cmd, opts, io.Discard); !errors.Is(err, ErrHardTimeout) { + t.Fatalf("Run without a Job Object got error %v, want ErrHardTimeout", err) + } +} diff --git a/supervisor/ui_test.go b/supervisor/ui_test.go new file mode 100644 index 0000000..7d42ff0 --- /dev/null +++ b/supervisor/ui_test.go @@ -0,0 +1,279 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package supervisor + +import ( + "errors" + "fmt" + "strings" + "testing" + "time" +) + +// uiTick is the synthetic poll interval used by the UI tests. +const uiTick = 2 * time.Second + +// dialog returns a candidate #32770 window. +func dialog(hwnd uintptr, title string) windowInfo { + return windowInfo{PID: 100, HWND: hwnd, Title: title, ClassName: "#32770", ExePath: "setup.exe"} +} + +// uiOptions returns unattended options with only UI detection enabled at the default grace +// period. +func uiOptions() Options { + return testOptions(Options{InactivityTimeout: -1, HardTimeout: -1, Unattended: true}) +} + +// TestUI_TransientDialogsSuccessive verifies that successive distinct dialogs, each shorter than +// the grace period, do not abort even though their total duration exceeds it. +func TestUI_TransientDialogsSuccessive(t *testing.T) { + tree := &fakeTree{winsAt: func(elapsed time.Duration, _ int) []windowInfo { + switch { + case elapsed < 20*time.Second: + return []windowInfo{dialog(101, "Transient Dialog 1")} + case elapsed >= 22*time.Second && elapsed < 42*time.Second: + return []windowInfo{dialog(102, "Transient Dialog 2")} + case elapsed >= 44*time.Second && elapsed < 64*time.Second: + return []windowInfo{dialog(103, "Transient Dialog 3")} + } + return nil + }} + if err := runFake(uiOptions(), tree, uiTick, 3*time.Minute); err != nil { + t.Fatalf("supervise got %v for successive transient dialogs, want nil", err) + } + if tree.windowReports < 27 { + t.Errorf("Detector reported windows on %d ticks, want at least 27; the dialogs were not exercised", tree.windowReports) + } +} + +// TestUI_TransientDialogFlapping verifies that a dialog that closes and reopens on every other +// poll does not abort. +func TestUI_TransientDialogFlapping(t *testing.T) { + tree := &fakeTree{winsAt: func(_ time.Duration, call int) []windowInfo { + if call%2 == 1 { + return []windowInfo{dialog(999, "Flapping Window")} + } + return nil + }} + if err := runFake(uiOptions(), tree, uiTick, 3*time.Minute); err != nil { + t.Fatalf("supervise got %v for a flapping dialog, want nil", err) + } + if tree.windowReports < 40 { + t.Errorf("Detector reported windows on %d ticks, want at least 40; the dialog was not exercised", tree.windowReports) + } +} + +// TestUI_ChangingHWND verifies that a dialog replaced by another before the grace period does +// not abort. +func TestUI_ChangingHWND(t *testing.T) { + tree := &fakeTree{winsAt: func(elapsed time.Duration, _ int) []windowInfo { + switch { + case elapsed < 20*time.Second: + return []windowInfo{dialog(1001, "Step 1")} + case elapsed < 40*time.Second: + return []windowInfo{dialog(1002, "Step 2")} + } + return nil + }} + if err := runFake(uiOptions(), tree, uiTick, 2*time.Minute); err != nil { + t.Fatalf("supervise got %v for successive changing dialogs, want nil", err) + } + if tree.windowReports < 19 { + t.Errorf("Detector reported windows on %d ticks, want at least 19", tree.windowReports) + } +} + +// TestUI_PersistentDialogTimingAndDiagnostics verifies that a persistent dialog aborts exactly +// one grace period after it was first seen, with full diagnostics. +func TestUI_PersistentDialogTimingAndDiagnostics(t *testing.T) { + tree := &fakeTree{winsAt: func(time.Duration, int) []windowInfo { + return []windowInfo{{PID: 100, HWND: 7777, Title: "Setup Error: Out of Disk Space", ClassName: "#32770", ExePath: `C:\Windows\Temp\setup.exe`}} + }} + err := runFake(uiOptions(), tree, uiTick, 5*time.Minute) + if !errors.Is(err, ErrInteractiveUIDetected) { + t.Fatalf("supervise got %v, want ErrInteractiveUIDetected", err) + } + if want := uiTick + defaultUIGracePeriod; tree.terminatedAt != want { + t.Errorf("Terminated at %v, want %v", tree.terminatedAt, want) + } + for _, want := range []string{`Title="Setup Error: Out of Disk Space"`, `Class="#32770"`, `Image="C:\\Windows\\Temp\\setup.exe"`, "HWND=0x1e61"} { + if !strings.Contains(err.Error(), want) { + t.Errorf("Error %q lacks %s", err, want) + } + } +} + +// TestUI_MultiWindowPersistentDialog verifies that a persistent dialog aborts regardless of its +// position in the enumerated window list. +func TestUI_MultiWindowPersistentDialog(t *testing.T) { + setupWin := windowInfo{PID: 100, HWND: 1001, Title: "Main Setup Window", ClassName: "SetupClass", ExePath: "setup.exe"} + status := windowInfo{PID: 100, HWND: 1002, Title: "Worker Status", ClassName: "ProgressClass", ExePath: "setup.exe"} + stuck := dialog(1003, "Modal Retry Dialog") + for _, tc := range []struct { + name string + winsAt func(time.Duration, int) []windowInfo + wantTitle string + }{ + {"ZOrderAlternation", func(_ time.Duration, call int) []windowInfo { + if call%2 == 1 { + return []windowInfo{setupWin, stuck} + } + return []windowInfo{stuck, setupWin} + }, ""}, + {"PersistentSecondaryWithChurningPrimary", func(_ time.Duration, call int) []windowInfo { + churn := windowInfo{PID: 100, HWND: uintptr(5000 + call), Title: fmt.Sprintf("Progress %d", call), ClassName: "ProgressClass"} + return []windowInfo{churn, stuck} + }, "Modal Retry Dialog"}, + {"ThreeWindowPermutations", func(_ time.Duration, call int) []windowInfo { + switch call % 3 { + case 0: + return []windowInfo{setupWin, status, stuck} + case 1: + return []windowInfo{status, stuck, setupWin} + default: + return []windowInfo{stuck, setupWin, status} + } + }, ""}, + } { + t.Run(tc.name, func(t *testing.T) { + tree := &fakeTree{winsAt: tc.winsAt} + err := runFake(uiOptions(), tree, uiTick, 5*time.Minute) + if !errors.Is(err, ErrInteractiveUIDetected) { + t.Fatalf("supervise got %v, want ErrInteractiveUIDetected", err) + } + if want := uiTick + defaultUIGracePeriod; tree.terminatedAt != want { + t.Errorf("Terminated at %v, want %v", tree.terminatedAt, want) + } + if tc.wantTitle != "" && !strings.Contains(err.Error(), fmt.Sprintf("Title=%q", tc.wantTitle)) { + t.Errorf("Error %q does not identify %q", err, tc.wantTitle) + } + }) + } +} + +// TestUISnifferState_Scenarios exercises the uiSnifferState state machine with synthetic time +// sequences and validates the diagnostics. +func TestUISnifferState_Scenarios(t *testing.T) { + baseTime := time.Date(2026, 9, 23, 12, 0, 0, 0, time.UTC) + grace := 30 * time.Second + + t.Run("EmptyWindowListResets", func(t *testing.T) { + var s uiSnifferState + w := []windowInfo{{HWND: 1, Title: "A", ClassName: "#32770", PID: 10, ExePath: "a.exe"}} + if abort, _ := s.check(w, grace, baseTime, noProgress); abort { + t.Fatal("Unexpected abort on first tick") + } + if abort, _ := s.check(nil, grace, baseTime.Add(10*time.Second), noProgress); abort { + t.Fatal("Unexpected abort on empty window list") + } + if len(s.tracked) != 0 { + t.Fatalf("State not reset after empty list: %+v", s) + } + if abort, _ := s.check(w, grace, baseTime.Add(20*time.Second), noProgress); abort { + t.Fatal("Unexpected abort on reappearance") + } + if abort, _ := s.check(w, grace, baseTime.Add(45*time.Second), noProgress); abort { + t.Fatal("Unexpected abort: elapsed 25s < grace 30s") + } + abort, details := s.check(w, grace, baseTime.Add(55*time.Second), noProgress) + if !abort { + t.Fatal("Expected abort when persisted for 35s >= 30s") + } + if !strings.Contains(details, `Title="A"`) || !strings.Contains(details, `Class="#32770"`) { + t.Errorf("Diagnostic details missing expected fields: %q", details) + } + }) + + t.Run("HWNDSwitchResetsTimer", func(t *testing.T) { + var s uiSnifferState + w1 := []windowInfo{{HWND: 1, Title: "A", ClassName: "#32770", PID: 10, ExePath: "a.exe"}} + w2 := []windowInfo{{HWND: 2, Title: "B", ClassName: "#32770", PID: 10, ExePath: "a.exe"}} + s.check(w1, grace, baseTime, noProgress) + if abort, _ := s.check(w1, grace, baseTime.Add(25*time.Second), noProgress); abort { + t.Fatal("Window 1 aborted prematurely at 25s") + } + if abort, _ := s.check(w2, grace, baseTime.Add(26*time.Second), noProgress); abort { + t.Fatal("Unexpected abort immediately on HWND switch") + } + if abort, _ := s.check(w2, grace, baseTime.Add(50*time.Second), noProgress); abort { + t.Fatal("Window 2 aborted prematurely at 24s after switch") + } + abort, details := s.check(w2, grace, baseTime.Add(57*time.Second), noProgress) + if !abort { + t.Fatal("Expected abort for Window 2 persisting > grace") + } + if !strings.Contains(details, `Title="B"`) { + t.Errorf("Diagnostic details expected Window B, got: %q", details) + } + }) + + t.Run("ZOrderAlternationDoesNotResetTimer", func(t *testing.T) { + var s uiSnifferState + w1 := windowInfo{HWND: 101, Title: "Installer Progress", ClassName: "SetupClass", PID: 1000, ExePath: "setup.exe"} + w2 := windowInfo{HWND: 102, Title: "Fatal Error Prompt", ClassName: "#32770", PID: 1000, ExePath: "setup.exe"} + for i := 0; i < 15; i++ { + list := []windowInfo{w1, w2} + if i%2 == 1 { + list = []windowInfo{w2, w1} + } + if abort, details := s.check(list, grace, baseTime.Add(time.Duration(i*2)*time.Second), noProgress); abort { + t.Fatalf("Unexpected abort at tick %d: %s", i, details) + } + } + if abort, _ := s.check([]windowInfo{w2, w1}, grace, baseTime.Add(30*time.Second), noProgress); !abort { + t.Fatal("Expected abort at t=30s for alternating Z-order windows") + } + }) + + t.Run("PersistentSecondaryWithChurningPrimary", func(t *testing.T) { + var s uiSnifferState + modal := windowInfo{HWND: 999, Title: "Stuck Dialog", ClassName: "#32770", PID: 500, ExePath: "setup.exe"} + for i := 0; i < 15; i++ { + churn := windowInfo{HWND: uintptr(2000 + i), Title: fmt.Sprintf("Step %d", i), ClassName: "ProgressClass", PID: 500, ExePath: "setup.exe"} + list := []windowInfo{churn, modal} + if i%2 == 1 { + list = []windowInfo{modal, churn} + } + if abort, details := s.check(list, grace, baseTime.Add(time.Duration(i*2)*time.Second), noProgress); abort { + t.Fatalf("Unexpected abort at tick %d: %s", i, details) + } + if len(s.tracked) > 2 { + t.Fatalf("Tracked %d windows at tick %d, want at most the 2 visible ones", len(s.tracked), i) + } + } + final := windowInfo{HWND: 3000, Title: "Final Step", ClassName: "ProgressClass", PID: 500, ExePath: "setup.exe"} + abort, details := s.check([]windowInfo{final, modal}, grace, baseTime.Add(30*time.Second), noProgress) + if !abort { + t.Fatal("Expected abort at t=30s for the persistent modal dialog") + } + if !strings.Contains(details, `Title="Stuck Dialog"`) { + t.Errorf("Expected details to identify Stuck Dialog, got: %s", details) + } + }) + + t.Run("ProgressBeforeFirstSeenDoesNotReset", func(t *testing.T) { + var s uiSnifferState + w := []windowInfo{{HWND: 7, Title: "Prompt", ClassName: "#32770"}} + first := baseTime.Add(10 * time.Second) + // Progress ended at 8s, before the window appeared at 10s. + progressedSince := func(since time.Time) bool { return since.Before(baseTime.Add(8 * time.Second)) } + s.check(w, grace, first, progressedSince) + if abort, _ := s.check(w, grace, first.Add(20*time.Second), progressedSince); abort { + t.Fatal("Unexpected abort 20s after the window appeared") + } + if abort, _ := s.check(w, grace, first.Add(grace), progressedSince); !abort { + t.Fatal("Expected abort one grace period after the window appeared") + } + }) +} diff --git a/system/supervisor_options_test.go b/system/supervisor_options_test.go new file mode 100644 index 0000000..50705a7 --- /dev/null +++ b/system/supervisor_options_test.go @@ -0,0 +1,87 @@ +/* +Copyright 2026 Google Inc. All Rights Reserved. +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + http://www.apache.org/licenses/LICENSE-2.0 +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package system + +import ( + "errors" + "os" + "path/filepath" + "runtime" + "testing" + "time" + + "github.com/google/googet/v2/goolib" + "github.com/google/googet/v2/supervisor" +) + +// TestSupervisorOptions verifies that per-command overrides are mapped and that invalid +// overrides fall back to zero options. +func TestSupervisorOptions(t *testing.T) { + got := supervisorOptions(goolib.ExecFile{Path: "install.sh", Timeout: "3h", InactivityTimeout: "0"}) + if got.HardTimeout != 3*time.Hour || got.InactivityTimeout != -1 { + t.Errorf("supervisorOptions(valid) = %+v, want HardTimeout 3h and InactivityTimeout -1", got) + } + got = supervisorOptions(goolib.ExecFile{Path: "install.sh", Timeout: "forever"}) + if got.HardTimeout != 0 || got.InactivityTimeout != 0 || len(got.LogFiles) != 0 { + t.Errorf("supervisorOptions(invalid) = %+v, want zero options", got) + } +} + +// writeSleepScript writes an executable shell script that sleeps for a long time. +func writeSleepScript(t *testing.T, dir, name string) { + t.Helper() + if runtime.GOOS != "linux" { + t.Skip("Shell script execution is only exercised on Linux") + } + if _, err := os.Stat("/bin/sh"); err != nil { + t.Skipf("/bin/sh is unavailable: %v", err) + } + if err := os.WriteFile(filepath.Join(dir, name), []byte("#!/bin/sh\nsleep 30\n"), 0755); err != nil { + t.Fatalf("os.WriteFile: %v", err) + } +} + +// TestVerifyAppliesTimeoutOverride verifies that Verify passes the per-command Timeout to +// the supervisor instead of running with the process-wide defaults. +func TestVerifyAppliesTimeoutOverride(t *testing.T) { + dir := t.TempDir() + writeSleepScript(t, dir, "verify.sh") + ps := &goolib.PkgSpec{Name: "foo", Verify: goolib.ExecFile{Path: "verify.sh", Timeout: "500ms"}} + + start := time.Now() + err := Verify(dir, ps) + if !errors.Is(err, supervisor.ErrHardTimeout) { + t.Fatalf("Verify() = %v, want an error wrapping ErrHardTimeout", err) + } + if elapsed := time.Since(start); elapsed > 20*time.Second { + t.Errorf("Verify() took %v, want the 500ms override to terminate it early", elapsed) + } +} + +// TestInstallAppliesTimeoutOverride verifies that Install passes the per-command Timeout to +// the supervisor instead of running with the process-wide defaults. +func TestInstallAppliesTimeoutOverride(t *testing.T) { + dir := t.TempDir() + writeSleepScript(t, dir, "install.sh") + ps := &goolib.PkgSpec{Name: "foo", Install: goolib.ExecFile{Path: "install.sh", Timeout: "500ms"}} + + start := time.Now() + err := Install(dir, ps) + if !errors.Is(err, supervisor.ErrHardTimeout) { + t.Fatalf("Install() = %v, want an error wrapping ErrHardTimeout", err) + } + if elapsed := time.Since(start); elapsed > 20*time.Second { + t.Errorf("Install() took %v, want the 500ms override to terminate it early", elapsed) + } +} diff --git a/system/system.go b/system/system.go index 01fe527..da080a6 100644 --- a/system/system.go +++ b/system/system.go @@ -24,6 +24,7 @@ import ( "github.com/google/googet/v2/goolib" "github.com/google/googet/v2/oswrap" + "github.com/google/googet/v2/supervisor" "github.com/google/logger" ) @@ -44,7 +45,19 @@ func Verify(dir string, ps *goolib.PkgSpec) error { logger.Error(err) } }() - return goolib.Exec(filepath.Join(dir, v.Path), v.Args, v.ExitCodes, out) + return goolib.ExecWithOptions(filepath.Join(dir, v.Path), v.Args, v.ExitCodes, supervisorOptions(v), out) +} + +// supervisorOptions returns the supervisor options declared by ef. Invalid overrides are +// logged and ignored so that the process-wide defaults apply; VerifyPkgSpec already +// rejects them when a package is built. +func supervisorOptions(ef goolib.ExecFile) supervisor.Options { + opts, err := ef.SupervisorOptions() + if err != nil { + logger.Warningf("Ignoring invalid timeout overrides for %q: %v", ef.Path, err) + return supervisor.Options{} + } + return opts } // isLockFileStale checks if the lock file is older than maxAge. @@ -74,6 +87,11 @@ func readPID(lockFile string) (int, error) { } // killProcess kills the process with the given PID. +// +// On Windows only the GooGet PID needs to be killed: installers run by the +// supervisor package are assigned to a Job Object with +// JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE whose only handle is owned by that GooGet +// process, so its termination also terminates any installer process tree. func killProcess(pid int) error { p, err := os.FindProcess(pid) if err != nil { diff --git a/system/system_darwin.go b/system/system_darwin.go index 442cea0..ccf5f70 100644 --- a/system/system_darwin.go +++ b/system/system_darwin.go @@ -46,8 +46,8 @@ func Install(dir string, ps *goolib.PkgSpec) error { logger.Error(err) } }() - if err := goolib.Exec(filepath.Join(dir, in.Path), in.Args, in.ExitCodes, out); err != nil { - return fmt.Errorf("error running install: %v", err) + if err := goolib.ExecWithOptions(filepath.Join(dir, in.Path), in.Args, in.ExitCodes, supervisorOptions(in), out); err != nil { + return fmt.Errorf("error running install: %w", err) } return nil } @@ -70,7 +70,7 @@ func Uninstall(dir string, ps *client.PackageState) error { logger.Error(err) } }() - return goolib.Exec(filepath.Join(dir, un.Path), un.Args, un.ExitCodes, out) + return goolib.ExecWithOptions(filepath.Join(dir, un.Path), un.Args, un.ExitCodes, supervisorOptions(un), out) } // InstallableArchs returns a slice of archs supported by this machine. diff --git a/system/system_linux.go b/system/system_linux.go index 18fd239..867a731 100644 --- a/system/system_linux.go +++ b/system/system_linux.go @@ -46,8 +46,8 @@ func Install(dir string, ps *goolib.PkgSpec) error { logger.Error(err) } }() - if err := goolib.Exec(filepath.Join(dir, in.Path), in.Args, in.ExitCodes, out); err != nil { - return fmt.Errorf("error running install: %v", err) + if err := goolib.ExecWithOptions(filepath.Join(dir, in.Path), in.Args, in.ExitCodes, supervisorOptions(in), out); err != nil { + return fmt.Errorf("error running install: %w", err) } return nil } @@ -70,7 +70,7 @@ func Uninstall(dir string, ps *client.PackageState) error { logger.Error(err) } }() - return goolib.Exec(filepath.Join(dir, un.Path), un.Args, un.ExitCodes, out) + return goolib.ExecWithOptions(filepath.Join(dir, un.Path), un.Args, un.ExitCodes, supervisorOptions(un), out) } // InstallableArchs returns a slice of archs supported by this machine. diff --git a/system/system_windows.go b/system/system_windows.go index ce28313..f65456f 100644 --- a/system/system_windows.go +++ b/system/system_windows.go @@ -268,25 +268,27 @@ func Install(dir string, ps *goolib.PkgSpec) error { s := filepath.Join(dir, in.Path) msiLog := filepath.Join(dir, "msi_install.log") ec := append(msiSuccessCodes, in.ExitCodes...) + opts := supervisorOptions(in) switch filepath.Ext(s) { case ".msi": - args := append([]string{"/i", s, "/qn", "/norestart", "/log", msiLog}, in.Args...) - err = goolib.Run(exec.Command("msiexec", args...), ec, out) + args := append([]string{"/i", s, "/qn", "/norestart", "/l*v", msiLog}, in.Args...) + err = goolib.RunWithOptions(exec.Command("msiexec", args...), ec, opts, out) case ".msp": - args := append([]string{"/update", s, "/qn", "/norestart", "/log", msiLog}, in.Args...) - err = goolib.Run(exec.Command("msiexec", args...), ec, out) + args := append([]string{"/update", s, "/qn", "/norestart", "/l*v", msiLog}, in.Args...) + err = goolib.RunWithOptions(exec.Command("msiexec", args...), ec, opts, out) case ".msu": + // supervisor.Run applies its servicing policy to wusa, including CBS.log progress. args := append([]string{s, "/quiet", "/norestart"}, in.Args...) - err = goolib.Run(exec.Command("wusa", args...), ec, out) + err = goolib.RunWithOptions(exec.Command("wusa", args...), ec, opts, out) case ".exe": - err = goolib.Run(exec.Command(s, in.Args...), ec, out) + err = goolib.RunWithOptions(exec.Command(s, in.Args...), ec, opts, out) case ".msix", ".msixbundle": // Add-AppxProvisionedPackage will install for all users. installCmd := fmt.Sprintf("Add-AppxProvisionedPackage -online -PackagePath %v -SkipLicense", s) args := append([]string{installCmd}, in.Args...) - err = goolib.Run(exec.Command("powershell", args...), ec, out) + err = goolib.RunWithOptions(exec.Command("powershell", args...), ec, opts, out) default: - err = goolib.Exec(s, in.Args, in.ExitCodes, out) + err = goolib.ExecWithOptions(s, in.Args, in.ExitCodes, opts, out) } if err != nil { return err @@ -356,23 +358,25 @@ func Uninstall(dir string, state *client.PackageState) error { filePath = filepath.Join(dir, un.Path) } ec := append(msiSuccessCodes, un.ExitCodes...) + opts := supervisorOptions(un) switch filepath.Ext(filePath) { case ".msi": msiLog := filepath.Join(dir, "msi_uninstall.log") - args := append([]string{"/x", filePath, "/qn", "/norestart", "/log", msiLog}, un.Args...) - err = goolib.Run(exec.Command("msiexec", args...), ec, out) + args := append([]string{"/x", filePath, "/qn", "/norestart", "/l*v", msiLog}, un.Args...) + err = goolib.RunWithOptions(exec.Command("msiexec", args...), ec, opts, out) case ".msu": + // supervisor.Run applies its servicing policy to wusa, including CBS.log progress. args := append([]string{filePath, "/uninstall", "/quiet", "/norestart"}, un.Args...) - err = goolib.Run(exec.Command("wusa", args...), ec, out) + err = goolib.RunWithOptions(exec.Command("wusa", args...), ec, opts, out) case ".exe": - err = goolib.Run(exec.Command(filePath, un.Args...), ec, out) + err = goolib.RunWithOptions(exec.Command(filePath, un.Args...), ec, opts, out) case ".msix", ".msixbundle": s := strings.Split(filepath.Base(filePath), "_")[0] removeCmd := fmt.Sprintf(`Get-AppxProvisionedPackage -online | Where {$_.DisplayName -match "%v*"} | Remove-AppProvisionedPackage -online -AllUsers`, s) args := append([]string{removeCmd}, un.Args...) - err = goolib.Run(exec.Command("powershell", args...), ec, out) + err = goolib.RunWithOptions(exec.Command("powershell", args...), ec, opts, out) default: - err = goolib.Exec(filepath.Join(dir, un.Path), un.Args, un.ExitCodes, out) + err = goolib.ExecWithOptions(filepath.Join(dir, un.Path), un.Args, un.ExitCodes, opts, out) } if err != nil { return err From 986d77dae0c47f4cb417a5fadd17a08a60a38a8e Mon Sep 17 00:00:00 2001 From: John-Michael Mulesa Date: Mon, 28 Sep 2026 16:28:15 +0000 Subject: [PATCH 2/4] Raise default install hard cap from 60m to 4h Long-running installers (SQL Server, Visual Studio, Office) routinely exceed 60m while making steady progress. The inactivity (5m) and UI (30s) watchdogs still catch true hangs quickly; the hard cap is only a backstop for spin loops or log spam that masquerade as progress. --- googet.goospec | 2 +- supervisor/supervise_test.go | 2 +- supervisor/supervisor.go | 2 +- supervisor/supervisor_test.go | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) diff --git a/googet.goospec b/googet.goospec index 1b890ec..fd1a4ea 100644 --- a/googet.goospec +++ b/googet.goospec @@ -15,7 +15,7 @@ "path": "install.ps1" }, "releaseNotes": [ - "3.4.0 - Feat: Supervise installers in a Job Object; terminate hung (5m no progress), interactive (30s modal dialog without progress in unattended mode), or over-long (60m) installers. Configure via SupervisorMode/InactivityTimeout/InstallTimeout/UIGracePeriod/UIDetection in googet.conf or per package via ExecFile timeout/inactivityTimeout. Installers longer than 60m must set a timeout override.", + "3.4.0 - Feat: Supervise installers in a Job Object; terminate hung (5m no progress), interactive (30s modal dialog without progress in unattended mode), or over-long (4h) installers. Configure via SupervisorMode/InactivityTimeout/InstallTimeout/UIGracePeriod/UIDetection in googet.conf or per package via ExecFile timeout/inactivityTimeout. Installers that legitimately run longer than 4h must set a timeout override.", "3.4.0 - Change: Killing googet (e.g. an outer agent timeout) now also terminates the in-flight installer process tree (Job Object KILL_ON_JOB_CLOSE).", "3.4.0 - Change: .msu (wusa) installs have no hard timeout unless the package sets one; CBS.log growth counts as progress.", "3.4.0 - Feat: Download stall detection with automatic HTTP Range/GCS resume; header timeout for HTTP requests.", diff --git a/supervisor/supervise_test.go b/supervisor/supervise_test.go index ee5c046..2ff8dd6 100644 --- a/supervisor/supervise_test.go +++ b/supervisor/supervise_test.go @@ -180,7 +180,7 @@ func TestSupervise_ProgressPreventsInactivity(t *testing.T) { // TestSupervise_HardTimeout verifies that the hard cap terminates a tree making progress. func TestSupervise_HardTimeout(t *testing.T) { tree := &fakeTree{cpuAt: steadyCPU} - err := runFake(testOptions(Options{}), tree, 2*time.Second, 2*time.Hour) + err := runFake(testOptions(Options{}), tree, 2*time.Second, 5*time.Hour) if !errors.Is(err, ErrHardTimeout) { t.Fatalf("supervise got %v, want ErrHardTimeout", err) } diff --git a/supervisor/supervisor.go b/supervisor/supervisor.go index 701e9a8..cb09dd5 100644 --- a/supervisor/supervisor.go +++ b/supervisor/supervisor.go @@ -178,7 +178,7 @@ type Options struct { const ( defaultMode = ModeEnforce defaultInactivityTimeout = 5 * time.Minute - defaultHardTimeout = 60 * time.Minute + defaultHardTimeout = 4 * time.Hour defaultUIGracePeriod = 30 * time.Second defaultPollInterval = 2 * time.Second defaultProgressWindow = 30 * time.Second diff --git a/supervisor/supervisor_test.go b/supervisor/supervisor_test.go index b269076..b280e33 100644 --- a/supervisor/supervisor_test.go +++ b/supervisor/supervisor_test.go @@ -742,7 +742,7 @@ func TestServicingOptions(t *testing.T) { wantInactivity time.Duration }{ {"NothingConfiguredLiftsCap", Options{LogFiles: []string{"pkg.msu.log"}}, false, windir, -1, []string{"pkg.msu.log", cbs}, 0}, - {"AdminExplicit60mRespected", Options{}, true, windir, 0, []string{cbs}, 0}, + {"AdminExplicitHardTimeoutRespected", Options{}, true, windir, 0, []string{cbs}, 0}, {"PackageTimeoutRespected", Options{HardTimeout: 2 * time.Hour, InactivityTimeout: 20 * time.Minute}, false, windir, 2 * time.Hour, []string{cbs}, 20 * time.Minute}, {"PackageTimeoutWinsOverAdmin", Options{HardTimeout: 2 * time.Hour}, true, windir, 2 * time.Hour, []string{cbs}, 0}, {"EmptyWindirFallsBack", Options{}, false, "", -1, []string{filepath.Join(`C:\Windows`, "Logs", "CBS", "CBS.log")}, 0}, From 85212326369499fb27eace1241f52df52e81dbab Mon Sep 17 00:00:00 2001 From: John-Michael Mulesa Date: Thu, 1 Oct 2026 03:05:43 +0000 Subject: [PATCH 3/4] Harden progress output, stall cancellation and log discovery - Run watchdog termination messages and install rollback/commit logging under the new progress.Interrupt so the active spinner cannot redraw over error lines that logger writes to stderr. - StallReader now returns ErrDownloadStalled only after cancel has completed, fixing a race that made TestStallReader_ZeroByteReadsDoNotResetTimer flaky under -race. - enrichOptions now discovers InnoSetup /LOG=path and fully quoted "/log:path" arguments, and no longer treats a following switch such as /quiet as the log path. The log-only corner-case test is replaced with assertions. - Port upstream's TestPackageHTTP table, which the merge dropped, as TestPackageHTTP_ResumeFromDisk. --- client/stall.go | 28 ++++++++--- download/download_test.go | 98 +++++++++++++++++++++++++++++++++++++++ goolib/goolib.go | 32 +++++++++---- goolib/goolib_test.go | 87 +++++++++++++++++----------------- install/install.go | 7 ++- progress/progress.go | 12 +++++ progress/progress_test.go | 39 ++++++++++++++++ supervisor/msi.go | 5 +- supervisor/progress.go | 8 +++- 9 files changed, 251 insertions(+), 65 deletions(-) diff --git a/client/stall.go b/client/stall.go index 6c593e4..82061f1 100644 --- a/client/stall.go +++ b/client/stall.go @@ -69,6 +69,9 @@ type StallReader struct { timer *time.Timer stalled bool closed bool + // canceled is closed once onStall has invoked cancel, so that a Read + // returning ErrDownloadStalled guarantees the context is already done. + canceled chan struct{} } // NewStallReader returns a StallReader that wraps r with an idle read watchdog. @@ -81,9 +84,10 @@ func NewStallReader(r io.Reader, timeout time.Duration, cancel context.CancelFun timeout = DefaultStallTimeout } s := &StallReader{ - r: r, - timeout: timeout, - cancel: cancel, + r: r, + timeout: timeout, + cancel: cancel, + canceled: make(chan struct{}), } s.timer = time.AfterFunc(timeout, s.onStall) return s @@ -96,8 +100,11 @@ func (s *StallReader) onStall() { s.mu.Unlock() return } + // stalled is set before cancel so that a read unblocked by the + // cancellation reports ErrDownloadStalled rather than context.Canceled. s.stalled = true s.mu.Unlock() + defer close(s.canceled) // Cancel context outside the mutex to prevent deadlocks. if s.cancel != nil { @@ -105,6 +112,13 @@ func (s *StallReader) onStall() { } } +// stallErr waits for onStall to finish canceling and returns +// ErrDownloadStalled. It must be called without holding s.mu. +func (s *StallReader) stallErr() error { + <-s.canceled + return ErrDownloadStalled +} + // isStalled reports whether the idle timer has fired. func (s *StallReader) isStalled() bool { s.mu.Lock() @@ -118,7 +132,7 @@ func (s *StallReader) Read(p []byte) (int, error) { s.mu.Lock() if s.stalled { s.mu.Unlock() - return 0, ErrDownloadStalled + return 0, s.stallErr() } if s.closed { s.mu.Unlock() @@ -129,15 +143,15 @@ func (s *StallReader) Read(p []byte) (int, error) { n, err := s.r.Read(p) s.mu.Lock() - defer s.mu.Unlock() - if s.stalled { + s.mu.Unlock() // If the stream finished cleanly with io.EOF, prefer EOF over stall. if err == io.EOF { return n, io.EOF } - return 0, ErrDownloadStalled + return 0, s.stallErr() } + defer s.mu.Unlock() if n > 0 { // Forward progress made: reset the idle timer. diff --git a/download/download_test.go b/download/download_test.go index ac99bb0..43444c4 100644 --- a/download/download_test.go +++ b/download/download_test.go @@ -26,6 +26,7 @@ import ( "net/http/httptest" "os" "path/filepath" + "reflect" "sync" "sync/atomic" "syscall" @@ -273,6 +274,103 @@ func TestPackageHTTP_Normal(t *testing.T) { } } +// TestPackageHTTP_ResumeFromDisk ports upstream's TestPackageHTTP table. It +// checks which GET requests are sent for a partial or complete file left on +// disk by an earlier googet run. +func TestPackageHTTP_ResumeFromDisk(t *testing.T) { + t.Parallel() + payload, chksum := testPayload(1000) + for _, tc := range []struct { + desc string + existing []byte // Contents written to dst before the download. + honorRange bool + wantGETs []string + }{ + { + // An empty destination sends no Range header. + desc: "fresh download", + honorRange: true, + wantGETs: []string{"GET "}, + }, + { + desc: "resumed download", + existing: payload[:400], + honorRange: true, + wantGETs: []string{"GET bytes=400-"}, + }, + { + // A 200 in reply to a Range request restarts from byte zero + // within the same attempt. + desc: "server ignores range", + existing: payload[:400], + wantGETs: []string{"GET bytes=400-"}, + }, + { + desc: "already downloaded", + existing: payload, + honorRange: true, + wantGETs: nil, + }, + } { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + var ( + mu sync.Mutex + gets []string + ) + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodHead { + writeHead(w, len(payload)) + return + } + mu.Lock() + gets = append(gets, r.Method+" "+r.Header.Get("Range")) + mu.Unlock() + start := 0 + if tc.honorRange { + start = rangeStart(t, r.Header.Get("Range")) + } + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload)-start)) + if start > 0 { + w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, len(payload)-1, len(payload))) + w.WriteHeader(http.StatusPartialContent) + } else { + w.WriteHeader(http.StatusOK) + } + w.Write(payload[start:]) + })) + defer ts.Close() + + downloader, err := client.NewDownloader("") + if err != nil { + t.Fatalf("client.NewDownloader failed: %v", err) + } + dst := filepath.Join(t.TempDir(), "pkg.goo") + if tc.existing != nil { + if err := os.WriteFile(dst, tc.existing, 0644); err != nil { + t.Fatalf("os.WriteFile failed: %v", err) + } + } + if err := packageHTTP(context.Background(), ts.URL+"/pkg.goo", dst, chksum, downloader); err != nil { + t.Fatalf("packageHTTP() = %v, want nil", err) + } + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) + } + if !bytes.Equal(got, payload) { + t.Errorf("downloaded %d bytes, want %d bytes matching the payload", len(got), len(payload)) + } + mu.Lock() + defer mu.Unlock() + if !reflect.DeepEqual(gets, tc.wantGETs) { + t.Errorf("GET requests = %q, want %q", gets, tc.wantGETs) + } + }) + } +} + func TestPackageHTTP_SlowTrickle(t *testing.T) { // Verify that slow trickle downloads complete without being aborted by stall timer. diff --git a/goolib/goolib.go b/goolib/goolib.go index 425ba84..c476762 100644 --- a/goolib/goolib.go +++ b/goolib/goolib.go @@ -106,6 +106,20 @@ func isLogFlag(s string) bool { return true } +// looksLikeSwitch reports whether s is a command line switch such as /quiet +// or -norestart rather than a path. A leading slash followed by another path +// separator, as in /tmp/install.log, is treated as a path. +func looksLikeSwitch(s string) bool { + s = strings.Trim(s, `"'`) + switch { + case strings.HasPrefix(s, "-"): + return true + case strings.HasPrefix(s, "/"): + return !strings.ContainsAny(s[1:], `/\`) + } + return false +} + // enrichOptions inspects command line arguments and the output writer to auto-discover // log files whose growth indicates forward progress. Unattended mode is configured // process-wide via supervisor.Configure and is not inferred here. @@ -127,21 +141,21 @@ func enrichOptions(c *exec.Cmd, opts supervisor.Options, w io.Writer) supervisor addLog(f.Name()) } - // 2. Inspect command arguments for log flags (space-separated or colon-delimited). + // 2. Inspect command arguments for log flags (space-separated or delimited). if c != nil { for i := 0; i < len(c.Args); i++ { - arg := c.Args[i] + // Quoting may wrap the whole token, as in "/log:C:\install.log". + arg := strings.Trim(c.Args[i], `"'`) - // Check for colon-delimited log flags (e.g. /log:, /l*v:, -log:, -l:). - if parts := strings.SplitN(arg, ":", 2); len(parts) == 2 { - if isLogFlag(parts[0]) { - addLog(strings.Trim(strings.TrimSpace(parts[1]), `"'`)) - continue - } + // Check for delimited log flags, e.g. /log:, /l*v:, + // -log:, -l:, or InnoSetup's /LOG=. + if j := strings.IndexAny(arg, ":="); j > 0 && isLogFlag(arg[:j]) { + addLog(strings.Trim(strings.TrimSpace(arg[j+1:]), `"'`)) + continue } // Check for space-separated log flags (e.g. /log , /l*v , -log , -l ). - if isLogFlag(arg) && i+1 < len(c.Args) && !isLogFlag(c.Args[i+1]) { + if isLogFlag(arg) && i+1 < len(c.Args) && !looksLikeSwitch(c.Args[i+1]) { addLog(strings.Trim(strings.TrimSpace(c.Args[i+1]), `"'`)) i++ } diff --git a/goolib/goolib_test.go b/goolib/goolib_test.go index b5a25d4..62e7475 100644 --- a/goolib/goolib_test.go +++ b/goolib/goolib_test.go @@ -501,53 +501,52 @@ func TestEnrichOptionsNilAndWriter(t *testing.T) { } } -// TestAdversarialCornerCases documents empirical edge-case behavior and limitations of enrichOptions. -func TestAdversarialCornerCases(t *testing.T) { - // 1. InnoSetup-style /LOG= is currently unhandled by enrichOptions. - // When installers use /LOG=path, the colon splitter does not split on '='. - // Therefore, the log path is not auto-discovered. - cmdInno := &exec.Cmd{Args: []string{"setup.exe", `/VERYSILENT`, `/LOG=C:\Windows\Logs\inno.log`}} - optsInno := enrichOptions(cmdInno, supervisor.Options{}, nil) - if len(optsInno.LogFiles) != 0 { - t.Logf("Notice: /LOG= was unexpectedly parsed as %v", optsInno.LogFiles) - } else { - t.Logf("Empirically confirmed: InnoSetup /LOG= syntax is unhandled by enrichOptions (LogFiles is empty)") - } - - // 2. Fully-quoted argument "/log:path" has leading quote on the switch. - // strings.SplitN produces parts[0] == `"/log`, which fails isLogFlag. - cmdQuotedSwitch := &exec.Cmd{Args: []string{"setup.exe", `"/log:C:\install.log"`}} - optsQuotedSwitch := enrichOptions(cmdQuotedSwitch, supervisor.Options{}, nil) - if len(optsQuotedSwitch.LogFiles) != 0 { - t.Logf("Notice: Fully-quoted switch was parsed as %v", optsQuotedSwitch.LogFiles) - } else { - t.Logf("Empirically confirmed: Fully-quoted switch \"/log:...\" is not extracted (LogFiles is empty)") +// TestEnrichOptionsArgEdgeCases covers installer argument forms that log +// discovery must handle or deliberately ignore. +func TestEnrichOptionsArgEdgeCases(t *testing.T) { + tests := []struct { + name string + args []string + want []string + }{ + {"InnoSetup equals", []string{"setup.exe", "/VERYSILENT", `/LOG=C:\Windows\Logs\inno.log`}, []string{`C:\Windows\Logs\inno.log`}}, + {"InnoSetup quoted value", []string{"setup.exe", `/LOG="C:\inno.log"`}, []string{`C:\inno.log`}}, + {"fully quoted switch", []string{"setup.exe", `"/log:C:\install.log"`}, []string{`C:\install.log`}}, + {"forward slash path", []string{"setup.exe", "/log:C:/temp/install.log"}, []string{"C:/temp/install.log"}}, + {"next arg is slash switch", []string{"msiexec.exe", "/i", "pkg.msi", "/log", "/quiet"}, nil}, + {"next arg is dash switch", []string{"setup.exe", "-log", "-norestart"}, nil}, + {"next arg is unix path", []string{"installer", "--log", "/tmp/install.log"}, []string{"/tmp/install.log"}}, + {"flag is last arg", []string{"msiexec.exe", "/i", "pkg.msi", "/l*v"}, nil}, + {"non-log switch with colon", []string{"setup.exe", "/dir:C:\\app"}, nil}, } - - // 3. Space-separated log flag followed by another command switch (/quiet). - // Because /quiet is not a known log flag, isLogFlag("/quiet") returns false. - // This causes enrichOptions to treat "/quiet" as the log file path. - cmdNextSwitch := &exec.Cmd{Args: []string{"msiexec.exe", "/i", "pkg.msi", "/log", "/quiet"}} - optsNextSwitch := enrichOptions(cmdNextSwitch, supervisor.Options{}, nil) - if len(optsNextSwitch.LogFiles) == 1 && optsNextSwitch.LogFiles[0] == "/quiet" { - t.Logf("Empirically confirmed: /log followed by non-log switch treats switch (/quiet) as log path") + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := enrichOptions(&exec.Cmd{Args: tc.args}, supervisor.Options{}, nil).LogFiles + if !reflect.DeepEqual(got, tc.want) { + t.Errorf("enrichOptions(%q).LogFiles = %q, want %q", tc.args, got, tc.want) + } + }) } +} - // 4. Bare argument "noconfirm" without dash or slash prefix in os.Args. - // strings.TrimLeft(arg, "-/") strips leading prefixes, but if none exist, clean == "noconfirm". - // This inadvertently sets opts.Unattended = true. - oldArgs := os.Args - defer func() { os.Args = oldArgs }() - os.Args = []string{"googet", "install", "noconfirm"} - optsBare := enrichOptions(&exec.Cmd{Args: []string{"cmd.exe"}}, supervisor.Options{}, nil) - if optsBare.Unattended { - t.Logf("Empirically confirmed: Package named 'noconfirm' in os.Args triggers opts.Unattended = true") +func TestLooksLikeSwitch(t *testing.T) { + tests := []struct { + in string + want bool + }{ + {"/quiet", true}, + {"-norestart", true}, + {"--log", true}, + {`"/qn"`, true}, + {"/tmp/install.log", false}, + {`/c\install.log`, false}, + {`C:\install.log`, false}, + {"install.log", false}, + {"", false}, } - - // 5. Forward slash paths in Windows commands: /log:C:/temp/install.log. - cmdFwd := &exec.Cmd{Args: []string{"setup.exe", `/log:C:/temp/install.log`}} - optsFwd := enrichOptions(cmdFwd, supervisor.Options{}, nil) - if len(optsFwd.LogFiles) != 1 || optsFwd.LogFiles[0] != "C:/temp/install.log" { - t.Errorf("enrichOptions() with forward-slash path = %v, want [C:/temp/install.log]", optsFwd.LogFiles) + for _, tc := range tests { + if got := looksLikeSwitch(tc.in); got != tc.want { + t.Errorf("looksLikeSwitch(%q) = %v, want %v", tc.in, got, tc.want) + } } } diff --git a/install/install.go b/install/install.go index 006e02f..88cda8c 100644 --- a/install/install.go +++ b/install/install.go @@ -591,7 +591,10 @@ func installPkgInner(ops installOps, pkg string, ps *goolib.PkgSpec, dbOnly, for // trigger rollback. success := false - defer func() { + // The spinner started by installPkg is still running here, and rollback + // and commit log errors to stderr. Run them under progress.Interrupt so + // the spinner cannot redraw over those lines. + defer progress.Interrupt(func() { if !success { txn.rollback() logger.Errorf("install logs preserved at %s", dir) @@ -601,7 +604,7 @@ func installPkgInner(ops installOps, pkg string, ps *goolib.PkgSpec, dbOnly, for if err := oswrap.RemoveAll(dir); err != nil { logger.Error(err) } - }() + }) for src, dst := range ps.Files { dst = resolveDst(dst) diff --git a/progress/progress.go b/progress/progress.go index dca7dac..543e8b0 100644 --- a/progress/progress.go +++ b/progress/progress.go @@ -162,6 +162,18 @@ func Printf(format string, a ...any) { fmt.Fprintf(stdout, format, a...) } +// Interrupt clears any active bar or spinner line and then runs fn while +// holding the console lock, so text that fn writes directly to os.Stdout or +// os.Stderr, such as logger.Errorf output, starts at column zero and cannot be +// overwritten by a concurrent redraw. The bar or spinner redraws itself on its +// next update. fn must not call any function in this package. +func Interrupt(fn func()) { + mu.Lock() + defer mu.Unlock() + clearLocked() + fn() +} + // redrawLocked overwrites the current console line with s, blanking any // trailing characters left over from a longer previous line. It writes // nothing if s is already what is on the line. diff --git a/progress/progress_test.go b/progress/progress_test.go index e835c67..5a71f01 100644 --- a/progress/progress_test.go +++ b/progress/progress_test.go @@ -586,6 +586,45 @@ func TestPrintfClearsBar(t *testing.T) { } } +func TestInterruptClearsBar(t *testing.T) { + outBuf, _, _ := setup(t, true) + + NewBar("Title", 100, 0) + n := currentLastLen() + if n == 0 { + t.Fatal("lastLen after NewBar = 0, want > 0") + } + outBuf.Reset() + + ran := false + Interrupt(func() { + ran = true + // mu is held here, so lastLine can be read directly. + if lastLine != "" { + t.Errorf("lastLine inside Interrupt = %q, want empty", lastLine) + } + if got, want := outBuf.String(), "\r"+strings.Repeat(" ", n)+"\r"; got != want { + t.Errorf("out inside Interrupt = %q, want %q", got, want) + } + }) + if !ran { + t.Error("Interrupt did not run fn") + } +} + +func TestInterruptNoActiveLine(t *testing.T) { + outBuf, _, _ := setup(t, true) + + ran := false + Interrupt(func() { ran = true }) + if !ran { + t.Error("Interrupt did not run fn") + } + if got := outBuf.String(); got != "" { + t.Errorf("out after Interrupt with no active line = %q, want empty", got) + } +} + func TestPrintfClearsSpinner(t *testing.T) { outBuf, stdoutBuf, _ := setup(t, true) diff --git a/supervisor/msi.go b/supervisor/msi.go index c65b577..d17c32c 100644 --- a/supervisor/msi.go +++ b/supervisor/msi.go @@ -16,6 +16,7 @@ package supervisor import ( "time" + "github.com/google/googet/v2/progress" "github.com/google/logger" ) @@ -96,7 +97,9 @@ func waitForMSITransaction(opts Options, env msiWaitEnv) { idle := now.Sub(latest(obs.lastProgress, start, obs.gateOpened)) if !terminated && opts.InactivityTimeout > 0 && idle >= opts.InactivityTimeout && len(obs.terminable) > 0 { terminated = true - logger.Errorf("Windows Installer service tree observed inactive for %v; terminating processes %v created for this install. The service process is not terminated.", idle.Round(time.Second), obs.terminable) + progress.Interrupt(func() { + logger.Errorf("Windows Installer service tree observed inactive for %v; terminating processes %v created for this install. The service process is not terminated.", idle.Round(time.Second), obs.terminable) + }) env.terminate(obs.terminable) } if !now.Before(deadline) { diff --git a/supervisor/progress.go b/supervisor/progress.go index 7f1ecea..cc108ae 100644 --- a/supervisor/progress.go +++ b/supervisor/progress.go @@ -382,7 +382,9 @@ func (w *watchdog) shouldTerminate(d *abortDecision) bool { } return false } - logger.Errorf("Terminating installer process tree: %v %s", d.reason, d.details) + // logger.Errorf always writes to stderr, so keep an active spinner from + // overwriting the most important line of a hung install. + progress.Interrupt(func() { logger.Errorf("Terminating installer process tree: %v %s", d.reason, d.details) }) return true } @@ -488,7 +490,9 @@ func waitAfterTerminate(waitErr <-chan error, d *abortDecision) error { case <-waitErr: return fmt.Errorf("%w: %s", reason, details) case <-time.After(killWaitTimeout): - logger.Errorf("Installer process tree terminated but Wait did not return within %v; its stdout/stderr pipes are likely held by a process outside the supervised tree.", killWaitTimeout) + progress.Interrupt(func() { + logger.Errorf("Installer process tree terminated but Wait did not return within %v; its stdout/stderr pipes are likely held by a process outside the supervised tree.", killWaitTimeout) + }) return fmt.Errorf("%w: %s; wait abandoned after %v because stdout/stderr pipes are held by processes outside the supervised tree", reason, details, killWaitTimeout) } } From 3d1cd2e04340d194b14b23265c9c94cd314b043d Mon Sep 17 00:00:00 2001 From: John-Michael Mulesa Date: Thu, 1 Oct 2026 15:13:26 +0000 Subject: [PATCH 4/4] Fix UI detection in session 0 and Windows-only test failures Windows CI ran these code paths for the first time and found: - TestWindows_MessageBoxAborts hung until the job timeout. On a non-interactive window station, such as session 0 where googet runs as SYSTEM, a blocking #32770 dialog reports IsWindowVisible false, so the visibility filter hid every prompt and UI detection never fired. The filter now applies only on an interactive window station, detected via GetProcessWindowStation and the WSF_VISIBLE flag. The test also gets a hard timeout so a detection regression fails instead of hanging. - Three install tests asserted Unix permission bits (0640, 0604, 0750) that Windows does not model. They now compare against the mode the OS applied after Chmod. --- install/install_test.go | 38 ++++++++++++-------- supervisor/supervisor_windows.go | 51 ++++++++++++++++++++++----- supervisor/supervisor_windows_test.go | 10 +++--- 3 files changed, 72 insertions(+), 27 deletions(-) diff --git a/install/install_test.go b/install/install_test.go index 9721343..e5e3af9 100644 --- a/install/install_test.go +++ b/install/install_test.go @@ -1471,6 +1471,16 @@ func assertAbsent(t *testing.T, paths ...string) { } } +// statMode returns the permission bits of path, failing the test on error. +func statMode(t *testing.T, path string) os.FileMode { + t.Helper() + fi, err := os.Stat(path) + if err != nil { + t.Fatalf("Stat(%q): %v", path, err) + } + return fi.Mode().Perm() +} + // assertNoBackups fails the test if any backup file remains under root. func assertNoBackups(t *testing.T, root string) { t.Helper() @@ -1779,6 +1789,9 @@ func TestInstallPkg_BackupFallbackDeletedFileRestored(t *testing.T) { if err := os.Chmod(target, 0640); err != nil { t.Fatalf("Chmod: %v", err) } + // Windows only models the read-only bit, so compare against the + // mode the OS actually applied rather than the requested one. + wantMode := statMode(t, target) ops := failingOps() if tc.succeed { @@ -1816,12 +1829,8 @@ func TestInstallPkg_BackupFallbackDeletedFileRestored(t *testing.T) { } assertContents(t, map[string][]byte{target: content}) assertNoBackups(t, dstDir) - fi, err := os.Stat(target) - if err != nil { - t.Fatalf("Stat(%q): %v", target, err) - } - if got := fi.Mode().Perm(); got != 0640 { - t.Errorf("Mode of restored %q = %v, want %v", target, got, os.FileMode(0640)) + if got := statMode(t, target); got != wantMode { + t.Errorf("Mode of restored %q = %v, want %v", target, got, wantMode) } }) } @@ -1869,6 +1878,8 @@ func TestInstallPkg_RemovedEmptyDirRecreatedOnRollback(t *testing.T) { if err := os.Chmod(emptyDir, 0750); err != nil { t.Fatalf("Chmod: %v", err) } + // Windows ignores most permission bits; compare against what was applied. + wantDirMode := statMode(t, emptyDir) other := filepath.Join(dstDir, "other.txt") orig := map[string][]byte{other: []byte("original other")} writeFiles(t, orig) @@ -1893,8 +1904,8 @@ func TestInstallPkg_RemovedEmptyDirRecreatedOnRollback(t *testing.T) { if !fi.IsDir() { t.Fatalf("%q is not a directory after rollback: mode %v", emptyDir, fi.Mode()) } - if got := fi.Mode().Perm(); got != 0750 { - t.Errorf("Mode of recreated %q = %v, want %v", emptyDir, got, os.FileMode(0750)) + if got := fi.Mode().Perm(); got != wantDirMode { + t.Errorf("Mode of recreated %q = %v, want %v", emptyDir, got, wantDirMode) } entries, err := os.ReadDir(emptyDir) if err != nil { @@ -1964,6 +1975,9 @@ func TestCopyToBackup(t *testing.T) { if err := os.Chmod(path, 0604); err != nil { t.Fatalf("Chmod: %v", err) } + // Windows only models the read-only bit, so compare against the mode the + // OS actually applied rather than the requested one. + wantMode := statMode(t, path) seen := make(map[string]bool) for i := 0; i < 3; i++ { bak, err := copyToBackup(path) @@ -1978,12 +1992,8 @@ func TestCopyToBackup(t *testing.T) { } seen[bak] = true assertContents(t, map[string][]byte{path: content, bak: content}) - fi, err := os.Stat(bak) - if err != nil { - t.Fatalf("Stat(%q): %v", bak, err) - } - if got := fi.Mode().Perm(); got != 0604 { - t.Errorf("Mode of %q = %v, want %v", bak, got, os.FileMode(0604)) + if got := statMode(t, bak); got != wantMode { + t.Errorf("Mode of %q = %v, want %v", bak, got, wantMode) } } } diff --git a/supervisor/supervisor_windows.go b/supervisor/supervisor_windows.go index 63cc525..1c7aa2a 100644 --- a/supervisor/supervisor_windows.go +++ b/supervisor/supervisor_windows.go @@ -31,12 +31,14 @@ import ( ) var ( - user32 = windows.NewLazySystemDLL("user32.dll") - procEnumWindows = user32.NewProc("EnumWindows") - procGetWindowTextW = user32.NewProc("GetWindowTextW") - procGetWindow = user32.NewProc("GetWindow") - procGetWindowLongW = user32.NewProc("GetWindowLongW") - enumWindowsCallback = windows.NewCallback(enumWindowsProc) + user32 = windows.NewLazySystemDLL("user32.dll") + procEnumWindows = user32.NewProc("EnumWindows") + procGetWindowTextW = user32.NewProc("GetWindowTextW") + procGetWindow = user32.NewProc("GetWindow") + procGetWindowLongW = user32.NewProc("GetWindowLongW") + procGetProcessWindowStation = user32.NewProc("GetProcessWindowStation") + procGetUserObjectInformation = user32.NewProc("GetUserObjectInformationW") + enumWindowsCallback = windows.NewCallback(enumWindowsProc) ) const ( @@ -46,8 +48,36 @@ const ( gwOwner = 4 // wsExDlgModalFrame is the extended window style of windows with a modal dialog frame. wsExDlgModalFrame = 0x00000001 + // uoiFlags is the GetUserObjectInformation index that returns USEROBJECTFLAGS. + uoiFlags = 1 + // wsfVisible is the USEROBJECTFLAGS flag of a window station that has visible display + // surfaces, that is, an interactive window station. + wsfVisible = 0x0001 ) +// userObjectFlags mirrors Win32 USEROBJECTFLAGS from winuser.h. +type userObjectFlags struct { + Inherit int32 + Reserved int32 + Flags uint32 +} + +// interactiveStation caches whether this process runs on an interactive window station. +var interactiveStation = sync.OnceValue(func() bool { + ws, _, _ := procGetProcessWindowStation.Call() + if ws == 0 { + return true + } + var f userObjectFlags + var needed uint32 + r, _, _ := procGetUserObjectInformation.Call(ws, uoiFlags, uintptr(unsafe.Pointer(&f)), unsafe.Sizeof(f), uintptr(unsafe.Pointer(&needed))) + if r == 0 { + // Assume interactive so that the visibility filter stays in place. + return true + } + return f.Flags&wsfVisible != 0 +}) + // gwlExStyle is the GetWindowLong index of the extended window style. It is a variable because // the negative constant cannot be converted to uintptr directly. var gwlExStyle int32 = -20 @@ -195,7 +225,10 @@ func enumWindowsProc(hwnd uintptr, lParam uintptr) uintptr { if tid == 0 || err != nil || !ctx.jobPIDs[pid] { return 1 } - if !windows.IsWindowVisible(h) { + // On a non-interactive window station, such as session 0 where googet runs as + // SYSTEM, IsWindowVisible is false even for a dialog that is blocking on input, so + // the visibility filter only applies on interactive stations. + if interactiveStation() && !windows.IsWindowVisible(h) { return 1 } @@ -259,8 +292,8 @@ func isSession0() bool { return sessionID == 0 } -// detectWindowsWin32 enumerates visible top-level windows owned by pids and returns candidate -// interactive prompts. +// detectWindowsWin32 enumerates top-level windows owned by pids and returns candidate +// interactive prompts. Invisible windows are skipped only on an interactive window station. func detectWindowsWin32(pids []uint32) ([]windowInfo, error) { ctx := &windowEnumContext{jobPIDs: make(map[uint32]bool, len(pids))} for _, p := range pids { diff --git a/supervisor/supervisor_windows_test.go b/supervisor/supervisor_windows_test.go index 4f897fd..e6f5b5b 100644 --- a/supervisor/supervisor_windows_test.go +++ b/supervisor/supervisor_windows_test.go @@ -112,10 +112,12 @@ func TestWindows_MessageBoxAborts(t *testing.T) { grace := 2 * time.Second opts := Options{ InactivityTimeout: -1, - UIGracePeriod: grace, - progressWindow: time.Second, - pollInterval: 100 * time.Millisecond, - Unattended: true, + // HardTimeout bounds the test if the dialog is never detected. + HardTimeout: grace + 30*time.Second, + UIGracePeriod: grace, + progressWindow: time.Second, + pollInterval: 100 * time.Millisecond, + Unattended: true, } start := time.Now() err := Run(cmd, opts, io.Discard)