fix: harden DWH credential registry reads
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user