fix(auth): pin native auth storage operations

This commit is contained in:
2026-08-17 14:52:21 +02:00
parent 6d438f4c7e
commit cb0e873ed7
23 changed files with 3016 additions and 1385 deletions
+209 -144
View File
@@ -9,8 +9,6 @@ import (
"encoding/json"
"errors"
"io"
"os"
"path/filepath"
"regexp"
"strings"
"unicode"
@@ -123,21 +121,21 @@ func execute(input request) (response, error) {
return response{Version: protocolVersion, OK: true, Prepared: true}, nil
}
if input.Operation == "read-auth-config" {
root, err := existingPrivateRoot(input.Root)
if err != nil {
return response{}, errInvalid
}
contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(root, input.Filename), maximumAuthConfigBytes)
if err != nil {
return response{}, errInvalid
}
return contentResponse(true, contents), nil
return readRootPrivateRegular(input.Root, input.Filename)
}
if input.Operation == "read-local-users" {
return readRootPrivateRegular(input.Root, input.Filename)
}
if !validDirectory(input.Directory) {
return response{}, errInvalid
}
directory, err := storageDirectory(input.Root, input.Directory)
if err != nil {
layout, err := openStorageLayout(input.Root, true)
if err != nil || layout == nil {
return response{}, errInvalid
}
defer layout.Close()
directory := layout.directory(input.Directory)
if directory == nil {
return response{}, errInvalid
}
switch input.Operation {
@@ -162,7 +160,7 @@ func execute(input request) (response, error) {
if err != nil {
return response{}, errInvalid
}
if err := safeio.ReplaceCanonicalRegular(filepath.Join(directory, input.Filename), contents, 0o600); err != nil {
if err := directory.ReplaceRegular(input.Filename, contents); err != nil {
return response{}, errInvalid
}
return response{Version: protocolVersion, OK: true, Replaced: true}, nil
@@ -177,46 +175,32 @@ func execute(input request) (response, error) {
if limit == 0 {
limit = defaultMaximumEntries
}
afterName := ""
if input.Continuation {
page, err := safeio.ListCanonicalPrivateDirectoryPage(
directory,
limit,
input.AfterName,
digestFilename.MatchString,
)
if err != nil {
return response{}, errInvalid
}
more := page.More
return response{Version: protocolVersion, OK: true, Entries: &page.Entries, More: &more}, nil
afterName = input.AfterName
}
entries, err := safeio.ListCanonicalPrivateDirectory(directory, limit)
page, err := directory.ListPage(limit, afterName, recordListName(input.Directory), recordListLinks(input.Directory))
if err != nil {
return response{}, errInvalid
}
for _, entry := range entries {
if !digestFilename.MatchString(entry.Name) && !(input.Directory == "oidc" && (claimFilename.MatchString(entry.Name) || oidcSlotFilename.MatchString(entry.Name))) {
return response{}, errInvalid
}
if input.Continuation {
more := page.More
return response{Version: protocolVersion, OK: true, Entries: &page.Entries, More: &more}, nil
}
return response{Version: protocolVersion, OK: true, Entries: &entries}, nil
if page.More {
return response{}, errInvalid
}
return response{Version: protocolVersion, OK: true, Entries: &page.Entries}, nil
case "claim-consume":
return claimConsume(directory, input.Filename)
case "read-claim":
contents, found, err := safeio.ReadCanonicalPrivateClaim(
filepath.Join(directory, input.Filename),
filepath.Join(directory, asClaimFilename(input.Filename)),
maximumOIDCStateBytes,
)
contents, found, err := directory.ReadClaim(input.Filename, asClaimFilename(input.Filename), maximumOIDCStateBytes)
if err != nil {
return response{}, errInvalid
}
return contentResponse(found, contents), nil
case "remove-claim":
removed, err := safeio.RemoveCanonicalPrivateClaim(
filepath.Join(directory, input.Filename),
filepath.Join(directory, asClaimFilename(input.Filename)),
)
removed, err := directory.RemoveClaim(input.Filename, asClaimFilename(input.Filename))
if err != nil {
return response{}, errInvalid
}
@@ -234,7 +218,7 @@ func validOperationShape(input request) bool {
switch input.Operation {
case "validate-root", "ensure-layout":
return input.Directory == "" && input.Filename == "" && noContents && noMaximumEntries && noAfterName && noContinuation
case "read-auth-config":
case "read-auth-config", "read-local-users":
return input.Directory == "" && authConfigFilename.MatchString(input.Filename) && noContents && noMaximumEntries && noAfterName && noContinuation
case "create", "replace":
return noMaximumEntries && noAfterName && noContinuation && (digestFilename.MatchString(input.Filename) || (input.Operation == "create" && input.Directory == "oidc" && oidcSlotFilename.MatchString(input.Filename)))
@@ -252,78 +236,170 @@ func validOperationShape(input request) bool {
}
}
func preflightRoot(root string) (bool, error) {
if !filepath.IsAbs(root) || filepath.Clean(root) != root || strings.IndexFunc(root, unicode.IsControl) >= 0 {
return false, errInvalid
}
return safeio.PreflightPrivateDirectory(root)
type storageLayout struct {
root safeio.PrivateDirectoryHandle
sessions safeio.PrivateDirectoryHandle
oidc safeio.PrivateDirectoryHandle
}
func existingPrivateRoot(root string) (string, error) {
exists, err := preflightRoot(root)
if err != nil || !exists || safeio.ValidatePrivateDirectory(root) != nil {
return "", errInvalid
func (layout *storageLayout) Close() {
if layout == nil {
return
}
return root, nil
if layout.oidc != nil {
_ = layout.oidc.Close()
layout.oidc = nil
}
if layout.sessions != nil {
_ = layout.sessions.Close()
layout.sessions = nil
}
if layout.root != nil {
_ = layout.root.Close()
layout.root = nil
}
}
func (layout *storageLayout) closeChildren() {
if layout == nil {
return
}
if layout.oidc != nil {
_ = layout.oidc.Close()
layout.oidc = nil
}
if layout.sessions != nil {
_ = layout.sessions.Close()
layout.sessions = nil
}
}
func (layout *storageLayout) directory(name string) safeio.PrivateDirectoryHandle {
if layout == nil {
return nil
}
if name == "sessions" {
return layout.sessions
}
if name == "oidc" {
return layout.oidc
}
return nil
}
func validStorageRoot(root string) bool {
return strings.IndexFunc(root, unicode.IsControl) < 0
}
// openStorageLayout keeps the root descriptor/handle open from its initial canonical validation
// through every child observation. The side-effect-free path performs a second child pass after
// the testable race boundary and retains those final handles; it never validates an earlier,
// discarded exists result.
func openStorageLayout(root string, ensure bool) (*storageLayout, error) {
if !validStorageRoot(root) {
return nil, errInvalid
}
rootHandle, found, err := safeio.OpenPrivateDirectory(root, ensure)
if err != nil {
return nil, errInvalid
}
if !found {
return nil, nil
}
layout := &storageLayout{root: rootHandle}
failed := true
defer func() {
if failed {
layout.Close()
}
}()
safeio.NotifyPrivateDirectoryTestHookForTest("after-auth-root-open")
if err := layout.openChildren(ensure); err != nil {
return nil, errInvalid
}
safeio.NotifyPrivateDirectoryTestHookForTest("after-auth-layout-first-pass")
if !ensure {
layout.closeChildren()
if err := layout.openChildren(false); err != nil {
return nil, errInvalid
}
}
if layout.root.Validate() != nil || (layout.sessions != nil && layout.sessions.Validate() != nil) ||
(layout.oidc != nil && layout.oidc.Validate() != nil) {
return nil, errInvalid
}
failed = false
return layout, nil
}
func (layout *storageLayout) openChildren(ensure bool) error {
if layout == nil || layout.root == nil || layout.root.Validate() != nil {
return errInvalid
}
for _, name := range []string{"sessions", "oidc"} {
child, found, err := layout.root.OpenChild(name, ensure)
if err != nil || (ensure && !found) {
if child != nil {
_ = child.Close()
}
return errInvalid
}
if !found {
continue
}
if child.Validate() != nil {
_ = child.Close()
return errInvalid
}
if name == "sessions" {
layout.sessions = child
} else {
layout.oidc = child
}
}
if layout.root.Validate() != nil {
return errInvalid
}
return nil
}
func validateStorageLayout(root string) error {
exists, err := preflightRoot(root)
layout, err := openStorageLayout(root, false)
if err != nil {
return errInvalid
}
if !exists {
return nil
}
existingChildren := make([]string, 0, 2)
for _, directory := range []string{"sessions", "oidc"} {
path := filepath.Join(root, directory)
if filepath.Dir(path) != root {
return errInvalid
}
childExists, err := safeio.PreflightPrivateDirectory(path)
if err != nil {
return errInvalid
}
if childExists {
existingChildren = append(existingChildren, path)
}
}
// Close permission/identity races between the individual side-effect-free preflights.
if safeio.ValidatePrivateDirectory(root) != nil {
return errInvalid
}
for _, directory := range []string{"sessions", "oidc"} {
if _, err := safeio.PreflightPrivateDirectory(filepath.Join(root, directory)); err != nil {
return errInvalid
}
}
for _, path := range existingChildren {
if safeio.ValidatePrivateDirectory(path) != nil {
return errInvalid
}
if layout != nil {
layout.Close()
}
return nil
}
func ensureStorageLayout(root string) error {
if _, err := preflightRoot(root); err != nil || safeio.EnsurePrivateDirectory(root) != nil {
return errInvalid
}
for _, directory := range []string{"sessions", "oidc"} {
path := filepath.Join(root, directory)
if filepath.Dir(path) != root || safeio.EnsurePrivateDirectory(path) != nil {
return errInvalid
}
}
if safeio.ValidatePrivateDirectory(root) != nil ||
safeio.ValidatePrivateDirectory(filepath.Join(root, "sessions")) != nil ||
safeio.ValidatePrivateDirectory(filepath.Join(root, "oidc")) != nil {
layout, err := openStorageLayout(root, true)
if err != nil || layout == nil {
return errInvalid
}
layout.Close()
return nil
}
func readRootPrivateRegular(root, filename string) (response, error) {
if !validStorageRoot(root) {
return response{}, errInvalid
}
directory, found, err := safeio.OpenPrivateDirectory(root, false)
if err != nil || !found || directory == nil {
return response{}, errInvalid
}
defer directory.Close()
safeio.NotifyPrivateDirectoryTestHookForTest("after-auth-root-open")
contents, found, err := directory.ReadRegular(filename, maximumAuthConfigBytes)
if err != nil || !found || directory.Validate() != nil {
return response{}, errInvalid
}
return contentResponse(true, contents), nil
}
func contentResponse(found bool, contents []byte) response {
if !found {
return response{Version: protocolVersion, OK: true}
@@ -331,17 +407,6 @@ func contentResponse(found bool, contents []byte) response {
return response{Version: protocolVersion, OK: true, Found: true, ContentBase64: base64.StdEncoding.EncodeToString(contents)}
}
func storageDirectory(root, directory string) (string, error) {
if _, err := preflightRoot(root); err != nil || safeio.EnsurePrivateDirectory(root) != nil {
return "", errInvalid
}
path := filepath.Join(root, directory)
if filepath.Dir(path) != root || safeio.EnsurePrivateDirectory(path) != nil {
return "", errInvalid
}
return path, nil
}
func validDirectory(value string) bool {
return value == "sessions" || value == "oidc"
}
@@ -361,67 +426,67 @@ func decodeContents(input request) ([]byte, error) {
return contents, nil
}
func createPrivate(directory, filename string, contents []byte) (bool, error) {
path := filepath.Join(directory, filename)
if _, err := os.Lstat(path); err == nil {
if safeio.ValidatePrivateRegular(path) != nil {
return false, errInvalid
}
return false, nil
} else if !errors.Is(err, os.ErrNotExist) {
func createPrivate(directory safeio.PrivateDirectoryHandle, filename string, contents []byte) (bool, error) {
created, err := directory.CreateRegular(filename, contents)
if err != nil {
return false, errInvalid
}
if err := safeio.WriteCanonicalNewPrivateFile(path, contents, 0o600); err == nil {
return true, nil
}
if safeio.ValidatePrivateRegular(path) == nil {
return false, nil
}
return false, errInvalid
return created, nil
}
func readPrivate(directory, filename string, maximum int64) ([]byte, bool, error) {
path := filepath.Join(directory, filename)
if _, err := os.Lstat(path); errors.Is(err, os.ErrNotExist) {
return nil, false, nil
} else if err != nil {
return nil, false, errInvalid
}
contents, err := safeio.ReadCanonicalPrivateRegular(path, maximum)
func readPrivate(directory safeio.PrivateDirectoryHandle, filename string, maximum int64) ([]byte, bool, error) {
contents, found, err := directory.ReadRegular(filename, maximum)
if err != nil {
return nil, false, errInvalid
}
return contents, true, nil
return contents, found, nil
}
func removePrivate(directory, filename string) (bool, error) {
path := filepath.Join(directory, filename)
if _, err := os.Lstat(path); errors.Is(err, os.ErrNotExist) {
return false, nil
} else if err != nil {
func removePrivate(directory safeio.PrivateDirectoryHandle, filename string) (bool, error) {
removed, err := directory.RemoveRegular(filename)
if err != nil {
return false, errInvalid
}
if err := safeio.RemoveCanonicalPrivateRegular(path); err != nil {
return false, errInvalid
}
return true, nil
return removed, nil
}
func claimConsume(directory, filename string) (response, error) {
source := filepath.Join(directory, filename)
claim := filepath.Join(directory, asClaimFilename(filename))
claimed, err := safeio.ClaimCanonicalPrivateRegular(source, claim)
func recordListName(directory string) func(string) bool {
if directory == "sessions" {
return digestFilename.MatchString
}
return func(name string) bool {
return digestFilename.MatchString(name) || claimFilename.MatchString(name) || oidcSlotFilename.MatchString(name)
}
}
func recordListLinks(directory string) func(string, uint64) bool {
if directory == "sessions" {
return func(name string, links uint64) bool {
return digestFilename.MatchString(name) && links == 1
}
}
return func(name string, links uint64) bool {
if digestFilename.MatchString(name) || claimFilename.MatchString(name) {
return links == 1 || links == 2
}
return oidcSlotFilename.MatchString(name) && links == 1
}
}
func claimConsume(directory safeio.PrivateDirectoryHandle, filename string) (response, error) {
claim := asClaimFilename(filename)
claimed, err := directory.ClaimRegular(filename, claim)
if err != nil {
return response{}, errInvalid
}
if !claimed {
return response{Version: protocolVersion, OK: true}, nil
}
contents, found, err := safeio.ReadCanonicalPrivateClaim(source, claim, maximumOIDCStateBytes)
contents, found, err := directory.ReadClaim(filename, claim, maximumOIDCStateBytes)
if err != nil || !found {
return response{}, errInvalid
}
removed, err := safeio.RemoveCanonicalPrivateClaim(source, claim)
removed, err := directory.RemoveClaim(filename, claim)
if err != nil || !removed {
return response{}, errInvalid
}
@@ -167,6 +167,369 @@ func TestProtocolEnsuresTheCompletePrivateSessionLayout(t *testing.T) {
runRejected(t, request{Version: 1, Operation: "ensure-layout", Root: root, Directory: "sessions"})
}
func TestProtocolPinsTheOriginalPrivateRootBeforeEveryRecordMutation(t *testing.T) {
parent := privateTestRoot(t)
root := filepath.Join(parent, "auth")
replacement := filepath.Join(parent, "replacement")
moved := filepath.Join(parent, "auth-original")
filename := "eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee.json"
for _, directory := range []string{root, replacement} {
if err := safeio.EnsurePrivateDirectory(directory); err != nil {
t.Fatal(err)
}
for _, child := range []string{"sessions", "oidc"} {
if err := safeio.EnsurePrivateDirectory(filepath.Join(directory, child)); err != nil {
t.Fatal(err)
}
}
}
swapped := false
blocked := false
restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) {
if stage != "after-auth-root-open" || swapped || blocked {
return
}
if runtime.GOOS == "windows" {
if err := os.Rename(root, moved); err == nil {
t.Fatal("retained Windows root handle permitted rename")
}
blocked = true
return
}
if err := os.Rename(root, moved); err != nil {
t.Fatal(err)
}
if err := os.Rename(replacement, root); err != nil {
t.Fatal(err)
}
swapped = true
})
t.Cleanup(restore)
runRequest(t, request{
Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename,
ContentBase64: base64.StdEncoding.EncodeToString([]byte("pinned-root-record")),
})
if !swapped && !blocked {
t.Fatal("record operation did not expose the retained-root test hook")
}
if runtime.GOOS == "windows" {
if _, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(root, "sessions", filename), maximumSessionBytes); err != nil {
t.Fatalf("record missing from retained Windows root: %v", err)
}
return
}
contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(moved, "sessions", filename), maximumSessionBytes)
if err != nil || string(contents) != "pinned-root-record" {
t.Fatalf("pinned root contents = %q error = %v", contents, err)
}
if _, err := os.Lstat(filepath.Join(root, "sessions", filename)); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("replacement root was mutated: %v", err)
}
}
func TestProtocolPinsTheOriginalPrivateRootForEveryRecordOperation(t *testing.T) {
filename := "eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee.json"
otherFilename := "ffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff.json"
for _, operation := range []string{
"create", "read", "replace", "remove", "list", "claim-consume", "read-claim", "remove-claim",
} {
t.Run(operation, func(t *testing.T) {
parent := privateTestRoot(t)
root := filepath.Join(parent, "auth")
replacement := filepath.Join(parent, "replacement")
moved := filepath.Join(parent, "auth-original")
for _, directory := range []string{root, replacement} {
if err := safeio.EnsurePrivateDirectory(directory); err != nil {
t.Fatal(err)
}
for _, child := range []string{"sessions", "oidc"} {
if err := safeio.EnsurePrivateDirectory(filepath.Join(directory, child)); err != nil {
t.Fatal(err)
}
}
}
write := func(base, directory, name, contents string) {
t.Helper()
if err := safeio.WriteCanonicalNewPrivateFile(filepath.Join(base, directory, name), []byte(contents), 0o600); err != nil {
t.Fatal(err)
}
}
pair := func(base, name, contents string) {
t.Helper()
write(base, "oidc", name, contents)
if err := os.Link(filepath.Join(base, "oidc", name), filepath.Join(base, "oidc", asClaimFilename(name))); err != nil {
t.Fatal(err)
}
}
input := request{Version: 1, Operation: operation, Root: root}
switch operation {
case "create":
input.Directory, input.Filename = "sessions", filename
input.ContentBase64 = base64.StdEncoding.EncodeToString([]byte("original-created"))
case "read":
write(root, "sessions", filename, "original-read")
write(replacement, "sessions", filename, "replacement-read")
input.Directory, input.Filename = "sessions", filename
case "replace":
write(root, "sessions", filename, "original-before")
write(replacement, "sessions", filename, "replacement-before")
input.Directory, input.Filename = "sessions", filename
input.ContentBase64 = base64.StdEncoding.EncodeToString([]byte("original-replaced"))
case "remove":
write(root, "sessions", filename, "original-remove")
write(replacement, "sessions", filename, "replacement-remove")
input.Directory, input.Filename = "sessions", filename
case "list":
write(root, "sessions", filename, "original-list")
write(replacement, "sessions", otherFilename, "replacement-list")
input.Directory = "sessions"
case "claim-consume":
write(root, "oidc", filename, "original-claim")
write(replacement, "oidc", filename, "replacement-claim")
input.Directory, input.Filename = "oidc", filename
case "read-claim", "remove-claim":
pair(root, filename, "original-pair")
pair(replacement, filename, "replacement-pair")
input.Directory, input.Filename = "oidc", filename
}
swapped := false
blocked := false
restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) {
if stage != "after-auth-root-open" || swapped || blocked {
return
}
if runtime.GOOS == "windows" {
if err := os.Rename(root, moved); err == nil {
t.Fatal("retained Windows root handle permitted rename")
}
blocked = true
return
}
if err := os.Rename(root, moved); err != nil {
t.Fatal(err)
}
if err := os.Rename(replacement, root); err != nil {
t.Fatal(err)
}
swapped = true
})
t.Cleanup(restore)
output := runRequest(t, input)
if !swapped && !blocked {
t.Fatal("record operation did not expose the retained-root test hook")
}
original := root
replacementRoot := replacement
if swapped {
original = moved
replacementRoot = root
}
switch operation {
case "create":
contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(original, "sessions", filename), maximumSessionBytes)
if err != nil || string(contents) != "original-created" {
t.Fatalf("pinned create contents = %q error = %v", contents, err)
}
if _, err := os.Lstat(filepath.Join(replacementRoot, "sessions", filename)); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("replacement root was mutated: %v", err)
}
case "read":
if !output.Found || decodeContent(t, output) != "original-read" {
t.Fatalf("pinned read = %#v", output)
}
case "replace":
contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(original, "sessions", filename), maximumSessionBytes)
if err != nil || string(contents) != "original-replaced" {
t.Fatalf("pinned replace contents = %q error = %v", contents, err)
}
replacementContents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(replacementRoot, "sessions", filename), maximumSessionBytes)
if err != nil || string(replacementContents) != "replacement-before" {
t.Fatalf("replacement contents = %q error = %v", replacementContents, err)
}
case "remove":
if _, err := os.Lstat(filepath.Join(original, "sessions", filename)); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("pinned remove left original record: %v", err)
}
if _, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(replacementRoot, "sessions", filename), maximumSessionBytes); err != nil {
t.Fatalf("replacement record was mutated: %v", err)
}
case "list":
if output.Entries == nil || len(*output.Entries) != 1 || (*output.Entries)[0].Name != filename {
t.Fatalf("pinned list = %#v", output)
}
case "claim-consume":
if !output.Found || decodeContent(t, output) != "original-claim" {
t.Fatalf("pinned claim consume = %#v", output)
}
if _, err := os.Lstat(filepath.Join(original, "oidc", filename)); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("pinned claim source remains: %v", err)
}
if contents, err := safeio.ReadCanonicalPrivateRegular(filepath.Join(replacementRoot, "oidc", filename), maximumOIDCStateBytes); err != nil || string(contents) != "replacement-claim" {
t.Fatalf("replacement claim contents = %q error = %v", contents, err)
}
case "read-claim":
if !output.Found || decodeContent(t, output) != "original-pair" {
t.Fatalf("pinned claim read = %#v", output)
}
case "remove-claim":
for _, name := range []string{filename, asClaimFilename(filename)} {
if _, err := os.Lstat(filepath.Join(original, "oidc", name)); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("pinned claim remove left original %s: %v", name, err)
}
if _, err := os.Lstat(filepath.Join(replacementRoot, "oidc", name)); err != nil {
t.Fatalf("replacement claim pair was mutated: %v", err)
}
}
}
})
}
}
func TestProtocolPinsTheOriginalPrivateRootForLayoutCreation(t *testing.T) {
parent := privateTestRoot(t)
root := filepath.Join(parent, "auth")
replacement := filepath.Join(parent, "replacement")
moved := filepath.Join(parent, "auth-original")
for _, directory := range []string{root, replacement} {
if err := safeio.EnsurePrivateDirectory(directory); err != nil {
t.Fatal(err)
}
}
swapped := false
blocked := false
restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) {
if stage != "after-auth-root-open" || swapped || blocked {
return
}
if runtime.GOOS == "windows" {
if err := os.Rename(root, moved); err == nil {
t.Fatal("retained Windows root handle permitted rename")
}
blocked = true
return
}
if err := os.Rename(root, moved); err != nil {
t.Fatal(err)
}
if err := os.Rename(replacement, root); err != nil {
t.Fatal(err)
}
swapped = true
})
t.Cleanup(restore)
runRequest(t, request{Version: 1, Operation: "ensure-layout", Root: root})
if !swapped && !blocked {
t.Fatal("layout creation did not expose the retained-root test hook")
}
original := root
replacementRoot := replacement
if swapped {
original = moved
replacementRoot = root
}
for _, child := range []string{"sessions", "oidc"} {
if err := safeio.ValidatePrivateDirectory(filepath.Join(original, child)); err != nil {
t.Fatalf("pinned %s child: %v", child, err)
}
if _, err := os.Lstat(filepath.Join(replacementRoot, child)); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("replacement root received layout mutation: %v", err)
}
}
}
func TestProtocolValidateRootRechecksChildrenObservedMissingUnderTheRetainedRoot(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
if err := safeio.EnsurePrivateDirectory(root); err != nil {
t.Fatal(err)
}
installedUnsafeChild := false
restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) {
if stage != "after-auth-layout-first-pass" || installedUnsafeChild {
return
}
if err := os.Mkdir(filepath.Join(root, "sessions"), 0o755); err != nil {
t.Fatal(err)
}
installedUnsafeChild = true
})
t.Cleanup(restore)
runRejected(t, request{Version: 1, Operation: "validate-root", Root: root})
if !installedUnsafeChild {
t.Fatal("layout validation did not expose the missing-child recheck hook")
}
}
func TestProtocolValidateLayoutPinsTheOriginalPOSIXRootAcrossMissingAndUnsafeChildren(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("native Windows retained-handle coverage uses DACLs in storage_windows_test.go")
}
for _, scenario := range []struct {
name string
originalUnsafe bool
replacementUnsafe bool
wantAccepted bool
}{
{name: "rejects an unsafe retained child after a safe lexical replacement", originalUnsafe: true, wantAccepted: false},
{name: "accepts missing retained children despite an unsafe lexical replacement", replacementUnsafe: true, wantAccepted: true},
} {
t.Run(scenario.name, func(t *testing.T) {
parent := privateTestRoot(t)
root := filepath.Join(parent, "auth")
replacement := filepath.Join(parent, "replacement")
moved := filepath.Join(parent, "auth-original")
for _, directory := range []string{root, replacement} {
if err := safeio.EnsurePrivateDirectory(directory); err != nil {
t.Fatal(err)
}
}
if scenario.originalUnsafe {
if err := os.Mkdir(filepath.Join(root, "sessions"), 0o755); err != nil {
t.Fatal(err)
}
}
if scenario.replacementUnsafe {
if err := os.Mkdir(filepath.Join(replacement, "sessions"), 0o755); err != nil {
t.Fatal(err)
}
}
swapped := false
restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) {
if stage != "after-auth-root-open" || swapped {
return
}
if err := os.Rename(root, moved); err != nil {
t.Fatal(err)
}
if err := os.Rename(replacement, root); err != nil {
t.Fatal(err)
}
swapped = true
})
t.Cleanup(restore)
input := request{Version: 1, Operation: "validate-root", Root: root}
if scenario.wantAccepted {
runRequest(t, input)
} else {
runRejected(t, input)
}
if !swapped {
t.Fatal("layout validation did not retain and expose the original root")
}
for _, child := range []string{"sessions", "oidc"} {
if _, err := os.Lstat(filepath.Join(moved, child)); !errors.Is(err, os.ErrNotExist) && !(scenario.originalUnsafe && child == "sessions") {
t.Fatalf("validate-root mutated retained %s child: %v", child, err)
}
}
})
}
}
func TestProtocolReadsOnlyBoundedPrivateAuthConfig(t *testing.T) {
root := privateTestRoot(t)
filename := "auth.yaml"
@@ -199,6 +562,36 @@ func TestProtocolReadsOnlyBoundedPrivateAuthConfig(t *testing.T) {
runRejected(t, request{Version: 1, Operation: "read-auth-config", Root: root, Directory: "sessions", Filename: filename})
}
func TestProtocolReadsLocalUsersOnlyAsABoundedPrivateRegularFile(t *testing.T) {
root := privateTestRoot(t)
filename := "users.yaml"
path := filepath.Join(root, filename)
contents := []byte("version: 1\nusers: []\n")
if err := safeio.WriteCanonicalNewPrivateFile(path, contents, 0o600); err != nil {
t.Fatal(err)
}
read := runRequest(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: filename})
if !read.Found || decodeContent(t, read) != string(contents) {
t.Fatalf("read-local-users = %#v", read)
}
hardLink := filepath.Join(root, "users-copy.yaml")
if err := os.Link(path, hardLink); err != nil {
t.Fatal(err)
}
runRejected(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: filename})
if err := os.Remove(hardLink); err != nil {
t.Fatal(err)
}
if err := os.Remove(path); err != nil {
t.Fatal(err)
}
if err := safeio.WriteCanonicalNewPrivateFile(path, bytes.Repeat([]byte("x"), maximumAuthConfigBytes+1), 0o600); err != nil {
t.Fatal(err)
}
runRejected(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: filename})
}
func TestProtocolPermitsBoundedReservationSlotsOnlyForOIDCRecords(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
slot := "slot-00.json"
@@ -71,6 +71,96 @@ func TestProtocolValidatesAndCreatesTheCompleteWindowsSessionLayout(t *testing.T
runRejected(t, request{Version: 1, Operation: "validate-root", Root: root})
}
func TestProtocolReadsWindowsLocalUsersWithOwnerOnlyDACLAndNoReparseFallback(t *testing.T) {
root := filepath.Join(t.TempDir(), "auth")
if err := safeio.EnsurePrivateDirectory(root); err != nil {
t.Fatal(err)
}
path := filepath.Join(root, "users.yaml")
contents := []byte("version: 1\nusers: []\n")
if err := safeio.WriteCanonicalNewPrivateFile(path, contents, 0o600); err != nil {
t.Fatal(err)
}
read := runRequest(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: "users.yaml"})
if !read.Found || decodeContent(t, read) != string(contents) {
t.Fatalf("read-local-users = %#v", read)
}
if err := setPermissiveDACL(path); err != nil {
t.Fatal(err)
}
runRejected(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: "users.yaml"})
if err := os.Remove(path); err != nil {
t.Fatal(err)
}
if err := os.Symlink(filepath.Join(root, "missing-target.yaml"), path); err != nil {
t.Skipf("Windows host does not permit test symlink creation: %v", err)
}
runRejected(t, request{Version: 1, Operation: "read-local-users", Root: root, Filename: "users.yaml"})
}
func TestProtocolValidateLayoutPinsWindowsRootAcrossMissingAndUnsafeChildren(t *testing.T) {
for _, scenario := range []struct {
name string
originalUnsafe bool
replacementUnsafe bool
wantAccepted bool
}{
{name: "rejects unsafe retained child", originalUnsafe: true, wantAccepted: false},
{name: "accepts missing retained child while replacement is unsafe", replacementUnsafe: true, wantAccepted: true},
} {
t.Run(scenario.name, func(t *testing.T) {
parent := t.TempDir()
root := filepath.Join(parent, "auth")
replacement := filepath.Join(parent, "replacement")
moved := filepath.Join(parent, "auth-original")
for _, directory := range []string{root, replacement} {
if err := safeio.EnsurePrivateDirectory(directory); err != nil {
t.Fatal(err)
}
}
if scenario.originalUnsafe {
path := filepath.Join(root, "sessions")
if err := safeio.EnsurePrivateDirectory(path); err != nil {
t.Fatal(err)
}
if err := setPermissiveDACL(path); err != nil {
t.Fatal(err)
}
}
if scenario.replacementUnsafe {
path := filepath.Join(replacement, "sessions")
if err := safeio.EnsurePrivateDirectory(path); err != nil {
t.Fatal(err)
}
if err := setPermissiveDACL(path); err != nil {
t.Fatal(err)
}
}
blocked := false
restore := safeio.SetPrivateDirectoryTestHookForTest(func(stage string) {
if stage != "after-auth-root-open" || blocked {
return
}
if err := os.Rename(root, moved); err == nil {
t.Fatal("retained Windows root handle permitted rename")
}
blocked = true
})
t.Cleanup(restore)
input := request{Version: 1, Operation: "validate-root", Root: root}
if scenario.wantAccepted {
runRequest(t, input)
} else {
runRejected(t, input)
}
if !blocked {
t.Fatal("layout validation did not retain the Windows root handle")
}
})
}
}
func setPermissiveDACL(path string) error {
world, err := windows.StringToSid("S-1-1-0")
if err != nil {
+67
View File
@@ -0,0 +1,67 @@
package safeio
import (
"strings"
"sync"
)
// PrivateDirectoryHandle pins one private directory for the duration of a storage operation.
// Implementations use descriptor-relative operations on POSIX and retained no-delete handles on
// Windows. Callers must close every returned child before releasing its parent.
type PrivateDirectoryHandle interface {
Close() error
Validate() error
OpenChild(name string, ensure bool) (PrivateDirectoryHandle, bool, error)
CreateRegular(name string, contents []byte) (bool, error)
ReadRegular(name string, maximum int64) ([]byte, bool, error)
ReplaceRegular(name string, contents []byte) error
RemoveRegular(name string) (bool, error)
ListPage(maximumEntries int, afterName string, validName func(string) bool, validLinks func(string, uint64) bool) (PrivateDirectoryPage, error)
ClaimRegular(source, claim string) (bool, error)
ReadClaim(source, claim string, maximum int64) ([]byte, bool, error)
RemoveClaim(source, claim string) (bool, error)
}
// OpenPrivateDirectory validates or creates the final private directory while retaining the
// opened canonical directory handle. With ensure=false, found=false means the final component is
// absent but its already-opened parent proves the same operation could safely create it.
func OpenPrivateDirectory(path string, ensure bool) (PrivateDirectoryHandle, bool, error) {
if err := ValidateCanonicalPath(path); err != nil {
return nil, false, ErrUnsafeFile
}
return openPrivateDirectory(path, ensure)
}
func validPrivateLeafName(name string) bool {
return name != "" && name != "." && name != ".." && !strings.ContainsAny(name, "\\/:\x00")
}
var privateDirectoryTestHook struct {
sync.RWMutex
hook func(string)
}
// SetPrivateDirectoryTestHookForTest installs a deterministic race hook for internal package
// tests. It has no production effect unless a test explicitly installs one.
func SetPrivateDirectoryTestHookForTest(hook func(string)) func() {
privateDirectoryTestHook.Lock()
previous := privateDirectoryTestHook.hook
privateDirectoryTestHook.hook = hook
privateDirectoryTestHook.Unlock()
return func() {
privateDirectoryTestHook.Lock()
privateDirectoryTestHook.hook = previous
privateDirectoryTestHook.Unlock()
}
}
// NotifyPrivateDirectoryTestHookForTest marks an internal retained-root boundary. It is called
// only by storage code and lets tests install deterministic directory replacement races.
func NotifyPrivateDirectoryTestHookForTest(stage string) {
privateDirectoryTestHook.RLock()
hook := privateDirectoryTestHook.hook
privateDirectoryTestHook.RUnlock()
if hook != nil {
hook(stage)
}
}
@@ -0,0 +1,420 @@
//go:build !windows
package safeio
import (
"errors"
"io"
"os"
"sort"
"time"
"golang.org/x/sys/unix"
)
type unixPrivateDirectory struct {
descriptor int
}
func openPrivateDirectory(path string, ensure bool) (PrivateDirectoryHandle, bool, error) {
parents, err := openCanonicalUnixParent(path)
if err != nil {
return nil, false, ErrUnsafeFile
}
defer parents.Close()
return openPrivateUnixDirectoryAt(parents.parent, parents.target, ensure)
}
func openPrivateUnixDirectoryAt(parent int, name string, ensure bool) (PrivateDirectoryHandle, bool, error) {
if parent < 0 || !validPrivateLeafName(name) {
return nil, false, ErrUnsafeFile
}
for attempt := 0; attempt < 2; attempt++ {
descriptor, err := unix.Openat(parent, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0)
if err == nil {
value := &unixPrivateDirectory{descriptor: descriptor}
if value.Validate() != nil {
_ = value.Close()
return nil, false, ErrUnsafeFile
}
return value, true, nil
}
if !errors.Is(err, unix.ENOENT) {
return nil, false, ErrUnsafeFile
}
if !ensure {
if unix.Faccessat(parent, ".", unix.W_OK|unix.X_OK, unix.AT_EACCESS) != nil {
return nil, false, ErrUnsafeFile
}
return nil, false, nil
}
if err := unix.Mkdirat(parent, name, 0o700); err != nil && !errors.Is(err, unix.EEXIST) {
return nil, false, ErrUnsafeFile
}
}
return nil, false, ErrUnsafeFile
}
func (directory *unixPrivateDirectory) Close() error {
if directory == nil || directory.descriptor < 0 {
return nil
}
err := unix.Close(directory.descriptor)
directory.descriptor = -1
if err != nil {
return ErrUnsafeFile
}
return nil
}
func (directory *unixPrivateDirectory) Validate() error {
if directory == nil || directory.descriptor < 0 {
return ErrUnsafeFile
}
var stat unix.Stat_t
if err := unix.Fstat(directory.descriptor, &stat); err != nil || !privateUnixDirectoryStat(&stat) {
return ErrUnsafeFile
}
return nil
}
func (directory *unixPrivateDirectory) OpenChild(name string, ensure bool) (PrivateDirectoryHandle, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) {
return nil, false, ErrUnsafeFile
}
child, found, err := openPrivateUnixDirectoryAt(directory.descriptor, name, ensure)
if err != nil || directory.Validate() != nil {
if child != nil {
_ = child.Close()
}
return nil, false, ErrUnsafeFile
}
return child, found, nil
}
func privateUnixRootRegular(stat *unix.Stat_t, links uint64) bool {
return stat != nil && stat.Mode&unix.S_IFMT == unix.S_IFREG && uint64(stat.Nlink) == links &&
stat.Uid == uint32(os.Geteuid()) && stat.Mode&0o7777 == 0o600
}
func sameUnixRootRegular(left, right unix.Stat_t) bool {
return left.Dev == right.Dev && left.Ino == right.Ino && left.Size == right.Size &&
left.Mtim == right.Mtim && left.Ctim == right.Ctim && left.Mode == right.Mode && left.Nlink == right.Nlink
}
func requirePrivateUnixRootRegularAt(directory int, name string, links uint64) (unix.Stat_t, error) {
var stat unix.Stat_t
if err := unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW); err != nil || !privateUnixRootRegular(&stat, links) {
return unix.Stat_t{}, ErrUnsafeFile
}
return stat, nil
}
func privateUnixRootRegularAtAllowedLinks(directory int, name string, allowed ...uint64) error {
var stat unix.Stat_t
if err := unix.Fstatat(directory, name, &stat, unix.AT_SYMLINK_NOFOLLOW); err != nil {
return ErrUnsafeFile
}
for _, links := range allowed {
if privateUnixRootRegular(&stat, links) {
return nil
}
}
return ErrUnsafeFile
}
func requireSamePrivateUnixRootPairAt(directory int, source, claim string) error {
left, err := requirePrivateUnixRootRegularAt(directory, source, 2)
if err != nil {
return ErrUnsafeFile
}
right, err := requirePrivateUnixRootRegularAt(directory, claim, 2)
if err != nil || !sameUnixPrivateFile(left, right) {
return ErrUnsafeFile
}
return nil
}
func isPrivateUnixRootClaimAbsentOrOrphanAt(directory int, source, claim string) bool {
var sourceStat, claimStat unix.Stat_t
sourceErr := unix.Fstatat(directory, source, &sourceStat, unix.AT_SYMLINK_NOFOLLOW)
claimErr := unix.Fstatat(directory, claim, &claimStat, unix.AT_SYMLINK_NOFOLLOW)
if errors.Is(sourceErr, unix.ENOENT) && errors.Is(claimErr, unix.ENOENT) {
return true
}
return errors.Is(sourceErr, unix.ENOENT) && claimErr == nil && privateUnixRootRegular(&claimStat, 1)
}
func writeAllPrivateRoot(file *os.File, contents []byte) error {
for written := 0; written < len(contents); {
count, err := file.Write(contents[written:])
written += count
if err != nil {
return err
}
if count == 0 {
return io.ErrShortWrite
}
}
return file.Sync()
}
func (directory *unixPrivateDirectory) CreateRegular(name string, contents []byte) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 {
return false, ErrUnsafeFile
}
descriptor, err := unix.Openat(directory.descriptor, name,
unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0o600)
if errors.Is(err, unix.EEXIST) {
if _, existingErr := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); existingErr != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return false, nil
}
if err != nil {
return false, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-create")
if file == nil {
_ = unix.Close(descriptor)
_ = unix.Unlinkat(directory.descriptor, name, 0)
return false, ErrUnsafeFile
}
failed := true
defer func() {
if failed {
_ = unix.Unlinkat(directory.descriptor, name, 0)
}
}()
if unix.Fchmod(descriptor, 0o600) != nil {
_ = file.Close()
return false, ErrUnsafeFile
}
var stat unix.Stat_t
if unix.Fstat(descriptor, &stat) != nil || !privateUnixRootRegular(&stat, 1) || writeAllPrivateRoot(file, contents) != nil || file.Close() != nil {
return false, ErrUnsafeFile
}
if directory.Validate() != nil || unix.Fsync(directory.descriptor) != nil {
return false, ErrUnsafeFile
}
failed = false
return true, nil
}
func (directory *unixPrivateDirectory) ReadRegular(name string, maximum int64) ([]byte, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) || maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, false, ErrUnsafeFile
}
before, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1)
if err != nil {
var stat unix.Stat_t
if statErr := unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW); errors.Is(statErr, unix.ENOENT) {
return nil, false, nil
}
return nil, false, ErrUnsafeFile
}
descriptor, err := unix.Openat(directory.descriptor, name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0)
if err != nil {
return nil, false, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-read")
if file == nil {
_ = unix.Close(descriptor)
return nil, false, ErrUnsafeFile
}
contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1))
var opened unix.Stat_t
statErr := unix.Fstat(descriptor, &opened)
closeErr := file.Close()
var current unix.Stat_t
currentErr := unix.Fstatat(directory.descriptor, name, &current, unix.AT_SYMLINK_NOFOLLOW)
if readErr != nil || statErr != nil || closeErr != nil || currentErr != nil || int64(len(contents)) > maximum ||
!privateUnixRootRegular(&opened, 1) || !privateUnixRootRegular(&current, 1) ||
!sameUnixRootRegular(before, opened) || !sameUnixRootRegular(opened, current) || directory.Validate() != nil {
return nil, false, ErrUnsafeFile
}
return contents, true, nil
}
func (directory *unixPrivateDirectory) ReplaceRegular(name string, contents []byte) error {
if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 {
return ErrUnsafeFile
}
if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil {
return ErrUnsafeFile
}
temporary, err := writePrivateTemporaryAt(directory.descriptor, contents)
if err != nil {
return ErrUnsafeFile
}
defer func() { _ = unix.Unlinkat(directory.descriptor, temporary, 0) }()
if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil || directory.Validate() != nil {
return ErrUnsafeFile
}
if err := unix.Renameat(directory.descriptor, temporary, directory.descriptor, name); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil {
return ErrUnsafeFile
}
return nil
}
func (directory *unixPrivateDirectory) RemoveRegular(name string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) {
return false, ErrUnsafeFile
}
if _, err := requirePrivateUnixRootRegularAt(directory.descriptor, name, 1); err != nil {
var stat unix.Stat_t
if statErr := unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW); errors.Is(statErr, unix.ENOENT) {
return false, nil
}
return false, ErrUnsafeFile
}
if err := unix.Unlinkat(directory.descriptor, name, 0); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func (directory *unixPrivateDirectory) ListPage(
maximumEntries int,
afterName string,
validName func(string) bool,
validLinks func(string, uint64) bool,
) (PrivateDirectoryPage, error) {
if directory.Validate() != nil || maximumEntries < 1 || maximumEntries > 4096 || validName == nil || validLinks == nil || (afterName != "" && !validName(afterName)) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
var before unix.Stat_t
if unix.Fstat(directory.descriptor, &before) != nil || !privateUnixDirectoryStat(&before) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
duplicate, err := unix.Dup(directory.descriptor)
if err != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
file := os.NewFile(uintptr(duplicate), "tht-safeio-private-root-list")
if file == nil {
_ = unix.Close(duplicate)
return PrivateDirectoryPage{}, ErrUnsafeFile
}
defer file.Close()
seen := make(map[string]struct{}, maximumEntries+1)
selected := make([]PrivateDirectoryEntry, 0, maximumEntries+1)
scanned := 0
for {
entries, readErr := file.ReadDir(1)
if readErr != nil && !errors.Is(readErr, io.EOF) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if len(entries) == 0 {
break
}
if len(entries) != 1 {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
scanned++
if scanned > maximumPrivateDirectoryPageScanEntries {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
name := entries[0].Name()
if !validPrivateLeafName(name) || !validName(name) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if _, duplicate := seen[name]; duplicate {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
seen[name] = struct{}{}
var stat unix.Stat_t
if unix.Fstatat(directory.descriptor, name, &stat, unix.AT_SYMLINK_NOFOLLOW) != nil ||
!privateUnixRootRegular(&stat, uint64(stat.Nlink)) || !validLinks(name, uint64(stat.Nlink)) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if name > afterName {
selected = appendBoundedPrivateDirectoryEntry(selected, PrivateDirectoryEntry{
Name: name, ModifiedUnixMs: time.Unix(stat.Mtim.Sec, stat.Mtim.Nsec).UnixMilli(),
}, maximumEntries+1)
}
if errors.Is(readErr, io.EOF) {
break
}
}
var after unix.Stat_t
if unix.Fstat(directory.descriptor, &after) != nil || !privateUnixDirectoryStat(&after) ||
before.Dev != after.Dev || before.Ino != after.Ino || before.Mode != after.Mode || before.Mtim != after.Mtim || before.Ctim != after.Ctim || directory.Validate() != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
sort.Slice(selected, func(left, right int) bool { return selected[left].Name < selected[right].Name })
more := len(selected) > maximumEntries
if more {
selected = selected[:maximumEntries]
}
return PrivateDirectoryPage{Entries: selected, More: more}, nil
}
func (directory *unixPrivateDirectory) ClaimRegular(source, claim string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) {
return false, ErrUnsafeFile
}
if err := privateUnixRootRegularAtAllowedLinks(directory.descriptor, source, 1); err != nil {
if requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim) == nil || isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) {
return false, nil
}
return false, ErrUnsafeFile
}
if err := unix.Linkat(directory.descriptor, source, directory.descriptor, claim, 0); err != nil {
if errors.Is(err, unix.EEXIST) && privateUnixRootRegularAtAllowedLinks(directory.descriptor, claim, 1, 2) == nil {
return false, nil
}
return false, ErrUnsafeFile
}
if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func (directory *unixPrivateDirectory) ReadClaim(source, claim string, maximum int64) ([]byte, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) || maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, false, ErrUnsafeFile
}
if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil {
if isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) {
return nil, false, nil
}
return nil, false, ErrUnsafeFile
}
descriptor, err := unix.Openat(directory.descriptor, source, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0)
if err != nil {
return nil, false, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-private-root-claim")
if file == nil {
_ = unix.Close(descriptor)
return nil, false, ErrUnsafeFile
}
contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1))
var after unix.Stat_t
statErr := unix.Fstat(descriptor, &after)
closeErr := file.Close()
if readErr != nil || statErr != nil || closeErr != nil || int64(len(contents)) > maximum || !privateUnixRootRegular(&after, 2) ||
requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim) != nil || directory.Validate() != nil {
return nil, false, ErrUnsafeFile
}
return contents, true, nil
}
func (directory *unixPrivateDirectory) RemoveClaim(source, claim string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) {
return false, ErrUnsafeFile
}
if err := requireSamePrivateUnixRootPairAt(directory.descriptor, source, claim); err != nil {
if isPrivateUnixRootClaimAbsentOrOrphanAt(directory.descriptor, source, claim) {
return false, nil
}
return false, ErrUnsafeFile
}
if err := unix.Unlinkat(directory.descriptor, source, 0); err != nil || privateUnixRootRegularAtAllowedLinks(directory.descriptor, claim, 1) != nil ||
unix.Unlinkat(directory.descriptor, claim, 0) != nil || unix.Fsync(directory.descriptor) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}
@@ -0,0 +1,921 @@
//go:build windows
package safeio
import (
"errors"
"io"
"os"
"runtime"
"sort"
"time"
"unsafe"
"golang.org/x/sys/windows"
)
// windowsPrivateDirectory keeps the root's canonical component handles alive, then uses NT
// RootDirectory-relative opens for every descendant. Unlike a lexical child path, an NT relative
// object name is resolved by the already-open directory handle and cannot be redirected by a
// rename, replacement, or reparse point at the original root path.
type windowsPrivateDirectory struct {
anchors *windowsParentHandles
parent windows.Handle
handle windows.Handle
info windows.ByHandleFileInformation
}
type windowsPrivateRegularAt struct {
handle windows.Handle
info windows.ByHandleFileInformation
}
func openPrivateDirectory(path string, ensure bool) (PrivateDirectoryHandle, bool, error) {
anchors, target, err := openCanonicalWindowsParent(path)
if err != nil || anchors == nil || len(anchors.handles) == 0 {
if anchors != nil {
anchors.Close()
}
return nil, false, ErrUnsafeFile
}
parent := anchors.handles[len(anchors.handles)-1]
handle, found, err := openWindowsPrivateDirectoryAt(parent, target, false)
if err != nil {
anchors.Close()
return nil, false, ErrUnsafeFile
}
if !found {
// The final root is absent. The retained canonical parent must prove that the same
// ensure operation could create it; validation itself must remain side-effect free.
// FILE_APPEND_DATA is the Win32 spelling of directory FILE_ADD_SUBDIRECTORY.
probe, probeErr := openWindowsComponentWithAccess(anchors.directory, true, windows.FILE_APPEND_DATA)
if probeErr != nil {
anchors.Close()
return nil, false, ErrUnsafeFile
}
_ = windows.CloseHandle(probe)
if !ensure {
anchors.Close()
return nil, false, nil
}
// The canonical parent chain remains pinned by anchors; this extra handle supplies
// FILE_ADD_SUBDIRECTORY for the one initial root creation without reopening a child
// beneath the private root lexically.
writableParent, writableErr := openWindowsComponentWithAccess(anchors.directory, true, windows.FILE_APPEND_DATA)
if writableErr != nil {
anchors.Close()
return nil, false, ErrUnsafeFile
}
handle, found, err = openWindowsPrivateDirectoryAt(writableParent, target, true)
_ = windows.CloseHandle(writableParent)
if err != nil || !found {
anchors.Close()
return nil, false, ErrUnsafeFile
}
}
value := &windowsPrivateDirectory{anchors: anchors, handle: handle}
if value.captureAndValidate() != nil {
_ = value.Close()
return nil, false, ErrUnsafeFile
}
return value, true, nil
}
func openWindowsPrivateDirectoryAt(parent windows.Handle, name string, ensure bool) (windows.Handle, bool, error) {
if parent == 0 || !validPrivateLeafName(name) {
return 0, false, ErrUnsafeFile
}
for attempt := 0; attempt < 2; attempt++ {
handle, err := openWindowsRelativeDirectory(parent, name)
if err == nil {
return handle, true, nil
}
if !isWindowsRelativeNotFound(err) {
return 0, false, ErrUnsafeFile
}
if !ensure {
return 0, false, nil
}
handle, err = createWindowsRelativePrivateDirectory(parent, name)
if err == nil {
return handle, true, nil
}
}
return 0, false, ErrUnsafeFile
}
func openWindowsRelativeDirectory(parent windows.Handle, name string) (windows.Handle, error) {
handle, err := openWindowsRelativeObject(
parent,
name,
// The retained directory handle is also the RootDirectory for create, rename,
// hard-link, and delete operations below, so it needs the owner's full private
// directory capability rather than a read-only probe handle.
windows.FILE_GENERIC_READ|windows.FILE_GENERIC_WRITE|windows.DELETE,
windows.FILE_OPEN,
windows.FILE_DIRECTORY_FILE|windows.FILE_SYNCHRONOUS_IO_NONALERT|windows.FILE_OPEN_REPARSE_POINT,
nil,
)
if err != nil {
return 0, err
}
if _, err := privateWindowsDirectoryInfo(handle); err != nil {
_ = windows.CloseHandle(handle)
return 0, ErrUnsafeFile
}
return handle, nil
}
func createWindowsRelativePrivateDirectory(parent windows.Handle, name string) (windows.Handle, error) {
security, err := newOwnerOnlySecurityDescriptor()
if err != nil {
return 0, ErrUnsafeFile
}
defer security.Close()
handle, err := openWindowsRelativeObject(
parent,
name,
windows.FILE_GENERIC_READ|windows.FILE_GENERIC_WRITE|windows.DELETE,
windows.FILE_CREATE,
windows.FILE_DIRECTORY_FILE|windows.FILE_SYNCHRONOUS_IO_NONALERT|windows.FILE_OPEN_REPARSE_POINT,
security,
)
if err != nil {
return 0, err
}
if _, err := privateWindowsDirectoryInfo(handle); err != nil {
_ = markWindowsHandleForDelete(handle)
_ = windows.CloseHandle(handle)
return 0, ErrUnsafeFile
}
return handle, nil
}
func openWindowsRelativeObject(
parent windows.Handle,
name string,
access uint32,
disposition uint32,
options uint32,
security *ownerOnlySecurityDescriptor,
) (windows.Handle, error) {
if parent == 0 || !validPrivateLeafName(name) {
return 0, ErrUnsafeFile
}
objectName, err := windows.NewNTUnicodeString(name)
if err != nil {
return 0, ErrUnsafeFile
}
attributes := &windows.OBJECT_ATTRIBUTES{
Length: uint32(unsafe.Sizeof(windows.OBJECT_ATTRIBUTES{})),
RootDirectory: parent,
ObjectName: objectName,
Attributes: windows.OBJ_CASE_INSENSITIVE,
SecurityDescriptor: nil,
}
if security != nil {
attributes.SecurityDescriptor = security.descriptor
}
var (
handle windows.Handle
status windows.IO_STATUS_BLOCK
allocationSize int64
)
err = windows.NtCreateFile(
&handle,
access,
attributes,
&status,
&allocationSize,
windows.FILE_ATTRIBUTE_NORMAL,
windowsRetainedHandleShareMode,
disposition,
options,
0,
0,
)
runtime.KeepAlive(objectName)
runtime.KeepAlive(security)
if err != nil {
return 0, err
}
return handle, nil
}
func privateWindowsDirectoryInfo(handle windows.Handle) (windows.ByHandleFileInformation, error) {
var info windows.ByHandleFileInformation
if handle == 0 || windows.GetFileInformationByHandle(handle, &info) != nil ||
info.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 ||
info.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY == 0 ||
validateOwnerOnlyDACL(handle) != nil {
return windows.ByHandleFileInformation{}, ErrUnsafeFile
}
return info, nil
}
func privateWindowsRegularInfo(handle windows.Handle, allowedLinks ...uint32) (windows.ByHandleFileInformation, error) {
var info windows.ByHandleFileInformation
if handle == 0 || windows.GetFileInformationByHandle(handle, &info) != nil ||
info.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 ||
info.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY != 0 ||
validateOwnerOnlyDACL(handle) != nil {
return windows.ByHandleFileInformation{}, ErrUnsafeFile
}
for _, links := range allowedLinks {
if info.NumberOfLinks == links {
return info, nil
}
}
return windows.ByHandleFileInformation{}, ErrUnsafeFile
}
func isWindowsRelativeNotFound(err error) bool {
return errors.Is(err, windows.ERROR_FILE_NOT_FOUND) ||
errors.Is(err, windows.ERROR_PATH_NOT_FOUND) ||
errors.Is(err, windows.STATUS_NO_SUCH_FILE) ||
errors.Is(err, windows.STATUS_OBJECT_NAME_NOT_FOUND) ||
errors.Is(err, windows.STATUS_OBJECT_PATH_NOT_FOUND)
}
func isWindowsRelativeCollision(err error) bool {
return errors.Is(err, windows.ERROR_FILE_EXISTS) ||
errors.Is(err, windows.ERROR_ALREADY_EXISTS) ||
errors.Is(err, windows.STATUS_OBJECT_NAME_COLLISION)
}
func (directory *windowsPrivateDirectory) captureAndValidate() error {
if directory == nil || directory.handle == 0 {
return ErrUnsafeFile
}
info, err := privateWindowsDirectoryInfo(directory.handle)
if err != nil {
return ErrUnsafeFile
}
directory.info = info
return nil
}
func (directory *windowsPrivateDirectory) Close() error {
if directory == nil {
return nil
}
var result error
if directory.handle != 0 {
if err := windows.CloseHandle(directory.handle); err != nil {
result = ErrUnsafeFile
}
directory.handle = 0
}
if directory.parent != 0 {
if err := windows.CloseHandle(directory.parent); err != nil {
result = ErrUnsafeFile
}
directory.parent = 0
}
if directory.anchors != nil {
directory.anchors.Close()
directory.anchors = nil
}
return result
}
func sameWindowsPrivateDirectoryIdentity(left, right windows.ByHandleFileInformation) bool {
return left.VolumeSerialNumber == right.VolumeSerialNumber && left.FileIndexHigh == right.FileIndexHigh &&
left.FileIndexLow == right.FileIndexLow && left.FileAttributes == right.FileAttributes
}
func sameWindowsPrivateDirectorySnapshot(left, right windows.ByHandleFileInformation) bool {
return sameWindowsPrivateDirectoryIdentity(left, right) && left.LastWriteTime == right.LastWriteTime
}
func (directory *windowsPrivateDirectory) Validate() error {
if directory == nil || directory.handle == 0 || (directory.anchors == nil && directory.parent == 0) {
return ErrUnsafeFile
}
if directory.anchors != nil && len(directory.anchors.handles) == 0 {
return ErrUnsafeFile
}
if directory.parent != 0 {
if _, err := privateWindowsDirectoryInfo(directory.parent); err != nil {
return ErrUnsafeFile
}
}
current, err := privateWindowsDirectoryInfo(directory.handle)
if err != nil || !sameWindowsPrivateDirectoryIdentity(directory.info, current) {
return ErrUnsafeFile
}
return nil
}
func duplicateWindowsRetainedHandle(handle windows.Handle) (windows.Handle, error) {
if handle == 0 {
return 0, ErrUnsafeFile
}
var duplicate windows.Handle
process := windows.CurrentProcess()
if err := windows.DuplicateHandle(process, handle, process, &duplicate, 0, false, windows.DUPLICATE_SAME_ACCESS); err != nil {
return 0, ErrUnsafeFile
}
return duplicate, nil
}
func (directory *windowsPrivateDirectory) OpenChild(name string, ensure bool) (PrivateDirectoryHandle, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) {
return nil, false, ErrUnsafeFile
}
parent, err := duplicateWindowsRetainedHandle(directory.handle)
if err != nil {
return nil, false, ErrUnsafeFile
}
handle, found, err := openWindowsPrivateDirectoryAt(parent, name, ensure)
if err != nil || !found {
_ = windows.CloseHandle(parent)
if err != nil {
return nil, false, ErrUnsafeFile
}
return nil, false, nil
}
child := &windowsPrivateDirectory{parent: parent, handle: handle}
if child.captureAndValidate() != nil || directory.Validate() != nil {
_ = child.Close()
return nil, false, ErrUnsafeFile
}
return child, true, nil
}
func openWindowsPrivateRegularAt(
parent windows.Handle,
name string,
access uint32,
allowedLinks ...uint32,
) (*windowsPrivateRegularAt, error) {
handle, err := openWindowsRelativeObject(
parent,
name,
access,
windows.FILE_OPEN,
windows.FILE_NON_DIRECTORY_FILE|windows.FILE_SYNCHRONOUS_IO_NONALERT|windows.FILE_OPEN_REPARSE_POINT,
nil,
)
if err != nil {
return nil, err
}
info, err := privateWindowsRegularInfo(handle, allowedLinks...)
if err != nil {
_ = windows.CloseHandle(handle)
return nil, ErrUnsafeFile
}
return &windowsPrivateRegularAt{handle: handle, info: info}, nil
}
func openWindowsPrivateRegularAtAllowedLinks(
parent windows.Handle,
name string,
access uint32,
allowedLinks ...uint32,
) (*windowsPrivateRegularAt, error) {
var (
lastError error
unsafeFound bool
)
for _, links := range allowedLinks {
value, err := openWindowsPrivateRegularAt(parent, name, access, links)
if err == nil {
return value, nil
}
lastError = err
if !isWindowsRelativeNotFound(err) {
unsafeFound = true
}
}
if unsafeFound {
return nil, ErrUnsafeFile
}
return nil, lastError
}
func createWindowsPrivateRegularAt(parent windows.Handle, name string) (*windowsPrivateRegularAt, error) {
security, err := newOwnerOnlySecurityDescriptor()
if err != nil {
return nil, ErrUnsafeFile
}
defer security.Close()
handle, err := openWindowsRelativeObject(
parent,
name,
windows.FILE_GENERIC_READ|windows.FILE_GENERIC_WRITE|windows.DELETE,
windows.FILE_CREATE,
windows.FILE_NON_DIRECTORY_FILE|windows.FILE_SYNCHRONOUS_IO_NONALERT|windows.FILE_OPEN_REPARSE_POINT,
security,
)
if err != nil {
return nil, err
}
info, err := privateWindowsRegularInfo(handle, 1)
if err != nil {
_ = markWindowsHandleForDelete(handle)
_ = windows.CloseHandle(handle)
return nil, ErrUnsafeFile
}
return &windowsPrivateRegularAt{handle: handle, info: info}, nil
}
func (value *windowsPrivateRegularAt) Close() error {
if value == nil || value.handle == 0 {
return nil
}
handle := value.handle
value.handle = 0
if err := windows.CloseHandle(handle); err != nil {
return ErrUnsafeFile
}
return nil
}
func writeWindowsPrivateRegular(value *windowsPrivateRegularAt, contents []byte) error {
if value == nil || value.handle == 0 || len(contents) == 0 {
return ErrUnsafeFile
}
for remaining := contents; len(remaining) > 0; {
var written uint32
if err := windows.WriteFile(value.handle, remaining, &written, nil); err != nil || written == 0 || int(written) > len(remaining) {
return ErrUnsafeFile
}
remaining = remaining[written:]
}
if err := windows.FlushFileBuffers(value.handle); err != nil {
return ErrUnsafeFile
}
info, err := privateWindowsRegularInfo(value.handle, 1)
if err != nil {
return ErrUnsafeFile
}
value.info = info
return nil
}
func readWindowsPrivateRegular(value *windowsPrivateRegularAt, maximum int64, links uint32) ([]byte, error) {
if value == nil || value.handle == 0 || maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, ErrUnsafeFile
}
contents := make([]byte, 0, 4096)
buffer := make([]byte, 4096)
for {
var read uint32
err := windows.ReadFile(value.handle, buffer, &read, nil)
if read > 0 {
if int64(len(contents))+int64(read) > maximum {
return nil, ErrUnsafeFile
}
contents = append(contents, buffer[:read]...)
}
if err != nil {
if errors.Is(err, windows.ERROR_HANDLE_EOF) {
break
}
return nil, ErrUnsafeFile
}
if read == 0 {
break
}
}
after, err := privateWindowsRegularInfo(value.handle, links)
if err != nil || !sameWindowsPrivateFile(value.info, after) {
return nil, ErrUnsafeFile
}
value.info = after
return contents, nil
}
func markWindowsHandleForDelete(handle windows.Handle) error {
if handle == 0 {
return ErrUnsafeFile
}
buffer := [1]byte{1}
var status windows.IO_STATUS_BLOCK
if err := windows.NtSetInformationFile(handle, &status, &buffer[0], uint32(len(buffer)), windows.FileDispositionInformation); err != nil {
return ErrUnsafeFile
}
return nil
}
func closeAndDeleteWindowsPrivateRegular(value *windowsPrivateRegularAt) error {
if value == nil || value.handle == 0 || markWindowsHandleForDelete(value.handle) != nil || value.Close() != nil {
return ErrUnsafeFile
}
return nil
}
func (directory *windowsPrivateDirectory) CreateRegular(name string, contents []byte) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 {
return false, ErrUnsafeFile
}
value, err := createWindowsPrivateRegularAt(directory.handle, name)
if err != nil {
existing, existingErr := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1)
if existingErr == nil {
_ = existing.Close()
return false, nil
}
return false, ErrUnsafeFile
}
published := false
defer func() {
if !published {
_ = closeAndDeleteWindowsPrivateRegular(value)
}
}()
if writeWindowsPrivateRegular(value, contents) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
if value.Close() != nil {
return false, ErrUnsafeFile
}
published = true
return true, nil
}
func (directory *windowsPrivateDirectory) ReadRegular(name string, maximum int64) ([]byte, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) || maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, false, ErrUnsafeFile
}
value, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1)
if isWindowsRelativeNotFound(err) {
return nil, false, nil
}
if err != nil {
return nil, false, ErrUnsafeFile
}
defer value.Close()
contents, err := readWindowsPrivateRegular(value, maximum, 1)
if err != nil || directory.Validate() != nil {
return nil, false, ErrUnsafeFile
}
return contents, true, nil
}
type windowsRelativeNameInformation struct {
Flags uint32
RootDirectory windows.Handle
FileNameLength uint32
FileName [1]uint16
}
func setWindowsRelativeNameInformation(
handle windows.Handle,
parent windows.Handle,
name string,
class uint32,
flags uint32,
) error {
if handle == 0 || parent == 0 || !validPrivateLeafName(name) {
return ErrUnsafeFile
}
encoded, err := windows.UTF16FromString(name)
if err != nil || len(encoded) < 2 {
return ErrUnsafeFile
}
nameBytes := (len(encoded) - 1) * 2
var header windowsRelativeNameInformation
size := int(unsafe.Offsetof(header.FileName)) + nameBytes
buffer := make([]byte, size)
value := (*windowsRelativeNameInformation)(unsafe.Pointer(&buffer[0]))
value.Flags = flags
value.RootDirectory = parent
value.FileNameLength = uint32(nameBytes)
copy(unsafe.Slice(&value.FileName[0], len(encoded)-1), encoded[:len(encoded)-1])
var status windows.IO_STATUS_BLOCK
if err := windows.NtSetInformationFile(handle, &status, &buffer[0], uint32(len(buffer)), class); err != nil {
return ErrUnsafeFile
}
runtime.KeepAlive(encoded)
runtime.KeepAlive(buffer)
return nil
}
func renameWindowsPrivateRegularAt(handle windows.Handle, parent windows.Handle, name string) error {
return setWindowsRelativeNameInformation(handle, parent, name, windows.FileRenameInformation, windows.FILE_RENAME_REPLACE_IF_EXISTS)
}
func linkWindowsPrivateRegularAt(handle windows.Handle, parent windows.Handle, name string) error {
return setWindowsRelativeNameInformation(handle, parent, name, windows.FileLinkInformation, 0)
}
func createWindowsPrivateTemporaryAt(parent windows.Handle, contents []byte) (*windowsPrivateRegularAt, error) {
for attempt := 0; attempt < 16; attempt++ {
name, err := randomTemporaryName()
if err != nil {
return nil, ErrUnsafeFile
}
value, err := createWindowsPrivateRegularAt(parent, name)
if err != nil {
if isWindowsRelativeCollision(err) {
continue
}
return nil, ErrUnsafeFile
}
if writeWindowsPrivateRegular(value, contents) == nil {
return value, nil
}
_ = closeAndDeleteWindowsPrivateRegular(value)
return nil, ErrUnsafeFile
}
return nil, ErrUnsafeFile
}
func (directory *windowsPrivateDirectory) ReplaceRegular(name string, contents []byte) error {
if directory.Validate() != nil || !validPrivateLeafName(name) || len(contents) == 0 {
return ErrUnsafeFile
}
existing, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1)
if err != nil {
return ErrUnsafeFile
}
if existing.Close() != nil {
return ErrUnsafeFile
}
temporary, err := createWindowsPrivateTemporaryAt(directory.handle, contents)
if err != nil {
return ErrUnsafeFile
}
renamed := false
defer func() {
if !renamed {
_ = closeAndDeleteWindowsPrivateRegular(temporary)
}
}()
current, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1)
if err != nil || current.Close() != nil || directory.Validate() != nil {
return ErrUnsafeFile
}
if renameWindowsPrivateRegularAt(temporary.handle, directory.handle, name) != nil || temporary.Close() != nil {
return ErrUnsafeFile
}
renamed = true
replaced, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ, 1)
if err != nil || replaced.Close() != nil || directory.Validate() != nil {
return ErrUnsafeFile
}
return nil
}
func (directory *windowsPrivateDirectory) RemoveRegular(name string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(name) {
return false, ErrUnsafeFile
}
value, err := openWindowsPrivateRegularAt(directory.handle, name, windows.FILE_GENERIC_READ|windows.DELETE, 1)
if isWindowsRelativeNotFound(err) {
return false, nil
}
if err != nil {
return false, ErrUnsafeFile
}
if closeAndDeleteWindowsPrivateRegular(value) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func (directory *windowsPrivateDirectory) ListPage(
maximumEntries int,
afterName string,
validName func(string) bool,
validLinks func(string, uint64) bool,
) (PrivateDirectoryPage, error) {
if directory.Validate() != nil || maximumEntries < 1 || maximumEntries > 4096 || validName == nil || validLinks == nil ||
(afterName != "" && !validName(afterName)) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
var before windows.ByHandleFileInformation
if err := windows.GetFileInformationByHandle(directory.handle, &before); err != nil ||
!sameWindowsPrivateDirectoryIdentity(directory.info, before) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
duplicate, err := duplicateWindowsRetainedHandle(directory.handle)
if err != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
file := os.NewFile(uintptr(duplicate), "tht-safeio-private-root-list")
if file == nil {
_ = windows.CloseHandle(duplicate)
return PrivateDirectoryPage{}, ErrUnsafeFile
}
defer file.Close()
seen := make(map[string]struct{}, maximumEntries+1)
selected := make([]PrivateDirectoryEntry, 0, maximumEntries+1)
scanned := 0
for {
entries, readErr := file.ReadDir(1)
if readErr != nil && !errors.Is(readErr, io.EOF) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if len(entries) == 0 {
break
}
if len(entries) != 1 {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
scanned++
if scanned > maximumPrivateDirectoryPageScanEntries {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
name := entries[0].Name()
if !validPrivateLeafName(name) || !validName(name) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if _, duplicate := seen[name]; duplicate {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
seen[name] = struct{}{}
value, valueErr := openWindowsPrivateRegularAtAllowedLinks(directory.handle, name, windows.FILE_GENERIC_READ, 1, 2)
if valueErr != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
info := value.info
if value.Close() != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if !validLinks(name, uint64(info.NumberOfLinks)) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if name > afterName {
selected = appendBoundedPrivateDirectoryEntry(selected, PrivateDirectoryEntry{
Name: name, ModifiedUnixMs: time.Unix(0, info.LastWriteTime.Nanoseconds()).UnixMilli(),
}, maximumEntries+1)
}
if errors.Is(readErr, io.EOF) {
break
}
}
var after windows.ByHandleFileInformation
if windows.GetFileInformationByHandle(directory.handle, &after) != nil || directory.Validate() != nil ||
!sameWindowsPrivateDirectorySnapshot(before, after) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
sort.Slice(selected, func(left, right int) bool { return selected[left].Name < selected[right].Name })
more := len(selected) > maximumEntries
if more {
selected = selected[:maximumEntries]
}
return PrivateDirectoryPage{Entries: selected, More: more}, nil
}
func sameWindowsRelativeClaim(source, claim *windowsPrivateRegularAt) bool {
return source != nil && claim != nil && sameWindowsPrivateFile(source.info, claim.info)
}
func windowsRelativeClaimPairExists(directory *windowsPrivateDirectory, source, claim string) (bool, error) {
left, err := openWindowsPrivateRegularAt(directory.handle, source, windows.FILE_GENERIC_READ, 2)
if isWindowsRelativeNotFound(err) {
return false, nil
}
if err != nil {
return false, ErrUnsafeFile
}
defer left.Close()
right, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 2)
if isWindowsRelativeNotFound(err) {
return false, nil
}
if err != nil {
return false, ErrUnsafeFile
}
defer right.Close()
return sameWindowsRelativeClaim(left, right), nil
}
func windowsRelativeClaimAbsentOrOrphan(directory *windowsPrivateDirectory, source, claim string) (bool, error) {
current, err := openWindowsPrivateRegularAtAllowedLinks(directory.handle, source, windows.FILE_GENERIC_READ, 1, 2)
if err == nil {
_ = current.Close()
return false, nil
}
if !isWindowsRelativeNotFound(err) {
return false, ErrUnsafeFile
}
orphan, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 1)
if err == nil {
_ = orphan.Close()
return true, nil
}
if isWindowsRelativeNotFound(err) {
return true, nil
}
return false, ErrUnsafeFile
}
func (directory *windowsPrivateDirectory) ClaimRegular(source, claim string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) {
return false, ErrUnsafeFile
}
value, err := openWindowsPrivateRegularAt(
directory.handle,
source,
windows.FILE_GENERIC_READ|windows.FILE_GENERIC_WRITE|windows.DELETE,
1,
)
if err != nil {
pair, pairErr := windowsRelativeClaimPairExists(directory, source, claim)
if pairErr == nil && pair {
return false, nil
}
orphan, orphanErr := windowsRelativeClaimAbsentOrOrphan(directory, source, claim)
if orphanErr == nil && orphan {
return false, nil
}
return false, ErrUnsafeFile
}
defer value.Close()
if linkWindowsPrivateRegularAt(value.handle, directory.handle, claim) != nil {
existing, existingErr := openWindowsPrivateRegularAtAllowedLinks(directory.handle, claim, windows.FILE_GENERIC_READ, 1, 2)
if existingErr == nil {
_ = existing.Close()
return false, nil
}
return false, ErrUnsafeFile
}
after, err := privateWindowsRegularInfo(value.handle, 2)
if err != nil {
return false, ErrUnsafeFile
}
value.info = after
claimed, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 2)
if err != nil {
return false, ErrUnsafeFile
}
defer claimed.Close()
if !sameWindowsRelativeClaim(value, claimed) || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func (directory *windowsPrivateDirectory) ReadClaim(source, claim string, maximum int64) ([]byte, bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) ||
maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, false, ErrUnsafeFile
}
value, err := openWindowsPrivateRegularAt(directory.handle, source, windows.FILE_GENERIC_READ, 2)
if err != nil {
orphan, orphanErr := windowsRelativeClaimAbsentOrOrphan(directory, source, claim)
if orphanErr == nil && orphan {
return nil, false, nil
}
return nil, false, ErrUnsafeFile
}
defer value.Close()
claimed, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 2)
if isWindowsRelativeNotFound(err) {
return nil, false, nil
}
if err != nil {
return nil, false, ErrUnsafeFile
}
defer claimed.Close()
if !sameWindowsRelativeClaim(value, claimed) {
return nil, false, ErrUnsafeFile
}
contents, err := readWindowsPrivateRegular(value, maximum, 2)
if err != nil {
return nil, false, ErrUnsafeFile
}
afterClaim, err := privateWindowsRegularInfo(claimed.handle, 2)
if err != nil || !sameWindowsPrivateFile(value.info, afterClaim) || directory.Validate() != nil {
return nil, false, ErrUnsafeFile
}
return contents, true, nil
}
func (directory *windowsPrivateDirectory) RemoveClaim(source, claim string) (bool, error) {
if directory.Validate() != nil || !validPrivateLeafName(source) || !validPrivateLeafName(claim) {
return false, ErrUnsafeFile
}
value, err := openWindowsPrivateRegularAt(directory.handle, source, windows.FILE_GENERIC_READ|windows.DELETE, 2)
if err != nil {
orphan, orphanErr := windowsRelativeClaimAbsentOrOrphan(directory, source, claim)
if orphanErr == nil && orphan {
return false, nil
}
return false, ErrUnsafeFile
}
claimed, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ, 2)
if isWindowsRelativeNotFound(err) {
_ = value.Close()
return false, nil
}
if err != nil || !sameWindowsRelativeClaim(value, claimed) {
_ = value.Close()
if claimed != nil {
_ = claimed.Close()
}
return false, ErrUnsafeFile
}
if claimed.Close() != nil || closeAndDeleteWindowsPrivateRegular(value) != nil {
return false, ErrUnsafeFile
}
remaining, err := openWindowsPrivateRegularAt(directory.handle, claim, windows.FILE_GENERIC_READ|windows.DELETE, 1)
if err != nil || closeAndDeleteWindowsPrivateRegular(remaining) != nil || directory.Validate() != nil {
return false, ErrUnsafeFile
}
return true, nil
}