Files

330 lines
6.7 KiB
Go

package files
import (
"errors"
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"time"
)
const DefaultMaxReadBytes = 1 << 20 // 1MB
var (
ErrPathTraversal = errors.New("path escapes data root")
ErrSymlink = errors.New("symlinks are not allowed")
ErrNotFound = errors.New("file not found")
ErrNotDirectory = errors.New("path is not a directory")
ErrIsDirectory = errors.New("path is a directory")
)
// DirEntry describes a file or directory entry returned by ListDir.
type DirEntry struct {
Name string
Path string
IsDir bool
Size int64
ModifiedAt time.Time
}
// ResolvePath resolves relativePath under dataRoot and rejects path traversal.
func ResolvePath(dataRoot, relativePath string) (string, error) {
root, err := filepath.Abs(filepath.Clean(dataRoot))
if err != nil {
return "", fmt.Errorf("resolve data root: %w", err)
}
rel := strings.TrimPrefix(relativePath, "/")
rel = strings.TrimPrefix(rel, `\`)
rel = filepath.Clean(rel)
if rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
return "", ErrPathTraversal
}
resolved := filepath.Join(root, rel)
resolved, err = filepath.Abs(resolved)
if err != nil {
return "", fmt.Errorf("resolve path: %w", err)
}
if resolved != root && !strings.HasPrefix(resolved, root+string(filepath.Separator)) {
return "", ErrPathTraversal
}
return resolved, nil
}
// ListDir lists entries in a directory relative to dataRoot.
func ListDir(dataRoot, relativePath string) ([]DirEntry, error) {
dirPath, err := ResolvePath(dataRoot, relativePath)
if err != nil {
return nil, err
}
if err := rejectSymlinks(dataRoot, dirPath); err != nil {
return nil, err
}
info, err := os.Lstat(dirPath)
if err != nil {
if os.IsNotExist(err) {
return nil, ErrNotFound
}
return nil, err
}
if !info.IsDir() {
return nil, ErrNotDirectory
}
entries, err := os.ReadDir(dirPath)
if err != nil {
return nil, err
}
root, err := filepath.Abs(filepath.Clean(dataRoot))
if err != nil {
return nil, err
}
result := make([]DirEntry, 0, len(entries))
for _, entry := range entries {
entryPath := filepath.Join(dirPath, entry.Name())
if err := rejectSymlinks(root, entryPath); err != nil {
continue
}
entryInfo, err := entry.Info()
if err != nil {
return nil, err
}
rel, err := filepath.Rel(root, entryPath)
if err != nil {
return nil, err
}
rel = filepath.ToSlash(rel)
result = append(result, DirEntry{
Name: entry.Name(),
Path: rel,
IsDir: entry.IsDir(),
Size: entryInfo.Size(),
ModifiedAt: entryInfo.ModTime().UTC(),
})
}
sort.Slice(result, func(i, j int) bool {
if result[i].IsDir != result[j].IsDir {
return result[i].IsDir
}
return strings.ToLower(result[i].Name) < strings.ToLower(result[j].Name)
})
return result, nil
}
// ReadFile reads file content relative to dataRoot up to maxBytes.
func ReadFile(dataRoot, relativePath string, maxBytes int) (content string, truncated bool, err error) {
if maxBytes <= 0 {
maxBytes = DefaultMaxReadBytes
}
filePath, err := ResolvePath(dataRoot, relativePath)
if err != nil {
return "", false, err
}
if err := rejectSymlinks(dataRoot, filePath); err != nil {
return "", false, err
}
info, err := os.Lstat(filePath)
if err != nil {
if os.IsNotExist(err) {
return "", false, ErrNotFound
}
return "", false, err
}
if info.IsDir() {
return "", false, ErrIsDirectory
}
if info.Mode()&os.ModeSymlink != 0 {
return "", false, ErrSymlink
}
f, err := os.Open(filePath)
if err != nil {
return "", false, err
}
defer f.Close()
limited := io.LimitReader(f, int64(maxBytes)+1)
data, err := io.ReadAll(limited)
if err != nil {
return "", false, err
}
if len(data) > maxBytes {
truncated = true
data = data[:maxBytes]
}
return string(data), truncated, nil
}
// WriteFile atomically writes content to a path relative to dataRoot.
func WriteFile(dataRoot, relativePath, content string, create bool) error {
filePath, err := ResolvePath(dataRoot, relativePath)
if err != nil {
return err
}
if _, err := os.Lstat(filePath); err != nil {
if os.IsNotExist(err) {
if !create {
return ErrNotFound
}
} else {
return err
}
} else if !create {
if err := rejectSymlinks(dataRoot, filePath); err != nil {
return err
}
}
dir := filepath.Dir(filePath)
if err := os.MkdirAll(dir, 0o755); err != nil {
return err
}
tmp, err := os.CreateTemp(dir, ".hexahost-write-*")
if err != nil {
return err
}
tmpPath := tmp.Name()
cleanup := func() {
_ = tmp.Close()
_ = os.Remove(tmpPath)
}
if _, err := tmp.WriteString(content); err != nil {
cleanup()
return err
}
if err := tmp.Sync(); err != nil {
cleanup()
return err
}
if err := tmp.Close(); err != nil {
cleanup()
return err
}
if err := os.Rename(tmpPath, filePath); err != nil {
cleanup()
return err
}
return nil
}
// RemovePath deletes a file or directory relative to dataRoot.
func RemovePath(dataRoot, relativePath string) error {
target, err := ResolvePath(dataRoot, relativePath)
if err != nil {
return err
}
if err := rejectSymlinks(dataRoot, target); err != nil {
return err
}
info, err := os.Lstat(target)
if err != nil {
if os.IsNotExist(err) {
return ErrNotFound
}
return err
}
if info.IsDir() {
return os.RemoveAll(target)
}
return os.Remove(target)
}
// DeleteFile removes a file relative to dataRoot.
func DeleteFile(dataRoot, relativePath string) error {
filePath, err := ResolvePath(dataRoot, relativePath)
if err != nil {
return err
}
if err := rejectSymlinks(dataRoot, filePath); err != nil {
return err
}
info, err := os.Lstat(filePath)
if err != nil {
if os.IsNotExist(err) {
return ErrNotFound
}
return err
}
if info.IsDir() {
return ErrIsDirectory
}
return os.Remove(filePath)
}
func rejectSymlinks(dataRoot, target string) error {
root, err := filepath.Abs(filepath.Clean(dataRoot))
if err != nil {
return err
}
target, err = filepath.Abs(target)
if err != nil {
return err
}
if target == root {
info, err := os.Lstat(target)
if err != nil {
if os.IsNotExist(err) {
return nil
}
return err
}
if info.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("%w: %s", ErrSymlink, target)
}
return nil
}
rel, err := filepath.Rel(root, target)
if err != nil {
return err
}
current := root
for _, part := range strings.Split(rel, string(filepath.Separator)) {
if part == "" || part == "." {
continue
}
current = filepath.Join(current, part)
info, err := os.Lstat(current)
if os.IsNotExist(err) {
return nil
}
if err != nil {
return err
}
if info.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("%w: %s", ErrSymlink, current)
}
}
return nil
}