//go:build linux package authprojection import ( "errors" "fmt" "os" "path/filepath" "strings" "sync" "syscall" "testing" "golang.org/x/sys/unix" ) func testSpec(t *testing.T) Spec { t.Helper() root := t.TempDir() if err := os.Chmod(root, runtimeDirectoryMode); err != nil { t.Fatal(err) } return Spec{RuntimeRoot: root, UID: uint32(os.Geteuid()), GID: uint32(os.Getegid())} } func testSnapshot(t *testing.T, suffix string) Snapshot { t.Helper() snapshot, err := NewSnapshot("local", []byte("auth-"+suffix), []byte("users-"+suffix)) if err != nil { t.Fatal(err) } return snapshot } func commitSnapshot(t *testing.T, spec Spec, snapshot Snapshot) Status { t.Helper() transaction, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } defer func() { _ = transaction.Close() }() status, err := transaction.Commit(snapshot) if err != nil { t.Fatal(err) } return status } func TestBeginCommitPublishesOneVerifiedGeneration(t *testing.T) { spec := testSpec(t) snapshot := testSnapshot(t, "first") status := commitSnapshot(t, spec, snapshot) if status.Selector.State != "ready" || status.Selector.Generation != snapshot.Generation { t.Fatalf("status = %#v", status.Selector) } got, err := Inspect(spec) if err != nil { t.Fatal(err) } if got.Snapshot.Generation != snapshot.Generation || string(got.Snapshot.Auth) != string(snapshot.Auth) { t.Fatalf("snapshot = %#v", got.Snapshot) } } func TestInspectRejectsMissingBlockedMalformedAndTamperedCurrent(t *testing.T) { spec := testSpec(t) if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { t.Fatalf("missing error = %v", err) } transaction, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } if _, err := Inspect(spec); !errors.Is(err, ErrBlocked) { t.Fatalf("blocked error = %v", err) } if err := transaction.Close(); err != nil { t.Fatal(err) } if err := os.WriteFile(filepath.Join(spec.RuntimeRoot, "CURRENT"), []byte("not-json"), runtimeFileMode); err != nil { t.Fatal(err) } if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { t.Fatalf("malformed error = %v", 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) } if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { t.Fatalf("tampered error = %v", err) } } func TestInspectIsReadOnlyWhenProjectionIsIncomplete(t *testing.T) { spec := testSpec(t) if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { t.Fatalf("Inspect error = %v", err) } entries, err := os.ReadDir(spec.RuntimeRoot) if err != nil { t.Fatal(err) } if len(entries) != 0 { t.Fatalf("Inspect created runtime entries: %#v", entries) } } func TestDescriptorDirectoryScansDoNotConsumeCallerOffset(t *testing.T) { spec := testSpec(t) commitSnapshot(t, spec, testSnapshot(t, "scan-offset")) root, err := openRoot(spec) if err != nil { t.Fatal(err) } defer unix.Close(root) first, err := scanDirectoryNames(root, spec) if err != nil { t.Fatal(err) } second, err := scanDirectoryNames(root, spec) if err != nil { t.Fatal(err) } if len(first) != 2 || len(second) != 2 || !second["CURRENT"] || !second["generations"] { t.Fatalf("directory scans = %#v then %#v", first, second) } } func TestInspectRejectsWrongModeOwnerHardlinkSymlinkAndUnexpectedEntry(t *testing.T) { cases := []struct { name string mutate func(t *testing.T, spec Spec, status Status) }{ { name: "mode", mutate: func(t *testing.T, spec Spec, status Status) { t.Helper() if err := os.Chmod(filepath.Join(spec.RuntimeRoot, "CURRENT"), 0o644); err != nil { t.Fatal(err) } }, }, { name: "generation directory mode", mutate: func(t *testing.T, spec Spec, status Status) { t.Helper() if err := os.Chmod(filepath.Join(spec.RuntimeRoot, "generations", status.Snapshot.Generation), 0o755); err != nil { t.Fatal(err) } }, }, { name: "owner", mutate: func(t *testing.T, spec Spec, status Status) { t.Helper() if os.Geteuid() != 0 { t.Skip("requires root to change owner") } if err := os.Chown(filepath.Join(spec.RuntimeRoot, "CURRENT"), 10002, 10002); err != nil { t.Fatal(err) } }, }, { name: "hardlink", mutate: func(t *testing.T, spec Spec, status Status) { t.Helper() source := filepath.Join(spec.RuntimeRoot, "generations", status.Snapshot.Generation, "auth.yaml") if err := os.Link(source, filepath.Join(spec.RuntimeRoot, "linked")); err != nil { t.Fatal(err) } }, }, { name: "symlink", mutate: func(t *testing.T, spec Spec, status Status) { t.Helper() if err := os.Symlink("CURRENT", filepath.Join(spec.RuntimeRoot, "link")); err != nil { t.Fatal(err) } }, }, { name: "unexpected", mutate: func(t *testing.T, spec Spec, status Status) { t.Helper() if err := os.WriteFile(filepath.Join(spec.RuntimeRoot, "unexpected"), []byte("x"), runtimeFileMode); err != nil { t.Fatal(err) } }, }, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { spec := testSpec(t) status := commitSnapshot(t, spec, testSnapshot(t, tc.name)) tc.mutate(t, spec, status) if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { t.Fatalf("Inspect error = %v", err) } }) } } func TestInspectRejectsUnexpectedGenerationEntry(t *testing.T) { spec := testSpec(t) status := commitSnapshot(t, spec, testSnapshot(t, "generation-entry")) path := filepath.Join(spec.RuntimeRoot, "generations", status.Snapshot.Generation, "unexpected") if err := os.WriteFile(path, []byte("x"), runtimeFileMode); err != nil { t.Fatal(err) } if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { t.Fatalf("Inspect error = %v", err) } } 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 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") }}}, } { t.Run(tc.name, func(t *testing.T) { spec := testSpec(t) transaction, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } defer transaction.Close() restore := setTestHooksForTest(tc.hooks) defer restore() if _, err := transaction.Commit(testSnapshot(t, tc.name)); !errors.Is(err, ErrIntegrity) || strings.Contains(fmt.Sprint(err), "sentinel") { t.Fatalf("Commit error = %v", err) } if _, err := Inspect(spec); !errors.Is(err, ErrBlocked) { t.Fatalf("Inspect error = %v", err) } }) } } func TestRestoreIfUnchangedRestoresPriorReadyOnlyForEqualCanonicalSnapshot(t *testing.T) { spec := testSpec(t) prior := testSnapshot(t, "prior") commitSnapshot(t, spec, prior) transaction, err := Begin(spec, &prior, true) if err != nil { t.Fatal(err) } defer transaction.Close() if err := transaction.RestoreIfUnchanged(prior); err != nil { t.Fatal(err) } if got, err := Inspect(spec); err != nil || got.Snapshot.Generation != prior.Generation { t.Fatalf("restored = %#v, %v", got, err) } if err := transaction.Close(); err != nil { t.Fatal(err) } transaction, err = Begin(spec, &prior, true) if err != nil { t.Fatal(err) } defer transaction.Close() if err := transaction.RestoreIfUnchanged(testSnapshot(t, "changed")); !errors.Is(err, ErrIntegrity) { t.Fatalf("changed restore error = %v", err) } if _, err := Inspect(spec); !errors.Is(err, ErrBlocked) { t.Fatalf("Inspect error = %v", err) } } func TestRecoveryRemovesOnlyRecordedSafeStageAndStrictCurrentTemporary(t *testing.T) { spec := testSpec(t) first, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } transactionID := first.transactionID if err := first.Close(); err != nil { t.Fatal(err) } stage := filepath.Join(spec.RuntimeRoot, "generations", ".stage-"+transactionID+"-"+testSnapshot(t, "stage").Generation) if err := os.MkdirAll(stage, runtimeDirectoryMode); err != nil { t.Fatal(err) } currentTmp := filepath.Join(spec.RuntimeRoot, ".current-"+transactionID+".tmp") currentData, err := encodeSelector(Selector{Version: SchemaVersion, State: "blocked", Transaction: transactionID}) if err != nil { t.Fatal(err) } if err := os.WriteFile(currentTmp, currentData, runtimeFileMode); err != nil { t.Fatal(err) } second, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } if err := second.Close(); err != nil { t.Fatal(err) } if _, err := os.Lstat(stage); !errors.Is(err, os.ErrNotExist) { t.Fatalf("stage remains: %v", err) } if _, err := os.Lstat(currentTmp); !errors.Is(err, os.ErrNotExist) { t.Fatalf("temporary remains: %v", err) } } func TestRecoveryRefusesUnrelatedOrUnsafeTemporaryEntries(t *testing.T) { spec := testSpec(t) transaction, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } if err := transaction.Close(); err != nil { t.Fatal(err) } other := strings.Repeat("a", 32) if other == transaction.transactionID { other = strings.Repeat("b", 32) } stage := filepath.Join(spec.RuntimeRoot, "generations", ".stage-"+other+"-"+testSnapshot(t, "unrelated").Generation) if err := os.Mkdir(stage, runtimeDirectoryMode); 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(stage); err != nil { t.Fatalf("unrelated stage removed: %v", err) } } func TestRecoveryRefusesUnsafeCurrentTemporaryEntries(t *testing.T) { for _, tc := range []struct { name string make func(t *testing.T, root, name string) }{ {"symlink", func(t *testing.T, root, name string) { if err := os.Symlink("CURRENT", filepath.Join(root, name)); err != nil { t.Fatal(err) } }}, {"hardlink", func(t *testing.T, root, name string) { if err := os.Link(filepath.Join(root, "CURRENT"), filepath.Join(root, name)); err != nil { t.Fatal(err) } }}, {"oversized", func(t *testing.T, root, name string) { if err := os.WriteFile(filepath.Join(root, name), []byte(strings.Repeat("x", maximumSelectorBytes+1)), runtimeFileMode); err != nil { t.Fatal(err) } }}, } { t.Run(tc.name, func(t *testing.T) { spec := testSpec(t) transaction, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } transactionID := transaction.transactionID if err := transaction.Close(); err != nil { t.Fatal(err) } name := ".current-" + transactionID + ".tmp" tc.make(t, spec.RuntimeRoot, name) if _, err := Begin(spec, nil, false); !errors.Is(err, ErrIntegrity) { t.Fatalf("Begin error = %v", err) } if _, err := os.Lstat(filepath.Join(spec.RuntimeRoot, name)); err != nil { t.Fatalf("unsafe temporary was removed: %v", err) } }) } } func TestRetentionKeepsCurrentAndTwoPredecessors(t *testing.T) { spec := testSpec(t) var snapshots []Snapshot for index := 0; index < 4; index++ { snapshot := testSnapshot(t, fmt.Sprint(index)) commitSnapshot(t, spec, snapshot) snapshots = append(snapshots, snapshot) } status, err := Inspect(spec) if err != nil { t.Fatal(err) } if len(status.Selector.PreviousGenerations) != retainedPredecessors { t.Fatalf("history = %#v", status.Selector.PreviousGenerations) } for _, snapshot := range snapshots[1:] { if _, err := os.Stat(filepath.Join(spec.RuntimeRoot, "generations", snapshot.Generation)); err != nil { t.Fatalf("retained generation %s: %v", snapshot.Generation, err) } } if _, err := os.Stat(filepath.Join(spec.RuntimeRoot, "generations", snapshots[0].Generation)); !errors.Is(err, os.ErrNotExist) { t.Fatalf("old generation still present: %v", err) } } func TestBeginCompletesSafeRetentionAfterReadyPublicationInterruptedBeforeCleanup(t *testing.T) { spec := testSpec(t) var snapshots []Snapshot for index := 0; index < 3; index++ { snapshot := testSnapshot(t, fmt.Sprintf("interrupted-%d", index)) commitSnapshot(t, spec, snapshot) snapshots = append(snapshots, snapshot) } fourth := testSnapshot(t, "interrupted-3") transaction, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } restore := setTestHooksForTest(testHooks{beforeRetention: func() error { return errors.New("sentinel") }}) if _, err := transaction.Commit(fourth); !errors.Is(err, ErrIntegrity) || strings.Contains(fmt.Sprint(err), "sentinel") { t.Fatalf("interrupted Commit error = %v", err) } restore() if err := transaction.Close(); err != nil { t.Fatal(err) } if _, err := Inspect(spec); !errors.Is(err, ErrIntegrity) { t.Fatalf("Inspect after interrupted retention error = %v", err) } fifth := testSnapshot(t, "interrupted-4") recovery, err := Begin(spec, &fourth, true) if err != nil { t.Fatal(err) } defer recovery.Close() status, err := recovery.Commit(fifth) if err != nil { t.Fatal(err) } if status.Snapshot.Generation != fifth.Generation { t.Fatalf("generation = %s", status.Snapshot.Generation) } if _, err := Inspect(spec); err != nil { t.Fatal(err) } if _, err := os.Lstat(filepath.Join(spec.RuntimeRoot, "generations", snapshots[0].Generation)); !errors.Is(err, os.ErrNotExist) { t.Fatalf("surplus generation remains: %v", err) } } func TestConcurrentTransactionsNeverPublishMixedGeneration(t *testing.T) { spec := testSpec(t) first := testSnapshot(t, "one") second := testSnapshot(t, "two") var group sync.WaitGroup errorsByWorker := make(chan error, 2) for _, snapshot := range []Snapshot{first, second} { group.Add(1) go func(snapshot Snapshot) { defer group.Done() transaction, err := Begin(spec, nil, false) if err == nil { _, err = transaction.Commit(snapshot) closeErr := transaction.Close() if err == nil { err = closeErr } } errorsByWorker <- err }(snapshot) } group.Wait() close(errorsByWorker) for err := range errorsByWorker { if err != nil { t.Fatal(err) } } got, err := Inspect(spec) if err != nil { t.Fatal(err) } if got.Snapshot.Generation != first.Generation && got.Snapshot.Generation != second.Generation { t.Fatalf("mixed generation = %s", got.Snapshot.Generation) } } func TestBeginCommitUsesNumericUIDGID10001WhenRoot(t *testing.T) { if os.Geteuid() != 0 { t.Skip("requires root") } spec := testSpec(t) spec.UID, spec.GID = 10001, 10001 if err := os.Chown(spec.RuntimeRoot, int(spec.UID), int(spec.GID)); err != nil { t.Fatal(err) } status := commitSnapshot(t, spec, testSnapshot(t, "numeric")) for _, path := range []string{spec.RuntimeRoot, filepath.Join(spec.RuntimeRoot, "generations"), filepath.Join(spec.RuntimeRoot, "generations", status.Snapshot.Generation), filepath.Join(spec.RuntimeRoot, "CURRENT"), filepath.Join(spec.RuntimeRoot, "generations", status.Snapshot.Generation, "auth.yaml"), filepath.Join(spec.RuntimeRoot, "generations", status.Snapshot.Generation, "users.yaml"), filepath.Join(spec.RuntimeRoot, "generations", status.Snapshot.Generation, "manifest.json")} { info, err := os.Lstat(path) if err != nil { t.Fatal(err) } stat := info.Sys().(*syscall.Stat_t) if stat.Uid != 10001 || stat.Gid != 10001 { t.Fatalf("%s ownership = %d:%d", path, stat.Uid, stat.Gid) } } } func TestConcurrentBeginSerializesAcrossCurrentRenameWithoutCreatingLockEntry(t *testing.T) { spec := testSpec(t) first, err := Begin(spec, nil, false) if err != nil { t.Fatal(err) } 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() { second, err := Begin(spec, nil, false) if err == nil { err = second.Close() } finished <- err }() if err := <-checked; err != nil { t.Fatal(err) } select { case err := <-finished: t.Fatalf("second Begin did not block: %v", err) default: } if err := first.Close(); err != nil { t.Fatal(err) } if err := <-finished; err != nil { t.Fatal(err) } entries, err := os.ReadDir(spec.RuntimeRoot) if err != nil { t.Fatal(err) } for _, entry := range entries { if strings.Contains(entry.Name(), "lock") { t.Fatalf("unexpected lock entry %q", entry.Name()) } } }