Files
ThothII/tools/dwh-auth/internal/registry/store_test.go

872 lines
24 KiB
Go

//go:build linux
package registry
import (
"bytes"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"errors"
"os"
"path/filepath"
"reflect"
"sort"
"sync"
"testing"
"time"
"github.com/aritmolab/thothii/tools/dwh-auth/internal/record"
)
func TestAddFindAndListPublishCanonicalProtectedRecord(t *testing.T) {
root := t.TempDir()
store, err := Open(root)
if err != nil {
t.Fatalf("Open() error = %v", err)
}
record := syntheticRecord(1)
if err := store.Add(record); err != nil {
t.Fatalf("Add() error = %v", err)
}
path := filepath.Join(root, "active", record.KeyID+".json")
gotBytes, err := os.ReadFile(path)
if err != nil {
t.Fatalf("ReadFile() error = %v", err)
}
wantBytes, err := json.Marshal(record)
if err != nil {
t.Fatalf("Marshal() error = %v", err)
}
wantBytes = append(wantBytes, '\n')
if !bytes.Equal(gotBytes, wantBytes) {
t.Fatalf("published record = %s, want %s", gotBytes, wantBytes)
}
info, err := os.Stat(path)
if err != nil {
t.Fatalf("Stat() error = %v", err)
}
if got, want := info.Mode().Perm(), os.FileMode(0o640); got != want {
t.Fatalf("record mode = %04o, want %04o", got, want)
}
found, err := store.Find(record.KeyID)
if err != nil {
t.Fatalf("Find() error = %v", err)
}
if found != record {
t.Fatalf("Find() = %#v, want %#v", found, record)
}
listed, err := store.List()
if err != nil {
t.Fatalf("List() error = %v", err)
}
if got, want := len(listed), 1; got != want {
t.Fatalf("List() length = %d, want %d", got, want)
}
if got, want := listed[0], publicRecord(record, StateActive); got != want {
t.Fatalf("List()[0] = %#v, want %#v", got, want)
}
encoded, err := json.Marshal(listed[0])
if err != nil {
t.Fatalf("Marshal(PublicRecord) error = %v", err)
}
if bytes.Contains(encoded, []byte(record.SecretSHA256)) {
t.Fatalf("PublicRecord JSON exposed secret digest: %s", encoded)
}
}
func TestOpenRejectsSymlinkedRegistryRootAndDirectories(t *testing.T) {
t.Run("root", func(t *testing.T) {
parent := t.TempDir()
target := t.TempDir()
root := filepath.Join(parent, "registry")
if err := os.Symlink(target, root); err != nil {
t.Fatalf("Symlink() error = %v", err)
}
if _, err := Open(root); err == nil {
t.Fatal("Open() error = nil, want symlink refusal")
}
})
t.Run("active directory", func(t *testing.T) {
root := t.TempDir()
if err := os.Symlink(t.TempDir(), filepath.Join(root, "active")); err != nil {
t.Fatalf("Symlink() error = %v", err)
}
if _, err := Open(root); err == nil {
t.Fatal("Open() error = nil, want active-directory symlink refusal")
}
})
t.Run("revoked directory", func(t *testing.T) {
root := t.TempDir()
if err := os.Symlink(t.TempDir(), filepath.Join(root, "revoked")); err != nil {
t.Fatalf("Symlink() error = %v", err)
}
if _, err := Open(root); err == nil {
t.Fatal("Open() error = nil, want revoked-directory symlink refusal")
}
})
}
func TestOpenReadOnlyRequiresCompletePreprovisionedRegistryAndNeverCreates(t *testing.T) {
root := t.TempDir()
before, err := os.ReadDir(root)
if err != nil {
t.Fatalf("ReadDir(before) error = %v", err)
}
if _, err := OpenReadOnly(root); err == nil {
t.Fatal("OpenReadOnly(incomplete) error = nil, want refusal")
}
after, err := os.ReadDir(root)
if err != nil {
t.Fatalf("ReadDir(after) error = %v", err)
}
if !reflect.DeepEqual(before, after) {
t.Fatalf("OpenReadOnly(incomplete) changed root entries: before=%v after=%v", before, after)
}
admin := openStore(t, root)
item := syntheticRecord(13)
if err := admin.Add(item); err != nil {
t.Fatalf("Add() error = %v", err)
}
if err := admin.Close(); err != nil {
t.Fatalf("Close(admin) error = %v", err)
}
for _, path := range []string{root, filepath.Join(root, "active"), filepath.Join(root, "revoked")} {
if err := os.Chmod(path, 0o750|os.ModeSetgid); err != nil {
t.Fatalf("Chmod(%q, 2750) error = %v", path, err)
}
}
if info, err := os.Stat(filepath.Join(root, registryLockName)); err != nil || info.Mode().Perm() != 0o640 {
t.Fatalf("registry lock = (%v, %v), want 0640", info, err)
}
entriesBefore := registryEntryModes(t, root)
reader, err := OpenReadOnly(root)
if err != nil {
t.Fatalf("OpenReadOnly() error = %v", err)
}
defer reader.Close()
if found, err := reader.Find(item.KeyID); err != nil || found.KeyID != item.KeyID {
t.Fatalf("Find() = (%#v, %v), want protected record", found, err)
}
if err := reader.Check(); err != nil {
t.Fatalf("Check() error = %v", err)
}
if err := reader.Add(syntheticRecord(14)); err == nil {
t.Fatal("read-only Add() error = nil, want refusal")
}
if err := reader.Revoke(item.KeyID, "synthetic", item.CreatedAt.Add(time.Hour)); err == nil {
t.Fatal("read-only Revoke() error = nil, want refusal")
}
if got := registryEntryModes(t, root); !reflect.DeepEqual(got, entriesBefore) {
t.Fatalf("read-only operations changed registry entries: before=%v after=%v", entriesBefore, got)
}
}
func registryEntryModes(t *testing.T, root string) []string {
t.Helper()
var result []string
for _, directory := range []string{root, filepath.Join(root, "active"), filepath.Join(root, "revoked")} {
entries, err := os.ReadDir(directory)
if err != nil {
t.Fatalf("ReadDir(%q) error = %v", directory, err)
}
for _, entry := range entries {
info, err := entry.Info()
if err != nil {
t.Fatalf("Info(%q) error = %v", entry.Name(), err)
}
result = append(result, filepath.Join(directory, entry.Name())+":"+info.Mode().String())
}
}
sort.Strings(result)
return result
}
func TestCheckRejectsUnsafeModesAndSymlinkedRecords(t *testing.T) {
t.Run("unsafe active mode", func(t *testing.T) {
root := t.TempDir()
if err := os.Mkdir(filepath.Join(root, "active"), 0o770); err != nil {
t.Fatalf("Mkdir(active) error = %v", err)
}
if err := os.Chmod(filepath.Join(root, "active"), 0o770); err != nil {
t.Fatalf("Chmod(active) error = %v", err)
}
if _, err := Open(root); err == nil {
t.Fatal("Open() error = nil, want unsafe directory-mode refusal")
}
})
t.Run("unsafe record mode", func(t *testing.T) {
root := t.TempDir()
store := openStore(t, root)
record := syntheticRecord(2)
writeRecord(t, root, StateActive, record.KeyID+".json", marshalRecord(t, record), 0o660)
if err := store.Check(); err == nil {
t.Fatal("Check() error = nil, want unsafe record-mode refusal")
}
})
t.Run("symlinked record", func(t *testing.T) {
root := t.TempDir()
store := openStore(t, root)
record := syntheticRecord(3)
target := filepath.Join(root, "target.json")
writeRegistryFile(t, target, marshalRecord(t, record), 0o640)
if err := os.Symlink(target, filepath.Join(root, "active", record.KeyID+".json")); err != nil {
t.Fatalf("Symlink() error = %v", err)
}
if _, err := store.Find(record.KeyID); err == nil {
t.Fatal("Find() error = nil, want symlink refusal")
}
})
}
func TestCheckRejectsInvalidRecordJSONAndFilenameMismatches(t *testing.T) {
cases := []struct {
name string
fileName func(record.Record) string
mutate func([]byte) []byte
}{
{
name: "unknown JSON field",
fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" },
mutate: func(data []byte) []byte {
data = bytes.TrimSuffix(data, []byte{'\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",
fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" },
mutate: func(data []byte) []byte {
return append([]byte(`{"schema_version":1,`), data[1:]...)
},
},
{
name: "trailing JSON value",
fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" },
mutate: func(data []byte) []byte {
return append(data, []byte(`{}`)...)
},
},
{
name: "partial JSON",
fileName: func(record.Record) string { return "AAAAAAAAAAAAAAAA.json" },
mutate: func([]byte) []byte { return []byte(`{"schema_version":`) },
},
{
name: "filename mismatch",
fileName: func(record.Record) string { return "wrong.json" },
mutate: func(data []byte) []byte { return data },
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
root := t.TempDir()
store := openStore(t, root)
record := syntheticRecord(0)
writeRecord(t, root, StateActive, tc.fileName(record), tc.mutate(marshalRecord(t, record)), 0o640)
if err := store.Check(); err == nil {
t.Fatal("Check() error = nil, want integrity refusal")
}
})
}
}
func TestRevokePublishesRevokedStateAndFindRevokedWins(t *testing.T) {
root := t.TempDir()
store := openStore(t, root)
record := syntheticRecord(4)
if err := store.Add(record); err != nil {
t.Fatalf("Add() error = %v", err)
}
revokedAt := record.CreatedAt.Add(2 * time.Hour)
if err := store.Revoke(record.KeyID, "synthetic rotation", revokedAt); err != nil {
t.Fatalf("Revoke() error = %v", err)
}
if _, err := os.Stat(filepath.Join(root, "active", record.KeyID+".json")); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("active record stat error = %v, want not-exist", err)
}
if _, err := os.Stat(filepath.Join(root, "revoked", record.KeyID+".json")); err != nil {
t.Fatalf("revoked record stat error = %v", err)
}
if _, err := store.Find(record.KeyID); !errors.Is(err, ErrRevoked) {
t.Fatalf("Find() error = %v, want ErrRevoked", err)
}
// A crash after revoked publication but before active removal must still deny the key.
writeRecord(t, root, StateActive, record.KeyID+".json", marshalRecord(t, record), 0o640)
if _, err := store.Find(record.KeyID); !errors.Is(err, ErrRevoked) {
t.Fatalf("Find() with active and revoked files error = %v, want ErrRevoked", err)
}
listed, err := store.List()
if err != nil {
t.Fatalf("List() error = %v", err)
}
if got, want := listed, []PublicRecord{publicRecord(revokedRecord(record, revokedAt, "synthetic rotation"), StateRevoked)}; !equalPublicRecords(got, want) {
t.Fatalf("List() = %#v, want %#v", got, want)
}
}
func TestFindLegacyAllowsOneAndRejectsMultipleRecords(t *testing.T) {
t.Run("one active legacy", func(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)
}
got, err := store.FindLegacy()
if err != nil {
t.Fatalf("FindLegacy() error = %v", err)
}
if got != legacy {
t.Fatalf("FindLegacy() = %#v, want %#v", got, legacy)
}
})
t.Run("multiple legacy records", func(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)
}
writeRecord(t, root, StateActive, "duplicate.json", marshalRecord(t, legacy), 0o640)
if _, err := store.FindLegacy(); err == nil {
t.Fatal("FindLegacy() error = nil, want multiple-legacy integrity refusal")
}
})
t.Run("active and revoked legacy records", func(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)
}
revokedAt := legacy.CreatedAt.Add(time.Hour)
writeRecord(t, root, StateRevoked, legacy.KeyID+".json", marshalRecord(t, revokedRecord(legacy, revokedAt, "synthetic")), 0o640)
if _, err := store.FindLegacy(); !errors.Is(err, ErrIntegrity) {
t.Fatalf("FindLegacy() error = %v, want ErrIntegrity", err)
}
if err := store.Check(); !errors.Is(err, ErrIntegrity) {
t.Fatalf("Check() error = %v, want ErrIntegrity", err)
}
})
}
func TestConcurrentAddDoesNotOverwriteAndFindNeverReadsPartialRecord(t *testing.T) {
root := t.TempDir()
left := openStore(t, root)
right := openStore(t, root)
first := syntheticRecord(5)
second := first
second.Description = "second synthetic writer"
start := make(chan struct{})
errs := make(chan error, 2)
var writers sync.WaitGroup
for _, candidate := range []record.Record{first, second} {
writers.Add(1)
go func(candidate record.Record) {
defer writers.Done()
<-start
errs <- left.Add(candidate)
}(candidate)
}
close(start)
writers.Wait()
close(errs)
successes := 0
for err := range errs {
if err == nil {
successes++
}
}
if got, want := successes, 1; got != want {
t.Fatalf("concurrent Add() successes = %d, want %d", got, want)
}
found, err := right.Find(first.KeyID)
if err != nil {
t.Fatalf("Find() after concurrent Add() error = %v", err)
}
if found.Description != first.Description && found.Description != second.Description {
t.Fatalf("Find() description = %q, want one published writer", found.Description)
}
readStore := right
readersDone := make(chan struct{})
readErrs := make(chan error, 16)
var readers sync.WaitGroup
for range 16 {
readers.Add(1)
go func() {
defer readers.Done()
for range 100 {
got, err := readStore.Find(first.KeyID)
if err == nil {
if got.KeyID != first.KeyID {
readErrs <- errors.New("Find returned an unexpected key ID")
}
continue
}
if !errors.Is(err, ErrRevoked) {
readErrs <- err
}
}
}()
}
if err := left.Revoke(first.KeyID, "synthetic rotation", first.CreatedAt.Add(time.Hour)); err != nil {
t.Fatalf("Revoke() during concurrent reads error = %v", err)
}
readers.Wait()
close(readersDone)
close(readErrs)
for err := range readErrs {
if err != nil {
t.Fatalf("concurrent Find() error = %v", err)
}
}
if _, err := readStore.Find(first.KeyID); !errors.Is(err, ErrRevoked) {
t.Fatalf("Find() after Revoke() error = %v, want ErrRevoked", err)
}
}
func syntheticRecord(seed byte) record.Record {
keyBytes := bytes.Repeat([]byte{seed}, 12)
digest := sha256.Sum256([]byte{seed, 's', 'y', 'n', 't', 'h', 'e', 't', 'i', 'c'})
return record.Record{
SchemaVersion: record.SchemaVersion,
Kind: record.KindV1,
KeyID: base64.RawURLEncoding.EncodeToString(keyBytes),
InstallationID: "test-installation",
Description: "synthetic test credential",
SecretSHA256: base64.RawURLEncoding.EncodeToString(digest[:]),
CreatedAt: time.Date(2026, 8, 20, 12, 0, int(seed), 0, time.UTC),
}
}
func syntheticLegacyRecord() record.Record {
value := syntheticRecord(9)
value.Kind = record.KindLegacyRaw
value.KeyID = record.LegacyKeyID
value.InstallationID = "legacy-shared"
return value
}
func revokedRecord(value record.Record, at time.Time, reason string) record.Record {
value.RevokedAt = &at
value.RevocationReason = reason
return value
}
func publicRecord(value record.Record, state State) PublicRecord {
return PublicRecord{
Kind: value.Kind,
KeyID: value.KeyID,
InstallationID: value.InstallationID,
Description: value.Description,
CreatedAt: value.CreatedAt,
ExpiresAt: value.ExpiresAt,
State: state,
RevokedAt: value.RevokedAt,
RevocationReason: value.RevocationReason,
}
}
func equalPublicRecords(got, want []PublicRecord) bool {
if len(got) != len(want) {
return false
}
for i := range got {
if !reflect.DeepEqual(got[i], want[i]) {
return false
}
}
return true
}
func openStore(t *testing.T, root string) *Store {
t.Helper()
store, err := Open(root)
if err != nil {
t.Fatalf("Open() error = %v", err)
}
return store
}
func marshalRecord(t *testing.T, value record.Record) []byte {
t.Helper()
data, err := json.Marshal(value)
if err != nil {
t.Fatalf("Marshal() error = %v", err)
}
return append(data, '\n')
}
func writeRecord(t *testing.T, root string, state State, name string, data []byte, mode os.FileMode) {
t.Helper()
writeRegistryFile(t, filepath.Join(root, string(state), name), data, mode)
}
func writeRegistryFile(t *testing.T, path string, data []byte, mode os.FileMode) {
t.Helper()
if err := os.WriteFile(path, data, mode); err != nil {
t.Fatalf("WriteFile(%q) error = %v", path, err)
}
if err := os.Chmod(path, mode); err != nil {
t.Fatalf("Chmod(%q) error = %v", path, err)
}
}
func TestListOrdersRecordsByKeyID(t *testing.T) {
root := t.TempDir()
store := openStore(t, root)
for _, seed := range []byte{8, 1, 5} {
if err := store.Add(syntheticRecord(seed)); err != nil {
t.Fatalf("Add(%d) error = %v", seed, err)
}
}
listed, err := store.List()
if err != nil {
t.Fatalf("List() error = %v", err)
}
keys := make([]string, len(listed))
for i, item := range listed {
keys[i] = item.KeyID
}
if !sort.StringsAreSorted(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 TestCrossStoreReadersWaitAcrossRevokePublicationAndUnlink(t *testing.T) {
root := t.TempDir()
writer := openStore(t, root)
reader := openStore(t, root)
legacy := syntheticLegacyRecord()
if err := writer.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 <- writer.withWriterLock(func() error {
close(writerReady)
<-releaseWriter
revoked := revokedRecord(legacy, revokedAt, "synthetic rotation")
if err := writer.writeRecord(writer.revoked, revoked); err != nil {
return err
}
if err := writer.active.Remove(recordFileName(legacy.KeyID)); err != nil {
return err
}
return writer.active.Sync()
})
}()
<-writerReady
type readResult struct {
name string
err error
}
readers := []struct {
name string
run func() error
}{
{name: "Find", run: func() error { _, err := reader.Find(legacy.KeyID); return err }},
{name: "List", run: func() error { _, err := reader.List(); return err }},
{name: "Check", run: reader.Check},
{name: "FindLegacy", run: func() error { _, err := reader.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 TestCrossStoreScanReadersWaitForWriterTemporaryFile(t *testing.T) {
root := t.TempDir()
writer := openStore(t, root)
reader := openStore(t, root)
legacy := syntheticLegacyRecord()
if err := writer.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 <- writer.withWriterLock(func() error {
file, err := writer.active.CreateExclusive(temporary, 0o600)
if err != nil {
return err
}
if err := file.Close(); err != nil {
return err
}
close(writerReady)
<-releaseWriter
if err := writer.active.Remove(temporary); err != nil {
return err
}
return writer.active.Sync()
})
}()
<-writerReady
type readResult struct {
name string
err error
}
readers := []struct {
name string
run func() error
}{
{name: "List", run: func() error { _, err := reader.List(); return err }},
{name: "Check", run: reader.Check},
{name: "FindLegacy", run: func() error { _, err := reader.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)
}
}
}