diff --git a/code/go-project/base-farmework/cmd/app/main.go b/code/go-project/base-farmework/cmd/app/main.go index c512f273..4cd0658a 100644 --- a/code/go-project/base-farmework/cmd/app/main.go +++ b/code/go-project/base-farmework/cmd/app/main.go @@ -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 diff --git a/code/go-project/base-farmework/configs/app/database.yaml b/code/go-project/base-farmework/configs/app/database.yaml new file mode 100644 index 00000000..c84ae752 --- /dev/null +++ b/code/go-project/base-farmework/configs/app/database.yaml @@ -0,0 +1,6 @@ +g3cd: + host: 117.72.182.135 + port: 5432 + username: postgres + password: zhang520.. + dbname: test diff --git a/code/go-project/base-farmework/configs/app/datasource.yaml b/code/go-project/base-farmework/configs/app/datasource.yaml deleted file mode 100644 index 88df8c8a..00000000 --- a/code/go-project/base-farmework/configs/app/datasource.yaml +++ /dev/null @@ -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 \ No newline at end of file diff --git a/code/go-project/base-farmework/internal/app/handle/test.go b/code/go-project/base-farmework/internal/app/handle/test.go index a97a4c8f..03649b0d 100644 --- a/code/go-project/base-farmework/internal/app/handle/test.go +++ b/code/go-project/base-farmework/internal/app/handle/test.go @@ -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() } diff --git a/code/go-project/base-farmework/pkg/config/datasource.go b/code/go-project/base-farmework/pkg/config/database.go similarity index 86% rename from code/go-project/base-farmework/pkg/config/datasource.go rename to code/go-project/base-farmework/pkg/config/database.go index ea71d109..35028cdb 100644 --- a/code/go-project/base-farmework/pkg/config/datasource.go +++ b/code/go-project/base-farmework/pkg/config/database.go @@ -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 } diff --git a/code/go-project/base-farmework/pkg/middleware/auth.go b/code/go-project/base-farmework/pkg/middleware/auth.go index b4951fbb..36a030d8 100644 --- a/code/go-project/base-farmework/pkg/middleware/auth.go +++ b/code/go-project/base-farmework/pkg/middleware/auth.go @@ -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() } } diff --git a/code/go-project/base-farmework/pkg/middleware/recover.go b/code/go-project/base-farmework/pkg/middleware/recover.go index f37e2222..c67e28a6 100644 --- a/code/go-project/base-farmework/pkg/middleware/recover.go +++ b/code/go-project/base-farmework/pkg/middleware/recover.go @@ -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() diff --git a/code/go-project/base-farmework/pkg/router/context.go b/code/go-project/base-farmework/pkg/router/context.go index 077ba3dd..86ab8f7e 100644 --- a/code/go-project/base-farmework/pkg/router/context.go +++ b/code/go-project/base-farmework/pkg/router/context.go @@ -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) } diff --git a/code/go-project/base-farmework/pkg/utils/db/db.go b/code/go-project/base-farmework/pkg/utils/db/db.go new file mode 100644 index 00000000..f06a23c7 --- /dev/null +++ b/code/go-project/base-farmework/pkg/utils/db/db.go @@ -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} +} diff --git a/code/go-project/base-farmework/pkg/utils/db/delete.go b/code/go-project/base-farmework/pkg/utils/db/delete.go new file mode 100644 index 00000000..3a49c63e --- /dev/null +++ b/code/go-project/base-farmework/pkg/utils/db/delete.go @@ -0,0 +1 @@ +package db diff --git a/code/go-project/base-farmework/pkg/utils/db/insert.go b/code/go-project/base-farmework/pkg/utils/db/insert.go new file mode 100644 index 00000000..3a49c63e --- /dev/null +++ b/code/go-project/base-farmework/pkg/utils/db/insert.go @@ -0,0 +1 @@ +package db diff --git a/code/go-project/base-farmework/pkg/utils/db/select.go b/code/go-project/base-farmework/pkg/utils/db/select.go new file mode 100644 index 00000000..75f3396c --- /dev/null +++ b/code/go-project/base-farmework/pkg/utils/db/select.go @@ -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 +} diff --git a/code/go-project/base-farmework/pkg/utils/db/update.go b/code/go-project/base-farmework/pkg/utils/db/update.go new file mode 100644 index 00000000..3a49c63e --- /dev/null +++ b/code/go-project/base-farmework/pkg/utils/db/update.go @@ -0,0 +1 @@ +package db