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