From f6e4dbcae247bbe6f62b91b6b7a11a5ebb3894f3 Mon Sep 17 00:00:00 2001 From: mptyl Date: Sun, 16 Aug 2026 18:24:20 +0200 Subject: [PATCH] feat(auth): add safe Argon2id local user registry --- backend/test/fixtures/argon2id-vectors.json | 1 + tools/tht/go.mod | 1 + tools/tht/go.sum | 2 + tools/tht/internal/authconfig/password.go | 152 ++++++++++++++ .../tht/internal/authconfig/password_test.go | 103 +++++++++ tools/tht/internal/authconfig/store.go | 169 +++++++++++++++ tools/tht/internal/authconfig/store_test.go | 144 +++++++++++++ tools/tht/internal/authconfig/types.go | 80 +++++++ tools/tht/internal/authconfig/users.go | 196 ++++++++++++++++++ tools/tht/internal/authconfig/users_test.go | 117 +++++++++++ tools/tht/internal/safeio/files.go | 24 ++- tools/tht/internal/safeio/files_unix_test.go | 77 +++++++ tools/tht/internal/safeio/files_windows.go | 2 +- .../tht/internal/safeio/files_windows_test.go | 4 + tools/tht/internal/safeio/replace_unix.go | 114 ++++++++++ tools/tht/internal/safeio/replace_windows.go | 86 ++++++++ 16 files changed, 1269 insertions(+), 3 deletions(-) create mode 100644 backend/test/fixtures/argon2id-vectors.json create mode 100644 tools/tht/internal/authconfig/password.go create mode 100644 tools/tht/internal/authconfig/password_test.go create mode 100644 tools/tht/internal/authconfig/store.go create mode 100644 tools/tht/internal/authconfig/store_test.go create mode 100644 tools/tht/internal/authconfig/types.go create mode 100644 tools/tht/internal/authconfig/users.go create mode 100644 tools/tht/internal/authconfig/users_test.go create mode 100644 tools/tht/internal/safeio/replace_unix.go create mode 100644 tools/tht/internal/safeio/replace_windows.go diff --git a/backend/test/fixtures/argon2id-vectors.json b/backend/test/fixtures/argon2id-vectors.json new file mode 100644 index 00000000..c04a149a --- /dev/null +++ b/backend/test/fixtures/argon2id-vectors.json @@ -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"}] diff --git a/tools/tht/go.mod b/tools/tht/go.mod index 647f13db..36abe03a 100644 --- a/tools/tht/go.mod +++ b/tools/tht/go.mod @@ -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 ) diff --git a/tools/tht/go.sum b/tools/tht/go.sum index 1b6db1ef..73b086ab 100644 --- a/tools/tht/go.sum +++ b/tools/tht/go.sum @@ -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= diff --git a/tools/tht/internal/authconfig/password.go b/tools/tht/internal/authconfig/password.go new file mode 100644 index 00000000..ad6eb59f --- /dev/null +++ b/tools/tht/internal/authconfig/password.go @@ -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 +} diff --git a/tools/tht/internal/authconfig/password_test.go b/tools/tht/internal/authconfig/password_test.go new file mode 100644 index 00000000..bd02cad2 --- /dev/null +++ b/tools/tht/internal/authconfig/password_test.go @@ -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] +} diff --git a/tools/tht/internal/authconfig/store.go b/tools/tht/internal/authconfig/store.go new file mode 100644 index 00000000..7264c2b8 --- /dev/null +++ b/tools/tht/internal/authconfig/store.go @@ -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 +} diff --git a/tools/tht/internal/authconfig/store_test.go b/tools/tht/internal/authconfig/store_test.go new file mode 100644 index 00000000..5b32c929 --- /dev/null +++ b/tools/tht/internal/authconfig/store_test.go @@ -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" +} diff --git a/tools/tht/internal/authconfig/types.go b/tools/tht/internal/authconfig/types.go new file mode 100644 index 00000000..11c3ceb3 --- /dev/null +++ b/tools/tht/internal/authconfig/types.go @@ -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"` +} diff --git a/tools/tht/internal/authconfig/users.go b/tools/tht/internal/authconfig/users.go new file mode 100644 index 00000000..54b44895 --- /dev/null +++ b/tools/tht/internal/authconfig/users.go @@ -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 +} diff --git a/tools/tht/internal/authconfig/users_test.go b/tools/tht/internal/authconfig/users_test.go new file mode 100644 index 00000000..cd28cc19 --- /dev/null +++ b/tools/tht/internal/authconfig/users_test.go @@ -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) +} diff --git a/tools/tht/internal/safeio/files.go b/tools/tht/internal/safeio/files.go index 86553a16..8a0a0fa1 100644 --- a/tools/tht/internal/safeio/files.go +++ b/tools/tht/internal/safeio/files.go @@ -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 diff --git a/tools/tht/internal/safeio/files_unix_test.go b/tools/tht/internal/safeio/files_unix_test.go index a8911396..d0a966f6 100644 --- a/tools/tht/internal/safeio/files_unix_test.go +++ b/tools/tht/internal/safeio/files_unix_test.go @@ -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) + } +} diff --git a/tools/tht/internal/safeio/files_windows.go b/tools/tht/internal/safeio/files_windows.go index abb8952a..020a86c7 100644 --- a/tools/tht/internal/safeio/files_windows.go +++ b/tools/tht/internal/safeio/files_windows.go @@ -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 } diff --git a/tools/tht/internal/safeio/files_windows_test.go b/tools/tht/internal/safeio/files_windows_test.go index 745ea672..8f694595 100644 --- a/tools/tht/internal/safeio/files_windows_test.go +++ b/tools/tht/internal/safeio/files_windows_test.go @@ -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) { diff --git a/tools/tht/internal/safeio/replace_unix.go b/tools/tht/internal/safeio/replace_unix.go new file mode 100644 index 00000000..83888857 --- /dev/null +++ b/tools/tht/internal/safeio/replace_unix.go @@ -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 +} diff --git a/tools/tht/internal/safeio/replace_windows.go b/tools/tht/internal/safeio/replace_windows.go new file mode 100644 index 00000000..257c8060 --- /dev/null +++ b/tools/tht/internal/safeio/replace_windows.go @@ -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 +}