diff --git a/tools/dwh-auth/internal/registry/store.go b/tools/dwh-auth/internal/registry/store.go new file mode 100644 index 00000000..fd054b31 --- /dev/null +++ b/tools/dwh-auth/internal/registry/store.go @@ -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) +} diff --git a/tools/dwh-auth/internal/registry/store_test.go b/tools/dwh-auth/internal/registry/store_test.go new file mode 100644 index 00000000..f2edc5de --- /dev/null +++ b/tools/dwh-auth/internal/registry/store_test.go @@ -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) + } +} diff --git a/tools/dwh-auth/internal/securefile/securefile_linux.go b/tools/dwh-auth/internal/securefile/securefile_linux.go new file mode 100644 index 00000000..1f57fc65 --- /dev/null +++ b/tools/dwh-auth/internal/securefile/securefile_linux.go @@ -0,0 +1,535 @@ +//go:build linux + +// Package securefile provides Linux-only, no-follow filesystem primitives for +// protected credential material. +package securefile + +import ( + "errors" + "fmt" + "io" + "os" + "path/filepath" + "strings" + "syscall" +) + +var ( + // ErrUnsafe identifies a path, mode, or file type that cannot be trusted. + ErrUnsafe = errors.New("unsafe filesystem object") + // ErrTooLarge identifies a file that exceeds its caller-provided bound. + ErrTooLarge = errors.New("file exceeds size bound") + // ErrChanged identifies a file whose size changed while it was read. + ErrChanged = errors.New("file changed during read") +) + +// Dir is an open, protected directory descriptor. Operations are rooted at the +// descriptor rather than reopening attacker-controlled path prefixes. +type Dir struct { + fd int + path string +} + +// Lock serializes cooperating writers using an advisory lock on a protected +// regular file. +type Lock struct { + fd int +} + +// OpenDir opens an absolute directory without following any path component and +// rejects unsafe modes on the protected final directory. +func OpenDir(path string) (*Dir, error) { + clean, err := absolutePath(path) + if err != nil { + return nil, err + } + fd, err := openDirectoryPath(clean) + if err != nil { + return nil, err + } + if err := validateDirectoryFD(fd); err != nil { + _ = syscall.Close(fd) + return nil, err + } + if err := compareWithLstat(fd, clean, true); err != nil { + _ = syscall.Close(fd) + return nil, err + } + return &Dir{fd: fd, path: clean}, nil +} + +// Close closes the directory descriptor. +func (d *Dir) Close() error { + if d == nil || d.fd < 0 { + return nil + } + err := syscall.Close(d.fd) + d.fd = -1 + return err +} + +// OpenDir opens one direct child directory without following it. +func (d *Dir) OpenDir(name string) (*Dir, error) { + if err := validName(name); err != nil { + return nil, err + } + if err := d.check(); err != nil { + return nil, err + } + fd, err := syscall.Openat(d.fd, name, directoryOpenFlags, 0) + if err != nil { + return nil, err + } + if err := validateDirectoryFD(fd); err != nil { + _ = syscall.Close(fd) + return nil, err + } + path := filepath.Join(d.path, name) + if err := compareWithLstat(fd, path, true); err != nil { + _ = syscall.Close(fd) + return nil, err + } + return &Dir{fd: fd, path: path}, nil +} + +// OpenOrCreateDir opens a protected child directory or creates it with a +// restrictive mode. A collision is reopened and fully revalidated. +func (d *Dir) OpenOrCreateDir(name string, mode os.FileMode) (*Dir, error) { + child, err := d.OpenDir(name) + if err == nil { + return child, nil + } + if !errors.Is(err, syscall.ENOENT) { + return nil, err + } + if err := validMode(mode); err != nil { + return nil, err + } + if err := d.check(); err != nil { + return nil, err + } + if err := syscall.Mkdirat(d.fd, name, uint32(mode.Perm())); err != nil && !errors.Is(err, syscall.EEXIST) { + return nil, err + } + child, err = d.OpenDir(name) + if err != nil { + return nil, err + } + if err := d.Sync(); err != nil { + _ = child.Close() + return nil, err + } + return child, nil +} + +// ReadFile reads one protected regular child file, bounded by max bytes. It +// verifies file type and mode before and after the read and detects size races. +func (d *Dir) ReadFile(name string, max int) ([]byte, error) { + return d.readFile(name, max, nil) +} + +// ReadSecret reads an absolute secret input whose final parent is protected and +// whose file mode is exactly 0600. +func ReadSecret(path string, max int) ([]byte, error) { + clean, err := absolutePath(path) + if err != nil { + return nil, err + } + parent := filepath.Dir(clean) + name := filepath.Base(clean) + if err := validName(name); err != nil { + return nil, err + } + dir, err := OpenDir(parent) + if err != nil { + return nil, err + } + defer dir.Close() + mode := uint32(0o600) + return dir.readFile(name, max, &mode) +} + +// CreateSecret creates a new absolute output file with O_CREAT|O_EXCL and an +// exact mode of 0600. The caller owns and must close the returned file. +func CreateSecret(path string) (*os.File, error) { + clean, err := absolutePath(path) + if err != nil { + return nil, err + } + parent := filepath.Dir(clean) + name := filepath.Base(clean) + if err := validName(name); err != nil { + return nil, err + } + dir, err := OpenDir(parent) + if err != nil { + return nil, err + } + defer dir.Close() + return dir.CreateExclusive(name, 0o600) +} + +// CreateExclusive creates one direct child with O_CREAT|O_EXCL|O_NOFOLLOW and +// the exact requested restrictive mode. The caller owns and must close it. +func (d *Dir) CreateExclusive(name string, mode os.FileMode) (*os.File, error) { + if err := validName(name); err != nil { + return nil, err + } + if err := validMode(mode); err != nil { + return nil, err + } + if err := d.check(); err != nil { + return nil, err + } + fd, err := syscall.Openat(d.fd, name, syscall.O_WRONLY|syscall.O_CREAT|syscall.O_EXCL|syscall.O_NOFOLLOW|syscall.O_CLOEXEC, uint32(mode.Perm())) + if err != nil { + return nil, err + } + if err := syscall.Fchmod(fd, uint32(mode.Perm())); err != nil { + _ = syscall.Close(fd) + return nil, err + } + if err := validateRegularFD(fd, uint32(mode.Perm())); err != nil { + _ = syscall.Close(fd) + return nil, err + } + return os.NewFile(uintptr(fd), filepath.Join(d.path, name)), nil +} + +// Exists reports whether a direct protected regular child exists. An unsafe +// collision is an error rather than an absent file. +func (d *Dir) Exists(name string) (bool, error) { + if err := validName(name); err != nil { + return false, err + } + if err := d.check(); err != nil { + return false, err + } + fd, err := syscall.Openat(d.fd, name, syscall.O_RDONLY|syscall.O_NOFOLLOW|syscall.O_CLOEXEC, 0) + if errors.Is(err, syscall.ENOENT) { + return false, nil + } + if err != nil { + return false, err + } + defer syscall.Close(fd) + if err := validateRegularFD(fd, 0); err != nil { + return false, err + } + if err := compareWithLstat(fd, filepath.Join(d.path, name), false); err != nil { + return false, err + } + return true, nil +} + +// Rename atomically renames a direct child into another protected directory. +// Callers that need no-replace semantics must serialize writers and verify the +// destination is absent before calling Rename. +func (d *Dir) Rename(oldName string, destination *Dir, newName string) error { + if err := validName(oldName); err != nil { + return err + } + if err := validName(newName); err != nil { + return err + } + if destination == nil { + return fmt.Errorf("%w: nil destination", ErrUnsafe) + } + if err := d.check(); err != nil { + return err + } + if err := destination.check(); err != nil { + return err + } + return syscall.Renameat(d.fd, oldName, destination.fd, newName) +} + +// Remove removes one direct child file without following it. +func (d *Dir) Remove(name string) error { + if err := validName(name); err != nil { + return err + } + if err := d.check(); err != nil { + return err + } + return syscall.Unlinkat(d.fd, name) +} + +// Sync makes prior directory entry changes durable. +func (d *Dir) Sync() error { + if err := d.check(); err != nil { + return err + } + return syscall.Fsync(d.fd) +} + +// Names returns direct entry names from the protected directory descriptor. +func (d *Dir) Names() ([]string, error) { + if err := d.check(); err != nil { + return nil, err + } + fd, err := syscall.Openat(d.fd, ".", directoryOpenFlags, 0) + if err != nil { + return nil, err + } + file := os.NewFile(uintptr(fd), d.path) + entries, err := file.ReadDir(-1) + closeErr := file.Close() + if err != nil { + return nil, err + } + if closeErr != nil { + return nil, closeErr + } + names := make([]string, 0, len(entries)) + for _, entry := range entries { + names = append(names, entry.Name()) + } + return names, nil +} + +// Lock opens or creates a protected 0600 lock file and acquires an exclusive +// advisory lock. Close releases the lock and descriptor. +func (d *Dir) Lock(name string) (*Lock, error) { + if err := validName(name); err != nil { + return nil, err + } + if err := d.check(); err != nil { + return nil, err + } + for { + fd, err := syscall.Openat(d.fd, name, syscall.O_RDWR|syscall.O_NOFOLLOW|syscall.O_CLOEXEC, 0) + if errors.Is(err, syscall.ENOENT) { + fd, err = syscall.Openat(d.fd, name, syscall.O_RDWR|syscall.O_CREAT|syscall.O_EXCL|syscall.O_NOFOLLOW|syscall.O_CLOEXEC, 0o600) + if errors.Is(err, syscall.EEXIST) { + continue + } + if err != nil { + return nil, err + } + if err := syscall.Fchmod(fd, 0o600); err != nil { + _ = syscall.Close(fd) + return nil, err + } + if err := d.Sync(); err != nil { + _ = syscall.Close(fd) + return nil, err + } + } + if err != nil { + return nil, err + } + if err := validateRegularFD(fd, 0o600); err != nil { + _ = syscall.Close(fd) + return nil, err + } + if err := compareWithLstat(fd, filepath.Join(d.path, name), false); err != nil { + _ = syscall.Close(fd) + return nil, err + } + if err := syscall.Flock(fd, syscall.LOCK_EX); err != nil { + _ = syscall.Close(fd) + return nil, err + } + return &Lock{fd: fd}, nil + } +} + +// Close releases an advisory lock and closes its descriptor. +func (l *Lock) Close() error { + if l == nil || l.fd < 0 { + return nil + } + unlockErr := syscall.Flock(l.fd, syscall.LOCK_UN) + closeErr := syscall.Close(l.fd) + l.fd = -1 + if unlockErr != nil { + return unlockErr + } + return closeErr +} + +func (d *Dir) readFile(name string, max int, exactMode *uint32) ([]byte, error) { + if err := validName(name); err != nil { + return nil, err + } + if max < 0 { + return nil, fmt.Errorf("%w: negative size bound", ErrUnsafe) + } + if err := d.check(); err != nil { + return nil, err + } + fd, err := syscall.Openat(d.fd, name, syscall.O_RDONLY|syscall.O_NOFOLLOW|syscall.O_CLOEXEC, 0) + if err != nil { + return nil, err + } + defer syscall.Close(fd) + + var before syscall.Stat_t + if err := syscall.Fstat(fd, &before); err != nil { + return nil, err + } + if err := validateRegularStat(&before, exactMode); err != nil { + return nil, err + } + if err := compareWithLstat(fd, filepath.Join(d.path, name), false); err != nil { + return nil, err + } + if before.Size < 0 || before.Size > int64(max) { + return nil, fmt.Errorf("%w: %d bytes", ErrTooLarge, before.Size) + } + size := int(before.Size) + data := make([]byte, size+1) + n := 0 + for n < len(data) { + read, readErr := syscall.Read(fd, data[n:]) + if read > 0 { + n += read + } + if errors.Is(readErr, syscall.EINTR) { + continue + } + if readErr != nil { + return nil, readErr + } + if read == 0 { + break + } + } + + var after syscall.Stat_t + if err := syscall.Fstat(fd, &after); err != nil { + return nil, err + } + if err := validateRegularStat(&after, exactMode); err != nil { + return nil, err + } + if before.Dev != after.Dev || before.Ino != after.Ino || before.Size != after.Size { + return nil, ErrChanged + } + if n != size { + return nil, ErrChanged + } + return data[:n], nil +} + +func (d *Dir) check() error { + if d == nil || d.fd < 0 { + return fmt.Errorf("%w: closed directory", ErrUnsafe) + } + return validateDirectoryFD(d.fd) +} + +func absolutePath(path string) (string, error) { + if !filepath.IsAbs(path) { + return "", fmt.Errorf("%w: path must be absolute", ErrUnsafe) + } + return filepath.Clean(path), nil +} + +func openDirectoryPath(path string) (int, error) { + fd, err := syscall.Open("/", directoryOpenFlags, 0) + if err != nil { + return -1, err + } + if path == "/" { + return fd, nil + } + for _, part := range strings.Split(strings.TrimPrefix(path, "/"), "/") { + next, err := syscall.Openat(fd, part, directoryOpenFlags, 0) + _ = syscall.Close(fd) + if err != nil { + return -1, err + } + var stat syscall.Stat_t + if err := syscall.Fstat(next, &stat); err != nil { + _ = syscall.Close(next) + return -1, err + } + if stat.Mode&syscall.S_IFMT != syscall.S_IFDIR { + _ = syscall.Close(next) + return -1, fmt.Errorf("%w: non-directory path component", ErrUnsafe) + } + fd = next + } + return fd, nil +} + +func compareWithLstat(fd int, path string, directory bool) error { + var opened syscall.Stat_t + if err := syscall.Fstat(fd, &opened); err != nil { + return err + } + var linked syscall.Stat_t + if err := syscall.Lstat(path, &linked); err != nil { + return err + } + if directory { + if linked.Mode&syscall.S_IFMT != syscall.S_IFDIR { + return fmt.Errorf("%w: path is not a directory", ErrUnsafe) + } + } else if linked.Mode&syscall.S_IFMT != syscall.S_IFREG { + return fmt.Errorf("%w: path is not a regular file", ErrUnsafe) + } + if opened.Dev != linked.Dev || opened.Ino != linked.Ino { + return fmt.Errorf("%w: path changed while opening", ErrUnsafe) + } + return nil +} + +func validateDirectoryFD(fd int) error { + var stat syscall.Stat_t + if err := syscall.Fstat(fd, &stat); err != nil { + return err + } + if stat.Mode&syscall.S_IFMT != syscall.S_IFDIR { + return fmt.Errorf("%w: not a directory", ErrUnsafe) + } + if stat.Mode&0o7022 != 0 { + return fmt.Errorf("%w: unsafe directory mode %04o", ErrUnsafe, stat.Mode&0o7777) + } + return nil +} + +func validateRegularFD(fd int, exactMode uint32) error { + var stat syscall.Stat_t + if err := syscall.Fstat(fd, &stat); err != nil { + return err + } + var required *uint32 + if exactMode != 0 { + required = &exactMode + } + return validateRegularStat(&stat, required) +} + +func validateRegularStat(stat *syscall.Stat_t, exactMode *uint32) error { + if stat.Mode&syscall.S_IFMT != syscall.S_IFREG { + return fmt.Errorf("%w: not a regular file", ErrUnsafe) + } + if stat.Mode&0o7022 != 0 { + return fmt.Errorf("%w: unsafe file mode %04o", ErrUnsafe, stat.Mode&0o7777) + } + if exactMode != nil && stat.Mode&0o777 != *exactMode { + return fmt.Errorf("%w: file mode %04o is not %04o", ErrUnsafe, stat.Mode&0o777, *exactMode) + } + return nil +} + +func validMode(mode os.FileMode) error { + if mode&^os.FileMode(0o777) != 0 || mode.Perm()&0o022 != 0 { + return fmt.Errorf("%w: unsafe creation mode %04o", ErrUnsafe, mode) + } + return nil +} + +func validName(name string) error { + if name == "" || name == "." || name == ".." || strings.Contains(name, "/") || strings.ContainsRune(name, 0) { + return fmt.Errorf("%w: invalid path component", ErrUnsafe) + } + return nil +} + +const directoryOpenFlags = syscall.O_RDONLY | syscall.O_DIRECTORY | syscall.O_NOFOLLOW | syscall.O_CLOEXEC + +var _ io.Closer = (*Dir)(nil) diff --git a/tools/dwh-auth/internal/securefile/securefile_linux_test.go b/tools/dwh-auth/internal/securefile/securefile_linux_test.go new file mode 100644 index 00000000..6a68140d --- /dev/null +++ b/tools/dwh-auth/internal/securefile/securefile_linux_test.go @@ -0,0 +1,215 @@ +//go:build linux + +package securefile + +import ( + "bytes" + "errors" + "os" + "path/filepath" + "testing" +) + +func TestOpenDirAndReadFileAcceptProtectedRegularFile(t *testing.T) { + root := t.TempDir() + path := filepath.Join(root, "record.json") + want := []byte(`{"record":"synthetic"}`) + writeFile(t, path, want, 0o640) + + dir, err := OpenDir(root) + if err != nil { + t.Fatalf("OpenDir() error = %v", err) + } + t.Cleanup(func() { _ = dir.Close() }) + + got, err := dir.ReadFile("record.json", 4096) + if err != nil { + t.Fatalf("ReadFile() error = %v", err) + } + if !bytes.Equal(got, want) { + t.Fatalf("ReadFile() = %q, want %q", got, want) + } +} + +func TestProtectedPathsRejectSymlinks(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 := OpenDir(root); err == nil { + t.Fatal("OpenDir() error = nil, want symlink refusal") + } + }) + + t.Run("directory", func(t *testing.T) { + root := t.TempDir() + target := t.TempDir() + if err := os.Symlink(target, filepath.Join(root, "active")); err != nil { + t.Fatalf("Symlink() error = %v", err) + } + dir, err := OpenDir(root) + if err != nil { + t.Fatalf("OpenDir() error = %v", err) + } + t.Cleanup(func() { _ = dir.Close() }) + if _, err := dir.OpenDir("active"); err == nil { + t.Fatal("OpenDir(active) error = nil, want symlink refusal") + } + }) + + t.Run("record", func(t *testing.T) { + root := t.TempDir() + dir, err := OpenDir(root) + if err != nil { + t.Fatalf("OpenDir() error = %v", err) + } + t.Cleanup(func() { _ = dir.Close() }) + target := filepath.Join(root, "target.json") + writeFile(t, target, []byte(`{"safe":false}`), 0o640) + if err := os.Symlink(target, filepath.Join(root, "record.json")); err != nil { + t.Fatalf("Symlink() error = %v", err) + } + if _, err := dir.ReadFile("record.json", 4096); err == nil { + t.Fatal("ReadFile() error = nil, want symlink refusal") + } + }) + + t.Run("secret", func(t *testing.T) { + root := t.TempDir() + target := filepath.Join(root, "target.secret") + writeFile(t, target, []byte("synthetic-secret"), 0o600) + secret := filepath.Join(root, "secret") + if err := os.Symlink(target, secret); err != nil { + t.Fatalf("Symlink() error = %v", err) + } + if _, err := ReadSecret(secret, 128); err == nil { + t.Fatal("ReadSecret() error = nil, want symlink refusal") + } + }) +} + +func TestProtectedPathsRejectUnsafeModes(t *testing.T) { + t.Run("root", func(t *testing.T) { + root := t.TempDir() + if err := os.Chmod(root, 0o770); err != nil { + t.Fatalf("Chmod() error = %v", err) + } + if _, err := OpenDir(root); err == nil { + t.Fatal("OpenDir() error = nil, want unsafe mode refusal") + } + }) + + t.Run("child directory", func(t *testing.T) { + root := t.TempDir() + child := filepath.Join(root, "active") + if err := os.Mkdir(child, 0o770); err != nil { + t.Fatalf("Mkdir() error = %v", err) + } + if err := os.Chmod(child, 0o770); err != nil { + t.Fatalf("Chmod() error = %v", err) + } + dir, err := OpenDir(root) + if err != nil { + t.Fatalf("OpenDir() error = %v", err) + } + t.Cleanup(func() { _ = dir.Close() }) + if _, err := dir.OpenDir("active"); err == nil { + t.Fatal("OpenDir(active) error = nil, want unsafe mode refusal") + } + }) + + t.Run("record", func(t *testing.T) { + root := t.TempDir() + path := filepath.Join(root, "record.json") + writeFile(t, path, []byte(`{"unsafe":true}`), 0o660) + dir, err := OpenDir(root) + if err != nil { + t.Fatalf("OpenDir() error = %v", err) + } + t.Cleanup(func() { _ = dir.Close() }) + if _, err := dir.ReadFile("record.json", 4096); err == nil { + t.Fatal("ReadFile() error = nil, want unsafe mode refusal") + } + }) + + t.Run("secret", func(t *testing.T) { + root := t.TempDir() + secret := filepath.Join(root, "secret") + writeFile(t, secret, []byte("synthetic-secret"), 0o640) + if _, err := ReadSecret(secret, 128); err == nil { + t.Fatal("ReadSecret() error = nil, want non-0600 refusal") + } + }) +} + +func TestReadFileRejectsOversizeContent(t *testing.T) { + root := t.TempDir() + path := filepath.Join(root, "record.json") + writeFile(t, path, bytes.Repeat([]byte{'a'}, 4097), 0o640) + dir, err := OpenDir(root) + if err != nil { + t.Fatalf("OpenDir() error = %v", err) + } + t.Cleanup(func() { _ = dir.Close() }) + + if _, err := dir.ReadFile("record.json", 4096); err == nil { + t.Fatal("ReadFile() error = nil, want bounded-read refusal") + } +} + +func TestReadSecretAndCreateSecretUse0600AndExclusiveCreate(t *testing.T) { + root := t.TempDir() + input := filepath.Join(root, "input.secret") + want := []byte("synthetic-secret") + writeFile(t, input, want, 0o600) + + got, err := ReadSecret(input, 128) + if err != nil { + t.Fatalf("ReadSecret() error = %v", err) + } + if !bytes.Equal(got, want) { + t.Fatalf("ReadSecret() = %q, want %q", got, want) + } + + output := filepath.Join(root, "output.secret") + file, err := CreateSecret(output) + if err != nil { + t.Fatalf("CreateSecret() error = %v", err) + } + if _, err := file.Write(want); err != nil { + _ = file.Close() + t.Fatalf("Write() error = %v", err) + } + if err := file.Sync(); err != nil { + _ = file.Close() + t.Fatalf("Sync() error = %v", err) + } + if err := file.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + info, err := os.Stat(output) + if err != nil { + t.Fatalf("Stat() error = %v", err) + } + if got, want := info.Mode().Perm(), os.FileMode(0o600); got != want { + t.Fatalf("output mode = %04o, want %04o", got, want) + } + if _, err := CreateSecret(output); err == nil { + t.Fatal("CreateSecret(existing) error = nil, want exclusive-create refusal") + } else if !errors.Is(err, os.ErrExist) { + t.Fatalf("CreateSecret(existing) error = %v, want os.ErrExist", err) + } +} + +func writeFile(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) + } +}