Files
libretechandClaude Opus 4.7 6c2e33a3af Implement email magic-link authentication
internal/auth/ provides:
- TokenStore: 32-byte cryptographically random one-time tokens.
  Only the SHA-256 hash is persisted (so a DB leak doesn't grant
  active sessions). Comparison uses subtle.ConstantTimeCompare.
  Single-use is enforced via UPDATE ... WHERE used_at IS NULL.
- Signer: HS256 JWTs with 24h lifetime, jwt.WithValidMethods to
  reject alg=none and other downgrade attacks.
- LogMailer (dev) and SMTPMailer (prod via net/smtp) behind a
  Mailer interface.
- RateLimiter: DB-backed fixed window per email; default 5 per
  15 min for the magic-link flow.
- Service: orchestrates RequestLogin (auto-creates user on first
  login, generates token, emails magic link) and Verify (consumes
  token, updates last_login, issues JWT).
- Handlers: POST /auth/login and GET/POST /auth/verify.
  HandleLogin returns 202 even on validation failure to avoid
  account enumeration; rate-limit hits surface as 429.

Schema additions: magic_tokens (with FK + cascade) and
login_attempts. UserStore.SetStoragePath added for completeness.

Tests cover: token issue/consume, single-use, expiry, rate limit,
JWT round-trip, alg=none rejection, signature tampering, purge,
HTTP handlers (login + verify, missing/invalid token paths).

Closes #9.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-28 22:16:25 +02:00

221 lines
5.8 KiB
Go

