new-api / image_task_compat_test.go
Codex
Add async image task compatibility
7cd5cb8
Raw
History Blame Contribute Delete
4.35 kB
package openai
import (
"bytes"
"context"
"io"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestPollOpenAIImageTaskReturnsFinalResponse(t *testing.T) {
attempts := 0
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
assert.Equal(t, "Bearer test-key", request.Header.Get("Authorization"))
attempts++
writer.Header().Set("Content-Type", "application/json")
if attempts == 1 {
writer.WriteHeader(http.StatusAccepted)
_, _ = writer.Write([]byte(`{"object":"image.task","status":"processing","task_id":"task-1","poll_url":"/v1/images/tasks/task-1","poll_after_ms":1}`))
return
}
_, _ = writer.Write([]byte(`{"created":123,"data":[{"url":"https://example.com/image.png"}],"usage":{"total_tokens":7}}`))
}))
defer server.Close()
request := httptest.NewRequest(http.MethodPost, server.URL+"/v1/images/edits", nil)
request.Header.Set("Authorization", "Bearer test-key")
initial := &http.Response{
StatusCode: http.StatusAccepted,
Header: make(http.Header),
Request: request,
Body: io.NopCloser(bytes.NewBufferString(
`{"object":"image.task","status":"queued","task_id":"task-1","poll_url":"/v1/images/tasks/task-1","poll_after_ms":1}`,
)),
}
resolved, stats, err := pollOpenAIImageTask(context.Background(), initial, server.Client(), time.Second, noImageTaskWait)
require.NoError(t, err)
require.NotNil(t, resolved)
assert.Equal(t, http.StatusOK, resolved.StatusCode)
assert.Equal(t, 2, stats.Attempts)
assert.Equal(t, "task-1", stats.TaskID)
resolvedBody, err := io.ReadAll(resolved.Body)
require.NoError(t, err)
assert.True(t, isFinalOpenAIImageResponse(resolvedBody))
var payload map[string]any
require.NoError(t, common.Unmarshal(resolvedBody, &payload))
assert.Equal(t, float64(123), payload["created"])
}
func TestPollOpenAIImageTaskRejectsCrossOriginPollURL(t *testing.T) {
request := httptest.NewRequest(http.MethodPost, "https://upstream.example/v1/images/edits", nil)
initial := &http.Response{
StatusCode: http.StatusAccepted,
Header: make(http.Header),
Request: request,
Body: io.NopCloser(bytes.NewBufferString(
`{"object":"image.task","status":"queued","task_id":"task-1","poll_url":"https://other.example/tasks/task-1"}`,
)),
}
_, _, err := pollOpenAIImageTask(context.Background(), initial, http.DefaultClient, time.Second, noImageTaskWait)
require.Error(t, err)
assert.Contains(t, err.Error(), "original upstream origin")
}
func TestPollOpenAIImageTaskReturnsFailure(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
writer.Header().Set("Content-Type", "application/json")
_, _ = writer.Write([]byte(`{"object":"image.task","status":"failed","task_id":"task-1","poll_url":"/v1/images/tasks/task-1","error":"generation failed"}`))
}))
defer server.Close()
request := httptest.NewRequest(http.MethodPost, server.URL+"/v1/images/edits", nil)
initial := &http.Response{
StatusCode: http.StatusAccepted,
Header: make(http.Header),
Request: request,
Body: io.NopCloser(bytes.NewBufferString(
`{"object":"image.task","status":"queued","task_id":"task-1","poll_url":"/v1/images/tasks/task-1"}`,
)),
}
_, _, err := pollOpenAIImageTask(context.Background(), initial, server.Client(), time.Second, noImageTaskWait)
require.Error(t, err)
assert.Contains(t, err.Error(), "generation failed")
}
func TestPollOpenAIImageTaskStopsWhenContextIsCancelled(t *testing.T) {
request := httptest.NewRequest(http.MethodPost, "https://upstream.example/v1/images/edits", nil)
initial := &http.Response{
StatusCode: http.StatusAccepted,
Header: make(http.Header),
Request: request,
Body: io.NopCloser(bytes.NewBufferString(
`{"object":"image.task","status":"queued","task_id":"task-1","poll_url":"/v1/images/tasks/task-1"}`,
)),
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
_, _, err := pollOpenAIImageTask(ctx, initial, http.DefaultClient, time.Second, waitForOpenAIImageTask)
require.Error(t, err)
assert.ErrorIs(t, err, context.Canceled)
}
func noImageTaskWait(context.Context, time.Duration) error {
return nil
}