Skip to content
Closed
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
97 changes: 97 additions & 0 deletions dto/openai_image.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package dto

import (
"encoding/json"
"fmt"
"reflect"
"strings"

Expand Down Expand Up @@ -182,3 +183,99 @@ type ImageData struct {
B64Json string `json:"b64_json"`
RevisedPrompt string `json:"revised_prompt"`
}

func (i *ImageRequest) InputImageSources() ([]types.FileSource, error) {
values := make([]string, 0)
for _, raw := range []json.RawMessage{i.Image, i.Images} {
parsed, err := parseImageSourceValues(raw)
if err != nil {
return nil, err
}
values = append(values, parsed...)
}
return fileSourcesFromImageValues(values), nil
}

func parseImageSourceValues(raw json.RawMessage) ([]string, error) {
if len(raw) == 0 || common.GetJsonType(raw) == "null" {
return nil, nil
}

var values []string
switch common.GetJsonType(raw) {
case "string":
var value string
if err := common.Unmarshal(raw, &value); err != nil {
return nil, err
}
values = append(values, value)
case "array":
var items []json.RawMessage
if err := common.Unmarshal(raw, &items); err != nil {
return nil, err
}
for _, item := range items {
itemValues, err := parseImageSourceValues(item)
if err != nil {
return nil, err
}
values = append(values, itemValues...)
}
case "object":
value, err := parseImageSourceObject(raw)
if err != nil {
return nil, err
}
if value != "" {
values = append(values, value)
}
default:
return nil, fmt.Errorf("image input must be a string, object, or array")
}

return values, nil
}

func fileSourcesFromImageValues(values []string) []types.FileSource {
sources := make([]types.FileSource, 0, len(values))
for _, value := range values {
value = strings.TrimSpace(value)
if value == "" {
continue
}
sources = append(sources, types.NewFileSourceFromData(value, ""))
}
return sources
}

func parseImageSourceObject(raw json.RawMessage) (string, error) {
var item map[string]json.RawMessage
if err := common.Unmarshal(raw, &item); err != nil {
return "", err
}

for _, key := range []string{"url", "image_url", "b64_json", "base64", "data"} {
rawValue, ok := item[key]
if !ok || common.GetJsonType(rawValue) == "null" {
continue
}
if key == "image_url" && common.GetJsonType(rawValue) == "object" {
if value, err := parseImageSourceObject(rawValue); err != nil || value != "" {
return value, err
}
continue
}
if common.GetJsonType(rawValue) != "string" {
continue
}
var value string
if err := common.Unmarshal(rawValue, &value); err != nil {
return "", err
}
if strings.TrimSpace(value) != "" {
return value, nil
}
}

Comment on lines +251 to +279

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🎯 Functional Correctness | 🟠 Major | ⚡ Quick win

Reject malformed recognized image-object fields instead of silently dropping them.

parseImageSourceObject() currently skips recognized keys when their value type is wrong, so inputs like {"image_url":{"url":123}} or {"b64_json":123} can be treated as “no image provided”. With a non-empty prompt, that turns an image-edit request into prompt-only generation instead of returning a 4xx validation error.

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@dto/openai_image.go` around lines 251 - 279, parseImageSourceObject currently
skips recognized image fields when they have the wrong JSON type, which can
silently downgrade malformed image-edit payloads into prompt-only requests.
Update parseImageSourceObject to validate each recognized key such as image_url,
url, b64_json, base64, and data and return an error when a present field is not
the expected string/object shape instead of continuing; keep the recursive
handling for nested image_url objects, but make malformed recognized fields fail
fast with a 4xx-style validation error.

return "", nil
}
51 changes: 51 additions & 0 deletions dto/openai_image_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
package dto

import (
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/types"
"github.com/stretchr/testify/require"
)

func TestImageRequestInputImageSources(t *testing.T) {
raw := []byte(`{
"model":"gemini-2.5-flash-image",
"prompt":"edit",
"image":"https://example.com/input.png",
"images":[
"data:image/png;base64,aW1hZ2U=",
{"url":"https://example.com/second.webp"},
{"image_url":{"url":"https://example.com/openai-style.jpg"}},
{"b64_json":"aW1hZ2Uy"}
]
}`)

var req ImageRequest
require.NoError(t, common.Unmarshal(raw, &req))

sources, err := req.InputImageSources()
require.NoError(t, err)
require.Len(t, sources, 5)

_, ok := sources[0].(*types.URLSource)
require.True(t, ok)
require.Equal(t, "https://example.com/input.png", sources[0].GetRawData())

_, ok = sources[1].(*types.Base64Source)
require.True(t, ok)
require.Equal(t, "data:image/png;base64,aW1hZ2U=", sources[1].GetRawData())

require.Equal(t, "https://example.com/second.webp", sources[2].GetRawData())
require.Equal(t, "https://example.com/openai-style.jpg", sources[3].GetRawData())
require.Equal(t, "aW1hZ2Uy", sources[4].GetRawData())
}

func TestImageRequestInputImageSourcesRejectsScalarJSON(t *testing.T) {
var req ImageRequest
require.NoError(t, common.Unmarshal([]byte(`{"model":"m","prompt":"p","image":123}`), &req))

_, err := req.InputImageSources()
require.Error(t, err)
require.Contains(t, err.Error(), "image input must be")
}
138 changes: 136 additions & 2 deletions relay/channel/gemini/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,11 +7,13 @@ import (
"net/http"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/relay/channel"
"github.com/QuantumNous/new-api/relay/channel/openai"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/constant"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/model_setting"
"github.com/QuantumNous/new-api/setting/reasoning"
"github.com/QuantumNous/new-api/types"
Expand Down Expand Up @@ -58,10 +60,26 @@ func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInf
}

func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest) (any, error) {
sources, err := request.InputImageSources()
if err != nil {
return nil, fmt.Errorf("invalid image input: %w", err)
}

if model_setting.IsGeminiModelSupportImagine(info.UpstreamModelName) {
return convertOpenAIImageRequestToGeminiGenerateContent(c, info, request, sources)
}

if len(sources) > 0 {
return nil, errors.New("input images are supported only by Gemini image generation models")
}

if !strings.HasPrefix(info.UpstreamModelName, "imagen") {
return nil, errors.New("not supported model for image generation, only imagen models are supported")
return nil, errors.New("not supported model for image generation, only imagen and Gemini image models are supported")
}
return convertOpenAIImageRequestToImagenPredict(request), nil
}

func convertOpenAIImageRequestToImagenPredict(request dto.ImageRequest) dto.GeminiImageRequest {
// convert size to aspect ratio but allow user to specify aspect ratio
aspectRatio := "1:1" // default aspect ratio
size := strings.TrimSpace(request.Size)
Expand Down Expand Up @@ -120,7 +138,119 @@ func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInf
geminiRequest.Parameters.ImageSize = imageSize
}

return geminiRequest, nil
return geminiRequest
}

func convertOpenAIImageRequestToGeminiGenerateContent(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest, sources []types.FileSource) (*dto.GeminiChatRequest, error) {
if request.Stream != nil && *request.Stream {
return nil, errors.New("streaming is not supported for Gemini image generation through images/generations")
}

parts := make([]dto.GeminiPart, 0, len(sources)+1)
if strings.TrimSpace(request.Prompt) != "" {
parts = append(parts, dto.GeminiPart{Text: request.Prompt})
}
for _, source := range sources {
base64Data, mimeType, err := service.GetBase64Data(c, source, "formatting image generation input for Gemini")
if err != nil {
return nil, fmt.Errorf("get image input from '%s' failed: %w", source.GetIdentifier(), err)
}
normalizedMimeType := strings.ToLower(mimeType)
if !strings.HasPrefix(normalizedMimeType, "image/") {
return nil, fmt.Errorf("mime type is not supported for Gemini image generation: '%s', url: '%s'", mimeType, source.GetIdentifier())
}
if _, ok := geminiSupportedMimeTypes[normalizedMimeType]; !ok {
return nil, fmt.Errorf("mime type is not supported by Gemini: '%s', url: '%s', supported types are: %v", mimeType, source.GetIdentifier(), getSupportedMimeTypesList())
}
parts = append(parts, dto.GeminiPart{
InlineData: &dto.GeminiInlineData{
MimeType: mimeType,
Data: base64Data,
},
})
}
if len(parts) == 0 {
return nil, errors.New("prompt or image is required")
}

if request.N != nil && *request.N > 1 {
return nil, errors.New("Gemini image generation supports only n=1")
}

info.RelayMode = constant.RelayModeImagesGenerations
geminiRequest := dto.GeminiChatRequest{
Contents: []dto.GeminiChatContent{
{
Role: "user",
Parts: parts,
},
},
GenerationConfig: dto.GeminiChatGenerationConfig{
ResponseModalities: []string{"TEXT", "IMAGE"},
},
SafetySettings: make([]dto.GeminiChatSafetySettings, 0, len(SafetySettingList)),
}

for _, category := range SafetySettingList {
geminiRequest.SafetySettings = append(geminiRequest.SafetySettings, dto.GeminiChatSafetySettings{
Category: category,
Threshold: model_setting.GetGeminiSafetySetting(category),
})
}

imageConfig := map[string]interface{}{}
if aspectRatio := geminiAspectRatioFromSize(request.Size); aspectRatio != "" {
imageConfig["aspectRatio"] = aspectRatio
}
if imageSize := geminiImageSizeFromQuality(request.Quality); imageSize != "" {
imageConfig["imageSize"] = imageSize
}
if len(imageConfig) > 0 {
imageConfigBytes, err := common.Marshal(imageConfig)
if err != nil {
return nil, fmt.Errorf("failed to marshal image config: %w", err)
}
geminiRequest.GenerationConfig.ImageConfig = imageConfigBytes
}

return &geminiRequest, nil
}

func geminiAspectRatioFromSize(size string) string {
size = strings.TrimSpace(size)
if size == "" {
return ""
}
if strings.Contains(size, ":") {
return size
}
switch size {
case "256x256", "512x512", "1024x1024":
return "1:1"
case "1536x1024":
return "3:2"
case "1024x1536":
return "2:3"
case "1024x1792":
return "9:16"
case "1792x1024":
return "16:9"
default:
return ""
}
}

func geminiImageSizeFromQuality(quality string) string {
switch strings.TrimSpace(quality) {
case "":
return ""
case "hd", "high", "2K":
return "2K"
case "standard", "medium", "low", "auto", "1K":
return "1K"
default:
return "1K"
}
}

func (a *Adaptor) Init(info *relaycommon.RelayInfo) {
Expand Down Expand Up @@ -263,6 +393,10 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom
return GeminiImageHandler(c, info, resp)
}

if info.RelayMode == constant.RelayModeImagesGenerations {
return GeminiGenerateContentImageHandler(c, info, resp)
}

// check if the model is an embedding model
if strings.HasPrefix(info.UpstreamModelName, "text-embedding") ||
strings.HasPrefix(info.UpstreamModelName, "embedding") ||
Expand Down
Loading