feat: add protected DWH credential registry
This commit is contained in:
@@ -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)
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user