diff --git a/tools/tht/internal/authprojection/projection_linux.go b/tools/tht/internal/authprojection/projection_linux.go index 246549a6..ceb6996e 100644 --- a/tools/tht/internal/authprojection/projection_linux.go +++ b/tools/tht/internal/authprojection/projection_linux.go @@ -16,6 +16,9 @@ import ( const maximumDirectoryEntries = 16 type testHooks struct { + beforeFlock func(int) + beforeStageWrite func() error + beforeStageFsync func() error beforeStageRename func() error beforeCurrentRename func() error beforeFinalVerify func() error @@ -49,6 +52,15 @@ func hook(selectHook func(testHooks) func() error) error { return nil } +func rendezvousBeforeFlock(root int) { + hookState.Lock() + callback := hookState.hooks.beforeFlock + hookState.Unlock() + if callback != nil { + callback(root) + } +} + func Inspect(spec Spec) (Status, error) { root, err := openRoot(spec) if err != nil { @@ -63,16 +75,12 @@ func Begin(spec Spec, before *Snapshot, requireReadyMatch bool) (*Transaction, e if err != nil { return nil, err } + rendezvousBeforeFlock(root) if unix.Flock(root, unix.LOCK_EX) != nil { unix.Close(root) return nil, ErrIntegrity } transaction := &Transaction{spec: spec, before: before, rootFD: root, lockFD: root} - if ensureGenerations(root, spec) != nil { - _ = transaction.Close() - return nil, ErrIntegrity - } - selector, selectorErr := readSelector(root, spec) switch { case selectorErr == nil && selector.State == "blocked": @@ -81,7 +89,7 @@ func Begin(spec Spec, before *Snapshot, requireReadyMatch bool) (*Transaction, e return nil, ErrIntegrity } case selectorErr == nil: - if validateRootEntries(root, spec) != nil { + if recoverCurrentTemporaries(root, spec) != nil || validateRootEntries(root, spec) != nil { _ = transaction.Close() return nil, ErrIntegrity } @@ -98,7 +106,10 @@ func Begin(spec Spec, before *Snapshot, requireReadyMatch bool) (*Transaction, e case currentExists(root): _ = transaction.Close() return nil, ErrIntegrity - case requireReadyMatch || validateInitialRootEntries(root, spec) != nil: + case requireReadyMatch || validateEmptyRootEntries(root, spec) != nil: + _ = transaction.Close() + return nil, ErrIntegrity + case ensureGenerations(root, spec) != nil: _ = transaction.Close() return nil, ErrIntegrity } @@ -135,7 +146,7 @@ func (transaction *Transaction) Commit(after Snapshot) (Status, error) { if writeSelector(transaction.rootFD, transaction.spec, selector) != nil { return Status{}, ErrIntegrity } - status, err := inspectFD(transaction.rootFD, transaction.spec) + status, err := inspectSelectedGeneration(transaction.rootFD, transaction.spec, selector) if err != nil || status.Snapshot.Generation != after.Generation { return Status{}, ErrIntegrity } @@ -146,6 +157,10 @@ func (transaction *Transaction) Commit(after Snapshot) (Status, error) { if retain(transaction.rootFD, transaction.spec, selector) != nil { return Status{}, ErrIntegrity } + status, err = inspectFD(transaction.rootFD, transaction.spec) + if err != nil || status.Snapshot.Generation != after.Generation { + return Status{}, ErrIntegrity + } return status, nil } @@ -313,9 +328,9 @@ func validateRootEntries(root int, spec Spec) error { } return nil } -func validateInitialRootEntries(root int, spec Spec) error { +func validateEmptyRootEntries(root int, spec Spec) error { names, err := scanDirectoryNames(root, spec) - if err != nil || len(names) != 1 || !names["generations"] { + if err != nil || len(names) != 0 { return ErrIntegrity } return nil @@ -332,6 +347,20 @@ func inspectFD(root int, spec Spec) (Status, error) { return inspectCapturedSelector(root, spec, selector) } func inspectCapturedSelector(root int, spec Spec, selector Selector) (Status, error) { + if selector.State == "blocked" { + return Status{Selector: selector}, ErrBlocked + } + status, err := inspectSelectedGeneration(root, spec, selector) + if err != nil { + return Status{}, err + } + if validateReadyGenerationNamespace(root, spec, selector) != nil { + return Status{}, ErrIntegrity + } + return status, nil +} + +func inspectSelectedGeneration(root int, spec Spec, selector Selector) (Status, error) { if selector.State == "blocked" { return Status{Selector: selector}, ErrBlocked } @@ -342,6 +371,34 @@ func inspectCapturedSelector(root int, spec Spec, selector Selector) (Status, er return Status{Selector: selector, Snapshot: snapshot}, nil } +func validateReadyGenerationNamespace(root int, spec Spec, selector Selector) error { + generations, err := openDirectoryAt(root, "generations", spec) + if err != nil { + return ErrIntegrity + } + defer unix.Close(generations) + names, err := scanDirectoryNames(generations, spec) + if err != nil { + return ErrIntegrity + } + keep := map[string]bool{selector.Generation: true} + for _, generation := range selector.PreviousGenerations { + keep[generation] = true + } + if len(names) != len(keep) { + return ErrIntegrity + } + for name := range names { + if !validHex(name, 64) || !keep[name] { + return ErrIntegrity + } + if _, err := readGeneration(root, spec, name); err != nil { + return ErrIntegrity + } + } + return nil +} + func readSelector(root int, spec Spec) (Selector, error) { data, err := readRegularAt(root, "CURRENT", spec, maximumSelectorBytes) if err != nil { @@ -516,7 +573,7 @@ func stageGeneration(root int, spec Spec, transactionID string, snapshot Snapsho if found, readErr := readGeneration(root, spec, snapshot.Generation); readErr == nil && found.Generation == snapshot.Generation { return nil } - if removeConfinedDirectory(generations, snapshot.Generation, spec, false) != nil { + if !currentMatchesBlockedTransaction(root, spec, transactionID) || removeConfinedDirectory(generations, snapshot.Generation, spec, false) != nil { return ErrIntegrity } } else if !errors.Is(err, unix.ENOENT) { @@ -546,7 +603,16 @@ func stageGeneration(root int, spec Spec, transactionID string, snapshot Snapsho if err != nil { return ErrIntegrity } - if writeRegularAt(stage, "auth.yaml", spec, snapshot.Auth) != nil || (snapshot.Mode == "local" && writeRegularAt(stage, "users.yaml", spec, snapshot.Users) != nil) || writeRegularAt(stage, "manifest.json", spec, manifestData) != nil || unix.Fsync(stage) != nil { + if err := hook(func(h testHooks) func() error { return h.beforeStageWrite }); err != nil { + return err + } + if writeRegularAt(stage, "auth.yaml", spec, snapshot.Auth) != nil || (snapshot.Mode == "local" && writeRegularAt(stage, "users.yaml", spec, snapshot.Users) != nil) || writeRegularAt(stage, "manifest.json", spec, manifestData) != nil { + return ErrIntegrity + } + if err := hook(func(h testHooks) func() error { return h.beforeStageFsync }); err != nil { + return err + } + if unix.Fsync(stage) != nil { return ErrIntegrity } if err := hook(func(h testHooks) func() error { return h.beforeStageRename }); err != nil { @@ -622,6 +688,13 @@ func recoverTemporary(root int, spec Spec, blocked Selector) error { return ErrIntegrity } } + if unix.Fsync(generations) != nil { + return ErrIntegrity + } + return recoverCurrentTemporaries(root, spec) +} + +func recoverCurrentTemporaries(root int, spec Spec) error { rootNames, err := scanDirectoryNames(root, spec) if err != nil || !rootNames["CURRENT"] || !rootNames["generations"] { return ErrIntegrity @@ -643,11 +716,16 @@ func recoverTemporary(root int, spec Spec, blocked Selector) error { return ErrIntegrity } } - if unix.Fsync(generations) != nil || unix.Fsync(root) != nil { + if unix.Fsync(root) != nil { return ErrIntegrity } return nil } + +func currentMatchesBlockedTransaction(root int, spec Spec, transactionID string) bool { + selector, err := readSelector(root, spec) + return err == nil && selector.State == "blocked" && selector.Transaction == transactionID +} func stageGenerationName(name, transactionID string) (string, bool) { prefix := ".stage-" + transactionID + "-" if len(name) != len(prefix)+64 || name[:len(prefix)] != prefix { diff --git a/tools/tht/internal/authprojection/projection_linux_test.go b/tools/tht/internal/authprojection/projection_linux_test.go index 18835066..2613ebc5 100644 --- a/tools/tht/internal/authprojection/projection_linux_test.go +++ b/tools/tht/internal/authprojection/projection_linux_test.go @@ -11,7 +11,6 @@ import ( "sync" "syscall" "testing" - "time" "golang.org/x/sys/unix" ) @@ -85,9 +84,7 @@ func TestInspectRejectsMissingBlockedMalformedAndTamperedCurrent(t *testing.T) { if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { t.Fatalf("malformed error = %v", err) } - if err := os.Remove(filepath.Join(spec.RuntimeRoot, "CURRENT")); err != nil { - t.Fatal(err) - } + spec = testSpec(t) commitSnapshot(t, spec, testSnapshot(t, "tampered")) if err := os.WriteFile(filepath.Join(spec.RuntimeRoot, "CURRENT"), []byte(`{"version":1,"state":"ready"}`), runtimeFileMode); err != nil { t.Fatal(err) @@ -220,12 +217,161 @@ func TestInspectRejectsUnexpectedGenerationEntry(t *testing.T) { } } +func TestInspectRejectsReadyNamespaceOutsideCurrentHistory(t *testing.T) { + for _, tc := range []struct { + name string + mutate func(t *testing.T, spec Spec, status Status) + }{ + {"extra generation", func(t *testing.T, spec Spec, _ Status) { + name := strings.Repeat("e", 64) + if err := os.Mkdir(filepath.Join(spec.RuntimeRoot, "generations", name), runtimeDirectoryMode); err != nil { + t.Fatal(err) + } + }}, + {"stage", func(t *testing.T, spec Spec, _ Status) { + name := ".stage-" + strings.Repeat("a", 32) + "-" + strings.Repeat("e", 64) + if err := os.Mkdir(filepath.Join(spec.RuntimeRoot, "generations", name), runtimeDirectoryMode); err != nil { + t.Fatal(err) + } + }}, + {"corrupt predecessor", func(t *testing.T, spec Spec, status Status) { + if len(status.Selector.PreviousGenerations) != 1 { + t.Fatalf("history = %#v", status.Selector.PreviousGenerations) + } + path := filepath.Join(spec.RuntimeRoot, "generations", status.Selector.PreviousGenerations[0], "auth.yaml") + if err := os.WriteFile(path, []byte("tampered"), runtimeFileMode); err != nil { + t.Fatal(err) + } + }}, + {"missing predecessor", func(t *testing.T, spec Spec, status Status) { + if len(status.Selector.PreviousGenerations) != 1 { + t.Fatalf("history = %#v", status.Selector.PreviousGenerations) + } + if err := os.RemoveAll(filepath.Join(spec.RuntimeRoot, "generations", status.Selector.PreviousGenerations[0])); err != nil { + t.Fatal(err) + } + }}, + } { + t.Run(tc.name, func(t *testing.T) { + spec := testSpec(t) + commitSnapshot(t, spec, testSnapshot(t, "prior-"+tc.name)) + status := commitSnapshot(t, spec, testSnapshot(t, "current-"+tc.name)) + tc.mutate(t, spec, status) + if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { + t.Fatalf("Inspect error = %v", err) + } + }) + } +} + +func TestBeginRecoversStrictCurrentTemporaryBesideReadySelector(t *testing.T) { + spec := testSpec(t) + prior := testSnapshot(t, "ready-temporary") + commitSnapshot(t, spec, prior) + temporaryID := strings.Repeat("a", 32) + temporary, err := encodeSelector(Selector{Version: SchemaVersion, State: "blocked", Transaction: temporaryID}) + if err != nil { + t.Fatal(err) + } + path := filepath.Join(spec.RuntimeRoot, ".current-"+temporaryID+".tmp") + if err := os.WriteFile(path, temporary, runtimeFileMode); err != nil { + t.Fatal(err) + } + transaction, err := Begin(spec, &prior, true) + if err != nil { + t.Fatal(err) + } + defer transaction.Close() + if transaction.prior == nil || transaction.prior.Snapshot.Generation != prior.Generation { + t.Fatalf("prior = %#v", transaction.prior) + } + if _, err := os.Lstat(path); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("ready temporary remains: %v", err) + } +} + +func TestBeginRejectsMalformedCurrentWithoutCreatingGenerations(t *testing.T) { + spec := testSpec(t) + if err := os.WriteFile(filepath.Join(spec.RuntimeRoot, "CURRENT"), []byte("malformed"), runtimeFileMode); err != nil { + t.Fatal(err) + } + if _, err := Begin(spec, nil, false); !errors.Is(err, ErrIntegrity) { + t.Fatalf("Begin error = %v", err) + } + if _, err := os.Lstat(filepath.Join(spec.RuntimeRoot, "generations")); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("Begin mutated malformed root: %v", err) + } +} + +func TestCommitRefusesCorruptTargetUnlessCurrentIsMatchingBlockedTransaction(t *testing.T) { + for _, tc := range []struct { + name string + current func(t *testing.T, root string, transaction string, snapshot Snapshot) + }{ + {"ready", func(t *testing.T, root, _ string, snapshot Snapshot) { + data, err := encodeSelector(Selector{Version: SchemaVersion, State: "ready", Transaction: strings.Repeat("b", 32), Generation: snapshot.Generation}) + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(root, "CURRENT"), data, runtimeFileMode); err != nil { + t.Fatal(err) + } + }}, + {"malformed", func(t *testing.T, root, _ string, _ Snapshot) { + if err := os.WriteFile(filepath.Join(root, "CURRENT"), []byte("malformed"), runtimeFileMode); err != nil { + t.Fatal(err) + } + }}, + {"other blocked transaction", func(t *testing.T, root, transaction string, _ Snapshot) { + other := strings.Repeat("b", 32) + if other == transaction { + other = strings.Repeat("c", 32) + } + data, err := encodeSelector(Selector{Version: SchemaVersion, State: "blocked", Transaction: other}) + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(root, "CURRENT"), data, runtimeFileMode); err != nil { + t.Fatal(err) + } + }}, + } { + t.Run(tc.name, func(t *testing.T) { + spec := testSpec(t) + snapshot := testSnapshot(t, "corrupt-"+tc.name) + transaction, err := Begin(spec, nil, false) + if err != nil { + t.Fatal(err) + } + defer transaction.Close() + target := filepath.Join(spec.RuntimeRoot, "generations", snapshot.Generation) + if err := os.Mkdir(target, runtimeDirectoryMode); err != nil { + t.Fatal(err) + } + for _, name := range []string{"auth.yaml", "users.yaml", "manifest.json"} { + if err := os.WriteFile(filepath.Join(target, name), []byte("corrupt"), runtimeFileMode); err != nil { + t.Fatal(err) + } + } + tc.current(t, spec.RuntimeRoot, transaction.transactionID, snapshot) + if _, err := transaction.Commit(snapshot); !errors.Is(err, ErrIntegrity) { + t.Fatalf("Commit error = %v", err) + } + if _, err := os.Lstat(target); err != nil { + t.Fatalf("corrupt target was removed: %v", err) + } + }) + } +} + func TestCommitLeavesBlockedAfterInjectedStageWriteFsyncRenameAndVerifyFailure(t *testing.T) { for _, tc := range []struct { name string hooks testHooks }{ - {"stage", testHooks{beforeStageRename: func() error { return errors.New("sentinel") }}}, + {"stage write", testHooks{beforeStageWrite: func() error { return errors.New("sentinel") }}}, + {"stage fsync", testHooks{beforeStageFsync: func() error { return errors.New("sentinel") }}}, + {"stage rename", testHooks{beforeStageRename: func() error { return errors.New("sentinel") }}}, {"current", testHooks{beforeCurrentRename: func() error { return errors.New("sentinel") }}}, {"verify", testHooks{beforeFinalVerify: func() error { return errors.New("sentinel") }}}, } { @@ -474,21 +620,30 @@ func TestConcurrentBeginSerializesAcrossCurrentRenameWithoutCreatingLockEntry(t if err != nil { t.Fatal(err) } - started := make(chan struct{}) + checked := make(chan error, 1) + restore := setTestHooksForTest(testHooks{beforeFlock: func(fd int) { + if err := unix.Flock(fd, unix.LOCK_EX|unix.LOCK_NB); !errors.Is(err, unix.EWOULDBLOCK) { + checked <- fmt.Errorf("nonblocking flock error = %v", err) + return + } + checked <- nil + }}) + defer restore() finished := make(chan error, 1) go func() { - close(started) second, err := Begin(spec, nil, false) if err == nil { err = second.Close() } finished <- err }() - <-started + if err := <-checked; err != nil { + t.Fatal(err) + } select { case err := <-finished: t.Fatalf("second Begin did not block: %v", err) - case <-time.After(50 * time.Millisecond): + default: } if err := first.Close(); err != nil { t.Fatal(err)