Files
ThothII/tools/tht/internal/authstorage/storage_test.go
T

277 lines
12 KiB
Go

package authstorage
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"os"
"path/filepath"
"strings"
"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 listed.Entries == nil || 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 TestProtocolPermitsBoundedReservationSlotsOnlyForOIDCRecords(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
slot := "slot-00.json"
contents := base64.StdEncoding.EncodeToString([]byte("reservation"))
if created := runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "oidc", Filename: slot, ContentBase64: contents}); !created.Created {
t.Fatal("OIDC slot create did not report success")
}
if read := runRequest(t, request{Version: 1, Operation: "read", Root: root, Directory: "oidc", Filename: slot}); !read.Found || decodeContent(t, read) != "reservation" {
t.Fatalf("OIDC slot read = %#v", read)
}
listed := runRequest(t, request{Version: 1, Operation: "list", Root: root, Directory: "oidc"})
if listed.Entries == nil || len(*listed.Entries) != 1 || (*listed.Entries)[0].Name != slot {
t.Fatalf("OIDC slot list = %#v", listed.Entries)
}
if removed := runRequest(t, request{Version: 1, Operation: "remove", Root: root, Directory: "oidc", Filename: slot}); !removed.Removed {
t.Fatal("OIDC slot remove did not report success")
}
runRejected(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: slot, ContentBase64: contents})
runRejected(t, request{Version: 1, Operation: "create", Root: root, Directory: "oidc", Filename: "slot-64.json", ContentBase64: contents})
}
func TestProtocolListSerializesLowerCamelBridgeDTO(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
filename := "cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc.json"
runRequest(t, request{Version: 1, Operation: "create", Root: root, Directory: "sessions", Filename: filename, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))})
input, err := json.Marshal(request{Version: 1, Operation: "list", Root: root, Directory: "sessions"})
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())
}
if !strings.Contains(stdout.String(), `"name":"`+filename+`"`) || strings.Contains(stdout.String(), `"Name":`) {
t.Fatalf("list bridge JSON = %q, want lower-camel entry fields", stdout.String())
}
}
func TestProtocolListAlwaysSerializesAnEmptyEntriesArray(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
input, err := json.Marshal(request{Version: 1, Operation: "list", Root: root, Directory: "sessions"})
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())
}
if !strings.Contains(stdout.String(), `"entries":[]`) {
t.Fatalf("empty list bridge JSON = %q, want entries array", stdout.String())
}
}
func TestProtocolPermitsClaimNamesOnlyForOIDCRemove(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
digest := "abababababababababababababababababababababababababababababababab.json"
claim := strings.TrimSuffix(digest, ".json") + ".claim"
for _, directory := range []string{"sessions", "oidc"} {
runRequest(t, request{Version: 1, Operation: "list", Root: root, Directory: directory})
}
if err := safeio.WriteCanonicalNewPrivateFile(filepath.Join(root, "sessions", claim), []byte("orphan"), 0o600); err != nil {
t.Fatal(err)
}
if err := safeio.WriteCanonicalNewPrivateFile(filepath.Join(root, "oidc", claim), []byte("orphan"), 0o600); err != nil {
t.Fatal(err)
}
runRejected(t, request{Version: 1, Operation: "remove", Root: root, Directory: "sessions", Filename: claim})
runRejected(t, request{Version: 1, Operation: "list", Root: root, Directory: "sessions", Filename: claim})
runRejected(t, request{Version: 1, Operation: "read", Root: root, Directory: "oidc", Filename: claim})
runRejected(t, request{Version: 1, Operation: "create", Root: root, Directory: "oidc", Filename: claim, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))})
runRejected(t, request{Version: 1, Operation: "replace", Root: root, Directory: "oidc", Filename: claim, ContentBase64: base64.StdEncoding.EncodeToString([]byte("record"))})
runRejected(t, request{Version: 1, Operation: "remove", Root: root, Directory: "oidc", Filename: "../" + claim})
runRejected(t, request{Version: 1, Operation: "remove", Root: root, Directory: "oidc", Filename: strings.TrimSuffix(claim, ".claim") + ".claim.bak"})
removed := runRequest(t, request{Version: 1, Operation: "remove", Root: root, Directory: "oidc", Filename: claim})
if !removed.Removed {
t.Fatal("OIDC orphan claim removal did not report success")
}
}
func TestRunRejectsProtocolOverflowAndTrailingJSONValues(t *testing.T) {
root := filepath.Join(privateTestRoot(t), "auth")
valid, err := json.Marshal(request{Version: 1, Operation: "list", Root: root, Directory: "sessions"})
if err != nil {
t.Fatal(err)
}
if len(valid) >= maximumProtocolBytes {
t.Fatal("test request unexpectedly consumes protocol bound")
}
overflowWhitespace := append(append([]byte(nil), valid...), bytes.Repeat([]byte(" "), maximumProtocolBytes+1-len(valid))...)
runRawRejected(t, overflowWhitespace)
withinBound := append(append([]byte(nil), valid...), []byte("{}")...)
runRawRejected(t, withinBound)
crossingBound := append(append([]byte(nil), valid...), bytes.Repeat([]byte(" "), maximumProtocolBytes-len(valid)-1)...)
crossingBound = append(crossingBound, []byte("{}")...)
runRawRejected(t, crossingBound)
}
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 runRawRejected(t *testing.T, input []byte) {
t.Helper()
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
}