package auth
import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"time"
"git.librete.ch/public/librenotes/internal/storage"
)
func newTestService(t *testing.T) (*Service, *bytes.Buffer) {
t.Helper()
db, err := storage.Open(filepath.Join(t.TempDir(), "auth.db"))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
users := storage.NewUserStore(db)
tokens := NewTokenStore(db)
limiter := NewRateLimiter(db, time.Minute, 3)
signer := NewSigner([]byte("test-secret-32-bytes-of-keymaterial!!"))
mailbox := &bytes.Buffer{}
svc, err := NewService(Config{
Users: users,
Tokens: tokens,
Limiter: limiter,
Mailer: LogMailer{W: mailbox},
Signer: signer,
BaseURL: "https://test.example",
DataDir: "/tmp/data",
})
if err != nil {
t.Fatal(err)
}
return svc, mailbox
}
func extractToken(t *testing.T, mailbox string) string {
t.Helper()
idx := strings.Index(mailbox, "token=")
if idx < 0 {
t.Fatalf("no token in mailbox: %q", mailbox)
}
rest := mailbox[idx+len("token="):]
end := strings.IndexAny(rest, " \n")
if end < 0 {
end = len(rest)
}
return strings.TrimSpace(rest[:end])
}
func TestRequestAndVerify(t *testing.T) {
svc, mailbox := newTestService(t)
ctx := context.Background()
if err := svc.RequestLogin(ctx, "Bob@example.com"); err != nil {
t.Fatalf("request: %v", err)
}
tok := extractToken(t, mailbox.String())
res, err := svc.Verify(ctx, tok)
if err != nil {
t.Fatalf("verify: %v", err)
}
if res.JWT == "" || res.User.Email != "bob@example.com" {
t.Errorf("bad result: %+v", res)
}
// Single use.
if _, err := svc.Verify(ctx, tok); !errors.Is(err, ErrInvalidToken) {
t.Errorf("expected ErrInvalidToken on reuse, got %v", err)
}
}
func TestVerifyExpiredToken(t *testing.T) {
svc, mailbox := newTestService(t)
ctx := context.Background()
if err := svc.RequestLogin(ctx, "exp@example.com"); err != nil {
t.Fatal(err)
}
tok := extractToken(t, mailbox.String())
// Force expiry by rewriting expires_at directly.
if _, err := svc.tokens.db.ExecContext(ctx,
`UPDATE magic_tokens SET expires_at = 0`); err != nil {
t.Fatal(err)
}
if _, err := svc.Verify(ctx, tok); !errors.Is(err, ErrInvalidToken) {
t.Errorf("expected ErrInvalidToken, got %v", err)
}
}
func TestRateLimit(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
for i := 0; i < 3; i++ {
if err := svc.RequestLogin(ctx, "rl@example.com"); err != nil {
t.Fatalf("attempt %d: %v", i, err)
}
}
if err := svc.RequestLogin(ctx, "rl@example.com"); !errors.Is(err, ErrRateLimited) {
t.Errorf("expected ErrRateLimited, got %v", err)
}
}
func TestInvalidEmail(t *testing.T) {
svc, _ := newTestService(t)
cases := []string{"", "noatsign", "@nohost", "no@host"}
for _, c := range cases {
if err := svc.RequestLogin(context.Background(), c); err == nil {
t.Errorf("expected error for %q", c)
}
}
}
func TestJWTRoundTrip(t *testing.T) {
s := NewSigner([]byte("k0123456789012345678901234567890"))
tok, err := s.Issue("u-1", "u@example.com")
if err != nil {
t.Fatal(err)
}
claims, err := s.Verify(tok)
if err != nil {
t.Fatal(err)
}
if claims.UserID != "u-1" || claims.Email != "u@example.com" {
t.Errorf("bad claims: %+v", claims)
}
}
func TestJWTRejectsForgedAlg(t *testing.T) {
s := NewSigner([]byte("k0123456789012345678901234567890"))
// alg=none token
bad := "eyJhbGciOiJub25lIiwidHlwIjoiSldUIn0.eyJzdWIiOiJ1LTEifQ."
if _, err := s.Verify(bad); err == nil {
t.Errorf("expected error on alg=none")
}
}
func TestJWTRejectsTamperedSignature(t *testing.T) {
s := NewSigner([]byte("k0123456789012345678901234567890"))
tok, _ := s.Issue("u-1", "u@example.com")
tampered := tok[:len(tok)-2] + "AA"
if _, err := s.Verify(tampered); err == nil {
t.Errorf("expected error on tampered token")
}
}
func TestPurgeExpired(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
_ = svc.RequestLogin(ctx, "purge@example.com")
if _, err := svc.tokens.db.ExecContext(ctx,
`UPDATE magic_tokens SET expires_at = 0`); err != nil {
t.Fatal(err)
}
if err := svc.tokens.PurgeExpired(ctx, time.Minute); err != nil {
t.Fatal(err)
}
var n int
_ = svc.tokens.db.QueryRowContext(ctx,
`SELECT COUNT(*) FROM magic_tokens`).Scan(&n)
if n != 0 {
t.Errorf("expected 0 rows, got %d", n)
}
}
func TestHandleLoginAndVerify(t *testing.T) {
svc, mailbox := newTestService(t)
h := Handlers{Service: svc}
body := strings.NewReader(`{"email":"http@example.com"}`)
req := httptest.NewRequest(http.MethodPost, "/auth/login", body)
rec := httptest.NewRecorder()
h.HandleLogin(rec, req)
if rec.Code != http.StatusAccepted {
t.Fatalf("login status: %d body=%s", rec.Code, rec.Body)
}
tok := extractToken(t, mailbox.String())
req2 := httptest.NewRequest(http.MethodGet, "/auth/verify?token="+tok, nil)
rec2 := httptest.NewRecorder()
h.HandleVerify(rec2, req2)
if rec2.Code != http.StatusOK {
t.Fatalf("verify status: %d body=%s", rec2.Code, rec2.Body)
}
var resp VerifyResponse
if err := json.NewDecoder(rec2.Body).Decode(&resp); err != nil {
t.Fatal(err)
}
if resp.JWT == "" || resp.UserID == "" {
t.Errorf("missing fields: %+v", resp)
}
}
func TestHandleVerifyMissingToken(t *testing.T) {
svc, _ := newTestService(t)
h := Handlers{Service: svc}
req := httptest.NewRequest(http.MethodGet, "/auth/verify", nil)
rec := httptest.NewRecorder()
h.HandleVerify(rec, req)
if rec.Code != http.StatusBadRequest {
t.Errorf("got %d", rec.Code)
}
}
func TestHandleVerifyInvalidToken(t *testing.T) {
svc, _ := newTestService(t)
h := Handlers{Service: svc}
req := httptest.NewRequest(http.MethodGet, "/auth/verify?token=deadbeef", nil)
rec := httptest.NewRecorder()
h.HandleVerify(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Errorf("got %d", rec.Code)
}
}