u
This commit is contained in:
1 parent
7552670e4a
commit
01ea998605
13 files changed
+173
-45
No files matched your search
@@ -14,7 +14,7 @@ func main() {
|
||||
log.Fatalf("加载配置失败: %v", err)
|
||||
}
|
||||
|
||||
err = config.InitDBConfig("./configs/app/datasource.yaml")
|
||||
err = config.InitDBConfig("./configs/app/database.yaml")
|
||||
if err != nil {
|
||||
log.Fatalf("加载数据库配置失败: %v", err)
|
||||
return
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
g3cd:
|
||||
host: 117.72.182.135
|
||||
port: 5432
|
||||
username: postgres
|
||||
password: zhang520..
|
||||
dbname: test
|
||||
@@ -1,12 +0,0 @@
|
||||
test1:
|
||||
host: 127.0.0.1
|
||||
port: 5432
|
||||
username: admin
|
||||
password: 123456
|
||||
dbname: test1
|
||||
test2:
|
||||
host: 127.0.0.1
|
||||
port: 5432
|
||||
username: admin
|
||||
password: 123456
|
||||
dbname: test1
|
||||
@@ -2,9 +2,27 @@ package handle
|
||||
|
||||
import (
|
||||
"base-framework/pkg/router"
|
||||
"base-framework/pkg/utils/db"
|
||||
"base-framework/pkg/utils/response"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
func Test(c *router.Context) {
|
||||
response.Success(c).Send()
|
||||
dbClient := db.NewClient(c)
|
||||
if dbClient == nil {
|
||||
return
|
||||
}
|
||||
|
||||
sqlStr := "SELECT id, name FROM user"
|
||||
rows, err := dbClient.QueryRows(sqlStr)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
return
|
||||
}
|
||||
|
||||
for _, row := range rows {
|
||||
fmt.Println(row["id"], row["name"])
|
||||
}
|
||||
|
||||
response.Success(c).Data(rows).Send()
|
||||
}
|
||||
+25
-16
@@ -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
|
||||
Reference in new issue
Block a user