From a5e25b444f5d381809bb48955d94083e4d942403 Mon Sep 17 00:00:00 2001 From: wangjia <809946525@qq.com> Date: Sat, 13 Jun 2026 12:19:10 +0800 Subject: [PATCH] =?UTF-8?q?feat(auth):=20=E9=AA=8C=E8=AF=81=E7=A0=81/?= =?UTF-8?q?=E6=B3=A8=E5=86=8C/=E7=99=BB=E5=BD=95/JWT=20=E9=89=B4=E6=9D=83?= =?UTF-8?q?=E6=A8=A1=E5=9D=97=20(tsk=5F2PFfyviECIXh)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 实现 server/internal/auth 控制面鉴权底座,对应 doc/02 §4.1、doc/06 §3: - POST /v1/auth/code:IP+邮箱双维度 Redis 滑窗限频(同邮箱 1/min、同 IP 10/h) + 一次性邮箱域黑名单(embed 词表)→ 6 位数字码写 auth:code:{email} TTL 10min → 异步发信(SMTP / 开发态 log mailer)。 - POST /v1/auth/register:验码(一次性,删 key;超次数烧码)→ 单事务建号 (uuid + dp_uuid 应用层生成、argon2id)+ 7 天 PRO 试用(source='trial')→ 签发 JWT。 - POST /v1/auth/login:argon2id 校验(失败恒定时、未知用户走 dummy hash); 失败计数限流 rl:login:{email} + 锁定;banned 拒绝。 - POST /v1/auth/refresh:RS256,access 15min + refresh 30d 落 Redis 白名单 jwt:refresh:{jti},旋转时删旧写新;JWT 头带 kid,验证端接受新旧公钥支持轮换。 - RequireAuth 中间件:解析 Bearer access,注入 user id(复用 codes.CtxKeyUserID 避免循环依赖)+ uuid + claims 到 context。 - 错误体统一走 internal/apierr 的 {code, message_zh, message_en},文案双语脱敏。 - 文件拆分:handler/service/password/token/ratelimit/emailcheck/store/mailer/ middleware/errors/keyloader;config 增加可选 JWT PEM 路径/kid 加载。 测试:password/token/ratelimit/emailcheck/service/handler/middleware 单测 (miniredis,含注册全流程、重复邮箱 409、验证码错误/过期/复用/烧码、限频 429 带 Retry-After、登录失败锁定、banned 拒绝、refresh 旋转后旧 token 失效), go test -race 通过;集成测试 integration_test.go(testcontainers MySQL8+Redis, register→login→refresh→受保护接口全链路,构建 tag integration)与 codes 模块同构。 Co-Authored-By: Claude Opus 4.8 --- server/go.mod | 6 +- server/internal/auth/disposable_domains.txt | 63 +++++ server/internal/auth/emailcheck.go | 68 +++++ server/internal/auth/emailcheck_test.go | 39 +++ server/internal/auth/errors.go | 86 ++++++ server/internal/auth/handler.go | 191 +++++++++++++ server/internal/auth/handler_test.go | 148 ++++++++++ server/internal/auth/helpers_test.go | 140 ++++++++++ server/internal/auth/integration_test.go | 217 +++++++++++++++ server/internal/auth/keyloader.go | 49 ++++ server/internal/auth/mailer.go | 91 ++++++ server/internal/auth/middleware.go | 86 ++++++ server/internal/auth/password.go | 111 ++++++++ server/internal/auth/password_test.go | 58 ++++ server/internal/auth/ratelimit.go | 161 +++++++++++ server/internal/auth/ratelimit_test.go | 97 +++++++ server/internal/auth/service.go | 292 ++++++++++++++++++++ server/internal/auth/service_test.go | 271 ++++++++++++++++++ server/internal/auth/store.go | 133 +++++++++ server/internal/auth/token.go | 290 +++++++++++++++++++ server/internal/auth/token_test.go | 127 +++++++++ server/internal/config/config.go | 35 +++ 22 files changed, 2756 insertions(+), 3 deletions(-) create mode 100644 server/internal/auth/disposable_domains.txt create mode 100644 server/internal/auth/emailcheck.go create mode 100644 server/internal/auth/emailcheck_test.go create mode 100644 server/internal/auth/errors.go create mode 100644 server/internal/auth/handler.go create mode 100644 server/internal/auth/handler_test.go create mode 100644 server/internal/auth/helpers_test.go create mode 100644 server/internal/auth/integration_test.go create mode 100644 server/internal/auth/keyloader.go create mode 100644 server/internal/auth/mailer.go create mode 100644 server/internal/auth/middleware.go create mode 100644 server/internal/auth/password.go create mode 100644 server/internal/auth/password_test.go create mode 100644 server/internal/auth/ratelimit.go create mode 100644 server/internal/auth/ratelimit_test.go create mode 100644 server/internal/auth/service.go create mode 100644 server/internal/auth/service_test.go create mode 100644 server/internal/auth/store.go create mode 100644 server/internal/auth/token.go create mode 100644 server/internal/auth/token_test.go diff --git a/server/go.mod b/server/go.mod index 3a95d70..3280982 100644 --- a/server/go.mod +++ b/server/go.mod @@ -6,12 +6,15 @@ require ( github.com/alicebob/miniredis/v2 v2.38.0 github.com/go-chi/chi/v5 v5.2.1 github.com/go-sql-driver/mysql v1.8.1 + github.com/golang-jwt/jwt/v4 v4.5.2 github.com/golang-migrate/migrate/v4 v4.19.1 + github.com/google/uuid v1.6.0 github.com/oapi-codegen/oapi-codegen/v2 v2.7.1 github.com/redis/go-redis/v9 v9.20.1 github.com/testcontainers/testcontainers-go v0.34.0 github.com/testcontainers/testcontainers-go/modules/mysql v0.34.0 github.com/testcontainers/testcontainers-go/modules/redis v0.34.0 + golang.org/x/crypto v0.48.0 google.golang.org/grpc v1.81.1 ) @@ -100,7 +103,6 @@ require ( github.com/gocql/gocql v0.0.0-20210515062232-b7ef815b4556 // indirect github.com/godbus/dbus v0.0.0-20190726142602-4481cbc300e2 // indirect github.com/gogo/protobuf v1.3.2 // indirect - github.com/golang-jwt/jwt/v4 v4.5.2 // indirect github.com/golang-sql/civil v0.0.0-20190719163853-cb61b32ac6fe // indirect github.com/golang-sql/sqlexp v0.1.0 // indirect github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da // indirect @@ -109,7 +111,6 @@ require ( github.com/google/go-github/v39 v39.2.0 // indirect github.com/google/go-querystring v1.1.0 // indirect github.com/google/s2a-go v0.1.9 // indirect - github.com/google/uuid v1.6.0 // indirect github.com/googleapis/enterprise-certificate-proxy v0.3.6 // indirect github.com/googleapis/gax-go/v2 v2.15.0 // indirect github.com/gsterjov/go-libsecret v0.0.0-20161001094733-a6f4afe4910c // indirect @@ -208,7 +209,6 @@ require ( go.opentelemetry.io/otel/trace v1.43.0 // indirect go.uber.org/atomic v1.11.0 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect - golang.org/x/crypto v0.48.0 // indirect golang.org/x/mod v0.33.0 // indirect golang.org/x/net v0.51.0 // indirect golang.org/x/oauth2 v0.36.0 // indirect diff --git a/server/internal/auth/disposable_domains.txt b/server/internal/auth/disposable_domains.txt new file mode 100644 index 0000000..9ab0c68 --- /dev/null +++ b/server/internal/auth/disposable_domains.txt @@ -0,0 +1,63 @@ +# Disposable / temporary email domains blocklist. +# One domain per line, lowercase, no leading dot. Lines starting with '#' and +# blank lines are ignored. Extend as new throwaway providers appear. +0-mail.com +10minutemail.com +10minutemail.net +20minutemail.com +33mail.com +guerrillamail.com +guerrillamail.net +guerrillamail.org +guerrillamailblock.com +sharklasers.com +grr.la +spam4.me +mailinator.com +mailinator.net +mailinator2.com +maildrop.cc +mintemail.com +mohmal.com +temp-mail.org +tempmail.com +tempmailo.com +tempr.email +tempinbox.com +throwawaymail.com +trashmail.com +trashmail.net +trashmail.de +dispostable.com +fakeinbox.com +getnada.com +nada.email +yopmail.com +yopmail.net +yopmail.fr +mailnesia.com +emailondeck.com +moakt.com +mytemp.email +tmpmail.org +tmpmail.net +tmpeml.com +inboxkitten.com +burnermail.io +spamgourmet.com +mailcatch.com +fakemailgenerator.com +mvrht.com +discard.email +discardmail.com +mail-temp.com +luxusmail.org +1secmail.com +1secmail.net +1secmail.org +emailfake.com +tempmailaddress.com +mailpoof.com +harakirimail.com +spambox.us +jetable.org diff --git a/server/internal/auth/emailcheck.go b/server/internal/auth/emailcheck.go new file mode 100644 index 0000000..4a1d21d --- /dev/null +++ b/server/internal/auth/emailcheck.go @@ -0,0 +1,68 @@ +package auth + +import ( + _ "embed" + "strings" +) + +//go:embed disposable_domains.txt +var disposableList string + +// disposableDomains is the parsed blocklist, built once at package init. +var disposableDomains = parseDisposable(disposableList) + +// parseDisposable turns the embedded wordlist into a lookup set. +func parseDisposable(raw string) map[string]struct{} { + set := make(map[string]struct{}) + for _, line := range strings.Split(raw, "\n") { + line = strings.TrimSpace(strings.ToLower(line)) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + set[line] = struct{}{} + } + return set +} + +// NormalizeEmail lowercases and trims the address for consistent storage and +// rate-limit keying. It does not alter the local part beyond trimming. +func NormalizeEmail(email string) string { + return strings.ToLower(strings.TrimSpace(email)) +} + +// emailDomain returns the lowercased domain part of an email, or "" if the +// address has no single '@'. +func emailDomain(email string) string { + at := strings.LastIndex(email, "@") + if at < 0 || at == len(email)-1 { + return "" + } + return strings.ToLower(email[at+1:]) +} + +// ValidEmail performs a cheap structural sanity check: exactly one '@', a +// non-empty local part, and a domain containing a dot. It is deliberately +// permissive — true deliverability is proven by the verification code. +func ValidEmail(email string) bool { + email = strings.TrimSpace(email) + if len(email) < 3 || len(email) > 254 { + return false + } + at := strings.IndexByte(email, '@') + if at <= 0 || at != strings.LastIndexByte(email, '@') { + return false + } + domain := email[at+1:] + return strings.Contains(domain, ".") && !strings.HasPrefix(domain, ".") && !strings.HasSuffix(domain, ".") +} + +// IsDisposable reports whether the email's domain is on the disposable-domain +// blocklist. +func IsDisposable(email string) bool { + domain := emailDomain(email) + if domain == "" { + return false + } + _, ok := disposableDomains[domain] + return ok +} diff --git a/server/internal/auth/emailcheck_test.go b/server/internal/auth/emailcheck_test.go new file mode 100644 index 0000000..d6777fa --- /dev/null +++ b/server/internal/auth/emailcheck_test.go @@ -0,0 +1,39 @@ +package auth + +import "testing" + +func TestValidEmail(t *testing.T) { + valid := []string{"a@b.com", "user.name+tag@sub.example.co", "x@y.io"} + for _, e := range valid { + if !ValidEmail(e) { + t.Errorf("expected %q to be valid", e) + } + } + invalid := []string{"", "no-at", "a@", "@b.com", "a@b", "a@@b.com", "a@.com", "a@b."} + for _, e := range invalid { + if ValidEmail(e) { + t.Errorf("expected %q to be invalid", e) + } + } +} + +func TestIsDisposable(t *testing.T) { + if !IsDisposable("foo@mailinator.com") { + t.Error("mailinator.com should be disposable") + } + if !IsDisposable("foo@MAILINATOR.com") { // case-insensitive + t.Error("disposable check should be case-insensitive") + } + if IsDisposable("foo@gmail.com") { + t.Error("gmail.com should not be disposable") + } + if IsDisposable("no-domain") { + t.Error("address without domain should not be flagged disposable") + } +} + +func TestNormalizeEmail(t *testing.T) { + if got := NormalizeEmail(" User@Example.COM "); got != "user@example.com" { + t.Errorf("NormalizeEmail = %q", got) + } +} diff --git a/server/internal/auth/errors.go b/server/internal/auth/errors.go new file mode 100644 index 0000000..cce6842 --- /dev/null +++ b/server/internal/auth/errors.go @@ -0,0 +1,86 @@ +package auth + +import "github.com/wangjia/pangolin/server/internal/apierr" + +// Auth-module error values. Codes use the dotted `auth.*` namespace documented +// in the OpenAPI contract (components.schemas.Error). All messages are bilingual +// and desensitised (no destination / traffic wording, no enumeration leaks). +var ( + // ErrInvalidRequest — malformed body or failed field validation. + ErrInvalidRequest = &apierr.Error{ + Code: "auth.invalid_request", + MessageZH: "请求参数有误,请检查后重试", + MessageEn: "Invalid request parameters, please verify and try again", + } + + // ErrEmailDisposable — the email domain is on the disposable-domain blocklist. + ErrEmailDisposable = &apierr.Error{ + Code: "auth.email_disposable", + MessageZH: "暂不支持该邮箱服务商,请更换邮箱", + MessageEn: "This email provider is not supported, please use another address", + } + + // ErrCodeInvalid — verification code wrong, expired, or already used. + // Intentionally generic to avoid distinguishing the three cases. + ErrCodeInvalid = &apierr.Error{ + Code: "auth.code_invalid", + MessageZH: "验证码无效或已过期,请重新获取", + MessageEn: "Verification code is invalid or expired, please request a new one", + } + + // ErrEmailExists — the email is already registered. + ErrEmailExists = &apierr.Error{ + Code: "auth.email_exists", + MessageZH: "该邮箱已注册,请直接登录", + MessageEn: "This email is already registered, please sign in", + } + + // ErrRateLimited — rate-limit window exceeded. Carries a Retry-After header. + ErrRateLimited = &apierr.Error{ + Code: "auth.rate_limited", + MessageZH: "操作过于频繁,请稍后再试", + MessageEn: "Too many attempts, please try again later", + } + + // ErrInvalidCredentials — wrong email or password (login). + ErrInvalidCredentials = &apierr.Error{ + Code: "auth.invalid_credentials", + MessageZH: "邮箱或密码不正确", + MessageEn: "Incorrect email or password", + } + + // ErrAccountLocked — too many failed logins; temporarily locked. + ErrAccountLocked = &apierr.Error{ + Code: "auth.account_locked", + MessageZH: "登录失败次数过多,账户已临时锁定,请稍后再试", + MessageEn: "Too many failed sign-in attempts, account temporarily locked", + } + + // ErrAccountBanned — the account has been disabled. + ErrAccountBanned = &apierr.Error{ + Code: "auth.account_banned", + MessageZH: "该账户已被停用", + MessageEn: "This account has been disabled", + } + + // ErrInvalidToken — refresh token absent from whitelist, malformed, or expired. + ErrInvalidToken = &apierr.Error{ + Code: "auth.invalid_token", + MessageZH: "登录已失效,请重新登录", + MessageEn: "Your session has expired, please sign in again", + } + + // ErrUnauthorized — missing or invalid bearer access token (middleware). + ErrUnauthorized = &apierr.Error{ + Code: "auth.unauthorized", + MessageZH: "请先登录", + MessageEn: "Authentication required", + } + + // ErrInternal — unexpected server-side failure. + ErrInternal = &apierr.Error{ + Code: "auth.internal", + MessageZH: "服务器内部错误,请稍后重试", + MessageEn: "Internal server error, please try again later", + } +) diff --git a/server/internal/auth/handler.go b/server/internal/auth/handler.go new file mode 100644 index 0000000..99c0d71 --- /dev/null +++ b/server/internal/auth/handler.go @@ -0,0 +1,191 @@ +package auth + +import ( + "encoding/json" + "net" + "net/http" + "strconv" + "strings" + "time" + + "github.com/go-chi/chi/v5" + "github.com/wangjia/pangolin/server/internal/apierr" +) + +// Handler adapts the auth Service to HTTP. Route shapes follow the OpenAPI +// contract (POST /v1/auth/{code,register,login,refresh}). +type Handler struct { + svc *Service +} + +// NewHandler builds a Handler. +func NewHandler(svc *Service) *Handler { return &Handler{svc: svc} } + +// RegisterRoutes mounts the four public auth endpoints onto r. The caller is +// expected to mount this group WITHOUT the bearer-auth middleware. +func (h *Handler) RegisterRoutes(r chi.Router) { + r.Post("/auth/code", h.SendCode) + r.Post("/auth/register", h.Register) + r.Post("/auth/login", h.Login) + r.Post("/auth/refresh", h.Refresh) +} + +// ---- request/response bodies (mirror openapi.yaml) ---- + +type sendCodeRequest struct { + Email string `json:"email"` +} + +type registerRequest struct { + Email string `json:"email"` + Code string `json:"code"` + Password string `json:"password"` +} + +type loginRequest struct { + Email string `json:"email"` + Password string `json:"password"` +} + +type refreshRequest struct { + RefreshToken string `json:"refresh_token"` +} + +type tokenPairResponse struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ExpiresIn int `json:"expires_in"` +} + +// SendCode handles POST /v1/auth/code. +func (h *Handler) SendCode(w http.ResponseWriter, r *http.Request) { + var req sendCodeRequest + if !decodeJSON(w, r, &req) { + return + } + retryAfter, apiErr := h.svc.SendCode(r.Context(), req.Email, clientIP(r)) + if apiErr != nil { + writeAPIErr(w, apiErr, retryAfter) + return + } + w.WriteHeader(http.StatusNoContent) +} + +// Register handles POST /v1/auth/register. +func (h *Handler) Register(w http.ResponseWriter, r *http.Request) { + var req registerRequest + if !decodeJSON(w, r, &req) { + return + } + pair, apiErr := h.svc.Register(r.Context(), req.Email, req.Code, req.Password) + if apiErr != nil { + writeAPIErr(w, apiErr, 0) + return + } + writeTokenPair(w, pair) +} + +// Login handles POST /v1/auth/login. +func (h *Handler) Login(w http.ResponseWriter, r *http.Request) { + var req loginRequest + if !decodeJSON(w, r, &req) { + return + } + pair, retryAfter, apiErr := h.svc.Login(r.Context(), req.Email, req.Password, clientIP(r)) + if apiErr != nil { + writeAPIErr(w, apiErr, retryAfter) + return + } + writeTokenPair(w, pair) +} + +// Refresh handles POST /v1/auth/refresh. +func (h *Handler) Refresh(w http.ResponseWriter, r *http.Request) { + var req refreshRequest + if !decodeJSON(w, r, &req) { + return + } + pair, apiErr := h.svc.Refresh(r.Context(), req.RefreshToken) + if apiErr != nil { + writeAPIErr(w, apiErr, 0) + return + } + writeTokenPair(w, pair) +} + +// ---- helpers ---- + +// decodeJSON decodes the request body, writing a 400 on malformed input. +// Returns false if the caller should stop (error already written). +func decodeJSON(w http.ResponseWriter, r *http.Request, dst interface{}) bool { + dec := json.NewDecoder(http.MaxBytesReader(w, r.Body, 1<<16)) + dec.DisallowUnknownFields() + if err := dec.Decode(dst); err != nil { + writeAPIErr(w, ErrInvalidRequest, 0) + return false + } + return true +} + +// writeTokenPair writes a 200 TokenPair body. +func writeTokenPair(w http.ResponseWriter, pair *TokenPair) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(http.StatusOK) + _ = json.NewEncoder(w).Encode(tokenPairResponse{ + AccessToken: pair.AccessToken, + RefreshToken: pair.RefreshToken, + ExpiresIn: pair.ExpiresIn, + }) +} + +// writeAPIErr maps an apierr.Error to the right HTTP status, attaching a +// Retry-After header when provided. +func writeAPIErr(w http.ResponseWriter, e *apierr.Error, retryAfter time.Duration) { + status := statusFor(e) + if retryAfter > 0 { + secs := int(retryAfter.Seconds()) + if secs < 1 { + secs = 1 + } + w.Header().Set("Retry-After", strconv.Itoa(secs)) + } + apierr.WriteJSON(w, status, e) +} + +// statusFor maps auth error codes to HTTP status codes. +func statusFor(e *apierr.Error) int { + switch e.Code { + case ErrInvalidRequest.Code, ErrEmailDisposable.Code, ErrCodeInvalid.Code: + return http.StatusBadRequest + case ErrEmailExists.Code: + return http.StatusConflict + case ErrRateLimited.Code, ErrAccountLocked.Code: + return http.StatusTooManyRequests + case ErrInvalidCredentials.Code, ErrInvalidToken.Code, ErrUnauthorized.Code: + return http.StatusUnauthorized + case ErrAccountBanned.Code: + return http.StatusForbidden + default: + return http.StatusInternalServerError + } +} + +// clientIP extracts the best-effort client IP, honouring X-Forwarded-For and +// X-Real-IP set by the edge proxy. Used only as a rate-limit key — never logged. +func clientIP(r *http.Request) string { + if xff := r.Header.Get("X-Forwarded-For"); xff != "" { + // First entry is the original client. + if i := strings.IndexByte(xff, ','); i >= 0 { + return strings.TrimSpace(xff[:i]) + } + return strings.TrimSpace(xff) + } + if xr := r.Header.Get("X-Real-IP"); xr != "" { + return strings.TrimSpace(xr) + } + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + return r.RemoteAddr + } + return host +} diff --git a/server/internal/auth/handler_test.go b/server/internal/auth/handler_test.go new file mode 100644 index 0000000..275a1ef --- /dev/null +++ b/server/internal/auth/handler_test.go @@ -0,0 +1,148 @@ +package auth + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "testing" + "time" + + "github.com/go-chi/chi/v5" + "github.com/wangjia/pangolin/server/internal/codes" +) + +// newTestHandler builds a Handler + chi router for HTTP-level tests. +func newTestHandler(t *testing.T, cfg ServiceConfig) (*Service, http.Handler) { + t.Helper() + svc, _, _ := newService(t, cfg) + h := NewHandler(svc) + r := chi.NewRouter() + h.RegisterRoutes(r) + return svc, r +} + +func doJSON(t *testing.T, h http.Handler, method, path string, body interface{}) *httptest.ResponseRecorder { + t.Helper() + var buf bytes.Buffer + if body != nil { + _ = json.NewEncoder(&buf).Encode(body) + } + req := httptest.NewRequest(method, path, &buf) + req.RemoteAddr = "203.0.113.5:1234" + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + return rec +} + +func TestHandler_SendCode_204(t *testing.T) { + _, h := newTestHandler(t, ServiceConfig{}) + rec := doJSON(t, h, http.MethodPost, "/auth/code", map[string]string{"email": "u@example.com"}) + if rec.Code != http.StatusNoContent { + t.Fatalf("status = %d, want 204 (body %s)", rec.Code, rec.Body) + } +} + +func TestHandler_SendCode_RateLimited_RetryAfter(t *testing.T) { + _, h := newTestHandler(t, ServiceConfig{EmailPerMinute: 1}) + _ = doJSON(t, h, http.MethodPost, "/auth/code", map[string]string{"email": "u@example.com"}) + rec := doJSON(t, h, http.MethodPost, "/auth/code", map[string]string{"email": "u@example.com"}) + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("status = %d, want 429", rec.Code) + } + ra := rec.Header().Get("Retry-After") + if ra == "" { + t.Fatal("missing Retry-After header") + } + if n, err := strconv.Atoi(ra); err != nil || n <= 0 { + t.Fatalf("bad Retry-After %q", ra) + } +} + +func TestHandler_RegisterLoginRefresh_HTTP(t *testing.T) { + svc, h := newTestHandler(t, ServiceConfig{}) + const email = "flow@example.com" + + // Send code, read it from Redis. + _ = doJSON(t, h, http.MethodPost, "/auth/code", map[string]string{"email": email}) + code := codeInRedis(t, svc, email) + + // Register. + rec := doJSON(t, h, http.MethodPost, "/auth/register", map[string]string{ + "email": email, "code": code, "password": "password123", + }) + if rec.Code != http.StatusOK { + t.Fatalf("register status = %d (body %s)", rec.Code, rec.Body) + } + var reg tokenPairResponse + _ = json.Unmarshal(rec.Body.Bytes(), ®) + if reg.AccessToken == "" || reg.ExpiresIn != 900 { + t.Fatalf("bad register response: %+v", reg) + } + + // Login. + rec = doJSON(t, h, http.MethodPost, "/auth/login", map[string]string{ + "email": email, "password": "password123", + }) + if rec.Code != http.StatusOK { + t.Fatalf("login status = %d (body %s)", rec.Code, rec.Body) + } + var login tokenPairResponse + _ = json.Unmarshal(rec.Body.Bytes(), &login) + + // Refresh. + rec = doJSON(t, h, http.MethodPost, "/auth/refresh", map[string]string{ + "refresh_token": login.RefreshToken, + }) + if rec.Code != http.StatusOK { + t.Fatalf("refresh status = %d (body %s)", rec.Code, rec.Body) + } +} + +func TestMiddleware_RequireAuth(t *testing.T) { + rdb, _ := newMiniRedis(t) + tm := newTokenManager(t, rdb, time.Now) + + // Protected handler that echoes the injected user id. + var gotUID int64 + var gotOK bool + protected := RequireAuth(tm)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotUID, gotOK = UserIDFromContext(r.Context()) + // Confirm the codes module reads the same value via its exported key. + if v, ok := r.Context().Value(codes.CtxKeyUserID).(int64); !ok || v != gotUID { + t.Errorf("codes.CtxKeyUserID mismatch: %v", v) + } + w.WriteHeader(http.StatusOK) + })) + + // No token → 401. + rec := httptest.NewRecorder() + protected.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/me", nil)) + if rec.Code != http.StatusUnauthorized { + t.Fatalf("no-token status = %d, want 401", rec.Code) + } + + // Valid token → 200 with injected uid. + pair, _ := tm.Issue(context.Background(), 77, "uuid-77") + req := httptest.NewRequest(http.MethodGet, "/me", nil) + req.Header.Set("Authorization", "Bearer "+pair.AccessToken) + rec = httptest.NewRecorder() + protected.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("valid-token status = %d, want 200", rec.Code) + } + if !gotOK || gotUID != 77 { + t.Fatalf("injected uid = %d ok=%v, want 77", gotUID, gotOK) + } + + // Refresh token must be rejected by access middleware. + req = httptest.NewRequest(http.MethodGet, "/me", nil) + req.Header.Set("Authorization", "Bearer "+pair.RefreshToken) + rec = httptest.NewRecorder() + protected.ServeHTTP(rec, req) + if rec.Code != http.StatusUnauthorized { + t.Fatalf("refresh-as-access status = %d, want 401", rec.Code) + } +} diff --git a/server/internal/auth/helpers_test.go b/server/internal/auth/helpers_test.go new file mode 100644 index 0000000..ce23a51 --- /dev/null +++ b/server/internal/auth/helpers_test.go @@ -0,0 +1,140 @@ +package auth + +import ( + "context" + "crypto/rand" + "crypto/rsa" + "sync" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/google/uuid" + "github.com/redis/go-redis/v9" +) + +// newMiniRedis spins up an in-memory Redis and returns a connected client. +func newMiniRedis(t *testing.T) (*redis.Client, *miniredis.Miniredis) { + t.Helper() + mr, err := miniredis.Run() + if err != nil { + t.Fatalf("miniredis: %v", err) + } + t.Cleanup(mr.Close) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + return rdb, mr +} + +// newRSAKey generates a 2048-bit RSA key for signing test tokens. +func newRSAKey(t *testing.T) *rsa.PrivateKey { + t.Helper() + k, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("rsa key: %v", err) + } + return k +} + +// newTokenManager builds a TokenManager with the given clock and a single kid. +func newTokenManager(t *testing.T, rdb *redis.Client, now func() time.Time) *TokenManager { + t.Helper() + key := newRSAKey(t) + tm, err := NewTokenManager(rdb, TokenConfig{ + SignKey: key, + SignKID: "k1", + AccessTTL: 15 * time.Minute, + RefreshTTL: 30 * 24 * time.Hour, + Now: now, + }) + if err != nil { + t.Fatalf("NewTokenManager: %v", err) + } + return tm +} + +// -------------------------------------------------------------------------- +// fakeStore — in-memory UserStore for unit tests. +// -------------------------------------------------------------------------- + +type trialRecord struct { + plan string + expiresAt time.Time + source string +} + +type fakeStore struct { + mu sync.Mutex + byEmail map[string]*User + trials map[int64]trialRecord + nextID int64 +} + +func newFakeStore() *fakeStore { + return &fakeStore{byEmail: map[string]*User{}, trials: map[int64]trialRecord{}} +} + +func (f *fakeStore) CreateUserWithTrial(_ context.Context, email, pwHash string, trialDays int) (*User, error) { + f.mu.Lock() + defer f.mu.Unlock() + if _, ok := f.byEmail[email]; ok { + return nil, ErrEmailTaken + } + f.nextID++ + u := &User{ + ID: f.nextID, + UUID: uuid.NewString(), + Email: email, + PwHash: pwHash, + DpUUID: uuid.NewString(), + Status: "active", + } + f.byEmail[email] = u + f.trials[u.ID] = trialRecord{ + plan: "pro", + expiresAt: time.Now().UTC().AddDate(0, 0, trialDays), + source: "trial", + } + return u, nil +} + +func (f *fakeStore) GetUserByEmail(_ context.Context, email string) (*User, error) { + f.mu.Lock() + defer f.mu.Unlock() + u, ok := f.byEmail[email] + if !ok { + return nil, ErrNotFound + } + cp := *u + return &cp, nil +} + +// setStatus mutates a stored user's status (e.g. to "banned") for tests. +func (f *fakeStore) setStatus(email, status string) { + f.mu.Lock() + defer f.mu.Unlock() + if u, ok := f.byEmail[email]; ok { + u.Status = status + } +} + +// captureMailer records the last code it was asked to send. +type captureMailer struct { + mu sync.Mutex + last map[string]string +} + +func newCaptureMailer() *captureMailer { return &captureMailer{last: map[string]string{}} } + +func (m *captureMailer) SendCode(_ context.Context, to, code string) error { + m.mu.Lock() + defer m.mu.Unlock() + m.last[to] = code + return nil +} + +func (m *captureMailer) codeFor(to string) string { + m.mu.Lock() + defer m.mu.Unlock() + return m.last[to] +} diff --git a/server/internal/auth/integration_test.go b/server/internal/auth/integration_test.go new file mode 100644 index 0000000..eb11010 --- /dev/null +++ b/server/internal/auth/integration_test.go @@ -0,0 +1,217 @@ +//go:build integration + +package auth + +import ( + "context" + "database/sql" + "net/http" + "net/http/httptest" + "testing" + "time" + + _ "github.com/go-sql-driver/mysql" + "github.com/redis/go-redis/v9" + "github.com/testcontainers/testcontainers-go" + tcmysql "github.com/testcontainers/testcontainers-go/modules/mysql" + tcredis "github.com/testcontainers/testcontainers-go/modules/redis" + + "github.com/wangjia/pangolin/server/internal/codes" +) + +func setupMySQL(t *testing.T) *sql.DB { + t.Helper() + ctx := context.Background() + ctr, err := tcmysql.Run(ctx, "mysql:8.0", + tcmysql.WithDatabase("pangolin_test"), + tcmysql.WithUsername("root"), + tcmysql.WithPassword("test"), + ) + testcontainers.CleanupContainer(t, ctr) + if err != nil { + t.Fatalf("mysql container: %v", err) + } + dsn, err := ctr.ConnectionString(ctx, "parseTime=true", "loc=UTC", "time_zone='+00:00'") + if err != nil { + t.Fatalf("mysql dsn: %v", err) + } + db, err := sql.Open("mysql", dsn) + if err != nil { + t.Fatalf("open mysql: %v", err) + } + t.Cleanup(func() { db.Close() }) + if err := applyAuthSchema(db); err != nil { + t.Fatalf("schema: %v", err) + } + return db +} + +func setupRedis(t *testing.T) *redis.Client { + t.Helper() + ctx := context.Background() + ctr, err := tcredis.Run(ctx, "redis:7-alpine") + testcontainers.CleanupContainer(t, ctr) + if err != nil { + t.Fatalf("redis container: %v", err) + } + addr, err := ctr.ConnectionString(ctx) + if err != nil { + t.Fatalf("redis addr: %v", err) + } + for _, p := range []string{"redis://", "rediss://"} { + if len(addr) > len(p) && addr[:len(p)] == p { + addr = addr[len(p):] + } + } + rdb := redis.NewClient(&redis.Options{Addr: addr}) + t.Cleanup(func() { rdb.Close() }) + return rdb +} + +func applyAuthSchema(db *sql.DB) error { + stmts := []string{ + `CREATE TABLE IF NOT EXISTS plans ( + id BIGINT UNSIGNED AUTO_INCREMENT PRIMARY KEY, + code ENUM('free','pro','team') NOT NULL UNIQUE, + max_devices INT NOT NULL DEFAULT 1, + daily_minutes INT NULL, + ad_gate BOOLEAN NOT NULL DEFAULT FALSE + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + `CREATE TABLE IF NOT EXISTS users ( + id BIGINT UNSIGNED AUTO_INCREMENT PRIMARY KEY, + uuid CHAR(36) NOT NULL UNIQUE, + email VARCHAR(255) NOT NULL UNIQUE, + pw_hash VARCHAR(255) NOT NULL, + dp_uuid CHAR(36) NOT NULL, + status ENUM('active','banned') NOT NULL DEFAULT 'active', + created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + `CREATE TABLE IF NOT EXISTS subscriptions ( + id BIGINT UNSIGNED AUTO_INCREMENT PRIMARY KEY, + user_id BIGINT UNSIGNED NOT NULL, + plan_id BIGINT UNSIGNED NOT NULL, + expires_at DATETIME(6) NOT NULL, + source ENUM('trial','code') NOT NULL, + created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6), + FOREIGN KEY (plan_id) REFERENCES plans(id), + INDEX idx_user_exp (user_id, expires_at) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + `INSERT IGNORE INTO plans (code, max_devices, daily_minutes, ad_gate) + VALUES ('free',1,10,TRUE),('pro',5,NULL,FALSE),('team',10,NULL,FALSE)`, + } + for _, s := range stmts { + if _, err := db.Exec(s); err != nil { + return err + } + } + return nil +} + +func newIntegrationService(t *testing.T, db *sql.DB, rdb *redis.Client) *Service { + t.Helper() + store := NewSQLStore(db) + rl := NewRateLimiter(rdb, nil) + key := newRSAKey(t) + tm, err := NewTokenManager(rdb, TokenConfig{SignKey: key, SignKID: "k1"}) + if err != nil { + t.Fatalf("token manager: %v", err) + } + return NewService(store, rdb, rl, tm, NewLogMailer(nil), ServiceConfig{}, nil) +} + +// TestIntegration_FullChain exercises register → login → refresh → protected +// route against real MySQL 8 and Redis containers. +func TestIntegration_FullChain(t *testing.T) { + db := setupMySQL(t) + rdb := setupRedis(t) + svc := newIntegrationService(t, db, rdb) + ctx := context.Background() + const email = "integration@example.com" + const pw = "password-integration" + + // 1. Send code (read it back from Redis to simulate the user). + if _, e := svc.SendCode(ctx, email, "198.51.100.7"); e != nil { + t.Fatalf("SendCode: %v", e) + } + code, err := rdb.Get(ctx, codeKey(email)).Result() + if err != nil { + t.Fatalf("code not stored: %v", err) + } + + // 2. Register → trial subscription must exist for 7 days. + pair, apiErr := svc.Register(ctx, email, code, pw) + if apiErr != nil { + t.Fatalf("Register: %v", apiErr) + } + + var plan, source string + var expires time.Time + err = db.QueryRowContext(ctx, + `SELECT p.code, s.source, s.expires_at + FROM subscriptions s JOIN plans p ON p.id = s.plan_id + JOIN users u ON u.id = s.user_id + WHERE u.email = ?`, email).Scan(&plan, &source, &expires) + if err != nil { + t.Fatalf("query trial: %v", err) + } + if plan != "pro" || source != "trial" { + t.Fatalf("trial = %s/%s, want pro/trial", plan, source) + } + days := time.Until(expires).Hours() / 24 + if days < 6.5 || days > 7.1 { + t.Fatalf("trial length = %.2f days, want ~7", days) + } + + // 3. Duplicate email → 409. + if _, e := svc.SendCode(ctx, email, ""); e != nil && e.Code != ErrRateLimited.Code { + t.Fatalf("second SendCode: %v", e) + } + // Force a fresh code regardless of rate limit. + _ = rdb.Set(ctx, codeKey(email), code, 10*time.Minute).Err() + if _, e := svc.Register(ctx, email, code, pw); e == nil || e.Code != ErrEmailExists.Code { + t.Fatalf("want email_exists, got %v", e) + } + + // 4. Login. + loginPair, _, apiErr := svc.Login(ctx, email, pw, "198.51.100.7") + if apiErr != nil { + t.Fatalf("Login: %v", apiErr) + } + + // 5. Refresh rotates. + rotated, apiErr := svc.Refresh(ctx, loginPair.RefreshToken) + if apiErr != nil { + t.Fatalf("Refresh: %v", apiErr) + } + if _, e := svc.Refresh(ctx, loginPair.RefreshToken); e == nil { + t.Fatal("old refresh token must be rejected after rotation") + } + + // 6. Access a protected route with the rotated access token. + protected := RequireAuth(svc.tokens)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + uid, ok := UserIDFromContext(r.Context()) + if !ok || uid == 0 { + w.WriteHeader(http.StatusUnauthorized) + return + } + // Confirm interop with the codes module's context key. + if _, ok := r.Context().Value(codes.CtxKeyUserID).(int64); !ok { + w.WriteHeader(http.StatusInternalServerError) + return + } + w.WriteHeader(http.StatusOK) + })) + + req := httptest.NewRequest(http.MethodGet, "/me", nil) + req.Header.Set("Authorization", "Bearer "+rotated.AccessToken) + rec := httptest.NewRecorder() + protected.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("protected route status = %d, want 200", rec.Code) + } + + // Pair returned at registration is also a valid access token. + if _, e := svc.tokens.ParseAccess(pair.AccessToken); e != nil { + t.Fatalf("register access token invalid: %v", e) + } +} diff --git a/server/internal/auth/keyloader.go b/server/internal/auth/keyloader.go new file mode 100644 index 0000000..9a65912 --- /dev/null +++ b/server/internal/auth/keyloader.go @@ -0,0 +1,49 @@ +package auth + +import ( + "crypto/rsa" + "fmt" + "os" +) + +// LoadTokenConfig assembles a TokenConfig from PEM files on disk: one private +// signing key (privateKeyPath, identified by signKID) and a set of public keys +// (publicKeyPaths: kid -> path) accepted for verification, supporting key +// rotation. The signing key's public part is added automatically if its kid is +// absent from publicKeyPaths. +func LoadTokenConfig(privateKeyPath, signKID string, publicKeyPaths map[string]string) (TokenConfig, error) { + if privateKeyPath == "" { + return TokenConfig{}, fmt.Errorf("auth: JWT private key path is empty") + } + if signKID == "" { + return TokenConfig{}, fmt.Errorf("auth: JWT key id (kid) is empty") + } + + privPEM, err := os.ReadFile(privateKeyPath) + if err != nil { + return TokenConfig{}, fmt.Errorf("auth: read private key: %w", err) + } + priv, err := LoadPrivateKeyPEM(privPEM) + if err != nil { + return TokenConfig{}, err + } + + verify := map[string]*rsa.PublicKey{} + for kid, path := range publicKeyPaths { + pubPEM, err := os.ReadFile(path) + if err != nil { + return TokenConfig{}, fmt.Errorf("auth: read public key %q: %w", kid, err) + } + pub, err := LoadPublicKeyPEM(pubPEM) + if err != nil { + return TokenConfig{}, fmt.Errorf("auth: parse public key %q: %w", kid, err) + } + verify[kid] = pub + } + + return TokenConfig{ + SignKey: priv, + SignKID: signKID, + VerifyKeys: verify, + }, nil +} diff --git a/server/internal/auth/mailer.go b/server/internal/auth/mailer.go new file mode 100644 index 0000000..b9aa9ac --- /dev/null +++ b/server/internal/auth/mailer.go @@ -0,0 +1,91 @@ +package auth + +import ( + "context" + "fmt" + "log" + "net/smtp" + "strings" +) + +// Mailer delivers a verification code to an email address. Implementations must +// be safe for concurrent use; Send is typically invoked from a goroutine. +type Mailer interface { + SendCode(ctx context.Context, to, code string) error +} + +// -------------------------------------------------------------------------- +// SMTP implementation +// -------------------------------------------------------------------------- + +// SMTPConfig configures the SMTP mailer. +type SMTPConfig struct { + Host string // SMTP host (no port) + Port int // SMTP port, e.g. 587 + Username string + Password string + From string // From header, e.g. "Pangolin " +} + +// SMTPMailer sends verification codes over SMTP with STARTTLS (PlainAuth). +type SMTPMailer struct { + cfg SMTPConfig +} + +// NewSMTPMailer builds an SMTPMailer. +func NewSMTPMailer(cfg SMTPConfig) *SMTPMailer { return &SMTPMailer{cfg: cfg} } + +// SendCode delivers the verification code. The message body intentionally +// contains no product-identifying or destination wording beyond a neutral +// account-verification notice. +func (m *SMTPMailer) SendCode(_ context.Context, to, code string) error { + subject := "Your verification code / 验证码" + body := fmt.Sprintf( + "Your verification code is %s. It expires in 10 minutes.\r\n"+ + "您的验证码为 %s,10 分钟内有效。请勿向他人泄露。\r\n", + code, code) + + msg := strings.Join([]string{ + "From: " + m.cfg.From, + "To: " + to, + "Subject: " + subject, + "MIME-Version: 1.0", + "Content-Type: text/plain; charset=UTF-8", + "", + body, + }, "\r\n") + + addr := fmt.Sprintf("%s:%d", m.cfg.Host, m.cfg.Port) + auth := smtp.PlainAuth("", m.cfg.Username, m.cfg.Password, m.cfg.Host) + if err := smtp.SendMail(addr, auth, m.cfg.From, []string{to}, []byte(msg)); err != nil { + return fmt.Errorf("auth: smtp send: %w", err) + } + return nil +} + +// -------------------------------------------------------------------------- +// Development log implementation +// -------------------------------------------------------------------------- + +// LogMailer is a development-only Mailer that prints the code to the process +// log instead of sending an email. It MUST NOT be used in production — it is +// the single deliberate exception to the no-secret-in-logs rule, scoped to +// local development where no real mailbox exists. +type LogMailer struct { + logger *log.Logger +} + +// NewLogMailer builds a LogMailer. If logger is nil, the standard logger is used. +func NewLogMailer(logger *log.Logger) *LogMailer { + return &LogMailer{logger: logger} +} + +// SendCode logs the code for local development. +func (m *LogMailer) SendCode(_ context.Context, to, code string) error { + if m.logger != nil { + m.logger.Printf("[dev-mailer] verification code for %s: %s", to, code) + } else { + log.Printf("[dev-mailer] verification code for %s: %s", to, code) + } + return nil +} diff --git a/server/internal/auth/middleware.go b/server/internal/auth/middleware.go new file mode 100644 index 0000000..e80b969 --- /dev/null +++ b/server/internal/auth/middleware.go @@ -0,0 +1,86 @@ +package auth + +import ( + "context" + "net/http" + "strings" + + "github.com/wangjia/pangolin/server/internal/codes" +) + +// ctxKey is this package's private context-key type for values other than the +// numeric user id (which is shared with the codes module — see below). +type ctxKey string + +const ( + // ctxKeyUserUUID stores the authenticated user's UUID (the JWT subject). + ctxKeyUserUUID ctxKey = "user_uuid" + // ctxKeyClaims stores the full *Claims for handlers that need jti/exp. + ctxKeyClaims ctxKey = "claims" +) + +// The numeric user id is injected under codes.CtxKeyUserID so that the codes +// module (and any other module sharing that exported key) reads the same value +// without an import cycle — codes deliberately exports the key for this purpose. + +// RequireAuth returns middleware that enforces a valid Bearer access token on +// every wrapped route. On success it injects the numeric user id, the user +// UUID, and the parsed claims into the request context. This is the auth base +// for all protected /v1 routes (everything except the auth group). +func RequireAuth(tm *TokenManager) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + token, ok := bearerToken(r) + if !ok { + writeAPIErr(w, ErrUnauthorized, 0) + return + } + claims, err := tm.ParseAccess(token) + if err != nil { + writeAPIErr(w, ErrUnauthorized, 0) + return + } + + ctx := r.Context() + ctx = context.WithValue(ctx, codes.CtxKeyUserID, claims.UID) + ctx = context.WithValue(ctx, ctxKeyUserUUID, claims.Subject) + ctx = context.WithValue(ctx, ctxKeyClaims, claims) + next.ServeHTTP(w, r.WithContext(ctx)) + }) + } +} + +// bearerToken extracts the token from an "Authorization: Bearer " header. +func bearerToken(r *http.Request) (string, bool) { + h := r.Header.Get("Authorization") + if h == "" { + return "", false + } + const prefix = "Bearer " + if len(h) <= len(prefix) || !strings.EqualFold(h[:len(prefix)], prefix) { + return "", false + } + token := strings.TrimSpace(h[len(prefix):]) + if token == "" { + return "", false + } + return token, true +} + +// UserIDFromContext returns the numeric user id injected by RequireAuth. +func UserIDFromContext(ctx context.Context) (int64, bool) { + v, ok := ctx.Value(codes.CtxKeyUserID).(int64) + return v, ok && v != 0 +} + +// UserUUIDFromContext returns the user UUID injected by RequireAuth. +func UserUUIDFromContext(ctx context.Context) (string, bool) { + v, ok := ctx.Value(ctxKeyUserUUID).(string) + return v, ok && v != "" +} + +// ClaimsFromContext returns the parsed access-token claims injected by RequireAuth. +func ClaimsFromContext(ctx context.Context) (*Claims, bool) { + v, ok := ctx.Value(ctxKeyClaims).(*Claims) + return v, ok +} diff --git a/server/internal/auth/password.go b/server/internal/auth/password.go new file mode 100644 index 0000000..05198d3 --- /dev/null +++ b/server/internal/auth/password.go @@ -0,0 +1,111 @@ +package auth + +import ( + "crypto/rand" + "crypto/subtle" + "encoding/base64" + "errors" + "fmt" + "strings" + + "golang.org/x/crypto/argon2" +) + +// argon2id parameters. Tuned for a control-plane login path: ~64 MiB memory, +// single pass, 4 lanes. Kept as constants so the encoded hash is self-describing +// and parameters can evolve without breaking existing hashes (params are parsed +// back out of the stored string at verify time). +const ( + argonMemory uint32 = 64 * 1024 // KiB → 64 MiB + argonTime uint32 = 1 + argonThreads uint8 = 4 + argonKeyLen uint32 = 32 + argonSaltLen = 16 +) + +// errInvalidHash is returned when a stored PHC string cannot be parsed. +var errInvalidHash = errors.New("auth: invalid argon2id hash format") + +// HashPassword derives an argon2id hash and returns it in the standard PHC +// string format: $argon2id$v=19$m=...,t=...,p=...$$. +func HashPassword(password string) (string, error) { + salt := make([]byte, argonSaltLen) + if _, err := rand.Read(salt); err != nil { + return "", fmt.Errorf("auth: read salt: %w", err) + } + return encodeArgon(password, salt, argonTime, argonMemory, argonThreads, argonKeyLen), nil +} + +// encodeArgon computes the hash and formats the PHC string. +func encodeArgon(password string, salt []byte, t, m uint32, p uint8, keyLen uint32) string { + key := argon2.IDKey([]byte(password), salt, t, m, p, keyLen) + return fmt.Sprintf( + "$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s", + argon2.Version, m, t, p, + base64.RawStdEncoding.EncodeToString(salt), + base64.RawStdEncoding.EncodeToString(key), + ) +} + +// VerifyPassword reports whether password matches the encoded argon2id hash. +// The comparison is constant-time. A malformed encoded hash returns false with +// an error so callers can distinguish "wrong password" from "corrupt record". +func VerifyPassword(encoded, password string) (bool, error) { + t, m, p, salt, key, err := decodeArgon(encoded) + if err != nil { + return false, err + } + other := argon2.IDKey([]byte(password), salt, t, m, p, uint32(len(key))) + if subtle.ConstantTimeCompare(key, other) == 1 { + return true, nil + } + return false, nil +} + +// decodeArgon parses a PHC argon2id string back into its parameters. +func decodeArgon(encoded string) (t, m uint32, p uint8, salt, key []byte, err error) { + parts := strings.Split(encoded, "$") + // ["", "argon2id", "v=19", "m=..,t=..,p=..", salt, key] + if len(parts) != 6 || parts[1] != "argon2id" { + return 0, 0, 0, nil, nil, errInvalidHash + } + + var version int + if _, err = fmt.Sscanf(parts[2], "v=%d", &version); err != nil { + return 0, 0, 0, nil, nil, errInvalidHash + } + if version != argon2.Version { + return 0, 0, 0, nil, nil, errInvalidHash + } + if _, err = fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &m, &t, &p); err != nil { + return 0, 0, 0, nil, nil, errInvalidHash + } + if salt, err = base64.RawStdEncoding.DecodeString(parts[4]); err != nil { + return 0, 0, 0, nil, nil, errInvalidHash + } + if key, err = base64.RawStdEncoding.DecodeString(parts[5]); err != nil { + return 0, 0, 0, nil, nil, errInvalidHash + } + return t, m, p, salt, key, nil +} + +// dummyHash is a pre-computed argon2id hash used to spend roughly the same CPU +// time when a login targets a non-existent account, keeping the login response +// time constant regardless of whether the email exists. Computed lazily once. +var dummyHash = mustDummyHash() + +func mustDummyHash() string { + h, err := HashPassword("pangolin-constant-time-placeholder") + if err != nil { + // Fall back to a static PHC string; verification will simply fail. + return "$argon2id$v=19$m=65536,t=1,p=4$AAAAAAAAAAAAAAAAAAAAAA$AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA" + } + return h +} + +// ConstantTimeReject performs a throwaway argon2id verification against a dummy +// hash. Login handlers call this when the account does not exist so the timing +// profile matches the "account exists, wrong password" branch. +func ConstantTimeReject(password string) { + _, _ = VerifyPassword(dummyHash, password) +} diff --git a/server/internal/auth/password_test.go b/server/internal/auth/password_test.go new file mode 100644 index 0000000..a87fd60 --- /dev/null +++ b/server/internal/auth/password_test.go @@ -0,0 +1,58 @@ +package auth + +import ( + "strings" + "testing" +) + +func TestHashPassword_RoundTrip(t *testing.T) { + const pw = "correct horse battery staple" + hash, err := HashPassword(pw) + if err != nil { + t.Fatalf("HashPassword: %v", err) + } + if !strings.HasPrefix(hash, "$argon2id$v=19$") { + t.Fatalf("unexpected PHC prefix: %s", hash) + } + + ok, err := VerifyPassword(hash, pw) + if err != nil { + t.Fatalf("VerifyPassword: %v", err) + } + if !ok { + t.Fatal("expected password to verify") + } +} + +func TestHashPassword_DistinctSalts(t *testing.T) { + h1, _ := HashPassword("same") + h2, _ := HashPassword("same") + if h1 == h2 { + t.Fatal("expected distinct hashes for equal passwords (random salt)") + } +} + +func TestVerifyPassword_Wrong(t *testing.T) { + hash, _ := HashPassword("right") + ok, err := VerifyPassword(hash, "wrong") + if err != nil { + t.Fatalf("VerifyPassword: %v", err) + } + if ok { + t.Fatal("expected wrong password to fail") + } +} + +func TestVerifyPassword_Malformed(t *testing.T) { + cases := []string{ + "", + "not-a-hash", + "$argon2id$v=19$bad", + "$bcrypt$v=19$m=1,t=1,p=1$aaaa$bbbb", + } + for _, c := range cases { + if _, err := VerifyPassword(c, "x"); err == nil { + t.Fatalf("expected error for malformed hash %q", c) + } + } +} diff --git a/server/internal/auth/ratelimit.go b/server/internal/auth/ratelimit.go new file mode 100644 index 0000000..d35184f --- /dev/null +++ b/server/internal/auth/ratelimit.go @@ -0,0 +1,161 @@ +package auth + +import ( + "context" + "crypto/rand" + "encoding/hex" + "fmt" + "time" + + "github.com/redis/go-redis/v9" +) + +// RateLimiter implements per-scope sliding-window rate limiting and failure +// counters on top of Redis. It stores nothing but short-lived counters — no +// connection or behaviour logs — per the no-log baseline (doc/06 §4). +// +// Sliding-window keys are ZSETs named `rl:{scope}:{key}` whose members are +// timestamped attempts; failure counters are plain integer keys. +type RateLimiter struct { + rdb *redis.Client + now func() time.Time +} + +// NewRateLimiter builds a RateLimiter. now may be nil (defaults to time.Now). +func NewRateLimiter(rdb *redis.Client, now func() time.Time) *RateLimiter { + if now == nil { + now = time.Now + } + return &RateLimiter{rdb: rdb, now: now} +} + +// slidingWindow is a single atomic Lua script: +// - drops members older than the window, +// - if the live count already reached the limit, returns {0, oldestScore}, +// - otherwise records the attempt and returns {1, 0}. +// +// All times are unix-milliseconds. +var slidingWindow = redis.NewScript(` +local key = KEYS[1] +local now = tonumber(ARGV[1]) +local window = tonumber(ARGV[2]) +local limit = tonumber(ARGV[3]) +local member = ARGV[4] +redis.call('ZREMRANGEBYSCORE', key, 0, now - window) +local count = redis.call('ZCARD', key) +if count >= limit then + local oldest = redis.call('ZRANGE', key, 0, 0, 'WITHSCORES') + return {0, oldest[2]} +end +redis.call('ZADD', key, now, member) +redis.call('PEXPIRE', key, window) +return {1, 0} +`) + +// rlKey builds the canonical `rl:{scope}:{key}` Redis key. +func rlKey(scope, key string) string { + return fmt.Sprintf("rl:%s:%s", scope, key) +} + +// Allow records an attempt in the sliding window for (scope, key). It returns +// allowed=false together with the Retry-After duration when the limit within +// window has been reached. limit is the maximum number of attempts permitted +// inside window. +func (rl *RateLimiter) Allow(ctx context.Context, scope, key string, limit int, window time.Duration) (allowed bool, retryAfter time.Duration, err error) { + nowMs := rl.now().UnixMilli() + winMs := window.Milliseconds() + + member, err := uniqueMember(nowMs) + if err != nil { + return false, 0, err + } + + res, err := slidingWindow.Run(ctx, rl.rdb, []string{rlKey(scope, key)}, + nowMs, winMs, limit, member).Result() + if err != nil { + return false, 0, fmt.Errorf("auth: ratelimit run: %w", err) + } + + vals, ok := res.([]interface{}) + if !ok || len(vals) != 2 { + return false, 0, fmt.Errorf("auth: ratelimit unexpected result %v", res) + } + ok1, _ := vals[0].(int64) + if ok1 == 1 { + return true, 0, nil + } + + // Denied: compute how long until the oldest attempt leaves the window. + oldest := toInt64(vals[1]) + ra := time.Duration(oldest+winMs-nowMs) * time.Millisecond + if ra < time.Second { + ra = time.Second + } + return false, ra, nil +} + +// failKey builds the failure-counter key, e.g. `rl:login:{email}`. +func failKey(scope, key string) string { + return rlKey(scope, key) +} + +// RecordFailure increments the failure counter for (scope, key), (re)setting +// its TTL to window on every increment so the lock slides forward while abuse +// continues. It returns the new count. +func (rl *RateLimiter) RecordFailure(ctx context.Context, scope, key string, window time.Duration) (int64, error) { + k := failKey(scope, key) + pipe := rl.rdb.Pipeline() + incr := pipe.Incr(ctx, k) + pipe.Expire(ctx, k, window) + if _, err := pipe.Exec(ctx); err != nil { + return 0, fmt.Errorf("auth: record failure: %w", err) + } + return incr.Val(), nil +} + +// FailureCount returns the current failure count and the remaining TTL (lock +// window) for (scope, key). count is 0 when no counter exists. +func (rl *RateLimiter) FailureCount(ctx context.Context, scope, key string) (count int64, ttl time.Duration, err error) { + k := failKey(scope, key) + pipe := rl.rdb.Pipeline() + get := pipe.Get(ctx, k) + pttl := pipe.PTTL(ctx, k) + if _, err := pipe.Exec(ctx); err != nil && err != redis.Nil { + return 0, 0, fmt.Errorf("auth: failure count: %w", err) + } + n, _ := get.Int64() + d := pttl.Val() + if d < 0 { + d = 0 + } + return n, d, nil +} + +// ClearFailures removes the failure counter after a successful authentication. +func (rl *RateLimiter) ClearFailures(ctx context.Context, scope, key string) error { + return rl.rdb.Del(ctx, failKey(scope, key)).Err() +} + +// uniqueMember returns a ZSET member that is unique even for attempts that share +// the same millisecond timestamp (timestamp prefix keeps ordering stable). +func uniqueMember(nowMs int64) (string, error) { + var b [8]byte + if _, err := rand.Read(b[:]); err != nil { + return "", fmt.Errorf("auth: ratelimit member: %w", err) + } + return fmt.Sprintf("%d-%s", nowMs, hex.EncodeToString(b[:])), nil +} + +// toInt64 coerces a Redis Lua return value (which may be int64 or string) to int64. +func toInt64(v interface{}) int64 { + switch t := v.(type) { + case int64: + return t + case string: + var n int64 + _, _ = fmt.Sscanf(t, "%d", &n) + return n + default: + return 0 + } +} diff --git a/server/internal/auth/ratelimit_test.go b/server/internal/auth/ratelimit_test.go new file mode 100644 index 0000000..4a89d1c --- /dev/null +++ b/server/internal/auth/ratelimit_test.go @@ -0,0 +1,97 @@ +package auth + +import ( + "context" + "testing" + "time" +) + +func TestRateLimiter_SlidingWindow(t *testing.T) { + rdb, _ := newMiniRedis(t) + base := time.Unix(1_700_000_000, 0) + clock := base + rl := NewRateLimiter(rdb, func() time.Time { return clock }) + ctx := context.Background() + + // limit 1 per minute. + ok, _, err := rl.Allow(ctx, "code:email", "a@b.com", 1, time.Minute) + if err != nil || !ok { + t.Fatalf("first attempt should pass: ok=%v err=%v", ok, err) + } + + // Second attempt within the window is denied with a Retry-After. + ok, ra, err := rl.Allow(ctx, "code:email", "a@b.com", 1, time.Minute) + if err != nil { + t.Fatalf("err: %v", err) + } + if ok { + t.Fatal("second attempt within window should be denied") + } + if ra <= 0 || ra > time.Minute { + t.Fatalf("retry-after = %v, want (0, 1m]", ra) + } + + // After the window passes, attempts are allowed again. + clock = base.Add(61 * time.Second) + ok, _, err = rl.Allow(ctx, "code:email", "a@b.com", 1, time.Minute) + if err != nil || !ok { + t.Fatalf("attempt after window should pass: ok=%v err=%v", ok, err) + } +} + +func TestRateLimiter_DistinctKeysIndependent(t *testing.T) { + rdb, _ := newMiniRedis(t) + rl := NewRateLimiter(rdb, nil) + ctx := context.Background() + + if ok, _, _ := rl.Allow(ctx, "code:email", "a@b.com", 1, time.Minute); !ok { + t.Fatal("a@b.com first should pass") + } + if ok, _, _ := rl.Allow(ctx, "code:email", "c@d.com", 1, time.Minute); !ok { + t.Fatal("c@d.com first should pass (independent key)") + } +} + +func TestRateLimiter_FailureCounter(t *testing.T) { + rdb, mr := newMiniRedis(t) + rl := NewRateLimiter(rdb, nil) + ctx := context.Background() + + for i := 1; i <= 3; i++ { + n, err := rl.RecordFailure(ctx, "login", "a@b.com", 15*time.Minute) + if err != nil { + t.Fatalf("RecordFailure: %v", err) + } + if n != int64(i) { + t.Fatalf("count = %d, want %d", n, i) + } + } + + count, ttl, err := rl.FailureCount(ctx, "login", "a@b.com") + if err != nil { + t.Fatalf("FailureCount: %v", err) + } + if count != 3 { + t.Fatalf("count = %d, want 3", count) + } + if ttl <= 0 { + t.Fatalf("ttl = %v, want > 0", ttl) + } + + // Clear resets the counter. + if err := rl.ClearFailures(ctx, "login", "a@b.com"); err != nil { + t.Fatalf("ClearFailures: %v", err) + } + count, _, _ = rl.FailureCount(ctx, "login", "a@b.com") + if count != 0 { + t.Fatalf("count after clear = %d, want 0", count) + } + + // Counter expires after its window. + _, _ = rl.RecordFailure(ctx, "login", "x@y.com", 15*time.Minute) + mr.FastForward(16 * time.Minute) + count, _, _ = rl.FailureCount(ctx, "login", "x@y.com") + if count != 0 { + t.Fatalf("count after TTL = %d, want 0", count) + } +} diff --git a/server/internal/auth/service.go b/server/internal/auth/service.go new file mode 100644 index 0000000..139e380 --- /dev/null +++ b/server/internal/auth/service.go @@ -0,0 +1,292 @@ +package auth + +import ( + "context" + "crypto/rand" + "crypto/subtle" + "errors" + "fmt" + "math/big" + "time" + + "github.com/redis/go-redis/v9" + "github.com/wangjia/pangolin/server/internal/apierr" +) + +// Redis key helpers for verification codes (doc/03 §4: auth:code:{email}). +func codeKey(email string) string { return "auth:code:" + email } +func codeAttemptsKey(email string) string { return "auth:code:attempts:" + email } + +// Rate-limit scopes. +const ( + scopeCodeEmail = "code:email" // per-email send limit + scopeCodeIP = "code:ip" // per-IP send limit + scopeLogin = "login" // per-email login-failure counter +) + +// ServiceConfig tunes the auth service. Zero values fall back to the documented +// defaults (doc/02 §4.1, doc/06 §3). +type ServiceConfig struct { + CodeTTL time.Duration // verification code lifetime (default 10m) + CodeMaxAttempts int // max verify attempts before code is burned (default 5) + TrialDays int // auto PRO trial length (default 7) + + EmailPerMinute int // code sends per email per minute (default 1) + IPPerHour int // code sends per IP per hour (default 10) + LoginFailMax int // failed logins before lock (default 5) + LoginLockWindow time.Duration // lock / failure-window length (default 15m) +} + +func (c *ServiceConfig) withDefaults() { + if c.CodeTTL <= 0 { + c.CodeTTL = 10 * time.Minute + } + if c.CodeMaxAttempts <= 0 { + c.CodeMaxAttempts = 5 + } + if c.TrialDays <= 0 { + c.TrialDays = 7 + } + if c.EmailPerMinute <= 0 { + c.EmailPerMinute = 1 + } + if c.IPPerHour <= 0 { + c.IPPerHour = 10 + } + if c.LoginFailMax <= 0 { + c.LoginFailMax = 5 + } + if c.LoginLockWindow <= 0 { + c.LoginLockWindow = 15 * time.Minute + } +} + +// Service is the auth business layer: code issuance, registration, login, and +// token refresh. It is safe for concurrent use. +type Service struct { + store UserStore + rdb *redis.Client + rl *RateLimiter + tokens *TokenManager + mailer Mailer + cfg ServiceConfig + now func() time.Time +} + +// NewService wires the auth service. now may be nil (defaults to time.Now). +func NewService(store UserStore, rdb *redis.Client, rl *RateLimiter, tokens *TokenManager, mailer Mailer, cfg ServiceConfig, now func() time.Time) *Service { + cfg.withDefaults() + if now == nil { + now = time.Now + } + return &Service{ + store: store, + rdb: rdb, + rl: rl, + tokens: tokens, + mailer: mailer, + cfg: cfg, + now: now, + } +} + +// SendCode applies rate-limiting and disposable-domain checks, generates a +// 6-digit code, stores it in Redis (TTL CodeTTL), and dispatches it +// asynchronously. retryAfter is non-zero only when a rate limit was hit. +func (s *Service) SendCode(ctx context.Context, rawEmail, ip string) (retryAfter time.Duration, apiErr *apierr.Error) { + email := NormalizeEmail(rawEmail) + if !ValidEmail(email) { + return 0, ErrInvalidRequest + } + if IsDisposable(email) { + return 0, ErrEmailDisposable + } + + // Per-email limit: 1/min by default. + ok, ra, err := s.rl.Allow(ctx, scopeCodeEmail, email, s.cfg.EmailPerMinute, time.Minute) + if err != nil { + return 0, ErrInternal + } + if !ok { + return ra, ErrRateLimited + } + // Per-IP limit: e.g. 10/h. Skipped when IP is unknown. + if ip != "" { + ok, ra, err = s.rl.Allow(ctx, scopeCodeIP, ip, s.cfg.IPPerHour, time.Hour) + if err != nil { + return 0, ErrInternal + } + if !ok { + return ra, ErrRateLimited + } + } + + code, err := genNumericCode(6) + if err != nil { + return 0, ErrInternal + } + + pipe := s.rdb.Pipeline() + pipe.Set(ctx, codeKey(email), code, s.cfg.CodeTTL) + pipe.Del(ctx, codeAttemptsKey(email)) // reset attempt counter for the new code + if _, err := pipe.Exec(ctx); err != nil { + return 0, ErrInternal + } + + // Dispatch asynchronously; the request must not block on SMTP. A detached + // context is used so request cancellation doesn't abort delivery. + go func(to, c string) { + sendCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + _ = s.mailer.SendCode(sendCtx, to, c) + }(email, code) + + return 0, nil +} + +// Register verifies the code (one-time), creates the account plus a 7-day PRO +// trial in a single transaction, and returns a fresh token pair. +func (s *Service) Register(ctx context.Context, rawEmail, code, password string) (*TokenPair, *apierr.Error) { + email := NormalizeEmail(rawEmail) + if !ValidEmail(email) || len(password) < 8 || len(code) != 6 { + return nil, ErrInvalidRequest + } + + // Verify the code with brute-force protection. + if apiErr := s.verifyCode(ctx, email, code); apiErr != nil { + return nil, apiErr + } + + pwHash, err := HashPassword(password) + if err != nil { + return nil, ErrInternal + } + + user, err := s.store.CreateUserWithTrial(ctx, email, pwHash, s.cfg.TrialDays) + if err != nil { + if errors.Is(err, ErrEmailTaken) { + return nil, ErrEmailExists + } + return nil, ErrInternal + } + + pair, err := s.tokens.Issue(ctx, user.ID, user.UUID) + if err != nil { + return nil, ErrInternal + } + return pair, nil +} + +// verifyCode checks the supplied code against Redis. The code is consumed +// (deleted) on success, and after CodeMaxAttempts failed tries it is burned to +// stop brute forcing the 6-digit space. All failure modes return ErrCodeInvalid +// to avoid distinguishing wrong / expired / used. +func (s *Service) verifyCode(ctx context.Context, email, code string) *apierr.Error { + stored, err := s.rdb.Get(ctx, codeKey(email)).Result() + if err == redis.Nil { + return ErrCodeInvalid + } + if err != nil { + return ErrInternal + } + + // Count this attempt (TTL bounded by the code lifetime). + attempts, err := s.rdb.Incr(ctx, codeAttemptsKey(email)).Result() + if err != nil { + return ErrInternal + } + if attempts == 1 { + _ = s.rdb.Expire(ctx, codeAttemptsKey(email), s.cfg.CodeTTL).Err() + } + if attempts > int64(s.cfg.CodeMaxAttempts) { + // Burn the code and the counter. + _ = s.rdb.Del(ctx, codeKey(email), codeAttemptsKey(email)).Err() + return ErrCodeInvalid + } + + if subtle.ConstantTimeCompare([]byte(stored), []byte(code)) != 1 { + return ErrCodeInvalid + } + + // Success: consume the code (one-time use). + _ = s.rdb.Del(ctx, codeKey(email), codeAttemptsKey(email)).Err() + return nil +} + +// Login authenticates email+password with constant-time behaviour and a +// failure-count lock. retryAfter is non-zero only when the account is locked. +func (s *Service) Login(ctx context.Context, rawEmail, password, ip string) (*TokenPair, time.Duration, *apierr.Error) { + _ = ip // IP reserved for future per-IP login throttling; not logged. + email := NormalizeEmail(rawEmail) + if email == "" || password == "" { + return nil, 0, ErrInvalidRequest + } + + // Lock check. + count, ttl, err := s.rl.FailureCount(ctx, scopeLogin, email) + if err != nil { + return nil, 0, ErrInternal + } + if count >= int64(s.cfg.LoginFailMax) && ttl > 0 { + return nil, ttl, ErrAccountLocked + } + + user, err := s.store.GetUserByEmail(ctx, email) + if err != nil { + if errors.Is(err, ErrNotFound) { + // Spend equivalent CPU so timing doesn't reveal account existence. + ConstantTimeReject(password) + _, _ = s.rl.RecordFailure(ctx, scopeLogin, email, s.cfg.LoginLockWindow) + return nil, 0, ErrInvalidCredentials + } + return nil, 0, ErrInternal + } + + valid, verr := VerifyPassword(user.PwHash, password) + if verr != nil || !valid { + _, _ = s.rl.RecordFailure(ctx, scopeLogin, email, s.cfg.LoginLockWindow) + return nil, 0, ErrInvalidCredentials + } + + if user.Status == "banned" { + return nil, 0, ErrAccountBanned + } + + // Success: clear the failure counter and issue tokens. + _ = s.rl.ClearFailures(ctx, scopeLogin, email) + pair, err := s.tokens.Issue(ctx, user.ID, user.UUID) + if err != nil { + return nil, 0, ErrInternal + } + return pair, 0, nil +} + +// Refresh validates and rotates a refresh token. +func (s *Service) Refresh(ctx context.Context, refreshToken string) (*TokenPair, *apierr.Error) { + if refreshToken == "" { + return nil, ErrInvalidRequest + } + pair, err := s.tokens.Refresh(ctx, refreshToken) + if err != nil { + if errors.Is(err, ErrInvalidTokenSentinel) { + return nil, ErrInvalidToken + } + // Parse / signature / expiry failures all map to an opaque invalid-token. + return nil, ErrInvalidToken + } + return pair, nil +} + +// genNumericCode returns an n-digit numeric string drawn from crypto/rand. +func genNumericCode(n int) (string, error) { + const digits = "0123456789" + b := make([]byte, n) + for i := range b { + idx, err := rand.Int(rand.Reader, big.NewInt(int64(len(digits)))) + if err != nil { + return "", fmt.Errorf("auth: gen code: %w", err) + } + b[i] = digits[idx.Int64()] + } + return string(b), nil +} diff --git a/server/internal/auth/service_test.go b/server/internal/auth/service_test.go new file mode 100644 index 0000000..30ea98b --- /dev/null +++ b/server/internal/auth/service_test.go @@ -0,0 +1,271 @@ +package auth + +import ( + "context" + "testing" + "time" +) + +// newService wires a Service over a fake store + miniredis for unit tests. +// It returns the service, the redis client (to read codes directly), the fake +// store, and the miniredis handle (to fast-forward TTLs). +func newService(t *testing.T, cfg ServiceConfig) (*Service, *fakeStore, *captureMailer) { + t.Helper() + rdb, _ := newMiniRedis(t) + rl := NewRateLimiter(rdb, nil) + tm := newTokenManager(t, rdb, time.Now) + store := newFakeStore() + mailer := newCaptureMailer() + svc := NewService(store, rdb, rl, tm, mailer, cfg, nil) + return svc, store, mailer +} + +// codeInRedis reads the active verification code straight from Redis. +func codeInRedis(t *testing.T, svc *Service, email string) string { + t.Helper() + c, err := svc.rdb.Get(context.Background(), codeKey(email)).Result() + if err != nil { + t.Fatalf("no code stored for %s: %v", email, err) + } + return c +} + +func TestService_RegisterFullFlow(t *testing.T) { + svc, store, _ := newService(t, ServiceConfig{}) + ctx := context.Background() + const email = "alice@example.com" + + if _, err := svc.SendCode(ctx, email, "1.1.1.1"); err != nil { + t.Fatalf("SendCode: %v", err) + } + code := codeInRedis(t, svc, email) + + pair, apiErr := svc.Register(ctx, email, code, "supersecret") + if apiErr != nil { + t.Fatalf("Register: %v", apiErr) + } + if pair.AccessToken == "" || pair.RefreshToken == "" || pair.ExpiresIn != 900 { + t.Fatalf("bad token pair: %+v", pair) + } + + // User exists. + u, err := store.GetUserByEmail(ctx, email) + if err != nil { + t.Fatalf("user not created: %v", err) + } + // Trial subscription exists, PRO, ~7 days. + tr, ok := store.trials[u.ID] + if !ok { + t.Fatal("trial subscription not created") + } + if tr.plan != "pro" || tr.source != "trial" { + t.Fatalf("trial = %+v, want pro/trial", tr) + } + days := time.Until(tr.expiresAt).Hours() / 24 + if days < 6.9 || days > 7.1 { + t.Fatalf("trial length = %.2f days, want ~7", days) + } + + // Issued access token authenticates. + claims, perr := svc.tokens.ParseAccess(pair.AccessToken) + if perr != nil { + t.Fatalf("ParseAccess: %v", perr) + } + if claims.UID != u.ID || claims.Subject != u.UUID { + t.Fatalf("claims mismatch: %+v", claims) + } +} + +func TestService_DuplicateEmailConflict(t *testing.T) { + // Higher send limit so the two registration attempts can each request a code. + svc, _, _ := newService(t, ServiceConfig{EmailPerMinute: 10}) + ctx := context.Background() + const email = "dup@example.com" + + // First registration. + _, _ = svc.SendCode(ctx, email, "") + if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), "password1"); e != nil { + t.Fatalf("first register: %v", e) + } + + // Second: new code, but the email is already taken → 409. + _, _ = svc.SendCode(ctx, email, "") + _, apiErr := svc.Register(ctx, email, codeInRedis(t, svc, email), "password2") + if apiErr == nil || apiErr.Code != ErrEmailExists.Code { + t.Fatalf("want email_exists, got %v", apiErr) + } +} + +func TestService_CodeWrong(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{}) + ctx := context.Background() + const email = "wrong@example.com" + _, _ = svc.SendCode(ctx, email, "") + + _, apiErr := svc.Register(ctx, email, "000000", "password1") + if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code { + t.Fatalf("want code_invalid, got %v", apiErr) + } +} + +func TestService_CodeExpired(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{CodeTTL: time.Minute}) + ctx := context.Background() + const email = "expired@example.com" + _, _ = svc.SendCode(ctx, email, "") + code := codeInRedis(t, svc, email) + + // Expire the code key. + svc.rdb.Del(ctx, codeKey(email)) + + _, apiErr := svc.Register(ctx, email, code, "password1") + if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code { + t.Fatalf("want code_invalid after expiry, got %v", apiErr) + } +} + +func TestService_CodeReuseRejected(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{}) + ctx := context.Background() + const email = "reuse@example.com" + _, _ = svc.SendCode(ctx, email, "") + code := codeInRedis(t, svc, email) + + if _, e := svc.Register(ctx, email, code, "password1"); e != nil { + t.Fatalf("first register: %v", e) + } + // Re-using the consumed code must fail. + _, apiErr := svc.Register(ctx, "other@example.com", code, "password1") + if apiErr == nil || apiErr.Code != ErrCodeInvalid.Code { + t.Fatalf("want code_invalid on reuse, got %v", apiErr) + } +} + +func TestService_CodeBruteForceBurned(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{CodeMaxAttempts: 3}) + ctx := context.Background() + const email = "brute@example.com" + _, _ = svc.SendCode(ctx, email, "") + good := codeInRedis(t, svc, email) + + // 3 wrong attempts burn the code. + for i := 0; i < 3; i++ { + if _, e := svc.Register(ctx, email, "999999", "password1"); e == nil { + t.Fatal("wrong code should fail") + } + } + // Even the correct code no longer works. + if _, e := svc.Register(ctx, email, good, "password1"); e == nil || e.Code != ErrCodeInvalid.Code { + t.Fatalf("burned code should reject correct value, got %v", e) + } +} + +func TestService_SendCodeRateLimited(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{EmailPerMinute: 1}) + ctx := context.Background() + const email = "rl@example.com" + + if _, e := svc.SendCode(ctx, email, "9.9.9.9"); e != nil { + t.Fatalf("first send: %v", e) + } + ra, apiErr := svc.SendCode(ctx, email, "9.9.9.9") + if apiErr == nil || apiErr.Code != ErrRateLimited.Code { + t.Fatalf("want rate_limited, got %v", apiErr) + } + if ra <= 0 { + t.Fatalf("expected positive retry-after, got %v", ra) + } +} + +func TestService_SendCodeDisposableBlocked(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{}) + _, apiErr := svc.SendCode(context.Background(), "x@mailinator.com", "") + if apiErr == nil || apiErr.Code != ErrEmailDisposable.Code { + t.Fatalf("want email_disposable, got %v", apiErr) + } +} + +func TestService_LoginAndLockout(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{LoginFailMax: 3, LoginLockWindow: 15 * time.Minute}) + ctx := context.Background() + const email = "login@example.com" + const pw = "rightpassword" + + _, _ = svc.SendCode(ctx, email, "") + if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), pw); e != nil { + t.Fatalf("register: %v", e) + } + + // Correct login works. + pair, _, apiErr := svc.Login(ctx, email, pw, "") + if apiErr != nil || pair == nil { + t.Fatalf("login should succeed: %v", apiErr) + } + + // 3 wrong attempts. + for i := 0; i < 3; i++ { + _, _, e := svc.Login(ctx, email, "wrong", "") + if e == nil || e.Code != ErrInvalidCredentials.Code { + t.Fatalf("attempt %d want invalid_credentials, got %v", i, e) + } + } + // Now locked, even with the correct password. + _, ra, e := svc.Login(ctx, email, pw, "") + if e == nil || e.Code != ErrAccountLocked.Code { + t.Fatalf("want account_locked, got %v", e) + } + if ra <= 0 { + t.Fatalf("expected positive retry-after on lock, got %v", ra) + } +} + +func TestService_LoginUnknownUser(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{}) + _, _, apiErr := svc.Login(context.Background(), "ghost@example.com", "whatever", "") + if apiErr == nil || apiErr.Code != ErrInvalidCredentials.Code { + t.Fatalf("want invalid_credentials for unknown user, got %v", apiErr) + } +} + +func TestService_BannedUserRejected(t *testing.T) { + svc, store, _ := newService(t, ServiceConfig{}) + ctx := context.Background() + const email = "banned@example.com" + const pw = "password1" + + _, _ = svc.SendCode(ctx, email, "") + if _, e := svc.Register(ctx, email, codeInRedis(t, svc, email), pw); e != nil { + t.Fatalf("register: %v", e) + } + store.setStatus(email, "banned") + + _, _, apiErr := svc.Login(ctx, email, pw, "") + if apiErr == nil || apiErr.Code != ErrAccountBanned.Code { + t.Fatalf("want account_banned, got %v", apiErr) + } +} + +func TestService_RefreshRotation(t *testing.T) { + svc, _, _ := newService(t, ServiceConfig{}) + ctx := context.Background() + const email = "refresh@example.com" + + _, _ = svc.SendCode(ctx, email, "") + pair, e := svc.Register(ctx, email, codeInRedis(t, svc, email), "password1") + if e != nil { + t.Fatalf("register: %v", e) + } + + rotated, apiErr := svc.Refresh(ctx, pair.RefreshToken) + if apiErr != nil { + t.Fatalf("refresh: %v", apiErr) + } + // Old refresh token now invalid. + if _, e := svc.Refresh(ctx, pair.RefreshToken); e == nil || e.Code != ErrInvalidToken.Code { + t.Fatalf("want invalid_token for rotated-out refresh, got %v", e) + } + // New one works. + if _, e := svc.Refresh(ctx, rotated.RefreshToken); e != nil { + t.Fatalf("new refresh should work: %v", e) + } +} diff --git a/server/internal/auth/store.go b/server/internal/auth/store.go new file mode 100644 index 0000000..b155dd8 --- /dev/null +++ b/server/internal/auth/store.go @@ -0,0 +1,133 @@ +package auth + +import ( + "context" + "database/sql" + "errors" + "fmt" + "strings" + "time" + + "github.com/google/uuid" +) + +// User is the subset of the users row the auth module needs. +type User struct { + ID int64 + UUID string + Email string + PwHash string + DpUUID string + Status string // "active" | "banned" +} + +// Sentinel store errors. Service maps these to API errors. +var ( + // ErrEmailTaken is returned by CreateUserWithTrial on a duplicate email. + ErrEmailTaken = errors.New("auth: email already registered") + // ErrNotFound is returned when a user lookup yields no row. + ErrNotFound = errors.New("auth: user not found") +) + +// UserStore is the persistence contract for the auth module. The MySQL +// implementation lives in this file; tests substitute an in-memory fake. +type UserStore interface { + // CreateUserWithTrial atomically inserts a new user and a 7-day PRO trial + // subscription (source='trial') in a single transaction. The email UNIQUE + // constraint guarantees a single trial per address; a duplicate returns + // ErrEmailTaken. trialDays controls the trial length. + CreateUserWithTrial(ctx context.Context, email, pwHash string, trialDays int) (*User, error) + // GetUserByEmail returns the user for login. ErrNotFound when absent. + GetUserByEmail(ctx context.Context, email string) (*User, error) +} + +// SQLStore is the MySQL-backed UserStore. +type SQLStore struct { + db *sql.DB +} + +// NewSQLStore builds a SQLStore. +func NewSQLStore(db *sql.DB) *SQLStore { return &SQLStore{db: db} } + +// CreateUserWithTrial implements UserStore. +func (s *SQLStore) CreateUserWithTrial(ctx context.Context, email, pwHash string, trialDays int) (*User, error) { + userUUID := uuid.NewString() + dpUUID := uuid.NewString() + + tx, err := s.db.BeginTx(ctx, &sql.TxOptions{Isolation: sql.LevelReadCommitted}) + if err != nil { + return nil, fmt.Errorf("auth: begin tx: %w", err) + } + committed := false + defer func() { + if !committed { + _ = tx.Rollback() + } + }() + + res, err := tx.ExecContext(ctx, + `INSERT INTO users (uuid, email, pw_hash, dp_uuid, status, created_at) + VALUES (?, ?, ?, ?, 'active', UTC_TIMESTAMP(6))`, + userUUID, email, pwHash, dpUUID) + if err != nil { + if isDuplicateKey(err) { + return nil, ErrEmailTaken + } + return nil, fmt.Errorf("auth: insert user: %w", err) + } + userID, err := res.LastInsertId() + if err != nil { + return nil, fmt.Errorf("auth: user last id: %w", err) + } + + // Resolve the PRO plan id and insert the trial subscription. + var proID int64 + if err := tx.QueryRowContext(ctx, `SELECT id FROM plans WHERE code='pro'`).Scan(&proID); err != nil { + return nil, fmt.Errorf("auth: lookup pro plan: %w", err) + } + expires := time.Now().UTC().AddDate(0, 0, trialDays) + if _, err := tx.ExecContext(ctx, + `INSERT INTO subscriptions (user_id, plan_id, expires_at, source, created_at) + VALUES (?, ?, ?, 'trial', UTC_TIMESTAMP(6))`, + userID, proID, expires); err != nil { + return nil, fmt.Errorf("auth: insert trial subscription: %w", err) + } + + if err := tx.Commit(); err != nil { + return nil, fmt.Errorf("auth: commit: %w", err) + } + committed = true + + return &User{ + ID: userID, + UUID: userUUID, + Email: email, + PwHash: pwHash, + DpUUID: dpUUID, + Status: "active", + }, nil +} + +// GetUserByEmail implements UserStore. +func (s *SQLStore) GetUserByEmail(ctx context.Context, email string) (*User, error) { + var u User + err := s.db.QueryRowContext(ctx, + `SELECT id, uuid, email, pw_hash, dp_uuid, status FROM users WHERE email = ?`, + email).Scan(&u.ID, &u.UUID, &u.Email, &u.PwHash, &u.DpUUID, &u.Status) + if err == sql.ErrNoRows { + return nil, ErrNotFound + } + if err != nil { + return nil, fmt.Errorf("auth: get user by email: %w", err) + } + return &u, nil +} + +// isDuplicateKey reports whether err is a MySQL duplicate-key (1062) error. +func isDuplicateKey(err error) bool { + if err == nil { + return false + } + msg := err.Error() + return strings.Contains(msg, "Duplicate entry") || strings.Contains(msg, "1062") +} diff --git a/server/internal/auth/token.go b/server/internal/auth/token.go new file mode 100644 index 0000000..8bff725 --- /dev/null +++ b/server/internal/auth/token.go @@ -0,0 +1,290 @@ +package auth + +import ( + "context" + "crypto/rsa" + "crypto/x509" + "encoding/pem" + "errors" + "fmt" + "time" + + "github.com/golang-jwt/jwt/v4" + "github.com/google/uuid" + "github.com/redis/go-redis/v9" +) + +// Token type discriminator carried in the `typ` claim so an access token can +// never be replayed as a refresh token and vice-versa. +const ( + typAccess = "access" + typRefresh = "refresh" +) + +// Default lifetimes (doc/02 §4.1): access 15 min, refresh 30 days. +const ( + DefaultAccessTTL = 15 * time.Minute + DefaultRefreshTTL = 30 * 24 * time.Hour +) + +// refreshKeyPrefix is the Redis whitelist prefix for refresh-token JTIs. +const refreshKeyPrefix = "jwt:refresh:" + +// Claims is the JWT payload. Subject holds the user UUID; UID carries the +// numeric primary key so downstream middleware/stores avoid a DB round-trip. +type Claims struct { + jwt.RegisteredClaims + UID int64 `json:"uid"` + Typ string `json:"typ"` +} + +// TokenPair is the issued access/refresh pair plus the access lifetime seconds. +type TokenPair struct { + AccessToken string + RefreshToken string + ExpiresIn int // access_token validity in seconds +} + +// TokenManager signs tokens with one RSA private key (identified by kid) and +// verifies with a set of public keys keyed by kid — accepting both the current +// and previous keys to support zero-downtime key rotation. Refresh tokens are +// whitelisted in Redis so logout/ban takes effect immediately. +type TokenManager struct { + signKey *rsa.PrivateKey + signKID string + verifyKeys map[string]*rsa.PublicKey // kid -> public key (current + old) + accessTTL time.Duration + refreshTTL time.Duration + rdb *redis.Client + now func() time.Time +} + +// TokenConfig configures a TokenManager. +type TokenConfig struct { + // SignKey is the active RSA private key used to sign new tokens. + SignKey *rsa.PrivateKey + // SignKID is the key id written into the JWT header. + SignKID string + // VerifyKeys maps kid -> public key. Must contain SignKID; may contain + // additional (older) keys still accepted during rotation. If nil, the + // public part of SignKey under SignKID is used. + VerifyKeys map[string]*rsa.PublicKey + AccessTTL time.Duration + RefreshTTL time.Duration + // Now is an optional clock override for tests. + Now func() time.Time +} + +// NewTokenManager constructs a TokenManager, validating the key material. +func NewTokenManager(rdb *redis.Client, cfg TokenConfig) (*TokenManager, error) { + if cfg.SignKey == nil { + return nil, errors.New("auth: token manager requires a signing key") + } + if cfg.SignKID == "" { + return nil, errors.New("auth: token manager requires a signing key id (kid)") + } + verify := cfg.VerifyKeys + if verify == nil { + verify = map[string]*rsa.PublicKey{} + } + if _, ok := verify[cfg.SignKID]; !ok { + verify[cfg.SignKID] = &cfg.SignKey.PublicKey + } + accessTTL := cfg.AccessTTL + if accessTTL <= 0 { + accessTTL = DefaultAccessTTL + } + refreshTTL := cfg.RefreshTTL + if refreshTTL <= 0 { + refreshTTL = DefaultRefreshTTL + } + now := cfg.Now + if now == nil { + now = time.Now + } + return &TokenManager{ + signKey: cfg.SignKey, + signKID: cfg.SignKID, + verifyKeys: verify, + accessTTL: accessTTL, + refreshTTL: refreshTTL, + rdb: rdb, + now: now, + }, nil +} + +// Issue mints a fresh access+refresh pair for the user and whitelists the +// refresh JTI in Redis with the refresh TTL. +func (tm *TokenManager) Issue(ctx context.Context, userID int64, userUUID string) (*TokenPair, error) { + now := tm.now() + + access, _, err := tm.sign(userID, userUUID, typAccess, tm.accessTTL, now) + if err != nil { + return nil, err + } + refresh, refreshJTI, err := tm.sign(userID, userUUID, typRefresh, tm.refreshTTL, now) + if err != nil { + return nil, err + } + if err := tm.whitelist(ctx, refreshJTI, userID); err != nil { + return nil, err + } + return &TokenPair{ + AccessToken: access, + RefreshToken: refresh, + ExpiresIn: int(tm.accessTTL.Seconds()), + }, nil +} + +// sign builds and signs one token, returning the compact string and its JTI. +func (tm *TokenManager) sign(userID int64, userUUID, typ string, ttl time.Duration, now time.Time) (string, string, error) { + jti := uuid.NewString() + claims := Claims{ + RegisteredClaims: jwt.RegisteredClaims{ + Subject: userUUID, + ID: jti, + IssuedAt: jwt.NewNumericDate(now), + ExpiresAt: jwt.NewNumericDate(now.Add(ttl)), + }, + UID: userID, + Typ: typ, + } + tok := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) + tok.Header["kid"] = tm.signKID + signed, err := tok.SignedString(tm.signKey) + if err != nil { + return "", "", fmt.Errorf("auth: sign token: %w", err) + } + return signed, jti, nil +} + +// keyfunc resolves the verification key from the token's kid header and rejects +// any algorithm other than RS256. +func (tm *TokenManager) keyfunc(t *jwt.Token) (interface{}, error) { + if _, ok := t.Method.(*jwt.SigningMethodRSA); !ok { + return nil, fmt.Errorf("auth: unexpected signing method %q", t.Header["alg"]) + } + kid, _ := t.Header["kid"].(string) + if kid == "" { + return nil, errors.New("auth: token missing kid") + } + pub, ok := tm.verifyKeys[kid] + if !ok { + return nil, fmt.Errorf("auth: unknown kid %q", kid) + } + return pub, nil +} + +// parse validates signature, expiry, and the expected token type. Expiry is +// checked manually against tm.now so tests can inject a clock without mutating +// the package-global jwt.TimeFunc. +func (tm *TokenManager) parse(tokenStr, wantTyp string) (*Claims, error) { + claims := &Claims{} + parser := jwt.NewParser( + jwt.WithValidMethods([]string{"RS256"}), + jwt.WithoutClaimsValidation(), // we validate exp/iat ourselves below + ) + if _, err := parser.ParseWithClaims(tokenStr, claims, tm.keyfunc); err != nil { + return nil, err + } + + now := tm.now() + if claims.ExpiresAt == nil || now.After(claims.ExpiresAt.Time) { + return nil, errors.New("auth: token expired") + } + if claims.IssuedAt != nil && now.Add(time.Minute).Before(claims.IssuedAt.Time) { + return nil, errors.New("auth: token used before issued") + } + if claims.Typ != wantTyp { + return nil, fmt.Errorf("auth: token type %q, want %q", claims.Typ, wantTyp) + } + return claims, nil +} + +// ParseAccess validates an access token and returns its claims. +func (tm *TokenManager) ParseAccess(tokenStr string) (*Claims, error) { + return tm.parse(tokenStr, typAccess) +} + +// Refresh validates a refresh token against the Redis whitelist, then rotates: +// the old JTI is deleted and a brand-new access+refresh pair is issued, so the +// presented refresh token can never be replayed. +func (tm *TokenManager) Refresh(ctx context.Context, refreshToken string) (*TokenPair, error) { + claims, err := tm.parse(refreshToken, typRefresh) + if err != nil { + return nil, err + } + + // Whitelist check + single-use rotation: DEL returns the number of keys + // removed; 0 means the JTI was absent (already rotated, logged out, or + // banned) → reject. + removed, err := tm.rdb.Del(ctx, refreshKeyPrefix+claims.ID).Result() + if err != nil { + return nil, fmt.Errorf("auth: refresh whitelist del: %w", err) + } + if removed == 0 { + return nil, ErrInvalidTokenSentinel + } + + return tm.Issue(ctx, claims.UID, claims.Subject) +} + +// Revoke removes a single refresh JTI from the whitelist (logout). +func (tm *TokenManager) Revoke(ctx context.Context, jti string) error { + return tm.rdb.Del(ctx, refreshKeyPrefix+jti).Err() +} + +// whitelist stores the refresh JTI with the refresh TTL. +func (tm *TokenManager) whitelist(ctx context.Context, jti string, userID int64) error { + if err := tm.rdb.Set(ctx, refreshKeyPrefix+jti, userID, tm.refreshTTL).Err(); err != nil { + return fmt.Errorf("auth: refresh whitelist set: %w", err) + } + return nil +} + +// ErrInvalidTokenSentinel is returned by Refresh when the token is structurally +// valid but no longer whitelisted. Callers map it to ErrInvalidToken. +var ErrInvalidTokenSentinel = errors.New("auth: refresh token not in whitelist") + +// -------------------------------------------------------------------------- +// PEM loading helpers (used by wiring/config to build a TokenManager) +// -------------------------------------------------------------------------- + +// LoadPrivateKeyPEM parses a PEM-encoded RSA private key (PKCS#1 or PKCS#8). +func LoadPrivateKeyPEM(pemBytes []byte) (*rsa.PrivateKey, error) { + block, _ := pem.Decode(pemBytes) + if block == nil { + return nil, errors.New("auth: no PEM block in private key") + } + if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil { + return key, nil + } + parsed, err := x509.ParsePKCS8PrivateKey(block.Bytes) + if err != nil { + return nil, fmt.Errorf("auth: parse private key: %w", err) + } + rsaKey, ok := parsed.(*rsa.PrivateKey) + if !ok { + return nil, errors.New("auth: private key is not RSA") + } + return rsaKey, nil +} + +// LoadPublicKeyPEM parses a PEM-encoded RSA public key (PKIX or PKCS#1). +func LoadPublicKeyPEM(pemBytes []byte) (*rsa.PublicKey, error) { + block, _ := pem.Decode(pemBytes) + if block == nil { + return nil, errors.New("auth: no PEM block in public key") + } + if pub, err := x509.ParsePKIXPublicKey(block.Bytes); err == nil { + if rsaPub, ok := pub.(*rsa.PublicKey); ok { + return rsaPub, nil + } + return nil, errors.New("auth: public key is not RSA") + } + rsaPub, err := x509.ParsePKCS1PublicKey(block.Bytes) + if err != nil { + return nil, fmt.Errorf("auth: parse public key: %w", err) + } + return rsaPub, nil +} diff --git a/server/internal/auth/token_test.go b/server/internal/auth/token_test.go new file mode 100644 index 0000000..b063446 --- /dev/null +++ b/server/internal/auth/token_test.go @@ -0,0 +1,127 @@ +package auth + +import ( + "context" + "crypto/rsa" + "testing" + "time" +) + +func TestToken_IssueAndParseAccess(t *testing.T) { + rdb, _ := newMiniRedis(t) + tm := newTokenManager(t, rdb, time.Now) + + pair, err := tm.Issue(context.Background(), 42, "uuid-42") + if err != nil { + t.Fatalf("Issue: %v", err) + } + if pair.ExpiresIn != 900 { + t.Errorf("ExpiresIn = %d, want 900", pair.ExpiresIn) + } + + claims, err := tm.ParseAccess(pair.AccessToken) + if err != nil { + t.Fatalf("ParseAccess: %v", err) + } + if claims.UID != 42 || claims.Subject != "uuid-42" { + t.Errorf("claims = %+v", claims) + } + if claims.Typ != typAccess { + t.Errorf("typ = %q", claims.Typ) + } +} + +func TestToken_RefreshCannotBeUsedAsAccess(t *testing.T) { + rdb, _ := newMiniRedis(t) + tm := newTokenManager(t, rdb, time.Now) + pair, _ := tm.Issue(context.Background(), 1, "u1") + if _, err := tm.ParseAccess(pair.RefreshToken); err == nil { + t.Fatal("refresh token must not parse as access") + } +} + +func TestToken_AccessExpiry(t *testing.T) { + rdb, _ := newMiniRedis(t) + base := time.Now() + clock := base + tm := newTokenManager(t, rdb, func() time.Time { return clock }) + + pair, _ := tm.Issue(context.Background(), 1, "u1") + clock = base.Add(16 * time.Minute) // past 15-min access TTL + if _, err := tm.ParseAccess(pair.AccessToken); err == nil { + t.Fatal("expected expired access token to fail") + } +} + +func TestToken_RefreshRotationInvalidatesOld(t *testing.T) { + rdb, _ := newMiniRedis(t) + tm := newTokenManager(t, rdb, time.Now) + ctx := context.Background() + + pair, _ := tm.Issue(ctx, 7, "u7") + + rotated, err := tm.Refresh(ctx, pair.RefreshToken) + if err != nil { + t.Fatalf("Refresh: %v", err) + } + if rotated.RefreshToken == pair.RefreshToken { + t.Fatal("rotation should produce a new refresh token") + } + + // Old refresh token is now invalid. + if _, err := tm.Refresh(ctx, pair.RefreshToken); err == nil { + t.Fatal("old refresh token must be rejected after rotation") + } + + // New refresh token still works. + if _, err := tm.Refresh(ctx, rotated.RefreshToken); err != nil { + t.Fatalf("new refresh token should work: %v", err) + } +} + +func TestToken_RevokeRefresh(t *testing.T) { + rdb, _ := newMiniRedis(t) + tm := newTokenManager(t, rdb, time.Now) + ctx := context.Background() + + pair, _ := tm.Issue(ctx, 3, "u3") + claims, err := tm.parse(pair.RefreshToken, typRefresh) + if err != nil { + t.Fatalf("parse refresh: %v", err) + } + if err := tm.Revoke(ctx, claims.ID); err != nil { + t.Fatalf("Revoke: %v", err) + } + if _, err := tm.Refresh(ctx, pair.RefreshToken); err == nil { + t.Fatal("revoked refresh token must be rejected") + } +} + +func TestToken_KeyRotationAcceptsOldKey(t *testing.T) { + rdb, _ := newMiniRedis(t) + + // Old manager signs with k1. + oldKey := newRSAKey(t) + oldTM, err := NewTokenManager(rdb, TokenConfig{SignKey: oldKey, SignKID: "k1"}) + if err != nil { + t.Fatal(err) + } + pair, _ := oldTM.Issue(context.Background(), 9, "u9") + + // New manager signs with k2 but still accepts k1 for verification. + newKey := newRSAKey(t) + newTM, err := NewTokenManager(rdb, TokenConfig{ + SignKey: newKey, + SignKID: "k2", + VerifyKeys: map[string]*rsa.PublicKey{ + "k1": &oldKey.PublicKey, + "k2": &newKey.PublicKey, + }, + }) + if err != nil { + t.Fatal(err) + } + if _, err := newTM.ParseAccess(pair.AccessToken); err != nil { + t.Fatalf("token signed with old key must still verify: %v", err) + } +} diff --git a/server/internal/config/config.go b/server/internal/config/config.go index fcfd5e5..93e45f9 100644 --- a/server/internal/config/config.go +++ b/server/internal/config/config.go @@ -4,6 +4,7 @@ package config import ( "fmt" "os" + "strings" "time" ) @@ -37,6 +38,16 @@ type Config struct { // WebhookNonceTTL is how long a webhook nonce is kept in Redis to prevent replay. // Should be > 2 * WebhookTimestampTolerance. Default: 15 minutes. WebhookNonceTTL time.Duration + + // ── Auth / JWT (RS256) ──────────────────────────────────────────────── + // JWTPrivateKeyPath is the PEM file holding the active RS256 signing key. + JWTPrivateKeyPath string + // JWTKeyID is the `kid` written into the JWT header (identifies the signing key). + JWTKeyID string + // JWTPublicKeys maps kid -> PEM public-key file path. It must include the + // active key's kid and may carry previous keys still accepted during + // rotation. Parsed from JWT_PUBLIC_KEYS="kid1:/path1,kid2:/path2". + JWTPublicKeys map[string]string } // FromEnv reads configuration from environment variables. @@ -52,6 +63,9 @@ func FromEnv() (*Config, error) { RedeemLockDuration: time.Hour, WebhookTimestampTolerance: 5 * time.Minute, WebhookNonceTTL: 15 * time.Minute, + JWTPrivateKeyPath: os.Getenv("JWT_PRIVATE_KEY_PATH"), + JWTKeyID: os.Getenv("JWT_KEY_ID"), + JWTPublicKeys: parseKeyMap(os.Getenv("JWT_PUBLIC_KEYS")), } if c.DSN == "" { @@ -69,3 +83,24 @@ func getEnvDefault(key, def string) string { } return def } + +// parseKeyMap parses a "kid1:/path1,kid2:/path2" string into a map. Empty input +// yields a nil map. Malformed entries (missing ':') are skipped. +func parseKeyMap(raw string) map[string]string { + if raw == "" { + return nil + } + m := map[string]string{} + for _, pair := range strings.Split(raw, ",") { + pair = strings.TrimSpace(pair) + if pair == "" { + continue + } + i := strings.IndexByte(pair, ':') + if i <= 0 || i == len(pair)-1 { + continue + } + m[strings.TrimSpace(pair[:i])] = strings.TrimSpace(pair[i+1:]) + } + return m +}