Files
2026-08-15 07:13:00 +00:00

408 lines
10 KiB
Go

package auth
import (
"context"
"errors"
"io"
"log/slog"
"path/filepath"
"sync"
"testing"
"time"
"github.com/iceBear67/simplepages/internal/store"
)
func testStore(t *testing.T) *store.DB {
t.Helper()
log := slog.New(slog.NewTextHandler(io.Discard, nil))
db, err := store.Open(context.Background(), filepath.Join(t.TempDir(), "pages.db"), log)
if err != nil {
t.Fatalf("store.Open: %v", err)
}
t.Cleanup(func() { db.Close() })
return db
}
// mintInto creates a real key in the database and returns its token.
func mintInto(t *testing.T, db *store.DB, scope store.Scope, projectID *int64) (token, keyID string) {
t.Helper()
token, keyID, hash, err := Mint()
if err != nil {
t.Fatal(err)
}
k := &store.APIKey{
ID: keyID,
SecretHash: hash[:],
Scope: scope,
ProjectID: projectID,
Name: "test",
}
if err := db.CreateKey(context.Background(), k); err != nil {
t.Fatalf("CreateKey: %v", err)
}
return token, keyID
}
func newVerifier(t *testing.T, db *store.DB) *Verifier {
t.Helper()
return NewVerifier(db, slog.New(slog.NewTextHandler(io.Discard, nil)), DefaultCacheTTL)
}
func TestVerifyAdminKey(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
token, keyID := mintInto(t, db, store.ScopeAdmin, nil)
id, err := v.Verify(ctx, token)
if err != nil {
t.Fatalf("Verify: %v", err)
}
if id.KeyID != keyID {
t.Errorf("KeyID = %q, want %q", id.KeyID, keyID)
}
if !id.IsAdmin() {
t.Error("admin key did not produce an admin identity")
}
if !id.Owns(1) || !id.Owns(999) {
t.Error("an admin must own every project")
}
if id.ProjectID != nil {
t.Errorf("admin identity carries project %v", *id.ProjectID)
}
}
func TestVerifyProjectKey(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
p := store.DefaultProject("demo")
if err := db.CreateProject(ctx, p); err != nil {
t.Fatal(err)
}
token, _ := mintInto(t, db, store.ScopeProject, &p.ID)
id, err := v.Verify(ctx, token)
if err != nil {
t.Fatalf("Verify: %v", err)
}
if id.IsAdmin() {
t.Error("project key produced an admin identity")
}
if !id.Owns(p.ID) {
t.Error("project key does not own its own project")
}
if id.Owns(p.ID + 1) {
t.Error("project key owns someone else's project")
}
}
func TestVerifyRejects(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
token, keyID := mintInto(t, db, store.ScopeAdmin, nil)
_, secret, err := Parse(token)
if err != nil {
t.Fatal(err)
}
// A well-formed token for a key that was never created.
otherToken, _, _, err := Mint()
if err != nil {
t.Fatal(err)
}
// The right key id with someone else's secret.
strangerToken, _, _, err := Mint()
if err != nil {
t.Fatal(err)
}
_, wrong, err := Parse(strangerToken)
if err != nil {
t.Fatal(err)
}
cases := []struct {
name string
token string
want error
}{
{"garbage", "not-a-token", ErrMalformedToken},
{"unknown key", otherToken, ErrUnknownKey},
{"wrong secret", Prefix + "_" + keyID + "_" + wrong, ErrBadSecret},
{"empty", "", ErrMalformedToken},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if _, err := v.Verify(ctx, tc.token); !errors.Is(err, tc.want) {
t.Errorf("Verify: got %v, want %v", err, tc.want)
}
})
}
// Sanity: the real token still works after all those failures.
if _, err := v.Verify(ctx, Prefix+"_"+keyID+"_"+secret); err != nil {
t.Errorf("valid token rejected after failed attempts: %v", err)
}
}
func TestVerifyRejectsExpiredKey(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
token, keyID, hash, err := Mint()
if err != nil {
t.Fatal(err)
}
past := time.Now().Add(-time.Hour)
if err := db.CreateKey(ctx, &store.APIKey{
ID: keyID, SecretHash: hash[:], Scope: store.ScopeAdmin, ExpiresAt: &past,
}); err != nil {
t.Fatal(err)
}
if _, err := v.Verify(ctx, token); !errors.Is(err, ErrExpired) {
t.Errorf("got %v, want ErrExpired", err)
}
}
// Expiry must be re-evaluated on every call, not frozen into the cache entry,
// or a key cached one second before it expires would stay valid for the whole
// cache TTL.
func TestCachedKeyStillExpires(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
now := time.Now()
fake := now
var mu sync.Mutex
v.now = func() time.Time {
mu.Lock()
defer mu.Unlock()
return fake
}
token, keyID, hash, err := Mint()
if err != nil {
t.Fatal(err)
}
exp := now.Add(30 * time.Second)
if err := db.CreateKey(ctx, &store.APIKey{
ID: keyID, SecretHash: hash[:], Scope: store.ScopeAdmin, ExpiresAt: &exp,
}); err != nil {
t.Fatal(err)
}
if _, err := v.Verify(ctx, token); err != nil {
t.Fatalf("key should be valid before expiry: %v", err)
}
mu.Lock()
fake = now.Add(31 * time.Second) // still inside the 60s cache TTL
mu.Unlock()
if _, err := v.Verify(ctx, token); !errors.Is(err, ErrExpired) {
t.Errorf("expired key still accepted from cache: %v", err)
}
}
func TestVerifyRejectsRevokedKey(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
token, keyID := mintInto(t, db, store.ScopeAdmin, nil)
if err := db.RevokeKey(ctx, keyID); err != nil {
t.Fatal(err)
}
if _, err := v.Verify(ctx, token); !errors.Is(err, ErrRevoked) {
t.Errorf("got %v, want ErrRevoked", err)
}
}
// Revocation has to take effect immediately in the process that performed it;
// this both proves that and demonstrates the cache is really being consulted.
func TestInvalidateMakesRevocationImmediate(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
token, keyID := mintInto(t, db, store.ScopeAdmin, nil)
if _, err := v.Verify(ctx, token); err != nil {
t.Fatal(err)
}
// Revoke behind the verifier's back. The cached entry is still live, so the
// key keeps working — which is exactly what makes the next assertion mean
// something.
if err := db.RevokeKey(ctx, keyID); err != nil {
t.Fatal(err)
}
if _, err := v.Verify(ctx, token); err != nil {
t.Fatalf("cache was not consulted (or the test is not measuring it): %v", err)
}
v.Invalidate()
if _, err := v.Verify(ctx, token); !errors.Is(err, ErrRevoked) {
t.Errorf("after Invalidate: got %v, want ErrRevoked", err)
}
}
// An unknown key id must not create a cache entry: the id space is
// attacker-chosen, so caching misses would be an unbounded memory sink.
func TestUnknownKeysAreNotCached(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
for i := 0; i < 200; i++ {
token, _, _, err := Mint()
if err != nil {
t.Fatal(err)
}
if _, err := v.Verify(ctx, token); !errors.Is(err, ErrUnknownKey) {
t.Fatalf("got %v, want ErrUnknownKey", err)
}
}
entries := 0
v.cache.Range(func(any, any) bool { entries++; return true })
if entries != 0 {
t.Errorf("%d unknown key ids were cached", entries)
}
}
// Same idea for revoked keys: a dead credential must not pin memory.
func TestRevokedKeysAreNotCached(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
token, keyID := mintInto(t, db, store.ScopeAdmin, nil)
if err := db.RevokeKey(ctx, keyID); err != nil {
t.Fatal(err)
}
for i := 0; i < 10; i++ {
if _, err := v.Verify(ctx, token); !errors.Is(err, ErrRevoked) {
t.Fatal(err)
}
}
entries := 0
v.cache.Range(func(any, any) bool { entries++; return true })
if entries != 0 {
t.Errorf("%d revoked keys were cached", entries)
}
}
func TestFlushTouches(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
token, keyID := mintInto(t, db, store.ScopeAdmin, nil)
if err := v.FlushTouches(ctx); err != nil {
t.Errorf("flushing an empty batch: %v", err)
}
before, err := db.KeyByID(ctx, keyID)
if err != nil {
t.Fatal(err)
}
if before.LastUsedAt != nil {
t.Error("last_used_at set before any use")
}
for i := 0; i < 5; i++ {
if _, err := v.Verify(ctx, token); err != nil {
t.Fatal(err)
}
}
// Nothing is written until the batch is flushed; that is the whole point of
// keeping the write connection out of the request path.
mid, err := db.KeyByID(ctx, keyID)
if err != nil {
t.Fatal(err)
}
if mid.LastUsedAt != nil {
t.Error("last_used_at written per request instead of in a batch")
}
if err := v.FlushTouches(ctx); err != nil {
t.Fatalf("FlushTouches: %v", err)
}
after, err := db.KeyByID(ctx, keyID)
if err != nil {
t.Fatal(err)
}
if after.LastUsedAt == nil {
t.Fatal("last_used_at still unset after flush")
}
// A flush drains the pending set, so a second one has nothing to do.
v.mu.Lock()
pending := len(v.touch)
v.mu.Unlock()
if pending != 0 {
t.Errorf("%d pending touches survived the flush", pending)
}
}
func TestRunFlusherFlushesOnShutdown(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
db := testStore(t)
v := newVerifier(t, db)
token, keyID := mintInto(t, db, store.ScopeAdmin, nil)
if _, err := v.Verify(ctx, token); err != nil {
t.Fatal(err)
}
done := make(chan struct{})
go func() {
defer close(done)
v.RunFlusher(ctx, time.Hour) // never ticks; only the shutdown path runs
}()
cancel()
<-done
k, err := db.KeyByID(context.Background(), keyID)
if err != nil {
t.Fatal(err)
}
if k.LastUsedAt == nil {
t.Error("shutdown flush dropped the pending timestamps")
}
}
func TestVerifyIsConcurrencySafe(t *testing.T) {
ctx := context.Background()
db := testStore(t)
v := newVerifier(t, db)
token, _ := mintInto(t, db, store.ScopeAdmin, nil)
bad, _, _, err := Mint()
if err != nil {
t.Fatal(err)
}
var wg sync.WaitGroup
for g := 0; g < 16; g++ {
wg.Add(1)
go func(g int) {
defer wg.Done()
for i := 0; i < 100; i++ {
if g%4 == 0 {
if _, err := v.Verify(ctx, bad); err == nil {
t.Error("bad token accepted")
}
continue
}
if _, err := v.Verify(ctx, token); err != nil {
t.Errorf("good token rejected: %v", err)
return
}
if i%25 == 0 {
v.Invalidate()
}
}
}(g)
}
wg.Wait()
}