Files

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