fix: harden DWH credential administration CLI
This commit is contained in:
@@ -18,6 +18,8 @@ import (
|
|||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
"unicode"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/credential"
|
"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/record"
|
||||||
@@ -34,6 +36,13 @@ const (
|
|||||||
commandServeMsg = "serve is not available in this release"
|
commandServeMsg = "serve is not available in this release"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
nowUTC = func() time.Time { return time.Now().UTC() }
|
||||||
|
generateCredential = func() (credential.Material, error) { return credential.Generate(rand.Reader) }
|
||||||
|
addRecord = func(store *registry.Store, value record.Record) error { return store.Add(value) }
|
||||||
|
closeStore = func(store *registry.Store) error { return store.Close() }
|
||||||
|
)
|
||||||
|
|
||||||
// Run executes one dwh-auth invocation and returns its stable process exit
|
// Run executes one dwh-auth invocation and returns its stable process exit
|
||||||
// status. It writes only redacted, deterministic diagnostics to stderr.
|
// status. It writes only redacted, deterministic diagnostics to stderr.
|
||||||
func Run(ctx context.Context, args []string, _ io.Reader, stdout, stderr io.Writer) int {
|
func Run(ctx context.Context, args []string, _ io.Reader, stdout, stderr io.Writer) int {
|
||||||
@@ -43,6 +52,9 @@ func Run(ctx context.Context, args []string, _ io.Reader, stdout, stderr io.Writ
|
|||||||
if stdout == nil || stderr == nil {
|
if stdout == nil || stderr == nil {
|
||||||
return exitIntegrity
|
return exitIntegrity
|
||||||
}
|
}
|
||||||
|
if len(args) > 0 && args[0] == "serve" {
|
||||||
|
return runServe(args[1:], stderr)
|
||||||
|
}
|
||||||
root, rest, err := parseRoot(args)
|
root, rest, err := parseRoot(args)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
writeLine(stderr, "unsafe invocation")
|
writeLine(stderr, "unsafe invocation")
|
||||||
@@ -58,8 +70,6 @@ func Run(ctx context.Context, args []string, _ io.Reader, stdout, stderr io.Writ
|
|||||||
return runKey(ctx, root, rest[1:], stdout, stderr)
|
return runKey(ctx, root, rest[1:], stdout, stderr)
|
||||||
case "check":
|
case "check":
|
||||||
return runCheck(ctx, root, rest[1:], stdout, stderr)
|
return runCheck(ctx, root, rest[1:], stdout, stderr)
|
||||||
case "serve":
|
|
||||||
return runServe(root, rest[1:], stderr)
|
|
||||||
default:
|
default:
|
||||||
writeLine(stderr, "unsafe invocation")
|
writeLine(stderr, "unsafe invocation")
|
||||||
return exitUsage
|
return exitUsage
|
||||||
@@ -105,6 +115,7 @@ type parsed struct {
|
|||||||
reason string
|
reason string
|
||||||
json bool
|
json bool
|
||||||
legacyRaw bool
|
legacyRaw bool
|
||||||
|
socket string
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseFlags(args []string, allowed map[string]bool) (parsed, error) {
|
func parseFlags(args []string, allowed map[string]bool) (parsed, error) {
|
||||||
@@ -147,6 +158,8 @@ func parseFlags(args []string, allowed map[string]bool) (parsed, error) {
|
|||||||
value.keyID = args[i]
|
value.keyID = args[i]
|
||||||
case "--reason":
|
case "--reason":
|
||||||
value.reason = args[i]
|
value.reason = args[i]
|
||||||
|
case "--socket":
|
||||||
|
value.socket = args[i]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return value, nil
|
return value, nil
|
||||||
@@ -159,11 +172,12 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i
|
|||||||
"--expires-at": true,
|
"--expires-at": true,
|
||||||
"--output": true,
|
"--output": true,
|
||||||
})
|
})
|
||||||
if err != nil || options.installationID == "" || options.output == "" || !filepath.IsAbs(options.output) || !validInstallationID(options.installationID) || !validMetadata(options.description) {
|
if err != nil || options.installationID == "" || options.output == "" || !filepath.IsAbs(options.output) || !validInstallationID(options.installationID) || !validMetadata(options.description) || containsCanonicalCredential(options.description) {
|
||||||
writeLine(stderr, "unsafe invocation")
|
writeLine(stderr, "unsafe invocation")
|
||||||
return exitUsage
|
return exitUsage
|
||||||
}
|
}
|
||||||
expiresAt, err := parseExpiry(options.expiresAt)
|
createdAt := nowUTC()
|
||||||
|
expiresAt, err := parseExpiry(options.expiresAt, createdAt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
writeLine(stderr, "unsafe invocation")
|
writeLine(stderr, "unsafe invocation")
|
||||||
return exitUsage
|
return exitUsage
|
||||||
@@ -173,7 +187,7 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i
|
|||||||
return exitIntegrity
|
return exitIntegrity
|
||||||
}
|
}
|
||||||
|
|
||||||
material, err := credential.Generate(rand.Reader)
|
material, err := generateCredential()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
writeLine(stderr, "integrity failure")
|
writeLine(stderr, "integrity failure")
|
||||||
return exitIntegrity
|
return exitIntegrity
|
||||||
@@ -210,7 +224,6 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i
|
|||||||
writeLine(stderr, "integrity failure")
|
writeLine(stderr, "integrity failure")
|
||||||
return exitIntegrity
|
return exitIntegrity
|
||||||
}
|
}
|
||||||
now := time.Now().UTC()
|
|
||||||
value := record.Record{
|
value := record.Record{
|
||||||
SchemaVersion: record.SchemaVersion,
|
SchemaVersion: record.SchemaVersion,
|
||||||
Kind: record.KindV1,
|
Kind: record.KindV1,
|
||||||
@@ -218,14 +231,20 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i
|
|||||||
InstallationID: options.installationID,
|
InstallationID: options.installationID,
|
||||||
Description: options.description,
|
Description: options.description,
|
||||||
SecretSHA256: base64.RawURLEncoding.EncodeToString(material.Digest[:]),
|
SecretSHA256: base64.RawURLEncoding.EncodeToString(material.Digest[:]),
|
||||||
CreatedAt: now,
|
CreatedAt: createdAt,
|
||||||
ExpiresAt: expiresAt,
|
ExpiresAt: expiresAt,
|
||||||
}
|
}
|
||||||
addErr := store.Add(value)
|
addErr := addRecord(store, value)
|
||||||
closeErr := store.Close()
|
if addErr != nil {
|
||||||
if addErr != nil || closeErr != nil {
|
lookupErr := lookupCreatedRecord(store, material.KeyID)
|
||||||
if cleanupOutput(options.output) != nil {
|
closeErr := closeStore(store)
|
||||||
writeCleanupFailure(stderr, options.output)
|
if errors.Is(lookupErr, registry.ErrNotFound) && closeErr == nil {
|
||||||
|
if cleanupOutput(options.output) != nil {
|
||||||
|
writeCleanupFailure(stderr, options.output)
|
||||||
|
return exitIntegrity
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
writeRecoveryFailure(stderr, options.output)
|
||||||
return exitIntegrity
|
return exitIntegrity
|
||||||
}
|
}
|
||||||
if errors.Is(addErr, registry.ErrConflict) || errors.Is(addErr, registry.ErrRevoked) {
|
if errors.Is(addErr, registry.ErrConflict) || errors.Is(addErr, registry.ErrRevoked) {
|
||||||
@@ -235,6 +254,10 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i
|
|||||||
writeLine(stderr, "integrity failure")
|
writeLine(stderr, "integrity failure")
|
||||||
return exitIntegrity
|
return exitIntegrity
|
||||||
}
|
}
|
||||||
|
if err := closeStore(store); err != nil {
|
||||||
|
writeRecoveryFailure(stderr, options.output)
|
||||||
|
return exitIntegrity
|
||||||
|
}
|
||||||
_, _ = fmt.Fprintf(stdout, "created key_id=%s installation_id=%s output=%s\n", material.KeyID, options.installationID, filepath.Clean(options.output))
|
_, _ = fmt.Fprintf(stdout, "created key_id=%s installation_id=%s output=%s\n", material.KeyID, options.installationID, filepath.Clean(options.output))
|
||||||
return exitOK
|
return exitOK
|
||||||
}
|
}
|
||||||
@@ -358,7 +381,7 @@ func runStatus(ctx context.Context, root string, args []string, stdout, stderr i
|
|||||||
|
|
||||||
func runRevoke(ctx context.Context, root string, args []string, stdout, stderr io.Writer) int {
|
func runRevoke(ctx context.Context, root string, args []string, stdout, stderr io.Writer) int {
|
||||||
options, err := parseFlags(args, map[string]bool{"--key-id": true, "--reason": true})
|
options, err := parseFlags(args, map[string]bool{"--key-id": true, "--reason": true})
|
||||||
if err != nil || options.keyID == "" || !validKeyID(options.keyID) || options.reason == "" || !validMetadata(options.reason) {
|
if err != nil || options.keyID == "" || !validKeyID(options.keyID) || options.reason == "" || !validMetadata(options.reason) || containsCanonicalCredential(options.reason) {
|
||||||
writeLine(stderr, "unsafe invocation")
|
writeLine(stderr, "unsafe invocation")
|
||||||
return exitUsage
|
return exitUsage
|
||||||
}
|
}
|
||||||
@@ -371,7 +394,7 @@ func runRevoke(ctx context.Context, root string, args []string, stdout, stderr i
|
|||||||
writeLine(stderr, "integrity failure")
|
writeLine(stderr, "integrity failure")
|
||||||
return exitIntegrity
|
return exitIntegrity
|
||||||
}
|
}
|
||||||
err = store.Revoke(options.keyID, options.reason, time.Now().UTC())
|
err = store.Revoke(options.keyID, options.reason, nowUTC())
|
||||||
closeErr := store.Close()
|
closeErr := store.Close()
|
||||||
if errors.Is(err, registry.ErrNotFound) {
|
if errors.Is(err, registry.ErrNotFound) {
|
||||||
writeLine(stderr, "not found")
|
writeLine(stderr, "not found")
|
||||||
@@ -414,27 +437,22 @@ func runCheck(ctx context.Context, root string, args []string, stdout, stderr io
|
|||||||
return exitOK
|
return exitOK
|
||||||
}
|
}
|
||||||
|
|
||||||
func runServe(root string, args []string, stderr io.Writer) int {
|
func runServe(args []string, stderr io.Writer) int {
|
||||||
options, err := parseFlags(args, map[string]bool{"--socket": true})
|
if len(args) != 4 || args[0] != "--registry-root" || args[2] != "--socket" || !filepath.IsAbs(args[1]) || !filepath.IsAbs(args[3]) {
|
||||||
if err != nil || options.output != "" {
|
|
||||||
writeLine(stderr, "unsafe invocation")
|
writeLine(stderr, "unsafe invocation")
|
||||||
return exitUsage
|
return exitUsage
|
||||||
}
|
}
|
||||||
// The Unix-socket service is intentionally introduced by Task 4. Keeping
|
// The Unix-socket service is intentionally introduced by Task 4. Keeping
|
||||||
// this grammar reserved prevents accidental plaintext or network fallback.
|
// this exact grammar reserved prevents accidental plaintext or network fallback.
|
||||||
for i := range args {
|
|
||||||
if args[i] == "--socket" {
|
|
||||||
if i+1 >= len(args) || !filepath.IsAbs(args[i+1]) {
|
|
||||||
writeLine(stderr, "unsafe invocation")
|
|
||||||
return exitUsage
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
_ = root
|
|
||||||
writeLine(stderr, commandServeMsg)
|
writeLine(stderr, commandServeMsg)
|
||||||
return exitIntegrity
|
return exitIntegrity
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func lookupCreatedRecord(store *registry.Store, keyID string) error {
|
||||||
|
_, err := store.Find(keyID)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
func openCheckedStore(ctx context.Context, root string) (*registry.Store, error) {
|
func openCheckedStore(ctx context.Context, root string) (*registry.Store, error) {
|
||||||
if err := ctx.Err(); err != nil {
|
if err := ctx.Err(); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -509,12 +527,16 @@ func writeCleanupFailure(stderr io.Writer, path string) {
|
|||||||
_, _ = fmt.Fprintf(stderr, "cleanup failed path=%s\n", filepath.Clean(path))
|
_, _ = fmt.Fprintf(stderr, "cleanup failed path=%s\n", filepath.Clean(path))
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseExpiry(value string) (*time.Time, error) {
|
func writeRecoveryFailure(stderr io.Writer, path string) {
|
||||||
|
_, _ = fmt.Fprintf(stderr, "publication uncertain path=%s\n", filepath.Clean(path))
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseExpiry(value string, createdAt time.Time) (*time.Time, error) {
|
||||||
if value == "" {
|
if value == "" {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
parsed, err := time.Parse(time.RFC3339, value)
|
parsed, err := time.Parse(time.RFC3339, value)
|
||||||
if err != nil || parsed.Location() != time.UTC {
|
if err != nil || parsed.Location() != time.UTC || !parsed.After(createdAt) {
|
||||||
return nil, errors.New("expiry must be UTC RFC3339")
|
return nil, errors.New("expiry must be UTC RFC3339")
|
||||||
}
|
}
|
||||||
return &parsed, nil
|
return &parsed, nil
|
||||||
@@ -535,17 +557,38 @@ func validInstallationID(value string) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func validMetadata(value string) bool {
|
func validMetadata(value string) bool {
|
||||||
if len([]rune(value)) > maxReasonBytes {
|
if !utf8.ValidString(value) || len([]rune(value)) > maxReasonBytes {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
for _, r := range value {
|
for _, r := range value {
|
||||||
if r < 0x20 || r == 0x7f {
|
if unicode.IsControl(r) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func containsCanonicalCredential(value string) bool {
|
||||||
|
const credentialLength = len(credential.Prefix) + 1 + credential.KeyIDEncodedLength + 1 + credential.SecretEncodedLength
|
||||||
|
needle := credential.Prefix + "."
|
||||||
|
for offset := 0; offset < len(value); {
|
||||||
|
index := strings.Index(value[offset:], needle)
|
||||||
|
if index < 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
index += offset
|
||||||
|
if index+credentialLength <= len(value) {
|
||||||
|
candidate := []byte(value[index : index+credentialLength])
|
||||||
|
sum := sha256.Sum256(candidate)
|
||||||
|
if _, ok := credential.VerifyV1(candidate, record.Digest(sum)); ok {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
offset = index + len(needle)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
func validKeyID(value string) bool {
|
func validKeyID(value string) bool {
|
||||||
if value == record.LegacyKeyID {
|
if value == record.LegacyKeyID {
|
||||||
return true
|
return true
|
||||||
|
|||||||
@@ -214,6 +214,151 @@ func TestCreateRejectsInvalidExpiryWithoutWritingOutput(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
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 TestServeAcceptsOnlyExactTaskFourReservation(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
socket := filepath.Join(t.TempDir(), "verify.sock")
|
||||||
|
stdout, stderr, code := run(t, "serve", "--registry-root", root, "--socket", socket)
|
||||||
|
if code != 4 || stdout != "" || stderr != "serve is not available in this release\n" {
|
||||||
|
t.Fatalf("valid serve reservation = (%d, %q, %q)", code, stdout, stderr)
|
||||||
|
}
|
||||||
|
invalid := [][]string{
|
||||||
|
{"serve", "--registry-root", "relative", "--socket", socket},
|
||||||
|
{"serve", "--registry-root", root},
|
||||||
|
{"serve", "--registry-root", root, "--socket", "relative"},
|
||||||
|
{"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 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 run(t *testing.T, args ...string) (string, string, int) {
|
func run(t *testing.T, args ...string) (string, string, int) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
var stdout, stderr bytes.Buffer
|
var stdout, stderr bytes.Buffer
|
||||||
|
|||||||
Reference in New Issue
Block a user