u
This commit is contained in:
1 parent
9af3234672
commit
ed8bee3d77
11 files changed
+171
-57
No files matched your search
@@ -0,0 +1,41 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/spf13/viper"
|
||||
)
|
||||
|
||||
type ServerConfig struct {
|
||||
Port int
|
||||
}
|
||||
|
||||
type JWTConfig struct {
|
||||
Secret string
|
||||
Expiry time.Duration
|
||||
}
|
||||
|
||||
var (
|
||||
Server ServerConfig
|
||||
JWT JWTConfig
|
||||
)
|
||||
|
||||
func InitApplicationConfig(configPath string) (error, error) {
|
||||
v := viper.New()
|
||||
v.SetConfigFile(configPath)
|
||||
v.SetConfigType("yaml")
|
||||
|
||||
if err := v.ReadInConfig(); err != nil {
|
||||
return err, nil
|
||||
}
|
||||
|
||||
if err := v.UnmarshalKey("server", &Server); err != nil {
|
||||
return err, nil
|
||||
}
|
||||
|
||||
if err := v.UnmarshalKey("jwt", &JWT); err != nil {
|
||||
return err, nil
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"sync"
|
||||
|
||||
"github.com/fsnotify/fsnotify"
|
||||
"github.com/spf13/viper"
|
||||
)
|
||||
|
||||
// DBConfig 数据库单个配置结构
|
||||
type DBConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
Password string
|
||||
Dbname string
|
||||
}
|
||||
|
||||
// dbConfigs 全局变量,存放所有数据库配置,key为配置名,如 test1, test2
|
||||
var (
|
||||
dbConfigs map[string]DBConfig
|
||||
mu sync.RWMutex
|
||||
v *viper.Viper
|
||||
)
|
||||
|
||||
// InitDBConfig 初始化并加载 db.yaml 配置,同时启动监听
|
||||
func InitDBConfig(configPath string) error {
|
||||
v = viper.New()
|
||||
v.SetConfigFile(configPath)
|
||||
v.SetConfigType("yaml")
|
||||
|
||||
if err := v.ReadInConfig(); err != nil {
|
||||
return fmt.Errorf("读取数据库配置失败: %w", err)
|
||||
}
|
||||
|
||||
if err := unmarshalConfigs(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 监听配置文件变化
|
||||
v.WatchConfig()
|
||||
v.OnConfigChange(func(e fsnotify.Event) {
|
||||
log.Printf("数据库配置文件发生变化: %s\n", e.Name)
|
||||
if err := unmarshalConfigs(); err != nil {
|
||||
log.Printf("重新加载数据库配置失败: %v\n", err)
|
||||
} else {
|
||||
log.Println("数据库配置已更新")
|
||||
}
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// unmarshalConfigs 解析配置到全局变量,内部加锁保证并发安全
|
||||
func unmarshalConfigs() error {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
temp := make(map[string]DBConfig)
|
||||
if err := v.Unmarshal(&temp); err != nil {
|
||||
return fmt.Errorf("解析数据库配置失败: %w", err)
|
||||
}
|
||||
dbConfigs = temp
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetDBConfigs 并发安全地返回当前的所有数据库配置副本
|
||||
func GetDBConfigs() map[string]DBConfig {
|
||||
mu.RLock()
|
||||
defer mu.RUnlock()
|
||||
|
||||
// 返回一份拷贝,避免外部修改内部数据
|
||||
res := make(map[string]DBConfig, len(dbConfigs))
|
||||
for k, v := range dbConfigs {
|
||||
res[k] = v
|
||||
}
|
||||
return res
|
||||
}
|
||||
@@ -1,9 +1,12 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"base-framework/pkg/config"
|
||||
"base-framework/pkg/router"
|
||||
"base-framework/pkg/utils"
|
||||
"net/http"
|
||||
"base-framework/pkg/utils/response"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
@@ -13,43 +16,50 @@ func Auth() router.HandlerFunc {
|
||||
userIdHeader := c.Header("user_id")
|
||||
orgIDHeader := c.Header("org_id")
|
||||
|
||||
if tokenHeader == "" {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "Authorization header missing"})
|
||||
return
|
||||
}
|
||||
if userIdHeader == "" {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "user_id header missing"})
|
||||
return
|
||||
}
|
||||
if orgIDHeader == "" {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "org_id header missing"})
|
||||
// 缺少登录信息
|
||||
if tokenHeader == "" || userIdHeader == "" || orgIDHeader == "" {
|
||||
response.Error(c).Code(response.CodeNoLogin).Send()
|
||||
return
|
||||
}
|
||||
|
||||
// Bearer token 格式校验
|
||||
parts := strings.Fields(tokenHeader)
|
||||
if len(parts) != 2 || strings.ToLower(parts[0]) != "bearer" {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "Authorization header format must be Bearer {token}"})
|
||||
if len(parts) != 2 || strings.ToLower(parts[0]) != "Bearer" {
|
||||
response.Error(c).Code(response.CodeInvalidToken).Send()
|
||||
return
|
||||
}
|
||||
|
||||
// 解析 token
|
||||
tokenStr := parts[1]
|
||||
claims, err := utils.VerifyToken(tokenStr)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "Invalid token: " + err.Error()})
|
||||
switch {
|
||||
case errors.Is(err, utils.ErrTokenExpired):
|
||||
response.Error(c).Code(response.CodeLoginExpired).Send()
|
||||
default:
|
||||
response.Error(c).Code(response.CodeInvalidToken).Send()
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// user_id 校验
|
||||
if claims.UserID != userIdHeader {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "user_id does not match token"})
|
||||
response.Error(c).Code(response.CodeInvalidToken).Send()
|
||||
return
|
||||
}
|
||||
|
||||
// org_id 校验
|
||||
if claims.OrgID != orgIDHeader {
|
||||
c.JSON(http.StatusUnauthorized, map[string]string{"error": "org_id does not match token"})
|
||||
response.Error(c).Code(response.CodeInvalidToken).Send()
|
||||
return
|
||||
}
|
||||
|
||||
// 认证通过,继续执行后续中间件或处理器
|
||||
configs := config.GetDBConfigs()
|
||||
|
||||
for k := range configs {
|
||||
fmt.Println(k)
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -1,23 +1,58 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// BodyCache 专门负责缓存请求体,保证只读一次
|
||||
type BodyCache struct {
|
||||
once sync.Once
|
||||
Data []byte
|
||||
Err error
|
||||
}
|
||||
|
||||
// Load 读取请求体并缓存,只执行一次
|
||||
func (b *BodyCache) Load(r *http.Request) ([]byte, error) {
|
||||
b.once.Do(func() {
|
||||
if r.Body == nil {
|
||||
b.Err = http.ErrBodyNotAllowed
|
||||
return
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
b.Err = func() error {
|
||||
_, err := io.Copy(&buf, r.Body)
|
||||
return err
|
||||
}()
|
||||
if b.Err != nil {
|
||||
return
|
||||
}
|
||||
b.Data = buf.Bytes()
|
||||
// 重新设置请求体方便后续读取
|
||||
r.Body = io.NopCloser(bytes.NewReader(b.Data))
|
||||
})
|
||||
return b.Data, b.Err
|
||||
}
|
||||
|
||||
// Context 自定义请求上下文,组合 BodyCache
|
||||
type Context struct {
|
||||
Writer http.ResponseWriter
|
||||
Request *http.Request
|
||||
Params map[string]string
|
||||
index int
|
||||
handlers []HandlerFunc
|
||||
Index int
|
||||
Handlers []HandlerFunc
|
||||
|
||||
BodyCache BodyCache // 请求体缓存
|
||||
}
|
||||
|
||||
// Next 执行下一个中间件或处理函数
|
||||
func (c *Context) Next() {
|
||||
c.index++
|
||||
if c.index < len(c.handlers) {
|
||||
c.handlers[c.index](c)
|
||||
c.Index++
|
||||
if c.Index < len(c.Handlers) {
|
||||
c.Handlers[c.Index](c)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -31,7 +66,7 @@ func (c *Context) Header(key string) string {
|
||||
return c.Request.Header.Get(key)
|
||||
}
|
||||
|
||||
// JSON 返回JSON格式响应
|
||||
// JSON 返回 JSON 格式响应
|
||||
func (c *Context) JSON(statusCode int, data interface{}) {
|
||||
c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
c.Writer.WriteHeader(statusCode)
|
||||
@@ -39,3 +74,25 @@ func (c *Context) JSON(statusCode int, data interface{}) {
|
||||
http.Error(c.Writer, err.Error(), http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
// Body 方便读取请求体,实际调用 BodyCache 的 Load 方法
|
||||
func (c *Context) Body() ([]byte, error) {
|
||||
return c.BodyCache.Load(c.Request)
|
||||
}
|
||||
|
||||
// BindJSON 反序列化 JSON 请求体到 obj
|
||||
func (c *Context) BindJSON(obj interface{}) error {
|
||||
body, err := c.Body()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return json.Unmarshal(body, obj)
|
||||
}
|
||||
|
||||
// PostForm 获取表单参数
|
||||
func (c *Context) PostForm(key string) string {
|
||||
if err := c.Request.ParseForm(); err != nil {
|
||||
return ""
|
||||
}
|
||||
return c.Request.FormValue(key)
|
||||
}
|
||||
@@ -1,22 +1,28 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
appConfig "base-framework/configs/app"
|
||||
commonConfig "base-framework/pkg/config"
|
||||
"errors"
|
||||
"github.com/golang-jwt/jwt"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt"
|
||||
)
|
||||
|
||||
// CustomClaims 定义自己的 payload 结构,可以根据需要扩展
|
||||
var (
|
||||
ErrTokenExpired = errors.New("token expired")
|
||||
ErrTokenInvalid = errors.New("token invalid")
|
||||
)
|
||||
|
||||
// CustomClaims 定义自己的 payload 结构
|
||||
type CustomClaims struct {
|
||||
OrgID string `json:"org_id"`
|
||||
UserID string `json:"user_id"`
|
||||
jwt.StandardClaims
|
||||
}
|
||||
|
||||
// CreateToken 创建一个JWT token
|
||||
// CreateToken 创建一个 JWT token
|
||||
func CreateToken(orgID, userID string) (string, error) {
|
||||
expireTime := time.Now().Add(appConfig.JWT.Expiry).Unix()
|
||||
expireTime := time.Now().Add(commonConfig.JWT.Expiry).Unix()
|
||||
|
||||
claims := CustomClaims{
|
||||
OrgID: orgID,
|
||||
@@ -28,29 +34,35 @@ func CreateToken(orgID, userID string) (string, error) {
|
||||
}
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
signKey := []byte(appConfig.JWT.Secret)
|
||||
signKey := []byte(commonConfig.JWT.Secret)
|
||||
return token.SignedString(signKey)
|
||||
}
|
||||
|
||||
// VerifyToken 验证并解析 JWT token
|
||||
func VerifyToken(tokenString string) (*CustomClaims, error) {
|
||||
signKey := []byte(appConfig.JWT.Secret)
|
||||
signKey := []byte(commonConfig.JWT.Secret)
|
||||
|
||||
token, err := jwt.ParseWithClaims(tokenString, &CustomClaims{}, func(token *jwt.Token) (interface{}, error) {
|
||||
// 校验签名算法
|
||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, errors.New("unexpected signing method")
|
||||
return nil, ErrTokenInvalid
|
||||
}
|
||||
return signKey, nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
// 判断是否过期错误
|
||||
if ve, ok := err.(*jwt.ValidationError); ok {
|
||||
if ve.Errors&jwt.ValidationErrorExpired != 0 {
|
||||
return nil, ErrTokenExpired
|
||||
}
|
||||
}
|
||||
return nil, ErrTokenInvalid
|
||||
}
|
||||
|
||||
if claims, ok := token.Claims.(*CustomClaims); ok && token.Valid {
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
return nil, errors.New("invalid token")
|
||||
return nil, ErrTokenInvalid
|
||||
}
|
||||
@@ -12,23 +12,27 @@ type Result struct {
|
||||
Data interface{} `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
// 常量状态码
|
||||
const (
|
||||
CodeSuccess = "0000"
|
||||
CodeFail = "9999"
|
||||
// 通用
|
||||
CodeSuccess = "0000" // 成功
|
||||
CodeFail = "0001" // 失败
|
||||
|
||||
// 鉴权 / 登录
|
||||
CodeNoLogin = "1001" // 未登录
|
||||
CodeLoginExpired = "1002" // 登录过期
|
||||
CodeInvalidToken = "1003" // Token 无效
|
||||
CodeInvalidOrgCode = "1004" // 机构码有误
|
||||
CodeInvalidAccount = "1005" // 账号或密码有误
|
||||
)
|
||||
|
||||
// code 对应默认提示
|
||||
var codeMessages = map[string]string{
|
||||
CodeSuccess: "请求成功",
|
||||
CodeFail: "请求失败",
|
||||
|
||||
"1001": "缺少用户ID",
|
||||
"1002": "未授权",
|
||||
"1003": "无权限访问",
|
||||
"1004": "资源不存在",
|
||||
|
||||
"9000": "系统内部错误",
|
||||
CodeSuccess: "请求成功",
|
||||
CodeFail: "请求失败",
|
||||
CodeNoLogin: "未登录",
|
||||
CodeLoginExpired: "登录已过期",
|
||||
CodeInvalidToken: "Token 无效",
|
||||
CodeInvalidOrgCode: "机构码有误",
|
||||
CodeInvalidAccount: "账号或密码有误",
|
||||
}
|
||||
|
||||
// 兜底提示
|
||||
@@ -65,9 +69,7 @@ func Error(c *router.Context) *Builder {
|
||||
// Code 设置状态码(自动填充默认提示,除非后面手动改)
|
||||
func (b *Builder) Code(code string) *Builder {
|
||||
b.result.Code = code
|
||||
if b.result.Message == "" || b.result.Message == getMessage(b.result.Code) {
|
||||
b.result.Message = getMessage(code)
|
||||
}
|
||||
b.result.Message = getMessage(code)
|
||||
return b
|
||||
}
|
||||
|
||||
|
||||
Reference in new issue
Block a user