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

264 lines
8.8 KiB
Go

// Package safeio reads and writes local files without following symlinked path components.
package safeio
import (
"crypto/rand"
"encoding/hex"
"errors"
"io"
"os"
"path/filepath"
"sort"
"strings"
"unicode/utf8"
)
var ErrUnsafeFile = errors.New("unsafe file")
// EnsurePrivateDirectory creates only the final canonical directory with the platform's
// owner-only protection, or validates an existing directory has that protection.
func EnsurePrivateDirectory(path string) error {
if err := ValidateCanonicalPath(path); err != nil {
return err
}
if err := requireCanonicalDirectory(filepath.Dir(path)); err != nil {
return err
}
if err := createPrivateDirectory(path); err != nil && !errors.Is(err, os.ErrExist) {
return ErrUnsafeFile
}
return ValidatePrivateDirectory(path)
}
// ValidateCanonicalPath rejects relative or lexically non-canonical paths before they are opened.
func ValidateCanonicalPath(path string) error {
if !filepath.IsAbs(path) || filepath.Clean(path) != path || strings.Contains(path, string(filepath.Separator)+".."+string(filepath.Separator)) {
return ErrUnsafeFile
}
return nil
}
func readBoundedRegularFile(path string, file *os.File, maximum int64) ([]byte, error) {
if maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, ErrUnsafeFile
}
info, err := file.Stat()
if err != nil || !info.Mode().IsRegular() || !hasSingleLink(info) {
return nil, ErrUnsafeFile
}
contents, err := io.ReadAll(io.LimitReader(file, maximum+1))
if err != nil || int64(len(contents)) > maximum {
return nil, ErrUnsafeFile
}
after, err := file.Stat()
if err != nil || !after.Mode().IsRegular() || !hasSingleLink(after) || !os.SameFile(info, after) {
return nil, ErrUnsafeFile
}
current, err := os.Lstat(path)
if err != nil || !current.Mode().IsRegular() || current.Mode()&os.ModeSymlink != 0 || !hasSingleLink(current) || !os.SameFile(info, current) {
return nil, ErrUnsafeFile
}
return contents, nil
}
func ReadCanonicalUTF8(path string, maximum int64) (string, error) {
contents, err := ReadCanonicalRegular(path, maximum)
if err != nil {
return "", err
}
if !utf8.Valid(contents) {
return "", ErrUnsafeFile
}
return string(contents), nil
}
// ReadCanonicalPrivateRegular reads one owner-only private record and revalidates its metadata
// after the bounded read. It intentionally rejects ordinary hard links; OIDC's explicitly named
// atomic claim pair uses ReadCanonicalPrivateClaim instead.
func ReadCanonicalPrivateRegular(path string, maximum int64) ([]byte, error) {
return readCanonicalPrivateRegular(path, maximum)
}
// PrivateDirectoryEntry is a bounded, untrusted directory listing item. Callers must still
// validate each filename and record before using it.
type PrivateDirectoryEntry struct {
Name string `json:"name"`
ModifiedUnixMs int64 `json:"modifiedUnixMs"`
}
// ListCanonicalPrivateDirectory lists regular, non-symlinked direct children from an owner-only
// directory. It returns no content and bounds the number of entries before allocating output.
func ListCanonicalPrivateDirectory(path string, maximumEntries int) ([]PrivateDirectoryEntry, error) {
if maximumEntries < 1 || maximumEntries > 4096 {
return nil, ErrUnsafeFile
}
if err := ValidatePrivateDirectory(path); err != nil {
return nil, ErrUnsafeFile
}
directory, err := os.Open(path)
if err != nil {
return nil, ErrUnsafeFile
}
defer directory.Close()
entries, err := directory.ReadDir(maximumEntries + 1)
if err != nil && !errors.Is(err, io.EOF) {
return nil, ErrUnsafeFile
}
if len(entries) > maximumEntries {
return nil, ErrUnsafeFile
}
result := make([]PrivateDirectoryEntry, 0, len(entries))
for _, entry := range entries {
if entry.Name() == "" || strings.Contains(entry.Name(), string(filepath.Separator)) {
return nil, ErrUnsafeFile
}
info, err := os.Lstat(filepath.Join(path, entry.Name()))
if err != nil || !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 {
return nil, ErrUnsafeFile
}
result = append(result, PrivateDirectoryEntry{Name: entry.Name(), ModifiedUnixMs: info.ModTime().UnixMilli()})
}
sort.Slice(result, func(left, right int) bool { return result[left].Name < result[right].Name })
if err := ValidatePrivateDirectory(path); err != nil {
return nil, ErrUnsafeFile
}
return result, nil
}
func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error {
return writeCanonicalNewFile(path, contents, mode, false)
}
// WriteCanonicalNewPrivateFile is the authentication-storage variant of exclusive file
// creation. It requires the final parent directory to already have platform-specific owner-only
// protection and preserves that check while the platform primitive opens the parent.
func WriteCanonicalNewPrivateFile(path string, contents []byte, mode os.FileMode) error {
return writeCanonicalNewFile(path, contents, mode, true)
}
func writeCanonicalNewFile(path string, contents []byte, mode os.FileMode, 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
}
var (
file *os.File
err error
)
if requirePrivateParent {
file, err = createCanonicalNewPrivateParentFile(path, mode)
} else {
file, err = createCanonicalNewPrivateFile(path, mode)
}
if err != nil {
return ErrUnsafeFile
}
if _, err := file.Write(contents); err != nil {
_ = file.Close()
_ = os.Remove(path)
return ErrUnsafeFile
}
if err := file.Sync(); err != nil {
_ = file.Close()
_ = os.Remove(path)
return ErrUnsafeFile
}
if err := file.Close(); err != nil {
_ = os.Remove(path)
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.
func ReplaceCanonicalRegular(path string, contents []byte, mode os.FileMode) error {
if err := ValidateCanonicalPath(path); err != nil || mode.Perm() != 0o600 || mode&os.ModeType != 0 {
return ErrUnsafeFile
}
return replaceCanonicalRegular(path, contents)
}
// RemoveCanonicalPrivateRegular removes one existing private regular file without following a
// symlinked path component. It is intended only for rolling back a file this process published.
func RemoveCanonicalPrivateRegular(path string) error {
if err := ValidatePrivateRegular(path); err != nil {
return ErrUnsafeFile
}
return removeCanonicalPrivateRegular(path)
}
// ClaimCanonicalPrivateRegular atomically creates a second, explicit private hard link to one
// existing record. It is used only for digest-named, single-use OIDC state claims.
func ClaimCanonicalPrivateRegular(source, claim string) (bool, error) {
if err := validateClaimPaths(source, claim); err != nil {
return false, ErrUnsafeFile
}
return claimCanonicalPrivateRegular(source, claim)
}
// ReadCanonicalPrivateClaim reads a verified two-link source/claim pair. found=false means the
// state has already been consumed or a winning process is between its two removal steps.
func ReadCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) {
if maximum < 0 || maximum == int64(^uint64(0)>>1) || validateClaimPaths(source, claim) != nil {
return nil, false, ErrUnsafeFile
}
return readCanonicalPrivateClaim(source, claim, maximum)
}
// RemoveCanonicalPrivateClaim removes exactly a verified two-link source/claim pair.
func RemoveCanonicalPrivateClaim(source, claim string) (bool, error) {
if err := validateClaimPaths(source, claim); err != nil {
return false, ErrUnsafeFile
}
return removeCanonicalPrivateClaim(source, claim)
}
func validateClaimPaths(source, claim string) error {
if ValidateCanonicalPath(source) != nil || ValidateCanonicalPath(claim) != nil || filepath.Dir(source) != filepath.Dir(claim) {
return ErrUnsafeFile
}
if err := ValidatePrivateDirectory(filepath.Dir(source)); err != nil {
return ErrUnsafeFile
}
return nil
}
func randomTemporaryName() (string, error) {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
return "", err
}
return ".tht-auth-" + hex.EncodeToString(bytes) + ".tmp", nil
}
func requireCanonicalDirectory(path string) error {
if err := ValidateCanonicalPath(path); err != nil {
return err
}
resolved, err := filepath.EvalSymlinks(path)
if err != nil || resolved != path {
return ErrUnsafeFile
}
info, err := os.Stat(path)
if err != nil || !info.IsDir() {
return ErrUnsafeFile
}
return nil
}