diff --git a/tools/dwh-auth/internal/command/command.go b/tools/dwh-auth/internal/command/command.go index 8ead69a1..df70c959 100644 --- a/tools/dwh-auth/internal/command/command.go +++ b/tools/dwh-auth/internal/command/command.go @@ -18,6 +18,8 @@ import ( "path/filepath" "strings" "time" + "unicode" + "unicode/utf8" "github.com/aritmolab/thothii/tools/dwh-auth/internal/credential" "github.com/aritmolab/thothii/tools/dwh-auth/internal/record" @@ -34,6 +36,13 @@ const ( 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 // status. It writes only redacted, deterministic diagnostics to stderr. 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 { return exitIntegrity } + if len(args) > 0 && args[0] == "serve" { + return runServe(args[1:], stderr) + } root, rest, err := parseRoot(args) if err != nil { 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) case "check": return runCheck(ctx, root, rest[1:], stdout, stderr) - case "serve": - return runServe(root, rest[1:], stderr) default: writeLine(stderr, "unsafe invocation") return exitUsage @@ -105,6 +115,7 @@ type parsed struct { reason string json bool legacyRaw bool + socket string } 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] case "--reason": value.reason = args[i] + case "--socket": + value.socket = args[i] } } return value, nil @@ -159,11 +172,12 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i "--expires-at": 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") return exitUsage } - expiresAt, err := parseExpiry(options.expiresAt) + createdAt := nowUTC() + expiresAt, err := parseExpiry(options.expiresAt, createdAt) if err != nil { writeLine(stderr, "unsafe invocation") return exitUsage @@ -173,7 +187,7 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i return exitIntegrity } - material, err := credential.Generate(rand.Reader) + material, err := generateCredential() if err != nil { writeLine(stderr, "integrity failure") return exitIntegrity @@ -210,7 +224,6 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i writeLine(stderr, "integrity failure") return exitIntegrity } - now := time.Now().UTC() value := record.Record{ SchemaVersion: record.SchemaVersion, Kind: record.KindV1, @@ -218,14 +231,20 @@ func runCreate(ctx context.Context, root string, args []string, stdout, stderr i InstallationID: options.installationID, Description: options.description, SecretSHA256: base64.RawURLEncoding.EncodeToString(material.Digest[:]), - CreatedAt: now, + CreatedAt: createdAt, ExpiresAt: expiresAt, } - addErr := store.Add(value) - closeErr := store.Close() - if addErr != nil || closeErr != nil { - if cleanupOutput(options.output) != nil { - writeCleanupFailure(stderr, options.output) + addErr := addRecord(store, value) + if addErr != nil { + lookupErr := lookupCreatedRecord(store, material.KeyID) + closeErr := closeStore(store) + 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 } 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") 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)) 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 { 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") return exitUsage } @@ -371,7 +394,7 @@ func runRevoke(ctx context.Context, root string, args []string, stdout, stderr i writeLine(stderr, "integrity failure") return exitIntegrity } - err = store.Revoke(options.keyID, options.reason, time.Now().UTC()) + err = store.Revoke(options.keyID, options.reason, nowUTC()) closeErr := store.Close() if errors.Is(err, registry.ErrNotFound) { writeLine(stderr, "not found") @@ -414,27 +437,22 @@ func runCheck(ctx context.Context, root string, args []string, stdout, stderr io return exitOK } -func runServe(root string, args []string, stderr io.Writer) int { - options, err := parseFlags(args, map[string]bool{"--socket": true}) - if err != nil || options.output != "" { +func runServe(args []string, stderr io.Writer) int { + if len(args) != 4 || args[0] != "--registry-root" || args[2] != "--socket" || !filepath.IsAbs(args[1]) || !filepath.IsAbs(args[3]) { writeLine(stderr, "unsafe invocation") return exitUsage } // The Unix-socket service is intentionally introduced by Task 4. Keeping - // this 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 + // this exact grammar reserved prevents accidental plaintext or network fallback. writeLine(stderr, commandServeMsg) 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) { if err := ctx.Err(); err != nil { 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)) } -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 == "" { return nil, nil } 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 &parsed, nil @@ -535,17 +557,38 @@ func validInstallationID(value string) bool { } func validMetadata(value string) bool { - if len([]rune(value)) > maxReasonBytes { + if !utf8.ValidString(value) || len([]rune(value)) > maxReasonBytes { return false } for _, r := range value { - if r < 0x20 || r == 0x7f { + if unicode.IsControl(r) { return false } } 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 { if value == record.LegacyKeyID { return true diff --git a/tools/dwh-auth/internal/command/command_test.go b/tools/dwh-auth/internal/command/command_test.go index d2d0ce24..a9c5eaf3 100644 --- a/tools/dwh-auth/internal/command/command_test.go +++ b/tools/dwh-auth/internal/command/command_test.go @@ -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) { t.Helper() var stdout, stderr bytes.Buffer