diff --git a/i18n/keys.go b/i18n/keys.go index 63e654b585e7..5be5dc8fe4c6 100644 --- a/i18n/keys.go +++ b/i18n/keys.go @@ -319,6 +319,7 @@ const ( MsgDistributorGroupAccessDenied = "distributor.group_access_denied" MsgDistributorGetChannelFailed = "distributor.get_channel_failed" MsgDistributorNoAvailableChannel = "distributor.no_available_channel" + MsgDistributorUnsupportedImageVariant = "distributor.unsupported_image_variant" MsgDistributorInvalidMidjourney = "distributor.invalid_midjourney_request" MsgDistributorInvalidParseModel = "distributor.invalid_request_parse_model" MsgDistributorUploadTimedOut = "distributor.upload_timed_out" diff --git a/i18n/locales/en.yaml b/i18n/locales/en.yaml index 085804bfbd92..9a2ab1af8a67 100644 --- a/i18n/locales/en.yaml +++ b/i18n/locales/en.yaml @@ -269,6 +269,7 @@ distributor.invalid_playground_request: "Invalid playground request: {{.Error}}" distributor.group_access_denied: "No permission to access this group" distributor.get_channel_failed: "Failed to get available channel for model {{.Model}} under group {{.Group}} (distributor): {{.Error}}" distributor.no_available_channel: "No available channel for model {{.Model}} under group {{.Group}} (distributor)" +distributor.unsupported_image_variant: "The requested image parameters are not supported for model {{.Model}} under group {{.Group}}" distributor.invalid_midjourney_request: "Invalid Midjourney request: {{.Error}}" distributor.invalid_request_parse_model: "Invalid request, unable to parse model" distributor.upload_timed_out: "Request body upload timed out before it completed. Please retry the request." diff --git a/i18n/locales/zh-CN.yaml b/i18n/locales/zh-CN.yaml index b72201ac1477..dba0a2f99b1a 100644 --- a/i18n/locales/zh-CN.yaml +++ b/i18n/locales/zh-CN.yaml @@ -270,6 +270,7 @@ distributor.invalid_playground_request: "无效的playground请求,{{.Error}}" distributor.group_access_denied: "无权访问该分组" distributor.get_channel_failed: "获取分组 {{.Group}} 下模型 {{.Model}} 的可用渠道失败(distributor):{{.Error}}" distributor.no_available_channel: "分组 {{.Group}} 下模型 {{.Model}} 无可用渠道(distributor)" +distributor.unsupported_image_variant: "分组 {{.Group}} 下模型 {{.Model}} 不支持请求的图片参数组合" distributor.invalid_midjourney_request: "无效的midjourney请求,{{.Error}}" distributor.invalid_request_parse_model: "无效的请求,无法解析模型" distributor.upload_timed_out: "请求体上传未完成并已超时,请重试。" diff --git a/i18n/locales/zh-TW.yaml b/i18n/locales/zh-TW.yaml index 67edc43f26e5..c4884fd0d077 100644 --- a/i18n/locales/zh-TW.yaml +++ b/i18n/locales/zh-TW.yaml @@ -270,6 +270,7 @@ distributor.invalid_playground_request: "無效的playground請求,{{.Error}}" distributor.group_access_denied: "無權存取該分組" distributor.get_channel_failed: "獲取分組 {{.Group}} 下模型 {{.Model}} 的可用管道失敗(distributor):{{.Error}}" distributor.no_available_channel: "分組 {{.Group}} 下模型 {{.Model}} 無可用管道(distributor)" +distributor.unsupported_image_variant: "分組 {{.Group}} 下模型 {{.Model}} 不支援請求的圖片參數組合" distributor.invalid_midjourney_request: "無效的midjourney請求,{{.Error}}" distributor.invalid_request_parse_model: "無效的請求,無法解析模型" distributor.upload_timed_out: "請求體上傳未完成並已逾時,請重試。" diff --git a/middleware/distributor.go b/middleware/distributor.go index f8b6858c49af..342ca553ac2d 100644 --- a/middleware/distributor.go +++ b/middleware/distributor.go @@ -234,6 +234,15 @@ func Distribute() func(c *gin.Context) { return } if channel == nil { + if imageRequirement != nil { + abortWithOpenAiMessage( + c, + http.StatusBadRequest, + i18n.T(c, i18n.MsgDistributorUnsupportedImageVariant, map[string]any{"Group": usingGroup, "Model": modelRequest.Model}), + types.ErrorCodeInvalidRequest, + ) + return + } abortWithOpenAiMessage(c, http.StatusServiceUnavailable, i18n.T(c, i18n.MsgDistributorNoAvailableChannel, map[string]any{"Group": usingGroup, "Model": modelRequest.Model}), types.ErrorCodeModelNotFound) return } diff --git a/middleware/distributor_image_routing_test.go b/middleware/distributor_image_routing_test.go new file mode 100644 index 000000000000..8fd7f674f650 --- /dev/null +++ b/middleware/distributor_image_routing_test.go @@ -0,0 +1,86 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/dto" + appI18n "github.com/QuantumNous/new-api/i18n" + "github.com/QuantumNous/new-api/model" + "github.com/QuantumNous/new-api/types" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDistributeRejectsUnsupportedVerifiedImageVariantAsInvalidRequest(t *testing.T) { + gin.SetMode(gin.TestMode) + require.NoError(t, appI18n.Init()) + oldMemoryCacheEnabled := common.MemoryCacheEnabled + common.MemoryCacheEnabled = true + model.ClearChannelCacheForTest() + t.Cleanup(func() { + model.ClearChannelCacheForTest() + common.MemoryCacheEnabled = oldMemoryCacheEnabled + }) + + priority := int64(10) + weight := uint(100) + channel := &model.Channel{ + Id: 31, + Type: constant.ChannelTypeOpenAI, + Status: common.ChannelStatusEnabled, + Name: "verified-auto-image", + Models: "gpt-image-2", + Group: "gpt pro", + Priority: &priority, + Weight: &weight, + } + channel.SetOtherSettings(dto.ChannelOtherSettings{ImageRouting: &dto.ImageRoutingConfig{ + Version: dto.ImageRoutingVersion1, + Profiles: []dto.ImageRoutingProfile{ + { + Model: "gpt-image-2", + Protocol: dto.ImageRoutingProtocolImagesGenerations, + UpstreamPath: "/v1/images/generations", + Operations: []dto.ImageOperation{dto.ImageOperationGeneration}, + Sizes: []string{"auto"}, + DefaultSize: "auto", + MaxOutputImages: 1, + AllowedCombinations: []dto.ImageRoutingCombination{{Operation: dto.ImageOperationGeneration, Size: "auto"}}, + VerificationStatus: dto.ImageRoutingVerificationProductionVerified, + }, + }, + }}) + model.SetChannelCacheForTest(map[int]*model.Channel{31: channel}, map[string]map[string][]int{ + "gpt pro": {"gpt-image-2": {31}}, + }) + + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Request = httptest.NewRequest( + http.MethodPost, + "/v1/images/generations", + strings.NewReader(`{"model":"gpt-image-2","prompt":"draw a cube","size":"1024x1024","n":1}`), + ) + ctx.Request.Header.Set("Content-Type", "application/json") + common.SetContextKey(ctx, constant.ContextKeyUsingGroup, "gpt pro") + common.SetContextKey(ctx, constant.ContextKeyUserGroup, "default") + + Distribute()(ctx) + + assert.Equal(t, http.StatusBadRequest, recorder.Code) + var response struct { + Error struct { + Code string `json:"code"` + } `json:"error"` + } + require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response)) + assert.Equal(t, string(types.ErrorCodeInvalidRequest), response.Error.Code) + _, selected := common.GetContextKey(ctx, constant.ContextKeyChannelId) + assert.False(t, selected) +}