This commit is contained in:
oneao committed 2025-08-13 17:30:57 +08:00
1 parent 7552670e4a
commit 01ea998605
13 files changed
+173 -45

No files matched your search

@@ -109,7 +109,7 @@ func buildDSN(cfg DBConfig) string {
cfg.Host, cfg.Port, cfg.Username, cfg.Password, cfg.Dbname)
}
// ReloadDataSources 支持增量更新数据源
// ReloadDataSources 支持增量更新数据源,并验证数据库连接
func ReloadDataSources(newConfigs map[string]DBConfig) error {
mu.Lock()
defer mu.Unlock()
@@ -120,30 +120,39 @@ func ReloadDataSources(newConfigs map[string]DBConfig) error {
newDSN := buildDSN(cfg)
oldDSN, exists := configsDSN[key]
if !exists {
// 新增数据源或配置变更
if !exists || oldDSN != newDSN {
// 关闭旧连接(如果存在)
if exists {
if oldDB := dataSources[key]; oldDB != nil {
_ = oldDB.Close()
}
}
db, err := sql.Open("postgres", newDSN)
if err != nil {
log.Printf("[db] 新增数据源 %s 失败: %v", key, err)
log.Printf("[db] 数据源 %s Open 失败: %v", key, err)
continue
}
// 尝试实际连接数据库
if err := db.Ping(); err != nil {
log.Printf("[db] 数据源 %s 连接失败: %v", key, err)
_ = db.Close()
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()
if !exists {
log.Printf("[db] 新增数据源 %s 成功", key)
} else {
log.Printf("[db] 更新数据源 %s 成功", key)
}
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
}
@@ -21,6 +21,7 @@ func Auth() router.HandlerFunc {
// 缺少登录信息
if tokenHeader == "" || userIdHeader == "" || orgIDHeader == "" {
response.Error(c).Code(response.CodeNoLogin).Send()
c.Abort()
return
}
@@ -28,6 +29,7 @@ func Auth() router.HandlerFunc {
parts := strings.Fields(tokenHeader)
if len(parts) != 2 || strings.ToLower(parts[0]) != "bearer" {
response.Error(c).Code(response.CodeInvalidToken).Send()
c.Abort()
return
}
@@ -41,18 +43,21 @@ func Auth() router.HandlerFunc {
default:
response.Error(c).Code(response.CodeInvalidToken).Send()
}
c.Abort()
return
}
// user_id 校验
if claims.UserID != userIdHeader {
response.Error(c).Code(response.CodeInvalidToken).Send()
c.Abort()
return
}
// org_id 校验
if claims.OrgID != orgIDHeader {
response.Error(c).Code(response.CodeInvalidToken).Send()
c.Abort()
return
}
@@ -70,6 +75,8 @@ func Auth() router.HandlerFunc {
return
}
c.Set("orgId", orgIDHeader)
c.Next()
}
}
@@ -3,7 +3,6 @@ package middleware
import (
"base-framework/pkg/router"
"log"
"net/http"
"runtime/debug"
)
@@ -12,9 +11,6 @@ func Recover() router.HandlerFunc {
defer func() {
if err := recover(); err != nil {
log.Printf("[PANIC] %v\n%s", err, debug.Stack())
c.JSON(http.StatusInternalServerError, map[string]string{
"error": "Internal Server Error",
})
}
}()
c.Next()
@@ -31,32 +31,42 @@ func (b *BodyCache) Load(r *http.Request) ([]byte, error) {
return
}
b.Data = buf.Bytes()
// 重新设置请求体方便后续读取
r.Body = io.NopCloser(bytes.NewReader(b.Data))
})
return b.Data, b.Err
}
// Context 自定义请求上下文,封装了请求信息、响应写入器、路由参数、中间件处理等功能
// Context 自定义请求上下文
type Context struct {
writer http.ResponseWriter // HTTP 响应写入器,用于构造响应数据
request *http.Request // HTTP 请求对象,包含请求相关信息
params map[string]string // 路由参数,如动态路径中的变量值
index int // 当前执行的中间件/处理函数索引,用于控制 Next 调用流程
handlers []HandlerFunc // 本次请求的中间件和最终处理函数链
writer http.ResponseWriter
request *http.Request
params map[string]string
index int
handlers []HandlerFunc
bodyCache BodyCache // 请求体缓存,确保请求体只读一次且可多次读取
keys map[string]interface{} // 用于存储请求生命周期内的自定义数据(如用户信息、orgId等)
bodyCache BodyCache
keys map[string]interface{}
aborted bool // 是否中止
}
// Next 执行下一个中间件或处理函数
func (c *Context) Next() {
c.index++
if c.index < len(c.handlers) {
for c.index < len(c.handlers) {
if c.aborted {
break
}
c.handlers[c.index](c)
c.index++
}
}
// Abort 中断后续中间件/处理函数
func (c *Context) Abort() {
c.aborted = true
}
// Param 获取路由参数
func (c *Context) Param(key string) string {
return c.params[key]
@@ -76,7 +86,7 @@ func (c *Context) JSON(statusCode int, data interface{}) {
}
}
// Body 方便读取请求体,实际调用 BodyCache 的 Load 方法
// Body 方便读取请求体
func (c *Context) Body() ([]byte, error) {
return c.bodyCache.Load(c.request)
}
@@ -0,0 +1,38 @@
package db
import (
"base-framework/pkg/config"
"base-framework/pkg/router"
"base-framework/pkg/utils/response"
"database/sql"
"fmt"
)
// Client 封装数据库连接对象
type Client struct {
Conn *sql.DB
}
// NewClient 根据 context 获取 orgId 并返回 Client 对象
func NewClient(c *router.Context) *Client {
val, ok := c.Get("orgId")
fmt.Println(val)
if !ok {
response.Error(c).Code(response.CodeInvalidOrgCode).Send()
return nil
}
orgId, ok := val.(string)
if !ok || orgId == "" {
response.Error(c).Code(response.CodeInvalidOrgCode).Send()
return nil
}
conn, ok := config.GetDB(orgId)
if !ok || conn == nil {
response.Error(c).Code(response.CodeInvalidOrgCode).Send()
return nil
}
return &Client{Conn: conn}
}
@@ -0,0 +1 @@
package db
@@ -0,0 +1 @@
package db
@@ -0,0 +1,53 @@
package db
import (
"database/sql"
"fmt"
)
func (c *Client) QueryRows(query string, args ...any) ([]map[string]any, error) {
if c.Conn == nil {
return nil, fmt.Errorf("database connection is nil")
}
rows, err := c.Conn.Query(query, args...)
if err != nil {
return nil, err
}
defer func(rows *sql.Rows) {
_ = rows.Close()
}(rows)
// 获取列名
columns, err := rows.Columns()
if err != nil {
return nil, err
}
results := make([]map[string]any, 0)
for rows.Next() {
// 为每列创建一个接口值的 slice
values := make([]any, len(columns))
valuePtrs := make([]any, len(columns))
for i := range columns {
valuePtrs[i] = &values[i]
}
if err := rows.Scan(valuePtrs...); err != nil {
return nil, err
}
rowMap := make(map[string]any)
for i, col := range columns {
rowMap[col] = values[i]
}
results = append(results, rowMap)
}
if err := rows.Err(); err != nil {
return nil, err
}
return results, nil
}
@@ -0,0 +1 @@
package db