fix(auth): add Windows session storage bridge

This commit is contained in:
2026-08-16 21:43:03 +02:00
parent 6bf8218fea
commit c9b02fc57e
14 changed files with 1630 additions and 8 deletions
+4
View File
@@ -15,6 +15,7 @@ import (
"strings"
"github.com/aritmolab/thothii/tools/tht/internal/authconfig"
"github.com/aritmolab/thothii/tools/tht/internal/authstorage"
"github.com/aritmolab/thothii/tools/tht/internal/backup"
"github.com/aritmolab/thothii/tools/tht/internal/compose"
"github.com/aritmolab/thothii/tools/tht/internal/config"
@@ -99,6 +100,9 @@ func main() {
}
func run(ctx context.Context, args []string, stdout, stderr io.Writer) int {
if len(args) > 0 && args[0] == "_auth-storage" {
return authstorage.Run(ctx, args[1:], os.Stdin, stdout, stderr)
}
if len(args) == 1 && (args[0] == "--help" || args[0] == "-h") {
fmt.Fprint(stdout, usage)
return 0
+291
View File
@@ -0,0 +1,291 @@
// Package authstorage implements tht's hidden, stdin/stdout-only bridge for protected browser
// session files. It deliberately has no operator-facing commands or installation configuration.
package authstorage
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"io"
"os"
"path/filepath"
"regexp"
"strings"
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
)
const (
protocolVersion = 1
maximumProtocolBytes = 64 * 1024
maximumSessionBytes = 16 * 1024
maximumOIDCStateBytes = 8 * 1024
maximumEntries = 256
)
var (
digestFilename = regexp.MustCompile(`^[a-f0-9]{64}\.json$`)
claimFilename = regexp.MustCompile(`^[a-f0-9]{64}\.claim$`)
errInvalid = errors.New("auth storage request invalid")
)
type request struct {
Version int `json:"version"`
Operation string `json:"operation"`
Root string `json:"root"`
Directory string `json:"directory"`
Filename string `json:"filename,omitempty"`
ContentBase64 string `json:"contentBase64,omitempty"`
}
type response struct {
Version int `json:"version"`
OK bool `json:"ok"`
Created bool `json:"created,omitempty"`
Replaced bool `json:"replaced,omitempty"`
Removed bool `json:"removed,omitempty"`
Found bool `json:"found,omitempty"`
Claimed bool `json:"claimed,omitempty"`
ContentBase64 string `json:"contentBase64,omitempty"`
Entries []safeio.PrivateDirectoryEntry `json:"entries,omitempty"`
}
// Run accepts exactly one strict JSON request on stdin and emits exactly one JSON response on
// stdout. All diagnostic text is fixed and goes only to stderr.
func Run(ctx context.Context, args []string, stdin io.Reader, stdout, stderr io.Writer) int {
if len(args) != 0 || ctx == nil || stdin == nil || stdout == nil || stderr == nil {
return fail(stderr)
}
if err := ctx.Err(); err != nil {
return fail(stderr)
}
decoder := json.NewDecoder(io.LimitReader(stdin, maximumProtocolBytes+1))
decoder.DisallowUnknownFields()
var input request
if err := decoder.Decode(&input); err != nil {
return fail(stderr)
}
var extra any
if err := decoder.Decode(&extra); err != io.EOF {
return fail(stderr)
}
if err := ctx.Err(); err != nil {
return fail(stderr)
}
result, err := execute(input)
if err != nil || ctx.Err() != nil {
return fail(stderr)
}
encoder := json.NewEncoder(stdout)
encoder.SetEscapeHTML(false)
if err := encoder.Encode(result); err != nil {
return fail(stderr)
}
return 0
}
func fail(stderr io.Writer) int {
_, _ = io.WriteString(stderr, "tht: auth storage request failed\n")
return 1
}
func execute(input request) (response, error) {
if input.Version != protocolVersion || !validDirectory(input.Directory) {
return response{}, errInvalid
}
directory, err := storageDirectory(input.Root, input.Directory)
if err != nil {
return response{}, errInvalid
}
switch input.Operation {
case "create":
contents, err := decodeContents(input)
if err != nil || !digestFilename.MatchString(input.Filename) {
return response{}, errInvalid
}
created, err := createPrivate(directory, input.Filename, contents)
if err != nil {
return response{}, errInvalid
}
return response{Version: protocolVersion, OK: true, Created: created}, nil
case "read":
if !digestFilename.MatchString(input.Filename) {
return response{}, errInvalid
}
contents, found, err := readPrivate(directory, input.Filename, recordMaximum(input.Directory))
if err != nil {
return response{}, errInvalid
}
return contentResponse(found, contents), nil
case "replace":
contents, err := decodeContents(input)
if err != nil || !digestFilename.MatchString(input.Filename) {
return response{}, errInvalid
}
if err := safeio.ReplaceCanonicalRegular(filepath.Join(directory, input.Filename), contents, 0o600); err != nil {
return response{}, errInvalid
}
return response{Version: protocolVersion, OK: true, Replaced: true}, nil
case "remove":
if !digestFilename.MatchString(input.Filename) && !claimFilename.MatchString(input.Filename) {
return response{}, errInvalid
}
removed, err := removePrivate(directory, input.Filename)
if err != nil {
return response{}, errInvalid
}
return response{Version: protocolVersion, OK: true, Removed: removed}, nil
case "list":
entries, err := safeio.ListCanonicalPrivateDirectory(directory, maximumEntries)
if err != nil {
return response{}, errInvalid
}
for _, entry := range entries {
if !digestFilename.MatchString(entry.Name) && !claimFilename.MatchString(entry.Name) {
return response{}, errInvalid
}
}
return response{Version: protocolVersion, OK: true, Entries: entries}, nil
case "claim-consume":
if input.Directory != "oidc" || !digestFilename.MatchString(input.Filename) {
return response{}, errInvalid
}
return claimConsume(directory, input.Filename)
case "read-claim":
if input.Directory != "oidc" || !digestFilename.MatchString(input.Filename) {
return response{}, errInvalid
}
contents, found, err := safeio.ReadCanonicalPrivateClaim(
filepath.Join(directory, input.Filename),
filepath.Join(directory, asClaimFilename(input.Filename)),
maximumOIDCStateBytes,
)
if err != nil {
return response{}, errInvalid
}
return contentResponse(found, contents), nil
case "remove-claim":
if input.Directory != "oidc" || !digestFilename.MatchString(input.Filename) {
return response{}, errInvalid
}
removed, err := safeio.RemoveCanonicalPrivateClaim(
filepath.Join(directory, input.Filename),
filepath.Join(directory, asClaimFilename(input.Filename)),
)
if err != nil {
return response{}, errInvalid
}
return response{Version: protocolVersion, OK: true, Removed: removed}, nil
default:
return response{}, errInvalid
}
}
func contentResponse(found bool, contents []byte) response {
if !found {
return response{Version: protocolVersion, OK: true}
}
return response{Version: protocolVersion, OK: true, Found: true, ContentBase64: base64.StdEncoding.EncodeToString(contents)}
}
func storageDirectory(root, directory string) (string, error) {
if !filepath.IsAbs(root) || filepath.Clean(root) != root || strings.ContainsRune(root, '\x00') || 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"
}
func recordMaximum(directory string) int64 {
if directory == "oidc" {
return maximumOIDCStateBytes
}
return maximumSessionBytes
}
func decodeContents(input request) ([]byte, error) {
contents, err := base64.StdEncoding.DecodeString(input.ContentBase64)
if err != nil || len(contents) == 0 || int64(len(contents)) > recordMaximum(input.Directory) {
return nil, errInvalid
}
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) {
return false, errInvalid
}
if err := safeio.WriteCanonicalNewFile(path, contents, 0o600); err == nil {
return true, nil
}
if safeio.ValidatePrivateRegular(path) == nil {
return false, nil
}
return false, errInvalid
}
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)
if err != nil {
return nil, false, errInvalid
}
return contents, true, 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 {
return false, errInvalid
}
if err := safeio.RemoveCanonicalPrivateRegular(path); err != nil {
return false, errInvalid
}
return true, 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)
if err != nil {
return response{}, errInvalid
}
if !claimed {
return response{Version: protocolVersion, OK: true}, nil
}
contents, found, err := safeio.ReadCanonicalPrivateClaim(source, claim, maximumOIDCStateBytes)
if err != nil || !found {
return response{}, errInvalid
}
removed, err := safeio.RemoveCanonicalPrivateClaim(source, claim)
if err != nil || !removed {
return response{}, errInvalid
}
return contentResponse(true, contents), nil
}
func asClaimFilename(filename string) string {
return strings.TrimSuffix(filename, ".json") + ".claim"
}
@@ -0,0 +1,164 @@
package authstorage
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"os"
"path/filepath"
"sync"
"testing"
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
"github.com/aritmolab/thothii/tools/tht/internal/testsupport"
)
func TestProtocolCreatesReadsReplacesListsAndRemovesPrivateRecord(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
filename := "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa.json"
created := runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("first"))})
if !created.Created {
t.Fatal("create did not report a new record")
}
if duplicate := runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("other"))}); duplicate.Created {
t.Fatal("duplicate exclusive create reported success")
}
read := runRequest(t, request{Version: 1, Operation: "read", Root: root, Directory: "sessions", Filename: filename})
if !read.Found || decodeContent(t, read) != "first" {
t.Fatalf("read = %#v, want private first record", read)
}
updated := runRequest(t, request{Version: 1, Operation: "replace", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("second"))})
if !updated.Replaced {
t.Fatal("replace did not report success")
}
listed := runRequest(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions"})
if len(listed.Entries) != 1 || listed.Entries[0].Name != filename {
t.Fatalf("list = %#v, want exactly %q", listed.Entries, filename)
}
removed := runRequest(t, request{Version: 1, Operation: "remove", Root: root, Directory: "sessions", Filename: filename})
if !removed.Removed {
t.Fatal("remove did not report success")
}
}
func TestProtocolClaimConsumeIsAtomicAcrossConcurrentRequests(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
filename := "dddddddddddddddddddddddddddddddddddddddddddddddddddddddddddddddd.json"
runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "oidc", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("oidc-record"))})
request := request{Version: 1, Operation: "claim-consume", Root: root, Directory: "oidc", Filename: filename}
results := make(chan response, 2)
errors := make(chan error, 2)
var group sync.WaitGroup
for range 2 {
group.Add(1)
go func() {
defer group.Done()
result, err := execute(request)
if err != nil {
errors <- err
return
}
results <- result
}()
}
group.Wait()
close(results)
close(errors)
for err := range errors {
t.Fatalf("concurrent claim error = %v", err)
}
found := 0
for result := range results {
if result.Found {
found++
}
}
if found != 1 {
t.Fatalf("winning claim count = %d, want 1", found)
}
}
func TestProtocolRejectsBoundsReparseAndUnexpectedStorageNames(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
filename := "eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee.json"
runRejected(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString(make([]byte, maximumSessionBytes+1))})
valid := request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))}
runRequest(t, valid)
if err := safeio.WriteCanonicalNewFile(filepath.Join(root, "sessions", "unexpected.txt"), []byte("junk"), 0o600); err != nil {
t.Fatal(err)
}
runRejected(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions"})
outer := privateTestRoot(t)
linkedRoot := filepath.Join(outer, "linked-auth")
testsupport.SymlinkOrSkip(t, filepath.Join(outer, "missing-real-auth"), linkedRoot)
runRejected(t, request{Version: 1, Operation: "create", Root: linkedRoot, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))})
}
func TestProtocolClaimConsumeHasOneWinner(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
filename := "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb.json"
runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "oidc", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("oidc-record"))})
first := runRequest(t, request{Version: 1, Operation: "claim-consume", Root: root, Directory: "oidc", Filename: filename})
second := runRequest(t, request{Version: 1, Operation: "claim-consume", Root: root, Directory: "oidc", Filename: filename})
if !first.Found || decodeContent(t, first) != "oidc-record" || second.Found {
t.Fatalf("claim results = %#v, %#v", first, second)
}
}
func runRequest(t *testing.T, value request) response {
t.Helper()
input, err := json.Marshal(value)
if err != nil {
t.Fatal(err)
}
var stdout, stderr bytes.Buffer
if code := Run(context.Background(), nil, bytes.NewReader(input), &stdout, &stderr); code != 0 {
t.Fatalf("Run() code = %d stderr = %q", code, stderr.String())
}
var output response
if err := json.Unmarshal(stdout.Bytes(), &output); err != nil || !output.OK || output.Version != 1 {
t.Fatalf("stdout = %q response = %#v error = %v", stdout.String(), output, err)
}
return output
}
func runRejected(t *testing.T, value request) {
t.Helper()
input, err := json.Marshal(value)
if err != nil {
t.Fatal(err)
}
var stdout, stderr bytes.Buffer
if code := Run(context.Background(), nil, bytes.NewReader(input), &stdout, &stderr); code == 0 || stdout.Len() != 0 || stderr.String() != "tht: auth storage request failed\n" {
t.Fatalf("Run() rejection code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String())
}
}
func decodeContent(t *testing.T, value response) string {
t.Helper()
contents, err := base64.StdEncoding.DecodeString(value.ContentBase64)
if err != nil {
t.Fatal(err)
}
return string(contents)
}
func privateTestRoot(t *testing.T) string {
t.Helper()
temporary, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporary, "tht-authstorage-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
return root
}
@@ -0,0 +1,71 @@
//go:build windows
package authstorage
import (
"encoding/base64"
"os"
"path/filepath"
"runtime"
"testing"
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
"golang.org/x/sys/windows"
)
func TestProtocolCreatesRecordsWithOwnerOnlyDACLAndRejectsReparseRoot(t *testing.T) {
root := filepath.Join(t.TempDir(), "auth")
filename := "cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc.json"
runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("private"))})
if err := safeio.ValidatePrivateDirectory(root); err != nil {
t.Fatalf("root DACL = %v", err)
}
if err := safeio.ValidatePrivateDirectory(filepath.Join(root, "sessions")); err != nil {
t.Fatalf("sessions DACL = %v", err)
}
if err := safeio.ValidatePrivateRegular(filepath.Join(root, "sessions", filename)); err != nil {
t.Fatalf("record DACL = %v", err)
}
}
func TestProtocolRejectsPermissiveDACLAndReparseRoot(t *testing.T) {
root := filepath.Join(t.TempDir(), "auth")
filename := "ffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff.json"
runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("private"))})
path := filepath.Join(root, "sessions", filename)
if err := setPermissiveDACL(path); err != nil {
t.Fatal(err)
}
runRejected(t, request{Version: 1, Operation: "read", Root: root, Directory: "sessions", Filename: filename})
linkedRoot := filepath.Join(t.TempDir(), "reparse-auth")
if err := os.Symlink(root, linkedRoot); err != nil {
t.Skipf("Windows host does not permit test symlink creation: %v", err)
}
runRejected(t, request{Version: 1, Operation: "create", Root: linkedRoot, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("private"))})
}
func setPermissiveDACL(path string) error {
world, err := windows.StringToSid("S-1-1-0")
if err != nil {
return err
}
var pinner runtime.Pinner
pinner.Pin(world)
defer pinner.Unpin()
acl, err := windows.ACLFromEntries([]windows.EXPLICIT_ACCESS{{
AccessPermissions: windows.GENERIC_READ | windows.GENERIC_WRITE,
AccessMode: windows.GRANT_ACCESS,
Trustee: windows.TRUSTEE{
TrusteeForm: windows.TRUSTEE_IS_SID,
TrusteeType: windows.TRUSTEE_IS_GROUP,
TrusteeValue: windows.TrusteeValueFromSID(world),
},
}}, nil)
if err != nil {
return err
}
return windows.SetNamedSecurityInfo(path, windows.SE_FILE_OBJECT,
windows.DACL_SECURITY_INFORMATION|windows.PROTECTED_DACL_SECURITY_INFORMATION,
nil, nil, acl, nil)
}
+169
View File
@@ -0,0 +1,169 @@
//go:build !windows
package safeio
import (
"errors"
"io"
"os"
"golang.org/x/sys/unix"
)
func claimCanonicalPrivateRegular(source, claim string) (bool, error) {
directory, sourceName, err := openCanonicalParentDirectory(source)
if err != nil {
return false, ErrUnsafeFile
}
defer unix.Close(directory)
claimName := claim[len(claim)-len(claimBaseName(claim)):]
if err := requirePrivateRegularAtLinks(directory, sourceName, 1); err != nil {
if requireSamePrivatePairAt(directory, sourceName, claimName) == nil {
return false, nil
}
if isPrivateClaimAbsentOrOrphanAt(directory, sourceName, claimName) {
return false, nil
}
return false, ErrUnsafeFile
}
if err := unix.Linkat(directory, sourceName, directory, claimName, 0); err != nil {
if errors.Is(err, unix.EEXIST) && privateRegularAtAllowedLinks(directory, claimName, 1, 2) == nil {
return false, nil
}
return false, ErrUnsafeFile
}
if err := requireSamePrivatePairAt(directory, sourceName, claimName); err != nil {
return false, ErrUnsafeFile
}
if err := unix.Fsync(directory); err != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func readCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) {
directory, sourceName, err := openCanonicalParentDirectory(source)
if err != nil {
return nil, false, ErrUnsafeFile
}
defer unix.Close(directory)
claimName := claim[len(claim)-len(claimBaseName(claim)):]
if err := requireSamePrivatePairAt(directory, sourceName, claimName); err != nil {
if isPrivateClaimAbsentOrOrphanAt(directory, sourceName, claimName) {
return nil, false, nil
}
return nil, false, ErrUnsafeFile
}
descriptor, err := unix.Openat(directory, sourceName, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0)
if err != nil {
return nil, false, ErrUnsafeFile
}
file := os.NewFile(uintptr(descriptor), "tht-safeio-oidc-claim")
if file == nil {
unix.Close(descriptor)
return nil, false, ErrUnsafeFile
}
defer file.Close()
var before unix.Stat_t
if err := unix.Fstat(descriptor, &before); err != nil || !privateRegularStatWithLinks(before, 2) {
return nil, false, ErrUnsafeFile
}
contents, err := io.ReadAll(io.LimitReader(file, maximum+1))
if err != nil || int64(len(contents)) > maximum {
return nil, false, ErrUnsafeFile
}
var after unix.Stat_t
if err := unix.Fstat(descriptor, &after); err != nil || !sameUnixPrivateFile(before, after) || !privateRegularStatWithLinks(after, 2) {
return nil, false, ErrUnsafeFile
}
if err := requireSamePrivatePairAt(directory, sourceName, claimName); err != nil {
return nil, false, ErrUnsafeFile
}
return contents, true, nil
}
func removeCanonicalPrivateClaim(source, claim string) (bool, error) {
directory, sourceName, err := openCanonicalParentDirectory(source)
if err != nil {
return false, ErrUnsafeFile
}
defer unix.Close(directory)
claimName := claim[len(claim)-len(claimBaseName(claim)):]
if err := requireSamePrivatePairAt(directory, sourceName, claimName); err != nil {
if isPrivateClaimAbsentOrOrphanAt(directory, sourceName, claimName) {
return false, nil
}
return false, ErrUnsafeFile
}
if err := unix.Unlinkat(directory, sourceName, 0); err != nil {
return false, ErrUnsafeFile
}
if err := requirePrivateRegularAtLinks(directory, claimName, 1); err != nil {
return false, ErrUnsafeFile
}
if err := unix.Unlinkat(directory, claimName, 0); err != nil {
return false, ErrUnsafeFile
}
if err := unix.Fsync(directory); err != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func claimBaseName(path string) string {
for index := len(path) - 1; index >= 0; index-- {
if path[index] == byte(os.PathSeparator) {
return path[index+1:]
}
}
return path
}
func privateRegularAtAllowedLinks(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 privateRegularStatWithLinks(stat, links) {
return nil
}
}
return ErrUnsafeFile
}
func requirePrivateRegularAtLinks(directory int, name string, links uint64) error {
return privateRegularAtAllowedLinks(directory, name, links)
}
func requireSamePrivatePairAt(directory int, source, claim string) error {
var left, right unix.Stat_t
if err := unix.Fstatat(directory, source, &left, unix.AT_SYMLINK_NOFOLLOW); err != nil || !privateRegularStatWithLinks(left, 2) {
return ErrUnsafeFile
}
if err := unix.Fstatat(directory, claim, &right, unix.AT_SYMLINK_NOFOLLOW); err != nil || !privateRegularStatWithLinks(right, 2) {
return ErrUnsafeFile
}
if !sameUnixPrivateFile(left, right) {
return ErrUnsafeFile
}
return nil
}
func isPrivateClaimAbsentOrOrphanAt(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 && privateRegularStatWithLinks(claimStat, 1)
}
func privateRegularStatWithLinks(stat unix.Stat_t, links uint64) bool {
return stat.Mode&unix.S_IFMT == unix.S_IFREG && uint64(stat.Nlink) == links && stat.Mode&0o7777 == 0o600
}
func sameUnixPrivateFile(left, right unix.Stat_t) bool {
return left.Dev == right.Dev && left.Ino == right.Ino && left.Size == right.Size && left.Mtim == right.Mtim
}
+210
View File
@@ -0,0 +1,210 @@
//go:build windows
package safeio
import (
"errors"
"io"
"os"
"path/filepath"
"golang.org/x/sys/windows"
)
type windowsPrivateRegular struct {
parents *windowsParentHandles
handle windows.Handle
path string
info windows.ByHandleFileInformation
}
func (value *windowsPrivateRegular) Close() {
if value.handle != 0 {
_ = windows.CloseHandle(value.handle)
value.handle = 0
}
if value.parents != nil {
value.parents.Close()
value.parents = nil
}
}
func openWindowsPrivateRegular(path string, links uint32) (*windowsPrivateRegular, error) {
parents, target, err := openCanonicalWindowsParent(path)
if err != nil || len(parents.handles) == 0 || validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil {
if parents != nil {
parents.Close()
}
return nil, ErrUnsafeFile
}
fullPath := filepath.Join(parents.directory, target)
handle, err := windows.CreateFile(
windows.StringToUTF16Ptr(fullPath),
windows.GENERIC_READ,
windowsRetainedHandleShareMode,
nil,
windows.OPEN_EXISTING,
windows.FILE_FLAG_OPEN_REPARSE_POINT|windows.FILE_ATTRIBUTE_NORMAL,
0,
)
if err != nil {
parents.Close()
return nil, err
}
var info windows.ByHandleFileInformation
if err := windows.GetFileInformationByHandle(handle, &info); err != nil || info.FileAttributes&windows.FILE_ATTRIBUTE_REPARSE_POINT != 0 || info.FileAttributes&windows.FILE_ATTRIBUTE_DIRECTORY != 0 || info.NumberOfLinks != links || validateOwnerOnlyDACL(handle) != nil {
_ = windows.CloseHandle(handle)
parents.Close()
return nil, ErrUnsafeFile
}
return &windowsPrivateRegular{parents: parents, handle: handle, path: fullPath, info: info}, nil
}
func claimCanonicalPrivateRegular(source, claim string) (bool, error) {
sourceFile, err := openWindowsPrivateRegular(source, 1)
if err != nil {
if windowsPrivateClaimPairExists(source, claim) {
return false, nil
}
if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) {
return false, nil
}
return false, ErrUnsafeFile
}
defer sourceFile.Close()
if err := windows.CreateHardLink(windows.StringToUTF16Ptr(claim), windows.StringToUTF16Ptr(sourceFile.path), 0); err != nil {
if errors.Is(err, windows.ERROR_FILE_EXISTS) || errors.Is(err, windows.ERROR_ALREADY_EXISTS) {
if existing, existingErr := openWindowsPrivateRegularWithAllowedLinks(claim, 1, 2); existingErr == nil {
existing.Close()
return false, nil
}
}
return false, ErrUnsafeFile
}
if err := windows.GetFileInformationByHandle(sourceFile.handle, &sourceFile.info); err != nil || sourceFile.info.NumberOfLinks != 2 {
return false, ErrUnsafeFile
}
claimFile, err := openWindowsPrivateRegular(claim, 2)
if err != nil {
return false, ErrUnsafeFile
}
defer claimFile.Close()
if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) {
return false, ErrUnsafeFile
}
return true, nil
}
func readCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) {
sourceFile, err := openWindowsPrivateRegular(source, 2)
if err != nil {
if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) {
return nil, false, nil
}
return nil, false, ErrUnsafeFile
}
defer sourceFile.Close()
claimFile, err := openWindowsPrivateRegular(claim, 2)
if err != nil {
if isWindowsNotFound(err) {
return nil, false, nil
}
return nil, false, ErrUnsafeFile
}
defer claimFile.Close()
if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) {
return nil, false, ErrUnsafeFile
}
file := os.NewFile(uintptr(sourceFile.handle), "tht-safeio-oidc-claim")
if file == nil {
return nil, false, ErrUnsafeFile
}
contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1))
file.Close()
sourceFile.handle = 0
if readErr != nil || int64(len(contents)) > maximum {
return nil, false, ErrUnsafeFile
}
if err := windows.GetFileInformationByHandle(claimFile.handle, &claimFile.info); err != nil || claimFile.info.NumberOfLinks != 2 || !sameWindowsPrivateFile(sourceFile.info, claimFile.info) {
return nil, false, ErrUnsafeFile
}
return contents, true, nil
}
func removeCanonicalPrivateClaim(source, claim string) (bool, error) {
sourceFile, err := openWindowsPrivateRegular(source, 2)
if err != nil {
if isWindowsNotFound(err) && windowsPrivateClaimAbsentOrOrphan(claim) {
return false, nil
}
return false, ErrUnsafeFile
}
claimFile, err := openWindowsPrivateRegular(claim, 2)
if err != nil {
sourceFile.Close()
if isWindowsNotFound(err) {
return false, nil
}
return false, ErrUnsafeFile
}
if !sameWindowsPrivateFile(sourceFile.info, claimFile.info) {
sourceFile.Close()
claimFile.Close()
return false, ErrUnsafeFile
}
sourceFile.Close()
claimFile.Close()
if err := windows.DeleteFile(windows.StringToUTF16Ptr(source)); err != nil {
return false, ErrUnsafeFile
}
remaining, err := openWindowsPrivateRegular(claim, 1)
if err != nil {
return false, ErrUnsafeFile
}
remaining.Close()
if err := windows.DeleteFile(windows.StringToUTF16Ptr(claim)); err != nil {
return false, ErrUnsafeFile
}
return true, nil
}
func openWindowsPrivateRegularWithAllowedLinks(path string, allowed ...uint32) (*windowsPrivateRegular, error) {
for _, links := range allowed {
value, err := openWindowsPrivateRegular(path, links)
if err == nil {
return value, nil
}
}
return nil, ErrUnsafeFile
}
func windowsPrivateClaimAbsentOrOrphan(claim string) bool {
claimFile, err := openWindowsPrivateRegular(claim, 1)
if err != nil {
return isWindowsNotFound(err)
}
claimFile.Close()
return true
}
func windowsPrivateClaimPairExists(source, claim string) bool {
sourceFile, sourceErr := openWindowsPrivateRegular(source, 2)
if sourceErr != nil {
return false
}
defer sourceFile.Close()
claimFile, claimErr := openWindowsPrivateRegular(claim, 2)
if claimErr != nil {
return false
}
defer claimFile.Close()
return sameWindowsPrivateFile(sourceFile.info, claimFile.info)
}
func isWindowsNotFound(err error) bool {
return errors.Is(err, windows.ERROR_FILE_NOT_FOUND) || errors.Is(err, windows.ERROR_PATH_NOT_FOUND)
}
func sameWindowsPrivateFile(left, right windows.ByHandleFileInformation) bool {
return left.VolumeSerialNumber == right.VolumeSerialNumber && left.FileIndexHigh == right.FileIndexHigh && left.FileIndexLow == right.FileIndexLow && left.FileSizeHigh == right.FileSizeHigh && left.FileSizeLow == right.FileSizeLow && left.LastWriteTime == right.LastWriteTime
}
+80
View File
@@ -71,6 +71,50 @@ func ReadCanonicalUTF8(path string, maximum int64) (string, error) {
return string(contents), nil
}
// ReadCanonicalPrivateRegular reads one owner-only private record and revalidates its metadata
// after the bounded read. It intentionally rejects ordinary hard links; OIDC's explicitly named
// atomic claim pair uses ReadCanonicalPrivateClaim instead.
func ReadCanonicalPrivateRegular(path string, maximum int64) ([]byte, error) {
return readCanonicalPrivateRegular(path, maximum)
}
// PrivateDirectoryEntry is a bounded, untrusted directory listing item. Callers must still
// validate each filename and record before using it.
type PrivateDirectoryEntry struct {
Name string
ModifiedUnixMs int64
}
// ListCanonicalPrivateDirectory lists regular, non-symlinked direct children from an owner-only
// directory. It returns no content and bounds the number of entries before allocating output.
func ListCanonicalPrivateDirectory(path string, maximumEntries int) ([]PrivateDirectoryEntry, error) {
if maximumEntries < 1 || maximumEntries > 4096 {
return nil, ErrUnsafeFile
}
if err := ValidatePrivateDirectory(path); err != nil {
return nil, ErrUnsafeFile
}
entries, err := os.ReadDir(path)
if err != nil || len(entries) > maximumEntries {
return nil, ErrUnsafeFile
}
result := make([]PrivateDirectoryEntry, 0, len(entries))
for _, entry := range entries {
if entry.Name() == "" || strings.Contains(entry.Name(), string(filepath.Separator)) {
return nil, ErrUnsafeFile
}
info, err := os.Lstat(filepath.Join(path, entry.Name()))
if err != nil || !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 {
return nil, ErrUnsafeFile
}
result = append(result, PrivateDirectoryEntry{Name: entry.Name(), ModifiedUnixMs: info.ModTime().UnixMilli()})
}
if err := ValidatePrivateDirectory(path); err != nil {
return nil, ErrUnsafeFile
}
return result, nil
}
func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error {
if err := ValidateCanonicalPath(path); err != nil {
return err
@@ -127,6 +171,42 @@ func RemoveCanonicalPrivateRegular(path string) error {
return removeCanonicalPrivateRegular(path)
}
// ClaimCanonicalPrivateRegular atomically creates a second, explicit private hard link to one
// existing record. It is used only for digest-named, single-use OIDC state claims.
func ClaimCanonicalPrivateRegular(source, claim string) (bool, error) {
if err := validateClaimPaths(source, claim); err != nil {
return false, ErrUnsafeFile
}
return claimCanonicalPrivateRegular(source, claim)
}
// ReadCanonicalPrivateClaim reads a verified two-link source/claim pair. found=false means the
// state has already been consumed or a winning process is between its two removal steps.
func ReadCanonicalPrivateClaim(source, claim string, maximum int64) ([]byte, bool, error) {
if maximum < 0 || maximum == int64(^uint64(0)>>1) || validateClaimPaths(source, claim) != nil {
return nil, false, ErrUnsafeFile
}
return readCanonicalPrivateClaim(source, claim, maximum)
}
// RemoveCanonicalPrivateClaim removes exactly a verified two-link source/claim pair.
func RemoveCanonicalPrivateClaim(source, claim string) (bool, error) {
if err := validateClaimPaths(source, claim); err != nil {
return false, ErrUnsafeFile
}
return removeCanonicalPrivateClaim(source, claim)
}
func validateClaimPaths(source, claim string) error {
if ValidateCanonicalPath(source) != nil || ValidateCanonicalPath(claim) != nil || filepath.Dir(source) != filepath.Dir(claim) {
return ErrUnsafeFile
}
if err := ValidatePrivateDirectory(filepath.Dir(source)); err != nil {
return ErrUnsafeFile
}
return nil
}
func randomTemporaryName() (string, error) {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
@@ -0,0 +1,17 @@
//go:build !windows
package safeio
func readCanonicalPrivateRegular(path string, maximum int64) ([]byte, error) {
if err := ValidatePrivateRegular(path); err != nil {
return nil, ErrUnsafeFile
}
contents, err := ReadCanonicalRegular(path, maximum)
if err != nil {
return nil, ErrUnsafeFile
}
if err := ValidatePrivateRegular(path); err != nil {
return nil, ErrUnsafeFile
}
return contents, nil
}
@@ -0,0 +1,44 @@
//go:build windows
package safeio
import (
"io"
"os"
"golang.org/x/sys/windows"
)
func readCanonicalPrivateRegular(path string, maximum int64) ([]byte, error) {
if maximum < 0 || maximum == int64(^uint64(0)>>1) {
return nil, ErrUnsafeFile
}
value, err := openWindowsPrivateRegular(path, 1)
if err != nil {
return nil, ErrUnsafeFile
}
defer value.Close()
file := os.NewFile(uintptr(value.handle), "tht-safeio-private-read")
if file == nil {
return nil, ErrUnsafeFile
}
before := value.info
contents, readErr := io.ReadAll(io.LimitReader(file, maximum+1))
var after = value.info
identityErr := windows.GetFileInformationByHandle(value.handle, &after)
daclErr := validateOwnerOnlyDACL(value.handle)
closeErr := file.Close()
value.handle = 0
if readErr != nil || identityErr != nil || daclErr != nil || !sameWindowsPrivateFile(before, after) || after.NumberOfLinks != 1 || closeErr != nil || int64(len(contents)) > maximum {
return nil, ErrUnsafeFile
}
current, err := openWindowsPrivateRegular(path, 1)
if err != nil {
return nil, ErrUnsafeFile
}
current.Close()
if !sameWindowsPrivateFile(before, current.info) {
return nil, ErrUnsafeFile
}
return contents, nil
}
+9 -1
View File
@@ -96,6 +96,14 @@ func ProtectPrivateRegular(path string) error {
// createCanonicalNewPrivateFile installs the owner-only protected DACL in the CreateFile call, so
// another mutation can never observe a newly-created lock with an inherited/default DACL.
func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, error) {
parents, target, err := openCanonicalWindowsParent(path)
if err != nil || len(parents.handles) == 0 || validateOwnerOnlyDACL(parents.handles[len(parents.handles)-1]) != nil {
if parents != nil {
parents.Close()
}
return nil, ErrUnsafeFile
}
defer parents.Close()
security, err := newOwnerOnlySecurityDescriptor()
if err != nil {
return nil, ErrUnsafeFile
@@ -106,7 +114,7 @@ func createCanonicalNewPrivateFile(path string, mode os.FileMode) (*os.File, err
SecurityDescriptor: security.descriptor,
}
handle, err := windows.CreateFile(
windows.StringToUTF16Ptr(path),
windows.StringToUTF16Ptr(filepath.Join(parents.directory, target)),
windows.GENERIC_WRITE,
windowsRetainedHandleShareMode,
attributes,