auth_handler.go 9.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330
  1. package handlers
  2. import (
  3. "crypto/rand"
  4. "crypto/sha256"
  5. "encoding/base64"
  6. "encoding/hex"
  7. "errors"
  8. "net/http"
  9. "strings"
  10. "time"
  11. "github.com/celestia-trace/backend/config"
  12. "github.com/celestia-trace/backend/models"
  13. "github.com/gin-gonic/gin"
  14. "github.com/golang-jwt/jwt/v5"
  15. "golang.org/x/crypto/bcrypt"
  16. "gorm.io/gorm"
  17. )
  18. type AuthHandler struct {
  19. DB *gorm.DB
  20. Cfg *config.Config
  21. }
  22. func respondError(c *gin.Context, status int, code int, message string) {
  23. c.JSON(status, models.ErrorResponse{Code: code, Message: message})
  24. }
  25. func respondSuccess(c *gin.Context, data interface{}) {
  26. c.JSON(http.StatusOK, gin.H{"code": 0, "message": "success", "data": data})
  27. }
  28. func (h *AuthHandler) generateAccessToken(userID, sessionID string) (string, time.Time, error) {
  29. now := time.Now().UTC()
  30. expiresAt := now.Add(h.Cfg.AccessTokenTTL)
  31. token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
  32. "sub": userID,
  33. "sid": sessionID,
  34. "iss": "celestia-trace",
  35. "iat": now.Unix(),
  36. "exp": expiresAt.Unix(),
  37. })
  38. signed, err := token.SignedString([]byte(h.Cfg.JWTSecret))
  39. return signed, expiresAt, err
  40. }
  41. func newRefreshToken() (string, string, error) {
  42. raw := make([]byte, 32)
  43. if _, err := rand.Read(raw); err != nil {
  44. return "", "", err
  45. }
  46. token := base64.RawURLEncoding.EncodeToString(raw)
  47. sum := sha256.Sum256([]byte(token))
  48. return token, hex.EncodeToString(sum[:]), nil
  49. }
  50. func hashRefreshToken(token string) string {
  51. sum := sha256.Sum256([]byte(token))
  52. return hex.EncodeToString(sum[:])
  53. }
  54. func (h *AuthHandler) issueTokens(tx *gorm.DB, user models.User) (models.AuthResponse, error) {
  55. refreshToken, refreshHash, err := newRefreshToken()
  56. if err != nil {
  57. return models.AuthResponse{}, err
  58. }
  59. session := models.AuthSession{
  60. UserID: user.ID,
  61. RefreshHash: refreshHash,
  62. ExpiresAt: time.Now().UTC().Add(h.Cfg.RefreshTokenTTL),
  63. }
  64. if err := tx.Create(&session).Error; err != nil {
  65. return models.AuthResponse{}, err
  66. }
  67. accessToken, expiresAt, err := h.generateAccessToken(user.ID, session.ID)
  68. if err != nil {
  69. return models.AuthResponse{}, err
  70. }
  71. return models.AuthResponse{
  72. User: user,
  73. Token: accessToken,
  74. RefreshToken: refreshToken,
  75. ExpiresAt: expiresAt,
  76. }, nil
  77. }
  78. func normalizeIdentifier(identifier string) string {
  79. identifier = strings.TrimSpace(identifier)
  80. if strings.Contains(identifier, "@") {
  81. return strings.ToLower(identifier)
  82. }
  83. return strings.ReplaceAll(identifier, " ", "")
  84. }
  85. func isEmail(identifier string) bool {
  86. return strings.Contains(identifier, "@")
  87. }
  88. func isDuplicateError(err error) bool {
  89. if err == nil {
  90. return false
  91. }
  92. message := strings.ToLower(err.Error())
  93. return strings.Contains(message, "duplicate key") || strings.Contains(message, "unique constraint")
  94. }
  95. // Register creates the server account and an immediately revocable login session.
  96. func (h *AuthHandler) Register(c *gin.Context) {
  97. var req models.RegisterRequest
  98. if err := c.ShouldBindJSON(&req); err != nil {
  99. respondError(c, http.StatusBadRequest, 400, "invalid input data")
  100. return
  101. }
  102. req.Username = strings.TrimSpace(req.Username)
  103. req.Identifier = normalizeIdentifier(req.Identifier)
  104. hash, err := bcrypt.GenerateFromPassword([]byte(req.Password), bcrypt.DefaultCost)
  105. if err != nil {
  106. respondError(c, http.StatusInternalServerError, 500, "failed to hash password")
  107. return
  108. }
  109. user := models.User{Username: req.Username, PasswordHash: string(hash)}
  110. if isEmail(req.Identifier) {
  111. user.Email = &req.Identifier
  112. } else {
  113. user.PhoneNumber = &req.Identifier
  114. }
  115. var response models.AuthResponse
  116. err = h.DB.Transaction(func(tx *gorm.DB) error {
  117. if err := tx.Create(&user).Error; err != nil {
  118. return err
  119. }
  120. issued, err := h.issueTokens(tx, user)
  121. response = issued
  122. return err
  123. })
  124. if err != nil {
  125. if isDuplicateError(err) {
  126. respondError(c, http.StatusConflict, 409, "username or identifier already exists")
  127. return
  128. }
  129. respondError(c, http.StatusInternalServerError, 500, "failed to create user")
  130. return
  131. }
  132. respondSuccess(c, response)
  133. }
  134. // Login supports username, email, or phone number.
  135. func (h *AuthHandler) Login(c *gin.Context) {
  136. var req models.LoginRequest
  137. if err := c.ShouldBindJSON(&req); err != nil {
  138. respondError(c, http.StatusBadRequest, 400, "invalid input data")
  139. return
  140. }
  141. identifier := normalizeIdentifier(req.Identifier)
  142. var user models.User
  143. if err := h.DB.Where("LOWER(username) = LOWER(?) OR LOWER(email) = LOWER(?) OR phone_number = ?", identifier, identifier, identifier).First(&user).Error; err != nil {
  144. if errors.Is(err, gorm.ErrRecordNotFound) {
  145. respondError(c, http.StatusUnauthorized, 401, "invalid credentials")
  146. return
  147. }
  148. respondError(c, http.StatusInternalServerError, 500, "database error")
  149. return
  150. }
  151. if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.Password)); err != nil {
  152. respondError(c, http.StatusUnauthorized, 401, "invalid credentials")
  153. return
  154. }
  155. response, err := h.issueTokens(h.DB, user)
  156. if err != nil {
  157. respondError(c, http.StatusInternalServerError, 500, "failed to create login session")
  158. return
  159. }
  160. respondSuccess(c, response)
  161. }
  162. // Refresh rotates a refresh token so a stolen token cannot be replayed.
  163. func (h *AuthHandler) Refresh(c *gin.Context) {
  164. var req models.RefreshRequest
  165. if err := c.ShouldBindJSON(&req); err != nil {
  166. respondError(c, http.StatusBadRequest, 400, "refresh token is required")
  167. return
  168. }
  169. var oldSession models.AuthSession
  170. err := h.DB.Where("refresh_hash = ? AND revoked_at IS NULL AND expires_at > ?", hashRefreshToken(req.RefreshToken), time.Now().UTC()).First(&oldSession).Error
  171. if err != nil {
  172. respondError(c, http.StatusUnauthorized, 401, "invalid or expired refresh token")
  173. return
  174. }
  175. var user models.User
  176. if err := h.DB.First(&user, "id = ?", oldSession.UserID).Error; err != nil {
  177. respondError(c, http.StatusUnauthorized, 401, "user no longer exists")
  178. return
  179. }
  180. var response models.AuthResponse
  181. err = h.DB.Transaction(func(tx *gorm.DB) error {
  182. now := time.Now().UTC()
  183. result := tx.Model(&models.AuthSession{}).
  184. Where("id = ? AND revoked_at IS NULL", oldSession.ID).
  185. Updates(map[string]interface{}{"revoked_at": &now, "last_used_at": &now})
  186. if result.Error != nil {
  187. return result.Error
  188. }
  189. if result.RowsAffected != 1 {
  190. return gorm.ErrRecordNotFound
  191. }
  192. issued, err := h.issueTokens(tx, user)
  193. response = issued
  194. return err
  195. })
  196. if err != nil {
  197. respondError(c, http.StatusUnauthorized, 401, "refresh token has already been used")
  198. return
  199. }
  200. respondSuccess(c, response)
  201. }
  202. func (h *AuthHandler) Logout(c *gin.Context) {
  203. sessionID := c.GetString("sessionID")
  204. if sessionID != "" {
  205. now := time.Now().UTC()
  206. h.DB.Model(&models.AuthSession{}).Where("id = ?", sessionID).Update("revoked_at", &now)
  207. }
  208. respondSuccess(c, nil)
  209. }
  210. func (h *AuthHandler) GetProfile(c *gin.Context) {
  211. userID := c.GetString("userID")
  212. var user models.User
  213. if err := h.DB.First(&user, "id = ?", userID).Error; err != nil {
  214. respondError(c, http.StatusNotFound, 404, "user not found")
  215. return
  216. }
  217. respondSuccess(c, user)
  218. }
  219. func (h *AuthHandler) UpdateProfile(c *gin.Context) {
  220. userID := c.GetString("userID")
  221. var req models.UpdateProfileRequest
  222. if err := c.ShouldBindJSON(&req); err != nil {
  223. respondError(c, http.StatusBadRequest, 400, "invalid input data")
  224. return
  225. }
  226. var user models.User
  227. if err := h.DB.First(&user, "id = ?", userID).Error; err != nil {
  228. respondError(c, http.StatusNotFound, 404, "user not found")
  229. return
  230. }
  231. if req.Username != nil {
  232. trimmed := strings.TrimSpace(*req.Username)
  233. if len([]rune(trimmed)) < 2 {
  234. respondError(c, http.StatusBadRequest, 400, "username must contain at least 2 characters")
  235. return
  236. }
  237. user.Username = trimmed
  238. }
  239. if req.Email != nil {
  240. value := strings.ToLower(strings.TrimSpace(*req.Email))
  241. if value == "" {
  242. user.Email = nil
  243. } else {
  244. user.Email = &value
  245. }
  246. }
  247. if req.PhoneNumber != nil {
  248. value := normalizeIdentifier(*req.PhoneNumber)
  249. if value == "" {
  250. user.PhoneNumber = nil
  251. } else {
  252. user.PhoneNumber = &value
  253. }
  254. }
  255. if req.AvatarURL != nil {
  256. value := strings.TrimSpace(*req.AvatarURL)
  257. if value == "" {
  258. user.AvatarURL = nil
  259. } else {
  260. user.AvatarURL = &value
  261. }
  262. }
  263. if err := h.DB.Save(&user).Error; err != nil {
  264. if isDuplicateError(err) {
  265. respondError(c, http.StatusConflict, 409, "username, email, or phone already exists")
  266. return
  267. }
  268. respondError(c, http.StatusInternalServerError, 500, "failed to update profile")
  269. return
  270. }
  271. respondSuccess(c, user)
  272. }
  273. func (h *AuthHandler) ChangePassword(c *gin.Context) {
  274. userID := c.GetString("userID")
  275. var req models.ChangePasswordRequest
  276. if err := c.ShouldBindJSON(&req); err != nil {
  277. respondError(c, http.StatusBadRequest, 400, "invalid input data")
  278. return
  279. }
  280. var user models.User
  281. if err := h.DB.First(&user, "id = ?", userID).Error; err != nil {
  282. respondError(c, http.StatusNotFound, 404, "user not found")
  283. return
  284. }
  285. if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.OldPassword)); err != nil {
  286. respondError(c, http.StatusUnauthorized, 401, "invalid old password")
  287. return
  288. }
  289. hash, err := bcrypt.GenerateFromPassword([]byte(req.NewPassword), bcrypt.DefaultCost)
  290. if err != nil {
  291. respondError(c, http.StatusInternalServerError, 500, "failed to hash new password")
  292. return
  293. }
  294. err = h.DB.Transaction(func(tx *gorm.DB) error {
  295. if err := tx.Model(&user).Update("password_hash", string(hash)).Error; err != nil {
  296. return err
  297. }
  298. now := time.Now().UTC()
  299. return tx.Model(&models.AuthSession{}).
  300. Where("user_id = ? AND id <> ? AND revoked_at IS NULL", userID, c.GetString("sessionID")).
  301. Update("revoked_at", &now).Error
  302. })
  303. if err != nil {
  304. respondError(c, http.StatusInternalServerError, 500, "failed to change password")
  305. return
  306. }
  307. respondSuccess(c, nil)
  308. }