Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 82 additions & 7 deletions apps/druid/adapters/cli/worker_pull.go
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@ func init() {
WorkerPullCommand.Flags().StringVar(&workerPullAction.MountPath, "root", "/scroll", "Mounted runtime root path")
WorkerPullCommand.Flags().StringVar(&workerPullAction.CallbackURL, "callback-url", "", "Daemon worker callback URL")
WorkerPullCommand.Flags().StringVar(&workerPullAction.TokenFile, "callback-token-file", "", "Projected ServiceAccount token file for callbacks")
WorkerPullCommand.Flags().BoolVar(&workerPullAction.PreserveReleaseManifest, "preserve-release-manifest", false, "Restore the release manifest.json carried by a backup")
WorkerPullCommand.Flags().StringVar(&workerPullMode, "mode", string(ports.RuntimeWorkerModeCreate), "Pull mode: create, update, or restore")
WorkerPullCommand.MarkFlagRequired("artifact")
WorkerPullCommand.MarkFlagRequired("runtime-id")
Expand Down Expand Up @@ -164,27 +165,101 @@ func pullWorkerUpdate(root string, artifact string, oci ports.OciRegistryInterfa
}

func pullWorkerRestore(root string, artifact string, oci ports.OciRegistryInterface) error {
tmp, err := os.MkdirTemp("", "druid-worker-restore-*")
if err := os.MkdirAll(root, 0755); err != nil {
return err
}
stage, err := os.MkdirTemp(root, ".druid-worker-restore-stage-*")
if err != nil {
return err
}
defer os.RemoveAll(tmp)
if err := coreservices.MaterializeScrollArtifact(artifact, tmp, oci, true); err != nil {
defer os.RemoveAll(stage)
puller, ok := oci.(interface {
PullSelectiveWithOptions(string, string, bool, *domain.SnapshotProgress, registry.TransferOptions) error
})
if !ok {
return fmt.Errorf("OCI registry does not support preserve-release-manifest")
}
if err := puller.PullSelectiveWithOptions(stage, artifact, true, nil, registry.TransferOptions{PreserveReleaseManifest: true}); err != nil {
return err
}
if err := os.MkdirAll(root, 0755); err != nil {
return replaceRestoredRoot(root, stage)
}

const restoreRootUnsafePrefix = "restore root may be partial:"

func replaceRestoredRoot(root string, stage string) error {
return replaceRestoredRootWithRename(root, stage, os.Rename)
}

// replaceRestoredRootWithRename swaps staged artifact entries into a runtime
// root without copying over a live tree. If an entry move fails, it restores
// the original entries before returning. An error with restoreRootUnsafePrefix
// means that rollback itself failed and callers must keep the runtime stopped.
func replaceRestoredRootWithRename(root string, stage string, rename func(string, string) error) error {
rollback, err := os.MkdirTemp(root, ".druid-worker-restore-rollback-*")
if err != nil {
return err
}
defer os.RemoveAll(rollback)

stageName := filepath.Base(stage)
rollbackName := filepath.Base(rollback)
entries, err := os.ReadDir(root)
if err != nil {
return err
}
original := make([]string, 0, len(entries))
for _, entry := range entries {
if err := os.RemoveAll(filepath.Join(root, entry.Name())); err != nil {
return err
if entry.Name() != stageName && entry.Name() != rollbackName {
original = append(original, entry.Name())
}
}
return copyPath(tmp, root)
stagedEntries, err := os.ReadDir(stage)
if err != nil {
return err
}

restoreOriginal := func(installed []string) error {
for _, name := range installed {
if err := os.RemoveAll(filepath.Join(root, name)); err != nil {
return err
}
}
for _, name := range original {
from := filepath.Join(rollback, name)
if _, err := os.Lstat(from); os.IsNotExist(err) {
continue
} else if err != nil {
return err
}
if err := rename(from, filepath.Join(root, name)); err != nil {
return err
}
}
return nil
}

for _, entry := range original {
if err := rename(filepath.Join(root, entry), filepath.Join(rollback, entry)); err != nil {
if rollbackErr := restoreOriginal(nil); rollbackErr != nil {
return fmt.Errorf("%s failed to restore original runtime after moving %s: %w", restoreRootUnsafePrefix, entry, rollbackErr)
}
return fmt.Errorf("restore transaction rolled back while moving original entry %s: %w", entry, err)
}
}

installed := make([]string, 0, len(stagedEntries))
for _, entry := range stagedEntries {
name := entry.Name()
if err := rename(filepath.Join(stage, name), filepath.Join(root, name)); err != nil {
if rollbackErr := restoreOriginal(installed); rollbackErr != nil {
return fmt.Errorf("%s failed to restore original runtime after staging %s: %w", restoreRootUnsafePrefix, name, rollbackErr)
}
return fmt.Errorf("restore transaction rolled back while staging %s: %w", name, err)
}
installed = append(installed, name)
}
return nil
}

func collectSkipUpdatePaths(out map[string]bool, parent string, chunks []*domain.Chunks) {
Expand Down
4 changes: 3 additions & 1 deletion apps/druid/adapters/cli/worker_push.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (

var workerPushArtifact string
var workerPushRoot string
var workerPushPreserveReleaseManifest bool

var WorkerPushCommand = &cobra.Command{
Use: "push",
Expand All @@ -29,7 +30,7 @@ var WorkerPushCommand = &cobra.Command{
}
repo, tag := utils.SplitArtifact(workerPushArtifact)
oci := registry.NewOciClient(loadWorkerRegistryStore())
_, err = oci.Push(workerPushRoot, repo, tag, nil, false, &scroll.File)
_, err = oci.PushWithOptions(workerPushRoot, repo, tag, nil, false, &scroll.File, registry.TransferOptions{PreserveReleaseManifest: workerPushPreserveReleaseManifest})
return err
},
}
Expand All @@ -38,5 +39,6 @@ func init() {
WorkerCommand.AddCommand(WorkerPushCommand)
WorkerPushCommand.Flags().StringVar(&workerPushArtifact, "artifact", "", "OCI artifact to push")
WorkerPushCommand.Flags().StringVar(&workerPushRoot, "root", "/scroll", "Mounted runtime root path")
WorkerPushCommand.Flags().BoolVar(&workerPushPreserveReleaseManifest, "preserve-release-manifest", false, "Include the installed release manifest.json in the backup payload")
WorkerPushCommand.MarkFlagRequired("artifact")
}
86 changes: 86 additions & 0 deletions apps/druid/adapters/cli/worker_test.go
Original file line number Diff line number Diff line change
@@ -1,14 +1,17 @@
package cli

import (
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"

"github.com/highcard-dev/daemon/internal/core/domain"
"github.com/highcard-dev/daemon/internal/core/ports"
"github.com/highcard-dev/daemon/internal/core/services/registry"
v1 "github.com/opencontainers/image-spec/specs-go/v1"
"github.com/spf13/cobra"
"oras.land/oras-go/v2/registry/remote"
Expand All @@ -33,6 +36,18 @@ func TestWorkerPullCommandRequiresRuntimeID(t *testing.T) {
}
}

func TestWorkerPullCommandSupportsReleaseManifestPreservation(t *testing.T) {
if flag := WorkerPullCommand.Flags().Lookup("preserve-release-manifest"); flag == nil {
t.Fatal("worker pull should expose --preserve-release-manifest")
}
}

func TestWorkerPushCommandSupportsReleaseManifestPreservation(t *testing.T) {
if flag := WorkerPushCommand.Flags().Lookup("preserve-release-manifest"); flag == nil {
t.Fatal("worker push should expose --preserve-release-manifest")
}
}

func TestReportWorkerResultUsesTokenOnlyWhenProvided(t *testing.T) {
tests := []struct {
name string
Expand Down Expand Up @@ -108,6 +123,67 @@ func TestWorkerRestoreStagesBeforeReplacingRoot(t *testing.T) {
}
}

func TestReplaceRestoredRootRollsBackAfterStagingBegins(t *testing.T) {
root := t.TempDir()
mustWrite(t, filepath.Join(root, "scroll.yaml"), "name: original\n")
mustWrite(t, filepath.Join(root, "data", "world.txt"), "original")
stage, err := os.MkdirTemp(root, ".druid-worker-restore-stage-*")
if err != nil {
t.Fatal(err)
}
defer os.RemoveAll(stage)
mustWrite(t, filepath.Join(stage, "a-new.txt"), "new")
mustWrite(t, filepath.Join(stage, "b-new.txt"), "new")

stagedMoves := 0
err = replaceRestoredRootWithRename(root, stage, func(from string, to string) error {
if strings.HasPrefix(from, stage+string(os.PathSeparator)) {
stagedMoves++
if stagedMoves == 2 {
return errors.New("injected staging failure")
}
}
return os.Rename(from, to)
})
if err == nil || !strings.Contains(err.Error(), "restore transaction rolled back") {
t.Fatalf("restore error = %v, want successful rollback error", err)
}
assertFile(t, filepath.Join(root, "scroll.yaml"), "name: original\n")
assertFile(t, filepath.Join(root, "data", "world.txt"), "original")
if _, err := os.Stat(filepath.Join(root, "a-new.txt")); !os.IsNotExist(err) {
t.Fatalf("staged file should be removed during rollback, stat err = %v", err)
}
}

func TestReplaceRestoredRootMarksFailedRollbackUnsafe(t *testing.T) {
root := t.TempDir()
mustWrite(t, filepath.Join(root, "original.txt"), "original")
stage, err := os.MkdirTemp(root, ".druid-worker-restore-stage-*")
if err != nil {
t.Fatal(err)
}
defer os.RemoveAll(stage)
mustWrite(t, filepath.Join(stage, "a-new.txt"), "new")
mustWrite(t, filepath.Join(stage, "b-new.txt"), "new")

stagedMoves := 0
err = replaceRestoredRootWithRename(root, stage, func(from string, to string) error {
if strings.HasPrefix(from, stage+string(os.PathSeparator)) {
stagedMoves++
if stagedMoves == 2 {
return errors.New("injected staging failure")
}
}
if strings.Contains(filepath.Base(filepath.Dir(from)), "restore-rollback-") {
return errors.New("injected rollback failure")
}
return os.Rename(from, to)
})
if err == nil || !strings.HasPrefix(err.Error(), restoreRootUnsafePrefix) {
t.Fatalf("restore error = %v, want unsafe rollback marker", err)
}
}

func TestWorkerCollectSkipUpdatePaths(t *testing.T) {
root := filepath.Join(t.TempDir(), "root")
mustWrite(t, filepath.Join(root, "scroll.yaml"), `name: skip-test
Expand Down Expand Up @@ -195,6 +271,16 @@ func (f fakeRestoreOCI) PullSelective(dir string, artifact string, includeData b
return nil
}

func (f fakeRestoreOCI) PullSelectiveWithOptions(dir string, artifact string, includeData bool, progress *domain.SnapshotProgress, options registry.TransferOptions) error {
if !options.PreserveReleaseManifest {
f.t.Fatal("restore must preserve the release manifest")
}
mustWrite(f.t, filepath.Join(dir, "scroll.yaml"), "name: restored\n")
mustWrite(f.t, filepath.Join(dir, "manifest.json"), `{"digest":"sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"}`)
mustWrite(f.t, filepath.Join(dir, "data", "logs", "latest.log"), "restored")
return nil
}

func (fakeRestoreOCI) FetchFile(string, string) ([]byte, error) {
return nil, os.ErrNotExist
}
Expand Down
41 changes: 39 additions & 2 deletions apps/druid/core/services/runtime_access.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ package services

import (
"context"
"fmt"
"strings"

"github.com/highcard-dev/daemon/internal/core/domain"
"github.com/highcard-dev/daemon/internal/core/ports"
Expand Down Expand Up @@ -56,31 +58,56 @@ func (s *RuntimeSupervisor) ApplyRouting(id string, assignments []domain.Runtime
}

func (s *RuntimeSupervisor) Backup(id string, artifact string, registryCredentials []domain.RegistryCredential) (*domain.RuntimeScroll, error) {
unlock := s.lockRuntimeOperation(id)
defer unlock()
session, err := s.sessionFor(id)
if err != nil {
return nil, err
}
wasRunning := sessionWasRunning(session)
if err := session.StopRuntimeForMaintenance(); err != nil {
session.markError(err)
return nil, err
}
if err := session.Backup(context.Background(), artifact, registryCredentials); err != nil {
if wasRunning {
if _, restartErr := s.startScroll(id); restartErr != nil {
return nil, fmt.Errorf("backup failed: %w; failed to restart prior runtime: %v", err, restartErr)
}
return nil, err
}
session.markError(err)
return nil, err
}
if wasRunning {
return s.startScroll(id)
}
return s.store.GetScroll(id)
}

func (s *RuntimeSupervisor) Restore(id string, artifact string, restart bool, registryCredentials []domain.RegistryCredential) (*domain.RuntimeScroll, error) {
unlock := s.lockRuntimeOperation(id)
defer unlock()
session, err := s.sessionFor(id)
if err != nil {
return nil, err
}
session.mu.Lock()
root := session.runtimeScroll.Root
session.mu.Unlock()
if err := session.StopRuntime(); err != nil {
wasRunning := sessionWasRunning(session)
if err := session.StopRuntimeForMaintenance(); err != nil {
session.markError(err)
return nil, err
}
materialized, err := s.runPullWorker(context.Background(), s.runtimeBackend, ports.RuntimeWorkerModeRestore, id, artifact, root, registryCredentials, "")
if err != nil {
if wasRunning && restoreFailureCanRestart(err) {
if _, restartErr := s.startScroll(id); restartErr != nil {
return nil, fmt.Errorf("restore failed: %w; failed to restart prior runtime: %v", err, restartErr)
}
return nil, err
}
session.markError(err)
return nil, err
}
Expand All @@ -89,11 +116,21 @@ func (s *RuntimeSupervisor) Restore(id string, artifact string, restart bool, re
return nil, err
}
if restart {
return s.StartScroll(id)
return s.startScroll(id)
}
return s.store.GetScroll(id)
}

func sessionWasRunning(session *RuntimeSession) bool {
session.mu.Lock()
defer session.mu.Unlock()
return session.runtimeScroll.Status == domain.RuntimeScrollStatusRunning
}

func restoreFailureCanRestart(err error) bool {
return !strings.Contains(err.Error(), "restore root may be partial:")
}

func (s *RuntimeSupervisor) ScrollFile(id string) (*domain.File, error) {
session, err := s.sessionFor(id)
if err != nil {
Expand Down
13 changes: 12 additions & 1 deletion apps/druid/core/services/runtime_lifecycle.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@ func (s *RuntimeSupervisor) Delete(id string) error {
}

func (s *RuntimeSupervisor) DeleteWithPolicy(id string, purgeData bool) error {
unlock := s.lockRuntimeOperation(id)
defer unlock()
s.mu.Lock()
session := s.sessions[id]
delete(s.sessions, id)
Expand All @@ -28,10 +30,17 @@ func (s *RuntimeSupervisor) DeleteWithPolicy(id string, purgeData bool) error {
}

func (s *RuntimeSupervisor) StartScroll(id string) (*domain.RuntimeScroll, error) {
unlock := s.lockRuntimeOperation(id)
defer unlock()
return s.startScroll(id)
}

func (s *RuntimeSupervisor) startScroll(id string) (*domain.RuntimeScroll, error) {
session, err := s.sessionFor(id)
if err != nil {
return nil, err
}
session.Start()
if err := session.AutoStartServe(); err != nil {
session.markError(err)
return nil, err
Expand All @@ -52,11 +61,13 @@ func (s *RuntimeSupervisor) StartScroll(id string) (*domain.RuntimeScroll, error
}

func (s *RuntimeSupervisor) Stop(id string) (*domain.RuntimeScroll, error) {
unlock := s.lockRuntimeOperation(id)
defer unlock()
session, err := s.detachSession(id)
if err != nil {
return nil, err
}
if err := session.StopRuntime(); err != nil {
if err := session.StopRuntimeForMaintenance(); err != nil {
session.markError(err)
return nil, err
}
Expand Down
Loading
Loading