This commit is contained in:
oneao committed 2026-01-13 22:18:01 +08:00
1 parent 6277fb2f9e
commit 8ecac0d7dd
270 files changed
+40438 -529

No files matched your search

@@ -0,0 +1,29 @@
package errorx
import (
"database/sql"
"github.com/cloudwego/hertz/pkg/app"
"github.com/jackc/pgx/v5"
"github.com/pkg/errors"
)
// Wrap 包装 error,附加调用栈
func Wrap(err error) error {
if err == nil {
return nil
}
return errors.WithStack(err)
}
// AddError 简化版,自动 wrap 并注册到 ctx
func AddError(c *app.RequestContext, err error) {
if err == nil {
return
}
_ = c.Error(Wrap(err)) // 自动 wrap 并注册
}
// IsNotFound 判断 err 是否因为数据库中没找到数据
func IsNotFound(err error) bool {
return errors.Is(err, sql.ErrNoRows) || errors.Is(err, pgx.ErrNoRows)
}
@@ -0,0 +1,83 @@
package idgen
import (
"strconv"
"time"
)
type DefaultIdGenerator struct {
Options *IdGeneratorOptions
SnowWorker ISnowWorker
IdGeneratorException IdGeneratorException
}
func NewDefaultIdGenerator(options *IdGeneratorOptions) *DefaultIdGenerator {
if options == nil {
panic("dig.Options error.")
}
// 1.BaseTime
minTime := int64(631123200000) // time.Now().AddDate(-30, 0, 0).UnixNano() / 1e6
if options.BaseTime < minTime || options.BaseTime > time.Now().UnixNano()/1e6 {
panic("BaseTime error.")
}
// 2.WorkerIdBitLength
if options.WorkerIdBitLength <= 0 {
panic("WorkerIdBitLength error.(range:[1, 21])")
}
if options.WorkerIdBitLength+options.SeqBitLength > 22 {
panic("error:WorkerIdBitLength + SeqBitLength <= 22")
}
// 3.WorkerId
maxWorkerIdNumber := uint16(1<<options.WorkerIdBitLength) - 1
if maxWorkerIdNumber == 0 {
maxWorkerIdNumber = 63
}
if options.WorkerId < 0 || options.WorkerId > maxWorkerIdNumber {
panic("WorkerId error. (range:[0, " + strconv.FormatUint(uint64(maxWorkerIdNumber), 10) + "]")
}
// 4.SeqBitLength
if options.SeqBitLength < 2 || options.SeqBitLength > 21 {
panic("SeqBitLength error. (range:[2, 21])")
}
// 5.MaxSeqNumber
maxSeqNumber := uint32(1<<options.SeqBitLength) - 1
if maxSeqNumber == 0 {
maxSeqNumber = 63
}
if options.MaxSeqNumber < 0 || options.MaxSeqNumber > maxSeqNumber {
panic("MaxSeqNumber error. (range:[1, " + strconv.FormatUint(uint64(maxSeqNumber), 10) + "]")
}
// 6.MinSeqNumber
if options.MinSeqNumber < 5 || options.MinSeqNumber > maxSeqNumber {
panic("MinSeqNumber error. (range:[5, " + strconv.FormatUint(uint64(maxSeqNumber), 10) + "]")
}
var snowWorker ISnowWorker
switch options.Method {
case 1:
snowWorker = NewSnowWorkerM1(options)
case 2:
snowWorker = NewSnowWorkerM2(options)
default:
snowWorker = NewSnowWorkerM1(options)
}
if options.Method == 1 {
time.Sleep(time.Duration(500) * time.Microsecond)
}
return &DefaultIdGenerator{
Options: options,
SnowWorker: snowWorker,
}
}
func (dig DefaultIdGenerator) NewLong() int64 {
return dig.SnowWorker.NextId()
}
@@ -0,0 +1,5 @@
package idgen
type IIdGenerator interface {
NewLong() uint64
}
@@ -0,0 +1,5 @@
package idgen
type ISnowWorker interface {
NextId() int64
}
@@ -0,0 +1,12 @@
package idgen
import "fmt"
type IdGeneratorException struct {
message string
error error
}
func (e IdGeneratorException) IdGeneratorException(message ...interface{}) {
fmt.Println(message)
}
@@ -0,0 +1,25 @@
package idgen
type IdGeneratorOptions struct {
Method uint16 // 雪花计算方法,(1-漂移算法|2-传统算法),默认1
BaseTime int64 // 基础时间(ms单位),不能超过当前系统时间
WorkerId uint16 // 机器码,必须由外部设定,最大值 2^WorkerIdBitLength-1
WorkerIdBitLength byte // 机器码位长,默认值6,取值范围 [1, 15](要求:序列数位长+机器码位长不超过22)
SeqBitLength byte // 序列数位长,默认值6,取值范围 [3, 21](要求:序列数位长+机器码位长不超过22)
MaxSeqNumber uint32 // 最大序列数(含),设置范围 [MinSeqNumber, 2^SeqBitLength-1],默认值0,表示最大序列数取最大值(2^SeqBitLength-1])
MinSeqNumber uint32 // 最小序列数(含),默认值5,取值范围 [5, MaxSeqNumber],每毫秒的前5个序列数对应编号0-4是保留位,其中1-4是时间回拨相应预留位,0是手工新值预留位
TopOverCostCount uint32 // 最大漂移次数(含),默认2000,推荐范围500-10000(与计算能力有关)
}
func NewIdGeneratorOptions(workerId uint16) *IdGeneratorOptions {
return &IdGeneratorOptions{
Method: 1,
WorkerId: workerId,
BaseTime: 1582136402000,
WorkerIdBitLength: 6,
SeqBitLength: 6,
MaxSeqNumber: 0,
MinSeqNumber: 5,
TopOverCostCount: 2000,
}
}
@@ -0,0 +1,19 @@
package idgen
type OverCostActionArg struct {
ActionType int32
TimeTick int64
WorkerId uint16
OverCostCountInOneTerm int32
GenCountInOneTerm int32
TermIndex int32
}
func (ocaa OverCostActionArg) OverCostActionArg(workerId uint16, timeTick int64, actionType int32, overCostCountInOneTerm int32, genCountWhenOverCost int32, index int32) {
ocaa.ActionType = actionType
ocaa.TimeTick = timeTick
ocaa.WorkerId = workerId
ocaa.OverCostCountInOneTerm = overCostCountInOneTerm
ocaa.GenCountInOneTerm = genCountWhenOverCost
ocaa.TermIndex = index
}
@@ -0,0 +1,243 @@
package idgen
import (
"sync"
"time"
)
// SnowWorkerM1 .
type SnowWorkerM1 struct {
BaseTime int64 //基础时间
WorkerId uint16 //机器码
WorkerIdBitLength byte //机器码位长
SeqBitLength byte //自增序列数位长
MaxSeqNumber uint32 //最大序列数(含)
MinSeqNumber uint32 //最小序列数(含)
TopOverCostCount uint32 //最大漂移次数
_TimestampShift byte
_CurrentSeqNumber uint32
_LastTimeTick int64
_TurnBackTimeTick int64
_TurnBackIndex byte
_IsOverCost bool
_OverCostCountInOneTerm uint32
_GenCountInOneTerm uint32
_TermIndex uint32
sync.Mutex
}
// NewSnowWorkerM1 .
func NewSnowWorkerM1(options *IdGeneratorOptions) ISnowWorker {
var workerIdBitLength byte
var seqBitLength byte
var maxSeqNumber uint32
// 1.BaseTime
var baseTime int64
if options.BaseTime != 0 {
baseTime = options.BaseTime
} else {
baseTime = 1582136402000
}
// 2.WorkerIdBitLength
if options.WorkerIdBitLength == 0 {
workerIdBitLength = 6
} else {
workerIdBitLength = options.WorkerIdBitLength
}
// 3.WorkerId
var workerId = options.WorkerId
// 4.SeqBitLength
if options.SeqBitLength == 0 {
seqBitLength = 6
} else {
seqBitLength = options.SeqBitLength
}
// 5.MaxSeqNumber
if options.MaxSeqNumber <= 0 {
maxSeqNumber = (1 << seqBitLength) - 1
} else {
maxSeqNumber = options.MaxSeqNumber
}
// 6.MinSeqNumber
var minSeqNumber = options.MinSeqNumber
// 7.Others
var topOverCostCount = options.TopOverCostCount
if topOverCostCount == 0 {
topOverCostCount = 2000
}
timestampShift := (byte)(workerIdBitLength + seqBitLength)
currentSeqNumber := minSeqNumber
return &SnowWorkerM1{
BaseTime: baseTime,
WorkerIdBitLength: workerIdBitLength,
WorkerId: workerId,
SeqBitLength: seqBitLength,
MaxSeqNumber: maxSeqNumber,
MinSeqNumber: minSeqNumber,
TopOverCostCount: topOverCostCount,
_TimestampShift: timestampShift,
_CurrentSeqNumber: currentSeqNumber,
_LastTimeTick: 0,
_TurnBackTimeTick: 0,
_TurnBackIndex: 0,
_IsOverCost: false,
_OverCostCountInOneTerm: 0,
_GenCountInOneTerm: 0,
_TermIndex: 0,
}
}
// DoGenIDAction .
func (m1 *SnowWorkerM1) DoGenIdAction(arg *OverCostActionArg) {
}
func (m1 *SnowWorkerM1) BeginOverCostAction(useTimeTick int64) {
}
func (m1 *SnowWorkerM1) EndOverCostAction(useTimeTick int64) {
if m1._TermIndex > 10000 {
m1._TermIndex = 0
}
}
func (m1 *SnowWorkerM1) BeginTurnBackAction(useTimeTick int64) {
}
func (m1 *SnowWorkerM1) EndTurnBackAction(useTimeTick int64) {
}
func (m1 *SnowWorkerM1) NextOverCostId() int64 {
currentTimeTick := m1.GetCurrentTimeTick()
if currentTimeTick > m1._LastTimeTick {
m1.EndOverCostAction(currentTimeTick)
m1._LastTimeTick = currentTimeTick
m1._CurrentSeqNumber = m1.MinSeqNumber
m1._IsOverCost = false
m1._OverCostCountInOneTerm = 0
m1._GenCountInOneTerm = 0
return m1.CalcId(m1._LastTimeTick)
}
if m1._OverCostCountInOneTerm >= m1.TopOverCostCount {
m1.EndOverCostAction(currentTimeTick)
m1._LastTimeTick = m1.GetNextTimeTick()
m1._CurrentSeqNumber = m1.MinSeqNumber
m1._IsOverCost = false
m1._OverCostCountInOneTerm = 0
m1._GenCountInOneTerm = 0
return m1.CalcId(m1._LastTimeTick)
}
if m1._CurrentSeqNumber > m1.MaxSeqNumber {
m1._LastTimeTick++
m1._CurrentSeqNumber = m1.MinSeqNumber
m1._IsOverCost = true
m1._OverCostCountInOneTerm++
m1._GenCountInOneTerm++
return m1.CalcId(m1._LastTimeTick)
}
m1._GenCountInOneTerm++
return m1.CalcId(m1._LastTimeTick)
}
// NextNormalID .
func (m1 *SnowWorkerM1) NextNormalId() int64 {
currentTimeTick := m1.GetCurrentTimeTick()
if currentTimeTick < m1._LastTimeTick {
if m1._TurnBackTimeTick < 1 {
m1._TurnBackTimeTick = m1._LastTimeTick - 1
m1._TurnBackIndex++
// 每毫秒序列数的前5位是预留位,0用于手工新值,1-4是时间回拨次序
// 最多4次回拨(防止回拨重叠)
if m1._TurnBackIndex > 4 {
m1._TurnBackIndex = 1
}
m1.BeginTurnBackAction(m1._TurnBackTimeTick)
}
// time.Sleep(time.Duration(1) * time.Millisecond)
return m1.CalcTurnBackId(m1._TurnBackTimeTick)
}
// 时间追平时,_TurnBackTimeTick清零
if m1._TurnBackTimeTick > 0 {
m1.EndTurnBackAction(m1._TurnBackTimeTick)
m1._TurnBackTimeTick = 0
}
if currentTimeTick > m1._LastTimeTick {
m1._LastTimeTick = currentTimeTick
m1._CurrentSeqNumber = m1.MinSeqNumber
return m1.CalcId(m1._LastTimeTick)
}
if m1._CurrentSeqNumber > m1.MaxSeqNumber {
m1.BeginOverCostAction(currentTimeTick)
m1._TermIndex++
m1._LastTimeTick++
m1._CurrentSeqNumber = m1.MinSeqNumber
m1._IsOverCost = true
m1._OverCostCountInOneTerm = 1
m1._GenCountInOneTerm = 1
return m1.CalcId(m1._LastTimeTick)
}
return m1.CalcId(m1._LastTimeTick)
}
// CalcID .
func (m1 *SnowWorkerM1) CalcId(useTimeTick int64) int64 {
result := int64(useTimeTick<<m1._TimestampShift) + int64(m1.WorkerId<<m1.SeqBitLength) + int64(m1._CurrentSeqNumber)
m1._CurrentSeqNumber++
return result
}
// CalcTurnBackID .
func (m1 *SnowWorkerM1) CalcTurnBackId(useTimeTick int64) int64 {
result := int64(useTimeTick<<m1._TimestampShift) + int64(m1.WorkerId<<m1.SeqBitLength) + int64(m1._TurnBackIndex)
m1._TurnBackTimeTick--
return result
}
// GetCurrentTimeTick .
func (m1 *SnowWorkerM1) GetCurrentTimeTick() int64 {
var millis = time.Now().UnixNano() / 1e6
return millis - m1.BaseTime
}
// GetNextTimeTick .
func (m1 *SnowWorkerM1) GetNextTimeTick() int64 {
tempTimeTicker := m1.GetCurrentTimeTick()
for tempTimeTicker <= m1._LastTimeTick {
tempTimeTicker = m1.GetCurrentTimeTick()
}
return tempTimeTicker
}
// NextId .
func (m1 *SnowWorkerM1) NextId() int64 {
m1.Lock()
defer m1.Unlock()
if m1._IsOverCost {
return m1.NextOverCostId()
} else {
return m1.NextNormalId()
}
}
@@ -0,0 +1,37 @@
package idgen
import (
"fmt"
"strconv"
)
type SnowWorkerM2 struct {
*SnowWorkerM1
}
func NewSnowWorkerM2(options *IdGeneratorOptions) ISnowWorker {
return &SnowWorkerM2{
NewSnowWorkerM1(options).(*SnowWorkerM1),
}
}
func (m2 SnowWorkerM2) NextId() int64 {
m2.Lock()
defer m2.Unlock()
currentTimeTick := m2.GetCurrentTimeTick()
if m2._LastTimeTick == currentTimeTick {
m2._CurrentSeqNumber++
if m2._CurrentSeqNumber > m2.MaxSeqNumber {
m2._CurrentSeqNumber = m2.MinSeqNumber
currentTimeTick = m2.GetNextTimeTick()
}
} else {
m2._CurrentSeqNumber = m2.MinSeqNumber
}
if currentTimeTick < m2._LastTimeTick {
fmt.Println("Time error for {0} milliseconds", strconv.FormatInt(m2._LastTimeTick-currentTimeTick, 10))
}
m2._LastTimeTick = currentTimeTick
result := int64(currentTimeTick<<m2._TimestampShift) + int64(m2.WorkerId<<m2.SeqBitLength) + int64(m2._CurrentSeqNumber)
return result
}
@@ -0,0 +1,29 @@
package idgen
import (
"sync"
)
var singletonMutex sync.Mutex
var idGenerator *DefaultIdGenerator
// SetIdGenerator .
func SetIdGenerator(options *IdGeneratorOptions) {
singletonMutex.Lock()
idGenerator = NewDefaultIdGenerator(options)
singletonMutex.Unlock()
}
// NextId .
func NextId() int64 {
if idGenerator == nil {
singletonMutex.Lock()
defer singletonMutex.Unlock()
if idGenerator == nil {
options := NewIdGeneratorOptions(1)
idGenerator = NewDefaultIdGenerator(options)
}
}
return idGenerator.NewLong()
}
+22
View File
@@ -0,0 +1,22 @@
package utils
import (
"crypto/rand"
"math/big"
)
// GenerateCode 生成指定长度的随机邀请码(大写字母+数字)
func GenerateCode(length int) (string, error) {
const charset = "ABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
code := make([]byte, length)
for i := 0; i < length; i++ {
num, err := rand.Int(rand.Reader, big.NewInt(int64(len(charset))))
if err != nil {
return "", err
}
code[i] = charset[num.Int64()]
}
return string(code), nil
}
+107
View File
@@ -0,0 +1,107 @@
package jwtx
import (
"allapp/conf"
"sync"
"time"
"github.com/golang-jwt/jwt/v4"
)
type CustomClaims struct {
UserID int64 `json:"user_id"`
jwt.RegisteredClaims
}
type TokenResult struct {
Claims *CustomClaims
IsValid bool
IsExpired bool
}
type jwtManager struct {
secret string
accessExpiry time.Duration
refreshExpiry time.Duration
}
var (
manager *jwtManager
once sync.Once
)
// InitJwt 初始化全局 JWT 配置,只能调用一次
func InitJwt() {
once.Do(func() {
manager = &jwtManager{
secret: conf.GetConf().JWT.Secret,
accessExpiry: conf.GetConf().JWT.AccessExpiry,
refreshExpiry: conf.GetConf().JWT.RefreshExpiry,
}
})
}
// CreateAccessToken 全局函数创建 Access Token
func CreateAccessToken(userID int64) (string, error) {
return manager.createToken(userID, manager.accessExpiry)
}
// CreateRefreshToken 全局函数创建 Refresh Token
func CreateRefreshToken(userID int64) (string, error) {
return manager.createToken(userID, manager.refreshExpiry)
}
// VerifyToken 全局函数验证 Token
func VerifyToken(tokenString string) TokenResult {
result := TokenResult{IsValid: false, IsExpired: false}
if tokenString == "" || manager == nil {
return result
}
token, err := jwt.ParseWithClaims(tokenString, &CustomClaims{}, func(token *jwt.Token) (interface{}, error) {
return []byte(manager.secret), nil
})
if token == nil {
return result
}
claims, ok := token.Claims.(*CustomClaims)
if !ok {
return result
}
if err != nil || !token.Valid {
return result
}
result.Claims = claims
result.IsValid = true
if claims.ExpiresAt != nil && time.Now().After(claims.ExpiresAt.Time) {
result.IsExpired = true
}
return result
}
// 内部方法,生成 Token
func (j *jwtManager) createToken(userID int64, expiry time.Duration) (string, error) {
claims := &CustomClaims{
UserID: userID,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(time.Now()),
},
}
switch {
case expiry > 0:
claims.ExpiresAt = jwt.NewNumericDate(time.Now().Add(expiry))
case expiry == -1:
claims.ExpiresAt = nil // 永不过期
default:
claims.ExpiresAt = jwt.NewNumericDate(time.Now().Add(time.Hour))
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return token.SignedString([]byte(j.secret))
}
+19
View File
@@ -0,0 +1,19 @@
package pathx
import (
"allapp/conf"
"strings"
)
// MatchPath 判断给定路径是否在白名单中
func MatchPath(path string, whitelist []string) bool {
baseUrl := conf.GetConf().Server.BaseUrl
for _, p := range whitelist {
fullPath := baseUrl + p
if strings.HasPrefix(path, fullPath) {
return true
}
}
return false
}
@@ -0,0 +1,125 @@
package pgtypex
import (
"github.com/jackc/pgx/v5/pgtype"
"github.com/shopspring/decimal"
"strings"
"time"
)
func NumericToString(n pgtype.Numeric) string {
if n.Int == nil || n.NaN {
return "0.00"
}
// 转成 decimal 保证精度
dec := decimal.NewFromBigInt(n.Int, int32(n.Exp))
return dec.StringFixed(2) // 保留两位小数
}
func NumericToDecimal(n pgtype.Numeric) decimal.Decimal {
if n.Int == nil || n.NaN {
return decimal.Zero
}
return decimal.NewFromBigInt(n.Int, int32(n.Exp))
}
// TimestampToDateString 转为 "YYYY-MM-DD"
func TimestampToDateString(ts pgtype.Timestamp) string {
if ts.Time.IsZero() {
return ""
}
return ts.Time.Format("2006-01-02")
}
// TimestampToDateTimeString 转为 "YYYY-MM-DD HH:MM:SS"
func TimestampToDateTimeString(ts pgtype.Timestamp) string {
if ts.Time.IsZero() {
return ""
}
return ts.Time.Format("2006-01-02 15:04:05")
}
// StringToNumeric 将字符串金额转换为 pgtype.Numeric
func StringToNumeric(amountStr string) pgtype.Numeric {
if strings.TrimSpace(amountStr) == "" {
return pgtype.Numeric{Valid: false}
}
dec, err := decimal.NewFromString(amountStr)
if err != nil {
return pgtype.Numeric{Valid: false}
}
// 保留两位小数,四舍五入
dec = dec.Round(2)
// 乘以 100 转整数
intVal := dec.Mul(decimal.NewFromInt(100)).BigInt()
return pgtype.Numeric{
Int: intVal,
Exp: -2, // 固定两位小数
NaN: false,
Valid: true,
}
}
// StringToText 将字符串转换为 pgtype.Text
// 空字符串返回零值 Text
func StringToText(s string) pgtype.Text {
if s == "" {
return pgtype.Text{Valid: false}
}
return pgtype.Text{
String: s,
Valid: true,
}
}
// StringToTimestamp 将字符串解析为 pgtype.Timestamp
// 支持格式:YYYY-MM-DD HH:MM:SS 或 YYYY-MM-DD
// 解析失败或空字符串返回零值 Timestamp
func StringToTimestamp(s string) pgtype.Timestamp {
if s == "" {
return pgtype.Timestamp{Valid: false}
}
var t time.Time
var err error
layouts := []string{
"2006-01-02 15:04:05", // 完整时间
"2006-01-02", // 仅日期也可解析
}
for _, layout := range layouts {
t, err = time.Parse(layout, s)
if err == nil {
return pgtype.Timestamp{
Time: t,
Valid: true,
}
}
}
return pgtype.Timestamp{Valid: false}
}
// StringToDate 将字符串解析为 pgtype.Date
// 支持格式:YYYY-MM-DD
// 解析失败或空字符串返回零值 Date
func StringToDate(s string) pgtype.Date {
if s == "" {
return pgtype.Date{Valid: false}
}
t, err := time.Parse("2006-01-02", s)
if err != nil {
return pgtype.Date{Valid: false}
}
return pgtype.Date{
Time: t,
Valid: true,
}
}
@@ -0,0 +1,86 @@
package response
import (
"net/http"
"github.com/cloudwego/hertz/pkg/app"
"github.com/hertz-contrib/requestid"
)
// Result 响应结构
type Result struct {
Code int `json:"code"`
Message string `json:"message"`
Data interface{} `json:"data"`
TrackId string `json:"trackId,omitempty"`
}
// HttpCode 响应状态码
var HttpCode = struct {
Success int
Unauthorized int
Fail int
RefreshToken int
}{
Success: 200,
Unauthorized: 401,
Fail: 500,
RefreshToken: 402,
}
// Builder 响应构建器
type Builder struct {
c *app.RequestContext
result Result
statusCode int // 可自定义 HTTP 状态码
}
// Success 构造成功响应
func Success(c *app.RequestContext) *Builder {
return &Builder{
c: c,
result: Result{Code: HttpCode.Success, Message: "请求成功"},
statusCode: http.StatusOK,
}
}
// Fail 构造失败响应
func Fail(c *app.RequestContext) *Builder {
return &Builder{
c: c,
result: Result{Code: HttpCode.Fail, Message: "请求失败"},
statusCode: http.StatusOK,
}
}
// Status 设置自定义 HTTP 状态码
func (b *Builder) Status(code int) *Builder {
b.statusCode = code
return b
}
func (b *Builder) Code(code int) *Builder {
b.result.Code = code
return b
}
func (b *Builder) Message(msg string) *Builder {
b.result.Message = msg
return b
}
func (b *Builder) Data(data interface{}) *Builder {
b.result.Data = data
return b
}
// Send 发送响应
func (b *Builder) Send() {
//if b.result.Data == nil {
// b.result.Data = ""
//}
reqID := requestid.Get(b.c)
b.result.TrackId = reqID
b.c.JSON(200, b.result)
}
@@ -0,0 +1,73 @@
package routinex
import (
"runtime"
"strconv"
"strings"
"sync"
)
// CurGID 获取 routine 的id
func CurGID() uint64 {
var buf [64]byte
n := runtime.Stack(buf[:], false)
line := strings.Fields(strings.TrimPrefix(string(buf[:n]), "goroutine "))[0]
gid, _ := strconv.ParseUint(line, 10, 64)
return gid
}
type RoutineLocal struct {
data sync.Map // gid -> map[key]any
}
var (
instance *RoutineLocal
instanceOnce sync.Once
)
// GetInstance 获取当前示例
func GetInstance() *RoutineLocal {
instanceOnce.Do(func() {
instance = &RoutineLocal{}
})
return instance
}
// Set 存储单条数据
func Set(key string, value any) {
gid := CurGID()
v, _ := GetInstance().data.LoadOrStore(gid, &sync.Map{})
m := v.(*sync.Map)
m.Store(key, value)
}
// Get 获取单条数据
func Get(key string) any {
gid := CurGID()
if v, ok := GetInstance().data.Load(gid); ok {
m := v.(*sync.Map)
if val, ok := m.Load(key); ok {
return val
}
}
return nil
}
// Clear 删除当前 goroutine 所有数据
func Clear() {
GetInstance().data.Delete(CurGID())
}
// Go 自动继承父 goroutine 数据,多 routine 的时候需要使用
func Go(f func()) {
parentGID := CurGID()
parentData, _ := GetInstance().data.Load(parentGID)
go func() {
if parentData != nil {
GetInstance().data.Store(CurGID(), parentData)
defer Clear()
}
f()
}()
}