This commit is contained in:
oneao committed 2026-04-24 22:37:52 +08:00
1 parent 46cbb2f50d
commit e3d765a9d4
15 files changed
+903 -190

No files matched your search

@@ -43,24 +43,40 @@ func InitPgsql(cfg *config.Config) (*pgxpool.Pool, error) {
}
// ======================
// 连接池配置
// 连接池配置(带默认值)
// ======================
// 最大连接数
// 默认最大连接数
maxConns := int32(20)
if pg.MaxOpenConns > 0 {
conf.MaxConns = int32(pg.MaxOpenConns)
maxConns = int32(pg.MaxOpenConns)
}
// 最小连接数(近似 idle)
// 默认最小连接数
minConns := int32(2)
if pg.MaxIdleConns > 0 {
conf.MinConns = int32(pg.MaxIdleConns)
minConns = int32(pg.MaxIdleConns)
}
// 连接最大生命周期
// 防止配置错误:min > max
if minConns > maxConns {
minConns = maxConns
}
// 应用配置
conf.MaxConns = maxConns
conf.MinConns = minConns
// 默认连接最大生命周期
if pg.ConnMaxLifetime > 0 {
conf.MaxConnLifetime = pg.ConnMaxLifetime
} else {
conf.MaxConnLifetime = time.Hour
}
// (推荐补充)空闲连接最大存活时间
conf.MaxConnIdleTime = 30 * time.Minute
// ======================
// 时区设置(正确方式 ⭐)
// ======================
@@ -3,6 +3,7 @@ package handle
import (
"allapp-go/internal/errors"
"allapp-go/internal/httpx"
"allapp-go/internal/middleware"
"allapp-go/internal/types"
"allapp-go/pkg/db"
"allapp-go/pkg/jwtx"
@@ -10,7 +11,9 @@ import (
"allapp-go/pkg/s3store"
"allapp-go/pkg/uniqueid"
"allapp-go/pkg/wechat"
"bytes"
"fmt"
"io"
"net/url"
"path/filepath"
"time"
@@ -161,6 +164,14 @@ func registerUser(
_ = stream.Close()
}()
// ✅ 只在这里转成可 seek
data, err := io.ReadAll(stream)
if err != nil {
return err
}
reader := bytes.NewReader(data)
// ✅ 从 avatar URL 获取扩展名(简化版)
u, _ := url.Parse(avatar)
ext := filepath.Ext(u.Path)
@@ -181,12 +192,12 @@ func registerUser(
// 上传到 S3
if err := s3store.UploadToRustFS(
c.Context(),
stream,
reader,
key,
size,
contentType,
); err != nil {
avatar = defaultAvatar
return errors.WithStack(err)
} else {
avatar = key
}
@@ -240,9 +251,160 @@ func registerUser(
return httpx.OK(c, vo)
}
// BindQq 绑定QQ
func BindQq(c fiber.Ctx) error {
var req types.LoginQqReq
if err := httpx.BindAndValidate(c, &req); err != nil {
return errors.WithStack(err)
}
dbClient := db.New()
// 检查 QQ 是否已经被其他账号绑定
var exist []map[string]any
_, err := dbClient.LoadDataBySQL(
c.Context(),
"SELECT id FROM b_user_oauth WHERE openid = ? AND type = 1",
[]any{req.Openid},
)
if err != nil {
return errors.WithStack(err)
}
if len(exist) > 0 {
return httpx.Fail(c, "该 QQ 已被其他账号绑定")
}
userID, ok := middleware.GetUserID(c.Context())
if !ok {
return httpx.Unauthorized(c, "账号异常")
}
var userBind []map[string]any
_, err = dbClient.LoadDataBySQL(
c.Context(),
"SELECT id FROM b_user_oauth WHERE user_id = ? AND type = 1",
[]any{userID},
)
if err != nil {
return errors.WithStack(err)
}
if len(userBind) > 0 {
return httpx.Fail(c, "该账号已绑定 QQ")
}
err = dbClient.WithTx(c.Context(), func(tx *db.Client) error {
err2 := tx.Insert(
c.Context(),
"b_user_oauth",
"id",
map[string]any{
"id": uniqueid.NextId(),
"user_id": userID,
"type": 1, // QQ
"openid": req.Openid,
"create_time": time.Now(),
"update_time": time.Now(),
},
)
if err2 != nil {
return err2
}
return nil
})
if err != nil {
return errors.WithStack(err)
}
return httpx.OK(c, "绑定成功")
}
// BindWechat 绑定微信
func BindWechat(c fiber.Ctx) error {
var req types.LoginWechatReq
if err := httpx.BindAndValidate(c, &req); err != nil {
return errors.WithStack(err)
}
dbClient := db.New()
// 2️⃣ 检查 微信 是否已经被其他账号绑定
var exist []map[string]any
_, err := dbClient.LoadDataBySQL(
c.Context(),
"SELECT id FROM b_user_oauth WHERE openid = ? AND type = 0",
[]any{req.Code},
)
if err != nil {
return errors.WithStack(err)
}
if len(exist) > 0 {
return httpx.Fail(c, "该 微信 已被其他账号绑定")
}
userID, ok := middleware.GetUserID(c.Context())
if !ok {
return httpx.Unauthorized(c, "账号异常")
}
var userBind []map[string]any
_, err = dbClient.LoadDataBySQL(
c.Context(),
"SELECT id FROM b_user_oauth WHERE user_id = ? AND type = 0",
[]any{userID},
)
if err != nil {
return errors.WithStack(err)
}
if len(userBind) > 0 {
return httpx.Fail(c, "该账号已绑定 微信")
}
openid, _, err := wechat.GetWechatAccess(req.Code)
if err != nil {
log.Errorw("微信获取Token失败", "code", req.Code, "error", err)
return httpx.Fail(c, "微信登录失败,请重试")
}
err = dbClient.WithTx(c.Context(), func(tx *db.Client) error {
err2 := tx.Insert(
c.Context(),
"b_user_oauth",
"id",
map[string]any{
"id": uniqueid.NextId(),
"user_id": userID,
"type": 0, // 微信
"openid": openid,
"create_time": time.Now(),
"update_time": time.Now(),
},
)
if err2 != nil {
return err2
}
return nil
})
if err != nil {
return errors.WithStack(err)
}
return httpx.OK(c, "绑定成功")
}
// ======================== DB 查询封装(去重复 SQL) ========================
func getUserByOpenID(dbClient *db.Client, c fiber.Ctx, openid string, loginType int16) (map[string]any, error) {
users, err := dbClient.LoadDataBySQL(
c.Context(),
`SELECT u.*
@@ -26,4 +26,8 @@ func SetupRouter(app *fiber.App, cfg *config.Config) {
// ==================== auth ====================
api.Post("/auth/login/qq", handle.LoginQq)
api.Post("/auth/login/wechat", handle.LoginWechat)
bind := api.Group("/bind", middleware.Auth())
bind.Post("/qq", handle.BindQq)
bind.Post("/wechat", handle.BindWechat)
}
+22 -14
View File
@@ -10,33 +10,41 @@ import (
var defaultClient *Client
type Executor interface {
Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error)
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
}
type Client struct {
pool *pgxpool.Pool
tx pgx.Tx
}
// Init 初始化
// Init 初始化(只调用一次)
func Init(pool *pgxpool.Pool) {
defaultClient = &Client{pool: pool}
}
// New 创建会话
// New 获取全局 client
func New() *Client {
return &Client{
pool: defaultClient.pool,
if defaultClient == nil {
panic("db not initialized, call db.Init(pool) first")
}
return defaultClient
}
// 内部获取执行器(关键)
func (c *Client) exec() Executor {
func (c *Client) Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error) {
if c.tx != nil {
return c.tx
return c.tx.Exec(ctx, sql, args...)
}
return c.pool
return c.pool.Exec(ctx, sql, args...)
}
func (c *Client) Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) {
if c.tx != nil {
return c.tx.Query(ctx, sql, args...)
}
return c.pool.Query(ctx, sql, args...)
}
func (c *Client) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row {
if c.tx != nil {
return c.tx.QueryRow(ctx, sql, args...)
}
return c.pool.QueryRow(ctx, sql, args...)
}
+8 -15
View File
@@ -7,32 +7,24 @@ import (
)
func (c *Client) Delete(ctx context.Context, table, keyField string, data map[string]any) error {
_, err := validateData(table, keyField, data)
keyVal, err := validateData(table, keyField, data)
if err != nil {
return err
}
where := make([]string, 0, len(data))
args := make([]any, 0, len(data))
i := 1
for k, v := range data {
where = append(where, fmt.Sprintf("%s = $%d", quoteCol(k), i))
args = append(args, v)
i++
}
sql := fmt.Sprintf(
"DELETE FROM %s WHERE %s",
"DELETE FROM %s WHERE %s = $1",
quoteTable(table),
strings.Join(where, " AND "),
quoteCol(keyField),
)
_, err = c.exec().Exec(ctx, sql, args...)
_, err = c.pool.Exec(ctx, sql, keyVal) // ⭐修复 exec
return err
}
func (c *Client) BatchDelete(ctx context.Context, table, keyField string, list []map[string]any) error {
if len(list) == 0 {
return fmt.Errorf("empty data")
}
@@ -44,6 +36,7 @@ func (c *Client) BatchDelete(ctx context.Context, table, keyField string, list [
)
for _, row := range list {
v, err := validateData(table, keyField, row)
if err != nil {
return err
@@ -61,6 +54,6 @@ func (c *Client) BatchDelete(ctx context.Context, table, keyField string, list [
strings.Join(in, ", "),
)
_, err := c.exec().Exec(ctx, sql, args...)
_, err := c.pool.Exec(ctx, sql, args...) // ⭐修复 exec
return err
}
+10 -4
View File
@@ -31,6 +31,7 @@ func buildInsertSQL(table string, data map[string]any) (string, []any) {
}
func (c *Client) Insert(ctx context.Context, table, keyField string, data map[string]any) error {
if _, err := validateData(table, keyField, data); err != nil {
return err
}
@@ -38,11 +39,13 @@ func (c *Client) Insert(ctx context.Context, table, keyField string, data map[st
data = applyMetaFields(ctx, table, data, true)
sql, args := buildInsertSQL(table, data)
_, err := c.exec().Exec(ctx, sql, args...)
_, err := c.pool.Exec(ctx, sql, args...)
return err
}
func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list []map[string]any) error {
if len(list) == 0 {
return fmt.Errorf("empty data")
}
@@ -54,6 +57,7 @@ func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list [
tableSQL := quoteTable(table)
// 固定字段顺序(稳定性关键)
var cols []string
for k := range first {
cols = append(cols, k)
@@ -67,9 +71,11 @@ func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list [
)
for _, row := range list {
row = applyMetaFields(ctx, table, row, true)
var place []string
for _, col := range cols {
v, ok := row[col]
if !ok {
@@ -85,8 +91,8 @@ func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list [
}
var quotedCols []string
for _, c := range cols {
quotedCols = append(quotedCols, quoteCol(c))
for _, col := range cols {
quotedCols = append(quotedCols, quoteCol(col))
}
sql := fmt.Sprintf(
@@ -96,6 +102,6 @@ func (c *Client) BatchInsert(ctx context.Context, table, keyField string, list [
strings.Join(values, ", "),
)
_, err := c.exec().Exec(ctx, sql, args...)
_, err := c.pool.Exec(ctx, sql, args...) // ⭐关键修复
return err
}
+13 -10
View File
@@ -34,7 +34,8 @@ func (c *Client) LoadData(
order = " ORDER BY " + orderBy
}
sql := fmt.Sprintf("SELECT %s FROM %s%s%s",
sql := fmt.Sprintf(
"SELECT %s FROM %s%s%s",
selectCols,
viewName,
where,
@@ -54,26 +55,25 @@ func (c *Client) LoadDataBySQL(
sql string,
args []any,
) ([]map[string]any, error) {
logger.FromCtx(ctx).Info("LoadDataBySQL",
zap.String("sql", sql),
zap.Any("args", args),
)
return c.query(ctx, sql, args...)
}
// ==========================
// 核心查询方法(已升级)
// ==========================
func (c *Client) query(ctx context.Context, sql string, args ...any) ([]map[string]any, error) {
sql = convertPlaceholder(sql)
rows, err := c.exec().Query(ctx, sql, args...)
rows, err := c.pool.Query(ctx, sql, args...)
if err != nil {
return nil, err
}
defer rows.Close()
// 获取字段名
fields := rows.FieldDescriptions()
result := make([]map[string]any, 0, 16)
@@ -86,19 +86,22 @@ func (c *Client) query(ctx context.Context, sql string, args ...any) ([]map[stri
row := make(map[string]any, len(values))
for i, f := range fields {
for i := range fields {
if i < len(values) {
row[string(f.Name)] = values[i]
row[string(fields[i].Name)] = values[i]
}
}
result = append(result, row)
}
return result, rows.Err()
if err := rows.Err(); err != nil {
return nil, err
}
return result, nil
}
// 高效 placeholder 转换
func convertPlaceholder(sql string) string {
var sb strings.Builder
sb.Grow(len(sql))
+14 -10
View File
@@ -2,12 +2,10 @@ package db
import (
"context"
"github.com/jackc/pgx/v5"
)
// WithTx 开启事务(在当前 client 上)
func (c *Client) WithTx(ctx context.Context, fn func(tx *Client) error) error {
tx, err := c.pool.Begin(ctx)
if err != nil {
return err
@@ -19,16 +17,22 @@ func (c *Client) WithTx(ctx context.Context, fn func(tx *Client) error) error {
tx: tx,
}
defer func(tx pgx.Tx, ctx context.Context) {
err := tx.Rollback(ctx)
if err != nil {
}
}(tx, ctx)
// rollback 兜底(防 panic / 提前 return)
defer func() {
_ = tx.Rollback(ctx)
}()
// 执行业务
if err := fn(txClient); err != nil {
_ = tx.Rollback(ctx)
return err
}
return tx.Commit(ctx)
// commit
if err := tx.Commit(ctx); err != nil {
_ = tx.Rollback(ctx)
return err
}
return nil
}
+12 -11
View File
@@ -8,6 +8,7 @@ import (
)
func (c *Client) Update(ctx context.Context, table, keyField string, data map[string]any) error {
keyVal, err := validateData(table, keyField, data)
if err != nil {
return err
@@ -15,9 +16,6 @@ func (c *Client) Update(ctx context.Context, table, keyField string, data map[st
data = applyMetaFields(ctx, table, data, false)
tableSQL := quoteTable(table)
keySQL := quoteCol(keyField)
var (
set []string
args []any
@@ -38,17 +36,18 @@ func (c *Client) Update(ctx context.Context, table, keyField string, data map[st
sql := fmt.Sprintf(
"UPDATE %s SET %s WHERE %s=$%d",
tableSQL,
quoteTable(table),
strings.Join(set, ", "),
keySQL,
quoteCol(keyField),
i,
)
_, err = c.exec().Exec(ctx, sql, args...)
_, err = c.pool.Exec(ctx, sql, args...)
return err
}
func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list []map[string]any) error {
if len(list) == 0 {
return fmt.Errorf("empty data")
}
@@ -75,6 +74,7 @@ func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list [
sets []string
)
// CASE 构建
for _, col := range cols {
colSQL := quoteCol(col)
@@ -84,7 +84,7 @@ func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list [
for _, row := range list {
row = applyMetaFields(ctx, table, row, false)
keyVal, _ := row[keyField]
keyVal := row[keyField]
val := row[col]
caseSQL.WriteString(fmt.Sprintf(
@@ -101,9 +101,10 @@ func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list [
sets = append(sets, caseSQL.String())
}
var where []string
// ⭐修复 IN 写法(关键)
var inPlaceholders []string
for _, row := range list {
where = append(where, fmt.Sprintf("$%d", argIndex))
inPlaceholders = append(inPlaceholders, fmt.Sprintf("$%d", argIndex))
args = append(args, row[keyField])
argIndex++
}
@@ -113,9 +114,9 @@ func (c *Client) BatchUpdate(ctx context.Context, table, keyField string, list [
tableSQL,
strings.Join(sets, ", "),
keySQL,
strings.Join(where, ", "),
strings.Join(inPlaceholders, ", "),
)
_, err := c.exec().Exec(ctx, sql, args...)
_, err := c.pool.Exec(ctx, sql, args...) // ⭐修复点
return err
}