From 8a8f2c21746be1133d7047b1df000d7d33c13dcc Mon Sep 17 00:00:00 2001 From: User Date: Fri, 21 Aug 2026 21:41:51 +0200 Subject: [PATCH] feat(auth): add Linux runtime projection primitives --- tools/tht/internal/authprojection/format.go | 269 +++++++ .../internal/authprojection/format_test.go | 243 ++++++ .../authprojection/projection_linux.go | 730 ++++++++++++++++++ .../authprojection/projection_linux_test.go | 508 ++++++++++++ .../authprojection/projection_unsupported.go | 9 + tools/tht/internal/authprojection/types.go | 61 ++ 6 files changed, 1820 insertions(+) create mode 100644 tools/tht/internal/authprojection/format.go create mode 100644 tools/tht/internal/authprojection/format_test.go create mode 100644 tools/tht/internal/authprojection/projection_linux.go create mode 100644 tools/tht/internal/authprojection/projection_linux_test.go create mode 100644 tools/tht/internal/authprojection/projection_unsupported.go create mode 100644 tools/tht/internal/authprojection/types.go diff --git a/tools/tht/internal/authprojection/format.go b/tools/tht/internal/authprojection/format.go new file mode 100644 index 00000000..423680ed --- /dev/null +++ b/tools/tht/internal/authprojection/format.go @@ -0,0 +1,269 @@ +package authprojection + +import ( + "bytes" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "strings" +) + +type manifest struct { + Version int `json:"version"` + Generation string `json:"generation"` + Mode string `json:"mode"` + CanonicalRevision string `json:"canonicalRevision"` + Files []manifestFile `json:"files"` +} + +type manifestFile struct { + Name string `json:"name"` + Size int `json:"size"` + SHA256 string `json:"sha256"` +} + +func NewSnapshot(mode string, auth, users []byte) (Snapshot, error) { + if (mode != "local" && mode != "oidc") || len(auth) > maximumAuthBytes || len(users) > maximumUsersBytes { + return Snapshot{}, ErrIntegrity + } + if mode == "local" && users == nil { + return Snapshot{}, ErrIntegrity + } + if mode == "oidc" && users != nil { + return Snapshot{}, ErrIntegrity + } + + authCopy := append([]byte(nil), auth...) + usersCopy := append([]byte(nil), users...) + authDigest := sha256.Sum256(authCopy) + usersDigest := sha256.Sum256(usersCopy) + authSHA256 := hex.EncodeToString(authDigest[:]) + usersSHA256 := "-" + if mode == "local" { + usersSHA256 = hex.EncodeToString(usersDigest[:]) + } + record := "thothii-auth-projection-v1\nmode=" + mode + "\nauth=" + authSHA256 + "\nusers=" + usersSHA256 + "\n" + generationDigest := sha256.Sum256([]byte(record)) + generation := hex.EncodeToString(generationDigest[:]) + return Snapshot{ + Mode: mode, + Auth: authCopy, + Users: usersCopy, + AuthSHA256: authSHA256, + UsersSHA256: usersSHA256, + Generation: generation, + CanonicalRevision: "sha256:" + generation, + }, nil +} + +func decodeSelector(data []byte) (Selector, error) { + if len(data) == 0 || len(data) > maximumSelectorBytes || hasDuplicateObjectKeys(data) { + return Selector{}, ErrIntegrity + } + var value Selector + if err := decodeExact(data, &value); err != nil || !validSelector(value) { + return Selector{}, ErrIntegrity + } + canonical, err := encodeSelector(value) + if err != nil || !bytes.Equal(data, canonical) { + return Selector{}, ErrIntegrity + } + return value, nil +} + +func encodeSelector(value Selector) ([]byte, error) { + if !validSelector(value) { + return nil, ErrIntegrity + } + encoded, err := json.Marshal(value) + if err != nil { + return nil, fmt.Errorf("%w: selector encoding", ErrIntegrity) + } + return append(encoded, '\n'), nil +} + +func validSelector(value Selector) bool { + if value.Version != SchemaVersion || !validHex(value.Transaction, 32) { + return false + } + switch value.State { + case "blocked": + return value.Generation == "" && len(value.PreviousGenerations) == 0 + case "ready": + if !validHex(value.Generation, 64) || len(value.PreviousGenerations) > retainedPredecessors { + return false + } + seen := map[string]struct{}{value.Generation: {}} + for _, generation := range value.PreviousGenerations { + if !validHex(generation, 64) { + return false + } + if _, duplicate := seen[generation]; duplicate { + return false + } + seen[generation] = struct{}{} + } + return true + default: + return false + } +} + +func newManifest(snapshot Snapshot) (manifest, error) { + if err := validateSnapshot(snapshot); err != nil { + return manifest{}, err + } + files := []manifestFile{{Name: "auth.yaml", Size: len(snapshot.Auth), SHA256: snapshot.AuthSHA256}} + if snapshot.Mode == "local" { + files = append(files, manifestFile{Name: "users.yaml", Size: len(snapshot.Users), SHA256: snapshot.UsersSHA256}) + } + return manifest{Version: SchemaVersion, Generation: snapshot.Generation, Mode: snapshot.Mode, CanonicalRevision: snapshot.CanonicalRevision, Files: files}, nil +} + +func decodeManifest(data []byte) (manifest, error) { + if len(data) == 0 || len(data) > maximumManifestBytes || hasDuplicateObjectKeys(data) { + return manifest{}, ErrIntegrity + } + var value manifest + if err := decodeExact(data, &value); err != nil || !validManifest(value) { + return manifest{}, ErrIntegrity + } + canonical, err := encodeManifest(value) + if err != nil || !bytes.Equal(data, canonical) { + return manifest{}, ErrIntegrity + } + return value, nil +} + +func encodeManifest(value manifest) ([]byte, error) { + if !validManifest(value) { + return nil, ErrIntegrity + } + encoded, err := json.Marshal(value) + if err != nil { + return nil, fmt.Errorf("%w: manifest encoding", ErrIntegrity) + } + return append(encoded, '\n'), nil +} + +func validateManifest(value manifest, snapshot Snapshot) error { + expected, err := newManifest(snapshot) + if err != nil || !manifestEqual(value, expected) { + return ErrIntegrity + } + return nil +} + +func manifestEqual(left, right manifest) bool { + if left.Version != right.Version || left.Generation != right.Generation || left.Mode != right.Mode || left.CanonicalRevision != right.CanonicalRevision || len(left.Files) != len(right.Files) { + return false + } + for index := range left.Files { + if left.Files[index] != right.Files[index] { + return false + } + } + return true +} + +func validManifest(value manifest) bool { + if value.Version != SchemaVersion || !validHex(value.Generation, 64) || (value.Mode != "local" && value.Mode != "oidc") || value.CanonicalRevision != "sha256:"+value.Generation { + return false + } + want := []string{"auth.yaml"} + if value.Mode == "local" { + want = append(want, "users.yaml") + } + if len(value.Files) != len(want) { + return false + } + for index, file := range value.Files { + maximum := maximumAuthBytes + if file.Name == "users.yaml" { + maximum = maximumUsersBytes + } + if file.Name != want[index] || file.Size < 0 || file.Size > maximum || !validHex(file.SHA256, 64) { + return false + } + } + return true +} + +func validateSnapshot(snapshot Snapshot) error { + expected, err := NewSnapshot(snapshot.Mode, snapshot.Auth, snapshot.Users) + if err != nil || expected.AuthSHA256 != snapshot.AuthSHA256 || expected.UsersSHA256 != snapshot.UsersSHA256 || expected.Generation != snapshot.Generation || expected.CanonicalRevision != snapshot.CanonicalRevision { + return ErrIntegrity + } + return nil +} + +func decodeExact(data []byte, destination any) error { + decoder := json.NewDecoder(bytes.NewReader(data)) + decoder.DisallowUnknownFields() + if err := decoder.Decode(destination); err != nil { + return err + } + if err := decoder.Decode(&struct{}{}); err != io.EOF { + return fmt.Errorf("multiple JSON values") + } + return nil +} + +func hasDuplicateObjectKeys(data []byte) bool { + decoder := json.NewDecoder(bytes.NewReader(data)) + return scanJSONValue(decoder, nil) != nil || decoder.Decode(&struct{}{}) != io.EOF +} + +func scanJSONValue(decoder *json.Decoder, seen map[string]struct{}) error { + token, err := decoder.Token() + if err != nil { + return err + } + delimiter, ok := token.(json.Delim) + if !ok { + return nil + } + switch delimiter { + case '{': + keys := make(map[string]struct{}) + for decoder.More() { + keyToken, err := decoder.Token() + if err != nil { + return err + } + key, ok := keyToken.(string) + if !ok { + return fmt.Errorf("object key") + } + if _, exists := keys[key]; exists { + return fmt.Errorf("duplicate object key") + } + keys[key] = struct{}{} + if err := scanJSONValue(decoder, keys); err != nil { + return err + } + } + _, err := decoder.Token() + return err + case '[': + for decoder.More() { + if err := scanJSONValue(decoder, nil); err != nil { + return err + } + } + _, err := decoder.Token() + return err + default: + return fmt.Errorf("unexpected delimiter") + } +} + +func validHex(value string, length int) bool { + if len(value) != length || strings.ToLower(value) != value { + return false + } + _, err := hex.DecodeString(value) + return err == nil +} diff --git a/tools/tht/internal/authprojection/format_test.go b/tools/tht/internal/authprojection/format_test.go new file mode 100644 index 00000000..e6f35f24 --- /dev/null +++ b/tools/tht/internal/authprojection/format_test.go @@ -0,0 +1,243 @@ +package authprojection + +import ( + "bytes" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "strings" + "testing" +) + +func TestNewSnapshotUsesDomainSeparatedGeneration(t *testing.T) { + got, err := NewSnapshot("local", []byte("auth\n"), []byte("users\n")) + if err != nil { + t.Fatal(err) + } + auth := sha256.Sum256([]byte("auth\n")) + users := sha256.Sum256([]byte("users\n")) + record := fmt.Sprintf("thothii-auth-projection-v1\nmode=local\nauth=%x\nusers=%x\n", auth, users) + generation := sha256.Sum256([]byte(record)) + if got.Generation != hex.EncodeToString(generation[:]) { + t.Fatalf("generation = %q", got.Generation) + } + if got.CanonicalRevision != "sha256:"+got.Generation { + t.Fatalf("revision = %q", got.CanonicalRevision) + } + if string(got.Auth) != "auth\n" || string(got.Users) != "users\n" { + t.Fatal("snapshot did not preserve supplied bytes") + } +} + +func TestNewSnapshotRejectsInvalidModeAndUsersShape(t *testing.T) { + for _, tc := range []struct { + mode string + users []byte + }{ + {mode: "unknown", users: nil}, + {mode: "local", users: nil}, + {mode: "oidc", users: []byte("unexpected")}, + } { + if _, err := NewSnapshot(tc.mode, []byte("auth"), tc.users); !errors.Is(err, ErrIntegrity) { + t.Fatalf("NewSnapshot(%q) error = %v", tc.mode, err) + } + } +} + +func TestNewSnapshotCopiesAndBoundsInput(t *testing.T) { + auth := []byte("auth") + users := []byte("users") + got, err := NewSnapshot("local", auth, users) + if err != nil { + t.Fatal(err) + } + auth[0] = 'x' + users[0] = 'x' + if string(got.Auth) != "auth" || string(got.Users) != "users" { + t.Fatal("snapshot aliases caller buffers") + } + if _, err := NewSnapshot("oidc", make([]byte, maximumAuthBytes+1), nil); !errors.Is(err, ErrIntegrity) { + t.Fatalf("oversized auth error = %v", err) + } + if _, err := NewSnapshot("local", []byte("auth"), make([]byte, maximumUsersBytes+1)); !errors.Is(err, ErrIntegrity) { + t.Fatalf("oversized users error = %v", err) + } +} + +func TestSelectorStrictlyDecodesReadyAndBlocked(t *testing.T) { + tx := strings.Repeat("a", 32) + gen := strings.Repeat("b", 64) + previous := strings.Repeat("c", 64) + for _, tc := range []struct { + name string + body string + want Selector + }{ + { + name: "ready", + body: fmt.Sprintf(`{"version":1,"state":"ready","transaction":"%s","generation":"%s","previousGenerations":["%s"]}`, tx, gen, previous), + want: Selector{Version: 1, State: "ready", Transaction: tx, Generation: gen, PreviousGenerations: []string{previous}}, + }, + { + name: "blocked", + body: fmt.Sprintf(`{"version":1,"state":"blocked","transaction":"%s"}`, tx), + want: Selector{Version: 1, State: "blocked", Transaction: tx}, + }, + } { + t.Run(tc.name, func(t *testing.T) { + got, err := decodeSelector(append([]byte(tc.body), '\n')) + if err != nil { + t.Fatal(err) + } + if fmt.Sprintf("%#v", got) != fmt.Sprintf("%#v", tc.want) { + t.Fatalf("selector = %#v, want %#v", got, tc.want) + } + encoded, err := encodeSelector(got) + if err != nil { + t.Fatal(err) + } + if !strings.HasSuffix(string(encoded), "\n") || strings.Count(string(encoded), "\n") != 1 { + t.Fatalf("encoded selector is not canonical newline-terminated JSON: %q", encoded) + } + }) + } +} + +func TestSelectorRejectsDuplicatesUnknownsAndInvalidIdentifiers(t *testing.T) { + tx := strings.Repeat("a", 32) + gen := strings.Repeat("b", 64) + cases := []string{ + fmt.Sprintf(`{"version":1,"version":1,"state":"blocked","transaction":"%s"}`, tx), + fmt.Sprintf(`{"version":1,"state":"blocked","transaction":"%s","extra":true}`, tx), + fmt.Sprintf(`{"version":1,"state":"ready","transaction":"%s","generation":"%s","previousGenerations":["%s","%s"]}`, tx, gen, gen, gen), + fmt.Sprintf(`{"version":1,"state":"ready","transaction":"%s","generation":"%s","previousGenerations":["%s","%s","%s"]}`, tx, gen, strings.Repeat("c", 64), strings.Repeat("d", 64), strings.Repeat("e", 64)), + fmt.Sprintf(`{"version":1,"state":"ready","transaction":"%s","generation":"%s"}`, strings.Repeat("A", 32), gen), + fmt.Sprintf(`{"version":1,"state":"blocked","transaction":"%s","generation":"%s"}`, tx, gen), + fmt.Sprintf(`{"version":1,"state":"ready","transaction":"%s","generation":"%s"} trailing`, tx, gen), + } + for _, body := range cases { + if _, err := decodeSelector([]byte(body)); !errors.Is(err, ErrIntegrity) { + t.Fatalf("decodeSelector(%q) error = %v", body, err) + } + } +} + +func TestManifestStrictlyValidatesExactModeDependentFiles(t *testing.T) { + snapshot, err := NewSnapshot("local", []byte("auth"), []byte("users")) + if err != nil { + t.Fatal(err) + } + manifest, err := newManifest(snapshot) + if err != nil { + t.Fatal(err) + } + encoded, err := encodeManifest(manifest) + if err != nil { + t.Fatal(err) + } + if !strings.HasSuffix(string(encoded), "\n") || strings.Count(string(encoded), "\n") != 1 { + t.Fatalf("manifest is not canonical newline-terminated JSON: %q", encoded) + } + got, err := decodeManifest(encoded) + if err != nil { + t.Fatal(err) + } + if err := validateManifest(got, snapshot); err != nil { + t.Fatal(err) + } + bad := strings.Replace(string(encoded), `"users.yaml"`, `"other.yaml"`, 1) + if _, err := decodeManifest([]byte(bad)); !errors.Is(err, ErrIntegrity) { + t.Fatalf("bad manifest error = %v", err) + } +} + +func TestSelectorAndManifestRequireCanonicalNewlineTerminatedJSON(t *testing.T) { + selector := Selector{Version: SchemaVersion, State: "blocked", Transaction: strings.Repeat("a", 32)} + selectorData, err := encodeSelector(selector) + if err != nil { + t.Fatal(err) + } + if _, err := decodeSelector(bytes.TrimSuffix(selectorData, []byte("\n"))); !errors.Is(err, ErrIntegrity) { + t.Fatalf("noncanonical selector error = %v", err) + } + snapshot, err := NewSnapshot("oidc", []byte("auth"), nil) + if err != nil { + t.Fatal(err) + } + value, err := newManifest(snapshot) + if err != nil { + t.Fatal(err) + } + manifestData, err := encodeManifest(value) + if err != nil { + t.Fatal(err) + } + if _, err := decodeManifest(bytes.TrimSuffix(manifestData, []byte("\n"))); !errors.Is(err, ErrIntegrity) { + t.Fatalf("noncanonical manifest error = %v", err) + } +} + +func TestManifestRejectsOIDCUsersAndSnapshotDigestOrSizeMismatch(t *testing.T) { + snapshot, err := NewSnapshot("oidc", []byte("auth"), nil) + if err != nil { + t.Fatal(err) + } + value, err := newManifest(snapshot) + if err != nil { + t.Fatal(err) + } + withUsers := value + withUsers.Files = append(withUsers.Files, manifestFile{Name: "users.yaml", Size: 1, SHA256: strings.Repeat("a", 64)}) + data, err := json.Marshal(withUsers) + if err != nil { + t.Fatal(err) + } + if _, err := decodeManifest(append(data, '\n')); !errors.Is(err, ErrIntegrity) { + t.Fatalf("oidc users error = %v", err) + } + for _, mutate := range []func(*manifest){ + func(value *manifest) { value.Files[0].Size++ }, + func(value *manifest) { value.Files[0].SHA256 = strings.Repeat("a", 64) }, + } { + changed := value + changed.Files = append([]manifestFile(nil), value.Files...) + mutate(&changed) + if err := validateManifest(changed, snapshot); !errors.Is(err, ErrIntegrity) { + t.Fatalf("changed manifest error = %v", err) + } + } +} + +func TestManifestRejectsFileSizesAboveReadBounds(t *testing.T) { + for _, tc := range []struct { + name string + mode string + users []byte + file int + tooLarge int + }{ + {name: "auth", mode: "oidc", file: 0, tooLarge: maximumAuthBytes + 1}, + {name: "users", mode: "local", users: []byte("users"), file: 1, tooLarge: maximumUsersBytes + 1}, + } { + t.Run(tc.name, func(t *testing.T) { + snapshot, err := NewSnapshot(tc.mode, []byte("auth"), tc.users) + if err != nil { + t.Fatal(err) + } + value, err := newManifest(snapshot) + if err != nil { + t.Fatal(err) + } + value.Files[tc.file].Size = tc.tooLarge + data, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + if _, err := decodeManifest(append(data, '\n')); !errors.Is(err, ErrIntegrity) { + t.Fatalf("oversized %s manifest error = %v", tc.name, err) + } + }) + } +} diff --git a/tools/tht/internal/authprojection/projection_linux.go b/tools/tht/internal/authprojection/projection_linux.go new file mode 100644 index 00000000..246549a6 --- /dev/null +++ b/tools/tht/internal/authprojection/projection_linux.go @@ -0,0 +1,730 @@ +//go:build linux + +package authprojection + +import ( + "crypto/rand" + "encoding/hex" + "errors" + "path/filepath" + "sync" + "syscall" + + "golang.org/x/sys/unix" +) + +const maximumDirectoryEntries = 16 + +type testHooks struct { + beforeStageRename func() error + beforeCurrentRename func() error + beforeFinalVerify func() error + beforeRetention func() error +} + +var hookState struct { + sync.Mutex + hooks testHooks +} + +func setTestHooksForTest(hooks testHooks) func() { + hookState.Lock() + previous := hookState.hooks + hookState.hooks = hooks + hookState.Unlock() + return func() { + hookState.Lock() + hookState.hooks = previous + hookState.Unlock() + } +} + +func hook(selectHook func(testHooks) func() error) error { + hookState.Lock() + callback := selectHook(hookState.hooks) + hookState.Unlock() + if callback != nil && callback() != nil { + return ErrIntegrity + } + return nil +} + +func Inspect(spec Spec) (Status, error) { + root, err := openRoot(spec) + if err != nil { + return Status{}, err + } + defer unix.Close(root) + return inspectFD(root, spec) +} + +func Begin(spec Spec, before *Snapshot, requireReadyMatch bool) (*Transaction, error) { + root, err := openRoot(spec) + if err != nil { + return nil, err + } + 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": + if requireReadyMatch || recoverTemporary(root, spec, selector) != nil { + _ = transaction.Close() + return nil, ErrIntegrity + } + case selectorErr == nil: + if validateRootEntries(root, spec) != nil { + _ = transaction.Close() + return nil, ErrIntegrity + } + status, inspectErr := inspectCapturedSelector(root, spec, selector) + if inspectErr != nil { + _ = transaction.Close() + return nil, ErrIntegrity + } + transaction.prior = &status + if requireReadyMatch && (before == nil || before.Generation != status.Snapshot.Generation) { + _ = transaction.Close() + return nil, ErrIntegrity + } + case currentExists(root): + _ = transaction.Close() + return nil, ErrIntegrity + case requireReadyMatch || validateInitialRootEntries(root, spec) != nil: + _ = transaction.Close() + return nil, ErrIntegrity + } + + transaction.transactionID, err = randomID() + if err != nil || writeSelector(root, spec, Selector{Version: SchemaVersion, State: "blocked", Transaction: transaction.transactionID}) != nil { + _ = transaction.Close() + return nil, ErrIntegrity + } + transaction.blocked = true + return transaction, nil +} + +func (transaction *Transaction) Commit(after Snapshot) (Status, error) { + if transaction == nil || transaction.closed || !transaction.blocked || validateSnapshot(after) != nil { + return Status{}, ErrIntegrity + } + if stageGeneration(transaction.rootFD, transaction.spec, transaction.transactionID, after) != nil { + return Status{}, ErrIntegrity + } + if err := hook(func(h testHooks) func() error { return h.beforeFinalVerify }); err != nil { + return Status{}, err + } + previous := make([]string, 0, retainedPredecessors) + if transaction.prior != nil { + previous = append(previous, transaction.prior.Selector.Generation) + previous = append(previous, transaction.prior.Selector.PreviousGenerations...) + } + previous = uniqueGenerations(previous, after.Generation) + if len(previous) > retainedPredecessors { + previous = previous[:retainedPredecessors] + } + selector := Selector{Version: SchemaVersion, State: "ready", Transaction: transaction.transactionID, Generation: after.Generation, PreviousGenerations: previous} + if writeSelector(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 + } + transaction.blocked = false + if err := hook(func(h testHooks) func() error { return h.beforeRetention }); err != nil { + return Status{}, err + } + if retain(transaction.rootFD, transaction.spec, selector) != nil { + return Status{}, ErrIntegrity + } + return status, nil +} + +func (transaction *Transaction) RestoreIfUnchanged(current Snapshot) error { + if transaction == nil || transaction.closed || !transaction.blocked || transaction.before == nil || transaction.prior == nil || validateSnapshot(current) != nil || current.Generation != transaction.before.Generation { + return ErrIntegrity + } + if writeSelector(transaction.rootFD, transaction.spec, transaction.prior.Selector) != nil { + return ErrIntegrity + } + status, err := inspectFD(transaction.rootFD, transaction.spec) + if err != nil || status.Snapshot.Generation != transaction.before.Generation { + return ErrIntegrity + } + transaction.blocked = false + return nil +} + +func (transaction *Transaction) Close() error { + if transaction == nil || transaction.closed { + return nil + } + transaction.closed = true + err := unix.Flock(transaction.lockFD, unix.LOCK_UN) + if closeErr := unix.Close(transaction.rootFD); err == nil { + err = closeErr + } + if err != nil { + return ErrIntegrity + } + return nil +} + +func openRoot(spec Spec) (int, error) { + if !filepath.IsAbs(spec.RuntimeRoot) || filepath.Clean(spec.RuntimeRoot) != spec.RuntimeRoot { + return -1, ErrIntegrity + } + fd, err := unix.Open(spec.RuntimeRoot, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0) + if err != nil { + return -1, ErrIntegrity + } + if validateDirectoryFD(fd, spec) != nil { + unix.Close(fd) + return -1, ErrIntegrity + } + return fd, nil +} + +func validDirectoryStat(stat unix.Stat_t, spec Spec) bool { + return stat.Mode&unix.S_IFMT == unix.S_IFDIR && stat.Uid == spec.UID && stat.Gid == spec.GID && stat.Mode&0o7777 == runtimeDirectoryMode +} +func validFileStat(stat unix.Stat_t, spec Spec) bool { + return stat.Mode&unix.S_IFMT == unix.S_IFREG && stat.Nlink == 1 && stat.Uid == spec.UID && stat.Gid == spec.GID && stat.Mode&0o7777 == runtimeFileMode +} +func sameStat(left, right unix.Stat_t) bool { + return left.Dev == right.Dev && left.Ino == right.Ino && left.Mode == right.Mode && left.Nlink == right.Nlink && left.Uid == right.Uid && left.Gid == right.Gid && left.Size == right.Size && left.Mtim.Sec == right.Mtim.Sec && left.Mtim.Nsec == right.Mtim.Nsec && left.Ctim.Sec == right.Ctim.Sec && left.Ctim.Nsec == right.Ctim.Nsec +} +func validateDirectoryFD(fd int, spec Spec) error { + var stat unix.Stat_t + if unix.Fstat(fd, &stat) != nil || !validDirectoryStat(stat, spec) { + return ErrIntegrity + } + return nil +} +func validateFileFD(fd int, spec Spec) error { + var stat unix.Stat_t + if unix.Fstat(fd, &stat) != nil || !validFileStat(stat, spec) { + return ErrIntegrity + } + return nil +} +func statAt(directory int, name string) (unix.Stat_t, error) { + var stat unix.Stat_t + if unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW) != nil { + return unix.Stat_t{}, ErrIntegrity + } + return stat, nil +} +func currentExists(root int) bool { + var stat unix.Stat_t + return unix.Fstatat(root, "CURRENT", &stat, unix.AT_SYMLINK_NOFOLLOW) == nil +} + +func ensureGenerations(root int, spec Spec) error { + var stat unix.Stat_t + err := unix.Fstatat(root, "generations", &stat, unix.AT_SYMLINK_NOFOLLOW) + if errors.Is(err, unix.ENOENT) { + if unix.Mkdirat(root, "generations", runtimeDirectoryMode) != nil { + return ErrIntegrity + } + fd, openErr := unix.Openat(root, "generations", unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0) + if openErr != nil { + return ErrIntegrity + } + defer unix.Close(fd) + if unix.Fchown(fd, int(spec.UID), int(spec.GID)) != nil || unix.Fchmod(fd, runtimeDirectoryMode) != nil || validateDirectoryFD(fd, spec) != nil || unix.Fsync(fd) != nil || unix.Fsync(root) != nil { + return ErrIntegrity + } + return nil + } + if err != nil || !validDirectoryStat(stat, spec) { + return ErrIntegrity + } + return nil +} + +func openDirectoryAt(parent int, name string, spec Spec) (int, error) { + before, err := statAt(parent, name) + if err != nil || !validDirectoryStat(before, spec) { + return -1, ErrIntegrity + } + fd, err := unix.Openat(parent, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0) + if err != nil { + return -1, ErrIntegrity + } + var opened unix.Stat_t + if unix.Fstat(fd, &opened) != nil || !sameStat(before, opened) || !validDirectoryStat(opened, spec) { + unix.Close(fd) + return -1, ErrIntegrity + } + return fd, nil +} + +func scanDirectoryNames(directory int, spec Spec) (map[string]bool, error) { + var original unix.Stat_t + if unix.Fstat(directory, &original) != nil || !validDirectoryStat(original, spec) { + return nil, ErrIntegrity + } + copyFD, err := unix.Openat(directory, ".", unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0) + if err != nil { + return nil, ErrIntegrity + } + defer unix.Close(copyFD) + var copied unix.Stat_t + if unix.Fstat(copyFD, &copied) != nil || !sameStat(original, copied) || !validDirectoryStat(copied, spec) { + return nil, ErrIntegrity + } + names := make(map[string]bool) + buffer := make([]byte, 4096) + for { + read, err := syscall.ReadDirent(copyFD, buffer) + if err != nil { + return nil, ErrIntegrity + } + if read == 0 { + return names, nil + } + _, count, parsed := syscall.ParseDirent(buffer[:read], maximumDirectoryEntries+1, nil) + if count != len(parsed) || len(names)+len(parsed) > maximumDirectoryEntries { + return nil, ErrIntegrity + } + for _, name := range parsed { + if name == "" || name == "." || name == ".." || names[name] { + return nil, ErrIntegrity + } + names[name] = true + } + } +} + +func validateRootEntries(root int, spec Spec) error { + names, err := scanDirectoryNames(root, spec) + if err != nil || len(names) != 2 || !names["CURRENT"] || !names["generations"] { + return ErrIntegrity + } + return nil +} +func validateInitialRootEntries(root int, spec Spec) error { + names, err := scanDirectoryNames(root, spec) + if err != nil || len(names) != 1 || !names["generations"] { + return ErrIntegrity + } + return nil +} + +func inspectFD(root int, spec Spec) (Status, error) { + if validateRootEntries(root, spec) != nil { + return Status{}, ErrIntegrity + } + selector, err := readSelector(root, spec) + if err != nil { + return Status{}, err + } + return inspectCapturedSelector(root, spec, selector) +} +func inspectCapturedSelector(root int, spec Spec, selector Selector) (Status, error) { + if selector.State == "blocked" { + return Status{Selector: selector}, ErrBlocked + } + snapshot, err := readGeneration(root, spec, selector.Generation) + if err != nil { + return Status{}, err + } + return Status{Selector: selector, Snapshot: snapshot}, nil +} + +func readSelector(root int, spec Spec) (Selector, error) { + data, err := readRegularAt(root, "CURRENT", spec, maximumSelectorBytes) + if err != nil { + return Selector{}, err + } + return decodeSelector(data) +} +func readGeneration(root int, spec Spec, generation string) (Snapshot, error) { + if !validHex(generation, 64) { + return Snapshot{}, ErrIntegrity + } + generations, err := openDirectoryAt(root, "generations", spec) + if err != nil { + return Snapshot{}, err + } + defer unix.Close(generations) + directory, err := openDirectoryAt(generations, generation, spec) + if err != nil { + return Snapshot{}, err + } + defer unix.Close(directory) + manifestData, err := readRegularAt(directory, "manifest.json", spec, maximumManifestBytes) + if err != nil { + return Snapshot{}, err + } + manifest, err := decodeManifest(manifestData) + if err != nil { + return Snapshot{}, err + } + if !expectedGenerationEntries(directory, spec, manifest.Mode) { + return Snapshot{}, ErrIntegrity + } + auth, err := readRegularAt(directory, "auth.yaml", spec, maximumAuthBytes) + if err != nil { + return Snapshot{}, err + } + var users []byte + if manifest.Mode == "local" { + users, err = readRegularAt(directory, "users.yaml", spec, maximumUsersBytes) + if err != nil { + return Snapshot{}, err + } + } + snapshot, err := NewSnapshot(manifest.Mode, auth, users) + if err != nil || snapshot.Generation != generation || validateManifest(manifest, snapshot) != nil { + return Snapshot{}, ErrIntegrity + } + return snapshot, nil +} +func expectedGenerationEntries(directory int, spec Spec, mode string) bool { + names, err := scanDirectoryNames(directory, spec) + if err != nil { + return false + } + want := map[string]bool{"auth.yaml": true, "manifest.json": true} + if mode == "local" { + want["users.yaml"] = true + } + if len(names) != len(want) { + return false + } + for name := range want { + if !names[name] { + return false + } + } + return true +} + +func readRegularAt(directory int, name string, spec Spec, limit int) ([]byte, error) { + before, err := statAt(directory, name) + if err != nil || !validFileStat(before, spec) { + return nil, ErrIntegrity + } + fd, err := unix.Openat(directory, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0) + if err != nil { + return nil, ErrIntegrity + } + defer unix.Close(fd) + var opened unix.Stat_t + if unix.Fstat(fd, &opened) != nil || !sameStat(before, opened) || !validFileStat(opened, spec) { + return nil, ErrIntegrity + } + data, err := readBounded(fd, limit) + if err != nil { + return nil, ErrIntegrity + } + var after unix.Stat_t + namedAfter, statErr := statAt(directory, name) + if unix.Fstat(fd, &after) != nil || statErr != nil || !sameStat(opened, after) || !sameStat(before, namedAfter) || !validFileStat(after, spec) { + return nil, ErrIntegrity + } + return data, nil +} +func readBounded(fd, limit int) ([]byte, error) { + data, buffer := make([]byte, 0, limit), make([]byte, 32*1024) + for { + read, err := unix.Read(fd, buffer) + if read > 0 { + if len(data)+read > limit { + return nil, ErrIntegrity + } + data = append(data, buffer[:read]...) + } + if err == nil { + if read == 0 { + return data, nil + } + continue + } + if errors.Is(err, unix.EINTR) { + continue + } + return nil, ErrIntegrity + } +} + +func writeSelector(root int, spec Spec, selector Selector) error { + data, err := encodeSelector(selector) + if err != nil { + return ErrIntegrity + } + name := ".current-" + selector.Transaction + ".tmp" + if writeRegularAt(root, name, spec, data) != nil { + return ErrIntegrity + } + if err := hook(func(h testHooks) func() error { return h.beforeCurrentRename }); err != nil { + _ = unix.Unlinkat(root, name, 0) + return err + } + if unix.Renameat(root, name, root, "CURRENT") != nil || unix.Fsync(root) != nil { + return ErrIntegrity + } + return nil +} +func writeRegularAt(directory int, name string, spec Spec, data []byte) error { + fd, err := unix.Openat(directory, name, unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, runtimeFileMode) + if err != nil { + return ErrIntegrity + } + defer unix.Close(fd) + if unix.Fchown(fd, int(spec.UID), int(spec.GID)) != nil || unix.Fchmod(fd, runtimeFileMode) != nil || validateFileFD(fd, spec) != nil || writeAll(fd, data) != nil || unix.Fsync(fd) != nil || validateFileFD(fd, spec) != nil { + return ErrIntegrity + } + return nil +} +func writeAll(fd int, data []byte) error { + for len(data) > 0 { + written, err := unix.Write(fd, data) + if err != nil { + if errors.Is(err, unix.EINTR) { + continue + } + return ErrIntegrity + } + if written == 0 { + return ErrIntegrity + } + data = data[written:] + } + return nil +} + +func stageGeneration(root int, spec Spec, transactionID string, snapshot Snapshot) error { + generations, err := openDirectoryAt(root, "generations", spec) + if err != nil { + return err + } + defer unix.Close(generations) + var existing unix.Stat_t + if err := unix.Fstatat(generations, snapshot.Generation, &existing, unix.AT_SYMLINK_NOFOLLOW); err == nil { + if found, readErr := readGeneration(root, spec, snapshot.Generation); readErr == nil && found.Generation == snapshot.Generation { + return nil + } + if removeConfinedDirectory(generations, snapshot.Generation, spec, false) != nil { + return ErrIntegrity + } + } else if !errors.Is(err, unix.ENOENT) { + return ErrIntegrity + } + stageName := ".stage-" + transactionID + "-" + snapshot.Generation + if unix.Mkdirat(generations, stageName, runtimeDirectoryMode) != nil { + return ErrIntegrity + } + stage, err := unix.Openat(generations, stageName, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0) + if err != nil { + return ErrIntegrity + } + defer unix.Close(stage) + if unix.Fchown(stage, int(spec.UID), int(spec.GID)) != nil || unix.Fchmod(stage, runtimeDirectoryMode) != nil || validateDirectoryFD(stage, spec) != nil || unix.Fsync(stage) != nil || unix.Fsync(generations) != nil { + return ErrIntegrity + } + var namedStage, openedStage unix.Stat_t + if unix.Fstatat(generations, stageName, &namedStage, unix.AT_SYMLINK_NOFOLLOW) != nil || unix.Fstat(stage, &openedStage) != nil || !sameStat(namedStage, openedStage) || !validDirectoryStat(namedStage, spec) { + return ErrIntegrity + } + manifest, err := newManifest(snapshot) + if err != nil { + return ErrIntegrity + } + manifestData, err := encodeManifest(manifest) + 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 { + return ErrIntegrity + } + if err := hook(func(h testHooks) func() error { return h.beforeStageRename }); err != nil { + return err + } + if unix.Renameat(generations, stageName, generations, snapshot.Generation) != nil || unix.Fsync(generations) != nil { + return ErrIntegrity + } + verified, err := readGeneration(root, spec, snapshot.Generation) + if err != nil || verified.Generation != snapshot.Generation { + return ErrIntegrity + } + return nil +} + +func retain(root int, spec Spec, selector Selector) error { + generations, err := openDirectoryAt(root, "generations", spec) + if err != nil { + return err + } + 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 + } + for name := range names { + if !validHex(name, 64) { + return ErrIntegrity + } + if keep[name] { + if _, err := readGeneration(root, spec, name); err != nil { + return ErrIntegrity + } + continue + } + if removeConfinedDirectory(generations, name, spec, false) != nil { + return ErrIntegrity + } + } + if unix.Fsync(generations) != nil || unix.Fsync(root) != nil { + return ErrIntegrity + } + return nil +} + +func recoverTemporary(root int, spec Spec, blocked Selector) error { + if blocked.State != "blocked" || !validHex(blocked.Transaction, 32) { + return ErrIntegrity + } + generations, err := openDirectoryAt(root, "generations", spec) + if err != nil { + return err + } + defer unix.Close(generations) + names, err := scanDirectoryNames(generations, spec) + if err != nil { + return ErrIntegrity + } + for name := range names { + if validHex(name, 64) { + directory, openErr := openDirectoryAt(generations, name, spec) + if openErr != nil || unix.Close(directory) != nil { + return ErrIntegrity + } + continue + } + generation, ok := stageGenerationName(name, blocked.Transaction) + if !ok || !validHex(generation, 64) || removeConfinedDirectory(generations, name, spec, true) != nil { + return ErrIntegrity + } + } + rootNames, err := scanDirectoryNames(root, spec) + if err != nil || !rootNames["CURRENT"] || !rootNames["generations"] { + return ErrIntegrity + } + for name := range rootNames { + if name == "CURRENT" || name == "generations" { + continue + } + transactionID, ok := currentTemporaryName(name) + if !ok { + return ErrIntegrity + } + data, readErr := readRegularAt(root, name, spec, maximumSelectorBytes) + if readErr != nil { + return ErrIntegrity + } + selector, decodeErr := decodeSelector(data) + if decodeErr != nil || selector.Transaction != transactionID || unix.Unlinkat(root, name, 0) != nil { + return ErrIntegrity + } + } + if unix.Fsync(generations) != nil || unix.Fsync(root) != nil { + return ErrIntegrity + } + return nil +} +func stageGenerationName(name, transactionID string) (string, bool) { + prefix := ".stage-" + transactionID + "-" + if len(name) != len(prefix)+64 || name[:len(prefix)] != prefix { + return "", false + } + return name[len(prefix):], true +} +func currentTemporaryName(name string) (string, bool) { + const prefix, suffix = ".current-", ".tmp" + if len(name) != len(prefix)+32+len(suffix) || name[:len(prefix)] != prefix || name[len(name)-len(suffix):] != suffix { + return "", false + } + id := name[len(prefix) : len(prefix)+32] + return id, validHex(id, 32) +} + +func removeConfinedDirectory(parent int, name string, spec Spec, partial bool) error { + directory, err := openDirectoryAt(parent, name, spec) + if err != nil { + return ErrIntegrity + } + defer unix.Close(directory) + names, err := scanDirectoryNames(directory, spec) + if err != nil || len(names) > 3 { + return ErrIntegrity + } + allowed := map[string]bool{"auth.yaml": true, "users.yaml": true, "manifest.json": true} + if !partial && !completeGenerationNames(names) { + return ErrIntegrity + } + for entry := range names { + if !allowed[entry] || verifyRegularAt(directory, entry, spec) != nil || unix.Unlinkat(directory, entry, 0) != nil { + return ErrIntegrity + } + } + if unix.Fsync(directory) != nil || unix.Unlinkat(parent, name, unix.AT_REMOVEDIR) != nil || unix.Fsync(parent) != nil { + return ErrIntegrity + } + return nil +} +func completeGenerationNames(names map[string]bool) bool { + if len(names) == 2 { + return names["auth.yaml"] && names["manifest.json"] + } + return len(names) == 3 && names["auth.yaml"] && names["users.yaml"] && names["manifest.json"] +} +func verifyRegularAt(directory int, name string, spec Spec) error { + before, err := statAt(directory, name) + if err != nil || !validFileStat(before, spec) { + return ErrIntegrity + } + fd, err := unix.Openat(directory, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0) + if err != nil { + return ErrIntegrity + } + defer unix.Close(fd) + var after unix.Stat_t + if unix.Fstat(fd, &after) != nil || !sameStat(before, after) || !validFileStat(after, spec) { + return ErrIntegrity + } + return nil +} +func uniqueGenerations(values []string, current string) []string { + seen := map[string]bool{current: true} + result := make([]string, 0, retainedPredecessors) + for _, value := range values { + if validHex(value, 64) && !seen[value] { + seen[value] = true + result = append(result, value) + } + } + return result +} +func randomID() (string, error) { + value := make([]byte, 16) + if _, err := rand.Read(value); err != nil { + return "", ErrIntegrity + } + return hex.EncodeToString(value), nil +} diff --git a/tools/tht/internal/authprojection/projection_linux_test.go b/tools/tht/internal/authprojection/projection_linux_test.go new file mode 100644 index 00000000..18835066 --- /dev/null +++ b/tools/tht/internal/authprojection/projection_linux_test.go @@ -0,0 +1,508 @@ +//go:build linux + +package authprojection + +import ( + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "sync" + "syscall" + "testing" + "time" + + "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) + } + if err := os.Remove(filepath.Join(spec.RuntimeRoot, "CURRENT")); err != nil { + t.Fatal(err) + } + 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 TestCommitLeavesBlockedAfterInjectedStageWriteFsyncRenameAndVerifyFailure(t *testing.T) { + for _, tc := range []struct { + name string + hooks testHooks + }{ + {"stage", 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 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) + } + started := make(chan struct{}) + finished := make(chan error, 1) + go func() { + close(started) + second, err := Begin(spec, nil, false) + if err == nil { + err = second.Close() + } + finished <- err + }() + <-started + select { + case err := <-finished: + t.Fatalf("second Begin did not block: %v", err) + case <-time.After(50 * time.Millisecond): + } + 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()) + } + } +} diff --git a/tools/tht/internal/authprojection/projection_unsupported.go b/tools/tht/internal/authprojection/projection_unsupported.go new file mode 100644 index 00000000..e576c6d8 --- /dev/null +++ b/tools/tht/internal/authprojection/projection_unsupported.go @@ -0,0 +1,9 @@ +//go:build !linux + +package authprojection + +func Inspect(Spec) (Status, error) { return Status{}, ErrUnsupported } +func Begin(Spec, *Snapshot, bool) (*Transaction, error) { return nil, ErrUnsupported } +func (*Transaction) Commit(Snapshot) (Status, error) { return Status{}, ErrUnsupported } +func (*Transaction) RestoreIfUnchanged(Snapshot) error { return ErrUnsupported } +func (*Transaction) Close() error { return ErrUnsupported } diff --git a/tools/tht/internal/authprojection/types.go b/tools/tht/internal/authprojection/types.go new file mode 100644 index 00000000..55866ce2 --- /dev/null +++ b/tools/tht/internal/authprojection/types.go @@ -0,0 +1,61 @@ +package authprojection + +import "errors" + +const SchemaVersion = 1 + +const ( + maximumAuthBytes = 1 << 20 + maximumUsersBytes = 1 << 20 + maximumSelectorBytes = 4096 + maximumManifestBytes = 4096 + runtimeDirectoryMode = 0o700 + runtimeFileMode = 0o600 + retainedPredecessors = 2 +) + +var ( + ErrBlocked = errors.New("authentication runtime projection is blocked") + ErrIntegrity = errors.New("authentication runtime projection integrity failure") + ErrUnsupported = errors.New("authentication runtime projection is unsupported") +) + +type Spec struct { + RuntimeRoot string + UID uint32 + GID uint32 +} + +type Snapshot struct { + Mode string + Auth []byte + Users []byte + AuthSHA256 string + UsersSHA256 string + Generation string + CanonicalRevision string +} + +type Selector struct { + Version int `json:"version"` + State string `json:"state"` + Transaction string `json:"transaction"` + Generation string `json:"generation,omitempty"` + PreviousGenerations []string `json:"previousGenerations,omitempty"` +} + +type Status struct { + Selector Selector + Snapshot Snapshot +} + +type Transaction struct { + spec Spec + transactionID string + before *Snapshot + prior *Status + rootFD int + lockFD int + blocked bool + closed bool +}