165 lines
6.3 KiB
Go
165 lines
6.3 KiB
Go
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
|
|
}
|