diff --git a/cli/commands/write/write.go b/cli/commands/write/write.go index c29d900..6da1ae2 100644 --- a/cli/commands/write/write.go +++ b/cli/commands/write/write.go @@ -77,7 +77,6 @@ func init() { subcommands.Register(&writeCmd{name: "write"}, "") subcommands.Register(&writeCmd{name: "update", distro: "windows", track: "stable", update: true}, "") subcommands.Register(&writeCmd{name: "windows", distro: "windows", track: "stable"}, "") - subcommands.Register(&writeCmd{name: "windowsdev", distro: "windowsdev", track: "stable"}, "") subcommands.Register(&writeCmd{name: "windowsffu", distro: "windowsffu", track: "stable", ffu: true}, "") } @@ -245,7 +244,7 @@ func (c *writeCmd) SetFlags(f *flag.FlagSet) { f.BoolVar(&c.update, "update", c.update, "attempts to perform a device refresh only for non-admin users") f.StringVar(&c.distro, "distro", c.distro, "the os distribution to be provisioned, typically 'windows' or 'linux'") f.StringVar(&c.track, "track", c.track, "track (variant) of the installer to provision") - f.StringVar(&c.confTrack, "conf_track", c.track, "track (variant) of the configuration file to provision, only valid with FFU based distros") + f.StringVar(&c.confTrack, "conf_track", "", "track (variant) of the configuration file to provision") f.StringVar(&c.seedServer, "seed_server", "", "override the default server to use for obtaining seeds, only used for debugging") f.BoolVar(&c.info, "info", false, "display console messages with debugging information included") f.IntVar(&c.v, "v", 1, "controls the level of info log verbosity") @@ -314,12 +313,6 @@ func (c *writeCmd) Execute(_ context.Context, f *flag.FlagSet, _ ...interface{}) return subcommands.ExitFailure } - // FFU images are the only ones that use confTrack. Default confTrack = track for reusability. - if !c.ffu && c.confTrack != "" { - deck.InfofA("Ignoring confTrack flag %q, as this is only used for windowsffu", c.confTrack).With(deck.V(1)).Go() - c.confTrack = "" - } - // We now know we have a valid list of devices to provision, and we can // begin provisioning. if err := execute(c, f); err != nil { diff --git a/cli/commands/write/write_test.go b/cli/commands/write/write_test.go index 9edf70a..ac6f931 100644 --- a/cli/commands/write/write_test.go +++ b/cli/commands/write/write_test.go @@ -17,6 +17,7 @@ package write import ( "context" "errors" + "fmt" "os" "path/filepath" "runtime" @@ -135,10 +136,45 @@ func TestExecute(t *testing.T) { want: subcommands.ExitFailure, }, { - desc: "--conf_track passed on non ffu distro", - cmd: &writeCmd{}, - args: []string{"--track=stable", "--conf_track=stable", "1"}, - execute: func(c *writeCmd, f *flag.FlagSet) error { return nil }, + desc: "--conf_track passed on non ffu distro", + cmd: &writeCmd{}, + args: []string{"--track=stable", "--conf_track=stable", "1"}, + execute: func(c *writeCmd, f *flag.FlagSet) error { + if c.confTrack != "stable" { + return fmt.Errorf("c.confTrack got %q, want 'stable'", c.confTrack) + } + return nil + }, + logDir: filepath.Dir(filepath.Join(os.TempDir(), binaryName)), + verbose: false, + want: subcommands.ExitSuccess, + }, + { + // An empty conf_track is passed through unchanged; config.New + // defaults it to the image track. + desc: "--conf_track is passed through empty when unspecified", + cmd: &writeCmd{}, + args: []string{"--track=testing", "1"}, + execute: func(c *writeCmd, f *flag.FlagSet) error { + if c.confTrack != "" { + return fmt.Errorf("c.confTrack got %q, want empty", c.confTrack) + } + return nil + }, + logDir: filepath.Dir(filepath.Join(os.TempDir(), binaryName)), + verbose: false, + want: subcommands.ExitSuccess, + }, + { + desc: "--conf_track explicitly specified on non ffu distro", + cmd: &writeCmd{}, + args: []string{"--track=stable", "--conf_track=unstable", "1"}, + execute: func(c *writeCmd, f *flag.FlagSet) error { + if c.confTrack != "unstable" { + return fmt.Errorf("c.confTrack got %q, want 'unstable'", c.confTrack) + } + return nil + }, logDir: filepath.Dir(filepath.Join(os.TempDir(), binaryName)), verbose: false, want: subcommands.ExitSuccess, diff --git a/cli/config/config.go b/cli/config/config.go index 5b59c43..dff5a20 100644 --- a/cli/config/config.go +++ b/cli/config/config.go @@ -63,15 +63,19 @@ const ( type distribution struct { os OperatingSystem confFile string // The final name of the config file. - confServer string // The FFU configs are obtained here. + confServer string // Runtime boot configs are obtained here. imageServer string // The base image is obtained here. label string // If set, is used to set partition labels. name string // Friendly name: e.g. Corp Windows. seedDest string // The relative path where the seed should be written. seedFile string // This file is hashed when obtainng a seed. seedServer string // If set, a seed is obtained from here. - images map[string]string - configs map[string]string // Contains config file names. + // ffu reports whether the distribution supports FFU restoration. FFU + // mode is rejected for distributions where this is false, because writing + // the FFU config would send a normal install into FFU restoration. + ffu bool + images map[string]string + configs map[string]string // Contains config file names. } // Configuration represents the state of all flags and selections provided @@ -103,20 +107,29 @@ func New(cleanup, warning, eject, ffu, update bool, devices []string, os, track, } if len(devices) > 0 { if err := conf.addDeviceList(devices); err != nil { - return nil, fmt.Errorf("addDeviceList(%q) returned %v", devices, err) + return nil, fmt.Errorf("addDeviceList(%q) returned %w", devices, err) } } // Sanity check the chosen distribution and add it to the config. if err := conf.addDistro(os); err != nil { - return nil, fmt.Errorf("addDistro(%q) returned %v", os, err) + return nil, fmt.Errorf("addDistro(%q) returned %w", os, err) + } + // FFU mode writes a config that triggers FFU restoration at boot, so it is + // only allowed for distributions that support it. + if ffu && !conf.distro.ffu { + return nil, fmt.Errorf("%w: distribution %q does not support FFU", errInput, os) } var err error // Sanity check the image and configuration tracks and add them to the config. if conf.track, err = validateTrack(track, conf.distro.images); err != nil { return nil, err } - if ffu { - if conf.confTrack, err = validateTrack(confTrack, conf.distro.configs); err != nil { + if conf.NeedsConfig() { + ct := confTrack + if ct == "" { + ct = conf.track + } + if conf.confTrack, err = validateTrack(ct, conf.distro.configs); err != nil { return nil, err } } @@ -236,8 +249,9 @@ func (c *Configuration) Track() string { return c.track } -// ConfTrack returns the selected confTrack for FFU. This generally maps -// to one of default, unstable, testing, or stable. +// ConfTrack returns the selected track of the runtime boot config, or blank +// when no config is needed. This generally maps to one of default, unstable, +// testing, or stable. func (c *Configuration) ConfTrack() string { return c.confTrack } @@ -274,19 +288,38 @@ func (c *Configuration) FFU() bool { return c.ffu } +// HasConfig returns whether or not configuration files are defined for this distribution. +func (c *Configuration) HasConfig() bool { + return c.distro != nil && c.distro.confServer != "" && len(c.distro.configs) > 0 +} + +// NeedsConfig reports whether a runtime boot config must be fetched and written. +func (c *Configuration) NeedsConfig() bool { + return c.HasConfig() || c.FFU() +} + // ConfFile returns the final name of the configuration file. func (c *Configuration) ConfFile() string { return c.distro.confFile } -// FFUConfFile returns the name of the config file. +// FFUConfFile returns the name of the runtime config file for the selected +// config track, or "" when the distribution defines none. Despite its name it +// serves both FFU and non-FFU distributions. +// TODO(b/544866964): Rename once ConfFile is retired to avoid the collision. func (c *Configuration) FFUConfFile() string { - // Return the filename only. + if c.distro == nil || c.distro.configs[c.confTrack] == "" { + return "" + } return filepath.Base(c.distro.configs[c.confTrack]) } -// FFUConfPath returns the path to the config. +// FFUConfPath returns the download URL of the runtime config file for the +// selected config track, or "" when the distribution defines none. func (c *Configuration) FFUConfPath() string { + if c.distro == nil || c.distro.confServer == "" || c.distro.configs[c.confTrack] == "" { + return "" + } return fmt.Sprintf(`%s/%s`, c.distro.confServer, c.distro.configs[c.confTrack]) } diff --git a/cli/config/config_test.go b/cli/config/config_test.go index 715adb9..65c4998 100644 --- a/cli/config/config_test.go +++ b/cli/config/config_test.go @@ -17,6 +17,7 @@ package config import ( "errors" "fmt" + "reflect" "strings" "testing" ) @@ -38,10 +39,10 @@ var ( distroDefaults = distributions ) -// cmpConfig is a custom comparer for the Configuration struct. We use a custom -// comparer to inspect public-facing members of the two structs. Errors -// describing members that do not match are returned. When all checked fields -// are equal, nil is returned. +// cmpConfig is a custom comparer for the Configuration struct. It compares the +// unexported members set by New and its helpers, including the selected +// distribution. Errors describing members that do not match are returned. When +// all checked fields are equal, nil is returned. // https://godoc.org/github.com/google/go-cmp/cmp#Exporter func cmpConfig(got, want Configuration) error { if got.track != want.track { @@ -56,19 +57,18 @@ func cmpConfig(got, want Configuration) error { if got.warning != want.warning { return fmt.Errorf("configuration warning mismatch, got: %t, want: %t", got.warning, want.warning) } + if got.ffu != want.ffu { + return fmt.Errorf("configuration ffu mismatch, got: %t, want: %t", got.ffu, want.ffu) + } + if got.elevated != want.elevated { + return fmt.Errorf("configuration elevated mismatch, got: %t, want: %t", got.elevated, want.elevated) + } if !equal(got.devices, want.devices) { return fmt.Errorf("configuration devices mismatch, got: %v, want: %v", got.devices, want.devices) } - // If no distro was provided anywhere, we can return now. - if got.distro == nil && want.distro == nil { - return nil - } - // If either distro is nil at this point, we have a mismatch. - if got.distro == nil || want.distro == nil { + if !reflect.DeepEqual(got.distro, want.distro) { return fmt.Errorf("configuration distro mismatch, got: %+v\n want: %+v", got.distro, want.distro) } - // distro's are generally static in config, so if we get here, we can safely - // assume a match, and return. return nil } @@ -86,6 +86,50 @@ func equal(left, right []string) bool { } func TestNew(t *testing.T) { + // Swap the production distributions for hermetic fixtures. + configured := distribution{ + os: windows, + name: "Configured Distro", + imageServer: imageServer, + confServer: "https://config.host.com/folder", + images: map[string]string{ + "default": "default_installer.iso", + "stable": "stable_installer.iso", + "unstable": "unstable_installer.iso", + }, + configs: map[string]string{ + "default": "default_config.yaml", + "stable": "stable_config.yaml", + }, + } + unconfigured := distribution{ + os: linux, + name: "Unconfigured Distro", + imageServer: imageServer, + images: map[string]string{ + "default": "default_installer.iso", + "stable": "stable_installer.iso", + }, + } + seeded := configured + seeded.seedServer = "https://seed.host.com" + ffuCapable := configured + ffuCapable.name = "FFU Distro" + ffuCapable.ffu = true + ffuNoConfigs := unconfigured + ffuNoConfigs.name = "FFU Distro Without Configs" + ffuNoConfigs.ffu = true + + oldDistributions, oldIsElevatedCmd := distributions, IsElevatedCmd + t.Cleanup(func() { distributions, IsElevatedCmd = oldDistributions, oldIsElevatedCmd }) + distributions = map[string]distribution{ + "configured": configured, + "unconfigured": unconfigured, + "ffu": ffuCapable, + "ffu_no_configs": ffuNoConfigs, + } + elevated := func() (bool, error) { return true, nil } + tests := []struct { desc string fakeIsElevated func() (bool, error) @@ -112,7 +156,7 @@ func TestNew(t *testing.T) { { desc: "bad track", devices: []string{"disk1"}, - os: "windows", + os: "configured", track: "foo", want: errTrack, }, @@ -120,16 +164,16 @@ func TestNew(t *testing.T) { desc: "bad ffu track", devices: []string{"disk1"}, ffu: true, - os: "windowsffu", + os: "ffu", confTrack: "foo", track: "foo", - fakeIsElevated: func() (bool, error) { return true, nil }, + fakeIsElevated: elevated, want: errTrack, }, { desc: "bad seed server", devices: []string{"disk1"}, - os: "windows", + os: "configured", track: "stable", seedServer: "test.foo@bar.com", want: errSeed, @@ -137,58 +181,167 @@ func TestNew(t *testing.T) { { desc: "isElevated error", devices: []string{"disk1"}, - os: "windows", + os: "configured", track: "stable", fakeIsElevated: func() (bool, error) { return false, errors.New("error") }, want: errElevation, }, { - desc: "valid config", + desc: "configured distro with empty confTrack defaults to track", devices: []string{"disk1"}, - os: "windows", + os: "configured", track: "stable", - fakeIsElevated: func() (bool, error) { return true, nil }, + fakeIsElevated: elevated, out: &Configuration{ - distro: &goodDistro, - track: "stable", - devices: []string{"disk1"}, - elevated: true, + distro: &configured, + track: "stable", + confTrack: "stable", + devices: []string{"disk1"}, + elevated: true, }, - want: nil, }, { - desc: "valid config with ffu", + desc: "configured distro with empty track and confTrack uses default", devices: []string{"disk1"}, - os: "windowsffu", - ffu: true, - confTrack: "unstable", + os: "configured", + fakeIsElevated: elevated, + out: &Configuration{ + distro: &configured, + track: "default", + confTrack: "default", + devices: []string{"disk1"}, + elevated: true, + }, + }, + { + desc: "configured distro with explicit confTrack", + devices: []string{"disk1"}, + os: "configured", track: "unstable", - fakeIsElevated: func() (bool, error) { return true, nil }, + confTrack: "stable", + fakeIsElevated: elevated, out: &Configuration{ - distro: &goodDistro, + distro: &configured, track: "unstable", - confTrack: "unstable", + confTrack: "stable", + devices: []string{"disk1"}, + elevated: true, + }, + }, + { + desc: "configured distro with invalid confTrack", + devices: []string{"disk1"}, + os: "configured", + track: "stable", + confTrack: "invalid_track", + fakeIsElevated: elevated, + want: errTrack, + }, + { + desc: "configured distro with image track lacking a config", + devices: []string{"disk1"}, + os: "configured", + track: "unstable", + fakeIsElevated: elevated, + want: errTrack, + }, + { + desc: "ffu distro with ffu", + devices: []string{"disk1"}, + os: "ffu", + ffu: true, + track: "stable", + confTrack: "stable", + fakeIsElevated: elevated, + out: &Configuration{ + distro: &ffuCapable, + ffu: true, + track: "stable", + confTrack: "stable", + devices: []string{"disk1"}, + elevated: true, + }, + }, + { + desc: "configured non-FFU distro rejects ffu", + devices: []string{"disk1"}, + os: "configured", + ffu: true, + track: "stable", + fakeIsElevated: elevated, + want: errInput, + }, + { + desc: "configured distro with seed server override", + devices: []string{"disk1"}, + os: "configured", + track: "stable", + seedServer: "seed.host.com", + fakeIsElevated: elevated, + out: &Configuration{ + distro: &seeded, + track: "stable", + confTrack: "stable", devices: []string{"disk1"}, elevated: true, }, - want: nil, + }, + { + desc: "unconfigured non-FFU distro ignores confTrack", + devices: []string{"disk1"}, + os: "unconfigured", + track: "stable", + confTrack: "invalid_track", + fakeIsElevated: elevated, + out: &Configuration{ + distro: &unconfigured, + track: "stable", + devices: []string{"disk1"}, + elevated: true, + }, + }, + { + desc: "unconfigured non-FFU distro rejects ffu", + devices: []string{"disk1"}, + os: "unconfigured", + ffu: true, + track: "stable", + fakeIsElevated: elevated, + want: errInput, + }, + { + desc: "ffu distro without configs has no default config", + devices: []string{"disk1"}, + os: "ffu_no_configs", + ffu: true, + track: "stable", + fakeIsElevated: elevated, + want: errInput, + }, + { + desc: "ffu distro without ffu still validates confTrack", + devices: []string{"disk1"}, + os: "ffu", + track: "stable", + confTrack: "invalid_track", + fakeIsElevated: elevated, + want: errTrack, }, } for _, tt := range tests { - IsElevatedCmd = tt.fakeIsElevated - c, got := New(false, false, false, tt.ffu, false, tt.devices, tt.os, tt.track, tt.confTrack, tt.seedServer) - if got == tt.want { - continue - } - if c == tt.out { - continue - } - if !errors.Is(got, tt.want) { - t.Errorf("%s: New() got: '%v', want: '%v'", tt.desc, got, tt.want) - } - if err := cmpConfig(*c, *tt.out); err != nil { - t.Errorf("%s: %v", tt.desc, err) - } + t.Run(tt.desc, func(t *testing.T) { + IsElevatedCmd = tt.fakeIsElevated + c, err := New(false, false, false, tt.ffu, false, tt.devices, tt.os, tt.track, tt.confTrack, tt.seedServer) + if !errors.Is(err, tt.want) { + t.Fatalf("New() returned %v, want %v", err, tt.want) + } + if tt.out == nil { + return + } + if err := cmpConfig(*c, *tt.out); err != nil { + t.Error(err) + } + }) } } @@ -314,6 +467,8 @@ func TestAddDeviceList(t *testing.T) { } func TestAddSeedServer(t *testing.T) { + overridden := goodDistro + overridden.seedServer = "https://foo.bar.com" tests := []struct { desc string server string @@ -338,7 +493,7 @@ func TestAddSeedServer(t *testing.T) { desc: "good fqdn", server: "foo.bar.com", distro: goodDistro, - out: Configuration{distro: &goodDistro}, + out: Configuration{distro: &overridden}, want: nil, }, } @@ -603,3 +758,147 @@ func TestString(t *testing.T) { t.Errorf("String() got: %q, want contains: %q", got, want) } } + +func TestHasConfig(t *testing.T) { + tests := []struct { + desc string + distro *distribution + want bool + }{ + { + desc: "nil distro", + distro: nil, + want: false, + }, + { + desc: "empty confServer", + distro: &distribution{ + confServer: "", + configs: map[string]string{"default": "conf.yaml"}, + }, + want: false, + }, + { + desc: "empty configs", + distro: &distribution{ + confServer: "https://foo.bar.com/configs", + configs: map[string]string{}, + }, + want: false, + }, + { + desc: "confServer and configs set", + distro: &distribution{ + confServer: "https://foo.bar.com/configs", + configs: map[string]string{"default": "conf.yaml"}, + }, + want: true, + }, + } + for _, tt := range tests { + c := Configuration{distro: tt.distro} + if got := c.HasConfig(); got != tt.want { + t.Errorf("%s: HasConfig() got: %t, want: %t", tt.desc, got, tt.want) + } + } +} + +func TestNeedsConfig(t *testing.T) { + withConfig := &distribution{ + confServer: "https://config.host.com/folder", + configs: map[string]string{"default": "conf.yaml"}, + } + tests := []struct { + desc string + distro *distribution + ffu bool + want bool + }{ + {desc: "no config and no ffu", distro: &distribution{}, want: false}, + {desc: "nil distro with ffu", distro: nil, ffu: true, want: true}, + {desc: "config without ffu", distro: withConfig, want: true}, + {desc: "config with ffu", distro: withConfig, ffu: true, want: true}, + } + for _, tt := range tests { + c := Configuration{distro: tt.distro, ffu: tt.ffu} + if got := c.NeedsConfig(); got != tt.want { + t.Errorf("%s: NeedsConfig() got: %t, want: %t", tt.desc, got, tt.want) + } + } +} + +func TestFFUConfFileGuards(t *testing.T) { + tests := []struct { + desc string + confTrack string + distro *distribution + want string + }{ + { + desc: "nil distro", + confTrack: "default", + distro: nil, + want: "", + }, + { + desc: "empty configs", + confTrack: "default", + distro: &distribution{}, + want: "", + }, + { + desc: "unmatched track", + confTrack: "nonexistent", + distro: &distribution{ + configs: map[string]string{"default": "conf.yaml"}, + }, + want: "", + }, + } + for _, tt := range tests { + c := Configuration{confTrack: tt.confTrack, distro: tt.distro} + if got := c.FFUConfFile(); got != tt.want { + t.Errorf("%s: FFUConfFile() got: %q, want: %q", tt.desc, got, tt.want) + } + } +} + +func TestFFUConfPathGuards(t *testing.T) { + tests := []struct { + desc string + confTrack string + distro *distribution + want string + }{ + { + desc: "nil distro", + confTrack: "default", + distro: nil, + want: "", + }, + { + desc: "empty confServer", + confTrack: "default", + distro: &distribution{ + confServer: "", + configs: map[string]string{"default": "conf.yaml"}, + }, + want: "", + }, + { + desc: "unmatched track", + confTrack: "nonexistent", + distro: &distribution{ + confServer: "https://foo.bar.com", + configs: map[string]string{"default": "conf.yaml"}, + }, + want: "", + }, + } + for _, tt := range tests { + c := Configuration{confTrack: tt.confTrack, distro: tt.distro} + if got := c.FFUConfPath(); got != tt.want { + t.Errorf("%s: FFUConfPath() got: %q, want: %q", tt.desc, got, tt.want) + } + } +} diff --git a/cli/config/defaults.go b/cli/config/defaults.go index d61a77c..17cc3c2 100644 --- a/cli/config/defaults.go +++ b/cli/config/defaults.go @@ -41,6 +41,7 @@ var ( name: "windows", imageServer: "https://image.host.com/folder", confServer: "https://config.host.com/folder", + ffu: true, images: map[string]string{ "default": "installer_img.iso", "stable": "installer_img.iso", diff --git a/cli/installer/installer.go b/cli/installer/installer.go index f754fc0..4da2f00 100644 --- a/cli/installer/installer.go +++ b/cli/installer/installer.go @@ -46,7 +46,13 @@ import ( const ( oneGB = uint64(1073741824) seedDestFile = `seed.json` + // confDestFile is the config name for FFU distributions. WinPE treats its + // presence on the OCI volume as a request for FFU/OSD restoration. confDestFile = `startimage.yaml` + // bootConfDestFile is the config name for non-FFU distributions. Unified + // WinPE media reads it to select the OS at boot; older per-OS media ignore + // it and use the OS baked into the image. + bootConfDestFile = `bootconfig.yaml` ) var ( @@ -117,6 +123,7 @@ type Configuration interface { ImageFile() string Elevated() bool FFU() bool + NeedsConfig() bool PowerOff() bool SeedDest() string SeedFile() string @@ -279,8 +286,8 @@ func (i *Installer) retrieveFile(fileName, filePath string) (err error) { return downloadFile(client, filePath, f) } -// Retrieve passes the necessary parameters to retrieveFile -// depending on whether or not the distribution will be FFU based. +// Retrieve downloads the image file and, when the configuration needs a +// runtime boot config, the config file for the selected config track. func (i *Installer) Retrieve() (err error) { // Confirm that the Installer has what we need. if i.config.ImagePath() == "" { @@ -290,9 +297,9 @@ func (i *Installer) Retrieve() (err error) { return errCache } - // If FFU is false, retrieve only the image file. - // Otherwise retrieve the image file and FFU manifest. - if !i.config.FFU() { + // If no runtime config is needed, retrieve only the image file. + // Otherwise retrieve the image file and configuration manifest. + if !i.config.NeedsConfig() { return i.retrieveFile(i.config.ImageFile(), i.config.ImagePath()) } @@ -527,8 +534,8 @@ func (i *Installer) Provision(d Device) error { if _, err := os.Stat(path); err != nil { return fmt.Errorf("os.Stat(%q) returned %v: %w", path, err, errPath) } - // Check that the FFU config is already in the cache. - if i.config.FFU() { + // Check that the config is already in the cache. + if i.config.NeedsConfig() { deck.InfofA("Checking %q for existence of %q.", i.cache, i.config.FFUConfFile()).With(deck.V(2)).Go() path := filepath.Join(i.cache, i.config.FFUConfFile()) if _, err := os.Stat(path); err != nil { @@ -597,10 +604,10 @@ func (i *Installer) provisionISO(d Device) (err error) { return fmt.Errorf("writeISO() returned %v: %w", err, errProvision) } - // If FFU, write config to disk. - if i.config.FFU() { + // If configured or FFU, write config to disk. + if i.config.NeedsConfig() { if err := i.writeConfig(p); err != nil { - return fmt.Errorf("writeConfig() returned %v", err) + return fmt.Errorf("writeConfig() returned %w", err) } } @@ -727,7 +734,8 @@ func (i *Installer) writeSeed(h isoHandler, p partition) error { return nil } -// writeConfig writes the FFU config file to disk using SeedDest directory. +// writeConfig writes the distribution config file to the SeedDest directory. +// FFU distributions get startimage.yaml and all others get bootconfig.yaml. func (i *Installer) writeConfig(p partition) error { source := filepath.Join(i.cache, i.config.FFUConfFile()) var content []byte @@ -749,7 +757,18 @@ func (i *Installer) writeConfig(p partition) error { if err := os.MkdirAll(dest, 0755); err != nil { return fmt.Errorf("os.MkdirAll(%q, 0755) returned %v: %w", dest, err, errPerm) } - destFile := filepath.Join(dest, confDestFile) + destName, staleName := bootConfDestFile, confDestFile + if i.config.FFU() { + destName, staleName = confDestFile, bootConfDestFile + } + // Remove the counterpart first so a leftover startimage.yaml cannot send a + // non-FFU boot into FFU restoration, and vice versa. Removing before writing + // guarantees both files never coexist if either step fails. + staleFile := filepath.Join(dest, staleName) + if err := os.Remove(staleFile); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("os.Remove(%q) returned %v: %w", staleFile, err, errIO) + } + destFile := filepath.Join(dest, destName) deck.InfofA("Writing config: %q.", destFile).With(deck.V(2)).Go() // Permissions = owner:read/write, group:read" if err := ioutil.WriteFile(destFile, content, 0644); err != nil { diff --git a/cli/installer/installer_test.go b/cli/installer/installer_test.go index 8a04566..c22a5fb 100644 --- a/cli/installer/installer_test.go +++ b/cli/installer/installer_test.go @@ -43,12 +43,13 @@ type fakeConfig struct { // config.Configuration is embedded, fakeConfig inherits all its members. config.Configuration - dismount bool - eject bool - elevated bool - ffu bool - update bool - err error // the error returned when isElevated is called. + dismount bool + eject bool + elevated bool + ffu bool + hasConfig bool + update bool + err error // the error returned when isElevated is called. confFile string distroLabel string @@ -62,6 +63,10 @@ type fakeConfig struct { ffuConfPath string } +func (f *fakeConfig) NeedsConfig() bool { + return f.hasConfig || f.ffu +} + func (f *fakeConfig) ConfFile() string { return f.confFile } @@ -271,6 +276,50 @@ func TestRetrieve(t *testing.T) { download: func(client httpDoer, path string, w io.Writer) error { return nil }, want: nil, }, + { + desc: "non-ffu with hasConfig download success", + installer: &Installer{cache: fakeCache, config: &fakeConfig{ + imagePath: `https://foo.bar.com/test_installer.img`, + imageFile: `test_installer.img`, + hasConfig: true, + ffuConfPath: "https://foo.bar.com/config/startimage.yaml", + ffuConfFile: "startimage.yaml", + }}, + download: func(client httpDoer, path string, w io.Writer) error { return nil }, + want: nil, + }, + { + desc: "non-ffu with hasConfig missing yaml config", + installer: &Installer{cache: fakeCache, config: &fakeConfig{ + imagePath: `https://foo.bar.com/test_installer.img`, + imageFile: `test_installer.img`, + hasConfig: true, + ffuConfFile: "", + ffuConfPath: "", + }}, + want: errConfName, + }, + { + desc: "non-ffu with hasConfig missing yaml path", + installer: &Installer{cache: fakeCache, config: &fakeConfig{ + imagePath: `https://foo.bar.com/test_installer.img`, + imageFile: `test_installer.img`, + hasConfig: true, + ffuConfFile: "startimage.yaml", + ffuConfPath: "", + }}, + want: errConfPath, + }, + { + desc: "non-ffu without hasConfig downloads only image", + installer: &Installer{cache: fakeCache, config: &fakeConfig{ + imagePath: `https://foo.bar.com/test_installer.img`, + imageFile: `test_installer.img`, + hasConfig: false, + }}, + download: func(client httpDoer, path string, w io.Writer) error { return nil }, + want: nil, + }, } for _, tt := range tests { downloadFile = tt.download @@ -786,6 +835,11 @@ func TestProvision(t *testing.T) { if _, err := os.Create(fakeImagePath); err != nil { t.Fatalf("os.Create(%q) returned %v", fakeImagePath, err) } + fakeConfPath := filepath.Join(fakeCache, "fake_conf.yaml") + if _, err := os.Create(fakeConfPath); err != nil { + t.Fatalf("os.Create(%q) returned %v", fakeConfPath, err) + } + defer os.RemoveAll(fakeCache) tests := []struct { desc string @@ -825,6 +879,30 @@ func TestProvision(t *testing.T) { installer: &Installer{cache: "/fake/path", config: &fakeConfig{imageFile: "fake.iso"}}, want: errPath, }, + { + desc: "hasConfig config file missing from cache", + installer: &Installer{cache: fakeCache, config: &fakeConfig{ + imageFile: "fake.iso", + hasConfig: true, + ffuConfFile: "missing_conf.yaml", + }}, + want: errPath, + }, + { + desc: "hasConfig success", + installer: &Installer{cache: fakeCache, config: &fakeConfig{ + imageFile: "fake.iso", + hasConfig: true, + ffuConfFile: "fake_conf.yaml", + seedDest: "oci", + }}, + mount: func(string) (isoHandler, error) { return &fakeHandler{}, nil }, + selPart: func(Device, uint64, storage.FileSystem) (partition, error) { + return &fakePartition{label: "test", id: "testid", mount: fakeCache}, nil + }, + writeISO: func(isoHandler, partition) error { return nil }, + want: nil, + }, { desc: "success", installer: &Installer{cache: fakeCache, config: &fakeConfig{imageFile: "fake.iso"}}, @@ -860,6 +938,11 @@ func TestProvisionISO(t *testing.T) { if _, err := os.Create(fakeImagePath); err != nil { t.Fatalf("os.Create(%q) returned %v", fakeImagePath, err) } + fakeConfPath := filepath.Join(fakeCache, "fake_conf.yaml") + if _, err := os.Create(fakeConfPath); err != nil { + t.Fatalf("os.Create(%q) returned %v", fakeConfPath, err) + } + defer os.RemoveAll(fakeCache) tests := []struct { desc string @@ -904,6 +987,35 @@ func TestProvisionISO(t *testing.T) { writeISO: func(isoHandler, partition) error { return nil }, want: errIO, }, + { + desc: "writeConfig error with hasConfig", + installer: &Installer{cache: fakeCache, config: &fakeConfig{ + imageFile: "fake.iso", + hasConfig: true, + ffuConfFile: "missing.yaml", + }}, + mount: func(string) (isoHandler, error) { return &fakeHandler{}, nil }, + device: &fakeDevice{}, + selPart: func(Device, uint64, storage.FileSystem) (partition, error) { return &fakePartition{label: "test"}, nil }, + writeISO: func(isoHandler, partition) error { return nil }, + want: errIO, + }, + { + desc: "success with hasConfig", + installer: &Installer{cache: fakeCache, config: &fakeConfig{ + imageFile: "fake.iso", + hasConfig: true, + ffuConfFile: "fake_conf.yaml", + seedDest: "oci", + }}, + mount: func(string) (isoHandler, error) { return &fakeHandler{}, nil }, + device: &fakeDevice{}, + selPart: func(Device, uint64, storage.FileSystem) (partition, error) { + return &fakePartition{label: "test", mount: fakeCache}, nil + }, + writeISO: func(isoHandler, partition) error { return nil }, + want: nil, + }, { desc: "success", installer: &Installer{cache: fakeCache, config: &fakeConfig{imageFile: "fake.iso"}}, @@ -925,6 +1037,224 @@ func TestProvisionISO(t *testing.T) { } } +// TestWriteConfig verifies the destination, content, permissions, and failure +// modes of writeConfig, including removal of the counterpart config file. +func TestWriteConfig(t *testing.T) { + const confFileName = "test_config.yaml" + confContent := []byte("os_code: test-os-stable\ntrack: stable\nmanaged: true\n") + staleContent := []byte("stale: true\n") + + // writeFile creates a file and any missing parent directories. + writeFile := func(t *testing.T, path string, content []byte, perm os.FileMode) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + t.Fatalf("os.MkdirAll(%q) returned %v", filepath.Dir(path), err) + } + if err := os.WriteFile(path, content, perm); err != nil { + t.Fatalf("os.WriteFile(%q) returned %v", path, err) + } + } + + tests := []struct { + desc string + ffu bool + // seedDest overrides the default "/oci" destination when set. + seedDest string + // confFile overrides the default cached config file name when set. + confFile string + // emptySource caches an empty config file instead of confContent. + emptySource bool + // needsPermEnforcement marks cases that rely on file permission checks, + // which are not enforced on Windows or for root, so they are skipped there. + needsPermEnforcement bool + // setup prepares the destination directory before writeConfig runs. + setup func(t *testing.T, dest string) + wantErr error + // wantFile is the file expected to hold the config on success. + wantFile string + // wantAbsent lists files in dest that must not exist afterwards. + wantAbsent []string + }{ + { + desc: "FFU writes startimage.yaml to a fresh mount", + ffu: true, + wantFile: confDestFile, + wantAbsent: []string{bootConfDestFile}, + }, + { + desc: "non-FFU writes bootconfig.yaml to a fresh mount", + wantFile: bootConfDestFile, + wantAbsent: []string{confDestFile}, + }, + { + desc: "target directory already exists", + ffu: true, + setup: func(t *testing.T, dest string) { + if err := os.MkdirAll(dest, 0755); err != nil { + t.Fatalf("os.MkdirAll(%q) returned %v", dest, err) + } + }, + wantFile: confDestFile, + }, + { + desc: "target file with stale content is overwritten", + ffu: true, + setup: func(t *testing.T, dest string) { + writeFile(t, filepath.Join(dest, confDestFile), staleContent, 0644) + }, + wantFile: confDestFile, + }, + { + desc: "FFU removes stale bootconfig.yaml", + ffu: true, + setup: func(t *testing.T, dest string) { + writeFile(t, filepath.Join(dest, bootConfDestFile), staleContent, 0644) + }, + wantFile: confDestFile, + wantAbsent: []string{bootConfDestFile}, + }, + { + desc: "non-FFU removes stale startimage.yaml", + setup: func(t *testing.T, dest string) { + writeFile(t, filepath.Join(dest, confDestFile), staleContent, 0644) + }, + wantFile: bootConfDestFile, + wantAbsent: []string{confDestFile}, + }, + { + desc: "empty config file in cache", + ffu: true, + emptySource: true, + wantFile: confDestFile, + }, + { + desc: "deep nested destination directory", + ffu: true, + seedDest: "/deep/nested/custom/oci", + wantFile: confDestFile, + }, + { + desc: "source config missing from cache", + ffu: true, + confFile: "nonexistent.yaml", + wantErr: errIO, + wantAbsent: []string{confDestFile}, + }, + { + desc: "target directory path blocked by file", + ffu: true, + setup: func(t *testing.T, dest string) { + writeFile(t, dest, []byte("blocking file"), 0644) + }, + wantErr: errPerm, + }, + { + desc: "FFU stale bootconfig.yaml cannot be removed", + ffu: true, + setup: func(t *testing.T, dest string) { + // A non-empty directory at the stale path makes os.Remove fail on every OS. + writeFile(t, filepath.Join(dest, bootConfDestFile, "child"), staleContent, 0644) + }, + wantErr: errIO, + wantAbsent: []string{confDestFile}, + }, + { + desc: "non-FFU stale startimage.yaml cannot be removed", + setup: func(t *testing.T, dest string) { + // A non-empty directory at the stale path makes os.Remove fail on every OS. + writeFile(t, filepath.Join(dest, confDestFile, "child"), staleContent, 0644) + }, + wantErr: errIO, + wantAbsent: []string{bootConfDestFile}, + }, + { + desc: "read only destination directory", + ffu: true, + needsPermEnforcement: true, + setup: func(t *testing.T, dest string) { + if err := os.MkdirAll(dest, 0555); err != nil { + t.Fatalf("os.MkdirAll(%q) returned %v", dest, err) + } + t.Cleanup(func() { os.Chmod(dest, 0755) }) + }, + wantErr: errIO, + wantAbsent: []string{confDestFile}, + }, + { + desc: "read only destination file", + ffu: true, + needsPermEnforcement: true, + setup: func(t *testing.T, dest string) { + path := filepath.Join(dest, confDestFile) + writeFile(t, path, []byte("readonly"), 0444) + t.Cleanup(func() { os.Chmod(path, 0644) }) + }, + wantErr: errIO, + }, + } + for _, tt := range tests { + t.Run(tt.desc, func(t *testing.T) { + if tt.needsPermEnforcement && (runtime.GOOS == "windows" || os.Geteuid() == 0) { + t.Skip("permission checks are not enforced on Windows or for root") + } + cache := t.TempDir() + mount := t.TempDir() + source := confContent + if tt.emptySource { + source = []byte{} + } + writeFile(t, filepath.Join(cache, confFileName), source, 0644) + seedDest := "/oci" + if tt.seedDest != "" { + seedDest = tt.seedDest + } + confFile := confFileName + if tt.confFile != "" { + confFile = tt.confFile + } + dest := filepath.Join(mount, filepath.FromSlash(seedDest)) + if tt.setup != nil { + tt.setup(t, dest) + } + + inst := &Installer{cache: cache, config: &fakeConfig{ffu: tt.ffu, ffuConfFile: confFile, seedDest: seedDest}} + err := inst.writeConfig(&fakePartition{mount: mount}) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("writeConfig() returned %v, want %v", err, tt.wantErr) + } + + if tt.wantFile != "" { + path := filepath.Join(dest, tt.wantFile) + got, err := os.ReadFile(path) + if err != nil { + t.Fatalf("os.ReadFile(%q) returned %v", path, err) + } + if string(got) != string(source) { + t.Errorf("writeConfig() wrote %q to %s, want %q", got, tt.wantFile, source) + } + if runtime.GOOS != "windows" { + info, err := os.Stat(path) + if err != nil { + t.Fatalf("os.Stat(%q) returned %v", path, err) + } + // The process umask may clear group or other bits from the + // requested 0644 (for example 0640 under umask 027), so only + // require owner read/write and nothing beyond 0644. + if mode := info.Mode().Perm(); mode&0600 != 0600 || mode&^0644 != 0 { + t.Errorf("%s permissions = %v, want owner read/write and no bits beyond 0644", tt.wantFile, mode) + } + } + } + for _, name := range tt.wantAbsent { + path := filepath.Join(dest, name) + if _, err := os.Stat(path); !os.IsNotExist(err) { + t.Errorf("os.Stat(%q) returned %v, want the file to be absent", path, err) + } + } + }) + } +} + // fakeISO represents iso.Handler. It inherits all members of iso.Handler // through embedding. Unimplemented members send a clear signal during tests // because they will panic if called, allowing us to implement only the minimum @@ -1273,3 +1603,61 @@ func TestFinalize(t *testing.T) { } } } + +// TestUnconfiguredDistroSkipsConfig verifies that Retrieve and provisionISO do +// not fetch or write a configuration file when NeedsConfig is false. +func TestUnconfiguredDistroSkipsConfig(t *testing.T) { + oldDownloadFile, oldMount, oldWriteISO, oldSelectPart := downloadFile, mount, writeISOFunc, selectPart + defer func() { + downloadFile, mount, writeISOFunc, selectPart = oldDownloadFile, oldMount, oldWriteISO, oldSelectPart + }() + + fakeCache := t.TempDir() + fakeMount := t.TempDir() + + // Unconfigured distro: hasConfig=false, ffu=false. + downloadedPaths := []string{} + downloadFile = func(client httpDoer, path string, w io.Writer) error { + downloadedPaths = append(downloadedPaths, path) + return nil + } + + inst := &Installer{ + cache: fakeCache, + config: &fakeConfig{ + imagePath: "https://image.host.com/media/stable/test.iso", + imageFile: "test.iso", + hasConfig: false, + ffu: false, + }, + } + + if err := inst.Retrieve(); err != nil { + t.Fatalf("Retrieve() returned %v", err) + } + if len(downloadedPaths) != 1 { + t.Errorf("Retrieve() called download %d times, want exactly 1", len(downloadedPaths)) + } + if len(downloadedPaths) > 0 && downloadedPaths[0] != "https://image.host.com/media/stable/test.iso" { + t.Errorf("downloaded path got %q, want test.iso URL", downloadedPaths[0]) + } + + // Test provisionISO with an unconfigured distro. + mount = func(string) (isoHandler, error) { return &fakeHandler{}, nil } + writeISOFunc = func(isoHandler, partition) error { return nil } + selectPart = func(Device, uint64, storage.FileSystem) (partition, error) { + return &fakePartition{label: "test", mount: fakeMount}, nil + } + + if err := inst.provisionISO(&fakeDevice{}); err != nil { + t.Fatalf("provisionISO() returned %v", err) + } + + // Verify that no config file was created. + for _, name := range []string{confDestFile, bootConfDestFile} { + ociFile := filepath.Join(fakeMount, "oci", name) + if _, err := os.Stat(ociFile); !os.IsNotExist(err) { + t.Errorf("provisionISO() created %q for unconfigured distro, want no file", ociFile) + } + } +}