File size: 4,346 Bytes
7cd5cb8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
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
}