diff --git a/docs/config/config-spec.md b/docs/config/config-spec.md index a68648f..b38c40a 100644 --- a/docs/config/config-spec.md +++ b/docs/config/config-spec.md @@ -4,6 +4,9 @@ Kitout uses YAML for the MVP. +A config must contain exactly one YAML document. Additional documents, including +empty documents, are rejected instead of silently ignored. + Default path: ```txt @@ -469,6 +472,13 @@ Examples: object - same shell command name twice +Exclusive writers also cannot claim the same normalized target across resource +types: repositories, copies, symlinks (including expanded groups), SSH private +and generated public keys, and +asdf tool-version files. Errors identify both conflicting config fields. +Directory prerequisites and parent/child targets remain allowed; this check +detects exact target collisions, not overlapping directory contents. + ## Unknown fields Unknown top-level fields fail validation. diff --git a/docs/resources/copy.md b/docs/resources/copy.md index a5069fa..7ac51ca 100644 --- a/docs/resources/copy.md +++ b/docs/resources/copy.md @@ -71,9 +71,10 @@ different location than the configured target path. ## Implementation status -Implemented as `resources.CopyResource`. Status compares regular file contents -and directory trees recursively. Apply copies regular files and directories, and -dry-run plans never write files. +Implemented as `resources.CopyResource`. Status compares each directory tree +once and checks file sizes before comparing contents in bounded chunks. Apply +streams file contents rather than loading whole files into memory. Dry-run +plans never write files. ## Shared expectations diff --git a/internal/config/loader.go b/internal/config/loader.go index 1029c7c..3032abf 100644 --- a/internal/config/loader.go +++ b/internal/config/loader.go @@ -4,6 +4,7 @@ import ( "bytes" "errors" "fmt" + "io" "os" "path/filepath" "strings" @@ -144,6 +145,13 @@ func LoadFile(path string) (LoadedConfig, error) { if err := decoder.Decode(&cfg); err != nil { return LoadedConfig{}, ParseError{Path: resolvedPath, Err: err} } + var extra yaml.Node + if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) { + if err == nil { + err = errors.New("config must contain exactly one YAML document") + } + return LoadedConfig{}, ParseError{Path: resolvedPath, Err: err} + } cfg.topLevelCasksSet = topLevelCasksSet if err := validateDecodedConfig(cfg); err != nil { diff --git a/internal/config/loader_test.go b/internal/config/loader_test.go index f47985f..e545dd5 100644 --- a/internal/config/loader_test.go +++ b/internal/config/loader_test.go @@ -586,6 +586,79 @@ brew: } } +func TestLoadFileRequiresSingleYAMLDocument(t *testing.T) { + for _, suffix := range []string{"---\nversion: 1\n", "---\n", "---\nbroken: [\n"} { + t.Run(suffix, func(t *testing.T) { + _, err := LoadFile(writeConfigFile(t, "version: 1\n"+suffix)) + var parseError ParseError + if !errors.As(err, &parseError) { + t.Fatalf("LoadFile error = %v, want ParseError", err) + } + }) + } + if _, err := LoadFile(writeConfigFile(t, "version: 1\n...\n# trailing comment\n")); err != nil { + t.Fatalf("valid document terminator and comment rejected: %v", err) + } +} + +func TestLoadFileRejectsConflictingNormalizedTargetOwners(t *testing.T) { + writers := []struct{ field, yaml string }{ + {"repos[0].path", "repos:\n - path: ./target\n url: https://example.com/repo.git\n"}, + {"copies[0].target", "copies:\n - source: ./source\n target: ./nested/../target\n"}, + {"symlinks[0].target", "symlinks:\n - source: ./source\n target: ./target\n"}, + {"ssh.keys[0].path", "ssh:\n keys:\n - path: ./target\n type: ed25519\n"}, + {"asdf.tool_versions[0].path", "asdf:\n tool_versions:\n - path: ./target\n tools:\n nodejs: '22.0.0'\n"}, + } + for i, first := range writers { + for _, second := range writers[i+1:] { + t.Run(first.field+"/"+second.field, func(t *testing.T) { + _, err := LoadFile(writeConfigFile(t, "version: 1\n"+first.yaml+second.yaml)) + if err == nil || !strings.Contains(err.Error(), "conflicts with") || !strings.Contains(err.Error(), first.field) || !strings.Contains(err.Error(), second.field) { + t.Fatalf("LoadFile error = %v, want both conflicting fields", err) + } + }) + } + } + _, err := LoadFile(writeConfigFile(t, "version: 1\n"+writers[1].yaml+"symlink_groups:\n - source_root: ./source\n target_root: .\n paths: [target]\n")) + if err == nil || !strings.Contains(err.Error(), "symlink_groups[0].paths[0] conflicts with copies[0].target") { + t.Fatalf("LoadFile error = %v, want expanded group target conflict", err) + } +} + +func TestLoadFileRejectsSSHGeneratedPublicKeyTargetConflicts(t *testing.T) { + for _, other := range []string{ + "copies:\n - source: ./source\n target: ./id.pub\n", + "asdf:\n tool_versions:\n - path: ./id.pub\n tools:\n nodejs: '22.0.0'\n", + } { + _, err := LoadFile(writeConfigFile(t, "version: 1\nssh:\n keys:\n - path: ./id\n type: ed25519\n"+other)) + if err == nil || !strings.Contains(err.Error(), "conflicts with") || !strings.Contains(err.Error(), "ssh.keys[0].path (public key)") { + t.Fatalf("LoadFile error = %v, want generated public-key conflict", err) + } + } + _, err := LoadFile(writeConfigFile(t, "version: 1\nssh:\n keys:\n - path: ./id\n type: ed25519\n - path: ./id.pub\n type: ed25519\n")) + if err == nil || !strings.Contains(err.Error(), "ssh.keys[0].path (public key) conflicts with ssh.keys[1].path") { + t.Fatalf("LoadFile error = %v, want private/public key collision", err) + } +} + +func TestLoadFileAllowsDirectoriesAndNestedManagedTargets(t *testing.T) { + _, err := LoadFile(writeConfigFile(t, `version: 1 +directories: [./workspace, ./workspace/repo, ./copied] +repos: + - path: ./workspace/repo + url: https://example.com/repo.git +copies: + - source: ./source + target: ./copied +symlinks: + - source: ./source + target: ./workspace/repo/config +`)) + if err != nil { + t.Fatalf("compatible directory declarations and nested targets rejected: %v", err) + } +} + func writeConfigFile(t *testing.T, contents string) string { t.Helper() diff --git a/internal/config/validate.go b/internal/config/validate.go index f282342..ad8ed84 100644 --- a/internal/config/validate.go +++ b/internal/config/validate.go @@ -165,6 +165,7 @@ func validate(cfg Config, opts validationOptions) error { } if opts.checkPathDuplicates { errs.detectDuplicates(symlinkTargetKeys(cfg)) + errs.detectConflictingTargetOwners(cfg) } for i, item := range cfg.MacOSDefaults { @@ -304,6 +305,35 @@ func (key duplicateKey) display() string { return key.Value } +// Compare exclusive writers across resource types. Directory declarations can +// coexist with managed directories, and parent/child paths are not collisions. +func (errs *ValidationErrors) detectConflictingTargetOwners(cfg Config) { + owners := make(map[string]duplicateKey) + for _, keys := range [][]duplicateKey{ + repoPathKeys(cfg.Repos), + copyTargetKeys(cfg.Copies), + symlinkTargetKeys(cfg), + sshKeyPathKeys(cfg.SSH.Keys), + sshPublicKeyPathKeys(cfg.SSH.Keys), + asdfToolVersionPathKeys(cfg.ASDF.ToolVersions), + } { + for _, key := range keys { + if key.Value == "" { + continue + } + if owner, ok := owners[filepath.Clean(key.Value)]; ok { + errs.add(key.Field, fmt.Sprintf("conflicts with %s: both manage target %s", owner.Field, key.Value)) + } + } + // Same-type duplicates retain their existing, more specific diagnostics. + for _, key := range keys { + if key.Value != "" { + owners[filepath.Clean(key.Value)] = key + } + } + } +} + func repoPathKeys(repos []Repo) []duplicateKey { keys := make([]duplicateKey, 0, len(repos)) for i, repo := range repos { @@ -401,6 +431,17 @@ func macOSDefaultKeys(items []MacOSDefault) []duplicateKey { return keys } +func sshPublicKeyPathKeys(keys []SSHKey) []duplicateKey { + paths := sshKeyPathKeys(keys) + for i := range paths { + if paths[i].Value != "" { + paths[i].Value += ".pub" + paths[i].Field += " (public key)" + } + } + return paths +} + func sshKeyPathKeys(keys []SSHKey) []duplicateKey { duplicates := make([]duplicateKey, 0, len(keys)) for i, key := range keys { diff --git a/internal/resources/copy.go b/internal/resources/copy.go index f6f0fe6..5196e08 100644 --- a/internal/resources/copy.go +++ b/internal/resources/copy.go @@ -5,6 +5,7 @@ import ( "context" "errors" "fmt" + "io" "io/fs" "os" "path/filepath" @@ -311,16 +312,53 @@ func copyTargetsMatch(source, target string, sourceInfo, targetInfo fs.FileInfo) } } +const copyBufferSize = 32 * 1024 + func filesMatch(source, target string) (bool, error) { - sourceContents, err := os.ReadFile(source) + sourceFile, err := os.Open(source) if err != nil { return false, fmt.Errorf("could not read copy source %s: %w", source, err) } - targetContents, err := os.ReadFile(target) + defer sourceFile.Close() + targetFile, err := os.Open(target) if err != nil { return false, fmt.Errorf("could not read copy target %s: %w", target, err) } - return bytes.Equal(sourceContents, targetContents), nil + defer targetFile.Close() + sourceInfo, err := sourceFile.Stat() + if err != nil { + return false, fmt.Errorf("could not inspect copy source %s: %w", source, err) + } + targetInfo, err := targetFile.Stat() + if err != nil { + return false, fmt.Errorf("could not inspect copy target %s: %w", target, err) + } + if sourceInfo.Size() != targetInfo.Size() { + return false, nil + } + bufferSize := copyBufferSize + if sourceInfo.Size() < int64(bufferSize) { + // Keep a nonempty buffer for empty files and observe EOF in the same read + // for small files, without allocating large buffers for tiny dotfiles. + bufferSize = int(sourceInfo.Size()) + 1 + } + sourceBuffer, targetBuffer := make([]byte, bufferSize), make([]byte, bufferSize) + for { + sourceN, sourceErr := io.ReadFull(sourceFile, sourceBuffer) + if sourceErr != nil && sourceErr != io.EOF && sourceErr != io.ErrUnexpectedEOF { + return false, fmt.Errorf("could not read copy source %s: %w", source, sourceErr) + } + targetN, targetErr := io.ReadFull(targetFile, targetBuffer) + if targetErr != nil && targetErr != io.EOF && targetErr != io.ErrUnexpectedEOF { + return false, fmt.Errorf("could not read copy target %s: %w", target, targetErr) + } + if sourceN != targetN || !bytes.Equal(sourceBuffer[:sourceN], targetBuffer[:targetN]) { + return false, nil + } + if sourceErr != nil || targetErr != nil { + return sourceErr == targetErr, nil + } + } } func directoriesMatch(source, target string) (bool, error) { @@ -360,7 +398,14 @@ func directoriesMatch(source, target string) (bool, error) { if err != nil { return err } - entryMatches, err := copyTargetsMatch(path, targetPath, sourceInfo, targetInfo) + // WalkDir visits descendants itself; recurse only at the top-level dispatch. + entryMatches := targetInfo.IsDir() && targetInfo.Mode()&os.ModeSymlink == 0 + if !sourceInfo.IsDir() { + entryMatches = targetInfo.Mode().IsRegular() + if entryMatches { + entryMatches, err = filesMatch(path, targetPath) + } + } if err != nil { return err } @@ -446,14 +491,21 @@ func copyDirectory(source, target string) error { } func copyFile(source, target string, mode fs.FileMode) error { - contents, err := os.ReadFile(source) + sourceFile, err := os.Open(source) if err != nil { return err } + defer sourceFile.Close() if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { return err } - return os.WriteFile(target, contents, mode) + targetFile, err := os.OpenFile(target, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, mode) + if err != nil { + return err + } + _, copyErr := io.Copy(targetFile, sourceFile) + closeErr := targetFile.Close() + return errors.Join(copyErr, closeErr) } func (resource CopyResource) status(state engine.ResourceState, message string) engine.StatusResult { diff --git a/internal/resources/copy_benchmark_test.go b/internal/resources/copy_benchmark_test.go new file mode 100644 index 0000000..41b60dc --- /dev/null +++ b/internal/resources/copy_benchmark_test.go @@ -0,0 +1,72 @@ +package resources + +import ( + "bytes" + "context" + "fmt" + "os" + "path/filepath" + "testing" + + "github.com/vwall/kitout/internal/engine" +) + +func BenchmarkCopyStatusNested(b *testing.B) { + for _, depth := range []int{4, 8, 16} { + b.Run(fmt.Sprintf("depth-%d", depth), func(b *testing.B) { + dir := b.TempDir() + source, target := filepath.Join(dir, "source"), filepath.Join(dir, "target") + for _, root := range []string{source, target} { + for level := 0; level < depth; level++ { + root = filepath.Join(root, "nested") + if err := os.MkdirAll(root, 0o755); err != nil { + b.Fatal(err) + } + if err := os.WriteFile(filepath.Join(root, "file"), []byte("contents"), 0o644); err != nil { + b.Fatal(err) + } + } + } + resource := NewCopy(source, target, false) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + status, err := resource.Status(context.Background()) + if err != nil || status.State != engine.StateSatisfied { + b.Fatalf("Status = %+v, %v", status, err) + } + } + }) + } +} + +func BenchmarkCopyLargeFile(b *testing.B) { + dir := b.TempDir() + source, target := filepath.Join(dir, "source"), filepath.Join(dir, "target") + const size = 16 * 1024 * 1024 + if err := os.WriteFile(source, bytes.Repeat([]byte("x"), size), 0o644); err != nil { + b.Fatal(err) + } + if err := copyFile(source, target, 0o644); err != nil { + b.Fatal(err) + } + b.Run("compare", func(b *testing.B) { + b.SetBytes(size) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + matches, err := filesMatch(source, target) + if err != nil || !matches { + b.Fatalf("filesMatch = %v, %v", matches, err) + } + } + }) + b.Run("copy", func(b *testing.B) { + b.SetBytes(size) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + if err := copyFile(source, target, 0o644); err != nil { + b.Fatal(err) + } + } + }) +} diff --git a/internal/resources/copy_test.go b/internal/resources/copy_test.go index 935e65a..12467d8 100644 --- a/internal/resources/copy_test.go +++ b/internal/resources/copy_test.go @@ -1,7 +1,9 @@ package resources import ( + "bytes" "context" + "fmt" "os" "path/filepath" "testing" @@ -447,3 +449,53 @@ func TestCopyApplyRejectsCaseInsensitiveOverlap(t *testing.T) { }) } } + +func TestCopyStreamingFiles(t *testing.T) { + for _, size := range []int{0, copyBufferSize - 1, copyBufferSize, copyBufferSize + 1, 3*copyBufferSize + 17} { + t.Run(fmt.Sprint(size), func(t *testing.T) { + dir := t.TempDir() + source, target := filepath.Join(dir, "source"), filepath.Join(dir, "target") + contents := bytes.Repeat([]byte("x"), size) + if err := os.WriteFile(source, contents, 0o700); err != nil { + t.Fatal(err) + } + resource := NewCopy(source, target, true) + result, err := resource.Apply(context.Background()) + if err != nil || !result.Changed { + t.Fatalf("Apply = %+v, %v", result, err) + } + got, err := os.ReadFile(target) + if err != nil || !bytes.Equal(got, contents) { + t.Fatalf("copied contents differ: %v", err) + } + info, err := os.Stat(target) + if err != nil || info.Mode().Perm() != 0o700 { + t.Fatalf("copied mode = %v, %v", info, err) + } + status, err := resource.Status(context.Background()) + if err != nil || status.State != engine.StateSatisfied { + t.Fatalf("Status = %+v, %v", status, err) + } + if size > 0 { + for _, offset := range []int{0, size / 2, size - 1} { + changed := bytes.Clone(contents) + changed[offset] = 'y' + if err := os.WriteFile(target, changed, 0o700); err != nil { + t.Fatal(err) + } + status, err = resource.Status(context.Background()) + if err != nil || status.State != engine.StateChanged { + t.Fatalf("mismatch at %d: Status = %+v, %v", offset, status, err) + } + } + } + if err := os.WriteFile(target, append(contents, 'z'), 0o700); err != nil { + t.Fatal(err) + } + status, err = resource.Status(context.Background()) + if err != nil || status.State != engine.StateChanged { + t.Fatalf("size mismatch: Status = %+v, %v", status, err) + } + }) + } +}