Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 55 additions & 0 deletions internal/config/writer.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"fmt"
"os"
"path/filepath"
"runtime"
"sort"
"strings"

Expand Down Expand Up @@ -973,12 +974,58 @@ func writeConfigFile(path string, cfg FileConfig) error {
return writeConfigData(path, data)
}

// Sync seams allow tests to inject persistence failures at each barrier.
var syncConfigFileFn = (*os.File).Sync
var syncConfigDirFn = syncConfigDir

// Windows cannot sync directory handles, and a directory that cannot be opened
// (e.g. writable but not readable) cannot be synced either; rename durability
// is best effort in both cases, matching the sessions store. Only a failed
// Sync or Close is reported.
func syncConfigDir(dir string) error {
if runtime.GOOS == "windows" {
return nil
}
d, err := os.Open(dir)
if err != nil {
return nil
}
return errors.Join(d.Sync(), d.Close())
}

// missingConfigDirs lists dir and each ancestor that does not exist yet,
// deepest first, so the entries MkdirAll creates can be synced into their
// parents.
func missingConfigDirs(dir string) []string {
var missing []string
for d := dir; ; {
if _, err := os.Lstat(d); !errors.Is(err, os.ErrNotExist) {
return missing
}
missing = append(missing, d)
parent := filepath.Dir(d)
if parent == d {
return missing
}
d = parent
}
}

func writeConfigData(path string, data []byte) error {
dir := filepath.Dir(path)
if dir != "." && dir != "" {
created := missingConfigDirs(dir)
if err := os.MkdirAll(dir, 0o700); err != nil {
return fmt.Errorf("create config directory %s: %w", dir, err)
}
// Persist each new directory entry so a crash cannot lose the
// directory, and with it the config, after a successful write.
for _, d := range created {
parent := filepath.Dir(d)
if err := syncConfigDirFn(parent); err != nil {
return fmt.Errorf("sync config directory %s: %w", parent, err)
}
}
}
if len(data) == 0 || data[len(data)-1] != '\n' {
data = append(data, '\n')
Expand All @@ -999,11 +1046,19 @@ func writeConfigData(path string, data []byte) error {
_ = tmp.Close()
return fmt.Errorf("write config %s: %w", path, err)
}
// Persist the complete contents before making the replacement visible.
if err := syncConfigFileFn(tmp); err != nil {
return fmt.Errorf("sync config %s: %w", path, errors.Join(err, tmp.Close()))
}
if err := tmp.Close(); err != nil {
return fmt.Errorf("write config %s: %w", path, err)
}
if err := os.Rename(tmpPath, path); err != nil {
return fmt.Errorf("write config %s: %w", path, err)
}
// The replacement is already visible if this fails; do not roll it back.
if err := syncConfigDirFn(dir); err != nil {
return fmt.Errorf("sync config directory %s: %w", dir, err)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}
return nil
}
141 changes: 141 additions & 0 deletions internal/config/writer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package config
import (
"encoding/json"
"errors"
"io"
"io/fs"
"os"
"os/exec"
Expand All @@ -13,6 +14,146 @@ import (
"testing"
)

func TestWriteConfigDataSync(t *testing.T) {
failure := errors.New("injected sync failure")
for _, stage := range []string{"success", "file", "directory"} {
t.Run(stage, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "config.json")
oldData, newData := "{\"old\":true}\n", "{\"new\":true}\n"
if err := os.WriteFile(path, []byte(oldData), 0o600); err != nil {
t.Fatal(err)
}
checkData := func(path, want string) {
t.Helper()
got, err := os.ReadFile(path)
if err != nil || string(got) != want {
t.Fatalf("read %s = %q, %v; want %q", path, got, err, want)
}
}
fileSync, dirSync := syncConfigFileFn, syncConfigDirFn
t.Cleanup(func() { syncConfigFileFn, syncConfigDirFn = fileSync, dirSync })
var tmp *os.File
var calls []string
syncConfigFileFn = func(f *os.File) error {
tmp = f
calls = append(calls, "file")
checkData(f.Name(), newData)
checkData(path, oldData)
if stage == "file" {
return failure
}
return fileSync(f)
}
syncConfigDirFn = func(got string) error {
calls = append(calls, "directory")
if got != dir || tmp == nil {
t.Fatalf("directory sync = %q, temp = %v", got, tmp)
}
if _, err := tmp.Seek(0, io.SeekCurrent); !errors.Is(err, os.ErrClosed) {
t.Fatalf("temp must be closed before directory sync: %v", err)
}
checkData(path, newData)
if stage == "directory" {
return failure
}
return dirSync(got)
}
err := writeConfigData(path, []byte(strings.TrimSuffix(newData, "\n")))
if stage == "success" {
if err != nil {
t.Fatal(err)
}
} else if !errors.Is(err, failure) {
t.Fatalf("error = %v, want injected sync failure", err)
}
wantCalls := []string{"file", "directory"}
wantData := newData
if stage == "file" {
wantCalls, wantData = []string{"file"}, oldData
}
if !reflect.DeepEqual(calls, wantCalls) {
t.Fatalf("sync calls = %v, want %v", calls, wantCalls)
}
checkData(path, wantData)
if _, err := tmp.Seek(0, io.SeekCurrent); !errors.Is(err, os.ErrClosed) {
t.Fatalf("temp must be closed on return: %v", err)
}
entries, err := os.ReadDir(dir)
if err != nil || len(entries) != 1 || entries[0].Name() != "config.json" {
t.Fatalf("temporary file not cleaned up: %v, %v", entries, err)
}
})
}
}

