fix(restore): share retained staging capability

This commit is contained in:
2026-08-18 14:30:33 +02:00
parent 2d1670e390
commit 6474118ec3
3 changed files with 134 additions and 14 deletions
+87 -11
View File
@@ -19,6 +19,7 @@ import (
"path/filepath" "path/filepath"
"regexp" "regexp"
"strings" "strings"
"sync"
"github.com/aritmolab/thothii/tools/tht/internal/config" "github.com/aritmolab/thothii/tools/tht/internal/config"
"github.com/aritmolab/thothii/tools/tht/internal/safeio" "github.com/aritmolab/thothii/tools/tht/internal/safeio"
@@ -108,10 +109,51 @@ type verifiedArchive struct {
type stagedArchive struct { type stagedArchive struct {
file *os.File file *os.File
parent safeio.PrivateDirectoryHandle parent safeio.PrivateDirectoryHandle
lease *stagingRootLease
name string name string
path string path string
} }
type stagingRootLease struct {
mu sync.Mutex
root string
parent safeio.PrivateDirectoryHandle
references int
}
func (lease *stagingRootLease) retain(root string) (safeio.PrivateDirectoryHandle, error) {
if lease == nil {
return nil, errors.New("private restore staging root is unavailable")
}
lease.mu.Lock()
defer lease.mu.Unlock()
if lease.parent == nil || lease.root != root || lease.parent.Validate() != nil {
return nil, errors.New("private restore staging root is unavailable")
}
lease.references++
return lease.parent, nil
}
func (lease *stagingRootLease) release() error {
if lease == nil {
return errors.New("private restore staging root is unavailable")
}
lease.mu.Lock()
if lease.references <= 0 || lease.parent == nil {
lease.mu.Unlock()
return errors.New("private restore staging root is unavailable")
}
lease.references--
if lease.references != 0 {
lease.mu.Unlock()
return nil
}
parent := lease.parent
lease.parent = nil
lease.mu.Unlock()
return parent.Close()
}
type inspectedArchiveEntry struct { type inspectedArchiveEntry struct {
metadata ArchiveEntryMetadata metadata ArchiveEntryMetadata
member *zip.File member *zip.File
@@ -271,6 +313,20 @@ func (result PreflightResult) revalidateArchive(ctx context.Context) (*os.File,
// StageArchive revalidates the retained archive and copies its exact bytes into a private file // StageArchive revalidates the retained archive and copies its exact bytes into a private file
// immediately before extraction. Later writes to the source archive cannot affect extraction. // immediately before extraction. Later writes to the source archive cannot affect extraction.
func (result PreflightResult) StageArchive(ctx context.Context) (_ *stagedArchive, resultErr error) { func (result PreflightResult) StageArchive(ctx context.Context) (_ *stagedArchive, resultErr error) {
return result.stageArchive(ctx, nil)
}
// stageArchiveAlongside reserves another immutable staging file through an already-retained
// root capability. Windows no-delete handles intentionally prevent reopening that root while a
// candidate stage is live, so restore shares the capability without widening share flags.
func (result PreflightResult) stageArchiveAlongside(ctx context.Context, existing *stagedArchive) (_ *stagedArchive, resultErr error) {
if existing == nil || existing.lease == nil {
return nil, errors.New("private restore staging root is unavailable")
}
return result.stageArchive(ctx, existing.lease)
}
func (result PreflightResult) stageArchive(ctx context.Context, existing *stagingRootLease) (_ *stagedArchive, resultErr error) {
source, err := result.revalidateArchive(ctx) source, err := result.revalidateArchive(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -278,9 +334,36 @@ func (result PreflightResult) StageArchive(ctx context.Context) (_ *stagedArchiv
if result.stagingRoot == "" || result.freeBytes == nil { if result.stagingRoot == "" || result.freeBytes == nil {
return nil, errors.New("backup archive has no controlled staging reservation") return nil, errors.New("backup archive has no controlled staging reservation")
} }
var (
parent safeio.PrivateDirectoryHandle
lease *stagingRootLease
)
if existing == nil {
if err := safeio.EnsurePrivateDirectory(result.stagingRoot); err != nil { if err := safeio.EnsurePrivateDirectory(result.stagingRoot); err != nil {
return nil, fmt.Errorf("create private restore staging root: %w", err) return nil, fmt.Errorf("create private restore staging root: %w", err)
} }
var found bool
parent, found, err = safeio.OpenPrivateDirectory(result.stagingRoot, true)
if err != nil || !found {
if parent != nil {
_ = parent.Close()
}
return nil, errors.New("open private restore staging root")
}
lease = &stagingRootLease{root: result.stagingRoot, parent: parent, references: 1}
} else {
lease = existing
parent, err = lease.retain(result.stagingRoot)
if err != nil {
return nil, err
}
}
releaseLease := true
defer func() {
if releaseLease {
_ = lease.release()
}
}()
freeBytes, err := result.freeBytes(result.stagingRoot) freeBytes, err := result.freeBytes(result.stagingRoot)
if err != nil { if err != nil {
return nil, errors.New("check private restore staging capacity") return nil, errors.New("check private restore staging capacity")
@@ -288,13 +371,6 @@ func (result PreflightResult) StageArchive(ctx context.Context) (_ *stagedArchiv
if result.ArchiveSize < 0 || freeBytes < uint64(result.ArchiveSize) { if result.ArchiveSize < 0 || freeBytes < uint64(result.ArchiveSize) {
return nil, errors.New("insufficient free disk space for private restore staging archive") return nil, errors.New("insufficient free disk space for private restore staging archive")
} }
parent, found, err := safeio.OpenPrivateDirectory(result.stagingRoot, true)
if err != nil || !found {
if parent != nil {
_ = parent.Close()
}
return nil, errors.New("open private restore staging root")
}
var ( var (
file *os.File file *os.File
name string name string
@@ -302,13 +378,11 @@ func (result PreflightResult) StageArchive(ctx context.Context) (_ *stagedArchiv
for attempt := 0; attempt < 8; attempt++ { for attempt := 0; attempt < 8; attempt++ {
name, err = newStagingArchiveName() name, err = newStagingArchiveName()
if err != nil { if err != nil {
_ = parent.Close()
return nil, errors.New("create private restore staging archive") return nil, errors.New("create private restore staging archive")
} }
var created bool var created bool
file, created, err = parent.CreateRegularFile(name) file, created, err = parent.CreateRegularFile(name)
if err != nil { if err != nil {
_ = parent.Close()
return nil, errors.New("create private restore staging archive") return nil, errors.New("create private restore staging archive")
} }
if created { if created {
@@ -317,15 +391,16 @@ func (result PreflightResult) StageArchive(ctx context.Context) (_ *stagedArchiv
file = nil file = nil
} }
if file == nil { if file == nil {
_ = parent.Close()
return nil, errors.New("create private restore staging archive") return nil, errors.New("create private restore staging archive")
} }
staged := &stagedArchive{ staged := &stagedArchive{
file: file, file: file,
parent: parent, parent: parent,
lease: lease,
name: name, name: name,
path: filepath.Join(result.stagingRoot, name), path: filepath.Join(result.stagingRoot, name),
} }
releaseLease = false
completed := false completed := false
defer func() { defer func() {
if !completed { if !completed {
@@ -394,10 +469,11 @@ func (staged *stagedArchive) Close() error {
if err != nil || !removed { if err != nil || !removed {
failed = true failed = true
} }
if err := staged.parent.Close(); err != nil { if staged.lease == nil || staged.lease.release() != nil {
failed = true failed = true
} }
staged.parent = nil staged.parent = nil
staged.lease = nil
} }
staged.name = "" staged.name = ""
staged.path = "" staged.path = ""
@@ -45,6 +45,50 @@ func TestStageArchiveProtectsWindowsStagingArtifactsWithOwnerOnlyACLs(t *testing
} }
} }
func TestStageArchiveRetainsTwoFilesThroughOneWindowsRootCapability(t *testing.T) {
installation := preflightTestInstallation(t)
archives := []string{filepath.Join(t.TempDir(), "candidate.zip"), filepath.Join(t.TempDir(), "recovery.zip")}
for _, archive := range archives {
writePreflightArchive(t, archive, preflightArchiveSpec{
entries: []preflightArchiveEntry{{path: "configuration/operator.env", body: []byte("safe")}},
})
}
candidate, err := Preflight(context.Background(), installation, PreflightRequest{Archive: archives[0], Confirm: true}, permissivePreflightDependencies())
if err != nil {
t.Fatal(err)
}
defer candidate.CloseArchive()
recovery, err := Preflight(context.Background(), installation, PreflightRequest{Archive: archives[1], Confirm: true}, permissivePreflightDependencies())
if err != nil {
t.Fatal(err)
}
defer recovery.CloseArchive()
candidateStage, err := candidate.StageArchive(context.Background())
if err != nil {
t.Fatal(err)
}
recoveryStage, err := recovery.stageArchiveAlongside(context.Background(), candidateStage)
if err != nil {
_ = candidateStage.Close()
t.Fatal(err)
}
if candidateStage.parent != recoveryStage.parent {
t.Fatal("paired stages did not share the retained root capability")
}
if err := candidateStage.Close(); err != nil {
_ = recoveryStage.Close()
t.Fatal(err)
}
if _, found, err := recoveryStage.parent.ReadRegular(recoveryStage.name, 1<<20); err != nil || !found {
_ = recoveryStage.Close()
t.Fatalf("recovery stage after candidate cleanup = found:%t err:%v", found, err)
}
if err := recoveryStage.Close(); err != nil {
t.Fatal(err)
}
}
func TestStageArchiveCloseUsesPinnedRootAfterAncestorSwap(t *testing.T) { func TestStageArchiveCloseUsesPinnedRootAfterAncestorSwap(t *testing.T) {
installation := preflightTestInstallation(t) installation := preflightTestInstallation(t)
if err := os.MkdirAll(installation.ControlDirectory(), 0o700); err != nil { if err := os.MkdirAll(installation.ControlDirectory(), 0o700); err != nil {
+1 -1
View File
@@ -134,7 +134,7 @@ func restoreWithDependencies(ctx context.Context, installation config.Installati
resultErr = errors.Join(resultErr, closeErr) resultErr = errors.Join(resultErr, closeErr)
} }
}() }()
recoveryStage, err := recovery.StageArchive(ctx) recoveryStage, err := recovery.stageArchiveAlongside(ctx, candidateStage)
if err != nil { if err != nil {
cleanupErr := deps.cleanupCheckpoint(checkpoint.Path) cleanupErr := deps.cleanupCheckpoint(checkpoint.Path)
return result, errors.Join(err, cleanupErr) return result, errors.Join(err, cleanupErr)