u
This commit is contained in:
1 parent
ed8bee3d77
commit
7552670e4a
23 files changed
+380
-154
No files matched your search
Generated
+10
@@ -0,0 +1,10 @@
|
||||
# Default ignored files
|
||||
/shelf/
|
||||
/workspace.xml
|
||||
# Editor-based HTTP Client requests
|
||||
/httpRequests/
|
||||
# Environment-dependent path to Maven home directory
|
||||
/mavenHomeManager.xml
|
||||
# Datasource local storage ignored files
|
||||
/dataSources/
|
||||
/dataSources.local.xml
|
||||
Generated
+9
@@ -0,0 +1,9 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<module type="JAVA_MODULE" version="4">
|
||||
<component name="NewModuleRootManager" inherit-compiler-output="true">
|
||||
<exclude-output />
|
||||
<content url="file://$MODULE_DIR$" />
|
||||
<orderEntry type="inheritedJdk" />
|
||||
<orderEntry type="sourceFolder" forTests="false" />
|
||||
</component>
|
||||
</module>
|
||||
Generated
+6
@@ -0,0 +1,6 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<project version="4">
|
||||
<component name="ProjectRootManager">
|
||||
<output url="file://$PROJECT_DIR$/out" />
|
||||
</component>
|
||||
</project>
|
||||
Generated
+8
@@ -0,0 +1,8 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<project version="4">
|
||||
<component name="ProjectModuleManager">
|
||||
<modules>
|
||||
<module fileurl="file://$PROJECT_DIR$/.idea/base-project.iml" filepath="$PROJECT_DIR$/.idea/base-project.iml" />
|
||||
</modules>
|
||||
</component>
|
||||
</project>
|
||||
Generated
+6
@@ -0,0 +1,6 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<project version="4">
|
||||
<component name="VcsDirectoryMappings">
|
||||
<mapping directory="$PROJECT_DIR$/../.." vcs="Git" />
|
||||
</component>
|
||||
</project>
|
||||
@@ -2,21 +2,19 @@ package main
|
||||
|
||||
import (
|
||||
"base-framework/internal/app"
|
||||
commonConfig "base-framework/pkg/config"
|
||||
"base-framework/pkg/config"
|
||||
"log"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
_ "github.com/lib/pq"
|
||||
)
|
||||
|
||||
func main() {
|
||||
_, err := commonConfig.InitApplicationConfig("./configs/app/application.yaml")
|
||||
_, err := config.InitApplicationConfig("./configs/app/application.yaml")
|
||||
if err != nil {
|
||||
log.Fatalf("加载配置失败: %v", err)
|
||||
}
|
||||
|
||||
err = commonConfig.InitDBConfig("./configs/app/db.yaml")
|
||||
err = config.InitDBConfig("./configs/app/datasource.yaml")
|
||||
if err != nil {
|
||||
log.Fatalf("加载数据库配置失败: %v", err)
|
||||
return
|
||||
@@ -24,7 +22,7 @@ func main() {
|
||||
|
||||
r := app.InitAppRouter()
|
||||
|
||||
err = http.ListenAndServe(":"+strconv.Itoa(commonConfig.Server.Port), r)
|
||||
err = http.ListenAndServe(":"+strconv.Itoa(config.Server.Port), r)
|
||||
if err != nil {
|
||||
log.Fatalf("server start failed: %v", err)
|
||||
}
|
||||
|
||||
Whitespace-only changes.
@@ -2,4 +2,4 @@ server:
|
||||
port: 8082
|
||||
jwt:
|
||||
secret: 3Bde3BGEbYqtqyEUzW3ry8jKFcaPH17fRmTmqE7MDr05Lwj95uruRKrrkb44TJ4s
|
||||
expiry: 43200 # 12 * 60 * 60 秒过期
|
||||
expiry: 24h
|
||||
File renamed without changes.
Whitespace-only changes.
@@ -1 +1,11 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"base-framework/internal/app/handle"
|
||||
"base-framework/pkg/router"
|
||||
)
|
||||
|
||||
func InitAuth(r *router.Router) {
|
||||
group := r.Group("/auth")
|
||||
group.POST("/login", handle.Login)
|
||||
}
|
||||
@@ -1,31 +1,66 @@
|
||||
package handle
|
||||
|
||||
import "base-framework/pkg/router"
|
||||
import (
|
||||
"base-framework/pkg/config"
|
||||
"base-framework/pkg/router"
|
||||
"base-framework/pkg/utils/jwt"
|
||||
"base-framework/pkg/utils/response"
|
||||
)
|
||||
|
||||
// Login 登录接口
|
||||
func Login(c *router.Context) {
|
||||
// 直接用 BindJSON 绑定请求体 JSON 到结构体
|
||||
// 绑定请求体
|
||||
var req struct {
|
||||
OrgID string `json:"orgID"`
|
||||
UserID string `json:"userID"`
|
||||
OrgID string `json:"orgId"`
|
||||
UserID string `json:"userId"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
if err := c.BindJSON(&req); err != nil {
|
||||
c.JSON(400, map[string]string{"error": "请求体格式错误"})
|
||||
response.Error(c).Message("请求体不能为空").Send()
|
||||
return
|
||||
}
|
||||
|
||||
// 模拟校验
|
||||
if req.UserID == "" || req.Password == "" {
|
||||
c.JSON(400, map[string]string{"error": "账号或密码不能为空"})
|
||||
if req.OrgID == "" {
|
||||
response.Error(c).Message("机构码不能为空").Send()
|
||||
return
|
||||
}
|
||||
|
||||
// 登录成功,返回 token 示例
|
||||
c.JSON(200, map[string]interface{}{
|
||||
"code": "0000",
|
||||
"message": "登录成功",
|
||||
"data": map[string]string{
|
||||
"token": "这里是token字符串",
|
||||
},
|
||||
})
|
||||
// 参数校验
|
||||
if req.UserID == "" {
|
||||
response.Error(c).Message("账户不能为空").Send()
|
||||
return
|
||||
}
|
||||
|
||||
if req.Password == "" {
|
||||
response.Error(c).Message("密码不能为空").Send()
|
||||
return
|
||||
}
|
||||
|
||||
// 检查是否有该机构号
|
||||
hasOrg := false
|
||||
for k := range config.GetDBConfigs() {
|
||||
if k == req.OrgID {
|
||||
hasOrg = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !hasOrg {
|
||||
response.Error(c).Code(response.CodeInvalidOrgCode).Send()
|
||||
return
|
||||
}
|
||||
|
||||
// 生成Token
|
||||
token, err := jwt.CreateToken(req.OrgID, req.UserID)
|
||||
|
||||
if err != nil {
|
||||
response.Error(c).Message("生成令牌失败")
|
||||
return
|
||||
}
|
||||
|
||||
response.Success(c).Data(
|
||||
map[string]string{
|
||||
"token": token,
|
||||
}).Send()
|
||||
}
|
||||
@@ -10,14 +10,14 @@ func InitAppRouter() *router.Router {
|
||||
r := router.NewRouter()
|
||||
|
||||
r.Use(middleware.Recover())
|
||||
r.Use(middleware.Auth()).ExcludePaths("/login")
|
||||
r.Use(middleware.Auth()).ExcludePaths("/auth/login")
|
||||
r.Use(middleware.Logger())
|
||||
|
||||
initApi(r)
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
func initApi(r *router.Router) {
|
||||
api.InitTest(r)
|
||||
api.InitAuth(r)
|
||||
}
|
||||
Whitespace-only changes.
@@ -1,25 +1,28 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/spf13/viper"
|
||||
"time"
|
||||
)
|
||||
|
||||
type ServerConfig struct {
|
||||
// server 服务器配置结构体(包内私有)
|
||||
type server struct {
|
||||
Port int
|
||||
}
|
||||
|
||||
type JWTConfig struct {
|
||||
// jwt 配置结构体(包内私有)
|
||||
type jwt struct {
|
||||
Secret string
|
||||
Expiry time.Duration
|
||||
}
|
||||
|
||||
// 全局导出变量,指向私有结构体实例
|
||||
var (
|
||||
Server ServerConfig
|
||||
JWT JWTConfig
|
||||
Server *server
|
||||
JWT *jwt
|
||||
)
|
||||
|
||||
// InitApplicationConfig 初始化配置,传入配置文件路径
|
||||
func InitApplicationConfig(configPath string) (error, error) {
|
||||
v := viper.New()
|
||||
v.SetConfigFile(configPath)
|
||||
@@ -29,13 +32,17 @@ func InitApplicationConfig(configPath string) (error, error) {
|
||||
return err, nil
|
||||
}
|
||||
|
||||
if err := v.UnmarshalKey("server", &Server); err != nil {
|
||||
var s server
|
||||
if err := v.UnmarshalKey("server", &s); err != nil {
|
||||
return err, nil
|
||||
}
|
||||
Server = &s
|
||||
|
||||
if err := v.UnmarshalKey("jwt", &JWT); err != nil {
|
||||
var j jwt
|
||||
if err := v.UnmarshalKey("jwt", &j); err != nil {
|
||||
return err, nil
|
||||
}
|
||||
JWT = &j
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
@@ -0,0 +1,191 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/fsnotify/fsnotify"
|
||||
"github.com/spf13/viper"
|
||||
|
||||
_ "github.com/lib/pq" // PostgreSQL 驱动,根据实际数据库替换
|
||||
)
|
||||
|
||||
// DBConfig 数据库配置结构体
|
||||
type DBConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
Password string
|
||||
Dbname string
|
||||
MaxOpenConns int // 最大打开连接数
|
||||
MaxIdleConns int // 最大空闲连接数
|
||||
ConnMaxLifetime time.Duration // 连接最大生命周期
|
||||
}
|
||||
|
||||
// 默认连接池参数
|
||||
const (
|
||||
defaultMaxOpenConns = 10
|
||||
defaultMaxIdleConns = 5
|
||||
defaultConnMaxLifetime = time.Hour
|
||||
)
|
||||
|
||||
var (
|
||||
dbConfigs map[string]DBConfig // 配置副本
|
||||
dataSources = make(map[string]*sql.DB) // 连接池
|
||||
configsDSN = make(map[string]string) // key -> dsn
|
||||
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
|
||||
}
|
||||
|
||||
// 首次加载后,初始化数据源
|
||||
if err := ReloadDataSources(GetDBConfigs()); err != nil {
|
||||
log.Printf("首次加载数据源失败: %v\n", err)
|
||||
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 {
|
||||
if err := ReloadDataSources(GetDBConfigs()); 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
|
||||
}
|
||||
|
||||
// buildDSN 构建 PostgreSQL DSN
|
||||
func buildDSN(cfg DBConfig) string {
|
||||
return fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=disable",
|
||||
cfg.Host, cfg.Port, cfg.Username, cfg.Password, cfg.Dbname)
|
||||
}
|
||||
|
||||
// ReloadDataSources 支持增量更新数据源
|
||||
func ReloadDataSources(newConfigs map[string]DBConfig) error {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
retain := make(map[string]bool)
|
||||
|
||||
for key, cfg := range newConfigs {
|
||||
newDSN := buildDSN(cfg)
|
||||
oldDSN, exists := configsDSN[key]
|
||||
|
||||
if !exists {
|
||||
db, err := sql.Open("postgres", newDSN)
|
||||
if err != nil {
|
||||
log.Printf("[db] 新增数据源 %s 失败: %v", key, err)
|
||||
continue
|
||||
}
|
||||
applyDBConfig(db, cfg)
|
||||
dataSources[key] = db
|
||||
configsDSN[key] = newDSN
|
||||
log.Printf("[db] 新增数据源 %s 成功", key)
|
||||
} else if oldDSN != newDSN {
|
||||
if oldDB := dataSources[key]; oldDB != nil {
|
||||
_ = oldDB.Close()
|
||||
}
|
||||
db, err := sql.Open("postgres", newDSN)
|
||||
if err != nil {
|
||||
log.Printf("[db] 更新数据源 %s 失败: %v", key, err)
|
||||
continue
|
||||
}
|
||||
applyDBConfig(db, cfg)
|
||||
dataSources[key] = db
|
||||
configsDSN[key] = newDSN
|
||||
log.Printf("[db] 更新数据源 %s 成功", key)
|
||||
}
|
||||
// 配置没变,不用操作
|
||||
retain[key] = true
|
||||
}
|
||||
|
||||
// 关闭并移除不再需要的数据源
|
||||
for key, db := range dataSources {
|
||||
if !retain[key] {
|
||||
_ = db.Close()
|
||||
delete(dataSources, key)
|
||||
delete(configsDSN, key)
|
||||
log.Printf("[db] 移除数据源 %s", key)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// applyDBConfig 设置 sql.DB 连接池参数,并应用默认值
|
||||
func applyDBConfig(db *sql.DB, cfg DBConfig) {
|
||||
maxOpen := cfg.MaxOpenConns
|
||||
if maxOpen <= 0 {
|
||||
maxOpen = defaultMaxOpenConns
|
||||
}
|
||||
db.SetMaxOpenConns(maxOpen)
|
||||
|
||||
maxIdle := cfg.MaxIdleConns
|
||||
if maxIdle <= 0 {
|
||||
maxIdle = defaultMaxIdleConns
|
||||
}
|
||||
db.SetMaxIdleConns(maxIdle)
|
||||
|
||||
timeout := cfg.ConnMaxLifetime
|
||||
if timeout <= 0 {
|
||||
timeout = defaultConnMaxLifetime
|
||||
}
|
||||
db.SetConnMaxLifetime(timeout)
|
||||
}
|
||||
|
||||
// GetDB 获取数据库连接
|
||||
func GetDB(key string) (*sql.DB, bool) {
|
||||
mu.RLock()
|
||||
defer mu.RUnlock()
|
||||
db, ok := dataSources[key]
|
||||
return db, ok
|
||||
}
|
||||
@@ -1,80 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -3,18 +3,20 @@ package middleware
|
||||
import (
|
||||
"base-framework/pkg/config"
|
||||
"base-framework/pkg/router"
|
||||
"base-framework/pkg/utils"
|
||||
"base-framework/pkg/utils/jwt"
|
||||
"base-framework/pkg/utils/response"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func Auth() router.HandlerFunc {
|
||||
return func(c *router.Context) {
|
||||
log.Println("[Auth] 开始鉴权")
|
||||
|
||||
tokenHeader := c.Header("Authorization")
|
||||
userIdHeader := c.Header("user_id")
|
||||
orgIDHeader := c.Header("org_id")
|
||||
userIdHeader := c.Header("userId")
|
||||
orgIDHeader := c.Header("orgId")
|
||||
|
||||
// 缺少登录信息
|
||||
if tokenHeader == "" || userIdHeader == "" || orgIDHeader == "" {
|
||||
@@ -24,17 +26,17 @@ func Auth() router.HandlerFunc {
|
||||
|
||||
// Bearer token 格式校验
|
||||
parts := strings.Fields(tokenHeader)
|
||||
if len(parts) != 2 || strings.ToLower(parts[0]) != "Bearer" {
|
||||
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)
|
||||
claims, err := jwt.VerifyToken(tokenStr)
|
||||
if err != nil {
|
||||
switch {
|
||||
case errors.Is(err, utils.ErrTokenExpired):
|
||||
case errors.Is(err, jwt.ErrTokenExpired):
|
||||
response.Error(c).Code(response.CodeLoginExpired).Send()
|
||||
default:
|
||||
response.Error(c).Code(response.CodeInvalidToken).Send()
|
||||
@@ -54,10 +56,18 @@ func Auth() router.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
configs := config.GetDBConfigs()
|
||||
// 检查是否有该机构号
|
||||
hasOrg := false
|
||||
for k := range config.GetDBConfigs() {
|
||||
if k == orgIDHeader {
|
||||
hasOrg = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
for k := range configs {
|
||||
fmt.Println(k)
|
||||
if !hasOrg {
|
||||
response.Error(c).Code(response.CodeInvalidOrgCode).Send()
|
||||
return
|
||||
}
|
||||
|
||||
c.Next()
|
||||
|
||||
@@ -12,8 +12,6 @@ func Logger() router.HandlerFunc {
|
||||
c.Next()
|
||||
elapsed := time.Since(start)
|
||||
// 即使业务中断,这里也能执行,打印耗时
|
||||
path := c.Request.URL.Path
|
||||
method := c.Request.Method
|
||||
println(method, path, "耗时:", elapsed.String())
|
||||
println("耗时:", elapsed.String())
|
||||
}
|
||||
}
|
||||
Whitespace-only changes.
@@ -37,47 +37,48 @@ func (b *BodyCache) Load(r *http.Request) ([]byte, error) {
|
||||
return b.Data, b.Err
|
||||
}
|
||||
|
||||
// Context 自定义请求上下文,组合 BodyCache
|
||||
// Context 自定义请求上下文,封装了请求信息、响应写入器、路由参数、中间件处理等功能
|
||||
type Context struct {
|
||||
Writer http.ResponseWriter
|
||||
Request *http.Request
|
||||
Params map[string]string
|
||||
Index int
|
||||
Handlers []HandlerFunc
|
||||
writer http.ResponseWriter // HTTP 响应写入器,用于构造响应数据
|
||||
request *http.Request // HTTP 请求对象,包含请求相关信息
|
||||
params map[string]string // 路由参数,如动态路径中的变量值
|
||||
index int // 当前执行的中间件/处理函数索引,用于控制 Next 调用流程
|
||||
handlers []HandlerFunc // 本次请求的中间件和最终处理函数链
|
||||
|
||||
BodyCache BodyCache // 请求体缓存
|
||||
bodyCache BodyCache // 请求体缓存,确保请求体只读一次且可多次读取
|
||||
keys map[string]interface{} // 用于存储请求生命周期内的自定义数据(如用户信息、orgId等)
|
||||
}
|
||||
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
|
||||
// Param 获取路由参数
|
||||
func (c *Context) Param(key string) string {
|
||||
return c.Params[key]
|
||||
return c.params[key]
|
||||
}
|
||||
|
||||
// Header 获取请求头
|
||||
func (c *Context) Header(key string) string {
|
||||
return c.Request.Header.Get(key)
|
||||
return c.request.Header.Get(key)
|
||||
}
|
||||
|
||||
// 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)
|
||||
if err := json.NewEncoder(c.Writer).Encode(data); err != nil {
|
||||
http.Error(c.Writer, err.Error(), http.StatusInternalServerError)
|
||||
c.writer.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
c.writer.WriteHeader(statusCode)
|
||||
if err := json.NewEncoder(c.writer).Encode(data); err != nil {
|
||||
http.Error(c.writer, err.Error(), http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
// Body 方便读取请求体,实际调用 BodyCache 的 Load 方法
|
||||
func (c *Context) Body() ([]byte, error) {
|
||||
return c.BodyCache.Load(c.Request)
|
||||
return c.bodyCache.Load(c.request)
|
||||
}
|
||||
|
||||
// BindJSON 反序列化 JSON 请求体到 obj
|
||||
@@ -91,8 +92,25 @@ func (c *Context) BindJSON(obj interface{}) error {
|
||||
|
||||
// PostForm 获取表单参数
|
||||
func (c *Context) PostForm(key string) string {
|
||||
if err := c.Request.ParseForm(); err != nil {
|
||||
if err := c.request.ParseForm(); err != nil {
|
||||
return ""
|
||||
}
|
||||
return c.Request.FormValue(key)
|
||||
return c.request.FormValue(key)
|
||||
}
|
||||
|
||||
// Set 存储键值对
|
||||
func (c *Context) Set(key string, value interface{}) {
|
||||
if c.keys == nil {
|
||||
c.keys = make(map[string]interface{})
|
||||
}
|
||||
c.keys[key] = value
|
||||
}
|
||||
|
||||
// Get 取值
|
||||
func (c *Context) Get(key string) (interface{}, bool) {
|
||||
if c.keys == nil {
|
||||
return nil, false
|
||||
}
|
||||
val, ok := c.keys[key]
|
||||
return val, ok
|
||||
}
|
||||
@@ -209,9 +209,9 @@ func (r *Router) ServeHTTP(w http.ResponseWriter, req *http.Request) {
|
||||
handlers = append(handlers, n.handler)
|
||||
|
||||
c := &Context{
|
||||
Writer: w,
|
||||
Request: req,
|
||||
Params: params,
|
||||
writer: w,
|
||||
request: req,
|
||||
params: params,
|
||||
index: -1,
|
||||
handlers: handlers,
|
||||
}
|
||||
|
||||
+6
-6
@@ -1,4 +1,4 @@
|
||||
package utils
|
||||
package jwt
|
||||
|
||||
import (
|
||||
commonConfig "base-framework/pkg/config"
|
||||
@@ -22,20 +22,19 @@ type CustomClaims struct {
|
||||
|
||||
// CreateToken 创建一个 JWT token
|
||||
func CreateToken(orgID, userID string) (string, error) {
|
||||
expireTime := time.Now().Add(commonConfig.JWT.Expiry).Unix()
|
||||
expireAt := time.Now().Add(commonConfig.JWT.Expiry)
|
||||
|
||||
claims := CustomClaims{
|
||||
OrgID: orgID,
|
||||
UserID: userID,
|
||||
StandardClaims: jwt.StandardClaims{
|
||||
ExpiresAt: expireTime,
|
||||
ExpiresAt: expireAt.Unix(),
|
||||
IssuedAt: time.Now().Unix(),
|
||||
},
|
||||
}
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
signKey := []byte(commonConfig.JWT.Secret)
|
||||
return token.SignedString(signKey)
|
||||
return token.SignedString([]byte(commonConfig.JWT.Secret))
|
||||
}
|
||||
|
||||
// VerifyToken 验证并解析 JWT token
|
||||
@@ -52,7 +51,8 @@ func VerifyToken(tokenString string) (*CustomClaims, error) {
|
||||
|
||||
if err != nil {
|
||||
// 判断是否过期错误
|
||||
if ve, ok := err.(*jwt.ValidationError); ok {
|
||||
var ve *jwt.ValidationError
|
||||
if errors.As(err, &ve) {
|
||||
if ve.Errors&jwt.ValidationErrorExpired != 0 {
|
||||
return nil, ErrTokenExpired
|
||||
}
|
||||
Reference in new issue
Block a user