diff --git a/tools/dwh-auth/internal/command/command.go b/tools/dwh-auth/internal/command/command.go index df70c959..6d24bd52 100644 --- a/tools/dwh-auth/internal/command/command.go +++ b/tools/dwh-auth/internal/command/command.go @@ -236,9 +236,9 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i } addErr := addRecord(store, value) if addErr != nil { - lookupErr := lookupCreatedRecord(store, material.KeyID) + present, snapshotErr := snapshotContainsKey(store, material.KeyID) closeErr := closeStore(store) - if errors.Is(lookupErr, registry.ErrNotFound) && closeErr == nil { + if snapshotErr == nil && !present && closeErr == nil { if cleanupOutput(options.output) != nil { writeCleanupFailure(stderr, options.output) return exitIntegrity @@ -448,9 +448,17 @@ func runServe(args []string, stderr io.Writer) int { return exitIntegrity } -func lookupCreatedRecord(store *registry.Store, keyID string) error { - _, err := store.Find(keyID) - return err +func snapshotContainsKey(store *registry.Store, keyID string) (bool, error) { + values, err := store.List() + if err != nil { + return false, err + } + for _, value := range values { + if value.KeyID == keyID { + return true, nil + } + } + return false, nil } func openCheckedStore(ctx context.Context, root string) (*registry.Store, error) { diff --git a/tools/dwh-auth/internal/command/command_test.go b/tools/dwh-auth/internal/command/command_test.go index a9c5eaf3..1059c1c6 100644 --- a/tools/dwh-auth/internal/command/command_test.go +++ b/tools/dwh-auth/internal/command/command_test.go @@ -359,6 +359,47 @@ func TestCreateCleansOnlyWhenFailedPublicationProvesKeyAbsent(t *testing.T) { } } +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