Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions i18n/keys.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
1 change: 1 addition & 0 deletions i18n/locales/en.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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."
Expand Down
1 change: 1 addition & 0 deletions i18n/locales/zh-CN.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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: "请求体上传未完成并已超时,请重试。"
Expand Down
1 change: 1 addition & 0 deletions i18n/locales/zh-TW.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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: "請求體上傳未完成並已逾時,請重試。"
Expand Down
9 changes: 9 additions & 0 deletions middleware/distributor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
86 changes: 86 additions & 0 deletions middleware/distributor_image_routing_test.go
Original file line number Diff line number Diff line change
@@ -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)
}