feat: define DWH installation credentials
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
module github.com/aritmolab/thothii/tools/dwh-auth
|
||||
|
||||
go 1.26.0
|
||||
|
||||
toolchain go1.26.5
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user