From 971a0e66b1f96694569b684550a4ff02ec65de08 Mon Sep 17 00:00:00 2001 From: User Date: Fri, 21 Aug 2026 00:18:12 +0200 Subject: [PATCH] fix: harden DWH credential registry reads --- tools/dwh-auth/internal/registry/store.go | 241 ++++++++++--- .../dwh-auth/internal/registry/store_test.go | 324 +++++++++++++++++- 2 files changed, 506 insertions(+), 59 deletions(-) diff --git a/tools/dwh-auth/internal/registry/store.go b/tools/dwh-auth/internal/registry/store.go index fd054b31..e12aa4e0 100644 --- a/tools/dwh-auth/internal/registry/store.go +++ b/tools/dwh-auth/internal/registry/store.go @@ -17,6 +17,7 @@ import ( "os" "sort" "strings" + "sync" "syscall" "time" @@ -60,6 +61,8 @@ type PublicRecord struct { // Store owns protected descriptors for one registry root. type Store struct { + mu sync.RWMutex + now func() time.Time root *securefile.Dir active *securefile.Dir revoked *securefile.Dir @@ -88,7 +91,7 @@ func Open(root string) (*Store, error) { _ = rootDir.Close() return nil, integrity(err) } - return &Store{root: rootDir, active: active, revoked: revoked}, nil + return &Store{root: rootDir, active: active, revoked: revoked, now: time.Now}, nil } // Close closes descriptors held by the store. @@ -96,6 +99,8 @@ func (s *Store) Close() error { if s == nil { return nil } + s.mu.Lock() + defer s.mu.Unlock() var first error for _, dir := range []*securefile.Dir{s.revoked, s.active, s.root} { if dir == nil { @@ -115,29 +120,33 @@ func (s *Store) Add(value record.Record) error { return err } return s.withWriterLock(func() error { - active, revoked, err := s.scanAll() - if err != nil { - return err + return s.addUnlocked(value) + }) +} + +func (s *Store) addUnlocked(value record.Record) error { + active, revoked, err := s.scanAll() + if err != nil { + return err + } + for _, existing := range revoked { + if existing.record.KeyID == value.KeyID { + return ErrRevoked } - for _, existing := range revoked { - if existing.record.KeyID == value.KeyID { - return ErrRevoked - } + } + for _, existing := range active { + if existing.record.KeyID == value.KeyID { + return ErrConflict } - for _, existing := range active { - if existing.record.KeyID == value.KeyID { + } + if value.Kind == record.KindLegacyRaw { + for _, existing := range append(active, revoked...) { + if existing.record.Kind == record.KindLegacyRaw { return ErrConflict } } - if value.Kind == record.KindLegacyRaw { - for _, existing := range append(active, revoked...) { - if existing.record.Kind == record.KindLegacyRaw { - return ErrConflict - } - } - } - return s.writeRecord(s.active, value) - }) + } + return s.writeRecord(s.active, value) } // Find returns an active credential record. A revoked counterpart is checked @@ -146,17 +155,42 @@ func (s *Store) Find(keyID string) (record.Record, error) { if !validKeyID(keyID) { return record.Record{}, ErrNotFound } + if s == nil { + return record.Record{}, integrity(errors.New("uninitialized store")) + } + s.mu.RLock() + defer s.mu.RUnlock() + return s.findUnlocked(keyID) +} + +func (s *Store) findUnlocked(keyID string) (record.Record, error) { if _, err := s.load(s.revoked, StateRevoked, keyID); err == nil { return record.Record{}, ErrRevoked } else if !errors.Is(err, ErrNotFound) { return record.Record{}, err } - return s.load(s.active, StateActive, keyID) + value, err := s.load(s.active, StateActive, keyID) + if err != nil { + return record.Record{}, err + } + if s.expired(value) { + return record.Record{}, ErrNotFound + } + return value, nil } // FindLegacy returns the sole active legacy_raw record. A revoked legacy record // wins; any multiple-legacy condition is an integrity fault. func (s *Store) FindLegacy() (record.Record, error) { + if s == nil { + return record.Record{}, integrity(errors.New("uninitialized store")) + } + s.mu.RLock() + defer s.mu.RUnlock() + return s.findLegacyUnlocked() +} + +func (s *Store) findLegacyUnlocked() (record.Record, error) { active, revoked, err := s.scanAll() if err != nil { return record.Record{}, err @@ -171,6 +205,9 @@ func (s *Store) FindLegacy() (record.Record, error) { } for _, existing := range active { if existing.record.Kind == record.KindLegacyRaw { + if s.expired(existing.record) { + return record.Record{}, ErrNotFound + } return existing.record, nil } } @@ -180,6 +217,15 @@ func (s *Store) FindLegacy() (record.Record, error) { // List returns a stable, redacted inventory. A revoked record replaces any // same-key active record visible during a revoked-first transition. func (s *Store) List() ([]PublicRecord, error) { + if s == nil { + return nil, integrity(errors.New("uninitialized store")) + } + s.mu.RLock() + defer s.mu.RUnlock() + return s.listUnlocked() +} + +func (s *Store) listUnlocked() ([]PublicRecord, error) { active, revoked, err := s.scanAll() if err != nil { return nil, err @@ -213,47 +259,60 @@ func (s *Store) Revoke(keyID, reason string, at time.Time) error { return ErrNotFound } return s.withWriterLock(func() error { - active, revoked, err := s.scanAll() - if err != nil { - return err - } - for _, existing := range revoked { - if existing.record.KeyID == keyID { - return ErrRevoked - } - } - var target *record.Record - for _, existing := range active { - if existing.record.KeyID == keyID { - candidate := existing.record - target = &candidate - break - } - } - if target == nil { - return ErrNotFound - } - target.RevokedAt = &at - target.RevocationReason = reason - if err := validateForState(*target, StateRevoked); err != nil { - return err - } - if err := s.writeRecord(s.revoked, *target); err != nil { - return err - } - if err := s.active.Remove(recordFileName(keyID)); err != nil { - return integrity(err) - } - if err := s.active.Sync(); err != nil { - return integrity(err) - } - return nil + return s.revokeUnlocked(keyID, reason, at) }) } +func (s *Store) revokeUnlocked(keyID, reason string, at time.Time) error { + active, revoked, err := s.scanAll() + if err != nil { + return err + } + for _, existing := range revoked { + if existing.record.KeyID == keyID { + return ErrRevoked + } + } + var target *record.Record + for _, existing := range active { + if existing.record.KeyID == keyID { + candidate := existing.record + target = &candidate + break + } + } + if target == nil { + return ErrNotFound + } + target.RevokedAt = &at + target.RevocationReason = reason + if err := validateForState(*target, StateRevoked); err != nil { + return err + } + if err := s.writeRecord(s.revoked, *target); err != nil { + return err + } + if err := s.active.Remove(recordFileName(keyID)); err != nil { + return integrity(err) + } + if err := s.active.Sync(); err != nil { + return integrity(err) + } + return nil +} + // Check validates every record and protected directory without exposing any // digest data. func (s *Store) Check() error { + if s == nil { + return integrity(errors.New("uninitialized store")) + } + s.mu.RLock() + defer s.mu.RUnlock() + return s.checkUnlocked() +} + +func (s *Store) checkUnlocked() error { active, revoked, err := s.scanAll() if err != nil { return err @@ -262,7 +321,12 @@ func (s *Store) Check() error { } func (s *Store) withWriterLock(fn func() error) error { - if s == nil || s.root == nil || s.active == nil || s.revoked == nil { + if s == nil { + return integrity(errors.New("uninitialized store")) + } + s.mu.Lock() + defer s.mu.Unlock() + if s.root == nil || s.active == nil || s.revoked == nil { return integrity(errors.New("uninitialized store")) } lock, err := s.root.Lock(".writer.lock") @@ -410,6 +474,17 @@ func validateForState(value record.Record, state State) error { return nil } +func (s *Store) expired(value record.Record) bool { + return value.ExpiresAt != nil && !value.ExpiresAt.After(s.nowUTC()) +} + +func (s *Store) nowUTC() time.Time { + if s.now == nil { + return time.Now().UTC() + } + return s.now().UTC() +} + func validateLegacyMultiplicity(active, revoked []storedRecord) error { activeCount := 0 revokedCount := 0 @@ -452,9 +527,22 @@ func decodeRecord(data []byte) (record.Record, error) { return value, nil } +var recordJSONFields = map[string]struct{}{ + "schema_version": {}, + "credential_kind": {}, + "key_id": {}, + "installation_id": {}, + "description": {}, + "secret_sha256": {}, + "created_at": {}, + "expires_at": {}, + "revoked_at": {}, + "revocation_reason": {}, +} + func rejectDuplicateOrTrailingJSON(data []byte) error { decoder := json.NewDecoder(bytes.NewReader(data)) - if err := consumeJSONValue(decoder); err != nil { + if err := consumeRecordObject(decoder); err != nil { return err } var extra any @@ -467,6 +555,45 @@ func rejectDuplicateOrTrailingJSON(data []byte) error { return nil } +func consumeRecordObject(decoder *json.Decoder) error { + token, err := decoder.Token() + if err != nil { + return err + } + if token != json.Delim('{') { + return errors.New("record JSON is not an object") + } + seen := make(map[string]struct{}) + for decoder.More() { + keyToken, err := decoder.Token() + if err != nil { + return err + } + key, ok := keyToken.(string) + if !ok { + return errors.New("JSON object key is not a string") + } + if _, allowed := recordJSONFields[key]; !allowed { + return fmt.Errorf("unknown JSON field %q", key) + } + if _, duplicate := seen[key]; duplicate { + return fmt.Errorf("duplicate JSON field %q", key) + } + seen[key] = struct{}{} + if err := consumeJSONValue(decoder); err != nil { + return err + } + } + end, err := decoder.Token() + if err != nil { + return err + } + if end != json.Delim('}') { + return errors.New("unterminated JSON object") + } + return nil +} + func consumeJSONValue(decoder *json.Decoder) error { token, err := decoder.Token() if err != nil { diff --git a/tools/dwh-auth/internal/registry/store_test.go b/tools/dwh-auth/internal/registry/store_test.go index f2edc5de..560a7fd3 100644 --- a/tools/dwh-auth/internal/registry/store_test.go +++ b/tools/dwh-auth/internal/registry/store_test.go @@ -165,6 +165,27 @@ func TestCheckRejectsInvalidRecordJSONAndFilenameMismatches(t *testing.T) { return append(append(data[:len(data)-1], []byte(`,"unknown":true}`)...), '\n') }, }, + { + name: "uppercase JSON field alias", + fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" }, + mutate: func(data []byte) []byte { + return bytes.Replace(data, []byte(`"secret_sha256"`), []byte(`"SECRET_SHA256"`), 1) + }, + }, + { + name: "mixed-case JSON field alias", + fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" }, + mutate: func(data []byte) []byte { + return bytes.Replace(data, []byte(`"secret_sha256"`), []byte(`"Secret_SHA256"`), 1) + }, + }, + { + name: "schema JSON field alias", + fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" }, + mutate: func(data []byte) []byte { + return bytes.Replace(data, []byte(`"schema_version"`), []byte(`"Schema_Version"`), 1) + }, + }, { name: "duplicate JSON field", fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" }, @@ -310,7 +331,7 @@ func TestConcurrentAddDoesNotOverwriteAndFindNeverReadsPartialRecord(t *testing. t.Fatalf("Find() description = %q, want one published writer", found.Description) } - readStore := openStore(t, root) + readStore := right readersDone := make(chan struct{}) readErrs := make(chan error, 16) var readers sync.WaitGroup @@ -332,7 +353,7 @@ func TestConcurrentAddDoesNotOverwriteAndFindNeverReadsPartialRecord(t *testing. } }() } - if err := right.Revoke(first.KeyID, "synthetic rotation", first.CreatedAt.Add(time.Hour)); err != nil { + if err := readStore.Revoke(first.KeyID, "synthetic rotation", first.CreatedAt.Add(time.Hour)); err != nil { t.Fatalf("Revoke() during concurrent reads error = %v", err) } readers.Wait() @@ -455,3 +476,302 @@ func TestListOrdersRecordsByKeyID(t *testing.T) { t.Fatalf("List() key order = %v, want sorted order", keys) } } + +func TestFindAndFindLegacyRejectPastExpiration(t *testing.T) { + tests := []struct { + name string + legacy bool + }{ + {name: "v1"}, + {name: "legacy", legacy: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + root := t.TempDir() + store := openStore(t, root) + now := time.Date(2026, 8, 21, 13, 0, 0, 0, time.UTC) + store.now = func() time.Time { return now } + value := syntheticRecord(10) + if tc.legacy { + value = syntheticLegacyRecord() + } + expiresAt := now.Add(-time.Minute) + value.ExpiresAt = &expiresAt + if err := store.Add(value); err != nil { + t.Fatalf("Add() error = %v", err) + } + + var err error + if tc.legacy { + _, err = store.FindLegacy() + } else { + _, err = store.Find(value.KeyID) + } + if !errors.Is(err, ErrNotFound) { + t.Fatalf("expired lookup error = %v, want ErrNotFound", err) + } + }) + } +} + +func TestFindAndFindLegacyAcceptFutureExpiration(t *testing.T) { + tests := []struct { + name string + legacy bool + }{ + {name: "v1"}, + {name: "legacy", legacy: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + root := t.TempDir() + store := openStore(t, root) + now := time.Date(2026, 8, 21, 13, 0, 0, 0, time.UTC) + store.now = func() time.Time { return now } + value := syntheticRecord(11) + if tc.legacy { + value = syntheticLegacyRecord() + } + expiresAt := now.Add(time.Hour) + value.ExpiresAt = &expiresAt + if err := store.Add(value); err != nil { + t.Fatalf("Add() error = %v", err) + } + + var ( + got record.Record + err error + ) + if tc.legacy { + got, err = store.FindLegacy() + } else { + got, err = store.Find(value.KeyID) + } + if err != nil { + t.Fatalf("future lookup error = %v", err) + } + if got.KeyID != value.KeyID || got.ExpiresAt == nil || !got.ExpiresAt.Equal(expiresAt) { + t.Fatalf("future lookup = %#v, want record with key %q and expiry %s", got, value.KeyID, expiresAt) + } + }) + } +} + +func TestFindAndFindLegacyRejectExactExpiration(t *testing.T) { + now := time.Date(2026, 8, 21, 13, 0, 0, 0, time.UTC) + tests := []struct { + name string + legacy bool + }{ + {name: "v1"}, + {name: "legacy", legacy: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + root := t.TempDir() + store := openStore(t, root) + store.now = func() time.Time { return now } + value := syntheticRecord(12) + if tc.legacy { + value = syntheticLegacyRecord() + } + expiresAt := now + value.ExpiresAt = &expiresAt + if err := store.Add(value); err != nil { + t.Fatalf("Add() error = %v", err) + } + + var err error + if tc.legacy { + _, err = store.FindLegacy() + } else { + _, err = store.Find(value.KeyID) + } + if !errors.Is(err, ErrNotFound) { + t.Fatalf("equal-expiry lookup error = %v, want ErrNotFound", err) + } + }) + } +} + +func TestReadOperationsWaitAcrossRevokePublicationAndUnlink(t *testing.T) { + root := t.TempDir() + store := openStore(t, root) + legacy := syntheticLegacyRecord() + if err := store.Add(legacy); err != nil { + t.Fatalf("Add() error = %v", err) + } + + writerReady := make(chan struct{}) + releaseWriter := make(chan struct{}) + writerDone := make(chan error, 1) + revokedAt := legacy.CreatedAt.Add(time.Hour) + go func() { + writerDone <- store.withWriterLock(func() error { + close(writerReady) + <-releaseWriter + revoked := revokedRecord(legacy, revokedAt, "synthetic rotation") + if err := store.writeRecord(store.revoked, revoked); err != nil { + return err + } + if err := store.active.Remove(recordFileName(legacy.KeyID)); err != nil { + return err + } + return store.active.Sync() + }) + }() + <-writerReady + + type readResult struct { + name string + err error + } + readers := []struct { + name string + run func() error + }{ + {name: "Find", run: func() error { _, err := store.Find(legacy.KeyID); return err }}, + {name: "List", run: func() error { _, err := store.List(); return err }}, + {name: "Check", run: store.Check}, + {name: "FindLegacy", run: func() error { _, err := store.FindLegacy(); return err }}, + } + results := make(chan readResult, len(readers)) + var started sync.WaitGroup + started.Add(len(readers)) + for _, reader := range readers { + go func(reader struct { + name string + run func() error + }) { + started.Done() + results <- readResult{name: reader.name, err: reader.run()} + }(reader) + } + started.Wait() + + var early []readResult + timer := time.NewTimer(200 * time.Millisecond) +waitForReaders: + for len(early) < len(readers) { + select { + case result := <-results: + early = append(early, result) + case <-timer.C: + break waitForReaders + } + } + timer.Stop() + close(releaseWriter) + if err := <-writerDone; err != nil { + t.Fatalf("writer error = %v", err) + } + all := append([]readResult(nil), early...) + for len(all) < len(readers) { + all = append(all, <-results) + } + if len(early) != 0 { + t.Fatalf("read operations completed during revocation snapshot: %#v", early) + } + for _, result := range all { + switch result.name { + case "Find", "FindLegacy": + if !errors.Is(result.err, ErrRevoked) { + t.Fatalf("%s() error = %v, want ErrRevoked", result.name, result.err) + } + default: + if result.err != nil { + t.Fatalf("%s() error = %v", result.name, result.err) + } + } + } +} + +func TestScanReadersWaitForWriterTemporaryFile(t *testing.T) { + root := t.TempDir() + store := openStore(t, root) + legacy := syntheticLegacyRecord() + if err := store.Add(legacy); err != nil { + t.Fatalf("Add() error = %v", err) + } + + writerReady := make(chan struct{}) + releaseWriter := make(chan struct{}) + writerDone := make(chan error, 1) + const temporary = ".tmp-regression" + go func() { + writerDone <- store.withWriterLock(func() error { + file, err := store.active.CreateExclusive(temporary, 0o600) + if err != nil { + return err + } + if err := file.Close(); err != nil { + return err + } + close(writerReady) + <-releaseWriter + if err := store.active.Remove(temporary); err != nil { + return err + } + return store.active.Sync() + }) + }() + <-writerReady + + type readResult struct { + name string + err error + } + readers := []struct { + name string + run func() error + }{ + {name: "List", run: func() error { _, err := store.List(); return err }}, + {name: "Check", run: store.Check}, + {name: "FindLegacy", run: func() error { _, err := store.FindLegacy(); return err }}, + } + results := make(chan readResult, len(readers)) + var started sync.WaitGroup + started.Add(len(readers)) + for _, reader := range readers { + go func(reader struct { + name string + run func() error + }) { + started.Done() + results <- readResult{name: reader.name, err: reader.run()} + }(reader) + } + started.Wait() + + var early []readResult + timer := time.NewTimer(200 * time.Millisecond) +waitForReaders: + for len(early) < len(readers) { + select { + case result := <-results: + early = append(early, result) + case <-timer.C: + break waitForReaders + } + } + timer.Stop() + close(releaseWriter) + if err := <-writerDone; err != nil { + t.Fatalf("writer error = %v", err) + } + all := append([]readResult(nil), early...) + for len(all) < len(readers) { + all = append(all, <-results) + } + if len(early) != 0 { + t.Fatalf("scan operations observed a writer temporary file: %#v", early) + } + for _, result := range all { + if result.err != nil { + t.Fatalf("%s() error = %v", result.name, result.err) + } + } +}