| 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 |
| } |
|
|