agagsder commited on
Commit
dcc3f07
·
verified ·
1 Parent(s): ab16a12

Normalize automatic GPT image size

Browse files
Files changed (2) hide show
  1. adapter/main.go +83 -2
  2. adapter/main_test.go +67 -4
adapter/main.go CHANGED
@@ -25,6 +25,7 @@ import (
25
  const (
26
  defaultExternalPort = "7860"
27
  defaultInternalPort = "3000"
 
28
  maxImageResponseBytes = 96 << 20
29
  maxDownloadedImageSize = 64 << 20
30
  imageDownloadTimeout = 3 * time.Minute
@@ -64,6 +65,10 @@ func run() error {
64
  if err != nil {
65
  return err
66
  }
 
 
 
 
67
 
68
  child := exec.Command("/new-api")
69
  child.Dir = "/data"
@@ -83,7 +88,7 @@ func run() error {
83
 
84
  server := &http.Server{
85
  Addr: ":" + externalPort,
86
- Handler: newProxy(backend, newPublicImageClient(), defaultMode, maxConcurrency),
87
  ReadHeaderTimeout: 30 * time.Second,
88
  IdleTimeout: 2 * time.Minute,
89
  }
@@ -98,11 +103,12 @@ func run() error {
98
  defer signal.Stop(signalCh)
99
 
100
  log.Printf(
101
- "image response adapter listening on :%s; New API on 127.0.0.1:%s; default=%s; max_concurrency=%d",
102
  externalPort,
103
  internalPort,
104
  defaultMode,
105
  maxConcurrency,
 
106
  )
107
 
108
  select {
@@ -132,6 +138,7 @@ func newProxy(
132
  downloadClient *http.Client,
133
  defaultMode imageResponseMode,
134
  maxConcurrency int,
 
135
  ) *httputil.ReverseProxy {
136
  proxy := httputil.NewSingleHostReverseProxy(backend)
137
  originalDirector := proxy.Director
@@ -139,6 +146,9 @@ func newProxy(
139
  if isImageEndpoint(request.URL.Path) {
140
  mode := imageResponseModeForRequest(request, defaultMode)
141
  request.Header.Del(responseFormatHeader)
 
 
 
142
  requestWithMode := request.WithContext(context.WithValue(request.Context(), responseModeContextKey{}, mode))
143
  *request = *requestWithMode
144
  }
@@ -374,6 +384,52 @@ func imageResponseModeForRequest(request *http.Request, fallback imageResponseMo
374
  return fallback
375
  }
376
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
377
  func imageResponseModeFromContext(ctx context.Context) imageResponseMode {
378
  if mode, ok := ctx.Value(responseModeContextKey{}).(imageResponseMode); ok {
379
  return mode
@@ -394,6 +450,31 @@ func parseImageResponseMode(raw string) (imageResponseMode, error) {
394
  }
395
  }
396
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
397
  func decodeNonEmptyString(raw json.RawMessage) (string, bool) {
398
  if len(raw) == 0 || string(raw) == "null" {
399
  return "", false
 
25
  const (
26
  defaultExternalPort = "7860"
27
  defaultInternalPort = "3000"
28
+ defaultAutoImageSize = "1024x1024"
29
  maxImageResponseBytes = 96 << 20
30
  maxDownloadedImageSize = 64 << 20
31
  imageDownloadTimeout = 3 * time.Minute
 
65
  if err != nil {
66
  return err
67
  }
68
+ autoImageSize, err := parseAutoImageSize(envOrDefault("IMAGE_AUTO_SIZE", defaultAutoImageSize))
69
+ if err != nil {
70
+ return err
71
+ }
72
 
73
  child := exec.Command("/new-api")
74
  child.Dir = "/data"
 
88
 
89
  server := &http.Server{
90
  Addr: ":" + externalPort,
91
+ Handler: newProxy(backend, newPublicImageClient(), defaultMode, maxConcurrency, autoImageSize),
92
  ReadHeaderTimeout: 30 * time.Second,
93
  IdleTimeout: 2 * time.Minute,
94
  }
 
103
  defer signal.Stop(signalCh)
104
 
105
  log.Printf(
106
+ "image response adapter listening on :%s; New API on 127.0.0.1:%s; default=%s; max_concurrency=%d; auto_size=%s",
107
  externalPort,
108
  internalPort,
109
  defaultMode,
110
  maxConcurrency,
111
+ displayAutoImageSize(autoImageSize),
112
  )
113
 
114
  select {
 
138
  downloadClient *http.Client,
139
  defaultMode imageResponseMode,
140
  maxConcurrency int,
141
+ autoImageSize string,
142
  ) *httputil.ReverseProxy {
143
  proxy := httputil.NewSingleHostReverseProxy(backend)
144
  originalDirector := proxy.Director
 
146
  if isImageEndpoint(request.URL.Path) {
147
  mode := imageResponseModeForRequest(request, defaultMode)
148
  request.Header.Del(responseFormatHeader)
149
+ if normalizeAutoImageSize(request, autoImageSize) {
150
+ log.Printf("normalized auto image size to %s", autoImageSize)
151
+ }
152
  requestWithMode := request.WithContext(context.WithValue(request.Context(), responseModeContextKey{}, mode))
153
  *request = *requestWithMode
154
  }
 
384
  return fallback
385
  }
386
 
387
+ func normalizeAutoImageSize(request *http.Request, fallbackSize string) bool {
388
+ if request == nil || request.URL.Path != "/v1/images/generations" || request.Body == nil || fallbackSize == "" {
389
+ return false
390
+ }
391
+
392
+ body, complete := readAndRestoreRequestBody(request, maxModeRequestBody)
393
+ if !complete {
394
+ return false
395
+ }
396
+ var payload map[string]json.RawMessage
397
+ if err := json.Unmarshal(body, &payload); err != nil {
398
+ return false
399
+ }
400
+ model, ok := decodeNonEmptyString(payload["model"])
401
+ if !ok || !strings.HasPrefix(strings.ToLower(model), "gpt-image-") {
402
+ return false
403
+ }
404
+ size, hasSize := decodeNonEmptyString(payload["size"])
405
+ if hasSize && !strings.EqualFold(strings.TrimSpace(size), "auto") {
406
+ return false
407
+ }
408
+ encodedSize, err := json.Marshal(fallbackSize)
409
+ if err != nil {
410
+ return false
411
+ }
412
+ payload["size"] = encodedSize
413
+ updatedBody, err := json.Marshal(payload)
414
+ if err != nil {
415
+ return false
416
+ }
417
+ request.Body = io.NopCloser(bytes.NewReader(updatedBody))
418
+ request.ContentLength = int64(len(updatedBody))
419
+ request.Header.Set("Content-Length", strconv.Itoa(len(updatedBody)))
420
+ return true
421
+ }
422
+
423
+ func readAndRestoreRequestBody(request *http.Request, limit int64) ([]byte, bool) {
424
+ prefix, err := io.ReadAll(io.LimitReader(request.Body, limit+1))
425
+ request.Body = io.NopCloser(io.MultiReader(bytes.NewReader(prefix), request.Body))
426
+ if err != nil || int64(len(prefix)) > limit {
427
+ return nil, false
428
+ }
429
+ request.Body = io.NopCloser(bytes.NewReader(prefix))
430
+ return prefix, true
431
+ }
432
+
433
  func imageResponseModeFromContext(ctx context.Context) imageResponseMode {
434
  if mode, ok := ctx.Value(responseModeContextKey{}).(imageResponseMode); ok {
435
  return mode
 
450
  }
451
  }
452
 
453
+ func parseAutoImageSize(raw string) (string, error) {
454
+ normalized := strings.ToLower(strings.TrimSpace(raw))
455
+ switch normalized {
456
+ case "off", "disabled", "passthrough":
457
+ return "", nil
458
+ }
459
+ parts := strings.Split(normalized, "x")
460
+ if len(parts) != 2 {
461
+ return "", fmt.Errorf("IMAGE_AUTO_SIZE must be WIDTHxHEIGHT or off")
462
+ }
463
+ width, widthErr := strconv.Atoi(parts[0])
464
+ height, heightErr := strconv.Atoi(parts[1])
465
+ if widthErr != nil || heightErr != nil || width < 1 || height < 1 {
466
+ return "", fmt.Errorf("IMAGE_AUTO_SIZE must be WIDTHxHEIGHT or off")
467
+ }
468
+ return strconv.Itoa(width) + "x" + strconv.Itoa(height), nil
469
+ }
470
+
471
+ func displayAutoImageSize(size string) string {
472
+ if size == "" {
473
+ return "off"
474
+ }
475
+ return size
476
+ }
477
+
478
  func decodeNonEmptyString(raw json.RawMessage) (string, bool) {
479
  if len(raw) == 0 || string(raw) == "null" {
480
  return "", false
adapter/main_test.go CHANGED
@@ -130,6 +130,58 @@ func TestImageResponseModeForRequest(t *testing.T) {
130
  }
131
  }
132
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
133
  func TestAdaptImageResponseHonorsURLMode(t *testing.T) {
134
  body := `{"data":[{"url":"https://cdn.example/image.png"}]}`
135
  request := httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil)
@@ -162,14 +214,22 @@ func TestProxyConvertsDefaultImageResponse(t *testing.T) {
162
  defer imageServer.Close()
163
 
164
  var forwardedOverride string
 
165
  backend := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
166
  forwardedOverride = request.Header.Get(responseFormatHeader)
 
 
 
 
 
 
 
167
  writer.Header().Set("Content-Type", "application/json")
168
  _, _ = fmt.Fprintf(writer, `{"created":123,"data":[{"url":%q}]}`, imageServer.URL+"/image.png")
169
  }))
170
  defer backend.Close()
171
 
172
- proxy := httptest.NewServer(newProxy(mustParseURL(t, backend.URL), imageServer.Client(), imageResponseModeBase64, 2))
173
  defer proxy.Close()
174
 
175
  request, err := http.NewRequest(http.MethodPost, proxy.URL+"/v1/images/generations", strings.NewReader(`{"model":"gpt-image-2"}`))
@@ -199,6 +259,9 @@ func TestProxyConvertsDefaultImageResponse(t *testing.T) {
199
  if forwardedOverride != "" {
200
  t.Fatalf("override header leaked upstream: %q", forwardedOverride)
201
  }
 
 
 
202
  if len(payload.Data) != 1 || payload.Data[0].URL != "" || payload.Data[0].B64JSON != base64.StdEncoding.EncodeToString(imageBytes) {
203
  t.Fatalf("unexpected converted response: %+v", payload)
204
  }
@@ -219,7 +282,7 @@ func TestProxyHonorsURLResponseFormat(t *testing.T) {
219
  }))
220
  defer backend.Close()
221
 
222
- proxy := httptest.NewServer(newProxy(mustParseURL(t, backend.URL), imageServer.Client(), imageResponseModeBase64, 2))
223
  defer proxy.Close()
224
 
225
  response, err := proxy.Client().Post(
@@ -255,7 +318,7 @@ func TestProxyLeavesNonImageResponseUnchanged(t *testing.T) {
255
  }))
256
  defer backend.Close()
257
 
258
- proxy := httptest.NewServer(newProxy(mustParseURL(t, backend.URL), http.DefaultClient, imageResponseModeBase64, 2))
259
  defer proxy.Close()
260
  response, err := proxy.Client().Post(proxy.URL+"/v1/responses", "application/json", strings.NewReader(`{"model":"gpt-5"}`))
261
  if err != nil {
@@ -301,7 +364,7 @@ func TestProxyLimitsConcurrentImageDownloads(t *testing.T) {
301
  _, _ = fmt.Fprintf(writer, `{"data":[{"url":%q}]}`, imageServer.URL+"/image.png")
302
  }))
303
  defer backend.Close()
304
- proxy := httptest.NewServer(newProxy(mustParseURL(t, backend.URL), imageServer.Client(), imageResponseModeBase64, concurrency))
305
  defer proxy.Close()
306
 
307
  errors := make(chan error, requestCount)
 
130
  }
131
  }
132
 
133
+ func TestNormalizeAutoImageSize(t *testing.T) {
134
+ tests := []struct {
135
+ name string
136
+ body string
137
+ fallback string
138
+ wantChanged bool
139
+ wantSize string
140
+ }{
141
+ {name: "auto", body: `{"model":"gpt-image-2","size":"auto","prompt":"test"}`, fallback: defaultAutoImageSize, wantChanged: true, wantSize: defaultAutoImageSize},
142
+ {name: "missing", body: `{"model":"gpt-image-2","prompt":"test"}`, fallback: defaultAutoImageSize, wantChanged: true, wantSize: defaultAutoImageSize},
143
+ {name: "explicit", body: `{"model":"gpt-image-2","size":"1536x1024"}`, fallback: defaultAutoImageSize, wantSize: "1536x1024"},
144
+ {name: "other model", body: `{"model":"dall-e-3","size":"auto"}`, fallback: defaultAutoImageSize, wantSize: "auto"},
145
+ {name: "disabled", body: `{"model":"gpt-image-2","size":"auto"}`, wantSize: "auto"},
146
+ }
147
+ for _, test := range tests {
148
+ t.Run(test.name, func(t *testing.T) {
149
+ request := httptest.NewRequest(http.MethodPost, "/v1/images/generations", strings.NewReader(test.body))
150
+ changed := normalizeAutoImageSize(request, test.fallback)
151
+ if changed != test.wantChanged {
152
+ t.Fatalf("changed=%v want=%v", changed, test.wantChanged)
153
+ }
154
+ var payload struct {
155
+ Size string `json:"size"`
156
+ }
157
+ if err := json.NewDecoder(request.Body).Decode(&payload); err != nil {
158
+ t.Fatal(err)
159
+ }
160
+ if payload.Size != test.wantSize {
161
+ t.Fatalf("size=%q want=%q", payload.Size, test.wantSize)
162
+ }
163
+ })
164
+ }
165
+ }
166
+
167
+ func TestParseAutoImageSize(t *testing.T) {
168
+ for raw, want := range map[string]string{
169
+ "1024x1024": "1024x1024",
170
+ " 1536X1024 ": "1536x1024",
171
+ "off": "",
172
+ } {
173
+ got, err := parseAutoImageSize(raw)
174
+ if err != nil || got != want {
175
+ t.Fatalf("parseAutoImageSize(%q)=%q, %v; want %q", raw, got, err, want)
176
+ }
177
+ }
178
+ for _, raw := range []string{"auto", "1024", "0x1024", "wide x tall"} {
179
+ if _, err := parseAutoImageSize(raw); err == nil {
180
+ t.Fatalf("parseAutoImageSize(%q) should fail", raw)
181
+ }
182
+ }
183
+ }
184
+
185
  func TestAdaptImageResponseHonorsURLMode(t *testing.T) {
186
  body := `{"data":[{"url":"https://cdn.example/image.png"}]}`
187
  request := httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil)
 
214
  defer imageServer.Close()
215
 
216
  var forwardedOverride string
217
+ var forwardedSize string
218
  backend := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
219
  forwardedOverride = request.Header.Get(responseFormatHeader)
220
+ var payload struct {
221
+ Size string `json:"size"`
222
+ }
223
+ if err := json.NewDecoder(request.Body).Decode(&payload); err != nil {
224
+ t.Error(err)
225
+ }
226
+ forwardedSize = payload.Size
227
  writer.Header().Set("Content-Type", "application/json")
228
  _, _ = fmt.Fprintf(writer, `{"created":123,"data":[{"url":%q}]}`, imageServer.URL+"/image.png")
229
  }))
230
  defer backend.Close()
231
 
232
+ proxy := httptest.NewServer(newProxy(mustParseURL(t, backend.URL), imageServer.Client(), imageResponseModeBase64, 2, defaultAutoImageSize))
233
  defer proxy.Close()
234
 
235
  request, err := http.NewRequest(http.MethodPost, proxy.URL+"/v1/images/generations", strings.NewReader(`{"model":"gpt-image-2"}`))
 
259
  if forwardedOverride != "" {
260
  t.Fatalf("override header leaked upstream: %q", forwardedOverride)
261
  }
262
+ if forwardedSize != defaultAutoImageSize {
263
+ t.Fatalf("forwarded size=%q want=%q", forwardedSize, defaultAutoImageSize)
264
+ }
265
  if len(payload.Data) != 1 || payload.Data[0].URL != "" || payload.Data[0].B64JSON != base64.StdEncoding.EncodeToString(imageBytes) {
266
  t.Fatalf("unexpected converted response: %+v", payload)
267
  }
 
282
  }))
283
  defer backend.Close()
284
 
285
+ proxy := httptest.NewServer(newProxy(mustParseURL(t, backend.URL), imageServer.Client(), imageResponseModeBase64, 2, defaultAutoImageSize))
286
  defer proxy.Close()
287
 
288
  response, err := proxy.Client().Post(
 
318
  }))
319
  defer backend.Close()
320
 
321
+ proxy := httptest.NewServer(newProxy(mustParseURL(t, backend.URL), http.DefaultClient, imageResponseModeBase64, 2, defaultAutoImageSize))
322
  defer proxy.Close()
323
  response, err := proxy.Client().Post(proxy.URL+"/v1/responses", "application/json", strings.NewReader(`{"model":"gpt-5"}`))
324
  if err != nil {
 
364
  _, _ = fmt.Fprintf(writer, `{"data":[{"url":%q}]}`, imageServer.URL+"/image.png")
365
  }))
366
  defer backend.Close()
367
+ proxy := httptest.NewServer(newProxy(mustParseURL(t, backend.URL), imageServer.Client(), imageResponseModeBase64, concurrency, defaultAutoImageSize))
368
  defer proxy.Close()
369
 
370
  errors := make(chan error, requestCount)