diff --git a/middleware/distributor.go b/middleware/distributor.go index d6269414..238eb033 100644 --- a/middleware/distributor.go +++ b/middleware/distributor.go @@ -23,8 +23,15 @@ import ( ) type ModelRequest struct { - Model string `json:"model"` - Group string `json:"group,omitempty"` + Model string `json:"model"` + Group string `json:"group,omitempty"` + Metadata *MuseRoutingPolicyInput `json:"metadata,omitempty"` +} + +type MuseRoutingPolicyInput struct { + Scene string `json:"scene,omitempty"` + RoutingPolicy string `json:"routing_policy,omitempty"` + DegradePolicy string `json:"degrade_policy,omitempty"` } func Distribute() func(c *gin.Context) { @@ -36,6 +43,12 @@ func Distribute() func(c *gin.Context) { abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidRequest, map[string]any{"Error": err.Error()})) return } + policy := service.ResolveMuseRoutingPolicy( + firstNonBlank(common.GetContextKeyString(c, constant.ContextKeyMuseScene), modelRequest.MetadataValue("scene")), + modelRequest.MetadataValue("routing_policy"), + modelRequest.MetadataValue("degrade_policy"), + ) + service.SetMuseRoutingPolicy(c, policy) if ok { id, err := strconv.Atoi(channelId.(string)) if err != nil { @@ -98,6 +111,7 @@ func Distribute() func(c *gin.Context) { common.SetContextKey(c, constant.ContextKeyUsingGroup, usingGroup) } } + usingGroup = service.ApplyMuseRoutingPolicy(c, usingGroup) if preferredChannelID, found := service.GetPreferredChannelByAffinity(c, modelRequest.Model, usingGroup); found { preferred, err := model.CacheGetChannel(preferredChannelID) @@ -109,7 +123,7 @@ func Distribute() func(c *gin.Context) { } } else if usingGroup == "auto" { userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup) - autoGroups := service.GetUserAutoGroup(userGroup) + autoGroups := service.ResolveMusePreferredGroups(c, userGroup) for _, g := range autoGroups { if model.IsChannelEnabledForGroupModel(g, modelRequest.Model, preferred.Id) { selectGroup = g @@ -128,12 +142,31 @@ func Distribute() func(c *gin.Context) { } if channel == nil { - channel, selectGroup, err = service.CacheGetRandomSatisfiedChannel(&service.RetryParam{ + retryParam := &service.RetryParam{ Ctx: c, ModelName: modelRequest.Model, TokenGroup: usingGroup, Retry: common.GetPointer(0), - }) + } + channel, selectGroup, err = service.CacheGetRandomSatisfiedChannel(retryParam) + if channel == nil && err == nil { + candidates := service.MuseFallbackModelCandidates(modelRequest.Model, policy) + for _, candidateModel := range candidates[1:] { + candidateParam := &service.RetryParam{ + Ctx: c, + ModelName: candidateModel, + TokenGroup: usingGroup, + Retry: common.GetPointer(0), + } + channel, selectGroup, err = service.CacheGetRandomSatisfiedChannel(candidateParam) + if channel != nil || err != nil { + if channel != nil { + modelRequest.Model = candidateModel + } + break + } + } + } if err != nil { showGroup := usingGroup if usingGroup == "auto" { @@ -164,6 +197,31 @@ func Distribute() func(c *gin.Context) { } } +func (r *ModelRequest) MetadataValue(name string) string { + if r == nil || r.Metadata == nil { + return "" + } + switch name { + case "scene": + return r.Metadata.Scene + case "routing_policy": + return r.Metadata.RoutingPolicy + case "degrade_policy": + return r.Metadata.DegradePolicy + default: + return "" + } +} + +func firstNonBlank(values ...string) string { + for _, value := range values { + if strings.TrimSpace(value) != "" { + return strings.TrimSpace(value) + } + } + return "" +} + // getModelFromRequest 从请求中读取模型信息 // 根据 Content-Type 自动处理: // - application/json diff --git a/service/channel_select.go b/service/channel_select.go index a3710ef8..e87f02d7 100644 --- a/service/channel_select.go +++ b/service/channel_select.go @@ -90,7 +90,7 @@ func CacheGetRandomSatisfiedChannel(param *RetryParam) (*model.Channel, string, if len(setting.GetAutoGroups()) == 0 { return nil, selectGroup, errors.New("auto groups is not enabled") } - autoGroups := GetUserAutoGroup(userGroup) + autoGroups := ResolveMusePreferredGroups(param.Ctx, userGroup) // startGroupIndex: the group index to start searching from // startGroupIndex: 开始搜索的分组索引 diff --git a/service/muse_routing_policy.go b/service/muse_routing_policy.go new file mode 100644 index 00000000..860d1036 --- /dev/null +++ b/service/muse_routing_policy.go @@ -0,0 +1,146 @@ +package service + +import ( + "slices" + "strings" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/gin-gonic/gin" +) + +const ginKeyMuseRoutingPolicy = "muse_routing_policy" + +type MuseRoutingPolicy struct { + Scene string + AllowCrossGroupRetry bool + AllowModelDegrade bool + PreferredGroups []string +} + +func ResolveMuseRoutingPolicy(scene string, routingPolicy string, degradePolicy string) MuseRoutingPolicy { + scene = normalizeMusePolicyValue(scene) + routingPolicy = normalizeMusePolicyValue(routingPolicy) + degradePolicy = normalizeMusePolicyValue(degradePolicy) + + policy := MuseRoutingPolicy{ + Scene: scene, + PreferredGroups: defaultMusePreferredGroups(scene, routingPolicy), + } + + switch routingPolicy { + case "quality_first", "balanced": + policy.AllowCrossGroupRetry = true + } + + if degradePolicy == "allow_lower_tier" { + policy.AllowModelDegrade = true + } + + return policy +} + +func SetMuseRoutingPolicy(c *gin.Context, policy MuseRoutingPolicy) { + if c == nil { + return + } + if policy.Scene == "" && !policy.AllowCrossGroupRetry && !policy.AllowModelDegrade && len(policy.PreferredGroups) == 0 { + return + } + c.Set(ginKeyMuseRoutingPolicy, policy) +} + +func GetMuseRoutingPolicy(c *gin.Context) (MuseRoutingPolicy, bool) { + if c == nil { + return MuseRoutingPolicy{}, false + } + value, ok := c.Get(ginKeyMuseRoutingPolicy) + if !ok { + return MuseRoutingPolicy{}, false + } + policy, ok := value.(MuseRoutingPolicy) + if !ok { + return MuseRoutingPolicy{}, false + } + return policy, true +} + +func ApplyMuseRoutingPolicy(c *gin.Context, usingGroup string) string { + policy, ok := GetMuseRoutingPolicy(c) + if !ok { + return usingGroup + } + if policy.AllowCrossGroupRetry { + common.SetContextKey(c, constant.ContextKeyTokenCrossGroupRetry, true) + if usingGroup != "auto" { + usingGroup = "auto" + common.SetContextKey(c, constant.ContextKeyUsingGroup, usingGroup) + } + } + return usingGroup +} + +func ResolveMusePreferredGroups(c *gin.Context, userGroup string) []string { + policy, ok := GetMuseRoutingPolicy(c) + if !ok || len(policy.PreferredGroups) == 0 { + return GetUserAutoGroup(userGroup) + } + usableGroups := GetUserUsableGroups(userGroup) + filtered := make([]string, 0, len(policy.PreferredGroups)) + for _, group := range policy.PreferredGroups { + if _, allowed := usableGroups[group]; allowed { + filtered = append(filtered, group) + } + } + if len(filtered) == 0 { + return GetUserAutoGroup(userGroup) + } + return filtered +} + +func MuseFallbackModelCandidates(modelName string, policy MuseRoutingPolicy) []string { + candidates := []string{modelName} + if !policy.AllowModelDegrade { + return candidates + } + + appendUnique := func(candidate string) { + candidate = strings.TrimSpace(candidate) + if candidate == "" { + return + } + if !slices.Contains(candidates, candidate) { + candidates = append(candidates, candidate) + } + } + + switch { + case strings.HasPrefix(modelName, "gpt-4o") && modelName != "gpt-4o-mini": + appendUnique("gpt-4o-mini") + case strings.Contains(modelName, "gemini") && strings.Contains(modelName, "pro"): + appendUnique(strings.Replace(modelName, "pro", "flash", 1)) + case strings.Contains(modelName, "claude") && strings.Contains(modelName, "opus"): + appendUnique(strings.Replace(modelName, "opus", "sonnet", 1)) + case strings.Contains(modelName, "claude") && strings.Contains(modelName, "sonnet"): + appendUnique(strings.Replace(modelName, "sonnet", "haiku", 1)) + } + + return candidates +} + +func defaultMusePreferredGroups(scene string, routingPolicy string) []string { + switch routingPolicy { + case "quality_first": + return []string{"svip", "vip", "default"} + case "balanced": + return []string{"vip", "default"} + } + if scene == "suggestion_generation" { + return []string{"default"} + } + return nil +} + +func normalizeMusePolicyValue(value string) string { + return strings.TrimSpace(strings.ToLower(value)) +} diff --git a/service/muse_routing_policy_test.go b/service/muse_routing_policy_test.go new file mode 100644 index 00000000..d5e25772 --- /dev/null +++ b/service/muse_routing_policy_test.go @@ -0,0 +1,43 @@ +package service + +import ( + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestResolveMuseRoutingPolicy_QualityFirstAllowsFallback(t *testing.T) { + policy := ResolveMuseRoutingPolicy("suggestion_generation", "quality_first", "allow_lower_tier") + + require.Equal(t, "suggestion_generation", policy.Scene) + require.True(t, policy.AllowCrossGroupRetry) + require.True(t, policy.AllowModelDegrade) + require.Equal(t, []string{"svip", "vip", "default"}, policy.PreferredGroups) +} + +func TestResolveMuseRoutingPolicy_DefaultsWithoutPolicies(t *testing.T) { + policy := ResolveMuseRoutingPolicy("suggestion_generation", "", "") + + require.False(t, policy.AllowCrossGroupRetry) + require.False(t, policy.AllowModelDegrade) + require.Equal(t, []string{"default"}, policy.PreferredGroups) +} + +func TestMuseFallbackModelCandidates_AllowLowerTierAddsFallback(t *testing.T) { + policy := ResolveMuseRoutingPolicy("suggestion_generation", "quality_first", "allow_lower_tier") + + candidates := MuseFallbackModelCandidates("gpt-4o", policy) + + require.Equal(t, []string{"gpt-4o", "gpt-4o-mini"}, candidates) +} + +func TestResolveMusePreferredGroups_UsesSceneDrivenDefaultsFromContext(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(nil) + SetMuseRoutingPolicy(ctx, ResolveMuseRoutingPolicy("suggestion_generation", "", "")) + + groups := ResolveMusePreferredGroups(ctx, "default") + + require.Equal(t, []string{"default"}, groups) +}