Bound user-supplied count/duration parameters at request validation, route ratio multipliers through guarded setters, and use saturating int conversions in all quota math paths.
150 lines
4.8 KiB
Go
150 lines
4.8 KiB
Go
package helper
|
|
|
|
import (
|
|
"bytes"
|
|
"fmt"
|
|
"io"
|
|
"mime/multipart"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"testing"
|
|
|
|
"github.com/QuantumNous/new-api/common"
|
|
"github.com/QuantumNous/new-api/dto"
|
|
relayconstant "github.com/QuantumNous/new-api/relay/constant"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/stretchr/testify/require"
|
|
)
|
|
|
|
// TestGetAndValidOpenAIImageRequestMultipartStream verifies multipart image
|
|
// edit parsing: the stream field is parsed and validated, and the request body
|
|
// stays replayable for the upstream request.
|
|
func TestGetAndValidOpenAIImageRequestMultipartStream(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
newContext := func(t *testing.T, streamValue string, withImage bool) (*gin.Context, string) {
|
|
var body bytes.Buffer
|
|
writer := multipart.NewWriter(&body)
|
|
require.NoError(t, writer.WriteField("model", "gpt-image-1"))
|
|
require.NoError(t, writer.WriteField("prompt", "edit this image"))
|
|
require.NoError(t, writer.WriteField("stream", streamValue))
|
|
if withImage {
|
|
part, err := writer.CreateFormFile("image", "input.png")
|
|
require.NoError(t, err)
|
|
_, err = part.Write([]byte("fake image"))
|
|
require.NoError(t, err)
|
|
}
|
|
require.NoError(t, writer.Close())
|
|
originalBody := body.String()
|
|
|
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/edits", &body)
|
|
c.Request.Header.Set("Content-Type", writer.FormDataContentType())
|
|
return c, originalBody
|
|
}
|
|
|
|
t.Run("valid stream value keeps body replayable", func(t *testing.T) {
|
|
c, originalBody := newContext(t, "true", true)
|
|
|
|
req, err := GetAndValidOpenAIImageRequest(c, relayconstant.RelayModeImagesEdits)
|
|
require.NoError(t, err)
|
|
require.NotNil(t, req.Stream)
|
|
require.True(t, *req.Stream)
|
|
require.True(t, req.IsStream(c))
|
|
|
|
bodyAfterValidation, err := io.ReadAll(c.Request.Body)
|
|
require.NoError(t, err)
|
|
require.Equal(t, originalBody, string(bodyAfterValidation))
|
|
|
|
form, err := common.ParseMultipartFormReusable(c)
|
|
require.NoError(t, err)
|
|
require.Equal(t, "true", url.Values(form.Value).Get("stream"))
|
|
require.Len(t, form.File["image"], 1)
|
|
})
|
|
|
|
t.Run("invalid stream value is rejected", func(t *testing.T) {
|
|
c, _ := newContext(t, "notabool", false)
|
|
|
|
_, err := GetAndValidOpenAIImageRequest(c, relayconstant.RelayModeImagesEdits)
|
|
require.Error(t, err)
|
|
require.Contains(t, err.Error(), "invalid stream value")
|
|
})
|
|
}
|
|
|
|
// TestGetAndValidOpenAIImageRequestNBounds guards the billing invariant that
|
|
// the image generation count can never reach quota calculation with a value
|
|
// large enough to overflow int64 into a negative charge.
|
|
func TestGetAndValidOpenAIImageRequestNBounds(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
newJSONContext := func(t *testing.T, body string) *gin.Context {
|
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewBufferString(body))
|
|
c.Request.Header.Set("Content-Type", "application/json")
|
|
return c
|
|
}
|
|
|
|
boundErr := fmt.Sprintf("n must be an integer between 1 and %d", dto.MaxImageN)
|
|
|
|
tests := []struct {
|
|
name string
|
|
body string
|
|
wantErr string
|
|
wantN uint
|
|
}{
|
|
{
|
|
name: "overflowed uint64 n is rejected",
|
|
body: `{"model":"gpt-image-1","prompt":"a cat","n":18446744073686646784}`,
|
|
wantErr: boundErr,
|
|
},
|
|
{
|
|
name: "n above max is rejected",
|
|
body: fmt.Sprintf(`{"model":"gpt-image-1","prompt":"a cat","n":%d}`, dto.MaxImageN+1),
|
|
wantErr: boundErr,
|
|
},
|
|
{
|
|
name: "n at max is accepted",
|
|
body: fmt.Sprintf(`{"model":"gpt-image-1","prompt":"a cat","n":%d}`, dto.MaxImageN),
|
|
wantN: dto.MaxImageN,
|
|
},
|
|
{
|
|
name: "absent n defaults to 1",
|
|
body: `{"model":"gpt-image-1","prompt":"a cat"}`,
|
|
wantN: 1,
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
c := newJSONContext(t, tt.body)
|
|
req, err := GetAndValidOpenAIImageRequest(c, relayconstant.RelayModeImagesGenerations)
|
|
if tt.wantErr != "" {
|
|
require.Error(t, err)
|
|
require.Contains(t, err.Error(), tt.wantErr)
|
|
return
|
|
}
|
|
require.NoError(t, err)
|
|
require.NotNil(t, req.N)
|
|
require.Equal(t, tt.wantN, *req.N)
|
|
})
|
|
}
|
|
|
|
t.Run("negative multipart n is rejected", func(t *testing.T) {
|
|
var body bytes.Buffer
|
|
writer := multipart.NewWriter(&body)
|
|
require.NoError(t, writer.WriteField("model", "gpt-image-1"))
|
|
require.NoError(t, writer.WriteField("prompt", "edit this image"))
|
|
require.NoError(t, writer.WriteField("n", "-22904832"))
|
|
require.NoError(t, writer.Close())
|
|
|
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/edits", &body)
|
|
c.Request.Header.Set("Content-Type", writer.FormDataContentType())
|
|
|
|
_, err := GetAndValidOpenAIImageRequest(c, relayconstant.RelayModeImagesEdits)
|
|
require.Error(t, err)
|
|
require.Contains(t, err.Error(), boundErr)
|
|
})
|
|
}
|