mirror of
https://github.com/agent-substrate/substrate.git
synced 2026-10-02 03:24:42 +08:00
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:
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user