Snapshot path validation (#1950)

Fixes https://github.com/agent-substrate/substrate/pull/1426 

Rebases that PR and adjusts to changes in main since.
https://github.com/agent-substrate/substrate/pull/1426#issuecomment-5804276427

Addresses outstanding review comments in additional commits.

I want to get this in M3. 

> It's a good idea to open an issue first for discussion.

- [x] Tests pass
- [x] Appropriate changes to documentation are included in the PR

---------

Co-authored-by: Du Bin <dubin555@gmail.com>
This commit is contained in:
Benjamin Elder
2026-09-28 20:34:38 +00:00
committed by GitHub
co-authored by Du Bin
parent 2e57c900a3
commit 5b7417e85d
4 changed files with 327 additions and 41 deletions
+27
View File
@@ -121,6 +121,18 @@ func SendLocalFileToGCSWithZstd(ctx context.Context, client ObjectStorage, gsURL
return nil
}
// SendFileToGCSWithZstd compresses and uploads an already-open file. The
// caller retains ownership of localFile.
func SendFileToGCSWithZstd(ctx context.Context, client ObjectStorage, gsURL string, localFile *os.File) error {
ctx, span := tracer.Start(ctx, "sendFileToGCSWithZstd")
defer span.End()
if err := sendZstd(ctx, client, gsURL, localFile); err != nil {
return fmt.Errorf("in sendZstd: %w", err)
}
return nil
}
// sparseFilePutter marks a backend that can upload a sparse FILE directly, splitting
// compression and upload by file range instead of piping one compressed stream into a
// part splitter. Reports errSparseTooSmall when the file is not worth splitting, in
@@ -329,6 +341,21 @@ func FetchLocalFileFromGCSWithZstd(ctx context.Context, client ObjectStorage, gs
return nil
}
// FetchFileFromGCSWithZstd downloads and decompresses into an already-open
// file. The caller retains ownership of localFile.
func FetchFileFromGCSWithZstd(ctx context.Context, client ObjectStorage, gsURL string, localFile *os.File) error {
ctx, span := tracer.Start(ctx, "fetchFileFromGCSWithZstd")
defer span.End()
if err := localFile.Chmod(0o600); err != nil {
return fmt.Errorf("in localFile.Chmod(0o600): %w", err)
}
if err := fetchFromGCSWithZstd(ctx, client, gsURL, localFile); err != nil {
return fmt.Errorf("while fetching %q from GCS: %w", gsURL, err)
}
return nil
}
func fetchFromGCSWithZstd(ctx context.Context, client ObjectStorage, gsURL string, out io.Writer) (err error) {
bucket, object, err := parseGCSURL(gsURL)
if err != nil {
+144 -37
View File
@@ -569,20 +569,11 @@ func initSnapshotSizeMetric() error {
// recordSnapshotSize labels each image with the registry's file.name. That
// label used to be spelled "kind", which means the snapshot's provenance
// everywhere else in the ate.* namespace, not one of its files.
func recordSnapshotSize(ctx context.Context, file, path, templateAtespace, templateName string) {
func recordSnapshotSize(ctx context.Context, file string, size int64, templateAtespace, templateName string) {
if snapshotSizeBytes == nil {
return
}
fi, err := os.Stat(path)
if errors.Is(err, os.ErrNotExist) {
return
}
if err != nil {
slog.WarnContext(ctx, "Failed to stat snapshot image for size metric",
slog.String("file", file), slog.String("path", path), slog.Any("err", err))
return
}
snapshotSizeBytes.Record(ctx, fi.Size(), metric.WithAttributes(
snapshotSizeBytes.Record(ctx, size, metric.WithAttributes(
semconv.FileNameKey.String(file),
ateattr.TemplateAtespaceKey.String(templateAtespace),
ateattr.TemplateNameKey.String(templateName),
@@ -669,9 +660,9 @@ func (s *AteomHerder) Checkpoint(ctx context.Context, req *ateletpb.CheckpointRe
s.systemInfoVolumes.Deregister(actorUID)
sandboxRec.SnapshotFiles = resp.GetSnapshotFiles()
if len(sandboxRec.SnapshotFiles) == 0 && shouldHaveSnapshots(req) {
return nil, fmt.Errorf("ateom reported no snapshot files for checkpoint")
sandboxRec.SnapshotFiles, err = checkpointSnapshotFiles(resp, shouldHaveSnapshots(req))
if err != nil {
return nil, err
}
sandboxRec.Atespace = req.GetAtespace()
sandboxRec.ActorName = req.GetActorName()
@@ -708,7 +699,7 @@ func (s *AteomHerder) Checkpoint(ctx context.Context, req *ateletpb.CheckpointRe
return nil, fmt.Errorf("while uploading external snapshot: %w", err)
}
case ateletpb.CheckpointType_CHECKPOINT_TYPE_LOCAL:
if err := s.moveLocalCheckpoint(ctx, req, checkpointDir, sandboxRec); err != nil {
if err := s.moveLocalCheckpoint(ctx, req, sandboxRec); err != nil {
dPersist = time.Since(tPersist)
return nil, fmt.Errorf("while moving to local snapshot: %w", err)
}
@@ -729,6 +720,17 @@ func (s *AteomHerder) Checkpoint(ctx context.Context, req *ateletpb.CheckpointRe
return &ateletpb.CheckpointResponse{}, nil
}
func checkpointSnapshotFiles(resp *ateompb.CheckpointWorkloadResponse, required bool) ([]string, error) {
files := resp.GetSnapshotFiles()
if len(files) == 0 && required {
return nil, errors.New("ateom reported no snapshot files for checkpoint")
}
if err := validateSnapshotFiles(files); err != nil {
return nil, fmt.Errorf("ateom reported invalid snapshot files: %w", err)
}
return files, nil
}
func toAteomSnapshotScope(scope ateletpb.SnapshotScope) ateompb.SnapshotScope {
// assumption the request already been validated and scope is in the valid values set
switch scope {
@@ -741,19 +743,40 @@ func toAteomSnapshotScope(scope ateletpb.SnapshotScope) ateompb.SnapshotScope {
}
}
func (s *AteomHerder) moveLocalCheckpoint(ctx context.Context, req *ateletpb.CheckpointRequest, checkpointDir string, rec *sandboxAssetsRecord) error {
localCheckpointPath := ateletpath.LocalSnapshotDir(req.GetActorUid(), req.GetLocalConfig().GetSnapshotName())
if err := os.MkdirAll(localCheckpointPath, 0o700); err != nil {
func (s *AteomHerder) moveLocalCheckpoint(ctx context.Context, req *ateletpb.CheckpointRequest, rec *sandboxAssetsRecord) error {
actorDir := ateletpath.ActorPath(req.GetActorUid())
root, err := os.OpenRoot(actorDir)
if err != nil {
return fmt.Errorf("while opening actor directory: %w", err)
}
defer root.Close()
checkpointDir, err := filepath.Rel(actorDir, ateletpath.CheckpointStateDir(req.GetActorUid()))
if err != nil {
return err
}
localDir, err := filepath.Rel(actorDir, ateletpath.LocalSnapshotDir(req.GetActorUid(), req.GetLocalConfig().GetSnapshotName()))
if err != nil {
return err
}
if err := root.MkdirAll(localDir, 0o700); err != nil {
return fmt.Errorf("while creating local checkpoint directory: %w", err)
}
// Move exactly the files ateom reported.
for _, fileName := range rec.SnapshotFiles {
src := filepath.Join(checkpointDir, fileName)
dst := filepath.Join(localCheckpointPath, fileName)
recordSnapshotSize(ctx, fileName, src, req.GetActorTemplateAtespace(), req.GetActorTemplateName())
dst := filepath.Join(localDir, fileName)
info, err := root.Lstat(src)
if err != nil {
return fmt.Errorf("while inspecting checkpoint file %s: %w", fileName, err)
}
if !info.Mode().IsRegular() {
return fmt.Errorf("checkpoint file %s is not a regular file", fileName)
}
recordSnapshotSize(ctx, fileName, info.Size(), req.GetActorTemplateAtespace(), req.GetActorTemplateName())
if err := os.Rename(src, dst); err != nil {
if err := root.Rename(src, dst); err != nil {
return fmt.Errorf("failed to move %s to %s: %w", src, dst, err)
}
}
@@ -763,7 +786,7 @@ func (s *AteomHerder) moveLocalCheckpoint(ctx context.Context, req *ateletpb.Che
if err != nil {
return fmt.Errorf("while marshaling snapshot manifest: %w", err)
}
if err := os.WriteFile(filepath.Join(localCheckpointPath, sandboxManifestName), manifest, 0o600); err != nil {
if err := root.WriteFile(filepath.Join(localDir, sandboxManifestName), manifest, 0o600); err != nil {
return fmt.Errorf("while writing snapshot manifest: %w", err)
}
@@ -799,16 +822,34 @@ func (s *AteomHerder) uploadExternalCheckpoint(ctx context.Context, req *ateletp
// leaves only orphaned files, never a manifest pointing at files that never
// landed; retries overwrite the deterministic object names.
func (s *AteomHerder) uploadSnapshot(ctx context.Context, uri resources.SnapshotURI, srcDir string, rec *sandboxAssetsRecord, templateAtespace, templateName string) error {
root, err := os.OpenRoot(srcDir)
if err != nil {
return fmt.Errorf("while opening snapshot directory: %w", err)
}
defer root.Close()
g, gCtx := errgroup.WithContext(ctx)
for _, fileName := range rec.SnapshotFiles {
local := filepath.Join(srcDir, fileName)
recordSnapshotSize(ctx, fileName, local, templateAtespace, templateName)
g.Go(func() error {
local, err := root.Open(fileName)
if err != nil {
return fmt.Errorf("while opening %s in snapshot directory: %w", fileName, err)
}
defer local.Close()
info, err := local.Stat()
if err != nil {
return fmt.Errorf("while inspecting %s in snapshot directory: %w", fileName, err)
}
if !info.Mode().IsRegular() {
return fmt.Errorf("snapshot file %s is not a regular file", fileName)
}
recordSnapshotSize(ctx, fileName, info.Size(), templateAtespace, templateName)
objectURI, err := uri.ObjectURI(fileName + ".zstd")
if err != nil {
return fmt.Errorf("while addressing %s in GCS: %w", fileName, err)
}
if err := ategcs.SendLocalFileToGCSWithZstd(gCtx, s.gcsClient, objectURI, local); err != nil {
if err := ategcs.SendFileToGCSWithZstd(gCtx, s.gcsClient, objectURI, local); err != nil {
return fmt.Errorf("while uploading %s to GCS: %w", fileName, err)
}
return nil
@@ -890,7 +931,7 @@ func (s *AteomHerder) uploadLocalCheckpointDir(ctx context.Context, req *ateletp
return "", fmt.Errorf("while addressing snapshot manifest in GCS: %w", err)
}
manifest, err := os.ReadFile(filepath.Join(localDir, sandboxManifestName))
manifest, err := readSnapshotManifest(localDir)
if errors.Is(err, os.ErrNotExist) {
// The local snapshot is gone. A previous invocation may have uploaded
// and pruned it: the remote manifest is uploaded last, so its presence
@@ -938,6 +979,15 @@ func (s *AteomHerder) uploadLocalCheckpointDir(ctx context.Context, req *ateletp
return rec.SandboxClass, s.uploadSnapshot(ctx, uri, localDir, rec, req.GetActorTemplateAtespace(), req.GetActorTemplateName())
}
func readSnapshotManifest(dir string) ([]byte, error) {
root, err := os.OpenRoot(dir)
if err != nil {
return nil, err
}
defer root.Close()
return root.ReadFile(sandboxManifestName)
}
// narrowFullCaptureToData rewrites rec so a FULL capture uploads as a DATA
// snapshot. Each sandbox class owns one branch: micro-VM durable data is a
// self-contained tar that can be carved out of the full file set; gVisor's
@@ -1057,7 +1107,7 @@ func (s *AteomHerder) Restore(ctx context.Context, req *ateletpb.RestoreRequest)
return nil, fmt.Errorf("while unmarshalling sandbox record: %w", err)
}
case ateletpb.CheckpointType_CHECKPOINT_TYPE_LOCAL:
manifest, err := os.ReadFile(filepath.Join(ateletpath.LocalSnapshotDir(actorUID, req.GetLocalConfig().GetSnapshotName()), sandboxManifestName))
manifest, err := readSnapshotManifest(ateletpath.LocalSnapshotDir(actorUID, req.GetLocalConfig().GetSnapshotName()))
if err != nil {
return nil, wrapFileSystemErr("while reading local snapshot manifest", err)
}
@@ -1147,7 +1197,7 @@ func (s *AteomHerder) Restore(ctx context.Context, req *ateletpb.RestoreRequest)
// the golden's from object storage, concurrently.
gLocal, gLocalCtx := errgroup.WithContext(gctx)
gLocal.Go(func() error {
if err := s.copyLocalCheckpoint(gLocalCtx, req.GetLocalConfig().GetSnapshotName(), ateletpath.LocalCheckpointsDir(actorUID), checkpointDir, sandboxRec.SnapshotFiles); err != nil {
if err := s.copyLocalCheckpoint(gLocalCtx, ateletpath.ActorPath(actorUID), req.GetLocalConfig().GetSnapshotName(), ateletpath.LocalCheckpointsDir(actorUID), checkpointDir, sandboxRec.SnapshotFiles); err != nil {
return err
}
return nil
@@ -1321,13 +1371,37 @@ func (s *AteomHerder) Terminate(ctx context.Context, req *ateletpb.TerminateRequ
return &ateletpb.TerminateResponse{}, nil
}
func (s *AteomHerder) copyLocalCheckpoint(ctx context.Context, snapshotName string, srcDir, dstDir string, files []string) error {
// copyLocalCheckpoint stages files from the local checkpoint snapshotName under
// srcDir into dstDir. Both must be inside actorDir, which confines every access.
func (s *AteomHerder) copyLocalCheckpoint(ctx context.Context, actorDir, snapshotName, srcDir, dstDir string, files []string) error {
root, err := os.OpenRoot(actorDir)
if err != nil {
return fmt.Errorf("while opening actor directory: %w", err)
}
defer root.Close()
srcDir, err = filepath.Rel(actorDir, filepath.Join(srcDir, snapshotName))
if err != nil {
return err
}
dstDir, err = filepath.Rel(actorDir, dstDir)
if err != nil {
return err
}
for _, fileName := range files {
if ctx.Err() != nil {
return fmt.Errorf("context cancelled: %w", ctx.Err())
}
src := filepath.Join(srcDir, snapshotName, fileName)
src := filepath.Join(srcDir, fileName)
dst := filepath.Join(dstDir, fileName)
// A link to a symlink would be followed later, outside the root.
info, err := root.Lstat(src)
if err != nil {
return fmt.Errorf("while inspecting %s: %w", src, err)
}
if !info.Mode().IsRegular() {
return fmt.Errorf("%s is not a regular file", src)
}
// Link rather than copy. The local checkpoint lives under the same actor dir
// as the restore staging area, so this stages the memory image in constant
// time instead of re-writing its whole working set. Nothing rewrites the
@@ -1338,9 +1412,9 @@ func (s *AteomHerder) copyLocalCheckpoint(ctx context.Context, snapshotName stri
//
// EXDEV alone falls back to copying, so an unexpected link failure surfaces
// instead of silently reverting to the full copy this exists to remove. It
// also keeps sparsefile.CopyFile off a dst that is already a link to src, where its
// also keeps copyRootFile off a dst that is already a link to src, where its
// O_TRUNC would empty both and report a successful copy of the old size.
switch err := linkFile(src, dst); {
switch err := linkFile(root, src, dst); {
case err == nil:
continue
case !errors.Is(err, unix.EXDEV):
@@ -1348,7 +1422,7 @@ func (s *AteomHerder) copyLocalCheckpoint(ctx context.Context, snapshotName stri
}
slog.WarnContext(ctx, "local checkpoint and restore dir are on different filesystems; copying instead of linking",
slog.String("src", src), slog.String("dst", dst))
if _, err := sparsefile.CopyFile(src, dst); err != nil {
if _, err := copyRootFile(root, src, dst); err != nil {
return fmt.Errorf("failed to copy %s to %s: %w", src, dst, err)
}
}
@@ -1356,9 +1430,31 @@ func (s *AteomHerder) copyLocalCheckpoint(ctx context.Context, snapshotName stri
return nil
}
// linkFile is os.Link, indirected so a test can force the cross-filesystem
// linkFile is os.Root.Link, indirected so a test can force the cross-filesystem
// fallback in copyLocalCheckpoint without mounting a second filesystem.
var linkFile = os.Link
var linkFile = (*os.Root).Link
func copyRootFile(root *os.Root, src, dst string) (int64, error) {
source, err := root.Open(src)
if err != nil {
return 0, err
}
defer source.Close()
info, err := source.Stat()
if err != nil {
return 0, err
}
if !info.Mode().IsRegular() {
return 0, fmt.Errorf("%s is not a regular file", src)
}
destination, err := root.OpenFile(dst, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o666)
if err != nil {
return 0, err
}
nBytes, err := sparsefile.Copy(source, destination)
return nBytes, errors.Join(err, destination.Close())
}
// goldenOnlyFiles returns the golden snapshot files not shadowed by the
// actor's own snapshot: on a DATA_ON_GOLDEN restore the actor's files (the
@@ -1398,16 +1494,27 @@ func (s *AteomHerder) downloadExternalCheckpoint(ctx context.Context, snapshotUR
if err != nil {
return err
}
root, err := os.OpenRoot(dstDir)
if err != nil {
return fmt.Errorf("while opening restore directory: %w", err)
}
defer root.Close()
g, gCtx := errgroup.WithContext(ctx)
for _, fileName := range files {
fileName := fileName
local := filepath.Join(dstDir, fileName)
g.Go(func() error {
objectURI, err := uri.ObjectURI(fileName + ".zstd")
if err != nil {
return fmt.Errorf("while addressing %s in GCS: %w", fileName, err)
}
if err := ategcs.FetchLocalFileFromGCSWithZstd(gCtx, s.gcsClient, objectURI, local); err != nil {
local, err := root.OpenFile(fileName, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o600)
if err != nil {
return fmt.Errorf("while opening %s in restore directory: %w", fileName, err)
}
fetchErr := ategcs.FetchFileFromGCSWithZstd(gCtx, s.gcsClient, objectURI, local)
closeErr := local.Close()
if err := errors.Join(fetchErr, closeErr); err != nil {
return fmt.Errorf("while downloading %s from GCS: %w", fileName, err)
}
return nil
+133 -4
View File
@@ -116,7 +116,7 @@ func TestCopyLocalCheckpointLinks(t *testing.T) {
t.Run("links when it can", func(t *testing.T) {
srcDir, dstDir := newDirs(t)
s := &AteomHerder{}
if err := s.copyLocalCheckpoint(context.Background(), snapshot, srcDir, dstDir, []string{"memory-ranges"}); err != nil {
if err := s.copyLocalCheckpoint(context.Background(), filepath.Dir(srcDir), snapshot, srcDir, dstDir, []string{"memory-ranges"}); err != nil {
t.Fatalf("copyLocalCheckpoint: %v", err)
}
src := filepath.Join(srcDir, snapshot, "memory-ranges")
@@ -133,11 +133,11 @@ func TestCopyLocalCheckpointLinks(t *testing.T) {
srcDir, dstDir := newDirs(t)
// EXDEV stands in for the mount boundary a unit test cannot produce.
orig := linkFile
linkFile = func(string, string) error { return unix.EXDEV }
linkFile = func(*os.Root, string, string) error { return unix.EXDEV }
t.Cleanup(func() { linkFile = orig })
s := &AteomHerder{}
if err := s.copyLocalCheckpoint(context.Background(), snapshot, srcDir, dstDir, []string{"memory-ranges"}); err != nil {
if err := s.copyLocalCheckpoint(context.Background(), filepath.Dir(srcDir), snapshot, srcDir, dstDir, []string{"memory-ranges"}); err != nil {
t.Fatalf("copyLocalCheckpoint: %v", err)
}
dst := filepath.Join(dstDir, "memory-ranges")
@@ -160,7 +160,7 @@ func TestCopyLocalCheckpointLinks(t *testing.T) {
}
s := &AteomHerder{}
// dst already exists, so os.Link fails with EEXIST.
if err := s.copyLocalCheckpoint(context.Background(), snapshot, srcDir, dstDir, []string{"memory-ranges"}); err == nil {
if err := s.copyLocalCheckpoint(context.Background(), filepath.Dir(srcDir), snapshot, srcDir, dstDir, []string{"memory-ranges"}); err == nil {
t.Fatal("copyLocalCheckpoint accepted a non-EXDEV link failure, want an error")
}
if got, err := os.ReadFile(dst); err != nil || !bytes.Equal(got, want) {
@@ -223,6 +223,135 @@ func TestSnapshotManifestRequiresPauseImage(t *testing.T) {
}
}
func TestSnapshotManifestRejectsNonLocalFile(t *testing.T) {
manifest, err := json.Marshal(sandboxAssetsRecord{
SandboxClass: "gvisor",
PauseImage: testPauseImage,
SnapshotFiles: []string{"../outside"},
})
if err != nil {
t.Fatal(err)
}
if _, err := unmarshalSandboxRecord(manifest); err == nil {
t.Fatal("unmarshalSandboxRecord() accepted a path outside the checkpoint directory")
}
}
func TestCheckpointSnapshotFiles(t *testing.T) {
for _, tc := range []struct {
name string
files []string
required bool
wantErr bool
}{
{name: "required and present", files: []string{"checkpoint.img"}, required: true},
{name: "optional and empty", required: false},
{name: "required and empty", required: true, wantErr: true},
{name: "escapes the directory", files: []string{"../outside"}, required: true, wantErr: true},
{name: "nested", files: []string{"a/b"}, required: true, wantErr: true},
{name: "dot", files: []string{"."}, required: true, wantErr: true},
{name: "unclean alias", files: []string{"checkpoint.img", "./checkpoint.img"}, required: true, wantErr: true},
{name: "duplicate", files: []string{"checkpoint.img", "checkpoint.img"}, required: true, wantErr: true},
{name: "manifest name", files: []string{"checkpoint.img", sandboxManifestName}, required: true, wantErr: true},
} {
t.Run(tc.name, func(t *testing.T) {
files, err := checkpointSnapshotFiles(&ateompb.CheckpointWorkloadResponse{SnapshotFiles: tc.files}, tc.required)
if (err != nil) != tc.wantErr {
t.Fatalf("checkpointSnapshotFiles() = %v, %v; wantErr %v", files, err, tc.wantErr)
}
})
}
}
func TestUploadSnapshotRejectsSymlinkOutsideRoot(t *testing.T) {
parent := t.TempDir()
checkpointDir := filepath.Join(parent, "checkpoint-state")
if err := os.Mkdir(checkpointDir, 0o700); err != nil {
t.Fatal(err)
}
outside := filepath.Join(parent, "outside")
if err := os.WriteFile(outside, []byte("secret"), 0o600); err != nil {
t.Fatal(err)
}
if err := os.Symlink(outside, filepath.Join(checkpointDir, "checkpoint.img")); err != nil {
t.Fatal(err)
}
uri, err := resources.ParseSnapshotURI(testSnapshotURI)
if err != nil {
t.Fatal(err)
}
store := &recordingObjectStorage{}
err = (&AteomHerder{gcsClient: store}).uploadSnapshot(context.Background(), uri, checkpointDir,
&sandboxAssetsRecord{SnapshotFiles: []string{"checkpoint.img"}}, "test", "test")
if err == nil {
t.Fatal("uploadSnapshot() followed a symlink outside the checkpoint directory")
}
if got := store.keys(); len(got) != 0 {
t.Fatalf("uploaded objects = %v, want none", got)
}
}
func TestDownloadExternalCheckpointRejectsSymlinkOutsideRoot(t *testing.T) {
parent := t.TempDir()
restoreDir := filepath.Join(parent, "restore-state")
if err := os.Mkdir(restoreDir, 0o700); err != nil {
t.Fatal(err)
}
outside := filepath.Join(parent, "outside")
if err := os.WriteFile(outside, []byte("keep me"), 0o600); err != nil {
t.Fatal(err)
}
if err := os.Symlink(outside, filepath.Join(restoreDir, "checkpoint.img")); err != nil {
t.Fatal(err)
}
store := &recordingObjectStorage{}
payload := filepath.Join(parent, "payload")
if err := os.WriteFile(payload, []byte("replacement"), 0o600); err != nil {
t.Fatal(err)
}
if err := ategcs.SendLocalFileToGCSWithZstd(context.Background(), store, testSnapshotURI+"/checkpoint.img.zstd", payload); err != nil {
t.Fatal(err)
}
err := (&AteomHerder{gcsClient: store}).downloadExternalCheckpoint(
context.Background(), testSnapshotURI, restoreDir, []string{"checkpoint.img"})
if err == nil {
t.Fatal("downloadExternalCheckpoint() followed a symlink outside the restore directory")
}
if got, err := os.ReadFile(outside); err != nil || string(got) != "keep me" {
t.Fatalf("outside file = %q, %v; want unchanged", got, err)
}
}
func TestCopyLocalCheckpointRejectsSymlinkOutsideRoot(t *testing.T) {
parent := t.TempDir()
snapshotName := "pause-1"
snapshotDir := filepath.Join(parent, "local-checkpoint", snapshotName)
if err := os.MkdirAll(snapshotDir, 0o700); err != nil {
t.Fatal(err)
}
outside := filepath.Join(parent, "outside")
if err := os.WriteFile(outside, []byte("secret"), 0o600); err != nil {
t.Fatal(err)
}
if err := os.Symlink(outside, filepath.Join(snapshotDir, "checkpoint.img")); err != nil {
t.Fatal(err)
}
restoreDir := filepath.Join(parent, "restore-state")
if err := os.Mkdir(restoreDir, 0o700); err != nil {
t.Fatal(err)
}
err := (&AteomHerder{}).copyLocalCheckpoint(context.Background(), parent, snapshotName,
filepath.Join(parent, "local-checkpoint"), restoreDir, []string{"checkpoint.img"})
if err == nil {
t.Fatal("copyLocalCheckpoint() followed a symlink outside the local checkpoint directory")
}
if _, err := os.Stat(filepath.Join(restoreDir, "checkpoint.img")); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("restore file exists after rejected copy: %v", err)
}
}
func TestWriteFileAtomic(t *testing.T) {
dir := t.TempDir()
target := filepath.Join(dir, "actor-id")
+23
View File
@@ -490,9 +490,32 @@ func unmarshalSandboxRecord(data []byte) (*sandboxAssetsRecord, error) {
if rec.PauseImage == "" {
return nil, fmt.Errorf("sandbox record/manifest has no pauseImage")
}
if err := validateSnapshotFiles(rec.SnapshotFiles); err != nil {
return nil, fmt.Errorf("sandbox record/manifest has invalid snapshotFiles: %w", err)
}
return rec, nil
}
// validateSnapshotFiles requires each name to be a distinct plain file name in
// the checkpoint directory, other than the manifest atelet writes beside them.
// Actual file access must still use os.Root so symlinks cannot escape that
// directory.
func validateSnapshotFiles(files []string) error {
seen := make(map[string]bool, len(files))
for i, name := range files {
switch {
case name != filepath.Base(name) || !filepath.IsLocal(name) || name == ".":
return fmt.Errorf("snapshotFiles[%d] %q is not a file name in the checkpoint directory", i, name)
case name == sandboxManifestName:
return fmt.Errorf("snapshotFiles[%d] %q is reserved for the snapshot manifest", i, name)
case seen[name]:
return fmt.Errorf("snapshotFiles[%d] %q is duplicated", i, name)
}
seen[name] = true
}
return nil
}
func wrapFileSystemErr(msg string, err error) error {
return fmt.Errorf("%s: %w", msg, err)
}