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

458 lines
13 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 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)
}
}