This commit is contained in:
oneao committed 2025-08-25 17:22:54 +08:00
1 parent 76a6c81af2
commit b19d083031
68 files changed
+3889 -343

No files matched your search

@@ -0,0 +1,68 @@
package db
import (
"base-framework/pkg/config"
"base-framework/pkg/router"
"database/sql"
"errors"
"fmt"
)
type Client struct {
conn *sql.DB
tx *sql.Tx
}
// JdbcTemplate 创建 JdbcTemplate
func JdbcTemplate(c *router.Context) (*Client, error) {
val, ok := c.Get("orgId")
if !ok {
return nil, errors.New("missing orgId in context")
}
orgId, ok := val.(string)
if !ok || orgId == "" {
return nil, errors.New("invalid orgId in context")
}
conn, ok := config.GetDB(orgId)
if !ok || conn == nil {
return nil, fmt.Errorf("no db connection found for orgId=%s", orgId)
}
return &Client{conn: conn}, nil
}
// ------------------------ 内部方法 ------------------------
// 执行查询,返回 *sql.Rows
func (c *Client) query(query string, args ...any) (*sql.Rows, error) {
if c.tx != nil {
return c.tx.Query(query, args...)
}
return c.conn.Query(query, args...)
}
// 执行执行类语句(insert/update/delete)
func (c *Client) exec(query string, args ...any) (sql.Result, error) {
if c.tx != nil {
return c.tx.Exec(query, args...)
}
return c.conn.Exec(query, args...)
}
// WithTransaction 自动处理事务提交或回滚
func (c *Client) WithTransaction(fn func(txClient *Client) error) error {
if c.tx != nil {
// 已经在事务中,直接执行
return fn(c)
}
tx, err := c.conn.Begin()
if err != nil {
return err
}
txClient := &Client{conn: c.conn, tx: tx}
if err := fn(txClient); err != nil {
_ = tx.Rollback()
return err
}
return tx.Commit()
}
@@ -0,0 +1,5 @@
package db
import "database/sql"
func (c *Client) Delete(query string, args ...any) (sql.Result, error) { return c.exec(query, args...) }
@@ -0,0 +1,15 @@
package db
import "database/sql"
func (c *Client) Insert(query string, args ...any) (sql.Result, error) { return c.exec(query, args...) }
// BatchInsert 批量插入,传入多组参数
func (c *Client) BatchInsert(query string, params [][]any) error {
for _, args := range params {
if _, err := c.Insert(query, args...); err != nil {
return err
}
}
return nil
}
@@ -0,0 +1,63 @@
package db
import "database/sql"
func (c *Client) Select(query string, args ...any) ([]map[string]any, error) {
rows, err := c.query(query, args...)
if err != nil {
return nil, err
}
defer func(rows *sql.Rows) {
_ = rows.Close()
}(rows)
cols, _ := rows.Columns()
var result []map[string]any
for rows.Next() {
columns := make([]any, len(cols))
columnPointers := make([]any, len(cols))
for i := range columns {
columnPointers[i] = &columns[i]
}
if err := rows.Scan(columnPointers...); err != nil {
return nil, err
}
rowMap := make(map[string]any)
for i, colName := range cols {
rowMap[colName] = columns[i]
}
result = append(result, rowMap)
}
return result, nil
}
func (c *Client) SelectOne(query string, args ...any) (map[string]any, error) {
rows, err := c.query(query, args...)
if err != nil {
return nil, err
}
defer func(rows *sql.Rows) {
_ = rows.Close()
}(rows)
if !rows.Next() {
return nil, sql.ErrNoRows
}
cols, _ := rows.Columns()
columns := make([]any, len(cols))
columnPointers := make([]any, len(cols))
for i := range columns {
columnPointers[i] = &columns[i]
}
if err := rows.Scan(columnPointers...); err != nil {
return nil, err
}
rowMap := make(map[string]any)
for i, colName := range cols {
rowMap[colName] = columns[i]
}
return rowMap, nil
}
@@ -0,0 +1,5 @@
package db
import "database/sql"
func (c *Client) Update(query string, args ...any) (sql.Result, error) { return c.exec(query, args...) }
@@ -0,0 +1,70 @@
package jwt
import (
commonConfig "base-framework/pkg/config"
"errors"
"fmt"
"time"
"github.com/golang-jwt/jwt"
)
var (
ErrTokenExpired = errors.New("token expired")
ErrTokenInvalid = errors.New("token invalid")
)
// CustomClaims 定义自己的 payload 结构
type CustomClaims struct {
OrgID string `json:"org_id"`
UserID string `json:"user_id"`
jwt.StandardClaims
}
// CreateToken 创建一个 JWT token
func CreateToken(orgID, userID string) (string, error) {
expireAt := time.Now().Add(commonConfig.JWT.Expiry)
fmt.Println("Token 过期时间:", expireAt.Format(time.RFC3339))
claims := CustomClaims{
OrgID: orgID,
UserID: userID,
StandardClaims: jwt.StandardClaims{
ExpiresAt: expireAt.Unix(),
IssuedAt: time.Now().Unix(),
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return token.SignedString([]byte(commonConfig.JWT.Secret))
}
// VerifyToken 验证并解析 JWT token
func VerifyToken(tokenString string) (*CustomClaims, error) {
signKey := []byte(commonConfig.JWT.Secret)
token, err := jwt.ParseWithClaims(tokenString, &CustomClaims{}, func(token *jwt.Token) (interface{}, error) {
// 校验签名算法
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, ErrTokenInvalid
}
return signKey, nil
})
if err != nil {
// 判断是否过期错误
var ve *jwt.ValidationError
if errors.As(err, &ve) {
if ve.Errors&jwt.ValidationErrorExpired != 0 {
return nil, ErrTokenExpired
}
}
return nil, ErrTokenInvalid
}
if claims, ok := token.Claims.(*CustomClaims); ok && token.Valid {
return claims, nil
}
return nil, ErrTokenInvalid
}
@@ -0,0 +1,53 @@
package strutil
import (
"fmt"
"strconv"
)
// ToString 通用高效转换为 string
func ToString(v any) string {
switch n := v.(type) {
// 整数类型
case int:
return strconv.FormatInt(int64(n), 10)
case int8:
return strconv.FormatInt(int64(n), 10)
case int16:
return strconv.FormatInt(int64(n), 10)
case int32:
return strconv.FormatInt(int64(n), 10)
case int64:
return strconv.FormatInt(n, 10)
// 无符号整数
case uint:
return strconv.FormatUint(uint64(n), 10)
case uint8:
return strconv.FormatUint(uint64(n), 10)
case uint16:
return strconv.FormatUint(uint64(n), 10)
case uint32:
return strconv.FormatUint(uint64(n), 10)
case uint64:
return strconv.FormatUint(n, 10)
// 浮点数
case float32:
return strconv.FormatFloat(float64(n), 'f', -1, 32)
case float64:
return strconv.FormatFloat(n, 'f', -1, 64)
// 字符串
case string:
return n
// 布尔
case bool:
return strconv.FormatBool(n)
default:
// fallback
return fmt.Sprintf("%v", v)
}
}
@@ -0,0 +1,90 @@
package uid
import (
"base-framework/pkg/config"
"sync"
"sync/atomic"
"time"
)
const (
datacenterBits = uint64(5)
workerBits = uint64(5)
sequenceBits = uint64(12)
maxSequence = -1 ^ (-1 << sequenceBits)
workerShift = sequenceBits
datacenterShift = sequenceBits + workerBits
timestampShift = sequenceBits + workerBits + datacenterBits
epoch = int64(1672531200000) // 2023-01-01
)
type Snowflake struct {
lastStamp uint64 // 上一次时间戳 (毫秒)
sequence uint64 // 序列号
dcID uint64
workerID uint64
}
var (
sfInstance *Snowflake
once sync.Once
)
// lazy 初始化实例
func initInstance() {
var dcID, wID uint64
if config.Snowflake != nil {
dcID = config.Snowflake.DatacenterID
wID = config.Snowflake.WorkerID
}
sfInstance = &Snowflake{
dcID: dcID,
workerID: wID,
}
}
// NextID 生成全局唯一 ID,无锁版本
func NextID() uint64 {
once.Do(initInstance)
for {
now := uint64(time.Now().UnixNano() / 1e6)
last := atomic.LoadUint64(&sfInstance.lastStamp)
seq := atomic.LoadUint64(&sfInstance.sequence)
if now < last {
// 系统时钟回拨,等待
now = waitNextMillis(last)
}
if now == last {
seq = (seq + 1) & uint64(maxSequence)
if seq == 0 {
now = waitNextMillis(last)
}
} else {
seq = 0
}
if atomic.CompareAndSwapUint64(&sfInstance.lastStamp, last, now) {
atomic.StoreUint64(&sfInstance.sequence, seq)
id := ((now - uint64(epoch)) << timestampShift) |
(sfInstance.dcID << datacenterShift) |
(sfInstance.workerID << workerShift) |
seq
return id
}
}
}
// 等待下一个毫秒
func waitNextMillis(last uint64) uint64 {
now := uint64(time.Now().UnixNano() / 1e6)
for now <= last {
now = uint64(time.Now().UnixNano() / 1e6)
}
return now
}