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) } }