feat: add protected DWH credential registry

This commit is contained in:
User
2026-08-20 23:41:43 +02:00
parent 1e82fd33fd
commit 541ef45529
4 changed files with 1791 additions and 0 deletions
+584
View File
@@ -0,0 +1,584 @@
//go:build linux
// Package registry stores protected DWH credential records in active and
// revoked directories. It never exposes credential digests through its public
// listing type.
package registry
import (
"bytes"
"crypto/rand"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"sort"
"strings"
"syscall"
"time"
"github.com/aritmolab/thothii/tools/dwh-auth/internal/record"
"github.com/aritmolab/thothii/tools/dwh-auth/internal/securefile"
)
const maxRecordBytes = 4096
// State describes which registry directory owns a public record.
type State string
const (
StateActive State = "active"
StateRevoked State = "revoked"
)
var (
// ErrNotFound means no protected record exists for the requested key ID.
ErrNotFound = errors.New("credential record not found")
// ErrRevoked means a well-formed revoked record exists for the key ID.
ErrRevoked = errors.New("credential record is revoked")
// ErrConflict means an active record cannot be replaced or resurrected.
ErrConflict = errors.New("credential record already exists")
// ErrIntegrity means an unsafe or malformed registry object was observed.
ErrIntegrity = errors.New("registry integrity failure")
)
// PublicRecord is the redacted inventory form of a persisted record.
type PublicRecord struct {
Kind record.Kind `json:"credential_kind"`
KeyID string `json:"key_id"`
InstallationID string `json:"installation_id"`
Description string `json:"description,omitempty"`
CreatedAt time.Time `json:"created_at"`
ExpiresAt *time.Time `json:"expires_at,omitempty"`
State State `json:"state"`
RevokedAt *time.Time `json:"revoked_at,omitempty"`
RevocationReason string `json:"revocation_reason,omitempty"`
}
// Store owns protected descriptors for one registry root.
type Store struct {
root *securefile.Dir
active *securefile.Dir
revoked *securefile.Dir
}
type storedRecord struct {
record record.Record
state State
}
// Open opens one existing protected absolute registry root and creates only its
// active and revoked child directories when absent.
func Open(root string) (*Store, error) {
rootDir, err := securefile.OpenDir(root)
if err != nil {
return nil, integrity(err)
}
active, err := rootDir.OpenOrCreateDir(string(StateActive), 0o750)
if err != nil {
_ = rootDir.Close()
return nil, integrity(err)
}
revoked, err := rootDir.OpenOrCreateDir(string(StateRevoked), 0o750)
if err != nil {
_ = active.Close()
_ = rootDir.Close()
return nil, integrity(err)
}
return &Store{root: rootDir, active: active, revoked: revoked}, nil
}
// Close closes descriptors held by the store.
func (s *Store) Close() error {
if s == nil {
return nil
}
var first error
for _, dir := range []*securefile.Dir{s.revoked, s.active, s.root} {
if dir == nil {
continue
}
if err := dir.Close(); err != nil && first == nil {
first = err
}
}
return first
}
// Add validates and atomically publishes one active record. Existing active or
// revoked records cannot be overwritten or resurrected.
func (s *Store) Add(value record.Record) error {
if err := validateForState(value, StateActive); err != nil {
return err
}
return s.withWriterLock(func() 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 active {
if existing.record.KeyID == value.KeyID {
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)
})
}
// Find returns an active credential record. A revoked counterpart is checked
// first and always wins, including during the safe revoke-publication overlap.
func (s *Store) Find(keyID string) (record.Record, error) {
if !validKeyID(keyID) {
return record.Record{}, ErrNotFound
}
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)
}
// 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) {
active, revoked, err := s.scanAll()
if err != nil {
return record.Record{}, err
}
if err := validateLegacyMultiplicity(active, revoked); err != nil {
return record.Record{}, err
}
for _, existing := range revoked {
if existing.record.Kind == record.KindLegacyRaw {
return record.Record{}, ErrRevoked
}
}
for _, existing := range active {
if existing.record.Kind == record.KindLegacyRaw {
return existing.record, nil
}
}
return record.Record{}, ErrNotFound
}
// 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) {
active, revoked, err := s.scanAll()
if err != nil {
return nil, err
}
if err := validateLegacyMultiplicity(active, revoked); err != nil {
return nil, err
}
byKey := make(map[string]PublicRecord, len(active)+len(revoked))
for _, existing := range active {
byKey[existing.record.KeyID] = public(existing.record, StateActive)
}
for _, existing := range revoked {
byKey[existing.record.KeyID] = public(existing.record, StateRevoked)
}
keys := make([]string, 0, len(byKey))
for keyID := range byKey {
keys = append(keys, keyID)
}
sort.Strings(keys)
result := make([]PublicRecord, 0, len(keys))
for _, keyID := range keys {
result = append(result, byKey[keyID])
}
return result, nil
}
// Revoke publishes a validated revoked record and fsyncs it before removing the
// active record. If deletion then fails, Find still returns ErrRevoked.
func (s *Store) Revoke(keyID, reason string, at time.Time) error {
if !validKeyID(keyID) {
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
})
}
// Check validates every record and protected directory without exposing any
// digest data.
func (s *Store) Check() error {
active, revoked, err := s.scanAll()
if err != nil {
return err
}
return validateLegacyMultiplicity(active, revoked)
}
func (s *Store) withWriterLock(fn func() error) error {
if s == nil || s.root == nil || s.active == nil || s.revoked == nil {
return integrity(errors.New("uninitialized store"))
}
lock, err := s.root.Lock(".writer.lock")
if err != nil {
return integrity(err)
}
defer lock.Close()
return fn()
}
func (s *Store) scanAll() ([]storedRecord, []storedRecord, error) {
if s == nil || s.active == nil || s.revoked == nil {
return nil, nil, integrity(errors.New("uninitialized store"))
}
active, err := s.scan(s.active, StateActive)
if err != nil {
return nil, nil, err
}
revoked, err := s.scan(s.revoked, StateRevoked)
if err != nil {
return nil, nil, err
}
return active, revoked, nil
}
func (s *Store) scan(dir *securefile.Dir, state State) ([]storedRecord, error) {
names, err := dir.Names()
if err != nil {
return nil, integrity(err)
}
sort.Strings(names)
result := make([]storedRecord, 0, len(names))
for _, name := range names {
if !strings.HasSuffix(name, ".json") {
return nil, integrity(fmt.Errorf("unexpected registry entry %q", name))
}
keyID := strings.TrimSuffix(name, ".json")
if !validKeyID(keyID) {
return nil, integrity(fmt.Errorf("invalid registry filename %q", name))
}
value, err := s.load(dir, state, keyID)
if err != nil {
return nil, err
}
result = append(result, storedRecord{record: value, state: state})
}
return result, nil
}
func (s *Store) load(dir *securefile.Dir, state State, keyID string) (record.Record, error) {
data, err := dir.ReadFile(recordFileName(keyID), maxRecordBytes)
if errors.Is(err, syscall.ENOENT) {
return record.Record{}, ErrNotFound
}
if err != nil {
return record.Record{}, integrity(err)
}
value, err := decodeRecord(data)
if err != nil {
return record.Record{}, integrity(err)
}
if value.KeyID != keyID {
return record.Record{}, integrity(fmt.Errorf("record key ID does not match filename"))
}
if err := validateForState(value, state); err != nil {
return record.Record{}, integrity(err)
}
return value, nil
}
func (s *Store) writeRecord(dir *securefile.Dir, value record.Record) (err error) {
data, err := json.Marshal(value)
if err != nil {
return err
}
data = append(data, '\n')
if len(data) > maxRecordBytes {
return fmt.Errorf("record exceeds %d byte bound", maxRecordBytes)
}
temporary, err := temporaryName()
if err != nil {
return err
}
file, err := dir.CreateExclusive(temporary, 0o600)
if err != nil {
return integrity(err)
}
published := false
defer func() {
if file != nil {
_ = file.Close()
}
if !published {
_ = dir.Remove(temporary)
_ = dir.Sync()
}
}()
if err := writeAll(file, data); err != nil {
return err
}
if err := file.Chmod(0o640); err != nil {
return err
}
if err := file.Sync(); err != nil {
return err
}
if err := file.Close(); err != nil {
return err
}
file = nil
name := recordFileName(value.KeyID)
exists, err := dir.Exists(name)
if err != nil {
return integrity(err)
}
if exists {
return ErrConflict
}
if err := dir.Rename(temporary, dir, name); err != nil {
return integrity(err)
}
published = true
if err := dir.Sync(); err != nil {
return integrity(err)
}
return nil
}
func validateForState(value record.Record, state State) error {
if err := record.Validate(value); err != nil {
return err
}
switch state {
case StateActive:
if value.RevokedAt != nil || value.RevocationReason != "" {
return errors.New("active record carries revocation data")
}
case StateRevoked:
if value.RevokedAt == nil || value.RevocationReason == "" {
return errors.New("revoked record lacks revocation data")
}
default:
return errors.New("unknown registry state")
}
return nil
}
func validateLegacyMultiplicity(active, revoked []storedRecord) error {
activeCount := 0
revokedCount := 0
for _, existing := range active {
if existing.record.Kind == record.KindLegacyRaw {
activeCount++
}
}
for _, existing := range revoked {
if existing.record.Kind == record.KindLegacyRaw {
revokedCount++
}
}
if activeCount > 1 || revokedCount > 1 || activeCount+revokedCount > 2 {
return integrity(errors.New("multiple legacy records"))
}
return nil
}
func decodeRecord(data []byte) (record.Record, error) {
if err := rejectDuplicateOrTrailingJSON(data); err != nil {
return record.Record{}, err
}
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
var value record.Record
if err := decoder.Decode(&value); err != nil {
return record.Record{}, err
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
if err == nil {
return record.Record{}, errors.New("trailing JSON value")
}
return record.Record{}, err
}
if err := record.Validate(value); err != nil {
return record.Record{}, err
}
return value, nil
}
func rejectDuplicateOrTrailingJSON(data []byte) error {
decoder := json.NewDecoder(bytes.NewReader(data))
if err := consumeJSONValue(decoder); err != nil {
return err
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
if err == nil {
return errors.New("trailing JSON value")
}
return err
}
return nil
}
func consumeJSONValue(decoder *json.Decoder) error {
token, err := decoder.Token()
if err != nil {
return err
}
delimiter, isDelimiter := token.(json.Delim)
if !isDelimiter {
return nil
}
switch delimiter {
case '{':
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 _, 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")
}
case '[':
for decoder.More() {
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 array")
}
default:
return errors.New("unexpected JSON delimiter")
}
return nil
}
func public(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 recordFileName(keyID string) string {
return keyID + ".json"
}
func validKeyID(keyID string) bool {
if keyID == record.LegacyKeyID {
return true
}
if len(keyID) != 16 {
return false
}
decoded, err := base64.RawURLEncoding.DecodeString(keyID)
return err == nil && len(decoded) == 12 && base64.RawURLEncoding.EncodeToString(decoded) == keyID
}
func temporaryName() (string, error) {
bytes := make([]byte, 16)
if _, err := io.ReadFull(rand.Reader, bytes); err != nil {
return "", err
}
return ".tmp-" + hex.EncodeToString(bytes), nil
}
func writeAll(file *os.File, data []byte) error {
for len(data) > 0 {
written, err := file.Write(data)
if err != nil {
return err
}
if written == 0 {
return io.ErrShortWrite
}
data = data[written:]
}
return nil
}
func integrity(err error) error {
if err == nil {
return nil
}
if errors.Is(err, ErrIntegrity) {
return err
}
return fmt.Errorf("%w: %v", ErrIntegrity, err)
}
@@ -0,0 +1,457 @@
//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 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: "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")
}
})
}
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 := openStore(t, root)
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 := right.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)
}
}