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 }