110 lines
4.9 KiB
Go
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
|