Files
ThothII/tools/dwh-auth/internal/command/command_test.go

447 lines
19 KiB
Go

//go:build linux
package command
import (
"bytes"
"context"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"errors"
"io"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/aritmolab/thothii/tools/dwh-auth/internal/credential"
"github.com/aritmolab/thothii/tools/dwh-auth/internal/record"
"github.com/aritmolab/thothii/tools/dwh-auth/internal/registry"
"github.com/aritmolab/thothii/tools/dwh-auth/internal/service"
)
const sentinelSecret = "SENTINEL-DWH-SECRET-must-never-be-emitted"
func TestCreateWritesNew0600KeyAndOnlyPublicSummary(t *testing.T) {
root := t.TempDir()
output := filepath.Join(t.TempDir(), "new-key")
stdout, stderr, code := run(t, "--registry-root", root, "key", "create", "--installation-id", "mac-client", "--description", "synthetic", "--output", output)
if code != 0 {
t.Fatalf("Run() exit = %d, stdout=%q stderr=%q", code, stdout, stderr)
}
if stderr != "" {
t.Fatalf("stderr = %q, want empty", stderr)
}
if !strings.HasPrefix(stdout, "created key_id=") || !strings.HasSuffix(stdout, " output="+output+"\n") {
t.Fatalf("stdout = %q, want only creation summary", stdout)
}
if strings.Contains(stdout+stderr, sentinelSecret) {
t.Fatal("sentinel secret appeared in command output")
}
value, err := os.ReadFile(output)
if err != nil {
t.Fatalf("ReadFile(output) error = %v", err)
}
if len(value) == 0 || !strings.HasPrefix(string(value), credential.Prefix+".") {
t.Fatalf("created key does not have canonical prefix: %q", value)
}
info, err := os.Stat(output)
if err != nil {
t.Fatalf("Stat(output) error = %v", err)
}
if got, want := info.Mode().Perm(), os.FileMode(0o600); got != want {
t.Fatalf("output mode = %04o, want %04o", got, want)
}
if code := runOnly(t, "--registry-root", root, "key", "create", "--installation-id", "other", "--output", output); code != 4 {
t.Fatalf("second create exit = %d, want filesystem/integrity exit 4", code)
}
if got, err := os.ReadFile(output); err != nil || !bytes.Equal(got, value) {
t.Fatalf("existing output changed after rejected create: err=%v value=%q", err, got)
}
}
func TestCreateSupportsExpiryAndImportReads0600WithoutChangingSource(t *testing.T) {
root := t.TempDir()
keyFile := filepath.Join(t.TempDir(), "legacy-key")
if err := os.WriteFile(keyFile, []byte(sentinelSecret), 0o600); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
original, err := os.ReadFile(keyFile)
if err != nil {
t.Fatalf("ReadFile() error = %v", err)
}
stdout, stderr, code := run(t, "--registry-root", root, "key", "import", "--legacy-raw", "--installation-id", "legacy-shared", "--from-file", keyFile)
if code != 0 || stderr != "" || stdout != "imported key_id=legacy-shared installation_id=legacy-shared\n" {
t.Fatalf("import result = (%d, %q, %q)", code, stdout, stderr)
}
got, err := os.ReadFile(keyFile)
if err != nil {
t.Fatalf("ReadFile(source) error = %v", err)
}
if !bytes.Equal(got, original) {
t.Fatal("legacy source file was modified")
}
if strings.Contains(stdout+stderr, sentinelSecret) {
t.Fatal("sentinel secret appeared in import output")
}
if info, err := os.Stat(keyFile); err != nil || info.Mode().Perm() != 0o600 {
t.Fatalf("legacy source mode = %v, err=%v; want 0600", info.Mode().Perm(), err)
}
if code := runOnly(t, "--registry-root", root, "key", "import", "--legacy-raw", "--installation-id", "legacy-shared", "--from-file", keyFile); code != 4 {
t.Fatalf("duplicate import exit = %d, want 4", code)
}
}
func TestListAndStatusAreRedactedAndJSONIsPristine(t *testing.T) {
root := t.TempDir()
output := filepath.Join(t.TempDir(), "key")
if code := runOnly(t, "--registry-root", root, "key", "create", "--installation-id", "json-client", "--output", output); code != 0 {
t.Fatalf("create exit = %d", code)
}
stdout, stderr, code := run(t, "--registry-root", root, "key", "list", "--json")
if code != 0 || stderr != "" {
t.Fatalf("list result = (%d, %q, %q)", code, stdout, stderr)
}
var listed []registry.PublicRecord
if err := json.Unmarshal([]byte(stdout), &listed); err != nil {
t.Fatalf("list stdout is not pristine JSON: %v (%q)", err, stdout)
}
if len(listed) != 1 || listed[0].InstallationID != "json-client" {
t.Fatalf("list = %#v", listed)
}
if strings.Contains(stdout, "secret_sha256") || strings.Contains(stdout, sentinelSecret) {
t.Fatalf("list exposed secret material: %q", stdout)
}
keyID := listed[0].KeyID
stdout, stderr, code = run(t, "--registry-root", root, "key", "status", "--key-id", keyID, "--json")
if code != 0 || stderr != "" {
t.Fatalf("status result = (%d, %q, %q)", code, stdout, stderr)
}
var status registry.PublicRecord
if err := json.Unmarshal([]byte(stdout), &status); err != nil {
t.Fatalf("status stdout is not pristine JSON: %v", err)
}
if status.KeyID != keyID || strings.Contains(stdout, "secret_sha256") || strings.Contains(stdout, sentinelSecret) {
t.Fatalf("status exposed secret material: %q", stdout)
}
}
func TestRevokeRequiresReasonAndStatusShowsRevokedWithoutDigest(t *testing.T) {
root := t.TempDir()
output := filepath.Join(t.TempDir(), "key")
if code := runOnly(t, "--registry-root", root, "key", "create", "--installation-id", "revoke-client", "--output", output); code != 0 {
t.Fatalf("create exit = %d", code)
}
listOut, _, _ := run(t, "--registry-root", root, "key", "list", "--json")
var listed []registry.PublicRecord
if err := json.Unmarshal([]byte(listOut), &listed); err != nil || len(listed) != 1 {
t.Fatalf("list = %q, err=%v", listOut, err)
}
keyID := listed[0].KeyID
if code := runOnly(t, "--registry-root", root, "key", "revoke", "--key-id", keyID); code != 2 {
t.Fatalf("missing reason exit = %d, want usage 2", code)
}
stdout, stderr, code := run(t, "--registry-root", root, "key", "revoke", "--key-id", keyID, "--reason", "synthetic rotation")
if code != 0 || stdout != "revoked key_id="+keyID+"\n" || stderr != "" {
t.Fatalf("revoke result = (%d, %q, %q)", code, stdout, stderr)
}
statusOut, _, statusCode := run(t, "--registry-root", root, "key", "status", "--key-id", keyID, "--json")
if statusCode != 0 || strings.Contains(statusOut, "secret_sha256") || strings.Contains(statusOut, sentinelSecret) {
t.Fatalf("revoked status = (%d, %q)", statusCode, statusOut)
}
if !strings.Contains(statusOut, "\"state\":\"revoked\"") {
t.Fatalf("status did not show revoked: %q", statusOut)
}
}
func TestCheckJSONIsPristineAndNotFoundIsExitThree(t *testing.T) {
root := t.TempDir()
stdout, stderr, code := run(t, "--registry-root", root, "check", "--json")
if code != 0 || stderr != "" || stdout != "{\"status\":\"ok\"}\n" {
t.Fatalf("check result = (%d, %q, %q)", code, stdout, stderr)
}
_, stderr, code = run(t, "--registry-root", root, "key", "status", "--key-id", "AAAAAAAAAAAAAAAA")
if code != 3 || !strings.Contains(stderr, "not found") {
t.Fatalf("not-found result = (%d, %q)", code, stderr)
}
}
func TestUnsafeInvocationRejectsRelativeAndSecretValuedFlags(t *testing.T) {
cases := [][]string{
{"--registry-root", "relative", "check"},
{"--registry-root", "/tmp/registry", "key", "create", "--installation-id", sentinelSecret, "--output", "/tmp/key"},
{"--registry-root", "/tmp/registry", "key", "create", "--installation-id", "safe", "--output", "relative"},
{"--registry-root", "/tmp/registry", "key", "import", "--legacy-raw", "--installation-id", "legacy-shared", "--from-file", "relative"},
{"--registry-root", "/tmp/registry", "key", "status", "--key-id", sentinelSecret},
}
for _, args := range cases {
stdout, stderr, code := run(t, args...)
if code != 2 {
t.Errorf("args %q exit = %d, want 2 (stdout=%q stderr=%q)", args, code, stdout, stderr)
}
if strings.Contains(stdout+stderr, sentinelSecret) {
t.Errorf("args %q leaked sentinel secret in output", args)
}
}
}
func TestIntegrityFailureUsesExitFourWithoutDigest(t *testing.T) {
root := t.TempDir()
if err := os.Mkdir(filepath.Join(root, "active"), 0o750); err != nil {
t.Fatalf("Mkdir(active) error = %v", err)
}
if err := os.Mkdir(filepath.Join(root, "revoked"), 0o750); err != nil {
t.Fatalf("Mkdir(revoked) error = %v", err)
}
if err := os.WriteFile(filepath.Join(root, "active", "unexpected"), []byte(sentinelSecret), 0o600); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
stdout, stderr, code := run(t, "--registry-root", root, "check")
if code != 4 || stdout != "" || strings.Contains(stderr, sentinelSecret) || strings.Contains(stderr, "sha") {
t.Fatalf("integrity result = (%d, %q, %q)", code, stdout, stderr)
}
}
func TestCreateRejectsInvalidExpiryWithoutWritingOutput(t *testing.T) {
root := t.TempDir()
output := filepath.Join(t.TempDir(), "key")
if code := runOnly(t, "--registry-root", root, "key", "create", "--installation-id", "client", "--expires-at", "not-time", "--output", output); code != 2 {
t.Fatalf("invalid expiry exit = %d, want 2", code)
}
if _, err := os.Stat(output); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("invalid expiry output stat error = %v, want absent", err)
}
}
func TestMetadataRejectsEmbeddedCanonicalCredentialBeforePersistence(t *testing.T) {
material, err := credential.Generate(bytes.NewReader(bytes.Repeat([]byte{0x33}, 44)))
if err != nil {
t.Fatalf("Generate() error = %v", err)
}
sentinel := string(material.Value)
root := t.TempDir()
output := filepath.Join(t.TempDir(), "key")
stdout, stderr, code := run(t, "--registry-root", root, "key", "create", "--installation-id", "embedded", "--description", "prefix-"+sentinel+"-suffix", "--output", output)
if code != 2 || stdout != "" || stderr != "unsafe invocation\n" {
t.Fatalf("embedded description result = (%d, %q, %q)", code, stdout, stderr)
}
if strings.Contains(stdout+stderr, sentinel) {
t.Fatal("embedded credential appeared in create diagnostics")
}
if _, err := os.Stat(output); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("embedded description output stat error = %v, want absent", err)
}
cleanOutput := filepath.Join(t.TempDir(), "clean-key")
if code := runOnly(t, "--registry-root", root, "key", "create", "--installation-id", "embedded", "--output", cleanOutput); code != 0 {
t.Fatalf("clean create exit = %d", code)
}
listOut, _, listCode := run(t, "--registry-root", root, "key", "list", "--json")
if listCode != 0 || strings.Contains(listOut, sentinel) {
t.Fatalf("list after rejected metadata = (%d, %q)", listCode, listOut)
}
var listed []registry.PublicRecord
if err := json.Unmarshal([]byte(listOut), &listed); err != nil || len(listed) != 1 {
t.Fatalf("list = %q, err=%v", listOut, err)
}
keyID := listed[0].KeyID
stdout, stderr, code = run(t, "--registry-root", root, "key", "revoke", "--key-id", keyID, "--reason", "prefix-"+sentinel+"-suffix")
if code != 2 || stdout != "" || stderr != "unsafe invocation\n" || strings.Contains(stdout+stderr, sentinel) {
t.Fatalf("embedded reason result = (%d, %q, %q)", code, stdout, stderr)
}
statusOut, _, statusCode := run(t, "--registry-root", root, "key", "status", "--key-id", keyID, "--json")
if statusCode != 0 || strings.Contains(statusOut, sentinel) {
t.Fatalf("status after rejected reason = (%d, %q)", statusCode, statusOut)
}
}
func TestMetadataAndExpiryAreValidatedBeforeKeyGeneration(t *testing.T) {
oldNow := nowUTC
nowUTC = func() time.Time { return time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) }
t.Cleanup(func() { nowUTC = oldNow })
cases := []struct {
name string
value string
}{
{name: "invalid utf8", value: string([]byte{0xc3, 0x28})},
{name: "unicode control", value: "before\u0085after"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
output := filepath.Join(t.TempDir(), "key")
if code := runOnly(t, "--registry-root", t.TempDir(), "key", "create", "--installation-id", "client", "--description", tc.value, "--output", output); code != 2 {
t.Fatalf("invalid metadata exit = %d, want 2", code)
}
if _, err := os.Stat(output); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("invalid metadata output stat error = %v, want absent", err)
}
})
}
for _, expiry := range []string{"2025-12-31T23:59:59Z", "2026-01-01T00:00:00Z"} {
t.Run("expiry "+expiry, func(t *testing.T) {
output := filepath.Join(t.TempDir(), "key")
if code := runOnly(t, "--registry-root", t.TempDir(), "key", "create", "--installation-id", "client", "--expires-at", expiry, "--output", output); code != 2 {
t.Fatalf("invalid expiry exit = %d, want 2", code)
}
if _, err := os.Stat(output); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("invalid expiry output stat error = %v, want absent", err)
}
})
}
output := filepath.Join(t.TempDir(), "future-key")
if code := runOnly(t, "--registry-root", t.TempDir(), "key", "create", "--installation-id", "client", "--expires-at", "2026-01-01T00:00:01Z", "--output", output); code != 0 {
t.Fatalf("future expiry exit = %d, want 0", code)
}
}
func TestServeAcceptsOnlyExactCheckedServiceConfig(t *testing.T) {
root := t.TempDir()
socket := filepath.Join(t.TempDir(), "verify.sock")
oldListen := listenAndServe
t.Cleanup(func() { listenAndServe = oldListen })
called := false
listenAndServe = func(ctx context.Context, config service.Config) error {
called = true
if ctx == nil || config.RegistryRoot != root || config.Socket != socket || config.Logger == nil {
t.Fatalf("serve config = %#v", config)
}
return nil
}
stdout, stderr, code := run(t, "serve", "--registry-root", root, "--socket", socket)
if code != 0 || stdout != "" || stderr != "" || !called {
t.Fatalf("valid serve = (%d, %q, %q), called=%v", code, stdout, stderr, called)
}
invalid := [][]string{
{"serve", "--registry-root", "relative", "--socket", socket},
{"serve", "--registry-root", root + "/.", "--socket", socket},
{"serve", "--registry-root", root},
{"serve", "--registry-root", root, "--socket", "relative"},
{"serve", "--registry-root", root, "--socket", filepath.Dir(socket) + "/./verify.sock"},
{"serve", "--registry-root", root, "--socket", socket, "--socket", socket},
{"serve", "--socket", socket, "--registry-root", root},
{"serve", "--registry-root", root, "--socket", socket, "--unknown", "x"},
{"--registry-root", root, "serve", "--socket", socket},
}
for _, args := range invalid {
stdout, stderr, code = run(t, args...)
if code != 2 || stdout != "" || stderr != "unsafe invocation\n" {
t.Errorf("invalid serve %q = (%d, %q, %q)", args, code, stdout, stderr)
}
}
}
func TestServeSanitizesServiceStartupFailure(t *testing.T) {
root := t.TempDir()
socket := filepath.Join(t.TempDir(), "verify.sock")
oldListen := listenAndServe
t.Cleanup(func() { listenAndServe = oldListen })
listenAndServe = func(context.Context, service.Config) error { return errors.New(sentinelSecret) }
stdout, stderr, code := run(t, "serve", "--registry-root", root, "--socket", socket)
if code != 4 || stdout != "" || stderr != "integrity failure\n" || strings.Contains(stdout+stderr, sentinelSecret) {
t.Fatalf("service startup failure = (%d, %q, %q)", code, stdout, stderr)
}
}
func TestCreateCleansOnlyWhenFailedPublicationProvesKeyAbsent(t *testing.T) {
oldAdd, oldClose := addRecord, closeStore
t.Cleanup(func() { addRecord, closeStore = oldAdd, oldClose })
root := t.TempDir()
output := filepath.Join(t.TempDir(), "pre-publication")
addRecord = func(*registry.Store, record.Record) error { return errors.New("synthetic add failure") }
if code := runOnly(t, "--registry-root", root, "key", "create", "--installation-id", "client", "--output", output); code != 4 {
t.Fatalf("pre-publication failure exit = %d, want 4", code)
}
if _, err := os.Stat(output); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("pre-publication output stat error = %v, want absent", err)
}
output = filepath.Join(t.TempDir(), "post-publication")
addRecord = func(store *registry.Store, value record.Record) error {
if err := store.Add(value); err != nil {
return err
}
return errors.New("synthetic ambiguous post-publication failure")
}
if code := runOnly(t, "--registry-root", root, "key", "create", "--installation-id", "client", "--output", output); code != 4 {
t.Fatalf("post-publication failure exit = %d, want 4", code)
}
if _, err := os.Stat(output); err != nil {
t.Fatalf("post-publication output stat error = %v, want retained: %v", err, err)
}
closeStore = func(store *registry.Store) error {
_ = store.Close()
return errors.New("synthetic close failure")
}
output = filepath.Join(t.TempDir(), "close-failure")
addRecord = oldAdd
if code := runOnly(t, "--registry-root", root, "key", "create", "--installation-id", "client-two", "--output", output); code != 4 {
t.Fatalf("close failure exit = %d, want 4", code)
}
if _, err := os.Stat(output); err != nil {
t.Fatalf("close-failure output stat error = %v, want retained: %v", err, err)
}
}
func TestCreateRetainsOutputWhenSnapshotFindsUnrelatedIntegrityFailure(t *testing.T) {
root := t.TempDir()
output := filepath.Join(t.TempDir(), "corrupt-registry")
addRecord = func(*registry.Store, record.Record) error {
if err := os.WriteFile(filepath.Join(root, "active", "unexpected"), []byte(sentinelSecret), 0o600); err != nil {
t.Fatalf("WriteFile(corrupt entry) error = %v", err)
}
return errors.New("synthetic add failure")
}
t.Cleanup(func() { addRecord = func(store *registry.Store, value record.Record) error { return store.Add(value) } })
stdout, stderr, code := run(t, "--registry-root", root, "key", "create", "--installation-id", "corrupt-client", "--output", output)
if code != 4 || stdout != "" || !strings.Contains(stderr, "publication uncertain path=") || strings.Contains(stderr, sentinelSecret) {
t.Fatalf("corrupt snapshot result = (%d, %q, %q)", code, stdout, stderr)
}
if _, err := os.Stat(output); err != nil {
t.Fatalf("corrupt snapshot output stat error = %v, want retained: %v", err, err)
}
}
func TestCreateRetainsOutputWhenFailedPublicationRecordIsExpired(t *testing.T) {
oldNow := nowUTC
nowUTC = func() time.Time { return time.Date(2020, 1, 1, 0, 0, 0, 0, time.UTC) }
t.Cleanup(func() { nowUTC = oldNow })
root := t.TempDir()
output := filepath.Join(t.TempDir(), "expired-record")
addRecord = func(store *registry.Store, value record.Record) error {
if err := store.Add(value); err != nil {
return err
}
return errors.New("synthetic expired post-publication failure")
}
t.Cleanup(func() { addRecord = func(store *registry.Store, value record.Record) error { return store.Add(value) } })
stdout, stderr, code := run(t, "--registry-root", root, "key", "create", "--installation-id", "expired-client", "--expires-at", "2020-01-01T00:00:01Z", "--output", output)
if code != 4 || stdout != "" || !strings.Contains(stderr, "publication uncertain path=") {
t.Fatalf("expired publication result = (%d, %q, %q)", code, stdout, stderr)
}
if _, err := os.Stat(output); err != nil {
t.Fatalf("expired publication output stat error = %v, want retained: %v", err, err)
}
}
func run(t *testing.T, args ...string) (string, string, int) {
t.Helper()
var stdout, stderr bytes.Buffer
code := Run(context.Background(), args, strings.NewReader(""), &stdout, &stderr)
return stdout.String(), stderr.String(), code
}
func runOnly(t *testing.T, args ...string) int {
t.Helper()
_, _, code := run(t, args...)
return code
}
// Keep imports and fixtures honest if this test file is extended with direct records.
var _ = base64.RawURLEncoding
var _ = sha256.Sum256
var _ = io.EOF
var _ = time.Time{}
var _ record.Record