diff --git a/relay/channel/ali/adaptor.go b/relay/channel/ali/adaptor.go index 2cf4bf96..14d77156 100644 --- a/relay/channel/ali/adaptor.go +++ b/relay/channel/ali/adaptor.go @@ -110,15 +110,15 @@ func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) { case constant.RelayModeResponses: fullRequestURL = fmt.Sprintf("%s/api/v2/apps/protocols/compatible-mode/v1/responses", info.ChannelBaseUrl) case constant.RelayModeImagesGenerations: - if isSyncImageModel(info.OriginModelName) { + if isSyncImageModel(info.UpstreamModelName) { fullRequestURL = fmt.Sprintf("%s/api/v1/services/aigc/multimodal-generation/generation", info.ChannelBaseUrl) } else { fullRequestURL = fmt.Sprintf("%s/api/v1/services/aigc/text2image/image-synthesis", info.ChannelBaseUrl) } case constant.RelayModeImagesEdits: - if isOldWanModel(info.OriginModelName) { + if isOldWanModel(info.UpstreamModelName) { fullRequestURL = fmt.Sprintf("%s/api/v1/services/aigc/image2image/image-synthesis", info.ChannelBaseUrl) - } else if isWanModel(info.OriginModelName) { + } else if isWanModel(info.UpstreamModelName) { fullRequestURL = fmt.Sprintf("%s/api/v1/services/aigc/image-generation/generation", info.ChannelBaseUrl) } else { fullRequestURL = fmt.Sprintf("%s/api/v1/services/aigc/multimodal-generation/generation", info.ChannelBaseUrl) @@ -143,14 +143,14 @@ func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Header, info *rel req.Set("X-DashScope-Plugin", c.GetString("plugin")) } if info.RelayMode == constant.RelayModeImagesGenerations { - if isSyncImageModel(info.OriginModelName) { + if isSyncImageModel(info.UpstreamModelName) { } else { req.Set("X-DashScope-Async", "enable") } } if info.RelayMode == constant.RelayModeImagesEdits { - if isWanModel(info.OriginModelName) { + if isWanModel(info.UpstreamModelName) { req.Set("X-DashScope-Async", "enable") } req.Set("Content-Type", "application/json") @@ -183,7 +183,7 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest) (any, error) { if info.RelayMode == constant.RelayModeImagesGenerations { - if isSyncImageModel(info.OriginModelName) { + if isSyncImageModel(info.UpstreamModelName) { a.IsSyncImageModel = true } aliRequest, err := oaiImage2AliImageRequest(info, request, a.IsSyncImageModel) @@ -192,11 +192,11 @@ func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInf } return aliRequest, nil } else if info.RelayMode == constant.RelayModeImagesEdits { - if isOldWanModel(info.OriginModelName) { + if isOldWanModel(info.UpstreamModelName) { return oaiFormEdit2WanxImageEdit(c, info, request) } - if isSyncImageModel(info.OriginModelName) { - if isWanModel(info.OriginModelName) { + if isSyncImageModel(info.UpstreamModelName) { + if isWanModel(info.UpstreamModelName) { a.IsSyncImageModel = false } else { a.IsSyncImageModel = true diff --git a/relay/channel/ali/adaptor_test.go b/relay/channel/ali/adaptor_test.go index a8b87140..08bc959a 100644 --- a/relay/channel/ali/adaptor_test.go +++ b/relay/channel/ali/adaptor_test.go @@ -2,11 +2,13 @@ package ali import ( "encoding/json" + "net/http" "net/http/httptest" "testing" "github.com/QuantumNous/new-api/common" relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/relay/constant" relayhelper "github.com/QuantumNous/new-api/relay/helper" "github.com/QuantumNous/new-api/relaykit/dto" "github.com/gin-gonic/gin" @@ -126,3 +128,34 @@ func TestConvertOpenAIRequestPreservesExplicitZeroForMappedQwenModel(t *testing. assert.True(t, value.Exists()) assert.Equal(t, int64(0), value.Int()) } + +func TestMappedAliImageModelUsesUpstreamProtocol(t *testing.T) { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + + info := &relaycommon.RelayInfo{ + RelayMode: constant.RelayModeImagesGenerations, + OriginModelName: "customer-image-model", + ChannelMeta: &relaycommon.ChannelMeta{ + ChannelBaseUrl: "https://dashscope.aliyuncs.com", + UpstreamModelName: "qwen-image-3.0-pro", + }, + } + + adaptor := &Adaptor{} + url, err := adaptor.GetRequestURL(info) + require.NoError(t, err) + assert.Equal(t, "https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation", url) + + header := http.Header{} + require.NoError(t, adaptor.SetupRequestHeader(c, &header, info)) + assert.Empty(t, header.Get("X-DashScope-Async")) + + converted, err := adaptor.ConvertImageRequest(c, info, dto.ImageRequest{ + Model: info.UpstreamModelName, + Prompt: "poster", + }) + require.NoError(t, err) + assert.True(t, adaptor.IsSyncImageModel) + assert.IsType(t, &AliImageRequest{}, converted) +}