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>
221 lines
5.8 KiB
Go
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)
|
|
}
|
|
}
|