Files
OpenFlare/backend/plugins/domain/user/handlers.go
T
ryan df6aa9ff4d fix(auth): encode snowflake user ids as strings in session and /user-info
Registered users get snowflake ids above JS MAX_SAFE_INTEGER.
/user-info emitted them as JSON numbers and login stored uint64 in
the session. Both now use decimal strings. Tests cover admin vs
non-admin cookie access to /user/self, /user-info, and /upload/my.
2026-09-02 20:22:22 +08:00

493 lines
14 KiB
Go

// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package user
import (
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"context"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"net/http"
"strconv"
"sync"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
var (
authMu sync.RWMutex
authSvc contracts.AuthService
)
// SetAuthService binds the AuthService for cache synchronization.
func SetAuthService(s contracts.AuthService) {
authMu.Lock()
defer authMu.Unlock()
authSvc = s
}
func getAuthService() contracts.AuthService {
authMu.RLock()
defer authMu.RUnlock()
return authSvc
}
func getUserIDFromSession(c *gin.Context) uint64 {
defer func() { _ = recover() }()
session := sessions.Default(c)
val := session.Get(contracts.AuthUserIDKey)
if val == nil {
return 0
}
switch v := val.(type) {
case uint64:
return v
case int64:
if v < 0 {
return 0
}
return uint64(v)
case float64:
if v < 0 {
return 0
}
return uint64(v)
case string:
id, _ := strconv.ParseUint(v, 10, 64)
return id
default:
return 0
}
}
func invalidateUserCache(ctx context.Context, userID uint64) {
if s := getAuthService(); s != nil {
s.InvalidateCachedUser(ctx, userID)
}
}
func invalidateTokenCache(ctx context.Context, tokenHash string) {
if s := getAuthService(); s != nil {
s.InvalidateCachedToken(ctx, tokenHash)
}
}
// Login 用户密码登录
// @Summary 用户密码登录
// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。
// @Tags user
// @Accept json
// @Produce json
// @Param request body user.loginRequest true "登录请求参数"
// @Success 200 {object} response.Any "登录成功,返回用户信息"
// @Failure 400 {object} response.Any "用户名或密码错误"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/user/login [post]
func Login(c *gin.Context) {
var req loginRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidParams)
return
}
user, err := GetUserByUsername(c.Request.Context(), req.Username)
if err != nil {
response.AbortUnauthorized(c, errPasswordMismatch)
return
}
if !user.CheckPassword(req.Password) {
response.AbortUnauthorized(c, errPasswordMismatch)
return
}
sess := sessions.Default(c)
sess.Set(contracts.AuthUserIDKey, strconv.FormatUint(user.ID, 10))
sess.Set(contracts.AuthUserNameKey, user.Username)
needChange := user.NeedChangePassword || user.IsPlaintextPassword()
user.NeedChangePassword = needChange
sess.Set("need_change_password", needChange)
if err := sess.Save(); err != nil {
logger.ErrorF(c.Request.Context(), "save session failed on login: %v", err)
}
c.JSON(http.StatusOK, response.OK(user))
}
// Register 用户注册
// @Summary 用户注册
// @Description 使用用户名和密码注册新账号,注册成功后自动登录并建立 Session。
// @Tags user
// @Accept json
// @Produce json
// @Param request body user.registerRequest true "注册请求参数"
// @Success 200 {object} response.Any "注册并登录成功,返回用户信息"
// @Failure 400 {object} response.Any "参数错误、用户名已存在或注册已关闭"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/user/register [post]
func Register(c *gin.Context) {
var req registerRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidParams)
return
}
newUser := &User{
Username: req.Username,
Email: req.Email,
IsActive: true,
}
if err := newUser.SetEncryptedPassword(req.Password); err != nil {
response.AbortInternal(c, errPasswordEncryptFailed)
return
}
if err := CreateUser(c.Request.Context(), newUser); err != nil {
response.AbortBadRequest(c, errCreateUserFailed+err.Error())
return
}
sess := sessions.Default(c)
sess.Set(contracts.AuthUserIDKey, strconv.FormatUint(newUser.ID, 10))
sess.Set(contracts.AuthUserNameKey, newUser.Username)
if err := sess.Save(); err != nil {
logger.ErrorF(c.Request.Context(), "save session failed on register: %v", err)
}
c.JSON(http.StatusOK, response.OK(newUser))
}
// Logout 用户退出登录
// @Summary 用户退出登录
// @Description 清除用户登录 Session,完成退出
// @Tags user
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=string} "退出成功"
// @Failure 500 {object} response.Any "Session 清除失败"
// @Router /api/v1/user/logout [get]
func Logout(c *gin.Context) {
sess := sessions.Default(c)
sess.Options(sessions.Options{
Path: "/",
MaxAge: -1,
})
sess.Clear()
if err := sess.Save(); err != nil {
logger.ErrorF(c.Request.Context(), "clear session failed on logout: %v", err)
}
c.JSON(http.StatusOK, response.OKNil())
}
// SendEmailCode 发送邮箱验证码
// @Summary 发送邮箱验证码
// @Description 向指定邮箱发送验证码(用于注册场景)
// @Tags user
// @Accept json
// @Produce json
// @Param request body user.sendEmailCodeRequest true "目标邮箱"
// @Success 200 {object} response.Any "发送成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 500 {object} response.Any "发送失败"
// @Router /api/v1/user/send-email-code [post]
func SendEmailCode(c *gin.Context) {
var req sendEmailCodeRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidParams)
return
}
payload, err := json.Marshal(sendEmailCodePayload{Email: req.Email})
if err != nil {
response.AbortInternal(c, errSendEmailFailed)
return
}
ctx := c.Request.Context()
if taskSvc := getTaskService(ctx); taskSvc != nil {
if _, err := taskSvc.Dispatch(ctx, TaskTypeSendEmailCode, payload, contracts.TaskTriggerSystem); err != nil {
logger.ErrorF(ctx, "dispatch send_email_code failed: %v", err)
response.AbortInternal(c, errSendEmailFailed)
return
}
c.JSON(http.StatusOK, response.OK(gin.H{"sent": true}))
return
}
if _, err := (&SendEmailCodeHandler{}).Execute(ctx, payload); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(gin.H{"sent": true}))
}
// ChangePassword 修改用户密码
// @Summary 修改用户密码
// @Description 修改当前登录用户的密码。
// @Tags user
// @Accept json
// @Produce json
// @Param request body user.changePasswordRequest true "修改密码请求参数"
// @Success 200 {object} response.Any{data=string} "修改密码成功"
// @Failure 400 {object} response.Any "原密码错误或新密码不符合要求"
// @Failure 401 {object} response.Any "请先登录"
// @Router /api/v1/user/change-password [post]
func ChangePassword(c *gin.Context) {
var req changePasswordRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidParams)
return
}
userID := getUserIDFromSession(c)
user, err := GetUserByID(c.Request.Context(), userID)
if err != nil {
response.AbortNotFound(c, errUserNotFound)
return
}
if !user.CheckPassword(req.OldPassword) {
response.AbortBadRequest(c, errOldPasswordIncorrect)
return
}
if err := user.SetEncryptedPassword(req.NewPassword); err != nil {
response.AbortInternal(c, errPasswordUpdateFailed)
return
}
if err := UpdateUser(c.Request.Context(), user); err != nil {
logger.ErrorF(c.Request.Context(), "persist changed password failed: %v", err)
}
invalidateUserCache(c.Request.Context(), user.ID)
sess := sessions.Default(c)
sess.Set("need_change_password", false)
if err := sess.Save(); err != nil {
logger.ErrorF(c.Request.Context(), "save session failed on change-password: %v", err)
}
c.JSON(http.StatusOK, response.OKNil())
}
// Self 获取当前登录用户信息
// @Summary 获取当前登录用户信息
// @Description 返回当前登录用户的基本信息,需要登录。
// @Tags user
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "用户信息"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/user/self [get]
func Self(c *gin.Context) {
svc := getAuthService()
if svc == nil {
response.AbortUnauthorized(c, errUserNotFound)
return
}
user, err := svc.GetCurrentUser(c)
if err != nil {
logger.ErrorF(c.Request.Context(), "get current user failed: %v", err)
response.AbortUnauthorized(c, errUserNotFound)
return
}
c.JSON(http.StatusOK, response.OK(user))
}
// UpdateProfile 修改当前登录用户的个人资料
// @Summary 修改当前登录用户的个人资料
// @Description 修改当前登录用户的昵称、头像、简介、电话、性别、个人网站和所在地。
// @Tags user
// @Accept json
// @Produce json
// @Param request body user.updateProfileRequest true "更新请求参数"
// @Success 200 {object} response.Any "修改成功,返回更新后的用户信息"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/user/profile [put]
func UpdateProfile(c *gin.Context) {
var req updateProfileRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidParams)
return
}
userID := getUserIDFromSession(c)
user, err := GetUserByID(c.Request.Context(), userID)
if err != nil {
response.AbortNotFound(c, errUserNotFound)
return
}
user.Nickname = req.Nickname
user.AvatarURL = req.AvatarURL
user.Bio = req.Bio
user.Phone = req.Phone
user.Gender = req.Gender
user.Website = req.Website
user.Location = req.Location
if err := UpdateUser(c.Request.Context(), user); err != nil {
logger.ErrorF(c.Request.Context(), "persist updated profile failed: %v", err)
}
invalidateUserCache(c.Request.Context(), user.ID)
c.JSON(http.StatusOK, response.OK(user))
}
// ListAccessTokens 获取当前用户的 AccessToken 列表
// @Summary 获取当前用户的 AccessToken 列表
// @Description 返回当前登录用户的所有 active access tokens(脱敏后)
// @Tags user
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]user.AccessToken} "令牌列表"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/user/access-tokens [get]
func ListAccessTokens(c *gin.Context) {
userID := getUserIDFromSession(c)
tokens, err := listAccessTokensByUser(c.Request.Context(), userID)
if err != nil {
logger.ErrorF(c.Request.Context(), "list access tokens failed: %v", err)
}
c.JSON(http.StatusOK, response.OK(tokens))
}
const (
tokenEntropyByteLength = 24
tokenMaskMinLength = 8
)
// CreateAccessToken 创建一个新的 AccessToken
// @Summary 创建一个新的 AccessToken
// @Description 为当前用户新建一个 API 访问令牌,仅在此接口返回一次明文令牌值,请妥善保存。
// @Tags user
// @Accept json
// @Produce json
// @Param request body user.createAccessTokenRequest true "令牌名称"
// @Security SessionCookie
// @Success 200 {object} response.Any "新建令牌成功"
// @Failure 400 {object} response.Any "参数错误或超限"
// @Router /api/v1/user/access-tokens [post]
func CreateAccessToken(c *gin.Context) {
var req createAccessTokenRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidParams)
return
}
userID := getUserIDFromSession(c)
rawBytes := make([]byte, tokenEntropyByteLength)
_, _ = rand.Read(rawBytes)
rawToken := "wvt_" + hex.EncodeToString(rawBytes)
hash := sha256.Sum256([]byte(rawToken))
tokenHash := hex.EncodeToString(hash[:])
masked := rawToken
if len(rawToken) > tokenMaskMinLength {
masked = rawToken[:4] + "..." + rawToken[len(rawToken)-4:]
}
token := AccessToken{
UserID: userID,
Name: req.Name,
TokenHash: tokenHash,
MaskedToken: masked,
IsAdmin: req.IsAdmin,
}
if err := createAccessTokenRow(c.Request.Context(), &token); err != nil {
response.AbortInternal(c, errCreateTokenFailed)
return
}
c.JSON(http.StatusOK, response.OK(gin.H{
"token": token,
"raw_token": rawToken,
}))
}
// DeleteAccessToken 删除一个 AccessToken
// @Summary 删除一个 AccessToken
// @Description 撤销并删除一个属于当前用户的 API 访问令牌
// @Tags user
// @Produce json
// @Param id path string true "令牌ID"
// @Security SessionCookie
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Router /api/v1/user/access-tokens/{id} [delete]
func DeleteAccessToken(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, errInvalidParams)
return
}
userID := getUserIDFromSession(c)
token, err := getAccessTokenOfUser(c.Request.Context(), id, userID)
if err != nil {
response.AbortNotFound(c, errTokenNotFound)
return
}
if err := deleteAccessTokenRow(c.Request.Context(), token); err != nil {
logger.ErrorF(c.Request.Context(), "delete access token failed: %v", err)
}
invalidateTokenCache(c.Request.Context(), token.TokenHash)
c.JSON(http.StatusOK, response.OKNil())
}
// RotateAccessToken 轮换一个 AccessToken
// @Summary 轮换一个 AccessToken
// @Description 轮换(重新生成)一个属于当前用户的 API 访问令牌的密钥,旧令牌将立即失效
// @Tags user
// @Produce json
// @Param id path string true "令牌ID"
// @Security SessionCookie
// @Success 200 {object} response.Any "令牌轮换成功"
// @Failure 400 {object} response.Any "参数错误"
// @Router /api/v1/user/access-tokens/{id}/rotate [post]
func RotateAccessToken(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, errInvalidParams)
return
}
userID := getUserIDFromSession(c)
token, err := getAccessTokenOfUser(c.Request.Context(), id, userID)
if err != nil {
response.AbortNotFound(c, errTokenNotFound)
return
}
invalidateTokenCache(c.Request.Context(), token.TokenHash)
rawBytes := make([]byte, tokenEntropyByteLength)
_, _ = rand.Read(rawBytes)
rawToken := "wvt_" + hex.EncodeToString(rawBytes)
hash := sha256.Sum256([]byte(rawToken))
token.TokenHash = hex.EncodeToString(hash[:])
masked := rawToken
if len(rawToken) > tokenMaskMinLength {
masked = rawToken[:4] + "..." + rawToken[len(rawToken)-4:]
}
token.MaskedToken = masked
if err := saveAccessTokenRow(c.Request.Context(), token); err != nil {
logger.ErrorF(c.Request.Context(), "rotate access token failed: %v", err)
}
c.JSON(http.StatusOK, response.OK(gin.H{
"token": token,
"raw_token": rawToken,
}))
}