98 lines
4.6 KiB
Go
98 lines
4.6 KiB
Go
package provider
|
|
|
|
import (
|
|
"context"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"errors"
|
|
"strings"
|
|
"testing"
|
|
|
|
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
|
)
|
|
|
|
func TestImagesProtocolForcesOneResultAndGoogleAuthentication(t *testing.T) {
|
|
image := base64.StdEncoding.EncodeToString([]byte("png"))
|
|
httpClient := &fakeHTTP{response: HTTPResponse{StatusCode: 200, Body: []byte(`{"data":[{"b64_json":"` + image + `"}]}`)}}
|
|
client, err := NewOpenAI(httpClient, OpenAIConfig{BaseURL: "https://provider.test/v1", AuthType: AuthGoogleAPIKey, APIKey: "test-key"})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
outputs, err := client.Generate(context.Background(), Request{Kind: model.KindImage, APIType: model.APIImages, ModelID: "mock-image", RenderedPrompt: "draw"})
|
|
if err != nil || len(outputs) != 1 || string(outputs[0].Content) != "png" {
|
|
t.Fatalf("Generate() = %#v, %v", outputs, err)
|
|
}
|
|
if httpClient.request.URL != "https://provider.test/v1/images/generations" || httpClient.request.Header.Get("x-goog-api-key") != "test-key" || httpClient.request.Header.Get("Authorization") != "" {
|
|
t.Fatalf("request=%#v", httpClient.request)
|
|
}
|
|
var body map[string]any
|
|
if err := json.Unmarshal(httpClient.request.Body, &body); err != nil || body["n"] != float64(1) || body["response_format"] != "b64_json" {
|
|
t.Fatalf("fixed images request body=%s error=%v", httpClient.request.Body, err)
|
|
}
|
|
}
|
|
|
|
func TestImagesRejectsMultipleResultsAndOversizedBase64(t *testing.T) {
|
|
encoded := base64.StdEncoding.EncodeToString([]byte("image"))
|
|
for _, test := range []struct {
|
|
name string
|
|
body string
|
|
limit int64
|
|
}{
|
|
{"multiple", `{"data":[{"b64_json":"` + encoded + `"},{"b64_json":"` + encoded + `"}]}`, 1024},
|
|
{"too_large", `{"data":[{"b64_json":"` + encoded + `"}]}`, 2},
|
|
} {
|
|
t.Run(test.name, func(t *testing.T) {
|
|
client, err := NewOpenAI(&fakeHTTP{response: HTTPResponse{StatusCode: 200, Body: []byte(test.body)}}, OpenAIConfig{BaseURL: "https://provider.test/v1", MaxResponseBytes: test.limit})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
_, err = client.Generate(context.Background(), Request{Kind: model.KindImage, APIType: model.APIImages, ModelID: "m", RenderedPrompt: "p"})
|
|
var providerErr *Error
|
|
if !errors.As(err, &providerErr) || providerErr.Code != CodeUnknown {
|
|
t.Fatalf("error=%v", err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestGeminiProtocolEscapesModelAndUsesInlineData(t *testing.T) {
|
|
image := []byte("image-content")
|
|
response := `{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"` + base64.StdEncoding.EncodeToString(image) + `"}}]}}]}`
|
|
httpClient := &fakeHTTP{response: HTTPResponse{StatusCode: 200, Body: []byte(response)}}
|
|
client, err := NewOpenAI(httpClient, OpenAIConfig{BaseURL: "https://provider.test/v1beta", AuthType: AuthNone})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
outputs, err := client.Generate(context.Background(), Request{
|
|
Kind: model.KindImage, APIType: model.APIGemini, ModelID: "gemini/image", RenderedPrompt: "edit",
|
|
Inputs: []Input{{Role: model.RolePrimary, Position: 0, MIMEType: "image/png", Content: image}},
|
|
})
|
|
if err != nil || len(outputs) != 1 || string(outputs[0].Content) != string(image) {
|
|
t.Fatalf("Generate() = %#v, %v", outputs, err)
|
|
}
|
|
if httpClient.request.URL != "https://provider.test/v1beta/models/gemini%2Fimage:generateContent" || !strings.Contains(string(httpClient.request.Body), `"inlineData"`) || !strings.Contains(string(httpClient.request.Body), base64.StdEncoding.EncodeToString(image)) {
|
|
t.Fatalf("gemini request=%#v body=%s", httpClient.request, httpClient.request.Body)
|
|
}
|
|
}
|
|
|
|
func TestProtocolCapabilityAndReservedBodyValidation(t *testing.T) {
|
|
client, err := NewOpenAI(&fakeHTTP{}, OpenAIConfig{BaseURL: "https://provider.test/v1"})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, request := range []Request{
|
|
{Kind: model.KindImage, APIType: model.APIChat, ModelID: "m", RenderedPrompt: "p"},
|
|
{Kind: model.KindImage, APIType: model.APIImages, ModelID: "m", RenderedPrompt: "p", Inputs: []Input{{Role: model.RolePrimary, MIMEType: "image/png", Content: []byte("x")}}},
|
|
{Kind: model.KindImage, APIType: model.APIImagesEdits, ModelID: "m", RenderedPrompt: "p", Inputs: []Input{{Role: model.RoleReference, MIMEType: "image/png", Content: []byte("x")}}},
|
|
} {
|
|
if _, err := client.Generate(context.Background(), request); !errors.Is(err, ErrInvalidRequest) {
|
|
t.Fatalf("request=%#v error=%v", request, err)
|
|
}
|
|
}
|
|
for _, value := range []string{`{"n":2}`, `{"response_format":"url"}`, `{"authorization":"override"}`} {
|
|
if _, err := NewOpenAI(&fakeHTTP{}, OpenAIConfig{BaseURL: "https://provider.test/v1", ExtraBody: []byte(value)}); !errors.Is(err, ErrInvalidConfig) {
|
|
t.Fatalf("extra body %s error=%v", value, err)
|
|
}
|
|
}
|
|
}
|