From feb5f9d7f29e530cd91a39691cb2eb986a54adc7 Mon Sep 17 00:00:00 2001 From: zijiren <84728412+zijiren233@users.noreply.github.com> Date: Mon, 3 Mar 2025 16:41:09 +0800 Subject: [PATCH] feat: relay retry ignore forbidden channel (#5435) * feat: relay retry ignore forbidden channel * fix: dont reuse meta --- service/aiproxy/controller/relay.go | 88 ++++++++++++++++++++--------- service/aiproxy/middleware/utils.go | 6 +- service/aiproxy/relay/meta/meta.go | 27 ++++----- 3 files changed, 75 insertions(+), 46 deletions(-) diff --git a/service/aiproxy/controller/relay.go b/service/aiproxy/controller/relay.go index 3b0606c9c..078b60401 100644 --- a/service/aiproxy/controller/relay.go +++ b/service/aiproxy/controller/relay.go @@ -80,8 +80,8 @@ func RelayHelper(meta *meta.Meta, c *gin.Context, relayController RelayControlle return err, shouldRetry(c, err.StatusCode) } -func getChannelWithFallback(cache *dbmodel.ModelCaches, model string, failedChannelIDs ...int) (*dbmodel.Channel, error) { - channel, err := cache.GetRandomSatisfiedChannel(model, failedChannelIDs...) +func getChannelWithFallback(cache *dbmodel.ModelCaches, model string, ignoreChannelIDs ...int) (*dbmodel.Channel, error) { + channel, err := cache.GetRandomSatisfiedChannel(model, ignoreChannelIDs...) if err == nil { return channel, nil } @@ -110,17 +110,14 @@ func relay(c *gin.Context, mode int, relayController RelayController) { if err != nil { log.Errorf("get %s auto banned channels failed: %+v", requestModel, err) } - log.Debugf("%s model banned channels: %+v", requestModel, ids) - - failedChannelIDs := []int{} + ignoreChannelIDs := make([]int, 0, len(ids)) for _, id := range ids { - failedChannelIDs = append(failedChannelIDs, int(id)) + ignoreChannelIDs = append(ignoreChannelIDs, int(id)) } mc := middleware.GetModelCaches(c) - - channel, err := getChannelWithFallback(mc, requestModel, failedChannelIDs...) + channel, err := getChannelWithFallback(mc, requestModel, ignoreChannelIDs...) if err != nil { c.JSON(http.StatusServiceUnavailable, gin.H{ "error": &model.Error{ @@ -137,35 +134,52 @@ func relay(c *gin.Context, mode int, relayController RelayController) { if bizErr == nil { return } - failedChannelIDs = append(failedChannelIDs, channel.ID) - requestID := middleware.GetRequestID(c) - var retryTimes int64 - if retry { - retryTimes = config.GetRetryTimes() + if !retry { + bizErr.Error.Message = middleware.MessageWithRequestID(c, bizErr.Error.Message) + c.JSON(bizErr.StatusCode, bizErr) + return } + + var lastCanContinueChannel *dbmodel.Channel + + retryTimes := config.GetRetryTimes() + if !channelCanContinue(bizErr.StatusCode) { + ignoreChannelIDs = append(ignoreChannelIDs, channel.ID) + } else { + lastCanContinueChannel = channel + } + for i := retryTimes; i > 0; i-- { - newChannel, err := mc.GetRandomSatisfiedChannel(requestModel, failedChannelIDs...) + newChannel, err := mc.GetRandomSatisfiedChannel(requestModel, ignoreChannelIDs...) if err != nil { - if errors.Is(err, dbmodel.ErrChannelsNotFound) { + if !errors.Is(err, dbmodel.ErrChannelsExhausted) || + lastCanContinueChannel == nil { break } - // use first channel to retry - if !errors.Is(err, dbmodel.ErrChannelsExhausted) { - break - } - newChannel = channel + // use last can continue channel to retry + newChannel = lastCanContinueChannel } - log.Warnf("using channel %s(%d) to retry (remain times %d)", newChannel.Name, newChannel.ID, i) + log.Warnf("using channel %s (type: %d, id: %d) to retry (remain times %d)", + newChannel.Name, + newChannel.Type, + newChannel.ID, + i-1, + ) + requestBody, err := common.GetRequestBody(c.Request) if err != nil { log.Errorf("GetRequestBody failed: %+v", err) break } c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody)) - meta.Reset(newChannel) - //nolint:gosec - // random wait 1-2 seconds - time.Sleep(time.Duration(rand.Float64()*float64(time.Second)) + time.Second) + + if shouldDelay(bizErr.StatusCode) { + //nolint:gosec + // random wait 1-2 seconds + time.Sleep(time.Duration(rand.Float64()*float64(time.Second)) + time.Second) + } + + meta := middleware.NewMetaByContext(c, newChannel, requestModel, mode) bizErr, retry = RelayHelper(meta, c, relayController) if bizErr == nil { return @@ -173,10 +187,15 @@ func relay(c *gin.Context, mode int, relayController RelayController) { if !retry { break } - failedChannelIDs = append(failedChannelIDs, newChannel.ID) + if !channelCanContinue(bizErr.StatusCode) { + ignoreChannelIDs = append(ignoreChannelIDs, newChannel.ID) + } else { + lastCanContinueChannel = newChannel + } } + if bizErr != nil { - bizErr.Error.Message = middleware.MessageWithRequestID(bizErr.Error.Message, requestID) + bizErr.Error.Message = middleware.MessageWithRequestID(c, bizErr.Error.Message) c.JSON(bizErr.StatusCode, bizErr) } } @@ -195,6 +214,21 @@ func shouldRetry(_ *gin.Context, statusCode int) bool { return ok } +var channelCanContinueStatusCodesMap = map[int]struct{}{ + http.StatusTooManyRequests: {}, + http.StatusRequestTimeout: {}, + http.StatusGatewayTimeout: {}, +} + +func channelCanContinue(statusCode int) bool { + _, ok := channelCanContinueStatusCodesMap[statusCode] + return ok +} + +func shouldDelay(statusCode int) bool { + return statusCode == http.StatusTooManyRequests +} + // 仅当是channel错误时,才需要记录,用户请求参数错误时,不需要记录 func shouldErrorMonitor(statusCode int) bool { return statusCode != http.StatusBadRequest diff --git a/service/aiproxy/middleware/utils.go b/service/aiproxy/middleware/utils.go index 0d37992ad..c7264561e 100644 --- a/service/aiproxy/middleware/utils.go +++ b/service/aiproxy/middleware/utils.go @@ -11,8 +11,8 @@ const ( ErrorTypeAIPROXY = "aiproxy_error" ) -func MessageWithRequestID(message string, id string) string { - return fmt.Sprintf("%s (aiproxy: %s)", message, id) +func MessageWithRequestID(c *gin.Context, message string) string { + return fmt.Sprintf("%s (aiproxy: %s)", message, GetRequestID(c)) } func abortLogWithMessage(c *gin.Context, statusCode int, message string) { @@ -23,7 +23,7 @@ func abortLogWithMessage(c *gin.Context, statusCode int, message string) { func abortWithMessage(c *gin.Context, statusCode int, message string) { c.JSON(statusCode, gin.H{ "error": &model.Error{ - Message: MessageWithRequestID(message, GetRequestID(c)), + Message: MessageWithRequestID(c, message), Type: ErrorTypeAIPROXY, }, }) diff --git a/service/aiproxy/relay/meta/meta.go b/service/aiproxy/relay/meta/meta.go index bde574e3a..f0417aa84 100644 --- a/service/aiproxy/relay/meta/meta.go +++ b/service/aiproxy/relay/meta/meta.go @@ -92,27 +92,22 @@ func NewMeta( } if channel != nil { - meta.Reset(channel) + meta.Channel = &ChannelMeta{ + Name: channel.Name, + BaseURL: channel.BaseURL, + Key: channel.Key, + ID: channel.ID, + Type: channel.Type, + } + if channel.Config != nil { + meta.ChannelConfig = *channel.Config + } + meta.ActualModel, _ = GetMappedModelName(modelName, channel.ModelMapping) } return &meta } -func (m *Meta) Reset(channel *model.Channel) { - m.Channel = &ChannelMeta{ - Name: channel.Name, - BaseURL: channel.BaseURL, - Key: channel.Key, - ID: channel.ID, - Type: channel.Type, - } - if channel.Config != nil { - m.ChannelConfig = *channel.Config - } - m.ActualModel, _ = GetMappedModelName(m.OriginModel, channel.ModelMapping) - m.ClearValues() -} - func (m *Meta) ClearValues() { clear(m.values) }