This commit is contained in:
oneao committed 2025-11-06 17:32:51 +08:00
1 parent 6849579652
commit d60be4a75b
24 files changed
+463 -179

No files matched your search

@@ -13,18 +13,48 @@ func InitAuthRouter(r *router.Router) {
group := r.Group("auth")
group.POST("/login/default", loginDefault)
group.POST("/register/default", registerDefault)
group.POST("/refresh/token", RefreshAccessToken)
}
//func RefreshToken(c *router.Context) error {
// _, err := c.GetBodyWithRequired("accessToken", "refreshToken")
// if err != nil {
// return err
// }
// //accessToken := bodyData.GetString("accessToken")
// //refreshToken := bodyData.GetString("refreshToken")
//}
func RefreshAccessToken(c *router.Context) error {
// 从 Header 获取 Refresh Token
refreshToken := c.Header("refreshToken")
if refreshToken == "" {
return response.Fail(c).
Code(response.HttpCode.Unauthorized).
Message("未提供 Token").
Send()
}
// 验证 Refresh Token
tokenResult := jwtx.VerifyToken(refreshToken)
if !tokenResult.IsValid || tokenResult.IsExpired {
// Refresh Token 无效或过期,需要重新登录
return response.Fail(c).
Code(response.HttpCode.Unauthorized).
Message("登录已过期,请重新登录").
Send()
}
userID := tokenResult.Claims.UserID
// 生成新的 Access Token
newAccessToken, err := jwtx.CreateAccessToken(userID)
if err != nil {
return err
}
// 返回新的 Access Token
respData := map[string]string{
"accessToken": newAccessToken,
}
return response.Success(c).
Data(respData).
Send()
}
// 默认登录
// 默认登录
func loginDefault(c *router.Context) error {
// 获取请求体并校验必填字段
@@ -43,13 +73,12 @@ func loginDefault(c *router.Context) error {
)
user, err := db.FindOne("user_info", userQuery)
if err != nil {
return response.Success(c).Message("查询用户失败,请稍后重试").IsSuccess(false).Send()
return response.Success(c).Message("查询用户失败,请稍后重试").Send()
}
if user == nil {
// 账号或密码错误
return response.Success(c).
IsSuccess(false).
return response.Fail(c).
Message("账号或密码错误").
Send()
}
@@ -99,7 +128,7 @@ func registerDefault(c *router.Context) error {
}
if len(existingUsers) > 0 {
return response.Error(c).IsSuccess(false).Message("该账号已被注册").Send()
return response.Fail(c).Message("该账号已被注册").Send()
}
// 构造新用户数据
@@ -115,7 +144,7 @@ func registerDefault(c *router.Context) error {
return err
}
if insertedRows != 1 {
return response.Success(c).IsSuccess(false).Message("注册失败").Send()
return response.Fail(c).Message("注册失败").Send()
}
return response.Success(c).Message("注册成功").Send()
@@ -0,0 +1,39 @@
package middleware
import (
response "base-go-v2/internal/reponse"
"base-go-v2/internal/router"
"base-go-v2/internal/utils/jwtx"
)
func AuthMiddleware() router.HandlerFunc {
return func(c *router.Context) error {
token := c.Header("token")
if token == "" {
c.Abort()
return response.Fail(c).
Code(response.HttpCode.Unauthorized).
Message("未提供 Token").
Send()
}
tokenVerify := jwtx.VerifyToken(token)
// Access Token 有效
if tokenVerify.IsValid {
c.Set("user_id", tokenVerify.Claims.UserID)
return nil
}
// Access Token 失效 → 返回刷新提示
if !tokenVerify.IsValid {
c.Abort()
return response.Fail(c).
Code(response.HttpCode.RefreshToken).
Message("Access Token 已过期,请刷新").
Send()
}
return c.Next()
}
}
@@ -28,7 +28,7 @@ func ErrorMiddleware() router.HandlerFunc {
// be 是 nil 指针类型
if errors.As(err, &be) {
// 业务错误,不记录日志
return response.Error(c).Message(err.Error()).Send()
return response.Fail(c).Message(err.Error()).Send()
}
}
@@ -40,6 +40,6 @@ func ErrorMiddleware() router.HandlerFunc {
logx.String("error", err.Error()),
logx.String("stack", stack),
)
return response.Error(c).Message("服务器内部错误").Send()
return response.Fail(c).Message("服务器内部错误").Send()
}
}
@@ -21,7 +21,7 @@ func RecoveryMiddleware() router.HandlerFunc {
)
// 返回统一 JSON 响应
_ = response.Error(c).Message("系统 Panic").Send()
_ = response.Fail(c).Message("系统 Panic").Send()
}
}()
@@ -14,6 +14,7 @@ func InitAppRouter() *router.Router {
r.Use(middleware.CORSMiddleware())
r.Use(middleware.LoggerMiddleware())
r.Use(middleware.ErrorMiddleware())
r.Use(middleware.AuthMiddleware()).ExcludePaths("/auth/**")
auth.InitAuthRouter(r)
return r
@@ -8,128 +8,71 @@ import (
"net/http"
)
// Result 是最终响应结构体
type Result struct {
IsSuccess bool `json:"isSuccess"` // 是否成功
Code string `json:"code"` // 状态码
Message string `json:"message"` // 提示信息
Data interface{} `json:"data,omitempty"` // 数据
TrackId string `json:"trackId,omitempty"` // 请求追踪ID
Code int `json:"code"`
Message string `json:"message"`
Data interface{} `json:"data,omitempty"`
TrackId string `json:"trackId,omitempty"`
}
// Code 定义状态码
type Code struct {
Value string
Message string
var HttpCode = struct {
Success int
Unauthorized int
Fail int
RefreshToken int
}{
Success: 200,
Unauthorized: 401,
Fail: 500,
RefreshToken: 402,
}
// 常用状态码
var (
CodeSuccess = Code{"0000", "请求成功"}
CodeFail = Code{"0001", "请求失败"}
CodeNoLogin = Code{"1001", "未登录"}
CodeLoginExpired = Code{"1002", "登录已过期"}
CodeInvalidToken = Code{"1003", "Token 无效"}
CodeInvalidAccount = Code{"1005", "账号或密码错误"}
)
const defaultMessage = "未知错误"
// Builder 用于构建响应
type Builder struct {
c *router.Context
result Result
}
// ---------------- 构造函数 ----------------
// Success 创建成功响应
func Success(c *router.Context) *Builder {
return &Builder{
c: c,
result: Result{
IsSuccess: true,
Code: CodeSuccess.Value,
Message: CodeSuccess.Message,
Code: HttpCode.Success,
Message: "请求成功",
},
}
}
// Error 创建失败响应
func Error(c *router.Context) *Builder {
func Fail(c *router.Context) *Builder {
return &Builder{
c: c,
result: Result{
IsSuccess: false,
Code: CodeFail.Value,
Message: CodeFail.Message,
Code: HttpCode.Fail,
Message: "请求失败",
},
}
}
// ---------------- 链式方法 ----------------
// 设置 IsSuccess 字段
func (b *Builder) IsSuccess(IsSuccess bool) *Builder {
b.result.IsSuccess = IsSuccess
return b
}
// 设置 Code,并根据 Code 自动设置 IsSuccess
func (b *Builder) Code(code string) *Builder {
func (b *Builder) Code(code int) *Builder {
b.result.Code = code
b.result.Message = GetMessage(code)
b.result.IsSuccess = code == CodeSuccess.Value
return b
}
// 设置 Code 对象,并自动设置 IsSuccess
func (b *Builder) CodeObj(code Code) *Builder {
b.result.Code = code.Value
b.result.Message = code.Message
b.result.IsSuccess = code.Value == CodeSuccess.Value
return b
}
// 自定义提示信息
func (b *Builder) Message(msg string) *Builder {
b.result.Message = msg
return b
}
// 设置返回数据
func (b *Builder) Data(data interface{}) *Builder {
b.result.Data = data
return b
}
// 发送响应
func (b *Builder) Send() error {
// 默认 Data 为空
if b.result.Data == nil {
b.result.Data = ""
}
// 从 routinex 获取 trackId
trackId := routinex.Get(logx.TrackID)
if trackId != nil {
if trackId := routinex.Get(logx.TrackID); trackId != nil {
b.result.TrackId = strutil.ToString(trackId)
}
return b.c.JSON(http.StatusOK, b.result)
}
// ---------------- 辅助函数 ----------------
// 根据 Code 获取默认提示信息
func GetMessage(value string) string {
for _, c := range []Code{
CodeSuccess, CodeFail, CodeNoLogin,
CodeLoginExpired, CodeInvalidToken, CodeInvalidAccount,
} {
if c.Value == value {
return c.Message
}
}
return defaultMessage
}
@@ -67,8 +67,24 @@ func (n *node) search(parts []string, height int, params map[string]string) *nod
// 中间件条目
type middlewareEntry struct {
handler HandlerFunc
excludePaths map[string]struct{}
handler HandlerFunc
excludePaths map[string]struct{}
excludePrefix []string
}
// 判断路径是否被排除
func (m *middlewareEntry) IsExcluded(path string) bool {
if m.excludePaths != nil {
if _, ok := m.excludePaths[path]; ok {
return true
}
}
for _, prefix := range m.excludePrefix {
if strings.HasPrefix(path, prefix) {
return true
}
}
return false
}
// Router 主结构
@@ -78,7 +94,7 @@ type Router struct {
basePath string
}
// NewRouter 创建
// 创建 Router
func NewRouter(basePath ...string) *Router {
path := "/"
if len(basePath) > 0 && basePath[0] != "" {
@@ -90,11 +106,13 @@ func NewRouter(basePath ...string) *Router {
}
}
// MiddlewareHandle 用于链式排除路径
type MiddlewareHandle struct {
router *Router
entryIdx int
}
// Use 注册中间件
func (r *Router) Use(handler HandlerFunc) *MiddlewareHandle {
entry := middlewareEntry{handler: handler}
r.middleware = append(r.middleware, entry)
@@ -104,15 +122,27 @@ func (r *Router) Use(handler HandlerFunc) *MiddlewareHandle {
}
}
// ExcludePaths 支持 /** 前缀匹配
func (mh *MiddlewareHandle) ExcludePaths(paths ...string) *Router {
excludeMap := make(map[string]struct{}, len(paths))
excludeMap := make(map[string]struct{}, 0)
excludePrefix := make([]string, 0)
for _, p := range paths {
excludeMap[p] = struct{}{}
if strings.HasSuffix(p, "/**") {
base := strings.TrimSuffix(p, "/**")
excludePrefix = append(excludePrefix, base)
} else {
excludeMap[p] = struct{}{}
}
}
mh.router.middleware[mh.entryIdx].excludePaths = excludeMap
mh.router.middleware[mh.entryIdx].excludePrefix = excludePrefix
return mh.router
}
// Group 创建子路由组
func (r *Router) Group(prefix string, m ...HandlerFunc) *Router {
newMiddleware := make([]middlewareEntry, len(r.middleware))
copy(newMiddleware, r.middleware)
@@ -127,6 +157,7 @@ func (r *Router) Group(prefix string, m ...HandlerFunc) *Router {
}
}
// GET/POST
func (r *Router) GET(path string, handler HandlerFunc) *Router {
return r.handle("GET", path, handler)
}
@@ -145,7 +176,7 @@ func (r *Router) handle(method, path string, handler HandlerFunc) *Router {
return r
}
// Handler 入口函数(fasthttp)
// Handler 入口(fasthttp)
func (r *Router) Handler(ctx *fasthttp.RequestCtx) {
method := string(ctx.Method())
root := r.roots[method]
@@ -166,10 +197,8 @@ func (r *Router) Handler(ctx *fasthttp.RequestCtx) {
handlers := make([]HandlerFunc, 0)
path := string(ctx.Path())
for _, m := range r.middleware {
if m.excludePaths != nil {
if _, excluded := m.excludePaths[path]; excluded {
continue
}
if m.IsExcluded(path) {
continue
}
handlers = append(handlers, m.handler)
}
@@ -185,6 +214,7 @@ func (r *Router) Handler(ctx *fasthttp.RequestCtx) {
_ = c.Next()
}
// 辅助函数
func parsePattern(pattern string) []string {
vs := strings.Split(strings.Trim(pattern, "/"), "/")
parts := make([]string, 0, len(vs))
@@ -15,9 +15,9 @@ type CustomClaims struct {
// TokenResult 用于返回验证结果
type TokenResult struct {
Claims *CustomClaims
Valid bool
Expired bool
Claims *CustomClaims
IsValid bool
IsExpired bool
}
// CreateAccessToken 创建 Access Token,只返回 token 和 error
@@ -61,7 +61,7 @@ func createToken(userID int, expiry time.Duration) (string, error) {
// VerifyToken 验证 token 并返回结构化结果
func VerifyToken(tokenString string) TokenResult {
result := TokenResult{Valid: false}
result := TokenResult{IsValid: false, IsExpired: false}
token, _ := jwt.ParseWithClaims(tokenString, &CustomClaims{}, func(token *jwt.Token) (interface{}, error) {
return []byte(config.App.JWT.Secret), nil
@@ -73,9 +73,9 @@ func VerifyToken(tokenString string) TokenResult {
}
result.Claims = claims
result.Valid = true
result.IsValid = true
if claims.ExpiresAt != nil && time.Now().After(claims.ExpiresAt.Time) {
result.Expired = true
result.IsExpired = true
}
return result
@@ -36,25 +36,25 @@ func NotEmpty(data map[string]interface{}, fields ...string) error {
// collectErrors 递归收集字段错误
func collectErrors(v interface{}, path string, errs map[string]string) {
if v == nil {
errs[path] = fmt.Sprintf("字段 %s 不能为空", path)
errs[path] = "字段不能为空"
return
}
switch val := v.(type) {
case string:
if strings.TrimSpace(val) == "" {
errs[path] = fmt.Sprintf("字段 %s 不能为空", path)
errs[path] = "字段不能为空"
}
case []interface{}:
if len(val) == 0 {
errs[path] = fmt.Sprintf("字段 %s 不能为空", path)
errs[path] = "字段不能为空"
}
for i, elem := range val {
collectErrors(elem, fmt.Sprintf("%s[%d]", path, i), errs)
}
case map[string]interface{}:
if len(val) == 0 {
errs[path] = fmt.Sprintf("字段 %s 不能为空", path)
errs[path] = "字段不能为空"
}
for k, elem := range val {
collectErrors(elem, fmt.Sprintf("%s.%s", path, k), errs)