This commit is contained in:
oneao committed 2026-01-25 22:18:27 +08:00
1 parent 500fab7a80
commit 8f3f6982db
54 files changed
+2076 -1877

No files matched your search

@@ -3,65 +3,289 @@
package auth
import (
"allapp/biz/model/auth"
"allapp/conf"
"allapp/db"
"allapp/db/repo"
"allapp/logx"
"allapp/utils"
"allapp/utils/errorx"
"allapp/utils/httpx"
"allapp/utils/idgen"
"allapp/utils/jwtx"
"allapp/utils/response"
"context"
"allapp/biz/model/auth"
"fmt"
"github.com/cloudwego/hertz/pkg/app"
"github.com/cloudwego/hertz/pkg/protocol/consts"
)
// LoginDefault .
// @router /auth/login/default [POST]
func LoginDefault(ctx context.Context, c *app.RequestContext) {
var err error
var req auth.LoginDefaultReq
err = c.BindAndValidate(&req)
if err != nil {
// LoginQq QQ 登录
// @router /auth/login/qq [POST]
func LoginQq(ctx context.Context, c *app.RequestContext) {
var req auth.LoginQqReq
if err := c.BindAndValidate(&req); err != nil {
c.String(consts.StatusBadRequest, err.Error())
return
}
// 登录
account := req.Account
password := req.Password
resp := &auth.LoginResp{}
openid := req.Openid
nickname := req.Username
avatar := req.Avatar
params := repo.GetUserForLoginParams{
Account: account,
Password: password,
}
loginUser, err := db.Queries.GetUserForLogin(ctx, params)
// ---------- 通用逻辑:查询用户 ----------
user, err := db.Queries.FindUserByOpenID(ctx, repo.FindUserByOpenIDParams{
Openid: openid,
Type: 1,
})
isExist := true
if err != nil {
response.Fail(c).Message("账号或密码错误!").Send()
if errorx.IsNotFound(err) {
isExist = false
} else {
errorx.AddError(c, err)
return
}
}
if isExist {
if user.Status == 0 {
response.Fail(c).Message("账号已被禁用").Send()
return
}
_ = db.Queries.UpdateUserLastLoginTime(ctx, user.UserID)
token, _ := jwtx.CreateToken(user.UserID)
resp.Token = token
resp.UserId = user.UserID
resp.Username = user.Username
resp.Avatar = user.Avatar
response.Success(c).Data(resp).Send()
return
}
// 异步更新登录事件
go func(userID int64) {
err = db.Queries.UpdateUserLoginLastTime(ctx, userID)
if err != nil {
logx.CtxError(ctx, "更新用户登录时间失败,UserID=%d, err=%v", loginUser.ID, err)
// ---------- 用户不存在 → 注册 ----------
registerUser(ctx, c, openid, nickname, avatar, 1)
}
// LoginWx 微信登录
// @router /auth/login/wx [POST]
func LoginWx(ctx context.Context, c *app.RequestContext) {
var req auth.LoginWxReq
if err := c.BindAndValidate(&req); err != nil {
c.String(consts.StatusBadRequest, err.Error())
return
}
resp := &auth.LoginResp{}
openid := ""
nickname := ""
avatar := ""
accessToken := ""
// 微信登录:code换access_token + openid
status, body, err := httpx.Get(ctx, "https://api.weixin.qq.com/sns/oauth2/access_token", map[string][]any{
"appid": {conf.GetConf().Wx.AppId},
"secret": {conf.GetConf().Wx.AppSecret},
"code": {req.Code},
"grant_type": {"authorization_code"},
})
if err != nil || status != 200 {
errorx.AddError(c, fmt.Errorf("获取 openid 失败: %w, body: %s", err, body))
return
}
accessMap, err := httpx.JSONStringToMap(body)
if err != nil {
errorx.AddError(c, fmt.Errorf("解析 access_token 失败: %w", err))
return
}
accessTokenVal, ok1 := accessMap["access_token"].(string)
openidVal, ok2 := accessMap["openid"].(string)
if !ok1 || !ok2 || accessTokenVal == "" || openidVal == "" {
errorx.AddError(c, fmt.Errorf("access_token 或 openid 不存在, body: %s", body))
return
}
openid = openidVal
accessToken = accessTokenVal
// ---------- 通用逻辑:查询用户 ----------
user, err := db.Queries.FindUserByOpenID(ctx, repo.FindUserByOpenIDParams{
Openid: openid,
Type: 0,
})
// 是否已注册
isExist := true
if err != nil {
if errorx.IsNotFound(err) {
isExist = false
} else {
errorx.AddError(c, err)
return
}
}(loginUser.ID)
}
// 包装数据
token, err := jwtx.CreateAccessToken(loginUser.ID)
refreshToken, err := jwtx.CreateRefreshToken(loginUser.ID)
if isExist {
if user.Status == 0 {
response.Fail(c).Message("账号已被禁用").Send()
return
}
resp := new(auth.LoginDefaultResp)
_ = db.Queries.UpdateUserLastLoginTime(ctx, user.UserID)
token, _ := jwtx.CreateToken(user.UserID)
resp.Token = token
resp.RefreshToken = refreshToken
resp.UserId = loginUser.ID
resp.Gender = loginUser.Gender.Int32
resp.SeqNo = loginUser.SeqNo.Int32
resp.Username = loginUser.Username.String
resp.Token = token
resp.UserId = user.UserID
resp.Username = user.Username
resp.Avatar = user.Avatar
response.Success(c).Data(resp).Send()
return
}
// ---------- 用户不存在 → 注册 ----------
// 用 access_token 获取用户信息
status, body, err = httpx.Get(ctx, "https://api.weixin.qq.com/sns/userinfo", map[string][]any{
"access_token": {accessToken},
"openid": {openid},
})
if err != nil || status != 200 {
errorx.AddError(c, fmt.Errorf("获取用户信息失败: %w, body: %s", err, body))
return
}
userMap, err := httpx.JSONStringToMap(body)
if err != nil {
errorx.AddError(c, fmt.Errorf("解析用户信息失败: %w", err))
return
}
nickname, _ = userMap["nickname"].(string)
avatar, _ = userMap["headimgurl"].(string)
registerUser(ctx, c, openid, nickname, avatar, 0) // 0 = 微信
}
// registerUser 注册通用函数
func registerUser(ctx context.Context, c *app.RequestContext, openid, nickname, avatar string, loginType int32) {
userId := idgen.NextId()
spaceId := idgen.NextId()
if nickname == "" {
nickname = fmt.Sprintf("用户_%d", userId)
}
if avatar == "" {
avatar = "https://default-avatar-url.com/default.png"
}
err := db.WithTx(ctx, func(q *repo.Queries) error {
// 生成唯一邀请码
var inviteCode string
for {
code := utils.GenerateCode(6)
_, err := q.FindSpaceByInviteCode(ctx, code)
if errorx.IsNotFound(err) {
inviteCode = code
break
} else if err != nil {
return err
}
}
// 插入空间
if err := q.InsertSpace(ctx, repo.InsertSpaceParams{
ID: spaceId,
InviteCode: inviteCode,
}); err != nil {
return err
}
// 插入空间用户
if err := q.InsertSpaceUser(ctx, repo.InsertSpaceUserParams{
ID: idgen.NextId(),
SpaceID: spaceId,
UserID: userId,
}); err != nil {
return err
}
// 插入用户
if err := q.InsertUser(ctx, repo.InsertUserParams{
ID: userId,
Username: nickname,
Avatar: avatar,
Status: 1,
}); err != nil {
return err
}
// 插入第三方登录信息
if err := q.InsertUserOAuth(ctx, repo.InsertUserOAuthParams{
ID: idgen.NextId(),
UserID: userId,
Type: loginType,
Openid: openid,
}); err != nil {
return err
}
// 初始化记账分类
if err := initMoneyCategoryTx(ctx, q, userId); err != nil {
return fmt.Errorf("初始化记账分类失败: %w", err)
}
return nil
})
if err != nil {
errorx.AddError(c, err)
return
}
token, _ := jwtx.CreateToken(userId)
resp := &auth.LoginResp{
Token: token,
UserId: userId,
Username: nickname,
Avatar: avatar,
}
response.Success(c).Data(resp).Send()
}
func initMoneyCategoryTx(ctx context.Context, q *repo.Queries, userId int64) error {
sysCategories, err := q.ListMoneySysCategory(ctx)
if err != nil {
return err
}
if len(sysCategories) == 0 {
return nil
}
userCategories := make([]repo.BatchInsertMoneyUserCategoriesParams, 0, len(sysCategories))
for _, sysCat := range sysCategories {
userCategories = append(userCategories, repo.BatchInsertMoneyUserCategoriesParams{
ID: idgen.NextId(),
UserID: userId,
Name: sysCat.Name,
Icon: sysCat.Icon,
Type: sysCat.Type,
SortNumber: sysCat.SortNumber,
})
}
if err := q.DeleteMoneyCategoryByUserId(ctx, userId); err != nil {
return err
}
if _, err := q.BatchInsertMoneyUserCategories(ctx, userCategories); err != nil {
return err
}
return nil
}
@@ -1,138 +0,0 @@
// Code generated by hertz generator.
package auth
import (
auth "allapp/biz/model/auth"
"allapp/db"
"allapp/db/repo"
"allapp/utils"
"allapp/utils/errorx"
"allapp/utils/idgen"
"allapp/utils/pgtypex"
"allapp/utils/response"
"context"
"github.com/cloudwego/hertz/pkg/app"
"github.com/cloudwego/hertz/pkg/protocol/consts"
"github.com/jackc/pgx/v5"
)
// RegisterDefault .
// @router /auth/register/default [POST]
func RegisterDefault(ctx context.Context, c *app.RequestContext) {
var err error
var req auth.RegisterDefaultReq
err = c.BindAndValidate(&req)
if err != nil {
c.String(consts.StatusBadRequest, err.Error())
return
}
account := req.Account
password := req.Password
_, err = db.Queries.GetUserByAccount(ctx, account)
if err != nil {
if errorx.IsNotFound(err) {
// 忽略
} else {
errorx.AddError(c, err)
return
}
} else {
response.Fail(c).Message("该账号已被注册").Send()
return
}
// 开启事务
tx, err := db.DB.BeginTx(ctx, pgx.TxOptions{})
if err != nil {
errorx.AddError(c, err)
return
}
defer func() {
if err != nil {
err := tx.Rollback(ctx)
if err != nil {
return
}
} else {
err := tx.Commit(ctx)
if err != nil {
return
}
}
}()
generateCode, err := utils.GenerateCode(8)
if err != nil {
errorx.AddError(c, err)
return
}
// 插入用户
params := repo.InsertUserParams{
ID: idgen.NextId(),
Account: account,
Password: password,
Username: pgtypex.StringToText("a_" + generateCode),
}
insert, err := db.Queries.WithTx(tx).InsertUser(ctx, params)
if err != nil || insert == 0 {
errorx.AddError(c, err)
return
}
// 初始化用户分类
err = InitMoneyCategoryTx(ctx, tx, params.ID)
if err != nil {
errorx.AddError(c, err)
return
}
// 注册成功
response.Success(c).Message("注册成功").Send()
}
func InitMoneyCategoryTx(
ctx context.Context,
tx pgx.Tx,
userId int64,
) error {
q := db.Queries.WithTx(tx)
sysCategories, err := q.ListMoneySysCategory(ctx)
if err != nil {
return err
}
if len(sysCategories) == 0 {
return nil
}
userCategories := make([]repo.BatchInsertMoneyUserCategoriesParams, 0, len(sysCategories))
for _, sysCat := range sysCategories {
userCategories = append(userCategories, repo.BatchInsertMoneyUserCategoriesParams{
ID: idgen.NextId(),
UserID: userId,
Name: sysCat.Name,
Icon: sysCat.Icon,
Type: sysCat.Type,
SortNumber: sysCat.SortNumber,
})
}
// ⚠️ 注意:这里是 DeleteMoneyCategoryByUserId,不是 ById(你代码里疑似写错)
if err := q.DeleteMoneyCategoryByUserId(ctx, userId); err != nil {
return err
}
if _, err := q.BatchInsertMoneyUserCategories(ctx, userCategories); err != nil {
return err
}
return nil
}
@@ -1,60 +0,0 @@
// Code generated by hertz generator.
package auth
import (
"allapp/utils/errorx"
"allapp/utils/jwtx"
"allapp/utils/response"
"context"
auth "allapp/biz/model/auth"
"github.com/cloudwego/hertz/pkg/app"
"github.com/cloudwego/hertz/pkg/protocol/consts"
)
// RefreshAccessToken .
// @router /auth/refresh/token [GET]
func RefreshAccessToken(ctx context.Context, c *app.RequestContext) {
var err error
var req auth.RefreshAccessTokenReq
err = c.BindAndValidate(&req)
if err != nil {
c.String(consts.StatusBadRequest, err.Error())
return
}
refreshToken := string(c.GetHeader("refreshToken"))
if refreshToken == "" {
response.Fail(c).
Code(response.HttpCode.Unauthorized).
Send()
return
}
tokenResult := jwtx.VerifyToken(refreshToken)
if !tokenResult.IsValid || tokenResult.IsExpired {
// Refresh Token 无效或过期,需要重新登录
response.Fail(c).
Code(response.HttpCode.Unauthorized).
Send()
return
}
userID := tokenResult.Claims.UserID
// 生成新的 Access Token
newAccessToken, err := jwtx.CreateAccessToken(userID)
if err != nil {
errorx.AddError(c, err)
return
}
resp := new(auth.RefreshAccessTokenResp)
resp.AccessToken = newAccessToken
response.Success(c).
Data(resp).
Send()
}
@@ -6,9 +6,7 @@ import (
"allapp/biz/router/middleware"
"allapp/db"
"allapp/db/repo"
"allapp/utils"
"allapp/utils/errorx"
"allapp/utils/idgen"
"allapp/utils/response"
"context"
@@ -18,105 +16,6 @@ import (
"github.com/cloudwego/hertz/pkg/protocol/consts"
)
// CreateSpace .
// @router /space/create [POST]
func CreateSpace(ctx context.Context, c *app.RequestContext) {
var err error
var req space.CreateSpaceParams
err = c.BindAndValidate(&req)
if err != nil {
c.String(consts.StatusBadRequest, err.Error())
return
}
userID := middleware.GetUserID(ctx)
// 检查该用户是否已经加入空间了
_, err = db.Queries.GetSpaceByUserId(ctx, userID)
if err != nil {
if errorx.IsNotFound(err) {
// 忽略
} else {
errorx.AddError(c, err)
return
}
} else {
response.Fail(c).Message("请先退出已加入的空间").Send()
return
}
// 生成邀请码
var inviteCode string
for {
// 生成邀请码
code, err := utils.GenerateCode(6)
if err != nil {
errorx.AddError(c, err)
return
}
// 检查数据库是否存在
_, err = db.Queries.GetSpaceByInviteCode(ctx, code)
if err != nil {
if errorx.IsNotFound(err) {
// 不存在,唯一,直接使用
inviteCode = code
break
} else {
// 系统错误
errorx.AddError(c, err)
return
}
}
// 如果存在,继续循环重新生成
}
spaceId := idgen.NextId()
err = db.WithTx(ctx, func(q *repo.Queries) error {
// 创建空间
params := repo.InsertSpaceParams{
ID: spaceId,
InviteCode: inviteCode,
}
err = q.InsertSpace(ctx, params)
if err != nil {
return err
}
// 将创建用户默认添加到该空间内
userParams := repo.InsertSpaceUserParams{
ID: idgen.NextId(),
SpaceID: spaceId,
UserID: userID,
Role: 0,
}
err := q.InsertSpaceUser(ctx, userParams)
if err != nil {
return err
}
return nil
})
if err != nil {
errorx.AddError(c, err)
return
}
resp := &space.CreateSpaceResult{
Id: spaceId,
InviteCode: inviteCode,
}
response.Success(c).Data(resp).Send()
}
// DissolveSpace .
// @router /space/dissolve [POST]
func DissolveSpace(ctx context.Context, c *app.RequestContext) {
@@ -177,7 +177,7 @@ func JoinSpace(ctx context.Context, c *app.RequestContext) {
}
// 检查验证码是否有效
spaceByCode, err := db.Queries.GetSpaceByInviteCode(ctx, inviteCode)
spaceByCode, err := db.Queries.FindSpaceByInviteCode(ctx, inviteCode)
if err != nil {
if errorx.IsNotFound(err) {