Files
hammy-backend/internal/api/auth_test.go
T
2026-09-06 15:15:21 +02:00

179 lines
3.9 KiB
Go

package api
import (
"bytes"
"errors"
"strings"
"testing"
)
func TestGenerateRoundTrip(t *testing.T) {
for _, env := range []Environment{Live, Test} {
g, err := Generate(env)
if err != nil {
t.Fatalf("Generate(%s): %v", env, err)
}
got, err := Validate(g.Plaintext)
if err != nil {
t.Fatalf("Validate(%q): %v", g.Plaintext, err)
}
if got != env {
t.Errorf("environment round trip: got %q want %q", got, env)
}
if !strings.HasPrefix(g.Plaintext, Issuer+"_"+string(env)+"_") {
t.Errorf("prefix wrong: %q", g.Plaintext)
}
if len(g.Hash) != 32 {
t.Errorf("hash length %d, want 32 (the api.keys CHECK requires it)", len(g.Hash))
}
if !bytes.Equal(g.Hash, Hash(g.Plaintext)) {
t.Error("Generate's hash differs from Hash of its own plaintext")
}
if len(g.Display) != displayLen {
t.Errorf("display length %d, want %d", len(g.Display), displayLen)
}
if strings.Contains(g.Plaintext[len(g.Display):], "") && g.Display == g.Plaintext {
t.Error("display prefix is the whole key")
}
}
}
func TestGenerateIsUnique(t *testing.T) {
const n = 2000
seen := make(map[string]bool, n)
for i := 0; i < n; i++ {
g, err := Generate(Live)
if err != nil {
t.Fatalf("Generate: %v", err)
}
if seen[g.Plaintext] {
t.Fatalf("duplicate key after %d generations", i)
}
seen[g.Plaintext] = true
}
}
func TestGenerateRejectsUnknownEnvironment(t *testing.T) {
if _, err := Generate(Environment("prod")); !errors.Is(err, ErrEnvironment) {
t.Errorf("got %v, want ErrEnvironment", err)
}
}
func TestValidateRejectsMalformed(t *testing.T) {
good, err := Generate(Live)
if err != nil {
t.Fatal(err)
}
cases := []struct {
name string
key string
want error
}{
{"empty", "", ErrMalformed},
{"no separators", "notakey", ErrMalformed},
{"one separator", "hmy_live", ErrMalformed},
{"wrong issuer", "xxx_live_" + good.Plaintext[9:], ErrMalformed},
{"unknown env", "hmy_prod_" + good.Plaintext[9:], ErrEnvironment},
{"body too short", "hmy_live_abc", ErrMalformed},
{"body too long", good.Plaintext + "X", ErrMalformed},
{"body truncated by one", good.Plaintext[:len(good.Plaintext)-1], ErrMalformed},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
if _, err := Validate(c.key); !errors.Is(err, c.want) {
t.Errorf("Validate(%q) = %v, want %v", c.key, err, c.want)
}
})
}
}
// The point of the checksum: a key with a typo is rejected without a database
// round trip.
func TestValidateCatchesTampering(t *testing.T) {
g, err := Generate(Live)
if err != nil {
t.Fatal(err)
}
body := []byte(g.Plaintext)
caught := 0
// Flip one character at each position in the secret and confirm the
// checksum notices.
for i := len("hmy_live_"); i < len(body)-checksumLen; i++ {
orig := body[i]
if orig == 'A' {
body[i] = 'B'
} else {
body[i] = 'A'
}
if _, err := Validate(string(body)); errors.Is(err, ErrBadChecksum) {
caught++
}
body[i] = orig
}
total := len(body) - checksumLen - len("hmy_live_")
if caught != total {
t.Errorf("checksum caught %d/%d single-character corruptions", caught, total)
}
}
func TestHashIsStable(t *testing.T) {
const key = "hmy_live_0123456789abcdefghijklmnopqrstuvwxyzABCDEFGabcdef"
a, b := Hash(key), Hash(key)
if !bytes.Equal(a, b) {
t.Error("Hash is not deterministic")
}
if bytes.Equal(a, Hash(key+"x")) {
t.Error("Hash collided on different input")
}
}
func TestRedact(t *testing.T) {
g, err := Generate(Live)
if err != nil {
t.Fatal(err)
}
r := Redact(g.Plaintext)
if strings.Contains(g.Plaintext, r) && len(r) >= len(g.Plaintext) {
t.Error("Redact returned the whole key")
}
if len(r) > displayLen+3 {
t.Errorf("Redact returned %d chars, too much", len(r))
}
// A secret must never survive redaction.
secret := g.Plaintext[len("hmy_live_"):]
if strings.Contains(r, secret) {
t.Error("Redact leaked the secret")
}
if got := Redact("short"); got != "hmy_***" {
t.Errorf("Redact(short) = %q", got)
}
}