//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" "sync" "syscall" "time" "github.com/aritmolab/thothii/tools/dwh-auth/internal/record" "github.com/aritmolab/thothii/tools/dwh-auth/internal/securefile" ) const ( maxRecordBytes = 4096 registryLockName = ".writer.lock" ) // 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") // ErrReadOnly means a runtime reader was asked to mutate registry state. ErrReadOnly = errors.New("registry is read-only") ) // 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 { mu sync.RWMutex now func() time.Time root *securefile.Dir active *securefile.Dir revoked *securefile.Dir readOnly bool } 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) } lock, err := rootDir.Lock(registryLockName) if err != nil { _ = revoked.Close() _ = active.Close() _ = rootDir.Close() return nil, integrity(err) } if err := lock.Close(); err != nil { _ = revoked.Close() _ = active.Close() _ = rootDir.Close() return nil, integrity(err) } return &Store{root: rootDir, active: active, revoked: revoked, now: time.Now}, nil } // OpenReadOnly opens only a complete preprovisioned protected registry. It // neither creates registry paths nor permits mutations through the Store. func OpenReadOnly(root string) (*Store, error) { rootDir, err := securefile.OpenDir(root) if err != nil { return nil, integrity(err) } active, err := rootDir.OpenDir(string(StateActive)) if err != nil { _ = rootDir.Close() return nil, integrity(err) } revoked, err := rootDir.OpenDir(string(StateRevoked)) if err != nil { _ = active.Close() _ = rootDir.Close() return nil, integrity(err) } lock, err := rootDir.LockShared(registryLockName) if err != nil { _ = revoked.Close() _ = active.Close() _ = rootDir.Close() return nil, integrity(err) } if err := lock.Close(); err != nil { _ = revoked.Close() _ = active.Close() _ = rootDir.Close() return nil, integrity(err) } return &Store{root: rootDir, active: active, revoked: revoked, now: time.Now, readOnly: true}, nil } // Close closes descriptors held by the store. func (s *Store) Close() error { if s == nil { return nil } s.mu.Lock() defer s.mu.Unlock() 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 s != nil && s.readOnly { return ErrReadOnly } if err := validateForState(value, StateActive); err != nil { return err } return s.withWriterLock(func() error { return s.addUnlocked(value) }) } func (s *Store) addUnlocked(value record.Record) 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 } var value record.Record err := s.withReaderLock(func() error { var findErr error value, findErr = s.findUnlocked(keyID) return findErr }) return value, err } func (s *Store) findUnlocked(keyID string) (record.Record, error) { if _, err := s.load(s.revoked, StateRevoked, keyID); err == nil { return record.Record{}, ErrRevoked } else if !errors.Is(err, ErrNotFound) { return record.Record{}, err } value, err := s.load(s.active, StateActive, keyID) if err != nil { return record.Record{}, err } if s.expired(value) { return record.Record{}, ErrNotFound } return value, nil } // 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) { var value record.Record err := s.withReaderLock(func() error { var findErr error value, findErr = s.findLegacyUnlocked() return findErr }) return value, err } func (s *Store) findLegacyUnlocked() (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 { if s.expired(existing.record) { return record.Record{}, ErrNotFound } 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) { var records []PublicRecord err := s.withReaderLock(func() error { var listErr error records, listErr = s.listUnlocked() return listErr }) return records, err } func (s *Store) listUnlocked() ([]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 s != nil && s.readOnly { return ErrReadOnly } if !validKeyID(keyID) { return ErrNotFound } return s.withWriterLock(func() error { return s.revokeUnlocked(keyID, reason, at) }) } func (s *Store) revokeUnlocked(keyID, reason string, at time.Time) 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 { return s.withReaderLock(s.checkUnlocked) } func (s *Store) checkUnlocked() error { active, revoked, err := s.scanAll() if err != nil { return err } return validateLegacyMultiplicity(active, revoked) } func (s *Store) withReaderLock(fn func() error) error { if s == nil { return integrity(errors.New("uninitialized store")) } s.mu.RLock() defer s.mu.RUnlock() if s.root == nil || s.active == nil || s.revoked == nil { return integrity(errors.New("uninitialized store")) } lock, err := s.root.LockShared(registryLockName) if err != nil { return integrity(err) } defer lock.Close() return fn() } func (s *Store) withWriterLock(fn func() error) error { if s == nil { return integrity(errors.New("uninitialized store")) } s.mu.Lock() defer s.mu.Unlock() if s.root == nil || s.active == nil || s.revoked == nil { return integrity(errors.New("uninitialized store")) } lock, err := s.root.Lock(registryLockName) 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 (s *Store) expired(value record.Record) bool { return value.ExpiresAt != nil && !value.ExpiresAt.After(s.nowUTC()) } func (s *Store) nowUTC() time.Time { if s.now == nil { return time.Now().UTC() } return s.now().UTC() } 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 > 1 { 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 } var recordJSONFields = map[string]struct{}{ "schema_version": {}, "credential_kind": {}, "key_id": {}, "installation_id": {}, "description": {}, "secret_sha256": {}, "created_at": {}, "expires_at": {}, "revoked_at": {}, "revocation_reason": {}, } func rejectDuplicateOrTrailingJSON(data []byte) error { decoder := json.NewDecoder(bytes.NewReader(data)) if err := consumeRecordObject(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 consumeRecordObject(decoder *json.Decoder) error { token, err := decoder.Token() if err != nil { return err } if token != json.Delim('{') { return errors.New("record JSON is not an object") } 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 _, allowed := recordJSONFields[key]; !allowed { return fmt.Errorf("unknown JSON field %q", key) } 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") } 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) }