diff --git a/controller/metaproxy_provision.go b/controller/metaproxy_provision.go index 36decc712dbe..54ebf8f82937 100644 --- a/controller/metaproxy_provision.go +++ b/controller/metaproxy_provision.go @@ -11,6 +11,7 @@ import ( "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/model" + "github.com/QuantumNous/new-api/setting/billing_setting" "github.com/gin-gonic/gin" ) @@ -38,6 +39,8 @@ type metaproxyProvisionOptionsRequest struct { ModelRatio string `json:"model_ratio"` CompletionRatio string `json:"completion_ratio"` CacheRatio string `json:"cache_ratio"` + ModelBillingMode string `json:"model_billing_mode"` + ModelBillingExpr string `json:"model_billing_expr"` GroupRatio string `json:"group_ratio"` UserUsableGroups string `json:"user_usable_groups"` } @@ -78,6 +81,22 @@ func parseUserUsableGroups(raw string) (map[string]string, error) { return values, nil } +func parseStringMap(name, raw string) (map[string]string, error) { + values := make(map[string]string) + if strings.TrimSpace(raw) == "" { + return values, nil + } + if err := common.Unmarshal([]byte(raw), &values); err != nil { + return nil, fmt.Errorf("%s must be a JSON object of strings: %w", name, err) + } + for key, value := range values { + if strings.TrimSpace(key) == "" || strings.TrimSpace(value) == "" { + return nil, fmt.Errorf("%s contains an empty model or value", name) + } + } + return values, nil +} + func validateOptionalJSONObject(name, raw string) error { if raw == "" { return nil @@ -152,6 +171,31 @@ func validateMetaproxyProvisionRequest( if _, err := parseRatioMap("CacheRatio", request.Options.CacheRatio, true); err != nil { return err } + billingModes, err := parseStringMap("ModelBillingMode", request.Options.ModelBillingMode) + if err != nil { + return err + } + billingExprs, err := parseStringMap("ModelBillingExpr", request.Options.ModelBillingExpr) + if err != nil { + return err + } + for modelName, mode := range billingModes { + if mode != billing_setting.BillingModeTieredExpr { + return fmt.Errorf("ModelBillingMode contains unsupported mode %q for %q", mode, modelName) + } + expr, ok := billingExprs[modelName] + if !ok { + return fmt.Errorf("model %q uses tiered_expr but is missing from ModelBillingExpr", modelName) + } + if err := billing_setting.SmokeTestExpr(expr); err != nil { + return fmt.Errorf("ModelBillingExpr for %q is invalid: %w", modelName, err) + } + } + for modelName := range billingExprs { + if _, ok := billingModes[modelName]; !ok { + return fmt.Errorf("model %q has ModelBillingExpr but is missing from ModelBillingMode", modelName) + } + } groupRatios, err := parseRatioMap("GroupRatio", request.Options.GroupRatio, false) if err != nil { return err @@ -220,10 +264,11 @@ func validateMetaproxyProvisionRequest( } for _, modelName := range models { ratio, priced := modelRatios[modelName] - if !priced { - return fmt.Errorf("enabled model %q in an offered group is missing from ModelRatio", modelName) + _, expressionPriced := billingModes[modelName] + if !priced && !expressionPriced { + return fmt.Errorf("enabled model %q in an offered group is missing from ModelRatio and ModelBillingMode", modelName) } - if ratio == 0 { + if priced && ratio == 0 && !expressionPriced { return fmt.Errorf("enabled model %q in an offered group must have a positive ModelRatio", modelName) } } @@ -232,6 +277,12 @@ func validateMetaproxyProvisionRequest( } func toMetaproxyProvisionConfig(request metaproxyProvisionRequest) model.MetaproxyProvisionConfig { + if strings.TrimSpace(request.Options.ModelBillingMode) == "" { + request.Options.ModelBillingMode = "{}" + } + if strings.TrimSpace(request.Options.ModelBillingExpr) == "" { + request.Options.ModelBillingExpr = "{}" + } channels := make([]model.MetaproxyProvisionChannel, 0, len(request.Channels)) for _, channel := range request.Channels { models, _ := commaValues(channel.Models) @@ -259,6 +310,8 @@ func toMetaproxyProvisionConfig(request metaproxyProvisionRequest) model.Metapro ModelRatio: request.Options.ModelRatio, CompletionRatio: request.Options.CompletionRatio, CacheRatio: request.Options.CacheRatio, + ModelBillingMode: request.Options.ModelBillingMode, + ModelBillingExpr: request.Options.ModelBillingExpr, GroupRatio: request.Options.GroupRatio, UserUsableGroups: request.Options.UserUsableGroups, }, diff --git a/controller/metaproxy_provision_test.go b/controller/metaproxy_provision_test.go index ca4eb47eccab..adfd0f6af7f4 100644 --- a/controller/metaproxy_provision_test.go +++ b/controller/metaproxy_provision_test.go @@ -34,17 +34,56 @@ func validMetaproxyProvisionRequest() metaproxyProvisionRequest { ModelRatio: `{"model-one":1}`, CompletionRatio: `{"model-one":2}`, CacheRatio: `{"model-one":0.1}`, + ModelBillingMode: `{}`, + ModelBillingExpr: `{}`, GroupRatio: `{"standard":1}`, UserUsableGroups: `{"standard":"Standard"}`, }, } } +func TestValidateMetaproxyProvisionRequestAcceptsExpressionPricedModel(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Options.ModelRatio = `{}` + request.Options.ModelBillingMode = `{"model-one":"tiered_expr"}` + request.Options.ModelBillingExpr = `{"model-one":"(param(\"n\") == nil ? 1 : param(\"n\")) * (param(\"resolution\") == \"1k\" ? tier(\"1k\", 200000) : tier(\"2k\", 300000))"}` + require.NoError(t, validateMetaproxyProvisionRequest(request, request.Digest, "none")) +} + +func TestValidateMetaproxyProvisionRequestRejectsInvalidBillingExpression(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Options.ModelRatio = `{}` + request.Options.ModelBillingMode = `{"model-one":"tiered_expr"}` + request.Options.ModelBillingExpr = `{"model-one":"tier(\"broken\", -1"}` + err := validateMetaproxyProvisionRequest(request, request.Digest, "none") + require.ErrorContains(t, err, "ModelBillingExpr") +} + +func TestValidateMetaproxyProvisionRequestRejectsExpressionModeWithoutExpression(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Options.ModelRatio = `{}` + request.Options.ModelBillingMode = `{"model-one":"tiered_expr"}` + err := validateMetaproxyProvisionRequest(request, request.Digest, "none") + require.ErrorContains(t, err, "model-one") + require.ErrorContains(t, err, "ModelBillingExpr") +} + func TestValidateMetaproxyProvisionRequestAcceptsCompleteConfig(t *testing.T) { request := validMetaproxyProvisionRequest() require.NoError(t, validateMetaproxyProvisionRequest(request, request.Digest, "none")) } +func TestValidateMetaproxyProvisionRequestAcceptsLegacyRequestWithoutBillingExpressions(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Options.ModelBillingMode = "" + request.Options.ModelBillingExpr = "" + require.NoError(t, validateMetaproxyProvisionRequest(request, request.Digest, "none")) + + config := toMetaproxyProvisionConfig(request) + require.Equal(t, "{}", config.Options.ModelBillingMode) + require.Equal(t, "{}", config.Options.ModelBillingExpr) +} + func TestValidateMetaproxyProvisionRequestRejectsMismatchedIdempotencyKey(t *testing.T) { request := validMetaproxyProvisionRequest() err := validateMetaproxyProvisionRequest( @@ -122,6 +161,17 @@ func TestToMetaproxyProvisionConfigStoresNormalizedLists(t *testing.T) { require.Equal(t, "standard,archive", config.Channels[0].Group) } +func TestToMetaproxyProvisionConfigNormalizesWhitespaceBillingOptions(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Options.ModelBillingMode = " \t\n" + request.Options.ModelBillingExpr = "\r\n" + + config := toMetaproxyProvisionConfig(request) + + require.Equal(t, "{}", config.Options.ModelBillingMode) + require.Equal(t, "{}", config.Options.ModelBillingExpr) +} + func TestApplyMetaproxyProvisionRequiresMemoryCache(t *testing.T) { previous := common.MemoryCacheEnabled common.MemoryCacheEnabled = false diff --git a/dto/openai_image.go b/dto/openai_image.go index 275fd5559080..bd6f514cd32e 100644 --- a/dto/openai_image.go +++ b/dto/openai_image.go @@ -20,6 +20,8 @@ type ImageRequest struct { Prompt string `json:"prompt" binding:"required"` N *uint `json:"n,omitempty"` Size string `json:"size,omitempty"` + Resolution string `json:"resolution,omitempty"` + AspectRatio string `json:"aspect_ratio,omitempty"` Quality string `json:"quality,omitempty"` ResponseFormat string `json:"response_format,omitempty"` Style json.RawMessage `json:"style,omitempty"` diff --git a/dto/openai_image_test.go b/dto/openai_image_test.go new file mode 100644 index 000000000000..c02bf1f8cfcc --- /dev/null +++ b/dto/openai_image_test.go @@ -0,0 +1,20 @@ +package dto + +import ( + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/stretchr/testify/require" +) + +func TestImageRequestPreservesProviderResolutionFields(t *testing.T) { + raw := []byte(`{"model":"grok-imagine-image","prompt":"test","resolution":"2k","aspect_ratio":"16:9"}`) + var request ImageRequest + require.NoError(t, common.Unmarshal(raw, &request)) + require.Equal(t, "2k", request.Resolution) + require.Equal(t, "16:9", request.AspectRatio) + + encoded, err := common.Marshal(request) + require.NoError(t, err) + require.JSONEq(t, string(raw), string(encoded)) +} diff --git a/model/metaproxy_provision.go b/model/metaproxy_provision.go index 5f14c5ee0ba0..ed6a9785c7dc 100644 --- a/model/metaproxy_provision.go +++ b/model/metaproxy_provision.go @@ -45,6 +45,8 @@ type MetaproxyProvisionOptions struct { ModelRatio string `json:"model_ratio"` CompletionRatio string `json:"completion_ratio"` CacheRatio string `json:"cache_ratio"` + ModelBillingMode string `json:"model_billing_mode"` + ModelBillingExpr string `json:"model_billing_expr"` GroupRatio string `json:"group_ratio"` UserUsableGroups string `json:"user_usable_groups"` } @@ -54,6 +56,8 @@ func (options MetaproxyProvisionOptions) orderedValues() []Option { {Key: "ModelRatio", Value: options.ModelRatio}, {Key: "CompletionRatio", Value: options.CompletionRatio}, {Key: "CacheRatio", Value: options.CacheRatio}, + {Key: "billing_setting.billing_mode", Value: options.ModelBillingMode}, + {Key: "billing_setting.billing_expr", Value: options.ModelBillingExpr}, {Key: "GroupRatio", Value: options.GroupRatio}, {Key: "UserUsableGroups", Value: options.UserUsableGroups}, } diff --git a/model/metaproxy_provision_test.go b/model/metaproxy_provision_test.go index 9fe410d7aead..610959dc6f0a 100644 --- a/model/metaproxy_provision_test.go +++ b/model/metaproxy_provision_test.go @@ -85,6 +85,8 @@ func seedProvisionState(t *testing.T, db *gorm.DB) Channel { {Key: "ModelRatio", Value: `{"old-model":1}`}, {Key: "CompletionRatio", Value: `{"old-model":2}`}, {Key: "CacheRatio", Value: `{}`}, + {Key: "billing_setting.billing_mode", Value: `{}`}, + {Key: "billing_setting.billing_expr", Value: `{}`}, {Key: "GroupRatio", Value: `{"standard":1}`}, {Key: "UserUsableGroups", Value: `{"default":"standard"}`}, }).Error) @@ -102,6 +104,8 @@ func desiredProvisionConfig() MetaproxyProvisionConfig { ModelRatio: `{"new-model":1.5}`, CompletionRatio: `{"new-model":3}`, CacheRatio: `{"new-model":0.1}`, + ModelBillingMode: `{"image-model":"tiered_expr"}`, + ModelBillingExpr: `{"image-model":"tier(\"base\", 200000)"}`, GroupRatio: `{"standard":1}`, UserUsableGroups: `{"default":"standard"}`, }, @@ -140,6 +144,8 @@ func TestApplyMetaproxyProvisionReplacesManagedStateAtomically(t *testing.T) { require.Equal(t, desiredProvisionConfig().Digest, got[MetaproxyProvisionDigestOption]) require.Equal(t, desiredProvisionConfig().Revision, got[MetaproxyProvisionRevisionOption]) require.Equal(t, `{"new-model":1.5}`, got["ModelRatio"]) + require.Equal(t, `{"image-model":"tiered_expr"}`, got["billing_setting.billing_mode"]) + require.Equal(t, `{"image-model":"tier(\"base\", 200000)"}`, got["billing_setting.billing_expr"]) common.OptionMapRWMutex.RLock() require.Equal(t, "old-digest", common.OptionMap[MetaproxyProvisionDigestOption]) diff --git a/relay/channel/minimax/adaptor_test.go b/relay/channel/minimax/adaptor_test.go index 46d57c11ff72..342eec861161 100644 --- a/relay/channel/minimax/adaptor_test.go +++ b/relay/channel/minimax/adaptor_test.go @@ -84,6 +84,22 @@ func TestConvertImageRequest(t *testing.T) { } } +func TestConvertImageRequestPreservesExplicitAspectRatio(t *testing.T) { + t.Parallel() + + request := dto.ImageRequest{ + Model: "image-01", + Prompt: "a red fox in snowfall", + Size: "1536x1024", + AspectRatio: "16:9", + } + + got := oaiImage2MiniMaxImageRequest(request) + if got.AspectRatio != "16:9" { + t.Fatalf("aspect_ratio = %q, want %q", got.AspectRatio, "16:9") + } +} + func TestDoResponseForImageGeneration(t *testing.T) { t.Parallel() diff --git a/relay/channel/minimax/image.go b/relay/channel/minimax/image.go index 9b316bdc8230..9d851fefa92f 100644 --- a/relay/channel/minimax/image.go +++ b/relay/channel/minimax/image.go @@ -69,6 +69,9 @@ func oaiImage2MiniMaxImageRequest(request dto.ImageRequest) MiniMaxImageRequest } func aspectRatioFromImageRequest(request dto.ImageRequest) string { + if request.AspectRatio != "" { + return request.AspectRatio + } if raw, ok := request.Extra["aspect_ratio"]; ok { var aspectRatio string if err := common.Unmarshal(raw, &aspectRatio); err == nil && aspectRatio != "" {