fix(auth): paginate session maintenance safely

This commit is contained in:
2026-08-17 09:05:20 +02:00
parent 3e7bb11313
commit c86f01e886
12 changed files with 1021 additions and 79 deletions
+24 -5
View File
@@ -41,6 +41,8 @@ type request struct {
Filename string `json:"filename,omitempty"`
ContentBase64 string `json:"contentBase64,omitempty"`
MaximumEntries int `json:"maximumEntries,omitempty"`
AfterName string `json:"afterName,omitempty"`
Continuation bool `json:"continuation,omitempty"`
}
type response struct {
@@ -53,6 +55,7 @@ type response struct {
Claimed bool `json:"claimed,omitempty"`
ContentBase64 string `json:"contentBase64,omitempty"`
Entries *[]safeio.PrivateDirectoryEntry `json:"entries,omitempty"`
More *bool `json:"more,omitempty"`
}
// Run accepts exactly one strict JSON request on stdin and emits exactly one JSON response on
@@ -143,6 +146,19 @@ func execute(input request) (response, error) {
if limit == 0 {
limit = defaultMaximumEntries
}
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
}
entries, err := safeio.ListCanonicalPrivateDirectory(directory, limit)
if err != nil {
return response{}, errInvalid
@@ -182,17 +198,20 @@ func execute(input request) (response, error) {
func validOperationShape(input request) bool {
noContents := input.ContentBase64 == ""
noMaximumEntries := input.MaximumEntries == 0
noAfterName := input.AfterName == ""
noContinuation := !input.Continuation
switch input.Operation {
case "create", "replace":
return noMaximumEntries && (digestFilename.MatchString(input.Filename) || (input.Operation == "create" && input.Directory == "oidc" && oidcSlotFilename.MatchString(input.Filename)))
return noMaximumEntries && noAfterName && noContinuation && (digestFilename.MatchString(input.Filename) || (input.Operation == "create" && input.Directory == "oidc" && oidcSlotFilename.MatchString(input.Filename)))
case "read":
return noContents && noMaximumEntries && (digestFilename.MatchString(input.Filename) || (input.Directory == "oidc" && oidcSlotFilename.MatchString(input.Filename)))
return noContents && noMaximumEntries && noAfterName && noContinuation && (digestFilename.MatchString(input.Filename) || (input.Directory == "oidc" && oidcSlotFilename.MatchString(input.Filename)))
case "remove":
return noContents && noMaximumEntries && (digestFilename.MatchString(input.Filename) || (input.Directory == "oidc" && (claimFilename.MatchString(input.Filename) || oidcSlotFilename.MatchString(input.Filename))))
return noContents && noMaximumEntries && noAfterName && noContinuation && (digestFilename.MatchString(input.Filename) || (input.Directory == "oidc" && (claimFilename.MatchString(input.Filename) || oidcSlotFilename.MatchString(input.Filename))))
case "list":
return input.Filename == "" && noContents && input.MaximumEntries >= 0 && input.MaximumEntries <= maximumEntries
return input.Filename == "" && noContents && input.MaximumEntries >= 0 && input.MaximumEntries <= maximumEntries &&
((noContinuation && noAfterName) || (input.Continuation && input.Directory == "sessions" && input.MaximumEntries >= 1 && (noAfterName || digestFilename.MatchString(input.AfterName))))
case "claim-consume", "read-claim", "remove-claim":
return input.Directory == "oidc" && noContents && noMaximumEntries && digestFilename.MatchString(input.Filename)
return input.Directory == "oidc" && noContents && noMaximumEntries && noAfterName && noContinuation && digestFilename.MatchString(input.Filename)
default:
return false
}
@@ -34,6 +34,34 @@ func TestProtocolListUsesACallerSuppliedBound(t *testing.T) {
runRejected(t, request{Version: 1, Operation: "read", Root: root, Directory: "sessions", Filename: filename, MaximumEntries: 1})
}
func TestProtocolListPaginatesOrdinarySessionRecords(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
for index := 0; index < 513; index++ {
filename := fmt.Sprintf("%064x.json", index)
runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))})
}
first := runRequest(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions", MaximumEntries: 512, Continuation: true})
if first.Entries == nil || len(*first.Entries) != 512 || first.More == nil || !*first.More {
t.Fatalf("first continuation page = %#v", first)
}
second := runRequest(t, request{
Version: 1,
Operation: "list",
Root: root,
Directory: "sessions",
MaximumEntries: 512,
Continuation: true,
AfterName: (*first.Entries)[511].Name,
})
if second.Entries == nil || len(*second.Entries) != 1 || second.More == nil || *second.More || (*second.Entries)[0].Name != fmt.Sprintf("%064x.json", 512) {
t.Fatalf("second continuation page = %#v", second)
}
runRejected(t, request{Version: 1, Operation: "list", Root: root, Directory: "oidc", MaximumEntries: 1, Continuation: true, AfterName: fmt.Sprintf("%064x.json", 0)})
runRejected(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions", MaximumEntries: 1, Continuation: true, AfterName: "../unsafe.json"})
}
func TestProtocolCreatesReadsReplacesListsAndRemovesPrivateRecord(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
filename := "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa.json"
@@ -215,6 +243,7 @@ func TestProtocolRejectsBoundsReparseAndUnexpectedStorageNames(t *testing.T) {
t.Fatal(err)
}
runRejected(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions"})
runRejected(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions", MaximumEntries: 1, Continuation: true})
outer := privateTestRoot(t)
linkedRoot := filepath.Join(outer, "linked-auth")
+124
View File
@@ -15,6 +15,11 @@ import (
var ErrUnsafeFile = errors.New("unsafe file")
// Pagination must never allocate or iterate indefinitely over a hostile private directory. This
// is deliberately independent from an auth record page: it permits many bounded pages while
// putting a fixed upper bound on one sorted keyset scan.
const maximumPrivateDirectoryPageScanEntries = 16384
// EnsurePrivateDirectory creates only the final canonical directory with the platform's
// owner-only protection, or validates an existing directory has that protection.
func EnsurePrivateDirectory(path string) error {
@@ -86,6 +91,28 @@ type PrivateDirectoryEntry struct {
ModifiedUnixMs int64 `json:"modifiedUnixMs"`
}
// PrivateDirectoryPage is one lexically ordered, keyset-paginated directory page. More is true
// only when the page is full and another validated entry follows AfterName.
type PrivateDirectoryPage struct {
Entries []PrivateDirectoryEntry
More bool
}
func privateDirectoryScanSnapshot(path string) (os.FileInfo, error) {
if err := ValidatePrivateDirectory(path); err != nil {
return nil, ErrUnsafeFile
}
info, err := os.Lstat(path)
if err != nil || !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
return nil, ErrUnsafeFile
}
return info, nil
}
func samePrivateDirectoryScanSnapshot(left, right os.FileInfo) bool {
return left != nil && right != nil && os.SameFile(left, right) && left.Mode() == right.Mode() && left.ModTime().Equal(right.ModTime())
}
// 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) {
@@ -125,6 +152,103 @@ func ListCanonicalPrivateDirectory(path string, maximumEntries int) ([]PrivateDi
return result, nil
}
// ListCanonicalPrivateDirectoryPage scans a private directory without materializing a complete
// output listing, validates every encountered child, and keeps only the next bounded lexical page.
// The caller supplies the filename grammar because safeio is intentionally storage-format agnostic.
func ListCanonicalPrivateDirectoryPage(
path string,
maximumEntries int,
afterName string,
validName func(string) bool,
) (PrivateDirectoryPage, error) {
if maximumEntries < 1 || maximumEntries > 4096 || validName == nil || (afterName != "" && !validName(afterName)) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
before, err := privateDirectoryScanSnapshot(path)
if err != nil {
return PrivateDirectoryPage{}, err
}
directory, err := os.Open(path)
if err != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
defer directory.Close()
seen := make(map[string]struct{}, maximumEntries+1)
selected := make([]PrivateDirectoryEntry, 0, maximumEntries+1)
scanned := 0
for {
entries, readErr := directory.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 name == "" || strings.Contains(name, string(filepath.Separator)) || !validName(name) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if _, duplicate := seen[name]; duplicate {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
seen[name] = struct{}{}
entryPath := filepath.Join(path, name)
if err := ValidatePrivateRegular(entryPath); err != nil {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
info, err := os.Lstat(entryPath)
if err != nil || !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || !hasSingleLink(info) {
return PrivateDirectoryPage{}, ErrUnsafeFile
}
if name > afterName {
selected = appendBoundedPrivateDirectoryEntry(selected, PrivateDirectoryEntry{
Name: name, ModifiedUnixMs: info.ModTime().UnixMilli(),
}, maximumEntries+1)
}
if errors.Is(readErr, io.EOF) {
break
}
}
after, err := privateDirectoryScanSnapshot(path)
if err != nil || !samePrivateDirectoryScanSnapshot(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 appendBoundedPrivateDirectoryEntry(
entries []PrivateDirectoryEntry,
entry PrivateDirectoryEntry,
maximumEntries int,
) []PrivateDirectoryEntry {
if len(entries) < maximumEntries {
return append(entries, entry)
}
maximum := 0
for index := 1; index < len(entries); index++ {
if entries[index].Name > entries[maximum].Name {
maximum = index
}
}
if entry.Name < entries[maximum].Name {
entries[maximum] = entry
}
return entries
}
func WriteCanonicalNewFile(path string, contents []byte, mode os.FileMode) error {
return writeCanonicalNewFile(path, contents, mode, false)
}
+68
View File
@@ -5,7 +5,9 @@ import (
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/aritmolab/thothii/tools/tht/internal/testsupport"
)
@@ -125,3 +127,69 @@ func TestListCanonicalPrivateDirectoryBoundsAndSortsValidatedEntries(t *testing.
t.Fatalf("257-entry listing error = %v, want ErrUnsafeFile", err)
}
}
func TestListCanonicalPrivateDirectoryPageContinuesPastOneBoundedPage(t *testing.T) {
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-page-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
if err := ProtectPrivateDirectory(root); err != nil {
t.Fatal(err)
}
validName := func(name string) bool {
return len(name) == len(fmt.Sprintf("%064x.json", 0)) && strings.HasSuffix(name, ".json")
}
for index := 0; index < 513; index++ {
path := filepath.Join(root, fmt.Sprintf("%064x.json", index))
if err := WriteCanonicalNewFile(path, []byte("record"), 0o600); err != nil {
t.Fatalf("WriteCanonicalNewFile(%d) error = %v", index, err)
}
}
first, err := ListCanonicalPrivateDirectoryPage(root, 512, "", validName)
if err != nil || len(first.Entries) != 512 || !first.More || first.Entries[0].Name != fmt.Sprintf("%064x.json", 0) || first.Entries[511].Name != fmt.Sprintf("%064x.json", 511) {
t.Fatalf("first page = %#v error = %v", first, err)
}
second, err := ListCanonicalPrivateDirectoryPage(root, 512, first.Entries[511].Name, validName)
if err != nil || len(second.Entries) != 1 || second.More || second.Entries[0].Name != fmt.Sprintf("%064x.json", 512) {
t.Fatalf("second page = %#v error = %v", second, err)
}
}
func TestListCanonicalPrivateDirectoryPageRejectsInPlaceDirectoryMutation(t *testing.T) {
temporaryRoot, err := filepath.EvalSymlinks(os.TempDir())
if err != nil {
t.Fatal(err)
}
root, err := os.MkdirTemp(temporaryRoot, "tht-safeio-page-")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.RemoveAll(root) })
if err := ProtectPrivateDirectory(root); err != nil {
t.Fatal(err)
}
filename := fmt.Sprintf("%064x.json", 0)
if err := WriteCanonicalNewFile(filepath.Join(root, filename), []byte("record"), 0o600); err != nil {
t.Fatal(err)
}
changed := false
_, err = ListCanonicalPrivateDirectoryPage(root, 1, "", func(name string) bool {
if !changed {
changed = true
at := time.Unix(1_893_456_245, 0)
if changeErr := os.Chtimes(root, at, at); changeErr != nil {
t.Fatalf("Chtimes() error = %v", changeErr)
}
}
return name == filename
})
if !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("unstable page error = %v, want ErrUnsafeFile", err)
}
}