Files
workspace/code/allapp/allapp-go-v3/internal/handle/auth.go
T
2026-04-19 21:30:51 +08:00

217 lines
4.7 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package handle
import (
"allapp-go/internal/errors"
"allapp-go/internal/httpx"
"allapp-go/internal/types"
"allapp-go/pkg/db"
"allapp-go/pkg/jwtx"
"allapp-go/pkg/uniqueid"
"allapp-go/pkg/wechat"
"fmt"
"github.com/gofiber/fiber/v3"
"github.com/gofiber/fiber/v3/log"
)
// LoginQq ======================== QQ 登录 ========================
func LoginQq(c fiber.Ctx) error {
var req types.LoginQqDTO
vo := new(types.LoginVO)
if err := httpx.BindAndValidate(c, &req); err != nil {
return errors.WithStack(err)
}
return handleThirdLogin(c, req.Openid, req.Nickname, req.Avatar, 1, vo)
}
// LoginWechat ======================== 微信登录 ========================
func LoginWechat(c fiber.Ctx) error {
var req types.LoginWechatDTO
vo := new(types.LoginVO)
if err := httpx.BindAndValidate(c, &req); err != nil {
return errors.WithStack(err)
}
openid, token, err := wechat.GetWechatAccess(req.Code)
if err != nil {
log.Errorw("wechat access failed", "code", req.Code, "error", err)
return httpx.Fail(c, "微信登录失败,请重试")
}
nickname, avatar, err := wechat.GetWechatUserInfo(token, openid)
if err != nil {
log.Errorw("wechat userinfo failed", "openid", openid, "error", err)
return httpx.Fail(c, "微信登录失败,请重试")
}
return handleThirdLogin(c, openid, nickname, avatar, 0, vo)
}
// ======================== 第三方登录主流程 ========================
func handleThirdLogin(
c fiber.Ctx,
openid string,
nickname string,
avatar string,
loginType int16,
vo *types.LoginVO,
) error {
dbClient := db.New()
user, err := getUserByOpenID(dbClient, c, openid, loginType)
if err != nil {
return errors.WithStack(err)
}
// 不存在 -> 注册
if user == nil {
return registerUser(dbClient, c, openid, nickname, avatar, loginType, vo)
}
// ========= 校验状态 =========
id := toInt64(user["id"])
status := toInt64(user["status"])
if status == 0 {
return httpx.Fail(c, "账号已被禁用")
}
// ========= 更新登录时间 =========
if err := dbClient.Update(
c.Context(),
"user",
"id",
map[string]any{
"id": id,
"last_login_time": "NOW()",
},
); err != nil {
return errors.WithStack(err)
}
// ========= token =========
token, err := jwtx.CreateToken(c.Context(), jwtx.TokenData{
UserID: id,
})
if err != nil {
return errors.WithStack(err)
}
// ========= 返回 =========
vo.Token = token
vo.UserId = id
vo.Nickname, _ = user["nickname"].(string)
vo.Avatar, _ = user["avatar"].(string)
return httpx.OK(c, vo)
}
// ======================== 注册 ========================
func registerUser(
dbClient *db.Client,
c fiber.Ctx,
openid, nickname, avatar string,
loginType int16,
vo *types.LoginVO,
) error {
userId := uniqueid.NextId()
if nickname == "" {
nickname = fmt.Sprintf("用户_%d", userId)
}
if avatar == "" {
avatar = "https://default-avatar-url.com/default.png"
}
err := dbClient.WithTx(c.Context(), func(tx *db.Client) error {
if err := tx.Insert(c.Context(), "b_user", "id", map[string]any{
"id": userId,
"nickname": nickname,
"avatar": avatar,
"created_at": "NOW()",
"updated_at": "NOW()",
"last_login_time": "NOW()",
}); err != nil {
return err
}
if err := tx.Insert(c.Context(), "b_user_oauth", "id", map[string]any{
"id": uniqueid.NextId(),
"user_id": userId,
"type": loginType,
"openid": openid,
"created_at": "NOW()",
"updated_at": "NOW()",
}); err != nil {
return err
}
return nil
})
if err != nil {
return errors.WithStack(err)
}
token, err := jwtx.CreateToken(c.Context(), jwtx.TokenData{
UserID: userId,
})
if err != nil {
return errors.WithStack(err)
}
vo.Token = token
vo.UserId = userId
vo.Nickname = nickname
vo.Avatar = avatar
return httpx.OK(c, vo)
}
// ======================== DB 查询封装(去重复 SQL) ========================
func getUserByOpenID(dbClient *db.Client, c fiber.Ctx, openid string, loginType int16) (map[string]any, error) {
users, err := dbClient.LoadDataBySQL(
c.Context(),
`SELECT u.*
FROM b_user_oauth o
JOIN b_user u ON u.id = o.user_id
WHERE o.openid = $1 AND o.type = $2`,
[]any{openid, loginType},
)
if err != nil {
return nil, err
}
if len(users) == 0 {
return nil, nil
}
return users[0], nil
}
// ======================== 类型工具(极简版) ========================
func toInt64(v any) int64 {
switch val := v.(type) {
case int64:
return val
case int32:
return int64(val)
case int16:
return int64(val)
case int:
return int64(val)
case float64:
return int64(val)
default:
fmt.Printf("unknown type: %T, value=%v\n", v, v)
return 0
}
}