package service import ( "context" "sort" "strconv" "strings" "time" "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/setting/operation_setting" ) const ( userConcurrencyKeyPrefix = "user:concurrency:" channelConcurrencyKeyPrefix = "channel:concurrency:" channelWaitingKeyPrefix = "channel:concurrency_waiting:" userRPMKeyPrefix = "user:rpm:" concurrencyTTL = 5 * time.Minute rpmTTL = 2 * time.Minute waitingTTL = 2 * time.Minute ) type OpsConcurrencySnapshot struct { RedisEnabled bool `json:"redis_enabled"` Users []OpsUserConcurrencySnapshot `json:"users"` Channels []OpsChannelConcurrencySnapshot `json:"channels"` GeneratedAt int64 `json:"generated_at"` } type OpsUserConcurrencySnapshot struct { UserID int `json:"user_id"` CurrentConcurrency int `json:"current_concurrency"` CurrentRPM int `json:"current_rpm"` } type OpsChannelConcurrencySnapshot struct { ChannelID int `json:"channel_id"` CurrentConcurrency int `json:"current_concurrency"` } func AcquireUserConcurrency(userId, limit int) (bool, func()) { return acquireConcurrency(userConcurrencyKeyPrefix+strconv.Itoa(userId), limit) } func AcquireChannelConcurrency(channelId, limit int) (bool, func()) { return acquireConcurrency(channelConcurrencyKeyPrefix+strconv.Itoa(channelId), limit) } func GetChannelConcurrencyWaitPlan() operation_setting.ChannelConcurrencyWaitPlanSetting { return operation_setting.GetChannelConcurrencyWaitPlanSetting() } func TryEnterChannelConcurrencyWaitQueue(channelId, maxWaiting int) (bool, func()) { release := func() {} if channelId <= 0 || maxWaiting <= 0 { return false, release } if !redisAvailable() { return true, release } ctx := context.Background() key := channelWaitingKeyPrefix + strconv.Itoa(channelId) count, err := common.RDB.Incr(ctx, key).Result() if err != nil { return true, release } release = func() { if redisAvailable() { _ = common.RDB.Decr(context.Background(), key).Err() } } if err = common.RDB.Expire(ctx, key, waitingTTL).Err(); err != nil { return true, release } if count > int64(maxWaiting) { release() return false, func() {} } return true, release } func WaitForChannelConcurrencySlot(ctx context.Context, channelId, limit int, timeout time.Duration) (bool, func(), error) { release := func() {} if timeout <= 0 { return false, release, context.DeadlineExceeded } if ctx == nil { ctx = context.Background() } waitCtx, cancel := context.WithTimeout(ctx, timeout) defer cancel() ticker := time.NewTicker(100 * time.Millisecond) defer ticker.Stop() for { allowed, release := AcquireChannelConcurrency(channelId, limit) if allowed { return true, release, nil } select { case <-waitCtx.Done(): return false, func() {}, waitCtx.Err() case <-ticker.C: } } } func CheckUserRPM(userId, limit int) bool { if limit <= 0 || !redisAvailable() { return true } key := userRPMKeyPrefix + strconv.Itoa(userId) + ":" + time.Now().Format("200601021504") ctx := context.Background() count, err := common.RDB.Incr(ctx, key).Result() if err != nil { return true } if err = common.RDB.Expire(ctx, key, rpmTTL).Err(); err != nil { return true } return count <= int64(limit) } func GetOpsConcurrencySnapshot() OpsConcurrencySnapshot { now := time.Now() snapshot := OpsConcurrencySnapshot{ RedisEnabled: redisAvailable(), Users: []OpsUserConcurrencySnapshot{}, Channels: []OpsChannelConcurrencySnapshot{}, GeneratedAt: now.Unix(), } if !snapshot.RedisEnabled { return snapshot } ctx := context.Background() userConcurrency, ok := scanIntValues(ctx, userConcurrencyKeyPrefix+"*", parseUserConcurrencyKey) if !ok { snapshot.RedisEnabled = false return emptyOpsConcurrencySnapshot(snapshot) } currentRPMMinute := now.Format("200601021504") userRPM, ok := scanIntValues(ctx, userRPMKeyPrefix+"*:"+currentRPMMinute, func(key string) (int, bool) { return parseUserRPMKey(key, currentRPMMinute) }) if !ok { snapshot.RedisEnabled = false return emptyOpsConcurrencySnapshot(snapshot) } channelConcurrency, ok := scanIntValues(ctx, channelConcurrencyKeyPrefix+"*", parseChannelConcurrencyKey) if !ok { snapshot.RedisEnabled = false return emptyOpsConcurrencySnapshot(snapshot) } for userID, concurrency := range userConcurrency { snapshot.Users = append(snapshot.Users, OpsUserConcurrencySnapshot{ UserID: userID, CurrentConcurrency: concurrency, CurrentRPM: userRPM[userID], }) delete(userRPM, userID) } for userID, rpm := range userRPM { snapshot.Users = append(snapshot.Users, OpsUserConcurrencySnapshot{ UserID: userID, CurrentRPM: rpm, }) } for channelID, concurrency := range channelConcurrency { snapshot.Channels = append(snapshot.Channels, OpsChannelConcurrencySnapshot{ ChannelID: channelID, CurrentConcurrency: concurrency, }) } sort.Slice(snapshot.Users, func(i, j int) bool { return snapshot.Users[i].UserID < snapshot.Users[j].UserID }) sort.Slice(snapshot.Channels, func(i, j int) bool { return snapshot.Channels[i].ChannelID < snapshot.Channels[j].ChannelID }) return snapshot } func acquireConcurrency(key string, limit int) (bool, func()) { release := func() {} if limit <= 0 || !redisAvailable() { return true, release } ctx := context.Background() count, err := common.RDB.Incr(ctx, key).Result() if err != nil { return true, release } release = func() { if redisAvailable() { _ = common.RDB.Decr(context.Background(), key).Err() } } if err = common.RDB.Expire(ctx, key, concurrencyTTL).Err(); err != nil { return true, release } if count > int64(limit) { release() return false, func() {} } return true, release } func redisAvailable() bool { return common.RedisEnabled && common.RDB != nil } func emptyOpsConcurrencySnapshot(snapshot OpsConcurrencySnapshot) OpsConcurrencySnapshot { snapshot.Users = []OpsUserConcurrencySnapshot{} snapshot.Channels = []OpsChannelConcurrencySnapshot{} return snapshot } func scanIntValues(ctx context.Context, pattern string, parseKey func(string) (int, bool)) (map[int]int, bool) { values := make(map[int]int) var cursor uint64 for { keys, nextCursor, err := common.RDB.Scan(ctx, cursor, pattern, 100).Result() if err != nil { return nil, false } for _, key := range keys { id, ok := parseKey(key) if !ok { continue } value, err := common.RDB.Get(ctx, key).Int() if err != nil { continue } values[id] = value } if nextCursor == 0 { break } cursor = nextCursor } return values, true } func parseUserConcurrencyKey(key string) (int, bool) { if !strings.HasPrefix(key, userConcurrencyKeyPrefix) { return 0, false } return parsePositiveID(strings.TrimPrefix(key, userConcurrencyKeyPrefix)) } func parseChannelConcurrencyKey(key string) (int, bool) { if !strings.HasPrefix(key, channelConcurrencyKeyPrefix) { return 0, false } return parsePositiveID(strings.TrimPrefix(key, channelConcurrencyKeyPrefix)) } func parseUserRPMKey(key, minute string) (int, bool) { prefix := userRPMKeyPrefix suffix := ":" + minute if !strings.HasPrefix(key, prefix) || !strings.HasSuffix(key, suffix) { return 0, false } return parsePositiveID(strings.TrimSuffix(strings.TrimPrefix(key, prefix), suffix)) } func parsePositiveID(value string) (int, bool) { if value == "" { return 0, false } id, err := strconv.Atoi(value) if err != nil || id <= 0 { return 0, false } return id, true }