fix: harden DWH credential registry reads
This commit is contained in:
@@ -17,6 +17,7 @@ import (
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
@@ -60,6 +61,8 @@ type PublicRecord struct {
|
||||
|
||||
// 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
|
||||
@@ -88,7 +91,7 @@ func Open(root string) (*Store, error) {
|
||||
_ = rootDir.Close()
|
||||
return nil, integrity(err)
|
||||
}
|
||||
return &Store{root: rootDir, active: active, revoked: revoked}, nil
|
||||
return &Store{root: rootDir, active: active, revoked: revoked, now: time.Now}, nil
|
||||
}
|
||||
|
||||
// Close closes descriptors held by the store.
|
||||
@@ -96,6 +99,8 @@ 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 {
|
||||
@@ -115,29 +120,33 @@ func (s *Store) Add(value record.Record) error {
|
||||
return err
|
||||
}
|
||||
return s.withWriterLock(func() error {
|
||||
active, revoked, err := s.scanAll()
|
||||
if err != nil {
|
||||
return err
|
||||
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 revoked {
|
||||
if existing.record.KeyID == value.KeyID {
|
||||
return ErrRevoked
|
||||
}
|
||||
}
|
||||
for _, existing := range active {
|
||||
if existing.record.KeyID == value.KeyID {
|
||||
return ErrConflict
|
||||
}
|
||||
for _, existing := range active {
|
||||
if existing.record.KeyID == value.KeyID {
|
||||
}
|
||||
if value.Kind == record.KindLegacyRaw {
|
||||
for _, existing := range append(active, revoked...) {
|
||||
if existing.record.Kind == record.KindLegacyRaw {
|
||||
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)
|
||||
})
|
||||
}
|
||||
return s.writeRecord(s.active, value)
|
||||
}
|
||||
|
||||
// Find returns an active credential record. A revoked counterpart is checked
|
||||
@@ -146,17 +155,42 @@ func (s *Store) Find(keyID string) (record.Record, error) {
|
||||
if !validKeyID(keyID) {
|
||||
return record.Record{}, ErrNotFound
|
||||
}
|
||||
if s == nil {
|
||||
return record.Record{}, integrity(errors.New("uninitialized store"))
|
||||
}
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.findUnlocked(keyID)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
return s.load(s.active, StateActive, keyID)
|
||||
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) {
|
||||
if s == nil {
|
||||
return record.Record{}, integrity(errors.New("uninitialized store"))
|
||||
}
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.findLegacyUnlocked()
|
||||
}
|
||||
|
||||
func (s *Store) findLegacyUnlocked() (record.Record, error) {
|
||||
active, revoked, err := s.scanAll()
|
||||
if err != nil {
|
||||
return record.Record{}, err
|
||||
@@ -171,6 +205,9 @@ func (s *Store) FindLegacy() (record.Record, error) {
|
||||
}
|
||||
for _, existing := range active {
|
||||
if existing.record.Kind == record.KindLegacyRaw {
|
||||
if s.expired(existing.record) {
|
||||
return record.Record{}, ErrNotFound
|
||||
}
|
||||
return existing.record, nil
|
||||
}
|
||||
}
|
||||
@@ -180,6 +217,15 @@ func (s *Store) FindLegacy() (record.Record, error) {
|
||||
// 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) {
|
||||
if s == nil {
|
||||
return nil, integrity(errors.New("uninitialized store"))
|
||||
}
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.listUnlocked()
|
||||
}
|
||||
|
||||
func (s *Store) listUnlocked() ([]PublicRecord, error) {
|
||||
active, revoked, err := s.scanAll()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -213,47 +259,60 @@ func (s *Store) Revoke(keyID, reason string, at time.Time) error {
|
||||
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
|
||||
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 {
|
||||
if s == nil {
|
||||
return integrity(errors.New("uninitialized store"))
|
||||
}
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.checkUnlocked()
|
||||
}
|
||||
|
||||
func (s *Store) checkUnlocked() error {
|
||||
active, revoked, err := s.scanAll()
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -262,7 +321,12 @@ func (s *Store) Check() error {
|
||||
}
|
||||
|
||||
func (s *Store) withWriterLock(fn func() error) error {
|
||||
if s == nil || s.root == nil || s.active == nil || s.revoked == nil {
|
||||
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(".writer.lock")
|
||||
@@ -410,6 +474,17 @@ func validateForState(value record.Record, state State) error {
|
||||
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
|
||||
@@ -452,9 +527,22 @@ func decodeRecord(data []byte) (record.Record, error) {
|
||||
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 := consumeJSONValue(decoder); err != nil {
|
||||
if err := consumeRecordObject(decoder); err != nil {
|
||||
return err
|
||||
}
|
||||
var extra any
|
||||
@@ -467,6 +555,45 @@ func rejectDuplicateOrTrailingJSON(data []byte) error {
|
||||
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 {
|
||||
|
||||
Reference in New Issue
Block a user