fix(concurrency): wait when channel slots are saturated
Add a channel concurrency wait plan so saturated account/channel selection can distinguish no-wait, queue-full, and wait-for-slot outcomes while preserving fallback to other available channels first.
This commit is contained in:
parent
7a2fc8683c
commit
fa77659fe8
@ -1,6 +1,7 @@
|
|||||||
package middleware
|
package middleware
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
@ -16,6 +17,7 @@ import (
|
|||||||
"github.com/QuantumNous/new-api/model"
|
"github.com/QuantumNous/new-api/model"
|
||||||
relayconstant "github.com/QuantumNous/new-api/relay/constant"
|
relayconstant "github.com/QuantumNous/new-api/relay/constant"
|
||||||
"github.com/QuantumNous/new-api/service"
|
"github.com/QuantumNous/new-api/service"
|
||||||
|
"github.com/QuantumNous/new-api/setting/operation_setting"
|
||||||
"github.com/QuantumNous/new-api/setting/ratio_setting"
|
"github.com/QuantumNous/new-api/setting/ratio_setting"
|
||||||
"github.com/QuantumNous/new-api/types"
|
"github.com/QuantumNous/new-api/types"
|
||||||
|
|
||||||
@ -27,6 +29,21 @@ type ModelRequest struct {
|
|||||||
Group string `json:"group,omitempty"`
|
Group string `json:"group,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type distributorConcurrencyResponse struct {
|
||||||
|
statusCode int
|
||||||
|
message string
|
||||||
|
code types.ErrorCode
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
checkChannelQuotaForDistribute = model.CheckChannelQuota
|
||||||
|
acquireChannelConcurrencyForDistribute = service.AcquireChannelConcurrency
|
||||||
|
cacheGetRandomSatisfiedChannelForDistribute = service.CacheGetRandomSatisfiedChannel
|
||||||
|
getChannelConcurrencyWaitPlanForDistribute = service.GetChannelConcurrencyWaitPlan
|
||||||
|
tryEnterChannelConcurrencyWaitQueueForDistribute = service.TryEnterChannelConcurrencyWaitQueue
|
||||||
|
waitForChannelConcurrencySlotForDistribute = service.WaitForChannelConcurrencySlot
|
||||||
|
)
|
||||||
|
|
||||||
func Distribute() func(c *gin.Context) {
|
func Distribute() func(c *gin.Context) {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
var channel *model.Channel
|
var channel *model.Channel
|
||||||
@ -155,22 +172,28 @@ func Distribute() func(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
common.SetContextKey(c, constant.ContextKeyRequestStartTime, time.Now())
|
|
||||||
SetupContextForSelectedChannel(c, channel, modelRequest.Model)
|
|
||||||
if channel != nil {
|
if channel != nil {
|
||||||
if err := model.CheckChannelQuota(channel.Id); err != nil {
|
selectedChannel, releaseChannelConcurrency, response := acquireDistributedChannelConcurrency(
|
||||||
abortWithOpenAiMessage(c, http.StatusTooManyRequests, err.Error())
|
c,
|
||||||
|
channel,
|
||||||
|
modelRequest.Model,
|
||||||
|
common.GetContextKeyString(c, constant.ContextKeyUsingGroup),
|
||||||
|
)
|
||||||
|
if response != nil {
|
||||||
|
if response.code != "" {
|
||||||
|
abortWithOpenAiMessage(c, response.statusCode, response.message, response.code)
|
||||||
|
} else {
|
||||||
|
abortWithOpenAiMessage(c, response.statusCode, response.message)
|
||||||
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if channel.ConcurrencyLimit > 0 {
|
channel = selectedChannel
|
||||||
allowed, release := service.AcquireChannelConcurrency(channel.Id, channel.ConcurrencyLimit)
|
if releaseChannelConcurrency != nil {
|
||||||
if !allowed {
|
defer releaseChannelConcurrency()
|
||||||
abortWithOpenAiMessage(c, http.StatusTooManyRequests, "渠道并发请求数超过限制")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
defer release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
common.SetContextKey(c, constant.ContextKeyRequestStartTime, time.Now())
|
||||||
|
SetupContextForSelectedChannel(c, channel, modelRequest.Model)
|
||||||
c.Next()
|
c.Next()
|
||||||
if channel != nil && c.Writer != nil && c.Writer.Status() < http.StatusBadRequest {
|
if channel != nil && c.Writer != nil && c.Writer.Status() < http.StatusBadRequest {
|
||||||
service.RecordChannelAffinity(c, channel.Id)
|
service.RecordChannelAffinity(c, channel.Id)
|
||||||
@ -178,6 +201,113 @@ func Distribute() func(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func acquireDistributedChannelConcurrency(c *gin.Context, channel *model.Channel, modelName string, tokenGroup string) (*model.Channel, func(), *distributorConcurrencyResponse) {
|
||||||
|
release := func() {}
|
||||||
|
if channel == nil {
|
||||||
|
return nil, release, nil
|
||||||
|
}
|
||||||
|
if err := checkChannelQuotaForDistribute(channel.Id); err != nil {
|
||||||
|
return nil, nil, &distributorConcurrencyResponse{
|
||||||
|
statusCode: http.StatusTooManyRequests,
|
||||||
|
message: err.Error(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if channel.ConcurrencyLimit <= 0 {
|
||||||
|
return channel, release, nil
|
||||||
|
}
|
||||||
|
if allowed, release := acquireChannelConcurrencyForDistribute(channel.Id, channel.ConcurrencyLimit); allowed {
|
||||||
|
return channel, release, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
excludedChannelIDs := map[int]struct{}{channel.Id: {}}
|
||||||
|
if _, specificChannel := common.GetContextKey(c, constant.ContextKeyTokenSpecificChannelId); !specificChannel {
|
||||||
|
if selected, release, response := acquireAlternativeChannelConcurrency(c, modelName, tokenGroup, excludedChannelIDs); selected != nil || response != nil {
|
||||||
|
return selected, release, response
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
waitPlan := getChannelConcurrencyWaitPlanForDistribute()
|
||||||
|
if !channelConcurrencyWaitPlanEnabled(waitPlan) {
|
||||||
|
return nil, nil, &distributorConcurrencyResponse{
|
||||||
|
statusCode: http.StatusServiceUnavailable,
|
||||||
|
message: "No available accounts",
|
||||||
|
code: types.ErrorCodeGetChannelFailed,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
canWait, releaseWait := tryEnterChannelConcurrencyWaitQueueForDistribute(channel.Id, waitPlan.MaxWaitingRequests)
|
||||||
|
if !canWait {
|
||||||
|
return nil, nil, &distributorConcurrencyResponse{
|
||||||
|
statusCode: http.StatusTooManyRequests,
|
||||||
|
message: "Too many pending requests, please retry later",
|
||||||
|
code: types.ErrorCodeGetChannelFailed,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
defer releaseWait()
|
||||||
|
|
||||||
|
acquired, release, err := waitForChannelConcurrencySlotForDistribute(distributorRequestContext(c), channel.Id, channel.ConcurrencyLimit, time.Duration(waitPlan.TimeoutSeconds)*time.Second)
|
||||||
|
if err != nil || !acquired {
|
||||||
|
return nil, nil, &distributorConcurrencyResponse{
|
||||||
|
statusCode: http.StatusTooManyRequests,
|
||||||
|
message: "Concurrency limit exceeded for account, please retry later",
|
||||||
|
code: types.ErrorCodeGetChannelFailed,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return channel, release, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func acquireAlternativeChannelConcurrency(c *gin.Context, modelName string, tokenGroup string, excludedChannelIDs map[int]struct{}) (*model.Channel, func(), *distributorConcurrencyResponse) {
|
||||||
|
release := func() {}
|
||||||
|
for retry := 0; retry <= common.RetryTimes; retry++ {
|
||||||
|
for {
|
||||||
|
retryValue := retry
|
||||||
|
candidate, _, err := cacheGetRandomSatisfiedChannelForDistribute(&service.RetryParam{
|
||||||
|
Ctx: c,
|
||||||
|
ModelName: modelName,
|
||||||
|
TokenGroup: tokenGroup,
|
||||||
|
Retry: &retryValue,
|
||||||
|
ExcludedChannelIDs: excludedChannelIDs,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, &distributorConcurrencyResponse{
|
||||||
|
statusCode: http.StatusServiceUnavailable,
|
||||||
|
message: err.Error(),
|
||||||
|
code: types.ErrorCodeGetChannelFailed,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if candidate == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if _, excluded := excludedChannelIDs[candidate.Id]; excluded {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if err := checkChannelQuotaForDistribute(candidate.Id); err != nil {
|
||||||
|
excludedChannelIDs[candidate.Id] = struct{}{}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if candidate.ConcurrencyLimit <= 0 {
|
||||||
|
return candidate, release, nil
|
||||||
|
}
|
||||||
|
if allowed, release := acquireChannelConcurrencyForDistribute(candidate.Id, candidate.ConcurrencyLimit); allowed {
|
||||||
|
return candidate, release, nil
|
||||||
|
}
|
||||||
|
excludedChannelIDs[candidate.Id] = struct{}{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil, nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func channelConcurrencyWaitPlanEnabled(waitPlan operation_setting.ChannelConcurrencyWaitPlanSetting) bool {
|
||||||
|
return waitPlan.Enabled && waitPlan.TimeoutSeconds > 0 && waitPlan.MaxWaitingRequests > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func distributorRequestContext(c *gin.Context) context.Context {
|
||||||
|
if c != nil && c.Request != nil {
|
||||||
|
return c.Request.Context()
|
||||||
|
}
|
||||||
|
return context.Background()
|
||||||
|
}
|
||||||
|
|
||||||
// getModelFromRequest 从请求中读取模型信息
|
// getModelFromRequest 从请求中读取模型信息
|
||||||
// 根据 Content-Type 自动处理:
|
// 根据 Content-Type 自动处理:
|
||||||
// - application/json
|
// - application/json
|
||||||
|
|||||||
182
middleware/distributor_concurrency_test.go
Normal file
182
middleware/distributor_concurrency_test.go
Normal file
@ -0,0 +1,182 @@
|
|||||||
|
package middleware
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/model"
|
||||||
|
"github.com/QuantumNous/new-api/service"
|
||||||
|
"github.com/QuantumNous/new-api/setting/operation_setting"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func restoreDistributorConcurrencyStubs(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
previousCheckQuota := checkChannelQuotaForDistribute
|
||||||
|
previousAcquire := acquireChannelConcurrencyForDistribute
|
||||||
|
previousSelect := cacheGetRandomSatisfiedChannelForDistribute
|
||||||
|
previousWaitPlan := getChannelConcurrencyWaitPlanForDistribute
|
||||||
|
previousEnterWaitQueue := tryEnterChannelConcurrencyWaitQueueForDistribute
|
||||||
|
previousWaitForSlot := waitForChannelConcurrencySlotForDistribute
|
||||||
|
t.Cleanup(func() {
|
||||||
|
checkChannelQuotaForDistribute = previousCheckQuota
|
||||||
|
acquireChannelConcurrencyForDistribute = previousAcquire
|
||||||
|
cacheGetRandomSatisfiedChannelForDistribute = previousSelect
|
||||||
|
getChannelConcurrencyWaitPlanForDistribute = previousWaitPlan
|
||||||
|
tryEnterChannelConcurrencyWaitQueueForDistribute = previousEnterWaitQueue
|
||||||
|
waitForChannelConcurrencySlotForDistribute = previousWaitForSlot
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func newDistributorTestContext() *gin.Context {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
c, _ := gin.CreateTestContext(nil)
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAcquireDistributedChannelConcurrencyReturns503WithoutWaitPlan(t *testing.T) {
|
||||||
|
restoreDistributorConcurrencyStubs(t)
|
||||||
|
|
||||||
|
checkChannelQuotaForDistribute = func(channelId int) error { return nil }
|
||||||
|
acquireChannelConcurrencyForDistribute = func(channelId, limit int) (bool, func()) {
|
||||||
|
return false, func() {}
|
||||||
|
}
|
||||||
|
cacheGetRandomSatisfiedChannelForDistribute = func(param *service.RetryParam) (*model.Channel, string, error) {
|
||||||
|
return nil, param.TokenGroup, nil
|
||||||
|
}
|
||||||
|
getChannelConcurrencyWaitPlanForDistribute = func() operation_setting.ChannelConcurrencyWaitPlanSetting {
|
||||||
|
return operation_setting.ChannelConcurrencyWaitPlanSetting{Enabled: false}
|
||||||
|
}
|
||||||
|
|
||||||
|
selected, release, response := acquireDistributedChannelConcurrency(
|
||||||
|
newDistributorTestContext(),
|
||||||
|
&model.Channel{Id: 1, ConcurrencyLimit: 1},
|
||||||
|
"gpt-test",
|
||||||
|
"default",
|
||||||
|
)
|
||||||
|
|
||||||
|
require.Nil(t, selected)
|
||||||
|
require.Nil(t, release)
|
||||||
|
require.NotNil(t, response)
|
||||||
|
require.Equal(t, http.StatusServiceUnavailable, response.statusCode)
|
||||||
|
require.Equal(t, "No available accounts", response.message)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAcquireDistributedChannelConcurrencyReturns429WhenWaitQueueFull(t *testing.T) {
|
||||||
|
restoreDistributorConcurrencyStubs(t)
|
||||||
|
|
||||||
|
checkChannelQuotaForDistribute = func(channelId int) error { return nil }
|
||||||
|
acquireChannelConcurrencyForDistribute = func(channelId, limit int) (bool, func()) {
|
||||||
|
return false, func() {}
|
||||||
|
}
|
||||||
|
cacheGetRandomSatisfiedChannelForDistribute = func(param *service.RetryParam) (*model.Channel, string, error) {
|
||||||
|
return nil, param.TokenGroup, nil
|
||||||
|
}
|
||||||
|
getChannelConcurrencyWaitPlanForDistribute = func() operation_setting.ChannelConcurrencyWaitPlanSetting {
|
||||||
|
return operation_setting.ChannelConcurrencyWaitPlanSetting{
|
||||||
|
Enabled: true,
|
||||||
|
TimeoutSeconds: 30,
|
||||||
|
MaxWaitingRequests: 1,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tryEnterChannelConcurrencyWaitQueueForDistribute = func(channelId, maxWaiting int) (bool, func()) {
|
||||||
|
return false, func() {}
|
||||||
|
}
|
||||||
|
|
||||||
|
selected, release, response := acquireDistributedChannelConcurrency(
|
||||||
|
newDistributorTestContext(),
|
||||||
|
&model.Channel{Id: 1, ConcurrencyLimit: 1},
|
||||||
|
"gpt-test",
|
||||||
|
"default",
|
||||||
|
)
|
||||||
|
|
||||||
|
require.Nil(t, selected)
|
||||||
|
require.Nil(t, release)
|
||||||
|
require.NotNil(t, response)
|
||||||
|
require.Equal(t, http.StatusTooManyRequests, response.statusCode)
|
||||||
|
require.Equal(t, "Too many pending requests, please retry later", response.message)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAcquireDistributedChannelConcurrencyWaitsForSlot(t *testing.T) {
|
||||||
|
restoreDistributorConcurrencyStubs(t)
|
||||||
|
|
||||||
|
released := false
|
||||||
|
checkChannelQuotaForDistribute = func(channelId int) error { return nil }
|
||||||
|
acquireChannelConcurrencyForDistribute = func(channelId, limit int) (bool, func()) {
|
||||||
|
return false, func() {}
|
||||||
|
}
|
||||||
|
cacheGetRandomSatisfiedChannelForDistribute = func(param *service.RetryParam) (*model.Channel, string, error) {
|
||||||
|
return nil, param.TokenGroup, nil
|
||||||
|
}
|
||||||
|
getChannelConcurrencyWaitPlanForDistribute = func() operation_setting.ChannelConcurrencyWaitPlanSetting {
|
||||||
|
return operation_setting.ChannelConcurrencyWaitPlanSetting{
|
||||||
|
Enabled: true,
|
||||||
|
TimeoutSeconds: 30,
|
||||||
|
MaxWaitingRequests: 1,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tryEnterChannelConcurrencyWaitQueueForDistribute = func(channelId, maxWaiting int) (bool, func()) {
|
||||||
|
return true, func() {}
|
||||||
|
}
|
||||||
|
waitForChannelConcurrencySlotForDistribute = func(ctx context.Context, channelId, limit int, timeout time.Duration) (bool, func(), error) {
|
||||||
|
return true, func() { released = true }, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
selected, release, response := acquireDistributedChannelConcurrency(
|
||||||
|
newDistributorTestContext(),
|
||||||
|
&model.Channel{Id: 1, ConcurrencyLimit: 1},
|
||||||
|
"gpt-test",
|
||||||
|
"default",
|
||||||
|
)
|
||||||
|
|
||||||
|
require.Nil(t, response)
|
||||||
|
require.NotNil(t, selected)
|
||||||
|
require.Equal(t, 1, selected.Id)
|
||||||
|
require.NotNil(t, release)
|
||||||
|
release()
|
||||||
|
require.True(t, released)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAcquireDistributedChannelConcurrencyUsesAnotherAvailableChannel(t *testing.T) {
|
||||||
|
restoreDistributorConcurrencyStubs(t)
|
||||||
|
|
||||||
|
checkChannelQuotaForDistribute = func(channelId int) error {
|
||||||
|
if channelId == 2 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
acquireChannelConcurrencyForDistribute = func(channelId, limit int) (bool, func()) {
|
||||||
|
if channelId == 2 {
|
||||||
|
return true, func() {}
|
||||||
|
}
|
||||||
|
return false, func() {}
|
||||||
|
}
|
||||||
|
cacheGetRandomSatisfiedChannelForDistribute = func(param *service.RetryParam) (*model.Channel, string, error) {
|
||||||
|
if _, excluded := param.ExcludedChannelIDs[1]; excluded {
|
||||||
|
return &model.Channel{Id: 2, ConcurrencyLimit: 1}, param.TokenGroup, nil
|
||||||
|
}
|
||||||
|
return nil, param.TokenGroup, errors.New("expected saturated channel to be excluded")
|
||||||
|
}
|
||||||
|
getChannelConcurrencyWaitPlanForDistribute = func() operation_setting.ChannelConcurrencyWaitPlanSetting {
|
||||||
|
return operation_setting.ChannelConcurrencyWaitPlanSetting{Enabled: false}
|
||||||
|
}
|
||||||
|
|
||||||
|
selected, release, response := acquireDistributedChannelConcurrency(
|
||||||
|
newDistributorTestContext(),
|
||||||
|
&model.Channel{Id: 1, ConcurrencyLimit: 1},
|
||||||
|
"gpt-test",
|
||||||
|
"default",
|
||||||
|
)
|
||||||
|
|
||||||
|
require.Nil(t, response)
|
||||||
|
require.NotNil(t, selected)
|
||||||
|
require.Equal(t, 2, selected.Id)
|
||||||
|
require.NotNil(t, release)
|
||||||
|
release()
|
||||||
|
}
|
||||||
@ -104,6 +104,10 @@ func getChannelQuery(group string, model string, retry int) (*gorm.DB, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func GetChannel(group string, model string, retry int) (*Channel, error) {
|
func GetChannel(group string, model string, retry int) (*Channel, error) {
|
||||||
|
return GetChannelExcluding(group, model, retry, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetChannelExcluding(group string, model string, retry int, excluded map[int]struct{}) (*Channel, error) {
|
||||||
var abilities []Ability
|
var abilities []Ability
|
||||||
|
|
||||||
var err error = nil
|
var err error = nil
|
||||||
@ -111,6 +115,9 @@ func GetChannel(group string, model string, retry int) (*Channel, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if len(excluded) > 0 {
|
||||||
|
channelQuery = channelQuery.Where("channel_id NOT IN ?", excludedChannelIDList(excluded))
|
||||||
|
}
|
||||||
if common.UsingSQLite || common.UsingPostgreSQL {
|
if common.UsingSQLite || common.UsingPostgreSQL {
|
||||||
err = channelQuery.Order("weight DESC").Find(&abilities).Error
|
err = channelQuery.Order("weight DESC").Find(&abilities).Error
|
||||||
} else {
|
} else {
|
||||||
@ -143,6 +150,14 @@ func GetChannel(group string, model string, retry int) (*Channel, error) {
|
|||||||
return &channel, err
|
return &channel, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func excludedChannelIDList(excluded map[int]struct{}) []int {
|
||||||
|
ids := make([]int, 0, len(excluded))
|
||||||
|
for id := range excluded {
|
||||||
|
ids = append(ids, id)
|
||||||
|
}
|
||||||
|
return ids
|
||||||
|
}
|
||||||
|
|
||||||
func (channel *Channel) AddAbilities(tx *gorm.DB) error {
|
func (channel *Channel) AddAbilities(tx *gorm.DB) error {
|
||||||
models_ := strings.Split(channel.Models, ",")
|
models_ := strings.Split(channel.Models, ",")
|
||||||
groups_ := strings.Split(channel.Group, ",")
|
groups_ := strings.Split(channel.Group, ",")
|
||||||
|
|||||||
@ -94,9 +94,13 @@ func SyncChannelCache(frequency int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func GetRandomSatisfiedChannel(group string, model string, retry int) (*Channel, error) {
|
func GetRandomSatisfiedChannel(group string, model string, retry int) (*Channel, error) {
|
||||||
|
return GetRandomSatisfiedChannelExcluding(group, model, retry, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetRandomSatisfiedChannelExcluding(group string, model string, retry int, excluded map[int]struct{}) (*Channel, error) {
|
||||||
// if memory cache is disabled, get channel directly from database
|
// if memory cache is disabled, get channel directly from database
|
||||||
if !common.MemoryCacheEnabled {
|
if !common.MemoryCacheEnabled {
|
||||||
return GetChannel(group, model, retry)
|
return GetChannelExcluding(group, model, retry, excluded)
|
||||||
}
|
}
|
||||||
|
|
||||||
channelSyncLock.RLock()
|
channelSyncLock.RLock()
|
||||||
@ -111,6 +115,19 @@ func GetRandomSatisfiedChannel(group string, model string, retry int) (*Channel,
|
|||||||
channels = group2model2channels[group][normalizedModel]
|
channels = group2model2channels[group][normalizedModel]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if len(channels) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
if len(excluded) > 0 {
|
||||||
|
filtered := make([]int, 0, len(channels))
|
||||||
|
for _, channelId := range channels {
|
||||||
|
if _, ok := excluded[channelId]; !ok {
|
||||||
|
filtered = append(filtered, channelId)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
channels = filtered
|
||||||
|
}
|
||||||
|
|
||||||
if len(channels) == 0 {
|
if len(channels) == 0 {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@ -12,11 +12,12 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type RetryParam struct {
|
type RetryParam struct {
|
||||||
Ctx *gin.Context
|
Ctx *gin.Context
|
||||||
TokenGroup string
|
TokenGroup string
|
||||||
ModelName string
|
ModelName string
|
||||||
Retry *int
|
Retry *int
|
||||||
resetNextTry bool
|
ExcludedChannelIDs map[int]struct{}
|
||||||
|
resetNextTry bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *RetryParam) GetRetry() int {
|
func (p *RetryParam) GetRetry() int {
|
||||||
@ -115,7 +116,7 @@ func CacheGetRandomSatisfiedChannel(param *RetryParam) (*model.Channel, string,
|
|||||||
}
|
}
|
||||||
logger.LogDebug(param.Ctx, "Auto selecting group: %s, priorityRetry: %d", autoGroup, priorityRetry)
|
logger.LogDebug(param.Ctx, "Auto selecting group: %s, priorityRetry: %d", autoGroup, priorityRetry)
|
||||||
|
|
||||||
channel, _ = model.GetRandomSatisfiedChannel(autoGroup, param.ModelName, priorityRetry)
|
channel, _ = getRandomSatisfiedChannel(autoGroup, param.ModelName, priorityRetry, param.ExcludedChannelIDs)
|
||||||
if channel == nil {
|
if channel == nil {
|
||||||
// Current group has no available channel for this model, try next group
|
// Current group has no available channel for this model, try next group
|
||||||
// 当前分组没有该模型的可用渠道,尝试下一个分组
|
// 当前分组没有该模型的可用渠道,尝试下一个分组
|
||||||
@ -153,10 +154,17 @@ func CacheGetRandomSatisfiedChannel(param *RetryParam) (*model.Channel, string,
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
channel, err = model.GetRandomSatisfiedChannel(param.TokenGroup, param.ModelName, param.GetRetry())
|
channel, err = getRandomSatisfiedChannel(param.TokenGroup, param.ModelName, param.GetRetry(), param.ExcludedChannelIDs)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, param.TokenGroup, err
|
return nil, param.TokenGroup, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return channel, selectGroup, nil
|
return channel, selectGroup, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func getRandomSatisfiedChannel(group string, modelName string, retry int, excluded map[int]struct{}) (*model.Channel, error) {
|
||||||
|
if len(excluded) == 0 {
|
||||||
|
return model.GetRandomSatisfiedChannel(group, modelName, retry)
|
||||||
|
}
|
||||||
|
return model.GetRandomSatisfiedChannelExcluding(group, modelName, retry, excluded)
|
||||||
|
}
|
||||||
|
|||||||
@ -8,15 +8,18 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/QuantumNous/new-api/common"
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
"github.com/QuantumNous/new-api/setting/operation_setting"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
userConcurrencyKeyPrefix = "user:concurrency:"
|
userConcurrencyKeyPrefix = "user:concurrency:"
|
||||||
channelConcurrencyKeyPrefix = "channel:concurrency:"
|
channelConcurrencyKeyPrefix = "channel:concurrency:"
|
||||||
|
channelWaitingKeyPrefix = "channel:concurrency_waiting:"
|
||||||
userRPMKeyPrefix = "user:rpm:"
|
userRPMKeyPrefix = "user:rpm:"
|
||||||
|
|
||||||
concurrencyTTL = 5 * time.Minute
|
concurrencyTTL = 5 * time.Minute
|
||||||
rpmTTL = 2 * time.Minute
|
rpmTTL = 2 * time.Minute
|
||||||
|
waitingTTL = 2 * time.Minute
|
||||||
)
|
)
|
||||||
|
|
||||||
type OpsConcurrencySnapshot struct {
|
type OpsConcurrencySnapshot struct {
|
||||||
@ -45,6 +48,70 @@ func AcquireChannelConcurrency(channelId, limit int) (bool, func()) {
|
|||||||
return acquireConcurrency(channelConcurrencyKeyPrefix+strconv.Itoa(channelId), limit)
|
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 {
|
func CheckUserRPM(userId, limit int) bool {
|
||||||
if limit <= 0 || !redisAvailable() {
|
if limit <= 0 || !redisAvailable() {
|
||||||
return true
|
return true
|
||||||
|
|||||||
26
service/concurrency_wait_test.go
Normal file
26
service/concurrency_wait_test.go
Normal file
@ -0,0 +1,26 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestTryEnterChannelConcurrencyWaitQueueRejectsWhenPlanDisabled(t *testing.T) {
|
||||||
|
canWait, release := TryEnterChannelConcurrencyWaitQueue(1, 0)
|
||||||
|
|
||||||
|
require.False(t, canWait)
|
||||||
|
require.NotNil(t, release)
|
||||||
|
release()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWaitForChannelConcurrencySlotTimesOut(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
acquired, release, err := WaitForChannelConcurrencySlot(ctx, 1, 1, 0)
|
||||||
|
|
||||||
|
require.False(t, acquired)
|
||||||
|
require.NotNil(t, release)
|
||||||
|
require.Error(t, err)
|
||||||
|
release()
|
||||||
|
}
|
||||||
23
setting/operation_setting/channel_concurrency_wait_plan.go
Normal file
23
setting/operation_setting/channel_concurrency_wait_plan.go
Normal file
@ -0,0 +1,23 @@
|
|||||||
|
package operation_setting
|
||||||
|
|
||||||
|
import "github.com/QuantumNous/new-api/setting/config"
|
||||||
|
|
||||||
|
type ChannelConcurrencyWaitPlanSetting struct {
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
TimeoutSeconds int `json:"timeout_seconds"`
|
||||||
|
MaxWaitingRequests int `json:"max_waiting_requests"`
|
||||||
|
}
|
||||||
|
|
||||||
|
var channelConcurrencyWaitPlanSetting = ChannelConcurrencyWaitPlanSetting{
|
||||||
|
Enabled: false,
|
||||||
|
TimeoutSeconds: 30,
|
||||||
|
MaxWaitingRequests: 100,
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
config.GlobalConfig.Register("channel_concurrency_wait_plan", &channelConcurrencyWaitPlanSetting)
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetChannelConcurrencyWaitPlanSetting() ChannelConcurrencyWaitPlanSetting {
|
||||||
|
return channelConcurrencyWaitPlanSetting
|
||||||
|
}
|
||||||
Loading…
x
Reference in New Issue
Block a user