u
This commit is contained in:
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)
|
||||
}
|
||||
@@ -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...)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in new issue
Block a user