| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330 |
- package handlers
- import (
- "crypto/rand"
- "crypto/sha256"
- "encoding/base64"
- "encoding/hex"
- "errors"
- "net/http"
- "strings"
- "time"
- "github.com/celestia-trace/backend/config"
- "github.com/celestia-trace/backend/models"
- "github.com/gin-gonic/gin"
- "github.com/golang-jwt/jwt/v5"
- "golang.org/x/crypto/bcrypt"
- "gorm.io/gorm"
- )
- type AuthHandler struct {
- DB *gorm.DB
- Cfg *config.Config
- }
- func respondError(c *gin.Context, status int, code int, message string) {
- c.JSON(status, models.ErrorResponse{Code: code, Message: message})
- }
- func respondSuccess(c *gin.Context, data interface{}) {
- c.JSON(http.StatusOK, gin.H{"code": 0, "message": "success", "data": data})
- }
- func (h *AuthHandler) generateAccessToken(userID, sessionID string) (string, time.Time, error) {
- now := time.Now().UTC()
- expiresAt := now.Add(h.Cfg.AccessTokenTTL)
- token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
- "sub": userID,
- "sid": sessionID,
- "iss": "celestia-trace",
- "iat": now.Unix(),
- "exp": expiresAt.Unix(),
- })
- signed, err := token.SignedString([]byte(h.Cfg.JWTSecret))
- return signed, expiresAt, err
- }
- func newRefreshToken() (string, string, error) {
- raw := make([]byte, 32)
- if _, err := rand.Read(raw); err != nil {
- return "", "", err
- }
- token := base64.RawURLEncoding.EncodeToString(raw)
- sum := sha256.Sum256([]byte(token))
- return token, hex.EncodeToString(sum[:]), nil
- }
- func hashRefreshToken(token string) string {
- sum := sha256.Sum256([]byte(token))
- return hex.EncodeToString(sum[:])
- }
- func (h *AuthHandler) issueTokens(tx *gorm.DB, user models.User) (models.AuthResponse, error) {
- refreshToken, refreshHash, err := newRefreshToken()
- if err != nil {
- return models.AuthResponse{}, err
- }
- session := models.AuthSession{
- UserID: user.ID,
- RefreshHash: refreshHash,
- ExpiresAt: time.Now().UTC().Add(h.Cfg.RefreshTokenTTL),
- }
- if err := tx.Create(&session).Error; err != nil {
- return models.AuthResponse{}, err
- }
- accessToken, expiresAt, err := h.generateAccessToken(user.ID, session.ID)
- if err != nil {
- return models.AuthResponse{}, err
- }
- return models.AuthResponse{
- User: user,
- Token: accessToken,
- RefreshToken: refreshToken,
- ExpiresAt: expiresAt,
- }, nil
- }
- func normalizeIdentifier(identifier string) string {
- identifier = strings.TrimSpace(identifier)
- if strings.Contains(identifier, "@") {
- return strings.ToLower(identifier)
- }
- return strings.ReplaceAll(identifier, " ", "")
- }
- func isEmail(identifier string) bool {
- return strings.Contains(identifier, "@")
- }
- func isDuplicateError(err error) bool {
- if err == nil {
- return false
- }
- message := strings.ToLower(err.Error())
- return strings.Contains(message, "duplicate key") || strings.Contains(message, "unique constraint")
- }
- // Register creates the server account and an immediately revocable login session.
- func (h *AuthHandler) Register(c *gin.Context) {
- var req models.RegisterRequest
- if err := c.ShouldBindJSON(&req); err != nil {
- respondError(c, http.StatusBadRequest, 400, "invalid input data")
- return
- }
- req.Username = strings.TrimSpace(req.Username)
- req.Identifier = normalizeIdentifier(req.Identifier)
- hash, err := bcrypt.GenerateFromPassword([]byte(req.Password), bcrypt.DefaultCost)
- if err != nil {
- respondError(c, http.StatusInternalServerError, 500, "failed to hash password")
- return
- }
- user := models.User{Username: req.Username, PasswordHash: string(hash)}
- if isEmail(req.Identifier) {
- user.Email = &req.Identifier
- } else {
- user.PhoneNumber = &req.Identifier
- }
- var response models.AuthResponse
- err = h.DB.Transaction(func(tx *gorm.DB) error {
- if err := tx.Create(&user).Error; err != nil {
- return err
- }
- issued, err := h.issueTokens(tx, user)
- response = issued
- return err
- })
- if err != nil {
- if isDuplicateError(err) {
- respondError(c, http.StatusConflict, 409, "username or identifier already exists")
- return
- }
- respondError(c, http.StatusInternalServerError, 500, "failed to create user")
- return
- }
- respondSuccess(c, response)
- }
- // Login supports username, email, or phone number.
- func (h *AuthHandler) Login(c *gin.Context) {
- var req models.LoginRequest
- if err := c.ShouldBindJSON(&req); err != nil {
- respondError(c, http.StatusBadRequest, 400, "invalid input data")
- return
- }
- identifier := normalizeIdentifier(req.Identifier)
- var user models.User
- if err := h.DB.Where("LOWER(username) = LOWER(?) OR LOWER(email) = LOWER(?) OR phone_number = ?", identifier, identifier, identifier).First(&user).Error; err != nil {
- if errors.Is(err, gorm.ErrRecordNotFound) {
- respondError(c, http.StatusUnauthorized, 401, "invalid credentials")
- return
- }
- respondError(c, http.StatusInternalServerError, 500, "database error")
- return
- }
- if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.Password)); err != nil {
- respondError(c, http.StatusUnauthorized, 401, "invalid credentials")
- return
- }
- response, err := h.issueTokens(h.DB, user)
- if err != nil {
- respondError(c, http.StatusInternalServerError, 500, "failed to create login session")
- return
- }
- respondSuccess(c, response)
- }
- // Refresh rotates a refresh token so a stolen token cannot be replayed.
- func (h *AuthHandler) Refresh(c *gin.Context) {
- var req models.RefreshRequest
- if err := c.ShouldBindJSON(&req); err != nil {
- respondError(c, http.StatusBadRequest, 400, "refresh token is required")
- return
- }
- var oldSession models.AuthSession
- err := h.DB.Where("refresh_hash = ? AND revoked_at IS NULL AND expires_at > ?", hashRefreshToken(req.RefreshToken), time.Now().UTC()).First(&oldSession).Error
- if err != nil {
- respondError(c, http.StatusUnauthorized, 401, "invalid or expired refresh token")
- return
- }
- var user models.User
- if err := h.DB.First(&user, "id = ?", oldSession.UserID).Error; err != nil {
- respondError(c, http.StatusUnauthorized, 401, "user no longer exists")
- return
- }
- var response models.AuthResponse
- err = h.DB.Transaction(func(tx *gorm.DB) error {
- now := time.Now().UTC()
- result := tx.Model(&models.AuthSession{}).
- Where("id = ? AND revoked_at IS NULL", oldSession.ID).
- Updates(map[string]interface{}{"revoked_at": &now, "last_used_at": &now})
- if result.Error != nil {
- return result.Error
- }
- if result.RowsAffected != 1 {
- return gorm.ErrRecordNotFound
- }
- issued, err := h.issueTokens(tx, user)
- response = issued
- return err
- })
- if err != nil {
- respondError(c, http.StatusUnauthorized, 401, "refresh token has already been used")
- return
- }
- respondSuccess(c, response)
- }
- func (h *AuthHandler) Logout(c *gin.Context) {
- sessionID := c.GetString("sessionID")
- if sessionID != "" {
- now := time.Now().UTC()
- h.DB.Model(&models.AuthSession{}).Where("id = ?", sessionID).Update("revoked_at", &now)
- }
- respondSuccess(c, nil)
- }
- func (h *AuthHandler) GetProfile(c *gin.Context) {
- userID := c.GetString("userID")
- var user models.User
- if err := h.DB.First(&user, "id = ?", userID).Error; err != nil {
- respondError(c, http.StatusNotFound, 404, "user not found")
- return
- }
- respondSuccess(c, user)
- }
- func (h *AuthHandler) UpdateProfile(c *gin.Context) {
- userID := c.GetString("userID")
- var req models.UpdateProfileRequest
- if err := c.ShouldBindJSON(&req); err != nil {
- respondError(c, http.StatusBadRequest, 400, "invalid input data")
- return
- }
- var user models.User
- if err := h.DB.First(&user, "id = ?", userID).Error; err != nil {
- respondError(c, http.StatusNotFound, 404, "user not found")
- return
- }
- if req.Username != nil {
- trimmed := strings.TrimSpace(*req.Username)
- if len([]rune(trimmed)) < 2 {
- respondError(c, http.StatusBadRequest, 400, "username must contain at least 2 characters")
- return
- }
- user.Username = trimmed
- }
- if req.Email != nil {
- value := strings.ToLower(strings.TrimSpace(*req.Email))
- if value == "" {
- user.Email = nil
- } else {
- user.Email = &value
- }
- }
- if req.PhoneNumber != nil {
- value := normalizeIdentifier(*req.PhoneNumber)
- if value == "" {
- user.PhoneNumber = nil
- } else {
- user.PhoneNumber = &value
- }
- }
- if req.AvatarURL != nil {
- value := strings.TrimSpace(*req.AvatarURL)
- if value == "" {
- user.AvatarURL = nil
- } else {
- user.AvatarURL = &value
- }
- }
- if err := h.DB.Save(&user).Error; err != nil {
- if isDuplicateError(err) {
- respondError(c, http.StatusConflict, 409, "username, email, or phone already exists")
- return
- }
- respondError(c, http.StatusInternalServerError, 500, "failed to update profile")
- return
- }
- respondSuccess(c, user)
- }
- func (h *AuthHandler) ChangePassword(c *gin.Context) {
- userID := c.GetString("userID")
- var req models.ChangePasswordRequest
- if err := c.ShouldBindJSON(&req); err != nil {
- respondError(c, http.StatusBadRequest, 400, "invalid input data")
- return
- }
- var user models.User
- if err := h.DB.First(&user, "id = ?", userID).Error; err != nil {
- respondError(c, http.StatusNotFound, 404, "user not found")
- return
- }
- if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.OldPassword)); err != nil {
- respondError(c, http.StatusUnauthorized, 401, "invalid old password")
- return
- }
- hash, err := bcrypt.GenerateFromPassword([]byte(req.NewPassword), bcrypt.DefaultCost)
- if err != nil {
- respondError(c, http.StatusInternalServerError, 500, "failed to hash new password")
- return
- }
- err = h.DB.Transaction(func(tx *gorm.DB) error {
- if err := tx.Model(&user).Update("password_hash", string(hash)).Error; err != nil {
- return err
- }
- now := time.Now().UTC()
- return tx.Model(&models.AuthSession{}).
- Where("user_id = ? AND id <> ? AND revoked_at IS NULL", userID, c.GetString("sessionID")).
- Update("revoked_at", &now).Error
- })
- if err != nil {
- respondError(c, http.StatusInternalServerError, 500, "failed to change password")
- return
- }
- respondSuccess(c, nil)
- }
|