This commit is contained in:
oneao committed 2025-08-13 14:51:19 +08:00
1 parent ed8bee3d77
commit 7552670e4a
23 files changed
+380 -154

No files matched your search

+10
View File
@@ -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
+9
View File
@@ -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>
+6
View File
@@ -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>
+8
View File
@@ -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>
+6
View File
@@ -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
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,
}
@@ -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
}