diff --git a/tools/dwh-auth/go.mod b/tools/dwh-auth/go.mod new file mode 100644 index 00000000..34187196 --- /dev/null +++ b/tools/dwh-auth/go.mod @@ -0,0 +1,5 @@ +module github.com/aritmolab/thothii/tools/dwh-auth + +go 1.26.0 + +toolchain go1.26.5 diff --git a/tools/dwh-auth/internal/credential/credential.go b/tools/dwh-auth/internal/credential/credential.go new file mode 100644 index 00000000..8b544be4 --- /dev/null +++ b/tools/dwh-auth/internal/credential/credential.go @@ -0,0 +1,78 @@ +// Package credential generates and verifies DWH installation credentials. +package credential + +import ( + "crypto/sha256" + "crypto/subtle" + "encoding/base64" + "io" + "strings" + + "github.com/aritmolab/thothii/tools/dwh-auth/internal/record" +) + +const Prefix = "thtdwh_v1" +const KeyIDEncodedLength = 16 +const SecretEncodedLength = 43 +const MaxHeaderBytes = 128 + +type Material struct { + Value []byte + KeyID string + Digest record.Digest +} + +// Generate creates one canonical version 1 credential from random. +func Generate(random io.Reader) (Material, error) { + keyIDBytes := make([]byte, 12) + secretBytes := make([]byte, 32) + if _, err := io.ReadFull(random, keyIDBytes); err != nil { + return Material{}, err + } + if _, err := io.ReadFull(random, secretBytes); err != nil { + return Material{}, err + } + keyID := base64.RawURLEncoding.EncodeToString(keyIDBytes) + secret := base64.RawURLEncoding.EncodeToString(secretBytes) + value := []byte(Prefix + "." + keyID + "." + secret) + sum := sha256.Sum256(value) + return Material{Value: value, KeyID: keyID, Digest: record.Digest(sum)}, nil +} + +// VerifyV1 authenticates one canonical version 1 credential and returns its key ID. +func VerifyV1(value []byte, expected record.Digest) (string, bool) { + if len(value) > MaxHeaderBytes || len(value) != len(Prefix)+1+KeyIDEncodedLength+1+SecretEncodedLength { + return "", false + } + parts := strings.Split(string(value), ".") + if len(parts) != 3 || parts[0] != Prefix || !canonicalBase64(parts[1], 12, KeyIDEncodedLength) || !canonicalBase64(parts[2], 32, SecretEncodedLength) { + return "", false + } + sum := sha256.Sum256(value) + if subtle.ConstantTimeCompare(sum[:], expected[:]) != 1 { + return "", false + } + return parts[1], true +} + +// VerifyLegacy authenticates an opaque legacy credential without interpreting its syntax. +func VerifyLegacy(value []byte, expected record.Digest) bool { + if len(value) < 1 || len(value) > MaxHeaderBytes { + return false + } + for _, b := range value { + if b <= 0x1f || b == 0x7f { + return false + } + } + sum := sha256.Sum256(value) + return subtle.ConstantTimeCompare(sum[:], expected[:]) == 1 +} + +func canonicalBase64(value string, decodedLength, encodedLength int) bool { + if len(value) != encodedLength { + return false + } + decoded, err := base64.RawURLEncoding.DecodeString(value) + return err == nil && len(decoded) == decodedLength && base64.RawURLEncoding.EncodeToString(decoded) == value +} diff --git a/tools/dwh-auth/internal/credential/credential_test.go b/tools/dwh-auth/internal/credential/credential_test.go new file mode 100644 index 00000000..ba4de877 --- /dev/null +++ b/tools/dwh-auth/internal/credential/credential_test.go @@ -0,0 +1,116 @@ +package credential + +import ( + "bytes" + "crypto/sha256" + "encoding/base64" + "strings" + "testing" + + "github.com/aritmolab/thothii/tools/dwh-auth/internal/record" +) + +func TestGenerateProducesCanonicalV1Material(t *testing.T) { + random := bytes.NewReader(bytes.Repeat([]byte{0x7f}, 44)) + material, err := Generate(random) + if err != nil { + t.Fatalf("Generate() error = %v", err) + } + + if got, want := string(material.Value), "thtdwh_v1.f39_f39_f39_f39_.f39_f39_f39_f39_f39_f39_f39_f39_f39_f39_f38"; got != want { + t.Fatalf("Value = %q, want %q", got, want) + } + if got, want := len(material.Value), 70; got != want { + t.Fatalf("len(Value) = %d, want %d", got, want) + } + if got, want := material.KeyID, "f39_f39_f39_f39_"; got != want { + t.Fatalf("KeyID = %q, want %q", got, want) + } + + segments := strings.Split(string(material.Value), ".") + if got, want := len(segments), 3; got != want { + t.Fatalf("segment count = %d, want %d", got, want) + } + if got, want := segments[0], Prefix; got != want { + t.Fatalf("prefix = %q, want %q", got, want) + } + if got, want := len(segments[1]), KeyIDEncodedLength; got != want { + t.Fatalf("key ID length = %d, want %d", got, want) + } + if got, want := len(segments[2]), SecretEncodedLength; got != want { + t.Fatalf("secret length = %d, want %d", got, want) + } + secret, err := base64.RawURLEncoding.DecodeString(segments[2]) + if err != nil { + t.Fatalf("DecodeString(secret) error = %v", err) + } + if got, want := len(secret), 32; got != want { + t.Fatalf("decoded secret length = %d, want %d", got, want) + } + + wantDigest := sha256.Sum256(material.Value) + if material.Digest != record.Digest(wantDigest) { + t.Fatalf("Digest = %x, want %x", material.Digest, wantDigest) + } +} + +func TestVerifyV1AcceptsGeneratedMaterialAndRejectsChangedByte(t *testing.T) { + material, err := Generate(bytes.NewReader(bytes.Repeat([]byte{0x01}, 44))) + if err != nil { + t.Fatalf("Generate() error = %v", err) + } + + if keyID, ok := VerifyV1(material.Value, material.Digest); !ok || keyID != material.KeyID { + t.Fatalf("VerifyV1() = (%q, %t), want (%q, true)", keyID, ok, material.KeyID) + } + + changed := append([]byte(nil), material.Value...) + changed[len(changed)-1] ^= 1 + if keyID, ok := VerifyV1(changed, material.Digest); ok || keyID != "" { + t.Fatalf("VerifyV1(changed) = (%q, %t), want (\"\", false)", keyID, ok) + } +} + +func TestVerifyV1RejectsNonCanonicalValues(t *testing.T) { + material, err := Generate(bytes.NewReader(bytes.Repeat([]byte{0x02}, 44))) + if err != nil { + t.Fatalf("Generate() error = %v", err) + } + + for _, value := range [][]byte{ + append(append([]byte(nil), material.Value...), ','), + []byte(string(material.Value) + "," + string(material.Value)), + []byte(string(material.Value) + "." + material.KeyID), + []byte(string(material.Value) + "="), + []byte(" " + string(material.Value)), + append(bytes.Repeat([]byte{'a'}, MaxHeaderBytes+1)), + } { + if keyID, ok := VerifyV1(value, material.Digest); ok || keyID != "" { + t.Fatalf("VerifyV1(%q) = (%q, %t), want (\"\", false)", value, keyID, ok) + } + } +} + +func TestVerifyLegacyHashesOpaqueRawValueWithoutV1Parsing(t *testing.T) { + value := []byte("legacy opaque value /:[]") + sum := sha256.Sum256(value) + if !VerifyLegacy(value, record.Digest(sum)) { + t.Fatal("VerifyLegacy() = false, want true") + } + if VerifyLegacy([]byte("thtdwh_v1.not-a-valid-v1"), record.Digest(sum)) { + t.Fatal("VerifyLegacy() accepted a different raw value") + } +} + +func TestVerifyLegacyRejectsInvalidOpaqueValues(t *testing.T) { + sum := sha256.Sum256([]byte("legacy")) + for _, value := range [][]byte{ + nil, + bytes.Repeat([]byte{'a'}, MaxHeaderBytes+1), + []byte("bad\nvalue"), + } { + if VerifyLegacy(value, record.Digest(sum)) { + t.Fatalf("VerifyLegacy(%q) = true, want false", value) + } + } +} diff --git a/tools/dwh-auth/internal/record/record.go b/tools/dwh-auth/internal/record/record.go new file mode 100644 index 00000000..414e242b --- /dev/null +++ b/tools/dwh-auth/internal/record/record.go @@ -0,0 +1,133 @@ +// Package record defines the persisted credential-record contract. +package record + +import ( + "encoding/base64" + "errors" + "fmt" + "strings" + "time" + "unicode" + "unicode/utf8" +) + +type Kind string + +const ( + SchemaVersion = 1 + KindV1 Kind = "per_installation_v1" + KindLegacyRaw Kind = "legacy_raw" + LegacyKeyID = "legacy-shared" +) + +type Digest [32]byte + +type Record struct { + SchemaVersion int `json:"schema_version"` + Kind Kind `json:"credential_kind"` + KeyID string `json:"key_id"` + InstallationID string `json:"installation_id"` + Description string `json:"description,omitempty"` + SecretSHA256 string `json:"secret_sha256"` + CreatedAt time.Time `json:"created_at"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` + RevokedAt *time.Time `json:"revoked_at,omitempty"` + RevocationReason string `json:"revocation_reason,omitempty"` +} + +// Validate verifies that r conforms to the version 1 persisted-record contract. +func Validate(r Record) error { + if r.SchemaVersion != SchemaVersion { + return fmt.Errorf("unsupported schema version %d", r.SchemaVersion) + } + if err := validateKindAndKeyID(r.Kind, r.KeyID); err != nil { + return err + } + if !validInstallationID(r.InstallationID) { + return errors.New("invalid installation ID") + } + if err := validateMetadata("description", r.Description); err != nil { + return err + } + if !validDigest(r.SecretSHA256) { + return errors.New("invalid secret SHA-256 digest") + } + if !validTimestamp(r.CreatedAt) { + return errors.New("invalid creation timestamp") + } + if r.ExpiresAt != nil { + if !validTimestamp(*r.ExpiresAt) { + return errors.New("invalid expiry timestamp") + } + if !r.ExpiresAt.After(r.CreatedAt) { + return errors.New("expiry must be after creation") + } + } + if (r.RevokedAt == nil) != (r.RevocationReason == "") { + return errors.New("revocation timestamp and reason must be paired") + } + if r.RevokedAt != nil && !validTimestamp(*r.RevokedAt) { + return errors.New("invalid revocation timestamp") + } + return validateMetadata("revocation reason", r.RevocationReason) +} + +func validateKindAndKeyID(kind Kind, keyID string) error { + switch kind { + case KindV1: + if !validV1KeyID(keyID) { + return errors.New("invalid v1 key ID") + } + case KindLegacyRaw: + if keyID != LegacyKeyID { + return errors.New("legacy record must use the legacy key ID") + } + default: + return fmt.Errorf("invalid credential kind %q", kind) + } + return nil +} + +func validV1KeyID(value string) bool { + if len(value) != 16 { + return false + } + decoded, err := base64.RawURLEncoding.DecodeString(value) + return err == nil && len(decoded) == 12 && base64.RawURLEncoding.EncodeToString(decoded) == value +} + +func validInstallationID(value string) bool { + if len(value) < 1 || len(value) > 63 { + return false + } + for i := range len(value) { + c := value[i] + if (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || (i > 0 && c == '-') { + continue + } + return false + } + return true +} + +func validDigest(value string) bool { + decoded, err := base64.RawURLEncoding.DecodeString(value) + return err == nil && len(decoded) == 32 && base64.RawURLEncoding.EncodeToString(decoded) == value +} + +func validTimestamp(value time.Time) bool { + return !value.IsZero() && value.Location() == time.UTC +} + +func validateMetadata(name, value string) error { + if !utf8.ValidString(value) { + return fmt.Errorf("%s is not valid UTF-8", name) + } + if utf8.RuneCountInString(value) > 160 { + return fmt.Errorf("%s exceeds 160 characters", name) + } + if strings.IndexFunc(value, unicode.IsControl) >= 0 { + return fmt.Errorf("%s contains a control character", name) + } + return nil +} diff --git a/tools/dwh-auth/internal/record/record_test.go b/tools/dwh-auth/internal/record/record_test.go new file mode 100644 index 00000000..150e924a --- /dev/null +++ b/tools/dwh-auth/internal/record/record_test.go @@ -0,0 +1,120 @@ +package record + +import ( + "crypto/sha256" + "encoding/base64" + "strings" + "testing" + "time" +) + +func validRecord() Record { + created := time.Date(2026, 8, 20, 12, 0, 0, 0, time.UTC) + sum := sha256.Sum256([]byte("synthetic credential")) + return Record{ + SchemaVersion: SchemaVersion, + Kind: KindV1, + KeyID: "AAAAAAAAAAAAAAAA", + InstallationID: "psd-clinical", + Description: "Synthetic test credential", + SecretSHA256: base64.RawURLEncoding.EncodeToString(sum[:]), + CreatedAt: created, + } +} + +func TestValidateAcceptsV1AndLegacyRecords(t *testing.T) { + if err := Validate(validRecord()); err != nil { + t.Fatalf("Validate(v1) error = %v", err) + } + + legacy := validRecord() + legacy.Kind = KindLegacyRaw + legacy.KeyID = LegacyKeyID + if err := Validate(legacy); err != nil { + t.Fatalf("Validate(legacy) error = %v", err) + } +} + +func TestValidateRejectsInvalidKindsAndKeyRelationships(t *testing.T) { + for name, mutate := range map[string]func(*Record){ + "unknown kind": func(r *Record) { r.Kind = "other" }, + "v1 legacy key": func(r *Record) { r.KeyID = LegacyKeyID }, + "v1 malformed key": func(r *Record) { r.KeyID = "short" }, + "legacy nonlegacy key": func(r *Record) { r.Kind = KindLegacyRaw }, + } { + t.Run(name, func(t *testing.T) { + r := validRecord() + mutate(&r) + if err := Validate(r); err == nil { + t.Fatal("Validate() error = nil, want validation error") + } + }) + } +} + +func TestValidateRejectsInvalidSchemaIdentifierAndDigest(t *testing.T) { + for name, mutate := range map[string]func(*Record){ + "schema": func(r *Record) { r.SchemaVersion = 2 }, + "identifier upper": func(r *Record) { r.InstallationID = "PSD" }, + "identifier leading dash": func(r *Record) { r.InstallationID = "-psd" }, + "identifier too long": func(r *Record) { r.InstallationID = "a" + strings.Repeat("b", 63) }, + "digest padding": func(r *Record) { r.SecretSHA256 += "=" }, + "digest short": func(r *Record) { r.SecretSHA256 = "AAAA" }, + } { + t.Run(name, func(t *testing.T) { + r := validRecord() + mutate(&r) + if err := Validate(r); err == nil { + t.Fatal("Validate() error = nil, want validation error") + } + }) + } +} + +func TestValidateRejectsInvalidMetadata(t *testing.T) { + for name, mutate := range map[string]func(*Record){ + "description control": func(r *Record) { r.Description = "bad\nmetadata" }, + "description too long": func(r *Record) { r.Description = strings.Repeat("é", 161) }, + "reason control": func(r *Record) { + now := r.CreatedAt.Add(time.Hour) + r.RevokedAt = &now + r.RevocationReason = "bad\tmetadata" + }, + "reason too long": func(r *Record) { + now := r.CreatedAt.Add(time.Hour) + r.RevokedAt = &now + r.RevocationReason = strings.Repeat("é", 161) + }, + } { + t.Run(name, func(t *testing.T) { + r := validRecord() + mutate(&r) + if err := Validate(r); err == nil { + t.Fatal("Validate() error = nil, want validation error") + } + }) + } +} + +func TestValidateRejectsInvalidTimesExpiryAndRevocation(t *testing.T) { + for name, mutate := range map[string]func(*Record){ + "created non UTC": func(r *Record) { r.CreatedAt = r.CreatedAt.In(time.FixedZone("CET", 3600)) }, + "expiry non UTC": func(r *Record) { at := r.CreatedAt.Add(time.Hour).In(time.FixedZone("CET", 3600)); r.ExpiresAt = &at }, + "expiry equals created": func(r *Record) { at := r.CreatedAt; r.ExpiresAt = &at }, + "revoked without reason": func(r *Record) { at := r.CreatedAt.Add(time.Hour); r.RevokedAt = &at }, + "reason without revoked": func(r *Record) { r.RevocationReason = "synthetic" }, + "revoked non UTC": func(r *Record) { + at := r.CreatedAt.Add(time.Hour).In(time.FixedZone("CET", 3600)) + r.RevokedAt = &at + r.RevocationReason = "synthetic" + }, + } { + t.Run(name, func(t *testing.T) { + r := validRecord() + mutate(&r) + if err := Validate(r); err == nil { + t.Fatal("Validate() error = nil, want validation error") + } + }) + } +}