Files
ThothII/tools/tht/internal/safeio/replace_windows.go
T

88 lines
2.1 KiB
Go

//go:build windows
package safeio
import (
"errors"
"io"
"os"
"path/filepath"
"golang.org/x/sys/windows"
)
const windowsReplaceMoveFlags = windows.MOVEFILE_REPLACE_EXISTING | windows.MOVEFILE_WRITE_THROUGH
func replaceCanonicalRegular(path string, contents []byte) error {
parents, target, err := openCanonicalWindowsParent(path)
if err != nil {
return ErrUnsafeFile
}
defer parents.Close()
if len(parents.handles) == 0 || validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil || !safeExistingRegular(filepath.Join(parents.directory, target)) {
return ErrUnsafeFile
}
temporary, err := writePrivateTemporary(parents.directory, contents)
if err != nil {
return ErrUnsafeFile
}
defer func() { _ = os.Remove(temporary) }()
if !safeExistingRegular(filepath.Join(parents.directory, target)) {
return ErrUnsafeFile
}
from, err := windows.UTF16PtrFromString(temporary)
if err != nil {
return ErrUnsafeFile
}
to, err := windows.UTF16PtrFromString(filepath.Join(parents.directory, target))
if err != nil {
return ErrUnsafeFile
}
if err := windows.MoveFileEx(from, to, windowsReplaceMoveFlags); err != nil {
return ErrUnsafeFile
}
if err := ValidatePrivateRegular(filepath.Join(parents.directory, target)); err != nil {
return ErrUnsafeFile
}
return nil
}
func safeExistingRegular(path string) bool {
return ValidatePrivateRegular(path) == nil
}
func writePrivateTemporary(directory string, contents []byte) (string, error) {
for attempt := 0; attempt < 16; attempt++ {
name, err := randomTemporaryName()
if err != nil {
return "", err
}
path := filepath.Join(directory, name)
file, err := createCanonicalNewPrivateFile(path, 0o600)
if errors.Is(err, os.ErrExist) {
continue
}
if err != nil {
return "", err
}
var written int
written, err = file.Write(contents)
if err == nil && written != len(contents) {
err = io.ErrShortWrite
}
if err == nil {
err = file.Sync()
}
closeErr := file.Close()
if err == nil {
err = closeErr
}
if err != nil {
_ = os.Remove(path)
return "", err
}
return path, nil
}
return "", ErrUnsafeFile
}