diff --git a/backend/go.mod b/backend/go.mod index b86368117fb..5cfe682b514 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -50,12 +50,12 @@ require ( github.com/wechatpay-apiv3/wechatpay-go v0.2.21 github.com/zeromicro/go-zero v1.9.4 go.uber.org/zap v1.24.0 - golang.org/x/crypto v0.53.0 - golang.org/x/image v0.41.0 - golang.org/x/mod v0.37.0 - golang.org/x/net v0.56.0 - golang.org/x/sync v0.21.0 - golang.org/x/term v0.44.0 + golang.org/x/crypto v0.54.0 + golang.org/x/image v0.45.0 + golang.org/x/mod v0.38.0 + golang.org/x/net v0.57.0 + golang.org/x/sync v0.22.0 + golang.org/x/term v0.45.0 google.golang.org/grpc v1.82.1 google.golang.org/protobuf v1.36.11 gopkg.in/natefinch/lumberjack.v2 v2.2.1 @@ -202,10 +202,10 @@ require ( go.uber.org/multierr v1.9.0 // indirect golang.org/x/arch v0.3.0 // indirect golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect - golang.org/x/sys v0.46.0 // indirect - golang.org/x/text v0.39.0 // indirect + golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.41.0 // indirect golang.org/x/time v0.12.0 // indirect - golang.org/x/tools v0.47.0 // indirect + golang.org/x/tools v0.48.0 // indirect google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 // indirect gopkg.in/ini.v1 v1.67.0 // indirect modernc.org/libc v1.67.6 // indirect diff --git a/backend/go.sum b/backend/go.sum index cbb257f12f4..ed6150e5b99 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -562,13 +562,13 @@ golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDf golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOMs= golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8= golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM= -golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto= -golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio= +golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= +golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70= -golang.org/x/image v0.41.0 h1:8wS72eGJMJaBxK6okTzd4WaXumUlTVlb753MlsSvTCo= -golang.org/x/image v0.41.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA= +golang.org/x/image v0.45.0 h1:FMb1nTbH5H9vF55SriQHgFw5GnNL9Jg6L25BwXKzhB0= +golang.org/x/image v0.45.0/go.mod h1:n62x/7RqlwXDvGsSU4u6IUTUf6KghUZ9Bt7cG/T9Fx4= golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU= golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= @@ -578,8 +578,8 @@ golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c= golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c= -golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= -golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= +golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= +golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= @@ -601,8 +601,8 @@ golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44= golang.org/x/net v0.23.0/go.mod h1:JKghWKKOSdJwpW2GEx0Ja7fmaKnMsbu+MWVZTokSYmg= golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM= golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE= -golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o= -golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec= +golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= +golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -614,8 +614,8 @@ golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y= golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= -golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= -golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= @@ -648,8 +648,8 @@ golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= -golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw= -golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= @@ -662,8 +662,8 @@ golang.org/x/term v0.17.0/go.mod h1:lLRBjIVuehSbZlaOtGMbcMncT+aqLLLmKrsjNrUguwk= golang.org/x/term v0.18.0/go.mod h1:ILwASektA3OnRv7amZ1xhE/KTR+u50pbXfZ03+6Nx58= golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY= golang.org/x/term v0.21.0/go.mod h1:ooXLefLobQVslOqselCNF4SxFAaoS6KujMbsGzSDmX0= -golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc= -golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y= +golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0= +golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= @@ -674,8 +674,8 @@ golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE= golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= golang.org/x/text v0.16.0/go.mod h1:GhwF1Be+LQoKShO3cGOHzqOgRrGaYc9AvblQOmPVHnI= -golang.org/x/text v0.39.0 h1:UbZz4pLOvn600D6Oh6GGEI6VAmndrEBLv8/6BEXzyus= -golang.org/x/text v0.39.0/go.mod h1:3UwRclnC2g0TU9x8PZiyfOajCd1zaUNHF9cvqcQZ+ZM= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= golang.org/x/time v0.12.0 h1:ScB/8o8olJvc+CQPWrK3fPZNfh7qgwCrY0zJmoEQLSE= golang.org/x/time v0.12.0/go.mod h1:CDIdPxbZBQxdj6cxyCIdrNogrJKMJ7pr37NYpMcMDSg= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= @@ -690,8 +690,8 @@ golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU= golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58= golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk= -golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= -golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= +golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE= +golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index e3f19f335ee..b8d4c103188 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -764,6 +764,9 @@ func shouldFailoverOpenAIPassthroughResponse(account *Account, statusCode int, r if isOpenAIContextWindowError("", responseBody) { return false } + if isOpenAIShortInputPolicyError(statusCode, responseBody) { + return true + } if isOpenAIHTTPUpstreamAccessStateError(statusCode, "", responseBody) { return true } @@ -1480,6 +1483,9 @@ func openAIStreamFailedEventShouldFailover(payload []byte, message string) bool if isOpenAIContextWindowError(message, payload) { return false } + if isOpenAIShortInputPolicyError(http.StatusBadRequest, payload) { + return true + } if isOpenAIUpstreamAccessStateError(message, payload) { return true } @@ -1532,6 +1538,9 @@ func openAIStreamErrorEventShouldFailover(payload []byte, message string) bool { if isOpenAIContextWindowError(message, payload) { return false } + if isOpenAIShortInputPolicyError(http.StatusBadRequest, payload) { + return true + } if isOpenAIUpstreamAccessStateError(message, payload) { return true } diff --git a/backend/internal/service/openai_gateway_upstream_errors.go b/backend/internal/service/openai_gateway_upstream_errors.go index 70c000e3216..05cfa3612fa 100644 --- a/backend/internal/service/openai_gateway_upstream_errors.go +++ b/backend/internal/service/openai_gateway_upstream_errors.go @@ -260,6 +260,9 @@ func (s *OpenAIGatewayService) shouldFailoverOpenAIUpstreamResponse(statusCode i if isOpenAIContextWindowError(upstreamMsg, upstreamBody) { return false } + if isOpenAIShortInputPolicyError(statusCode, upstreamBody) { + return true + } if isOpenAIHTTPUpstreamAccessStateError(statusCode, upstreamMsg, upstreamBody) { return true } @@ -272,6 +275,27 @@ func (s *OpenAIGatewayService) shouldFailoverOpenAIUpstreamResponse(statusCode i return isOpenAITransientProcessingError(statusCode, upstreamMsg, upstreamBody) } +const ( + openAIShortInputPolicyReason = GatewayFailureReason("openai_short_input_policy") + openAIShortInputPolicyClientMessage = "Upstream does not support this request; please retry" +) + +// isOpenAIShortInputPolicyError recognizes a provider-specific policy that +// rejects otherwise valid short Responses requests. The structured code is +// authoritative; messages and arbitrary JSON fields may contain echoed user +// input and must not trigger account failover. +func isOpenAIShortInputPolicyError(statusCode int, upstreamBody []byte) bool { + if statusCode != http.StatusBadRequest || len(upstreamBody) == 0 || !gjson.ValidBytes(upstreamBody) { + return false + } + for _, path := range []string{"error.code", "response.error.code", "code"} { + if strings.EqualFold(strings.TrimSpace(gjson.GetBytes(upstreamBody, path).String()), "short_input_rejected") { + return true + } + } + return false +} + // OpenAIRequestBodyTooLargeClientMessage is the fixed downstream message used // after all account-specific request body limit failovers are exhausted. const OpenAIRequestBodyTooLargeClientMessage = "Request payload is too large" @@ -306,6 +330,15 @@ func newOpenAIUpstreamFailoverError( failoverErr.ClientStatusCode = http.StatusRequestEntityTooLarge failoverErr.ClientMessage = OpenAIRequestBodyTooLargeClientMessage } + if isOpenAIShortInputPolicyError(statusCode, responseBody) { + failoverErr.RetryableOnSameAccount = false + failoverErr.RequestScopedTransient = false + failoverErr.Scope = GatewayFailureScopeAccount + failoverErr.Reason = openAIShortInputPolicyReason + failoverErr.NextAccountAction = NextAccountRetry + failoverErr.ClientStatusCode = http.StatusBadGateway + failoverErr.ClientMessage = openAIShortInputPolicyClientMessage + } if isOpenAIHTTPUpstreamAccessStateError(statusCode, upstreamMsg, responseBody) { failoverErr.RetryableOnSameAccount = false failoverErr.RequestScopedTransient = false diff --git a/backend/internal/service/openai_image_mask.go b/backend/internal/service/openai_image_mask.go new file mode 100644 index 00000000000..6159b2c0151 --- /dev/null +++ b/backend/internal/service/openai_image_mask.go @@ -0,0 +1,138 @@ +package service + +import ( + "bytes" + "encoding/base64" + "fmt" + "image" + "image/color" + "image/png" + "strings" + + _ "image/jpeg" + + xdraw "golang.org/x/image/draw" + _ "golang.org/x/image/webp" +) + +const openAIImageMaskMaxPixels = 4096 * 4096 + +// openAIImageMaskCompositor enforces the Images API mask contract on OAuth +// responses. The upstream Responses image tool treats masks as guidance, so +// protected pixels must be restored before the generated image reaches the +// client. +type openAIImageMaskCompositor struct { + source image.Image + mask image.Image + width int + height int +} + +func newOpenAIImageMaskCompositor(parsed *OpenAIImagesRequest) (*openAIImageMaskCompositor, error) { + if parsed == nil || !parsed.IsEdits() || !parsed.Multipart || parsed.MaskUpload == nil { + return nil, nil + } + if len(parsed.Uploads) == 0 { + return nil, fmt.Errorf("masked image edits require an image file") + } + + source, sourceFormat, err := decodeOpenAIImageForMask(parsed.Uploads[0].Data, "image") + if err != nil { + return nil, err + } + if sourceFormat != "png" && sourceFormat != "jpeg" && sourceFormat != "webp" { + return nil, fmt.Errorf("unsupported image format %q for masked edit", sourceFormat) + } + mask, maskFormat, err := decodeOpenAIImageForMask(parsed.MaskUpload.Data, "mask") + if err != nil { + return nil, err + } + if maskFormat != "png" { + return nil, fmt.Errorf("mask must be a PNG image") + } + + width := source.Bounds().Dx() + height := source.Bounds().Dy() + if mask.Bounds().Dx() != width || mask.Bounds().Dy() != height { + return nil, fmt.Errorf("mask dimensions must match the source image") + } + return &openAIImageMaskCompositor{source: source, mask: mask, width: width, height: height}, nil +} + +func decodeOpenAIImageForMask(data []byte, field string) (image.Image, string, error) { + if len(data) == 0 { + return nil, "", fmt.Errorf("%s image is empty", field) + } + cfg, format, err := image.DecodeConfig(bytes.NewReader(data)) + if err != nil { + return nil, "", fmt.Errorf("decode %s image metadata: %w", field, err) + } + if cfg.Width <= 0 || cfg.Height <= 0 || cfg.Width > openAIImageMaskMaxPixels/cfg.Height { + return nil, "", fmt.Errorf("%s image exceeds the masked edit pixel limit", field) + } + decoded, _, err := image.Decode(bytes.NewReader(data)) + if err != nil { + return nil, "", fmt.Errorf("decode %s image: %w", field, err) + } + return decoded, strings.ToLower(strings.TrimSpace(format)), nil +} + +func (c *openAIImageMaskCompositor) applyResult(result *openAIResponsesImageResult) error { + if c == nil || result == nil { + return nil + } + raw := normalizeOpenAIImageBase64(result.Result) + generatedBytes, err := base64.StdEncoding.DecodeString(raw) + if err != nil { + return fmt.Errorf("decode generated image base64: %w", err) + } + generated, _, err := decodeOpenAIImageForMask(generatedBytes, "generated") + if err != nil { + return err + } + + generated = c.resizeGenerated(generated) + out := image.NewNRGBA(image.Rect(0, 0, c.width, c.height)) + sourceBounds := c.source.Bounds() + maskBounds := c.mask.Bounds() + generatedBounds := generated.Bounds() + for y := 0; y < c.height; y++ { + for x := 0; x < c.width; x++ { + src, ok := color.NRGBAModel.Convert(c.source.At(sourceBounds.Min.X+x, sourceBounds.Min.Y+y)).(color.NRGBA) + if !ok { + return fmt.Errorf("convert source pixel to NRGBA") + } + gen, ok := color.NRGBAModel.Convert(generated.At(generatedBounds.Min.X+x, generatedBounds.Min.Y+y)).(color.NRGBA) + if !ok { + return fmt.Errorf("convert generated pixel to NRGBA") + } + _, _, _, alpha16 := c.mask.At(maskBounds.Min.X+x, maskBounds.Min.Y+y).RGBA() + keep := uint32(alpha16 >> 8) + edit := uint32(255) - keep + out.SetNRGBA(x, y, color.NRGBA{ + R: uint8((uint32(src.R)*keep + uint32(gen.R)*edit + 127) / 255), + G: uint8((uint32(src.G)*keep + uint32(gen.G)*edit + 127) / 255), + B: uint8((uint32(src.B)*keep + uint32(gen.B)*edit + 127) / 255), + A: uint8((uint32(src.A)*keep + uint32(gen.A)*edit + 127) / 255), + }) + } + } + + var encoded bytes.Buffer + if err := png.Encode(&encoded, out); err != nil { + return fmt.Errorf("encode masked image: %w", err) + } + result.Result = base64.StdEncoding.EncodeToString(encoded.Bytes()) + result.OutputFormat = "png" + result.Size = fmt.Sprintf("%dx%d", c.width, c.height) + return nil +} + +func (c *openAIImageMaskCompositor) resizeGenerated(generated image.Image) image.Image { + if generated.Bounds().Dx() == c.width && generated.Bounds().Dy() == c.height { + return generated + } + resized := image.NewNRGBA(image.Rect(0, 0, c.width, c.height)) + xdraw.CatmullRom.Scale(resized, resized.Bounds(), generated, generated.Bounds(), xdraw.Over, nil) + return resized +} diff --git a/backend/internal/service/openai_image_mask_test.go b/backend/internal/service/openai_image_mask_test.go new file mode 100644 index 00000000000..63bcfd763ea --- /dev/null +++ b/backend/internal/service/openai_image_mask_test.go @@ -0,0 +1,190 @@ +package service + +import ( + "bytes" + "context" + "encoding/base64" + "image" + "image/color" + "image/png" + "testing" + + "github.com/stretchr/testify/require" +) + +func openAIImageMaskTestPNG(t *testing.T, width, height int, pixel func(x, y int) color.NRGBA) []byte { + t.Helper() + img := image.NewNRGBA(image.Rect(0, 0, width, height)) + for y := 0; y < height; y++ { + for x := 0; x < width; x++ { + img.SetNRGBA(x, y, pixel(x, y)) + } + } + var out bytes.Buffer + require.NoError(t, png.Encode(&out, img)) + return out.Bytes() +} + +func openAIImageMaskTestSolidPNG(t *testing.T, width, height int, pixel color.NRGBA) []byte { + t.Helper() + return openAIImageMaskTestPNG(t, width, height, func(_, _ int) color.NRGBA { return pixel }) +} + +func openAIImageMaskTestResult(t *testing.T, data []byte, format string) openAIResponsesImageResult { + t.Helper() + return openAIResponsesImageResult{ + Result: base64.StdEncoding.EncodeToString(data), + OutputFormat: format, + } +} + +func openAIImageMaskTestDecodeResult(t *testing.T, result openAIResponsesImageResult) image.Image { + t.Helper() + decoded, err := base64.StdEncoding.DecodeString(result.Result) + require.NoError(t, err) + img, format, err := image.Decode(bytes.NewReader(decoded)) + require.NoError(t, err) + require.Equal(t, "png", format) + return img +} + +func TestOpenAIImageMaskCompositorEnforcesAlphaMask(t *testing.T) { + source := openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{R: 200, G: 10, B: 20, A: 255}) + mask := openAIImageMaskTestPNG(t, 2, 2, func(x, y int) color.NRGBA { + alpha := uint8(255) + if x == 1 && y == 0 { + alpha = 0 + } + if x == 0 && y == 1 { + alpha = 128 + } + return color.NRGBA{A: alpha} + }) + parsed := &OpenAIImagesRequest{ + Endpoint: openAIImagesEditsEndpoint, + Multipart: true, + Uploads: []OpenAIImagesUpload{{Data: source}}, + MaskUpload: &OpenAIImagesUpload{ + Data: mask, + }, + } + compositor, err := newOpenAIImageMaskCompositor(parsed) + require.NoError(t, err) + + result := openAIImageMaskTestResult(t, openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{R: 10, G: 110, B: 220, A: 255}), "webp") + require.NoError(t, compositor.applyResult(&result)) + require.Equal(t, "png", result.OutputFormat) + require.Equal(t, "2x2", result.Size) + + got := openAIImageMaskTestDecodeResult(t, result) + require.Equal(t, color.NRGBA{R: 200, G: 10, B: 20, A: 255}, color.NRGBAModel.Convert(got.At(0, 0))) + require.Equal(t, color.NRGBA{R: 10, G: 110, B: 220, A: 255}, color.NRGBAModel.Convert(got.At(1, 0))) + require.Equal(t, color.NRGBA{R: 105, G: 60, B: 120, A: 255}, color.NRGBAModel.Convert(got.At(0, 1))) +} + +func TestOpenAIImageMaskCompositorRejectsDimensionMismatch(t *testing.T) { + parsed := &OpenAIImagesRequest{ + Endpoint: openAIImagesEditsEndpoint, + Multipart: true, + Uploads: []OpenAIImagesUpload{{ + Data: openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{A: 255}), + }}, + MaskUpload: &OpenAIImagesUpload{ + Data: openAIImageMaskTestSolidPNG(t, 1, 1, color.NRGBA{A: 255}), + }, + } + + compositor, err := newOpenAIImageMaskCompositor(parsed) + require.Nil(t, compositor) + require.ErrorContains(t, err, "mask dimensions must match") +} + +func TestOpenAIImageMaskCompositorResizesGeneratedImageToSource(t *testing.T) { + source := openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{R: 200, A: 255}) + mask := openAIImageMaskTestPNG(t, 2, 2, func(x, _ int) color.NRGBA { + if x == 0 { + return color.NRGBA{A: 255} + } + return color.NRGBA{A: 0} + }) + compositor, err := newOpenAIImageMaskCompositor(&OpenAIImagesRequest{ + Endpoint: openAIImagesEditsEndpoint, + Multipart: true, + Uploads: []OpenAIImagesUpload{{Data: source}}, + MaskUpload: &OpenAIImagesUpload{ + Data: mask, + }, + }) + require.NoError(t, err) + + result := openAIImageMaskTestResult(t, openAIImageMaskTestSolidPNG(t, 1, 1, color.NRGBA{B: 220, A: 255}), "png") + require.NoError(t, compositor.applyResult(&result)) + + got := openAIImageMaskTestDecodeResult(t, result) + require.Equal(t, image.Rect(0, 0, 2, 2), got.Bounds()) + require.Equal(t, color.NRGBA{R: 200, A: 255}, color.NRGBAModel.Convert(got.At(0, 1))) + require.Equal(t, color.NRGBA{B: 220, A: 255}, color.NRGBAModel.Convert(got.At(1, 1))) +} + +func TestOpenAIImageMaskCompositorUsesFirstImageForMultiImageEdit(t *testing.T) { + first := openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{R: 1, A: 255}) + second := openAIImageMaskTestSolidPNG(t, 3, 3, color.NRGBA{G: 1, A: 255}) + parsed := &OpenAIImagesRequest{ + Endpoint: openAIImagesEditsEndpoint, + Multipart: true, + Uploads: []OpenAIImagesUpload{{Data: first}, {Data: second}}, + MaskUpload: &OpenAIImagesUpload{ + Data: openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{A: 255}), + }, + } + + compositor, err := newOpenAIImageMaskCompositor(parsed) + require.NoError(t, err) + require.NotNil(t, compositor) + require.Equal(t, 2, compositor.width) + require.Equal(t, 2, compositor.height) +} + +func TestForwardOpenAIImagesOAuthRejectsInvalidMaskBeforeCredentials(t *testing.T) { + source := openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{A: 255}) + tests := []struct { + name string + mask []byte + wantErrMsg string + }{ + { + name: "invalid PNG", + mask: []byte("not-an-image"), + wantErrMsg: "decode mask image metadata", + }, + { + name: "dimension mismatch", + mask: openAIImageMaskTestSolidPNG(t, 1, 1, color.NRGBA{A: 255}), + wantErrMsg: "mask dimensions must match", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + parsed := &OpenAIImagesRequest{ + Endpoint: openAIImagesEditsEndpoint, + Model: "gpt-image-2", + Multipart: true, + Uploads: []OpenAIImagesUpload{{Data: source}}, + MaskUpload: &OpenAIImagesUpload{ + Data: tt.mask, + }, + } + result, err := (&OpenAIGatewayService{}).forwardOpenAIImagesOAuth( + context.Background(), + nil, + &Account{Type: AccountTypeOAuth}, + parsed, + "", + ) + + require.Nil(t, result) + require.ErrorContains(t, err, tt.wantErrMsg) + }) + } +} diff --git a/backend/internal/service/openai_images.go b/backend/internal/service/openai_images.go index c661aade367..276f3c041c8 100644 --- a/backend/internal/service/openai_images.go +++ b/backend/internal/service/openai_images.go @@ -1342,7 +1342,8 @@ func normalizeOpenAIImageBase64(raw string) string { } } raw = strings.TrimSpace(raw) - raw = strings.TrimRight(raw, "=") + strings.Repeat("=", (4-len(raw)%4)%4) + raw = strings.TrimRight(raw, "=") + raw += strings.Repeat("=", (4-len(raw)%4)%4) if raw == "" { return "" } diff --git a/backend/internal/service/openai_images_incomplete_test.go b/backend/internal/service/openai_images_incomplete_test.go index ebd3c0e5277..b89f252b83e 100644 --- a/backend/internal/service/openai_images_incomplete_test.go +++ b/backend/internal/service/openai_images_incomplete_test.go @@ -106,7 +106,7 @@ func TestImagesOAuthNonStreaming_CompletedNoImageTriggersSameAccountRetry(t *tes } svc := &OpenAIGatewayService{} - _, _, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2") + _, _, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2", nil) if err == nil { t.Fatal("completed-but-no-image should return an error") @@ -140,7 +140,7 @@ func TestImagesOAuthNonStreaming_ContentRefusalReturns400NoRetry(t *testing.T) { resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: io.NopCloser(strings.NewReader(upstreamSSE))} svc := &OpenAIGatewayService{} - _, _, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2") + _, _, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2", nil) if err == nil { t.Fatal("content refusal should return an error") @@ -176,7 +176,7 @@ func TestImagesOAuthNonStreaming_TextFallbackReturnsCapabilityError(t *testing.T resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: io.NopCloser(strings.NewReader(upstreamSSE))} svc := &OpenAIGatewayService{} - _, _, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2") + _, _, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2", nil) var imgErr *OpenAIImagesUpstreamError if !errors.As(err, &imgErr) { @@ -202,7 +202,7 @@ func TestImagesOAuthStreaming_TextFallbackReturnsCapabilityError(t *testing.T) { resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: io.NopCloser(strings.NewReader(upstreamSSE))} svc := &OpenAIGatewayService{} - _, _, _, _, err := svc.handleOpenAIImagesOAuthStreamingResponse(resp, c, time.Now(), "b64_json", "image_generation", "gpt-image-2") + _, _, _, _, err := svc.handleOpenAIImagesOAuthStreamingResponse(resp, c, time.Now(), "b64_json", "image_generation", "gpt-image-2", nil) var imgErr *OpenAIImagesUpstreamError if !errors.As(err, &imgErr) { @@ -233,7 +233,7 @@ func TestImagesOAuthStreaming_SplitSafetyRefusalReturns400(t *testing.T) { resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: io.NopCloser(strings.NewReader(upstreamSSE))} svc := &OpenAIGatewayService{} - _, _, _, _, err := svc.handleOpenAIImagesOAuthStreamingResponse(resp, c, time.Now(), "b64_json", "image_generation", "gpt-image-2") + _, _, _, _, err := svc.handleOpenAIImagesOAuthStreamingResponse(resp, c, time.Now(), "b64_json", "image_generation", "gpt-image-2", nil) var imgErr *OpenAIImagesUpstreamError if !errors.As(err, &imgErr) { diff --git a/backend/internal/service/openai_images_json_keepalive_test.go b/backend/internal/service/openai_images_json_keepalive_test.go index a7207b3adb7..809dac50f9f 100644 --- a/backend/internal/service/openai_images_json_keepalive_test.go +++ b/backend/internal/service/openai_images_json_keepalive_test.go @@ -140,7 +140,7 @@ func TestOpenAIImagesJSONKeepalive_KeepsOAuthNonStreamResponseValid(t *testing.T Body: reader, } svc := &OpenAIGatewayService{} - _, imageCount, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2") + _, imageCount, _, err := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2", nil) stop() require.NoError(t, err) diff --git a/backend/internal/service/openai_images_responses.go b/backend/internal/service/openai_images_responses.go index 9374a6f7c7a..782b3649911 100644 --- a/backend/internal/service/openai_images_responses.go +++ b/backend/internal/service/openai_images_responses.go @@ -1291,6 +1291,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthNonStreamingResponse( c *gin.Context, responseFormat string, fallbackModel string, + maskCompositor *openAIImageMaskCompositor, ) (OpenAIUsage, int, []string, error) { body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) if err != nil { @@ -1339,6 +1340,14 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthNonStreamingResponse( if strings.TrimSpace(firstMeta.Model) == "" { firstMeta.Model = strings.TrimSpace(fallbackModel) } + for i := range results { + if err := maskCompositor.applyResult(&results[i]); err != nil { + return OpenAIUsage{}, 0, nil, err + } + } + if len(results) > 0 { + mergeOpenAIResponsesImageMeta(&firstMeta, results[0]) + } responseBody, err := buildOpenAIImagesAPIResponse(results, createdAt, usageRaw, firstMeta, responseFormat) if err != nil { @@ -1356,6 +1365,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( responseFormat string, streamPrefix string, fallbackModel string, + maskCompositor *openAIImageMaskCompositor, ) (OpenAIUsage, int, []string, *int, error) { responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) c.Header("Content-Type", "text/event-stream") @@ -1433,6 +1443,17 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( OutputFormat: strings.TrimSpace(gjson.GetBytes(dataBytes, "output_format").String()), Background: strings.TrimSpace(gjson.GetBytes(dataBytes, "background").String()), }) + partialResult := openAIResponsesImageResult{Result: b64} + if err := maskCompositor.applyResult(&partialResult); err != nil { + processDataErr = err + processDataDone = true + return + } + if maskCompositor != nil { + b64 = partialResult.Result + partialMeta.OutputFormat = partialResult.OutputFormat + partialMeta.Size = partialResult.Size + } payload := buildOpenAIImagesStreamPartialPayload( eventName, b64, @@ -1483,6 +1504,13 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( mergeOpenAIResponsesImageMeta(&img, streamMeta) appendOpenAIResponsesImageResultDedup(&finalResults, finalSeen, "", img) } + for i := range finalResults { + if err := maskCompositor.applyResult(&finalResults[i]); err != nil { + processDataErr = err + processDataDone = true + return + } + } reconcileOpenAIResponsesImageResultSizes(finalResults, nil) if len(finalResults) == 0 { textFallbackErr := openAIImagesTextFallbackErrorForText(fallbackText.String()) @@ -1561,6 +1589,9 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( finalResults := append([]openAIResponsesImageResult(nil), pendingResults...) for i := range finalResults { mergeOpenAIResponsesImageMeta(&finalResults[i], streamMeta) + if err := maskCompositor.applyResult(&finalResults[i]); err != nil { + return err + } } reconcileOpenAIResponsesImageResultSizes(finalResults, nil) for _, img := range finalResults { @@ -1755,6 +1786,10 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth( if err := validateOpenAIImagesModel(requestModel); err != nil { return nil, err } + maskCompositor, err := newOpenAIImageMaskCompositor(parsed) + if err != nil { + return nil, err + } logger.LegacyPrintf( "service.openai_gateway", "[OpenAI] Images request routing request_model=%s endpoint=%s account_type=%s uploads=%d", @@ -1854,7 +1889,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth( // keepalive 心跳字节,避免 failover 第 2 轮起把上一轮心跳残留误判为已写响应。 writerSizeBeforeResponse := OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) if parsed.Stream { - usage, imageCount, imageOutputSizes, firstTokenMs, err = s.handleOpenAIImagesOAuthStreamingResponse(resp, c, startTime, parsed.ResponseFormat, openAIImagesStreamPrefix(parsed), requestModel) + usage, imageCount, imageOutputSizes, firstTokenMs, err = s.handleOpenAIImagesOAuthStreamingResponse(resp, c, startTime, parsed.ResponseFormat, openAIImagesStreamPrefix(parsed), requestModel, maskCompositor) if err != nil { if imageCount > 0 { return &OpenAIForwardResult{ @@ -1884,7 +1919,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth( ) } } else { - usage, imageCount, imageOutputSizes, err = s.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, parsed.ResponseFormat, requestModel) + usage, imageCount, imageOutputSizes, err = s.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, parsed.ResponseFormat, requestModel, maskCompositor) if err != nil { return nil, s.handleOpenAIImagesOAuthResponseError( upstreamCtx, diff --git a/backend/internal/service/openai_images_test.go b/backend/internal/service/openai_images_test.go index f06c213f464..9c1c7886277 100644 --- a/backend/internal/service/openai_images_test.go +++ b/backend/internal/service/openai_images_test.go @@ -3,8 +3,10 @@ package service import ( "bytes" "context" + "encoding/base64" "errors" "fmt" + "image/color" "io" "mime/multipart" "net/http" @@ -66,6 +68,26 @@ func TestOpenAIGatewayServiceParseOpenAIImagesRequest_JSON(t *testing.T) { require.False(t, parsed.Multipart) } +func TestNormalizeOpenAIImageBase64PreservesValidPadding(t *testing.T) { + tests := []struct { + name string + raw string + want string + }{ + {name: "single padding", raw: "aGk=", want: "aGk="}, + {name: "double padding", raw: "aA==", want: "aA=="}, + {name: "unpadded", raw: "aGk", want: "aGk="}, + {name: "data URL", raw: "data:image/png;base64,aA==", want: "aA=="}, + {name: "invalid", raw: "not base64!", want: ""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, normalizeOpenAIImageBase64(tt.raw)) + }) + } +} + func TestOpenAIGatewayServiceParseOpenAIImagesRequest_MultipartEdit(t *testing.T) { gin.SetMode(gin.TestMode) @@ -1156,7 +1178,7 @@ func TestOpenAIImagesOAuthBodyReadTransportErrorFailover(t *testing.T) { account := &Account{ID: 5400, Name: "openai-oauth", Platform: PlatformOpenAI, Type: AccountTypeOAuth} svc := &OpenAIGatewayService{} - _, _, _, readErr := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2") + _, _, _, readErr := svc.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, "b64_json", "gpt-image-2", nil) require.Error(t, readErr) err := svc.handleOpenAIImagesOAuthResponseError(context.Background(), c, account, "gpt-image-2", "https://api.openai.com/v1/responses", resp, OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c), readErr) @@ -1787,6 +1809,15 @@ func TestOpenAIGatewayServiceForwardImages_APIKeyStreamingDrainsAfterClientDisco func TestOpenAIGatewayServiceForwardImages_OAuthEditsMultipartUsesResponsesAPI(t *testing.T) { gin.SetMode(gin.TestMode) + sourcePNG := openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{R: 220, G: 20, B: 30, A: 255}) + maskPNG := openAIImageMaskTestPNG(t, 2, 2, func(x, y int) color.NRGBA { + if x == 1 && y == 1 { + return color.NRGBA{A: 0} + } + return color.NRGBA{A: 255} + }) + generatedPNG := openAIImageMaskTestSolidPNG(t, 2, 2, color.NRGBA{R: 10, G: 40, B: 230, A: 255}) + generatedB64 := base64.StdEncoding.EncodeToString(generatedPNG) var body bytes.Buffer writer := multipart.NewWriter(&body) @@ -1801,7 +1832,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthEditsMultipartUsesResponsesAPI(t imageHeader.Set("Content-Type", "image/png") imagePart, err := writer.CreatePart(imageHeader) require.NoError(t, err) - _, err = imagePart.Write([]byte("png-image-content")) + _, err = imagePart.Write(sourcePNG) require.NoError(t, err) maskHeader := make(textproto.MIMEHeader) @@ -1809,7 +1840,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthEditsMultipartUsesResponsesAPI(t maskHeader.Set("Content-Type", "image/png") maskPart, err := writer.CreatePart(maskHeader) require.NoError(t, err) - _, err = maskPart.Write([]byte("png-mask-content")) + _, err = maskPart.Write(maskPNG) require.NoError(t, err) require.NoError(t, writer.Close()) @@ -1833,7 +1864,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthEditsMultipartUsesResponsesAPI(t "X-Request-Id": []string{"req_img_edit_123"}, }, Body: io.NopCloser(strings.NewReader( - "data: {\"type\":\"response.completed\",\"response\":{\"created_at\":1710000002,\"usage\":{\"input_tokens\":13,\"output_tokens\":21,\"output_tokens_details\":{\"image_tokens\":8}},\"tool_usage\":{\"image_gen\":{\"images\":1}},\"output\":[{\"type\":\"image_generation_call\",\"result\":\"ZWRpdGVk\",\"revised_prompt\":\"replace background with aurora\",\"output_format\":\"webp\",\"quality\":\"high\"}]}}\n\n" + + fmt.Sprintf("data: {\"type\":\"response.completed\",\"response\":{\"created_at\":1710000002,\"usage\":{\"input_tokens\":13,\"output_tokens\":21,\"output_tokens_details\":{\"image_tokens\":8}},\"tool_usage\":{\"image_gen\":{\"images\":1}},\"output\":[{\"type\":\"image_generation_call\",\"result\":%q,\"revised_prompt\":\"replace background with aurora\",\"output_format\":\"webp\",\"quality\":\"high\"}]}}\n\n", generatedB64) + "data: [DONE]\n\n", )), }, @@ -1861,7 +1892,10 @@ func TestOpenAIGatewayServiceForwardImages_OAuthEditsMultipartUsesResponsesAPI(t require.True(t, strings.HasPrefix(gjson.GetBytes(upstream.lastBody, "input.0.content.1.image_url").String(), "data:image/png;base64,")) require.True(t, strings.HasPrefix(gjson.GetBytes(upstream.lastBody, "tools.0.input_image_mask.image_url").String(), "data:image/png;base64,")) require.Equal(t, "replace background with aurora", gjson.GetBytes(upstream.lastBody, "input.0.content.0.text").String()) - require.Equal(t, "ZWRpdGVk", gjson.Get(rec.Body.String(), "data.0.b64_json").String()) + masked := openAIImageMaskTestDecodeResult(t, openAIResponsesImageResult{Result: gjson.Get(rec.Body.String(), "data.0.b64_json").String()}) + require.Equal(t, color.NRGBA{R: 220, G: 20, B: 30, A: 255}, color.NRGBAModel.Convert(masked.At(0, 0))) + require.Equal(t, color.NRGBA{R: 10, G: 40, B: 230, A: 255}, color.NRGBAModel.Convert(masked.At(1, 1))) + require.Equal(t, "png", gjson.Get(rec.Body.String(), "output_format").String()) require.Equal(t, "replace background with aurora", gjson.Get(rec.Body.String(), "data.0.revised_prompt").String()) } diff --git a/backend/internal/service/openai_short_input_failover_test.go b/backend/internal/service/openai_short_input_failover_test.go new file mode 100644 index 00000000000..2680d00896b --- /dev/null +++ b/backend/internal/service/openai_short_input_failover_test.go @@ -0,0 +1,69 @@ +package service + +import ( + "net/http" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestOpenAIShortInputPolicyClassification(t *testing.T) { + policyBody := []byte(`{"error":{"code":"short_input_rejected","message":"Upstream rejected illegal short-input distillation or heartbeat probing.","type":"invalid_request_error"}}`) + streamBody := []byte(`{"type":"response.failed","response":{"error":{"code":"short_input_rejected","message":"request rejected","type":"invalid_request_error"}}}`) + svc := &OpenAIGatewayService{} + + require.True(t, isOpenAIShortInputPolicyError(http.StatusBadRequest, policyBody)) + require.True(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadRequest, "request rejected", policyBody)) + require.True(t, shouldFailoverOpenAIPassthroughResponse(&Account{Type: AccountTypeAPIKey}, http.StatusBadRequest, policyBody)) + require.True(t, openAIStreamFailedEventShouldFailover(streamBody, "request rejected")) + require.True(t, openAIStreamErrorEventShouldFailover(streamBody, "request rejected")) + + failoverErr := newOpenAIUpstreamFailoverError(http.StatusBadRequest, nil, policyBody, "request rejected", true) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, openAIShortInputPolicyReason, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.False(t, failoverErr.RetryableOnSameAccount) + require.False(t, failoverErr.RequestScopedTransient) + require.Equal(t, http.StatusBadGateway, failoverErr.ClientStatusCode) + require.Equal(t, openAIShortInputPolicyClientMessage, failoverErr.ClientMessage) +} + +func TestOpenAIShortInputPolicyClassificationRequiresStructuredCode(t *testing.T) { + tests := []struct { + name string + statusCode int + body string + }{ + { + name: "message only", + statusCode: http.StatusBadRequest, + body: `{"error":{"code":"invalid_request_error","message":"short_input_rejected"}}`, + }, + { + name: "echoed user input", + statusCode: http.StatusBadRequest, + body: `{"error":{"code":"unknown_parameter","message":"invalid input"},"echo":{"error":{"code":"short_input_rejected"}}}`, + }, + { + name: "wrong status", + statusCode: http.StatusServiceUnavailable, + body: `{"error":{"code":"short_input_rejected"}}`, + }, + { + name: "plain text", + statusCode: http.StatusBadRequest, + body: `short_input_rejected`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body := []byte(tt.body) + require.False(t, isOpenAIShortInputPolicyError(tt.statusCode, body)) + if tt.statusCode == http.StatusBadRequest { + require.False(t, (&OpenAIGatewayService{}).shouldFailoverOpenAIUpstreamResponse(tt.statusCode, "", body)) + require.False(t, shouldFailoverOpenAIPassthroughResponse(&Account{Type: AccountTypeAPIKey}, tt.statusCode, body)) + } + }) + } +}