diff --git a/tools/tht/internal/backup/preflight.go b/tools/tht/internal/backup/preflight.go index 6ca1c2a6..00b0203a 100644 --- a/tools/tht/internal/backup/preflight.go +++ b/tools/tht/internal/backup/preflight.go @@ -287,7 +287,7 @@ func (result PreflightResult) StageArchive(ctx context.Context) (_ *stagedArchiv return nil, errors.New("protect private restore staging directory") } path := filepath.Join(directory, "archive.zip") - file, err := os.OpenFile(path, os.O_RDWR|os.O_CREATE|os.O_EXCL, 0o600) + file, err := safeio.CreateCanonicalNewPrivateFile(path) if err != nil { _ = os.Remove(directory) return nil, errors.New("create private restore staging archive") diff --git a/tools/tht/internal/backup/preflight_test.go b/tools/tht/internal/backup/preflight_test.go index a720ac91..ae9c58dc 100644 --- a/tools/tht/internal/backup/preflight_test.go +++ b/tools/tht/internal/backup/preflight_test.go @@ -455,6 +455,37 @@ func TestPreflightStagesArchiveIntoImmutablePrivateBytes(t *testing.T) { } } +func TestStageArchiveCleansPrivateFileAfterStreamingFailure(t *testing.T) { + installation := preflightTestInstallation(t) + if err := os.MkdirAll(installation.ControlDirectory(), 0o700); err != nil { + t.Fatal(err) + } + archive := filepath.Join(t.TempDir(), "checked.zip") + writePreflightArchive(t, archive, preflightArchiveSpec{ + entries: []preflightArchiveEntry{{path: "configuration/operator.env", body: []byte("before")}}, + }) + + result, err := Preflight(context.Background(), installation, PreflightRequest{Archive: archive, Confirm: true}, permissivePreflightDependencies()) + if err != nil { + t.Fatal(err) + } + defer result.CloseArchive() + result.freeBytes = func(string) (uint64, error) { + result.archive.digest = "after-preflight-mismatch" + return 1024, nil + } + if _, err := result.StageArchive(context.Background()); err == nil || !strings.Contains(err.Error(), "changed") { + t.Fatalf("StageArchive() error = %v, want streaming size refusal", err) + } + entries, err := os.ReadDir(result.stagingRoot) + if err != nil { + t.Fatal(err) + } + if len(entries) != 0 { + t.Fatalf("private staging leftovers = %v, want none", entries) + } +} + func TestPreflightStagingRejectsInPlaceArchiveHashMutation(t *testing.T) { installation := preflightTestInstallation(t) archive := filepath.Join(t.TempDir(), "checked.zip") diff --git a/tools/tht/internal/backup/preflight_unix_test.go b/tools/tht/internal/backup/preflight_unix_test.go index 05c8d79e..ce691fe7 100644 --- a/tools/tht/internal/backup/preflight_unix_test.go +++ b/tools/tht/internal/backup/preflight_unix_test.go @@ -10,6 +10,7 @@ import ( "testing" "github.com/aritmolab/thothii/tools/tht/internal/config" + "github.com/aritmolab/thothii/tools/tht/internal/safeio" ) func TestStageArchiveRejectsSymlinkedInstallationAncestor(t *testing.T) { @@ -46,3 +47,35 @@ func TestStageArchiveRejectsSymlinkedInstallationAncestor(t *testing.T) { t.Fatalf("StageArchive() error = %v, want unsafe symlinked staging-root rejection", err) } } + +func TestStageArchiveCreatesUnixPrivateRegularFile(t *testing.T) { + installation := preflightTestInstallation(t) + if err := os.MkdirAll(installation.ControlDirectory(), 0o700); err != nil { + t.Fatal(err) + } + archive := filepath.Join(t.TempDir(), "valid.zip") + writePreflightArchive(t, archive, preflightArchiveSpec{ + entries: []preflightArchiveEntry{{path: "configuration/operator.env", body: []byte("safe")}}, + }) + result, err := Preflight(context.Background(), installation, PreflightRequest{Archive: archive, Confirm: true}, permissivePreflightDependencies()) + if err != nil { + t.Fatal(err) + } + defer result.CloseArchive() + + staged, err := result.StageArchive(context.Background()) + if err != nil { + t.Fatal(err) + } + defer staged.Close() + if err := safeio.ValidatePrivateRegular(staged.path); err != nil { + t.Fatalf("staged archive privacy = %v, want owner-private regular file", err) + } + info, err := os.Lstat(staged.path) + if err != nil { + t.Fatal(err) + } + if got := info.Mode().Perm(); got != 0o600 { + t.Fatalf("staged archive mode = %#o, want 0600", got) + } +} diff --git a/tools/tht/internal/backup/preflight_windows_test.go b/tools/tht/internal/backup/preflight_windows_test.go index 48cde248..5968e83b 100644 --- a/tools/tht/internal/backup/preflight_windows_test.go +++ b/tools/tht/internal/backup/preflight_windows_test.go @@ -11,7 +11,7 @@ import ( "github.com/aritmolab/thothii/tools/tht/internal/safeio" ) -func TestStageArchiveProtectsWindowsStagingDirectoriesWithOwnerOnlyACLs(t *testing.T) { +func TestStageArchiveProtectsWindowsStagingArtifactsWithOwnerOnlyACLs(t *testing.T) { installation := preflightTestInstallation(t) if err := os.MkdirAll(installation.ControlDirectory(), 0o700); err != nil { t.Fatal(err) @@ -37,4 +37,7 @@ func TestStageArchiveProtectsWindowsStagingDirectoriesWithOwnerOnlyACLs(t *testi if err := safeio.ValidatePrivateDirectory(staged.directory); err != nil { t.Fatalf("staged archive directory ACL = %v, want owner-only", err) } + if err := safeio.ValidatePrivateRegular(staged.path); err != nil { + t.Fatalf("staged archive ACL = %v, want owner-only", err) + } } diff --git a/tools/tht/internal/safeio/files.go b/tools/tht/internal/safeio/files.go index 33b0a7fd..36b9d11c 100644 --- a/tools/tht/internal/safeio/files.go +++ b/tools/tht/internal/safeio/files.go @@ -275,25 +275,24 @@ func WriteCanonicalNewPrivateFile(path string, contents []byte, mode os.FileMode return writeCanonicalNewFile(path, contents, mode, true) } +// CreateCanonicalNewPrivateFile exclusively creates an owner-private regular file under an +// already private parent and returns a read/write handle for streamed contents. Callers must close +// the returned handle and remove the file if their stream fails. +func CreateCanonicalNewPrivateFile(path string) (*os.File, error) { + if err := validateCanonicalNewFile(path, true); err != nil { + return nil, err + } + file, err := createCanonicalNewPrivateParentReadWriteFile(path, 0o600) + if err != nil { + return nil, ErrUnsafeFile + } + return file, nil +} + func writeCanonicalNewFile(path string, contents []byte, mode os.FileMode, requirePrivateParent bool) error { - if err := ValidateCanonicalPath(path); err != nil { + if err := validateCanonicalNewFile(path, requirePrivateParent); err != nil { return err } - parent := filepath.Dir(path) - if err := requireCanonicalDirectory(parent); err != nil { - return err - } - if requirePrivateParent && ValidatePrivateDirectory(parent) != nil { - return ErrUnsafeFile - } - if info, err := os.Lstat(path); err == nil { - if !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || info.Mode()&os.ModeType != 0 { - return ErrUnsafeFile - } - return ErrUnsafeFile - } else if !errors.Is(err, os.ErrNotExist) { - return ErrUnsafeFile - } var ( file *os.File err error @@ -323,6 +322,28 @@ func writeCanonicalNewFile(path string, contents []byte, mode os.FileMode, requi return nil } +func validateCanonicalNewFile(path string, requirePrivateParent bool) error { + if err := ValidateCanonicalPath(path); err != nil { + return err + } + parent := filepath.Dir(path) + if err := requireCanonicalDirectory(parent); err != nil { + return err + } + if requirePrivateParent && ValidatePrivateDirectory(parent) != nil { + return ErrUnsafeFile + } + if info, err := os.Lstat(path); err == nil { + if !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || info.Mode()&os.ModeType != 0 { + return ErrUnsafeFile + } + return ErrUnsafeFile + } else if !errors.Is(err, os.ErrNotExist) { + return ErrUnsafeFile + } + return nil +} + // ReplaceCanonicalRegular durably replaces one existing private regular file without following // symlinked path components. Platform implementations keep the temporary file in the target // directory and use the platform's atomic replace primitive. diff --git a/tools/tht/internal/safeio/files_test.go b/tools/tht/internal/safeio/files_test.go index 82170172..3bcd3efa 100644 --- a/tools/tht/internal/safeio/files_test.go +++ b/tools/tht/internal/safeio/files_test.go @@ -3,6 +3,7 @@ package safeio import ( "errors" "fmt" + "io" "os" "path/filepath" "strings" @@ -97,6 +98,57 @@ func TestWriteCanonicalNewFileRejectsExistingTargets(t *testing.T) { } } +func TestCreateCanonicalNewPrivateFileProvidesPrivateStreamingWriter(t *testing.T) { + temporaryRoot, err := filepath.EvalSymlinks(os.TempDir()) + if err != nil { + t.Fatal(err) + } + directory, err := os.MkdirTemp(temporaryRoot, "tht-safeio-private-stream-") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.RemoveAll(directory) }) + if err := EnsurePrivateDirectory(directory); err != nil { + t.Fatal(err) + } + + path := filepath.Join(directory, "archive.zip") + file, err := CreateCanonicalNewPrivateFile(path) + if err != nil { + t.Fatal(err) + } + if _, err := file.Write([]byte("streamed archive")); err != nil { + _ = file.Close() + t.Fatal(err) + } + if _, err := file.Seek(0, io.SeekStart); err != nil { + _ = file.Close() + t.Fatal(err) + } + streamed := make([]byte, len("streamed archive")) + if _, err := io.ReadFull(file, streamed); err != nil { + _ = file.Close() + t.Fatal(err) + } + if string(streamed) != "streamed archive" { + _ = file.Close() + t.Fatalf("streamed contents through open file = %q, want %q", streamed, "streamed archive") + } + if err := file.Close(); err != nil { + t.Fatal(err) + } + if err := ValidatePrivateRegular(path); err != nil { + t.Fatalf("ValidatePrivateRegular() = %v, want owner-private streamed file", err) + } + contents, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + if string(contents) != "streamed archive" { + t.Fatalf("streamed contents = %q, want %q", contents, "streamed archive") + } +} + func TestPreflightPrivateDirectoryAllowsOnlyAMissingFinalComponentWithoutMutation(t *testing.T) { temporaryRoot, err := filepath.EvalSymlinks(os.TempDir()) if err != nil { diff --git a/tools/tht/internal/safeio/private_unix.go b/tools/tht/internal/safeio/private_unix.go index 2e3d8769..322f1ce9 100644 --- a/tools/tht/internal/safeio/private_unix.go +++ b/tools/tht/internal/safeio/private_unix.go @@ -89,7 +89,11 @@ func ProtectPrivateRegular(path string) error { } func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, error) { - file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, mode) + return createCanonicalNewPrivateFileWithFlags(path, mode, os.O_WRONLY) +} + +func createCanonicalNewPrivateFileWithFlags(path string, mode os.FileMode, flags int) (*os.File, error) { + file, err := os.OpenFile(path, flags|os.O_CREATE|os.O_EXCL, mode) if err != nil { return nil, err } @@ -108,6 +112,13 @@ func createCanonicalNewPrivateParentFile(path string, mode os.FileMode) (*os.Fil return createCanonicalNewPrivateFile(path, mode) } +func createCanonicalNewPrivateParentReadWriteFile(path string, mode os.FileMode) (*os.File, error) { + if err := ValidatePrivateDirectory(filepath.Dir(path)); err != nil { + return nil, ErrUnsafeFile + } + return createCanonicalNewPrivateFileWithFlags(path, mode, os.O_RDWR) +} + // ValidatePrivateRegular requires a canonical, single-link private regular file. func ValidatePrivateRegular(path string) error { if err := ValidateCanonicalPath(path); err != nil { diff --git a/tools/tht/internal/safeio/private_windows.go b/tools/tht/internal/safeio/private_windows.go index dc916e83..18a30744 100644 --- a/tools/tht/internal/safeio/private_windows.go +++ b/tools/tht/internal/safeio/private_windows.go @@ -96,14 +96,18 @@ func ProtectPrivateRegular(path string) error { // createCanonicalNewPrivateFile installs the owner-only protected DACL in the CreateFile call, so // another mutation can never observe a newly-created lock with an inherited/default DACL. func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, error) { - return createCanonicalNewFile(path, mode, false) + return createCanonicalNewFile(path, mode, false, windows.GENERIC_WRITE) } func createCanonicalNewPrivateParentFile(path string, mode os.FileMode) (*os.File, error) { - return createCanonicalNewFile(path, mode, true) + return createCanonicalNewFile(path, mode, true, windows.GENERIC_WRITE) } -func createCanonicalNewFile(path string, mode os.FileMode, requirePrivateParent bool) (*os.File, error) { +func createCanonicalNewPrivateParentReadWriteFile(path string, mode os.FileMode) (*os.File, error) { + return createCanonicalNewFile(path, mode, true, windows.GENERIC_READ|windows.GENERIC_WRITE) +} + +func createCanonicalNewFile(path string, mode os.FileMode, requirePrivateParent bool, access uint32) (*os.File, error) { parents, target, err := openCanonicalWindowsParent(path) if err != nil || len(parents.handles) == 0 || (requirePrivateParent && validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil) { if parents != nil { @@ -123,7 +127,7 @@ func createCanonicalNewFile(path string, mode os.FileMode, requirePrivateParent } handle, err := windows.CreateFile( windows.StringToUTF16Ptr(filepath.Join(parents.directory, target)), - windows.GENERIC_WRITE, + access, windowsRetainedHandleShareMode, attributes, windows.CREATE_NEW,