feat(auth): add safe Argon2id local user registry
This commit is contained in:
+1
@@ -0,0 +1 @@
|
||||
[{"password":"correct horse battery staple","saltHex":"000102030405060708090a0b0c0d0e0f","memoryKiB":65536,"passes":3,"parallelism":1,"keyLength":32,"phc":"$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4"}]
|
||||
@@ -11,6 +11,7 @@ require (
|
||||
github.com/distribution/reference v0.6.0
|
||||
github.com/gofrs/flock v0.12.1
|
||||
github.com/sirupsen/logrus v1.9.1
|
||||
golang.org/x/crypto v0.55.0
|
||||
golang.org/x/sys v0.47.0
|
||||
)
|
||||
|
||||
|
||||
@@ -26,6 +26,8 @@ github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+
|
||||
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||
github.com/stretchr/testify v1.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg=
|
||||
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||
golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
package authconfig
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"io"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"golang.org/x/crypto/argon2"
|
||||
)
|
||||
|
||||
const (
|
||||
argon2MemoryKiB uint32 = 65536
|
||||
argon2Passes uint32 = 3
|
||||
argon2Parallel uint8 = 1
|
||||
argon2SaltBytes = 16
|
||||
argon2KeyBytes uint32 = 32
|
||||
|
||||
argon2MaximumMemoryKiB uint32 = 256 * 1024
|
||||
argon2MaximumPasses uint32 = 10
|
||||
argon2MaximumParallel uint8 = 4
|
||||
argon2MinimumSaltBytes = 16
|
||||
argon2MaximumSaltBytes = 64
|
||||
argon2MinimumKeyBytes uint32 = 16
|
||||
argon2MaximumKeyBytes uint32 = 64
|
||||
maximumPHCBytes = 256
|
||||
)
|
||||
|
||||
type argon2Parameters struct {
|
||||
memory uint32
|
||||
passes uint32
|
||||
parallelism uint8
|
||||
salt []byte
|
||||
digest []byte
|
||||
}
|
||||
|
||||
// HashPassword derives a fixed-parameter Argon2id v19 PHC string. random must be a
|
||||
// cryptographically secure source in production; it is injected only to permit the shared vector.
|
||||
func HashPassword(password []byte, random io.Reader) (string, error) {
|
||||
if !validPassword(password) {
|
||||
return "", errInvalidAuthenticationConfig
|
||||
}
|
||||
if random == nil {
|
||||
random = rand.Reader
|
||||
}
|
||||
salt := make([]byte, argon2SaltBytes)
|
||||
if _, err := io.ReadFull(random, salt); err != nil {
|
||||
return "", errInvalidAuthenticationConfig
|
||||
}
|
||||
digest := argon2.IDKey(password, salt, argon2Passes, argon2MemoryKiB, argon2Parallel, argon2KeyBytes)
|
||||
return "$argon2id$v=19$m=65536,t=3,p=1$" + base64.RawStdEncoding.EncodeToString(salt) + "$" + base64.RawStdEncoding.EncodeToString(digest), nil
|
||||
}
|
||||
|
||||
// VerifyPassword accepts only bounded, canonical Argon2id v19 PHC strings.
|
||||
func VerifyPassword(password []byte, encoded string) bool {
|
||||
if !validPassword(password) {
|
||||
return false
|
||||
}
|
||||
parameters, ok := parsePHC(encoded)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
derived := argon2.IDKey(password, parameters.salt, parameters.passes, parameters.memory, parameters.parallelism, uint32(len(parameters.digest)))
|
||||
return subtle.ConstantTimeCompare(derived, parameters.digest) == 1
|
||||
}
|
||||
|
||||
func validPassword(password []byte) bool {
|
||||
return len(password) >= passwordMinBytes && len(password) <= passwordMaxBytes && utf8.Valid(password)
|
||||
}
|
||||
|
||||
func parsePHC(encoded string) (argon2Parameters, bool) {
|
||||
if len(encoded) == 0 || len(encoded) > maximumPHCBytes {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
parts := strings.Split(encoded, "$")
|
||||
if len(parts) != 6 || parts[0] != "" || parts[1] != "argon2id" || parts[2] != "v=19" {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
parameters, ok := parsePHCParameters(parts[3])
|
||||
if !ok {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
salt, ok := decodePHCBase64(parts[4], argon2MinimumSaltBytes, argon2MaximumSaltBytes)
|
||||
if !ok {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
digest, ok := decodePHCBase64(parts[5], int(argon2MinimumKeyBytes), int(argon2MaximumKeyBytes))
|
||||
if !ok {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
parameters.salt = salt
|
||||
parameters.digest = digest
|
||||
return parameters, true
|
||||
}
|
||||
|
||||
func parsePHCParameters(value string) (argon2Parameters, bool) {
|
||||
parts := strings.Split(value, ",")
|
||||
if len(parts) != 3 || !strings.HasPrefix(parts[0], "m=") || !strings.HasPrefix(parts[1], "t=") || !strings.HasPrefix(parts[2], "p=") {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
memory, ok := parseDecimal(parts[0][2:], uint64(argon2MaximumMemoryKiB))
|
||||
if !ok || memory < 8 {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
passes, ok := parseDecimal(parts[1][2:], uint64(argon2MaximumPasses))
|
||||
if !ok || passes == 0 {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
parallelism, ok := parseDecimal(parts[2][2:], uint64(argon2MaximumParallel))
|
||||
if !ok || parallelism == 0 || memory < 8*parallelism {
|
||||
return argon2Parameters{}, false
|
||||
}
|
||||
return argon2Parameters{memory: uint32(memory), passes: uint32(passes), parallelism: uint8(parallelism)}, true
|
||||
}
|
||||
|
||||
func parseDecimal(value string, maximum uint64) (uint64, bool) {
|
||||
if value == "" || (len(value) > 1 && value[0] == '0') || len(value) > 10 {
|
||||
return 0, false
|
||||
}
|
||||
for _, character := range value {
|
||||
if character < '0' || character > '9' {
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
parsed, err := strconv.ParseUint(value, 10, 64)
|
||||
if err != nil || parsed > maximum {
|
||||
return 0, false
|
||||
}
|
||||
return parsed, true
|
||||
}
|
||||
|
||||
func decodePHCBase64(value string, minimum, maximum int) ([]byte, bool) {
|
||||
if value == "" || strings.ContainsRune(value, '=') || len(value) > base64.RawStdEncoding.EncodedLen(maximum) || len(value) < base64.RawStdEncoding.EncodedLen(minimum)-1 {
|
||||
return nil, false
|
||||
}
|
||||
decoded, err := base64.RawStdEncoding.DecodeString(value)
|
||||
if err != nil || len(decoded) < minimum || len(decoded) > maximum {
|
||||
return nil, false
|
||||
}
|
||||
return decoded, true
|
||||
}
|
||||
|
||||
func validatePasswordHash(encoded string) error {
|
||||
if _, ok := parsePHC(encoded); !ok {
|
||||
return errors.New("invalid password hash")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package authconfig
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type passwordVector struct {
|
||||
Password string `json:"password"`
|
||||
SaltHex string `json:"saltHex"`
|
||||
MemoryKiB uint32 `json:"memoryKiB"`
|
||||
Passes uint32 `json:"passes"`
|
||||
Parallelism uint8 `json:"parallelism"`
|
||||
KeyLength uint32 `json:"keyLength"`
|
||||
PHC string `json:"phc"`
|
||||
}
|
||||
|
||||
func TestHashPasswordMatchesSharedFixedArgon2idVector(t *testing.T) {
|
||||
vector := loadPasswordVector(t)
|
||||
salt, err := hex.DecodeString(vector.SaltHex)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
got, err := HashPassword([]byte(vector.Password), bytes.NewReader(salt))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != vector.PHC {
|
||||
t.Fatalf("HashPassword() = %q, want shared PHC vector %q", got, vector.PHC)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyPasswordAcceptsOnlyTheMatchingPassword(t *testing.T) {
|
||||
vector := loadPasswordVector(t)
|
||||
if !VerifyPassword([]byte(vector.Password), vector.PHC) {
|
||||
t.Fatal("VerifyPassword() rejected the shared fixed vector")
|
||||
}
|
||||
if VerifyPassword([]byte(vector.Password+"!"), vector.PHC) {
|
||||
t.Fatal("VerifyPassword() accepted a one-byte password change")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyPasswordRejectsMalformedAndOversizedPHCBeforeHashing(t *testing.T) {
|
||||
password := []byte("correct horse battery staple")
|
||||
for name, encoded := range map[string]string{
|
||||
"wrong algorithm": "$argon2i$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4",
|
||||
"wrong version": "$argon2id$v=18$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4",
|
||||
"wrong parameter order": "$argon2id$v=19$t=3,m=65536,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4",
|
||||
"padded base64": "$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw==$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4",
|
||||
"memory too large": "$argon2id$v=19$m=262145,t=1,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4",
|
||||
"passes too large": "$argon2id$v=19$m=65536,t=11,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4",
|
||||
"parallelism too large": "$argon2id$v=19$m=65536,t=3,p=5$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4",
|
||||
"short salt": "$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0O$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4",
|
||||
"long digest": "$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if VerifyPassword(password, encoded) {
|
||||
t.Fatal("VerifyPassword() accepted an invalid PHC string")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestHashPasswordRejectsPasswordsOutsideUTF8ByteBounds(t *testing.T) {
|
||||
for name, password := range map[string][]byte{
|
||||
"too short": bytes.Repeat([]byte("a"), 11),
|
||||
"too long": bytes.Repeat([]byte("a"), 1025),
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if _, err := HashPassword(password, bytes.NewReader(make([]byte, 16))); err == nil {
|
||||
t.Fatal("HashPassword() accepted an out-of-bounds password")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func loadPasswordVector(t *testing.T) passwordVector {
|
||||
t.Helper()
|
||||
_, testFile, _, ok := runtime.Caller(0)
|
||||
if !ok {
|
||||
t.Fatal("locate password test fixture")
|
||||
}
|
||||
path := filepath.Join(filepath.Dir(testFile), "..", "..", "..", "..", "backend", "test", "fixtures", "argon2id-vectors.json")
|
||||
contents, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var vectors []passwordVector
|
||||
if err := json.Unmarshal(contents, &vectors); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(vectors) != 1 || vectors[0].Password == "" || strings.Contains(vectors[0].PHC, "\n") {
|
||||
t.Fatal("shared Argon2id fixture is invalid")
|
||||
}
|
||||
return vectors[0]
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
package authconfig
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
|
||||
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
|
||||
"github.com/gofrs/flock"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
const (
|
||||
authFileName = "auth.yaml"
|
||||
usersFileName = "users.yaml"
|
||||
lockFileName = ".auth.lock"
|
||||
)
|
||||
|
||||
// Load reads auth.yaml and, for local authentication, the paired users.yaml registry. All reads
|
||||
// are bounded and reject non-private, linked, or symlinked files.
|
||||
func Load(directory string) (Config, Registry, error) {
|
||||
return load(directory)
|
||||
}
|
||||
|
||||
func load(directory string) (Config, Registry, error) {
|
||||
if err := requirePrivateDirectory(directory); err != nil {
|
||||
return Config{}, Registry{}, err
|
||||
}
|
||||
authContents, err := readPrivateFile(filepath.Join(directory, authFileName))
|
||||
if err != nil {
|
||||
return Config{}, Registry{}, err
|
||||
}
|
||||
var config Config
|
||||
if err := decodeStrictYAML(authContents, &config); err != nil || !validConfig(config) {
|
||||
return Config{}, Registry{}, errInvalidAuthenticationConfig
|
||||
}
|
||||
if config.Mode != "local" {
|
||||
return config, Registry{}, nil
|
||||
}
|
||||
usersContents, err := readPrivateFile(filepath.Join(directory, config.Local.UsersFile))
|
||||
if err != nil {
|
||||
return Config{}, Registry{}, err
|
||||
}
|
||||
var registry Registry
|
||||
if err := decodeStrictYAML(usersContents, ®istry); err != nil {
|
||||
return Config{}, Registry{}, errInvalidAuthenticationConfig
|
||||
}
|
||||
if err := validateRegistry(registry); err != nil {
|
||||
return Config{}, Registry{}, errInvalidAuthenticationConfig
|
||||
}
|
||||
return config, registry, nil
|
||||
}
|
||||
|
||||
// MutateUsers serializes the entire read-check-write transaction under the configuration lock.
|
||||
// The resulting registry is revalidated and atomically replaced only after all invariants hold.
|
||||
func MutateUsers(directory string, mutate func(*Registry) error) error {
|
||||
if mutate == nil {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
if err := requirePrivateDirectory(directory); err != nil {
|
||||
return err
|
||||
}
|
||||
lock, err := acquireLock(directory)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = lock.Unlock() }()
|
||||
|
||||
config, registry, err := load(directory)
|
||||
if err != nil || config.Mode != "local" {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
before := cloneRegistry(registry)
|
||||
if err := mutate(®istry); err != nil {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
if err := applyMutationInvariants(before, ®istry); err != nil {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
contents, err := yaml.Marshal(registry)
|
||||
if err != nil {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
contents = append(contents, '\n')
|
||||
if err := safeio.ReplaceCanonicalRegular(filepath.Join(directory, config.Local.UsersFile), contents, 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validConfig(config Config) bool {
|
||||
if config.Version != 1 || (config.Mode != "local" && config.Mode != "oidc") {
|
||||
return false
|
||||
}
|
||||
return config.Mode != "local" || config.Local.UsersFile == usersFileName
|
||||
}
|
||||
|
||||
func decodeStrictYAML(contents []byte, destination any) error {
|
||||
decoder := yaml.NewDecoder(bytes.NewReader(contents))
|
||||
decoder.KnownFields(true)
|
||||
if err := decoder.Decode(destination); err != nil {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
var extra any
|
||||
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func requirePrivateDirectory(directory string) error {
|
||||
if err := safeio.ValidateCanonicalPath(directory); err != nil {
|
||||
return err
|
||||
}
|
||||
info, err := os.Lstat(directory)
|
||||
if err != nil || !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
|
||||
return safeio.ErrUnsafeFile
|
||||
}
|
||||
resolved, err := filepath.EvalSymlinks(directory)
|
||||
if err != nil || resolved != directory {
|
||||
return safeio.ErrUnsafeFile
|
||||
}
|
||||
if runtime.GOOS != "windows" && info.Mode().Perm() != 0o700 {
|
||||
return safeio.ErrUnsafeFile
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func readPrivateFile(path string) ([]byte, error) {
|
||||
before, err := os.Lstat(path)
|
||||
if err != nil || !before.Mode().IsRegular() || before.Mode()&os.ModeSymlink != 0 || (runtime.GOOS != "windows" && before.Mode().Perm() != 0o600) {
|
||||
return nil, safeio.ErrUnsafeFile
|
||||
}
|
||||
contents, err := safeio.ReadCanonicalRegular(path, maxYAMLBytes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
after, err := os.Lstat(path)
|
||||
if err != nil || !after.Mode().IsRegular() || after.Mode()&os.ModeSymlink != 0 || (runtime.GOOS != "windows" && after.Mode().Perm() != 0o600) || !os.SameFile(before, after) {
|
||||
return nil, safeio.ErrUnsafeFile
|
||||
}
|
||||
return contents, nil
|
||||
}
|
||||
|
||||
func acquireLock(directory string) (*flock.Flock, error) {
|
||||
path := filepath.Join(directory, lockFileName)
|
||||
if _, err := os.Lstat(path); errors.Is(err, os.ErrNotExist) {
|
||||
if err := safeio.WriteCanonicalNewFile(path, nil, 0o600); err != nil {
|
||||
info, statErr := os.Lstat(path)
|
||||
if statErr != nil || !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || (runtime.GOOS != "windows" && info.Mode().Perm() != 0o600) {
|
||||
return nil, safeio.ErrUnsafeFile
|
||||
}
|
||||
}
|
||||
} else if err != nil {
|
||||
return nil, safeio.ErrUnsafeFile
|
||||
}
|
||||
lock := flock.New(path, flock.SetPermissions(0o600))
|
||||
if err := lock.Lock(); err != nil {
|
||||
return nil, errInvalidAuthenticationConfig
|
||||
}
|
||||
if _, err := readPrivateFile(path); err != nil {
|
||||
_ = lock.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
return lock, nil
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
package authconfig
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
|
||||
"github.com/aritmolab/thothii/tools/tht/internal/testsupport"
|
||||
)
|
||||
|
||||
const publicFixedHash = "$argon2id$v=19$m=65536,t=3,p=1$AAECAwQFBgcICQoLDA0ODw$DRo8ZSPI8G5OCvnFFapbVEjP69aDjy1Sw9i2743cPC4"
|
||||
|
||||
func TestLoadReadsTheStrictLocalConfigurationAndRegistry(t *testing.T) {
|
||||
directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(adminUserYAML("admin", "Admin", true, "admin")))
|
||||
|
||||
config, registry, err := Load(directory)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if config.Version != 1 || config.Mode != "local" || config.Local.UsersFile != "users.yaml" {
|
||||
t.Fatalf("Load() config = %#v, want local auth.yaml configuration", config)
|
||||
}
|
||||
if len(registry.Users) != 1 || registry.Users[0].Username != "admin" || registry.Users[0].PasswordHash != publicFixedHash {
|
||||
t.Fatalf("Load() registry = %#v, want one saved administrator", registry)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsUnsafeConfigurationAndRegistryFiles(t *testing.T) {
|
||||
for name, arrange := range map[string]func(t *testing.T, directory string){
|
||||
"symlinked directory": func(t *testing.T, directory string) {
|
||||
real := directory + "-real"
|
||||
if err := os.Rename(directory, real); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
testsupport.SymlinkOrSkip(t, real, directory)
|
||||
},
|
||||
"symlinked auth file": func(t *testing.T, directory string) {
|
||||
if err := os.Remove(filepath.Join(directory, "auth.yaml")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
testsupport.SymlinkOrSkip(t, filepath.Join(directory, "users.yaml"), filepath.Join(directory, "auth.yaml"))
|
||||
},
|
||||
"symlinked users file": func(t *testing.T, directory string) {
|
||||
if err := os.Remove(filepath.Join(directory, "users.yaml")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
testsupport.SymlinkOrSkip(t, filepath.Join(directory, "auth.yaml"), filepath.Join(directory, "users.yaml"))
|
||||
},
|
||||
"hard linked users file": func(t *testing.T, directory string) {
|
||||
linked := filepath.Join(directory, "users-linked.yaml")
|
||||
if err := os.Link(filepath.Join(directory, "users.yaml"), linked); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
},
|
||||
"world readable users file": func(t *testing.T, directory string) {
|
||||
if err := os.Chmod(filepath.Join(directory, "users.yaml"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(adminUserYAML("admin", "Admin", true, "admin")))
|
||||
arrange(t, directory)
|
||||
if _, _, err := Load(directory); err == nil {
|
||||
t.Fatal("Load() accepted an unsafe authentication path")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsDuplicateAndUnknownYAMLFields(t *testing.T) {
|
||||
for name, fixture := range map[string]struct{ auth, users string }{
|
||||
"duplicate auth field": {
|
||||
auth: "version: 1\nmode: local\nmode: oidc\npublicUrl: http://127.0.0.1:8080\nlocal:\n usersFile: users.yaml\n",
|
||||
users: registryYAML(adminUserYAML("admin", "Admin", true, "admin")),
|
||||
},
|
||||
"unknown auth field": {
|
||||
auth: strings.Replace(defaultAuthYAML, "version: 1", "version: 1\nunexpected: true", 1),
|
||||
users: registryYAML(adminUserYAML("admin", "Admin", true, "admin")),
|
||||
},
|
||||
"unknown user field": {
|
||||
auth: defaultAuthYAML,
|
||||
users: strings.Replace(registryYAML(adminUserYAML("admin", "Admin", true, "admin")), " enabled: true", " enabled: true\n unexpected: true", 1),
|
||||
},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
directory := writeAuthFiles(t, fixture.auth, fixture.users)
|
||||
if _, _, err := Load(directory); err == nil {
|
||||
t.Fatal("Load() accepted malformed YAML")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsOversizedYAMLBeforeParsing(t *testing.T) {
|
||||
directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(adminUserYAML("admin", "Admin", true, "admin")))
|
||||
if err := os.WriteFile(filepath.Join(directory, "users.yaml"), []byte(strings.Repeat("#", 1<<20)+"\n"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, _, err := Load(directory); !errors.Is(err, safeio.ErrUnsafeFile) {
|
||||
t.Fatalf("Load() error = %v, want bounded unsafe-file error", err)
|
||||
}
|
||||
}
|
||||
|
||||
const defaultAuthYAML = "version: 1\nmode: local\npublicUrl: http://127.0.0.1:8080\nlocal:\n usersFile: users.yaml\n"
|
||||
|
||||
func writeAuthFiles(t *testing.T, auth, users string) string {
|
||||
t.Helper()
|
||||
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
directory, err := os.MkdirTemp(temporaryRoot, "tht-authconfig-")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.RemoveAll(directory) })
|
||||
if err := os.Chmod(directory, 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for name, contents := range map[string]string{"auth.yaml": auth, "users.yaml": users} {
|
||||
if err := os.WriteFile(filepath.Join(directory, name), []byte(contents), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
return directory
|
||||
}
|
||||
|
||||
func registryYAML(users ...string) string {
|
||||
return "version: 1\nusers:\n" + strings.Join(users, "")
|
||||
}
|
||||
|
||||
func adminUserYAML(username, displayName string, enabled bool, role string) string {
|
||||
return " - id: 6ba7b810-9dad-4ed1-80b4-00c04fd430c8\n" +
|
||||
" username: " + username + "\n" +
|
||||
" displayName: " + displayName + "\n" +
|
||||
" passwordHash: " + publicFixedHash + "\n" +
|
||||
" roles:\n - " + role + "\n" +
|
||||
" enabled: " + map[bool]string{true: "true", false: "false"}[enabled] + "\n" +
|
||||
" authRevision: 1\n"
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
// Package authconfig safely reads and mutates the installation-local authentication files.
|
||||
package authconfig
|
||||
|
||||
import "errors"
|
||||
|
||||
const (
|
||||
maxYAMLBytes int64 = 1 << 20
|
||||
passwordMinBytes = 12
|
||||
passwordMaxBytes = 1024
|
||||
)
|
||||
|
||||
var errInvalidAuthenticationConfig = errors.New("authentication configuration is invalid")
|
||||
|
||||
// Role is one of the stable local authorization roles.
|
||||
type Role string
|
||||
|
||||
const (
|
||||
RoleUser Role = "user"
|
||||
RoleAdmin Role = "admin"
|
||||
)
|
||||
|
||||
// User is the persisted local identity. PasswordHash is deliberately omitted from JSON output.
|
||||
type User struct {
|
||||
ID string `yaml:"id" json:"id"`
|
||||
Username string `yaml:"username" json:"username"`
|
||||
DisplayName string `yaml:"displayName,omitempty" json:"displayName,omitempty"`
|
||||
PasswordHash string `yaml:"passwordHash" json:"-"`
|
||||
Roles []Role `yaml:"roles" json:"roles"`
|
||||
Enabled bool `yaml:"enabled" json:"enabled"`
|
||||
AuthRevision uint64 `yaml:"authRevision" json:"authRevision"`
|
||||
}
|
||||
|
||||
// Config is the strict, non-secret auth.yaml format. The OIDC fields are retained here so callers
|
||||
// can inspect a configuration without accepting arbitrary YAML fields.
|
||||
type Config struct {
|
||||
Version int `yaml:"version" json:"version"`
|
||||
Mode string `yaml:"mode" json:"mode"`
|
||||
PublicURL string `yaml:"publicUrl" json:"publicUrl"`
|
||||
Session SessionConfig `yaml:"session" json:"session"`
|
||||
Local LocalConfig `yaml:"local" json:"local"`
|
||||
OIDC OIDCConfig `yaml:"oidc" json:"oidc"`
|
||||
GroupCatalog GroupCatalogConfig `yaml:"groupCatalog" json:"groupCatalog"`
|
||||
Authorization AuthorizationConfig `yaml:"authorization" json:"authorization"`
|
||||
}
|
||||
|
||||
type SessionConfig struct {
|
||||
RegularTTLSeconds uint64 `yaml:"regularTtlSeconds" json:"regularTtlSeconds"`
|
||||
RegularIdleSeconds uint64 `yaml:"regularIdleSeconds" json:"regularIdleSeconds"`
|
||||
RememberTTLSeconds uint64 `yaml:"rememberTtlSeconds" json:"rememberTtlSeconds"`
|
||||
RememberIdleSeconds uint64 `yaml:"rememberIdleSeconds" json:"rememberIdleSeconds"`
|
||||
OIDCTTLSeconds uint64 `yaml:"oidcTtlSeconds" json:"oidcTtlSeconds"`
|
||||
}
|
||||
|
||||
type LocalConfig struct {
|
||||
UsersFile string `yaml:"usersFile" json:"usersFile"`
|
||||
}
|
||||
|
||||
type OIDCConfig struct {
|
||||
Issuer string `yaml:"issuer" json:"issuer"`
|
||||
ClientID string `yaml:"clientId" json:"clientId"`
|
||||
ClientSecretRef string `yaml:"clientSecretRef" json:"clientSecretRef"`
|
||||
Scopes []string `yaml:"scopes" json:"scopes"`
|
||||
GroupsClaim string `yaml:"groupsClaim" json:"groupsClaim"`
|
||||
}
|
||||
|
||||
type GroupCatalogConfig struct {
|
||||
Driver string `yaml:"driver" json:"driver"`
|
||||
BaseURL string `yaml:"baseUrl" json:"baseUrl"`
|
||||
APITokenRef string `yaml:"apiTokenRef" json:"apiTokenRef"`
|
||||
}
|
||||
|
||||
type AuthorizationConfig struct {
|
||||
GroupRoles map[string][]Role `yaml:"groupRoles" json:"groupRoles"`
|
||||
}
|
||||
|
||||
// Registry is the strict users.yaml format.
|
||||
type Registry struct {
|
||||
Version int `yaml:"version" json:"version"`
|
||||
Users []User `yaml:"users" json:"users"`
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package authconfig
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"regexp"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
var usernamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._@-]{2,63}$`)
|
||||
var uuidV4Pattern = regexp.MustCompile(`^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$`)
|
||||
|
||||
func newUser(username, displayName, passwordHash string, roles []Role) (User, error) {
|
||||
id, err := randomUUIDv4()
|
||||
if err != nil {
|
||||
return User{}, errInvalidAuthenticationConfig
|
||||
}
|
||||
user := User{
|
||||
ID: id,
|
||||
Username: username,
|
||||
DisplayName: displayName,
|
||||
PasswordHash: passwordHash,
|
||||
Roles: append([]Role(nil), roles...),
|
||||
Enabled: true,
|
||||
AuthRevision: 1,
|
||||
}
|
||||
if err := validateUser(user); err != nil {
|
||||
return User{}, errInvalidAuthenticationConfig
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func randomUUIDv4() (string, error) {
|
||||
bytes := make([]byte, 16)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
return "", err
|
||||
}
|
||||
bytes[6] = bytes[6]&0x0f | 0x40
|
||||
bytes[8] = bytes[8]&0x3f | 0x80
|
||||
encoded := hex.EncodeToString(bytes)
|
||||
return encoded[0:8] + "-" + encoded[8:12] + "-" + encoded[12:16] + "-" + encoded[16:20] + "-" + encoded[20:32], nil
|
||||
}
|
||||
|
||||
func validateRegistry(registry Registry) error {
|
||||
if registry.Version != 1 || len(registry.Users) == 0 {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
ids := make(map[string]struct{}, len(registry.Users))
|
||||
usernames := make(map[string]struct{}, len(registry.Users))
|
||||
enabledAdmin := false
|
||||
for _, user := range registry.Users {
|
||||
if err := validateUser(user); err != nil {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
if _, exists := ids[user.ID]; exists {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
ids[user.ID] = struct{}{}
|
||||
normalized := normalizeUsername(user.Username)
|
||||
if _, exists := usernames[normalized]; exists {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
usernames[normalized] = struct{}{}
|
||||
if user.Enabled && hasRole(user.Roles, RoleAdmin) {
|
||||
enabledAdmin = true
|
||||
}
|
||||
}
|
||||
if !enabledAdmin {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUser(user User) error {
|
||||
if !uuidV4Pattern.MatchString(user.ID) || !usernamePattern.MatchString(user.Username) || !utf8.ValidString(user.DisplayName) || containsControl(user.DisplayName) || user.AuthRevision == 0 {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
if err := validatePasswordHash(user.PasswordHash); err != nil {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
if len(user.Roles) == 0 {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
roles := make(map[Role]struct{}, len(user.Roles))
|
||||
for _, role := range user.Roles {
|
||||
if role != RoleUser && role != RoleAdmin {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
if _, exists := roles[role]; exists {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
roles[role] = struct{}{}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func containsControl(value string) bool {
|
||||
for _, character := range value {
|
||||
if unicode.IsControl(character) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func normalizeUsername(username string) string {
|
||||
var normalized strings.Builder
|
||||
normalized.Grow(len(username))
|
||||
for index := 0; index < len(username); index++ {
|
||||
character := username[index]
|
||||
if character >= 'A' && character <= 'Z' {
|
||||
character += 'a' - 'A'
|
||||
}
|
||||
normalized.WriteByte(character)
|
||||
}
|
||||
return normalized.String()
|
||||
}
|
||||
|
||||
func hasRole(roles []Role, wanted Role) bool {
|
||||
for _, role := range roles {
|
||||
if role == wanted {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// FindByUsername uses the configured ASCII case-insensitive lookup while returning the original
|
||||
// display spelling in the persisted record.
|
||||
func (registry *Registry) FindByUsername(username string) *User {
|
||||
normalized := normalizeUsername(username)
|
||||
for index := range registry.Users {
|
||||
if normalizeUsername(registry.Users[index].Username) == normalized {
|
||||
return ®istry.Users[index]
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func cloneRegistry(registry Registry) Registry {
|
||||
clone := Registry{Version: registry.Version, Users: make([]User, len(registry.Users))}
|
||||
for index, user := range registry.Users {
|
||||
clone.Users[index] = user
|
||||
clone.Users[index].Roles = append([]Role(nil), user.Roles...)
|
||||
}
|
||||
return clone
|
||||
}
|
||||
|
||||
func applyMutationInvariants(before Registry, after *Registry) error {
|
||||
previous := make(map[string]User, len(before.Users))
|
||||
for _, user := range before.Users {
|
||||
previous[user.ID] = user
|
||||
}
|
||||
for index := range after.Users {
|
||||
user := &after.Users[index]
|
||||
old, exists := previous[user.ID]
|
||||
if !exists {
|
||||
if user.AuthRevision != 1 {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
continue
|
||||
}
|
||||
if user.Username != old.Username || user.AuthRevision < old.AuthRevision {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
securityChanged := user.PasswordHash != old.PasswordHash || user.Enabled != old.Enabled || !sameRoles(user.Roles, old.Roles)
|
||||
explicitLogoutAll := user.AuthRevision == old.AuthRevision+1
|
||||
if user.AuthRevision != old.AuthRevision && !explicitLogoutAll {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
if securityChanged && !explicitLogoutAll {
|
||||
if old.AuthRevision == ^uint64(0) {
|
||||
return errInvalidAuthenticationConfig
|
||||
}
|
||||
user.AuthRevision = old.AuthRevision + 1
|
||||
}
|
||||
if !securityChanged && !explicitLogoutAll {
|
||||
user.AuthRevision = old.AuthRevision
|
||||
}
|
||||
}
|
||||
return validateRegistry(*after)
|
||||
}
|
||||
|
||||
func sameRoles(left, right []Role) bool {
|
||||
if len(left) != len(right) {
|
||||
return false
|
||||
}
|
||||
for _, role := range left {
|
||||
if !hasRole(right, role) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
package authconfig
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestMutateUsersRejectsRemovingTheLastEnabledAdministrator(t *testing.T) {
|
||||
directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(adminUserYAML("admin", "Admin", true, "admin")))
|
||||
|
||||
err := MutateUsers(directory, func(registry *Registry) error {
|
||||
registry.Users[0].Roles = []Role{RoleUser}
|
||||
return nil
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("MutateUsers() allowed removal of the last enabled administrator")
|
||||
}
|
||||
_, registry, err := Load(directory)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(registry.Users[0].Roles) != 1 || registry.Users[0].Roles[0] != RoleAdmin {
|
||||
t.Fatal("MutateUsers() wrote an invalid last-admin mutation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMutateUsersSerializesConcurrentReadCheckWriteTransactions(t *testing.T) {
|
||||
users := adminUserYAML("admin", "Admin", true, "admin")
|
||||
for index := 0; index < 12; index++ {
|
||||
users += userYAML(index)
|
||||
}
|
||||
directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(users))
|
||||
|
||||
var group sync.WaitGroup
|
||||
errors := make(chan error, 12)
|
||||
for index := 0; index < 12; index++ {
|
||||
index := index
|
||||
group.Add(1)
|
||||
go func() {
|
||||
defer group.Done()
|
||||
errors <- MutateUsers(directory, func(registry *Registry) error {
|
||||
registry.Users[index+1].DisplayName = fmt.Sprintf("Updated %d", index)
|
||||
return nil
|
||||
})
|
||||
}()
|
||||
}
|
||||
group.Wait()
|
||||
close(errors)
|
||||
for err := range errors {
|
||||
if err != nil {
|
||||
t.Fatalf("MutateUsers() concurrent mutation error = %v", err)
|
||||
}
|
||||
}
|
||||
_, registry, err := Load(directory)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for index := 0; index < 12; index++ {
|
||||
if got, want := registry.Users[index+1].DisplayName, fmt.Sprintf("Updated %d", index); got != want {
|
||||
t.Fatalf("user %d displayName = %q, want %q; mutation was lost", index, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMutateUsersIncrementsAuthRevisionForSecurityChanges(t *testing.T) {
|
||||
directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(
|
||||
adminUserYAML("admin", "Admin", true, "admin"),
|
||||
userYAML(0),
|
||||
))
|
||||
|
||||
if err := MutateUsers(directory, func(registry *Registry) error {
|
||||
registry.Users[1].Enabled = false
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, registry, err := Load(directory)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got, want := registry.Users[1].AuthRevision, uint64(2); got != want {
|
||||
t.Fatalf("security mutation authRevision = %d, want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegistryRejectsAmbiguousOrUnsafeUserRecords(t *testing.T) {
|
||||
secondAdmin := strings.Replace(adminUserYAML("Admin", "Admin", true, "admin"), "6ba7b810-9dad-4ed1-80b4-00c04fd430c8", "7ba7b810-9dad-4ed1-80b4-00c04fd430c8", 1)
|
||||
for name, users := range map[string]string{
|
||||
"duplicate ASCII case-insensitive username": secondAdmin,
|
||||
"invalid username": strings.Replace(userYAML(0), "username: user0", "username: _user", 1),
|
||||
"control display name": strings.Replace(userYAML(0), "displayName: User 0", "displayName: \"User\\t0\"", 1),
|
||||
"unknown role": strings.Replace(userYAML(0), "- user", "- operator", 1),
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
directory := writeAuthFiles(t, defaultAuthYAML, registryYAML(adminUserYAML("admin", "Admin", true, "admin")+users))
|
||||
if _, _, err := Load(directory); err == nil {
|
||||
t.Fatal("Load() accepted an ambiguous or unsafe user record")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewUserGeneratesAUUIDv4(t *testing.T) {
|
||||
user, err := newUser("operator", "Operator", publicFixedHash, []Role{RoleUser})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(user.ID) != 36 || user.ID[14] != '4' || user.ID[19] != '8' && user.ID[19] != '9' && user.ID[19] != 'a' && user.ID[19] != 'b' {
|
||||
t.Fatalf("newUser() ID = %q, want UUIDv4", user.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func userYAML(index int) string {
|
||||
return fmt.Sprintf(" - id: 8a0a0b0c-0d0e-4f1a-8b2c-%012d\n username: user%d\n displayName: User %d\n passwordHash: %s\n roles:\n - user\n enabled: true\n authRevision: 1\n", index+1, index, index, publicFixedHash)
|
||||
}
|
||||
@@ -2,6 +2,8 @@
|
||||
package safeio
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
@@ -36,8 +38,8 @@ func readBoundedRegularFile(path string, file *os.File, maximum int64) ([]byte,
|
||||
if err != nil || !after.Mode().IsRegular() || !hasSingleLink(after) || !os.SameFile(info, after) {
|
||||
return nil, ErrUnsafeFile
|
||||
}
|
||||
current, err := os.Stat(path)
|
||||
if err != nil || !os.SameFile(info, current) {
|
||||
current, err := os.Lstat(path)
|
||||
if err != nil || !current.Mode().IsRegular() || current.Mode()&os.ModeSymlink != 0 || !hasSingleLink(current) || !os.SameFile(info, current) {
|
||||
return nil, ErrUnsafeFile
|
||||
}
|
||||
return contents, nil
|
||||
@@ -86,6 +88,24 @@ func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error
|
||||
return nil
|
||||
}
|
||||
|
||||
// ReplaceCanonicalRegular durably replaces one existing private regular file without following
|
||||
// symlinked path components. Platform implementations keep the temporary file in the target
|
||||
// directory and use the platform's atomic replace primitive.
|
||||
func ReplaceCanonicalRegular(path string, contents []byte, mode os.FileMode) error {
|
||||
if err := ValidateCanonicalPath(path); err != nil || mode.Perm() != 0o600 || mode&os.ModeType != 0 {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
return replaceCanonicalRegular(path, contents)
|
||||
}
|
||||
|
||||
func randomTemporaryName() (string, error) {
|
||||
bytes := make([]byte, 16)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return ".tht-auth-" + hex.EncodeToString(bytes) + ".tmp", nil
|
||||
}
|
||||
|
||||
func requireCanonicalDirectory(path string) error {
|
||||
if err := ValidateCanonicalPath(path); err != nil {
|
||||
return err
|
||||
|
||||
@@ -30,3 +30,80 @@ func TestReadCanonicalRegularRejectsNamedPipeWithoutBlocking(t *testing.T) {
|
||||
t.Fatalf("named pipe error = %v, want ErrUnsafeFile", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplaceCanonicalRegularReplacesOnlySafeRegularTargets(t *testing.T) {
|
||||
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-replace-")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.RemoveAll(root) })
|
||||
path := filepath.Join(root, "users.yaml")
|
||||
if err := os.WriteFile(path, []byte("old"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err := ReplaceCanonicalRegular(path, []byte("new"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
contents, err := ReadCanonicalRegular(path, 1024)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(contents) != "new" {
|
||||
t.Fatalf("replacement content = %q, want new", contents)
|
||||
}
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if info.Mode().Perm() != 0o600 {
|
||||
t.Fatalf("replacement mode = %o, want 600", info.Mode().Perm())
|
||||
}
|
||||
|
||||
linked := filepath.Join(root, "linked.yaml")
|
||||
if err := os.Link(path, linked); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ReplaceCanonicalRegular(path, []byte("unsafe"), 0o600); !errors.Is(err, ErrUnsafeFile) {
|
||||
t.Fatalf("hard-linked replacement error = %v, want ErrUnsafeFile", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplaceCanonicalRegularRejectsSymlinkedPathComponents(t *testing.T) {
|
||||
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-replace-")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.RemoveAll(root) })
|
||||
realDirectory := filepath.Join(root, "real")
|
||||
if err := os.Mkdir(realDirectory, 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
target := filepath.Join(realDirectory, "users.yaml")
|
||||
if err := os.WriteFile(target, []byte("old"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
linkedDirectory := filepath.Join(root, "linked")
|
||||
if err := os.Symlink(realDirectory, linkedDirectory); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ReplaceCanonicalRegular(filepath.Join(linkedDirectory, "users.yaml"), []byte("new"), 0o600); !errors.Is(err, ErrUnsafeFile) {
|
||||
t.Fatalf("parent symlink replacement error = %v, want ErrUnsafeFile", err)
|
||||
}
|
||||
|
||||
linkedFile := filepath.Join(root, "linked-file.yaml")
|
||||
if err := os.Symlink(target, linkedFile); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ReplaceCanonicalRegular(linkedFile, []byte("new"), 0o600); !errors.Is(err, ErrUnsafeFile) {
|
||||
t.Fatalf("final symlink replacement error = %v, want ErrUnsafeFile", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -82,7 +82,7 @@ func openWindowsComponent(path string, directory bool) (windows.Handle, error) {
|
||||
}
|
||||
if information.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 ||
|
||||
(directory && information.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY == 0) ||
|
||||
(!directory && information.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY != 0) {
|
||||
(!directory && (information.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY != 0 || information.NumberOfLinks != 1)) {
|
||||
windows.CloseHandle(handle)
|
||||
return 0, ErrUnsafeFile
|
||||
}
|
||||
|
||||
@@ -12,10 +12,14 @@ import (
|
||||
|
||||
const expectedWindowsRetainedHandleShareMode = windows.FILE_SHARE_READ | windows.FILE_SHARE_WRITE
|
||||
|
||||
const expectedWindowsReplaceMoveFlags = windows.MOVEFILE_REPLACE_EXISTING | windows.MOVEFILE_WRITE_THROUGH
|
||||
|
||||
// Keep this contract compile-enforced so Windows cross-test compilation catches a future
|
||||
// FILE_SHARE_DELETE regression even when the tests are compiled on a non-Windows host.
|
||||
var _ [windowsRetainedHandleShareMode - expectedWindowsRetainedHandleShareMode]struct{}
|
||||
var _ [expectedWindowsRetainedHandleShareMode - windowsRetainedHandleShareMode]struct{}
|
||||
var _ [windowsReplaceMoveFlags - expectedWindowsReplaceMoveFlags]struct{}
|
||||
var _ [expectedWindowsReplaceMoveFlags - windowsReplaceMoveFlags]struct{}
|
||||
|
||||
func TestOpenWindowsComponentBlocksMutationWhileHandleIsRetained(t *testing.T) {
|
||||
t.Run("parent rename", func(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
//go:build !windows
|
||||
|
||||
package safeio
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
func replaceCanonicalRegular(path string, contents []byte) error {
|
||||
directory, target, err := openCanonicalParentDirectory(path)
|
||||
if err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
defer unix.Close(directory)
|
||||
if err := requireSingleRegularAt(directory, target); err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
|
||||
temporary, err := writePrivateTemporaryAt(directory, contents)
|
||||
if err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
defer func() {
|
||||
if temporary != "" {
|
||||
_ = unix.Unlinkat(directory, temporary, 0)
|
||||
}
|
||||
}()
|
||||
if err := requireSingleRegularAt(directory, target); err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
if err := unix.Renameat(directory, temporary, directory, target); err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
temporary = ""
|
||||
if err := unix.Fsync(directory); err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func openCanonicalParentDirectory(path string) (int, string, error) {
|
||||
components := strings.Split(strings.TrimPrefix(path, string(os.PathSeparator)), string(os.PathSeparator))
|
||||
if len(components) < 2 || components[0] == "" || components[len(components)-1] == "" {
|
||||
return -1, "", ErrUnsafeFile
|
||||
}
|
||||
directory, err := unix.Open(string(os.PathSeparator), unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY, 0)
|
||||
if err != nil {
|
||||
return -1, "", err
|
||||
}
|
||||
for _, component := range components[:len(components)-1] {
|
||||
next, err := unix.Openat(directory, component, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0)
|
||||
if err != nil {
|
||||
unix.Close(directory)
|
||||
return -1, "", err
|
||||
}
|
||||
unix.Close(directory)
|
||||
directory = next
|
||||
}
|
||||
return directory, components[len(components)-1], nil
|
||||
}
|
||||
|
||||
func requireSingleRegularAt(directory int, name string) error {
|
||||
var stat unix.Stat_t
|
||||
if err := unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW); err != nil || stat.Mode&unix.S_IFMT != unix.S_IFREG || stat.Nlink != 1 {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func writePrivateTemporaryAt(directory int, contents []byte) (string, error) {
|
||||
for attempt := 0; attempt < 16; attempt++ {
|
||||
name, err := randomTemporaryName()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
descriptor, err := unix.Openat(directory, name, unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0o600)
|
||||
if errors.Is(err, unix.EEXIST) {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
file := os.NewFile(uintptr(descriptor), "tht-safeio-replacement")
|
||||
if file == nil {
|
||||
unix.Close(descriptor)
|
||||
return "", ErrUnsafeFile
|
||||
}
|
||||
if err := file.Chmod(0o600); err == nil {
|
||||
var written int
|
||||
written, err = file.Write(contents)
|
||||
if err == nil && written != len(contents) {
|
||||
err = io.ErrShortWrite
|
||||
}
|
||||
}
|
||||
if err == nil {
|
||||
err = file.Sync()
|
||||
}
|
||||
closeErr := file.Close()
|
||||
if err == nil {
|
||||
err = closeErr
|
||||
}
|
||||
if err != nil {
|
||||
_ = unix.Unlinkat(directory, name, 0)
|
||||
return "", err
|
||||
}
|
||||
return name, nil
|
||||
}
|
||||
return "", ErrUnsafeFile
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
//go:build windows
|
||||
|
||||
package safeio
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
const windowsReplaceMoveFlags = windows.MOVEFILE_REPLACE_EXISTING | windows.MOVEFILE_WRITE_THROUGH
|
||||
|
||||
func replaceCanonicalRegular(path string, contents []byte) error {
|
||||
directory := filepath.Dir(path)
|
||||
if err := requireCanonicalDirectory(directory); err != nil || !safeExistingRegular(path) {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
temporary, err := writePrivateTemporary(directory, contents)
|
||||
if err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
defer func() { _ = os.Remove(temporary) }()
|
||||
if !safeExistingRegular(path) {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
from, err := windows.UTF16PtrFromString(temporary)
|
||||
if err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
to, err := windows.UTF16PtrFromString(path)
|
||||
if err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
if err := windows.MoveFileEx(from, to, windowsReplaceMoveFlags); err != nil {
|
||||
return ErrUnsafeFile
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func safeExistingRegular(path string) bool {
|
||||
handle, err := openWindowsComponent(path, false)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return windows.CloseHandle(handle) == nil
|
||||
}
|
||||
|
||||
func writePrivateTemporary(directory string, contents []byte) (string, error) {
|
||||
for attempt := 0; attempt < 16; attempt++ {
|
||||
name, err := randomTemporaryName()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
path := filepath.Join(directory, name)
|
||||
file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o600)
|
||||
if errors.Is(err, os.ErrExist) {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := file.Chmod(0o600); err == nil {
|
||||
var written int
|
||||
written, err = file.Write(contents)
|
||||
if err == nil && written != len(contents) {
|
||||
err = io.ErrShortWrite
|
||||
}
|
||||
}
|
||||
if err == nil {
|
||||
err = file.Sync()
|
||||
}
|
||||
closeErr := file.Close()
|
||||
if err == nil {
|
||||
err = closeErr
|
||||
}
|
||||
if err != nil {
|
||||
_ = os.Remove(path)
|
||||
return "", err
|
||||
}
|
||||
return path, nil
|
||||
}
|
||||
return "", ErrUnsafeFile
|
||||
}
|
||||
Reference in New Issue
Block a user