fix: harden DWH credential registry reads
This commit is contained in:
@@ -17,6 +17,7 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -60,6 +61,8 @@ type PublicRecord struct {
|
|||||||
|
|
||||||
// Store owns protected descriptors for one registry root.
|
// Store owns protected descriptors for one registry root.
|
||||||
type Store struct {
|
type Store struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
now func() time.Time
|
||||||
root *securefile.Dir
|
root *securefile.Dir
|
||||||
active *securefile.Dir
|
active *securefile.Dir
|
||||||
revoked *securefile.Dir
|
revoked *securefile.Dir
|
||||||
@@ -88,7 +91,7 @@ func Open(root string) (*Store, error) {
|
|||||||
_ = rootDir.Close()
|
_ = rootDir.Close()
|
||||||
return nil, integrity(err)
|
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.
|
// Close closes descriptors held by the store.
|
||||||
@@ -96,6 +99,8 @@ func (s *Store) Close() error {
|
|||||||
if s == nil {
|
if s == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
var first error
|
var first error
|
||||||
for _, dir := range []*securefile.Dir{s.revoked, s.active, s.root} {
|
for _, dir := range []*securefile.Dir{s.revoked, s.active, s.root} {
|
||||||
if dir == nil {
|
if dir == nil {
|
||||||
@@ -115,6 +120,11 @@ func (s *Store) Add(value record.Record) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return s.withWriterLock(func() error {
|
return s.withWriterLock(func() error {
|
||||||
|
return s.addUnlocked(value)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) addUnlocked(value record.Record) error {
|
||||||
active, revoked, err := s.scanAll()
|
active, revoked, err := s.scanAll()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -137,7 +147,6 @@ func (s *Store) Add(value record.Record) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
return s.writeRecord(s.active, value)
|
return s.writeRecord(s.active, value)
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Find returns an active credential record. A revoked counterpart is checked
|
// 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) {
|
if !validKeyID(keyID) {
|
||||||
return record.Record{}, ErrNotFound
|
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 {
|
if _, err := s.load(s.revoked, StateRevoked, keyID); err == nil {
|
||||||
return record.Record{}, ErrRevoked
|
return record.Record{}, ErrRevoked
|
||||||
} else if !errors.Is(err, ErrNotFound) {
|
} else if !errors.Is(err, ErrNotFound) {
|
||||||
return record.Record{}, err
|
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
|
// FindLegacy returns the sole active legacy_raw record. A revoked legacy record
|
||||||
// wins; any multiple-legacy condition is an integrity fault.
|
// wins; any multiple-legacy condition is an integrity fault.
|
||||||
func (s *Store) FindLegacy() (record.Record, error) {
|
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()
|
active, revoked, err := s.scanAll()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return record.Record{}, err
|
return record.Record{}, err
|
||||||
@@ -171,6 +205,9 @@ func (s *Store) FindLegacy() (record.Record, error) {
|
|||||||
}
|
}
|
||||||
for _, existing := range active {
|
for _, existing := range active {
|
||||||
if existing.record.Kind == record.KindLegacyRaw {
|
if existing.record.Kind == record.KindLegacyRaw {
|
||||||
|
if s.expired(existing.record) {
|
||||||
|
return record.Record{}, ErrNotFound
|
||||||
|
}
|
||||||
return existing.record, nil
|
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
|
// List returns a stable, redacted inventory. A revoked record replaces any
|
||||||
// same-key active record visible during a revoked-first transition.
|
// same-key active record visible during a revoked-first transition.
|
||||||
func (s *Store) List() ([]PublicRecord, error) {
|
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()
|
active, revoked, err := s.scanAll()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -213,6 +259,11 @@ func (s *Store) Revoke(keyID, reason string, at time.Time) error {
|
|||||||
return ErrNotFound
|
return ErrNotFound
|
||||||
}
|
}
|
||||||
return s.withWriterLock(func() error {
|
return s.withWriterLock(func() error {
|
||||||
|
return s.revokeUnlocked(keyID, reason, at)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) revokeUnlocked(keyID, reason string, at time.Time) error {
|
||||||
active, revoked, err := s.scanAll()
|
active, revoked, err := s.scanAll()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -248,12 +299,20 @@ func (s *Store) Revoke(keyID, reason string, at time.Time) error {
|
|||||||
return integrity(err)
|
return integrity(err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check validates every record and protected directory without exposing any
|
// Check validates every record and protected directory without exposing any
|
||||||
// digest data.
|
// digest data.
|
||||||
func (s *Store) Check() error {
|
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()
|
active, revoked, err := s.scanAll()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -262,7 +321,12 @@ func (s *Store) Check() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *Store) withWriterLock(fn func() error) 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"))
|
return integrity(errors.New("uninitialized store"))
|
||||||
}
|
}
|
||||||
lock, err := s.root.Lock(".writer.lock")
|
lock, err := s.root.Lock(".writer.lock")
|
||||||
@@ -410,6 +474,17 @@ func validateForState(value record.Record, state State) error {
|
|||||||
return nil
|
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 {
|
func validateLegacyMultiplicity(active, revoked []storedRecord) error {
|
||||||
activeCount := 0
|
activeCount := 0
|
||||||
revokedCount := 0
|
revokedCount := 0
|
||||||
@@ -452,9 +527,22 @@ func decodeRecord(data []byte) (record.Record, error) {
|
|||||||
return value, nil
|
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 {
|
func rejectDuplicateOrTrailingJSON(data []byte) error {
|
||||||
decoder := json.NewDecoder(bytes.NewReader(data))
|
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||||
if err := consumeJSONValue(decoder); err != nil {
|
if err := consumeRecordObject(decoder); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
var extra any
|
var extra any
|
||||||
@@ -467,6 +555,45 @@ func rejectDuplicateOrTrailingJSON(data []byte) error {
|
|||||||
return nil
|
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 {
|
func consumeJSONValue(decoder *json.Decoder) error {
|
||||||
token, err := decoder.Token()
|
token, err := decoder.Token()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -165,6 +165,27 @@ func TestCheckRejectsInvalidRecordJSONAndFilenameMismatches(t *testing.T) {
|
|||||||
return append(append(data[:len(data)-1], []byte(`,"unknown":true}`)...), '\n')
|
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",
|
name: "duplicate JSON field",
|
||||||
fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" },
|
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)
|
t.Fatalf("Find() description = %q, want one published writer", found.Description)
|
||||||
}
|
}
|
||||||
|
|
||||||
readStore := openStore(t, root)
|
readStore := right
|
||||||
readersDone := make(chan struct{})
|
readersDone := make(chan struct{})
|
||||||
readErrs := make(chan error, 16)
|
readErrs := make(chan error, 16)
|
||||||
var readers sync.WaitGroup
|
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)
|
t.Fatalf("Revoke() during concurrent reads error = %v", err)
|
||||||
}
|
}
|
||||||
readers.Wait()
|
readers.Wait()
|
||||||
@@ -455,3 +476,302 @@ func TestListOrdersRecordsByKeyID(t *testing.T) {
|
|||||||
t.Fatalf("List() key order = %v, want sorted order", keys)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user