103 lines
2.5 KiB
Go
103 lines
2.5 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 removeCanonicalPrivateRegular(path string) error {
|
|
parents, target, err := openCanonicalWindowsParent(path)
|
|
if err != nil {
|
|
return ErrUnsafeFile
|
|
}
|
|
defer parents.Close()
|
|
if err := ValidatePrivateRegular(path); err != nil {
|
|
return ErrUnsafeFile
|
|
}
|
|
if err := windows.DeleteFile(windows.StringToUTF16Ptr(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
|
|
}
|