Phase3
This commit is contained in:
280
apps/node-agent/internal/files/safe.go
Normal file
280
apps/node-agent/internal/files/safe.go
Normal file
@@ -0,0 +1,280 @@
|
||||
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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
94
apps/node-agent/internal/files/safe_test.go
Normal file
94
apps/node-agent/internal/files/safe_test.go
Normal file
@@ -0,0 +1,94 @@
|
||||
package files
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestResolvePath_BlocksTraversal(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
|
||||
tests := []string{
|
||||
"../outside",
|
||||
"../../etc/passwd",
|
||||
"foo/../../outside",
|
||||
}
|
||||
|
||||
for _, rel := range tests {
|
||||
t.Run(rel, func(t *testing.T) {
|
||||
_, err := ResolvePath(root, rel)
|
||||
if err == nil {
|
||||
t.Fatalf("expected path traversal error for %q", rel)
|
||||
}
|
||||
if err != ErrPathTraversal {
|
||||
t.Fatalf("expected ErrPathTraversal, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePath_AllowsNestedPaths(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
nested := filepath.Join(root, "world", "region")
|
||||
if err := os.MkdirAll(nested, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
resolved, err := ResolvePath(root, "world/region")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve nested path: %v", err)
|
||||
}
|
||||
if resolved != nested {
|
||||
t.Fatalf("resolved %q want %q", resolved, nested)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadFile_BlocksTraversal(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
secret := filepath.Join(filepath.Dir(root), "secret.txt")
|
||||
if err := os.WriteFile(secret, []byte("secret"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.Remove(secret) })
|
||||
|
||||
_, _, err := ReadFile(root, "../"+filepath.Base(secret))
|
||||
if err == nil {
|
||||
t.Fatal("expected path traversal error")
|
||||
}
|
||||
if err != ErrPathTraversal {
|
||||
t.Fatalf("expected ErrPathTraversal, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteFile_AtomicWrite(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
rel := "server.properties"
|
||||
|
||||
if err := WriteFile(root, rel, "motd=Hello", true); err != nil {
|
||||
t.Fatalf("write file: %v", err)
|
||||
}
|
||||
|
||||
content, truncated, err := ReadFile(root, rel, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("read file: %v", err)
|
||||
}
|
||||
if truncated {
|
||||
t.Fatal("did not expect truncated read")
|
||||
}
|
||||
if content != "motd=Hello" {
|
||||
t.Fatalf("content %q want %q", content, "motd=Hello")
|
||||
}
|
||||
}
|
||||
|
||||
func TestListDir_RejectsTraversal(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
|
||||
_, err := ListDir(root, "..")
|
||||
if err == nil {
|
||||
t.Fatal("expected path traversal error")
|
||||
}
|
||||
if err != ErrPathTraversal {
|
||||
t.Fatalf("expected ErrPathTraversal, got %v", err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user