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..82061f1 --- /dev/null +++ b/client/stall.go @@ -0,0 +1,190 @@ +/* +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 + // 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. +// 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, + canceled: make(chan struct{}), + } + 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 + } + // 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 { + s.cancel() + } +} + +// 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() + 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, s.stallErr() + } + if s.closed { + s.mu.Unlock() + return 0, io.ErrClosedPipe + } + s.mu.Unlock() + + n, err := s.r.Read(p) + + s.mu.Lock() + 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, s.stallErr() + } + defer s.mu.Unlock() + + 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 3b422f1..b9702c6 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" @@ -36,6 +42,7 @@ import ( "github.com/google/googet/v2/oswrap" "github.com/google/googet/v2/progress" "github.com/google/logger" + "google.golang.org/api/googleapi" ) // Package downloads a package from the given url, @@ -47,118 +54,575 @@ 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, the offset the stream actually starts at (which is either offset or +// 0 when the source ignored the resume request), and the total size of the +// object in bytes (or -1 if unknown). It returns an error wrapping +// errResumeRejected when the source cannot serve offset. +type opener func(ctx context.Context, offset int64) (io.ReadCloser, int64, 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} } - // restart discards any partial download so the response body is written - // from the beginning of the file. - restart := func() error { - if err := f.Truncate(0); err != nil { - return err - } - if _, err := f.Seek(0, 0); err != nil { - return 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 + } + 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 } - hash.Reset() - size = 0 - return nil } - 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)) - } else if err := restart(); err != nil { - return err + 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 { + 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++ } - resp, err := downloader.HTTPClient.Do(req) + 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 err + return nil, 0, err } - defer resp.Body.Close() - if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent { - return fmt.Errorf("downloading %s: unexpected status %s", url, resp.Status) + 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} } - if resp.StatusCode == http.StatusOK && size > 0 { - // The server ignored the Range header and is sending the whole file; - // appending it to the partial download would corrupt the file and - // overstate the progress total. - logger.Infof("server ignored range request for %s, restarting download", url) - if err := restart(); err != nil { + 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) + } + case start == 0: + logger.Infof("server did not resume download of %s, restarting from start", name) + if err := discardPartial(f); err != nil { return err } + h.Reset() + default: + return fmt.Errorf("%w: source resumed at offset %d, want %d or 0", errResumeRejected, start, size) + } + 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, total, err := open(attemptCtx, size) + if err != nil { + return size, size, err } - // The total is what is already on disk plus what the server will send. - total := int64(-1) - if resp.ContentLength >= 0 { - total = size + resp.ContentLength + sr := client.NewStallReader(body, stallTimeout, cancel) + defer sr.Close() + + if err := alignToStream(f, h, name, size, start); err != nil { + return size, size, err } - bar := progress.NewBar(fmt.Sprintf("Downloading %s", filepath.Base(dst)), total, size) - // Continue hashing the file as we download it. - n, err := io.Copy(io.MultiWriter(hash, f, bar), resp.Body) + bar := progress.NewBar(fmt.Sprintf("Downloading %s", filepath.Base(f.Name())), total, start) + n, err := io.Copy(io.MultiWriter(fileWriter{w: f}, h, bar), sr) if err != nil { bar.Abort() - return fmt.Errorf("downloading %s: %v", url, err) + } else { + bar.Finish() } - bar.Finish() - // 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) + 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) + } + 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) + } + } +} - r, err := client.Bucket(bucket).Object(object).NewReader(ctx) +// 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, 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, 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, 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, 0, err + } + total := int64(-1) + if resp.ContentLength >= 0 { + total = offset + resp.ContentLength + } + return resp.Body, offset, total, 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, resp.ContentLength, nil + case resp.StatusCode == http.StatusRequestedRangeNotSatisfiable && resume: + resp.Body.Close() + return nil, 0, 0, fmt.Errorf("%w: server returned %s for offset %d", errResumeRejected, resp.Status, offset) + default: + resp.Body.Close() + return nil, 0, 0, &statusError{code: resp.StatusCode, status: resp.Status} + } + } +} + +// 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, int64, error)) opener { + return func(ctx context.Context, offset int64) (io.ReadCloser, int64, int64, error) { + r, total, 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, 0, fmt.Errorf("%w: %v", errResumeRejected, err) + } + return nil, 0, 0, err + } + return r, offset, total, 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, int64, error) { + r, err := obj.NewRangeReader(ctx, offset, -1) + if err != nil { + return nil, 0, err + } + return r, r.Attrs.Size, 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, r.Attrs.Size, 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 @@ -186,39 +650,6 @@ func Latest(ctx context.Context, name, dir string, rm client.RepoMap, archs []st return FromRepo(ctx, rs, repo, dir, downloader) } -// download copies r to dst, verifying the SHA256 checksum, and renders a -// progress bar when enabled. -func download(r io.Reader, size int64, 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 - } - }() - - bar := progress.NewBar(fmt.Sprintf("Downloading %s", filepath.Base(dst)), size, 0) - hash := sha256.New() - tw := io.MultiWriter(f, hash, bar) - - b, err := io.Copy(tw, r) - if err != nil { - bar.Abort() - return err - } - bar.Finish() - - 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 381f13b..43444c4 100644 --- a/download/download_test.go +++ b/download/download_test.go @@ -18,101 +18,268 @@ import ( "bytes" "compress/gzip" "context" + "crypto/sha256" + "errors" + "fmt" "io" "net/http" "net/http/httptest" "os" - "path" "path/filepath" - "slices" - "strconv" - "strings" + "reflect" + "sync" + "sync/atomic" + "syscall" "testing" + "time" "github.com/google/googet/v2/client" - "github.com/google/googet/v2/goolib" "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, io.Discard) + // Tests must not wait for real backoff delays. + sleep = func(ctx context.Context, _ time.Duration) error { return ctx.Err() } +} + +// 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 TestDownload(t *testing.T) { - r := bytes.NewReader([]byte("some content")) +func TestExtractPkg(t *testing.T) { + t.Parallel() tempDir, err := os.MkdirTemp("", "") if err != nil { t.Fatalf("error creating temp directory: %v", err) } defer oswrap.RemoveAll(tempDir) + tempFile := filepath.Join(tempDir, "test.pkg") + f, err := oswrap.Create(tempFile) + if err != nil { + t.Fatalf("error creating temp file: %v", err) + } + gw := gzip.NewWriter(f) + tw := tar.NewWriter(gw) + + name := "foo/../test" + body := "this is a test file" + if err := tw.WriteHeader(&tar.Header{ + Name: name, + Mode: 0600, + Size: int64(len(body)), + }); err != nil { + t.Fatal(err) + } + if _, err := tw.Write([]byte(body)); err != nil { + t.Fatalf("error writing file: %v", err) + } + + if err := tw.Close(); err != nil { + t.Fatalf("error closing tar: %v", err) + } + if err := gw.Close(); err != nil { + t.Fatalf("error closing gzip: %v", err) + } + if err := f.Close(); err != nil { + t.Fatalf("error closing file: %v", err) + } - chksum := goolib.Checksum(r) - if _, err := r.Seek(0, 0); err != nil { - t.Errorf("error seeking to front of reader: %v", err) + dst, err := ExtractPkg(tempFile) + if err != nil { + t.Fatalf("error running ExtractPkg: %v", err) } - tempFile := path.Join(tempDir, "test") - if err := download(r, int64(r.Len()), tempFile, chksum); err != nil { - t.Errorf("error downloading and checking checksum: %v", err) + + cts, err := os.ReadFile(filepath.Join(dst, filepath.Clean(name))) + if err != nil { + t.Fatalf("error opening test file: %v", err) } - if err := download(r, int64(r.Len()), tempFile, "notachecksum"); err == nil { - t.Error("wanted but did not recieve checksum error") + if string(cts) != body { + t.Errorf("contents of extracted file does not match expected contents: got: %q, want: %q", string(cts), body) } } -func TestDownloadUnknownSize(t *testing.T) { - content := []byte("some content") - chksum := goolib.Checksum(bytes.NewReader(content)) - tempFile := path.Join(t.TempDir(), "test") - // A size of 0 means the total is unknown; the checksum must still verify. - if err := download(bytes.NewReader(content), 0, tempFile, chksum); err != nil { - t.Errorf("error downloading with unknown size: %v", err) +func TestExtractPkgPathTraversal(t *testing.T) { + t.Parallel() + tempDir, err := os.MkdirTemp("", "") + if err != nil { + t.Fatalf("error creating temp directory: %v", err) } - got, err := os.ReadFile(tempFile) + defer oswrap.RemoveAll(tempDir) + tempFile := filepath.Join(tempDir, "test.pkg") + f, err := oswrap.Create(tempFile) if err != nil { - t.Fatalf("error reading downloaded file: %v", err) + t.Fatalf("error creating temp file: %v", err) + } + gw := gzip.NewWriter(f) + tw := tar.NewWriter(gw) + + name := "foo/../../test" + body := "this is a test file" + if err := tw.WriteHeader(&tar.Header{ + Name: name, + Mode: 0600, + Size: int64(len(body)), + }); err != nil { + t.Fatal(err) + } + if _, err := tw.Write([]byte(body)); err != nil { + t.Fatalf("error writing file: %v", err) + } + + if err := tw.Close(); err != nil { + t.Fatalf("error closing tar: %v", err) + } + if err := gw.Close(); err != nil { + t.Fatalf("error closing gzip: %v", err) } - if !bytes.Equal(got, content) { - t.Errorf("downloaded contents = %q, want %q", got, content) + if err := f.Close(); err != nil { + t.Fatalf("error closing file: %v", err) + } + + if _, err := ExtractPkg(tempFile); err == nil { + t.Fatal("error expected because of path traversal") } } -// rangeServer serves payload, advertising range support on HEAD. When -// honorRange is false it answers ranged GETs with the whole file and 200 OK, -// as some proxies do. -func rangeServer(t *testing.T, payload []byte, honorRange bool, requests *[]string) *httptest.Server { - t.Helper() - return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - *requests = append(*requests, r.Method+" "+r.Header.Get("Range")) - w.Header().Set("Accept-Ranges", "bytes") +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("Content-Length", strconv.Itoa(len(payload))) + w.Header().Set("Accept-Ranges", "bytes") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) return } - if rng := r.Header.Get("Range"); rng != "" && honorRange { - start, err := strconv.Atoi(strings.TrimSuffix(strings.TrimPrefix(rng, "bytes="), "-")) - if err != nil { - t.Errorf("bad Range header %q: %v", rng, err) - w.WriteHeader(http.StatusBadRequest) - 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.Header().Set("Content-Length", strconv.Itoa(len(payload)-start)) - w.WriteHeader(http.StatusPartialContent) - w.Write(payload[start:]) + 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", strconv.Itoa(len(payload))) + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(payload))) + w.WriteHeader(http.StatusOK) w.Write(payload) })) -} + defer ts.Close() -func TestPackageHTTP(t *testing.T) { - payload := bytes.Repeat([]byte("0123456789"), 100) - chksum := goolib.Checksum(bytes.NewReader(payload)) downloader, err := client.NewDownloader("") if err != nil { - t.Fatalf("client.NewDownloader: %v", err) + 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)) } +} + +// 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. @@ -120,10 +287,10 @@ func TestPackageHTTP(t *testing.T) { wantGETs []string }{ { - // An empty destination still asks for bytes=0-; the server may - // answer 200 or 206 and either way nothing is on disk to keep. - desc: "fresh download", - wantGETs: []string{"GET bytes=0-"}, + // An empty destination sends no Range header. + desc: "fresh download", + honorRange: true, + wantGETs: []string{"GET "}, }, { desc: "resumed download", @@ -132,139 +299,1900 @@ func TestPackageHTTP(t *testing.T) { 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, - wantGETs: nil, + desc: "already downloaded", + existing: payload, + honorRange: true, + wantGETs: nil, }, } { t.Run(tc.desc, func(t *testing.T) { - var requests []string - srv := rangeServer(t, payload, tc.honorRange, &requests) - defer srv.Close() + 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("writing existing file: %v", err) + t.Fatalf("os.WriteFile failed: %v", err) } } - if err := packageHTTP(context.Background(), srv.URL+"/pkg.goo", dst, chksum, downloader); err != nil { + 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("reading downloaded file: %v", err) + 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)) } - var gets []string - for _, r := range requests { - if strings.HasPrefix(r, "GET") { - gets = append(gets, r) - } - } - if !slices.Equal(gets, tc.wantGETs) { + mu.Lock() + defer mu.Unlock() + if !reflect.DeepEqual(gets, tc.wantGETs) { t.Errorf("GET requests = %q, want %q", gets, tc.wantGETs) } }) } } -func TestExtractPkg(t *testing.T) { - tempDir, err := os.MkdirTemp("", "") +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("error creating temp directory: %v", err) + t.Fatalf("client.NewDownloader failed: %v", err) } - defer oswrap.RemoveAll(tempDir) - tempFile := filepath.Join(tempDir, "test.pkg") - f, err := oswrap.Create(tempFile) - if err != nil { - t.Fatalf("error creating temp file: %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) } - gw := gzip.NewWriter(f) - tw := tar.NewWriter(gw) - name := "foo/../test" - body := "this is a test file" - if err := tw.WriteHeader(&tar.Header{ - Name: name, - Mode: 0600, - Size: int64(len(body)), - }); err != nil { - t.Fatal(err) + got, err := os.ReadFile(dst) + if err != nil { + t.Fatalf("os.ReadFile failed: %v", err) } - if _, err := tw.Write([]byte(body)); err != nil { - t.Fatalf("error writing file: %v", err) + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) } +} - if err := tw.Close(); err != nil { - t.Fatalf("error closing tar: %v", err) - } - if err := gw.Close(); err != nil { - t.Fatalf("error closing gzip: %v", err) +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) } - if err := f.Close(); err != nil { - t.Fatalf("error closing file: %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) } - dst, err := ExtractPkg(tempFile) - if err != nil { - t.Fatalf("error running ExtractPkg: %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]) } - cts, err := os.ReadFile(filepath.Join(dst, filepath.Clean(name))) + got, err := os.ReadFile(dst) if err != nil { - t.Fatalf("error opening test file: %v", err) + t.Fatalf("os.ReadFile failed: %v", err) } - if string(cts) != body { - t.Errorf("contents of extracted file does not match expected contents: got: %q, want: %q", string(cts), body) + if !bytes.Equal(got, payload) { + t.Errorf("content mismatch: got %q, want %q", string(got), string(payload)) } } -func TestExtractPkgPathTraversal(t *testing.T) { - tempDir, err := os.MkdirTemp("", "") - if err != nil { - t.Fatalf("error creating temp directory: %v", err) - } - defer oswrap.RemoveAll(tempDir) - tempFile := filepath.Join(tempDir, "test.pkg") - f, err := oswrap.Create(tempFile) +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("error creating temp file: %v", err) + t.Fatalf("client.NewDownloader failed: %v", err) } - gw := gzip.NewWriter(f) - tw := tar.NewWriter(gw) + downloader.StallTimeout = 30 * time.Millisecond - name := "foo/../../test" - body := "this is a test file" - if err := tw.WriteHeader(&tar.Header{ - Name: name, - Mode: 0600, - Size: int64(len(body)), - }); err != nil { - t.Fatal(err) + 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 _, err := tw.Write([]byte(body)); err != nil { - t.Fatalf("error writing file: %v", err) + if !errors.Is(err, client.ErrDownloadStalled) { + t.Errorf("expected error wrapping client.ErrDownloadStalled, got: %v", err) } - if err := tw.Close(); err != nil { - t.Fatalf("error closing tar: %v", err) - } - if err := gw.Close(); err != nil { - t.Fatalf("error closing gzip: %v", err) - } - if err := f.Close(); err != nil { - t.Fatalf("error closing file: %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) } +} - if _, err := ExtractPkg(tempFile); err == nil { - t.Fatal("error expected because of path traversal") +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, int64, error) { + f.mu.Lock() + f.offsets = append(f.offsets, offset) + call := len(f.offsets) + f.mu.Unlock() + if f.openErr != nil { + return nil, 0, f.openErr + } + if f.rangeErrCalls[call] { + return nil, 0, &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])}, int64(len(f.payload)), nil + } + return io.NopCloser(bytes.NewReader(f.payload[offset:])), int64(len(f.payload)), 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, int64, error) { return nil, 0, 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 f467e24..4840d51 100644 --- a/googet.go +++ b/googet.go @@ -21,9 +21,11 @@ import ( "fmt" "os" + "github.com/google/googet/v2/client" "github.com/google/googet/v2/googetdb" "github.com/google/googet/v2/progress" "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" @@ -79,6 +81,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") @@ -171,6 +189,7 @@ func run(ctx context.Context) int { } }) progress.Init(wantProgress(settings.Progress, progressSet, *progressFlag, *verbose)) + 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) diff --git a/googet.goospec b/googet.goospec index 247d89a..77e2429 100644 --- a/googet.goospec +++ b/googet.goospec @@ -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 (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.", + "3.4.0 - Fix: Failed installs restore overwritten files and preserve installer logs; batch updates continue past unresolvable packages.", "3.4.0 - Feat: Add a download progress bar and install spinner on interactive terminals; installer output is still shown live.", " - Add the -progress flag and progress googet.conf option (default true) to disable it; -verbose also disables it.", " - Fix: Restart a resumed download from the beginning when the server ignores the Range request.", diff --git a/goolib/goolib.go b/goolib/goolib.go index 7de5a1e..c476762 100644 --- a/goolib/goolib.go +++ b/goolib/goolib.go @@ -18,8 +18,10 @@ import ( "compress/gzip" "crypto/sha256" "encoding/hex" + "errors" "fmt" "io" + "os" "os/exec" "path/filepath" "regexp" @@ -28,7 +30,7 @@ import ( "strings" "syscall" - "github.com/google/googet/v2/progress" + "github.com/google/googet/v2/supervisor" ) var interpreter = map[string]string{ @@ -52,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": @@ -75,7 +82,87 @@ 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 +} + +// 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. +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 delimited). + if c != nil { + for i := 0; i < len(c.Args); i++ { + // Quoting may wrap the whole token, as in "/log:C:\install.log". + arg := strings.Trim(c.Args[i], `"'`) + + // 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) && !looksLikeSwitch(c.Args[i+1]) { + addLog(strings.Trim(strings.TrimSpace(c.Args[i+1]), `"'`)) + i++ + } + } + } + + return opts } // Run runs a command. @@ -85,20 +172,38 @@ func Exec(s string, args []string, ec []int, w io.Writer) error { // produced, a line at a time after clearing the spinner line; nothing is // withheld or discarded. func Run(c *exec.Cmd, ec []int, w io.Writer) error { - c.Stdout = io.MultiWriter(progress.Stdout(), w) - c.Stderr = io.MultiWriter(progress.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 51a89b4..62e7475 100644 --- a/goolib/goolib_test.go +++ b/goolib/goolib_test.go @@ -20,10 +20,13 @@ import ( "math/rand" "os" "os/exec" + "reflect" "runtime" "strings" "sync" "testing" + + "github.com/google/googet/v2/supervisor" ) func TestScriptInterpreter(t *testing.T) { @@ -136,6 +139,247 @@ 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) + } + }) + } +} + // syncBuffer is a bytes.Buffer that is safe for concurrent writers. type syncBuffer struct { mu sync.Mutex @@ -222,3 +466,87 @@ func TestRun(t *testing.T) { }) } } + +// 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()) + } +} + +// 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}, + } + 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) + } + }) + } +} + +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}, + } + 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/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 62b5431..88cda8c 100644 --- a/install/install.go +++ b/install/install.go @@ -34,12 +34,9 @@ import ( "github.com/google/googet/v2/progress" "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) @@ -202,7 +199,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 } @@ -228,6 +225,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 } @@ -278,7 +280,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 } @@ -337,7 +339,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) } @@ -405,18 +407,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) @@ -424,38 +428,34 @@ func makeInstallFunction(src, dst string, insFiles map[string]string, dbOnly, fo progress.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 { @@ -475,10 +475,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 } } @@ -557,18 +557,28 @@ var errInstallInterrupted = errors.New("install interrupted") // installPkg extracts and installs a package, rendering a spinner on // interactive terminals for the duration of the install. -func installPkg(pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb.GooDB) (insFiles map[string]string, err error) { +func installPkg(ops installOps, pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb.GooDB) (insFiles map[string]string, err error) { sp := progress.NewSpinner(fmt.Sprintf("Installing %s.%s.%s", ps.Name, ps.Arch, ps.Version)) // The spinner is stopped by a deferred call so that no exit path leaves it // redrawing. A normal return overwrites err before the deferred call runs; // a panic leaves errInstallInterrupted in place, so it renders "failed". err = errInstallInterrupted defer func() { sp.Stop(err) }() - return installPkgInner(pkg, ps, dbOnly, force, db) + return installPkgInner(ops, pkg, ps, dbOnly, force, db) } -// installPkgInner extracts the package, copies its files and runs its install script. -func installPkgInner(pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *googetdb.GooDB) (map[string]string, error) { +// installPkgInner 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 installPkgInner returns a nil error. +func installPkgInner(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 @@ -576,39 +586,42 @@ func installPkgInner(pkg string, ps *goolib.PkgSpec, dbOnly, force bool, db *goo logger.Infof("Executing install of package %q", filepath.Base(dir)) - toRemove = []string{} - // Try to cleanup moved files after package is installed. - defer func() { - for _, fn := range toRemove { - oswrap.Remove(fn) - } - }() + txn := newInstallTxn(ops, dbOnly, force, conflictMap) + // success is set only on the final return so that errors and panics both + // trigger rollback. + success := false - conflictMap, err := buildConflictMap(db, ps.Name) - if err != nil { - return nil, err - } + // 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) + return + } + txn.commit() + if err := oswrap.RemoveAll(dir); err != nil { + logger.Error(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..e5e3af9 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,1435 @@ 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) + } + } +} + +// 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() + 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) + } + // 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 { + 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) + if got := statMode(t, target); got != wantMode { + t.Errorf("Mode of restored %q = %v, want %v", target, got, wantMode) + } + }) + } +} + +// 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) + } + // 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) + + 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 != wantDirMode { + t.Errorf("Mode of recreated %q = %v, want %v", emptyDir, got, wantDirMode) + } + 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) + } + // 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) + 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}) + if got := statMode(t, bak); got != wantMode { + t.Errorf("Mode of %q = %v, want %v", bak, got, wantMode) + } + } +} + +// 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/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/settings/settings.go b/settings/settings.go index 9e87ca5..7685c8e 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" @@ -35,6 +36,23 @@ var ( // interactive terminals; set from googet.conf (default true) and // overridden by an explicit -progress flag. Progress = true + // 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. @@ -89,7 +107,13 @@ type conf struct { AllowUnsafeURL bool StrictConflicts bool // Progress is a pointer so an absent key keeps the default of true. - Progress *bool + Progress *bool + SupervisorMode string + InactivityTimeout string + InstallTimeout string + UIGracePeriod string + UIDetection *bool + DownloadStallTimeout string } // unmarshalConfFile returns a conf from a YAML configuration file. @@ -152,4 +176,42 @@ func readConf(filename string) { if gc.Progress != nil { Progress = *gc.Progress } + + 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 505137c..ab430fe 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) { @@ -85,3 +86,80 @@ func TestProgressDefault(t *testing.T) { t.Errorf("settings.Progress with no progress key = false, want true") } } + +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..d17c32c --- /dev/null +++ b/supervisor/msi.go @@ -0,0 +1,116 @@ +/* +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/googet/v2/progress" + "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 + 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) { + 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..cc108ae --- /dev/null +++ b/supervisor/progress.go @@ -0,0 +1,498 @@ +/* +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/googet/v2/progress" + "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 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 +} + +// 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(progress.Stdout(), out) + c.Stderr = io.MultiWriter(progress.Stderr(), out) + } else { + if c.Stdout == nil { + c.Stdout = progress.Stdout() + } + if c.Stderr == nil { + c.Stderr = progress.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): + 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) + } +} diff --git a/supervisor/supervise_test.go b/supervisor/supervise_test.go new file mode 100644 index 0000000..2ff8dd6 --- /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, 5*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..cb09dd5 --- /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 = 4 * time.Hour + 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..b280e33 --- /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}, + {"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}, + } { + 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..1c7aa2a --- /dev/null +++ b/supervisor/supervisor_windows.go @@ -0,0 +1,627 @@ +//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") + procGetProcessWindowStation = user32.NewProc("GetProcessWindowStation") + procGetUserObjectInformation = user32.NewProc("GetUserObjectInformationW") + 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 + // 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 + +// 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 + } + // 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 + } + + 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 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 { + 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..e6f5b5b --- /dev/null +++ b/supervisor/supervisor_windows_test.go @@ -0,0 +1,269 @@ +//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, + // 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) + 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