fix: harden DWH credential administration CLI

This commit is contained in:
User
2026-08-21 01:19:13 +02:00
parent ebb360f700
commit 055dcabf36
2 changed files with 219 additions and 31 deletions
+74 -31
View File
@@ -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