872 lines
24 KiB
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)
|
|
}
|
|
}
|
|
}
|