Files
chorus/internal/core/provider/openai_test.go
T

110 lines
4.9 KiB
Go

package provider
import (
"context"
"encoding/base64"
"errors"
"net/http"
"strings"
"testing"
"git.ilapage.cn/OPC/chorus/internal/core/model"
)
type fakeHTTP struct {
response HTTPResponse
err error
fetched HTTPResponse
request HTTPRequest
fetchURL string
}
func (f *fakeHTTP) Do(_ context.Context, request HTTPRequest) (HTTPResponse, error) {
f.request = request
return f.response, f.err
}
func (f *fakeHTTP) Fetch(_ context.Context, rawURL string, _ int64) (HTTPResponse, error) {
f.fetchURL = rawURL
return f.fetched, f.err
}
func TestChatProtocol(t *testing.T) {
httpClient := &fakeHTTP{response: HTTPResponse{StatusCode: 200, ContentType: "application/json", Body: []byte(`{"choices":[{"message":{"content":"done"}}]}`)}}
client, err := NewOpenAI(httpClient, OpenAIConfig{BaseURL: "https://provider.test/v1", APIKey: "secret", ExtraBody: jsonBytes(`{"temperature":0.2}`)})
if err != nil {
t.Fatal(err)
}
outputs, err := client.Generate(context.Background(), Request{Kind: model.KindText, APIType: model.APIChat, ModelID: "mock-chat", RenderedPrompt: "write"})
if err != nil || len(outputs) != 1 || outputs[0].Text != "done" {
t.Fatalf("Generate() = %#v, %v", outputs, err)
}
if httpClient.request.URL != "https://provider.test/v1/chat/completions" || httpClient.request.Header.Get("Authorization") != "Bearer secret" || !strings.Contains(string(httpClient.request.Body), `"temperature":0.2`) {
t.Fatalf("request = %#v body=%s", httpClient.request, string(httpClient.request.Body))
}
}
func TestImagesEditsBase64AndURL(t *testing.T) {
pngData := []byte("image-bytes")
responseBody := `{"data":[{"b64_json":"` + base64.StdEncoding.EncodeToString(pngData) + `"},{"url":"https://result.test/output"}]}`
httpClient := &fakeHTTP{response: HTTPResponse{StatusCode: 200, Body: []byte(responseBody)}, fetched: HTTPResponse{StatusCode: 200, ContentType: "image/png", Body: pngData}}
client, _ := NewOpenAI(httpClient, OpenAIConfig{BaseURL: "https://provider.test/v1", ExtraBody: jsonBytes(`{"size":"1024x1024"}`)})
outputs, err := client.Generate(context.Background(), Request{Kind: model.KindImage, APIType: model.APIImagesEdits, ModelID: "mock-image", RenderedPrompt: "edit", Inputs: []Input{{Role: model.RolePrimary, MIMEType: "image/png", Content: pngData}}})
if err != nil || len(outputs) != 2 || string(outputs[0].Content) != "image-bytes" || httpClient.fetchURL != "https://result.test/output" {
t.Fatalf("Generate() = %#v, %v, fetch=%q", outputs, err, httpClient.fetchURL)
}
if !strings.Contains(httpClient.request.Header.Get("Content-Type"), "multipart/form-data") || !strings.Contains(string(httpClient.request.Body), "mock-image") {
t.Fatalf("multipart request missing fields")
}
}
func TestFailureClassificationMatrix(t *testing.T) {
tests := []struct {
name string
status int
body string
class FailureClass
code ErrorCode
}{
{"429", 429, "", FailureRateLimited, CodeRateLimited}, {"500", 500, "", FailureServer, CodeServerError},
{"400", 400, "", FailureBadRequest, CodeBadRequest}, {"401", 401, "", FailureUnauthorized, CodeUnauthorized},
{"policy", 400, `{"error":{"code":"content_policy_violation"}}`, FailurePolicyRejected, CodePolicyRejected},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
client, _ := NewOpenAI(&fakeHTTP{response: HTTPResponse{StatusCode: test.status, Body: []byte(test.body)}}, OpenAIConfig{BaseURL: "https://provider.test/v1"})
_, err := client.Generate(context.Background(), Request{Kind: model.KindText, APIType: model.APIChat, ModelID: "m", RenderedPrompt: "p"})
var providerErr *Error
if !errors.As(err, &providerErr) || providerErr.Class != test.class || providerErr.Code != test.code || (test.body != "" && strings.Contains(providerErr.Error(), test.body)) {
t.Fatalf("error = %#v", err)
}
})
}
for _, test := range []struct {
name string
err error
class FailureClass
}{{"timeout", context.DeadlineExceeded, FailureTimeout}, {"connection", errors.New("dial failed"), FailureConnection}} {
t.Run(test.name, func(t *testing.T) {
client, _ := NewOpenAI(&fakeHTTP{err: test.err}, OpenAIConfig{BaseURL: "https://provider.test/v1"})
_, err := client.Generate(context.Background(), Request{Kind: model.KindText, APIType: model.APIChat, ModelID: "m", RenderedPrompt: "p"})
var providerErr *Error
if !errors.As(err, &providerErr) || providerErr.Class != test.class {
t.Fatalf("error=%v", err)
}
})
}
}
func TestExtraBodyRejectsProtocolOverrides(t *testing.T) {
for _, value := range []string{`{"model":"override"}`, `{"messages":[]}`, `{"unknown":true}`} {
if _, err := NewOpenAI(&fakeHTTP{}, OpenAIConfig{BaseURL: "https://provider.test/v1", ExtraBody: jsonBytes(value)}); !errors.Is(err, ErrInvalidConfig) {
t.Fatalf("extra_body %s error=%v", value, err)
}
}
}
func jsonBytes(value string) []byte { return []byte(value) }
var _ HTTPClient = (*fakeHTTP)(nil)
var _ = http.MethodPost