func TestWriteConfigDataSyncsCreatedDirectories(t *testing.T) {
root := t.TempDir()
outer := filepath.Join(root, "a")
dir := filepath.Join(outer, "b")
path := filepath.Join(dir, "config.json")
failure := errors.New("injected sync failure")
dirSync := syncConfigDirFn
t.Cleanup(func() { syncConfigDirFn = dirSync })

for _, fail := range []bool{false, true} {
var calls []string
syncConfigDirFn = func(got string) error {
calls = append(calls, got)
if fail && got == root {
return failure
}
return dirSync(got)
}
if err := os.RemoveAll(outer); err != nil {
t.Fatal(err)
}
err := writeConfigData(path, []byte(`{}`))
if fail {
if !errors.Is(err, failure) {
t.Fatalf("error = %v, want injected sync failure", err)
}
if want := []string{outer, root}; !reflect.DeepEqual(calls, want) {
t.Fatalf("sync calls = %v, want %v", calls, want)
}
if _, err := os.Stat(path); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("config must not be written after a failed directory sync: %v", err)
}
continue
}
if err != nil {
t.Fatal(err)
}
// Each new directory's entry is synced into its parent, then the
// config's own directory after the rename.
if want := []string{outer, root, dir}; !reflect.DeepEqual(calls, want) {
t.Fatalf("sync calls = %v, want %v", calls, want)
}
}
}

func TestSyncConfigDir(t *testing.T) {
dir := t.TempDir()
if err := syncConfigDir(dir); err != nil {
t.Fatal(err)
}
// A directory that cannot be opened is best effort, as in the sessions
// store: the rename has already happened, so the save is not a failure.
if err := syncConfigDir(filepath.Join(dir, "missing")); err != nil {
t.Fatalf("unopenable directory must be best effort: %v", err)
}
if runtime.GOOS != "windows" {
unreadable := filepath.Join(dir, "unreadable")
if err := os.Mkdir(unreadable, 0o300); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.Chmod(unreadable, 0o700) })
if err := syncConfigDir(unreadable); err != nil {
t.Fatalf("write-only directory must be best effort: %v", err)
}
}
}

func TestSetActiveProviderSwitchesConfiguredProvider(t *testing.T) {
path := filepath.Join(t.TempDir(), "zero.json")
writeConfigFixture(t, path, FileConfig{
Expand Down
Loading