307 lines
14 KiB
Go
307 lines
14 KiB
Go
package providers
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"io"
|
|
"net/http"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
type roundTripFunc func(*http.Request) (*http.Response, error)
|
|
|
|
func (f roundTripFunc) Do(r *http.Request) (*http.Response, error) { return f(r) }
|
|
|
|
func TestHTTPAdaptersMapRequestsAndResponses(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
adapter Adapter
|
|
wantSubmitPath, wantQueryPath, response string
|
|
}{
|
|
{"evolink", NewEvoLink(Config{BaseURL: "https://e.test", APIKey: "secret", Model: "gpt-image-2"}, nil), "/v1/images/generations", "/v1/tasks/task-1", `{"id":"task-1","status":"completed","results":[{"url":"https://cdn.test/a.png"}]}`},
|
|
{"bailian", NewBailian(Config{BaseURL: "https://b.test", APIKey: "secret", Model: "wan"}, nil), "/api/v1/services/aigc/image-generation/generation", "/api/v1/tasks/task-1", `{"output":{"task_id":"task-1","task_status":"SUCCEEDED","results":[{"url":"https://cdn.test/a.png"}]}}`},
|
|
{"seedance", NewSeedance(Config{BaseURL: "https://s.test/api/v3", APIKey: "secret", Model: "seed"}, nil), "/api/v3/contents/generations/tasks", "/api/v3/contents/generations/tasks/task-1", `{"id":"task-1","status":"succeeded","content":{"video_url":"https://cdn.test/a.mp4"}}`},
|
|
}
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
var calls int
|
|
client := roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
|
calls++
|
|
want := tt.wantSubmitPath
|
|
if calls == 2 {
|
|
want = tt.wantQueryPath
|
|
}
|
|
if r.URL.Path != want {
|
|
t.Fatalf("path=%s want=%s", r.URL.Path, want)
|
|
}
|
|
if !strings.HasPrefix(r.Header.Get("Authorization"), "Bearer ") {
|
|
t.Fatal("missing bearer")
|
|
}
|
|
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(tt.response)), Header: http.Header{}}, nil
|
|
})
|
|
switch a := tt.adapter.(type) {
|
|
case *EvoLink:
|
|
a.client = client
|
|
case *Bailian:
|
|
a.client = client
|
|
case *Seedance:
|
|
a.client = client
|
|
}
|
|
submitted, err := tt.adapter.Submit(context.Background(), Request{Capability: "image.generate", Prompt: "hello"})
|
|
if err != nil || submitted.TaskID != "task-1" {
|
|
t.Fatalf("Submit=%#v,%v", submitted, err)
|
|
}
|
|
queried, err := tt.adapter.Query(context.Background(), "task-1")
|
|
if err != nil || queried.Status != "succeeded" || len(queried.OutputURLs) != 1 {
|
|
t.Fatalf("Query=%#v,%v", queried, err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestVolcengineSignsSubmitAndMapsResult(t *testing.T) {
|
|
var request *http.Request
|
|
client := roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
|
request = r
|
|
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"data":{"task_id":"task-1","status":"done","image_urls":["https://cdn.test/a.png"]}}`)), Header: http.Header{}}, nil
|
|
})
|
|
a := NewVolcengine(Config{BaseURL: "https://visual.test", AccessKeyID: "ak", SecretAccessKey: "sk", Region: "cn-north-1", Service: "cv", Model: "jimeng"}, client, func() time.Time { return time.Date(2026, 8, 13, 0, 0, 0, 0, time.UTC) })
|
|
got, err := a.Submit(context.Background(), Request{Capability: "image.generate", Prompt: "hello"})
|
|
if err != nil || got.TaskID != "task-1" {
|
|
t.Fatalf("Submit=%#v,%v", got, err)
|
|
}
|
|
if request.URL.Query().Get("Action") != "CVSync2AsyncSubmitTask" || !strings.Contains(request.Header.Get("Authorization"), "Credential=ak/") || request.Header.Get("X-Date") != "20260813T000000Z" {
|
|
t.Fatalf("request=%#v headers=%#v", request.URL, request.Header)
|
|
}
|
|
}
|
|
|
|
func TestVolcenginePayloadsMatchJimengSubmitAndQueryProtocols(t *testing.T) {
|
|
var bodies []map[string]any
|
|
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
|
var body map[string]any
|
|
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
bodies = append(bodies, body)
|
|
response := `{"data":{"task_id":"task-1","status":"queued"}}`
|
|
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(response)), Header: http.Header{}}, nil
|
|
})
|
|
adapter := NewVolcengine(Config{BaseURL: "https://visual.test", AccessKeyID: "ak", SecretAccessKey: "sk", Model: "jimeng"}, client, func() time.Time { return time.Date(2026, 8, 13, 0, 0, 0, 0, time.UTC) })
|
|
if _, err := adapter.Submit(context.Background(), Request{Prompt: "draw", InputURLs: []string{"https://cdn.test/ref.png"}, Settings: map[string]any{"width": 1024, "height": 768, "force_single": true, "ignored": "value"}}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := adapter.Query(context.Background(), "task-1"); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if bodies[0]["width"] != float64(1024) || bodies[0]["height"] != float64(768) || bodies[0]["force_single"] != true || bodies[0]["ignored"] != nil {
|
|
t.Fatalf("submit body=%#v", bodies[0])
|
|
}
|
|
queryJSON, ok := bodies[1]["req_json"].(string)
|
|
if !ok || queryJSON == "" {
|
|
t.Fatalf("query body=%#v", bodies[1])
|
|
}
|
|
var queryOptions map[string]any
|
|
if err := json.Unmarshal([]byte(queryJSON), &queryOptions); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
logo := queryOptions["logo_info"].(map[string]any)
|
|
if queryOptions["return_url"] != true || logo["add_logo"] != false || logo["opacity"] != float64(1) {
|
|
t.Fatalf("query options=%#v", queryOptions)
|
|
}
|
|
}
|
|
|
|
func TestVolcengineQueryUsesPersistedTaskModelInsteadOfAdapterFallback(t *testing.T) {
|
|
var body map[string]any
|
|
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
|
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"data":{"task_id":"task-1","status":"queued"}}`)), Header: http.Header{}}, nil
|
|
})
|
|
adapter := NewVolcengine(Config{BaseURL: "https://visual.test", AccessKeyID: "ak", SecretAccessKey: "sk", Model: "fallback-model-b"}, client, func() time.Time {
|
|
return time.Date(2026, 8, 13, 0, 0, 0, 0, time.UTC)
|
|
})
|
|
|
|
if _, err := adapter.QueryModel(context.Background(), "task-1", "persisted-model-a"); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if body["req_key"] != "persisted-model-a" || body["task_id"] != "task-1" {
|
|
t.Fatalf("query body=%#v", body)
|
|
}
|
|
}
|
|
|
|
func TestProviderErrorsAreGenericAndDoNotLeakSecrets(t *testing.T) {
|
|
client := roundTripFunc(func(*http.Request) (*http.Response, error) {
|
|
return &http.Response{StatusCode: 401, Body: io.NopCloser(strings.NewReader(`{"message":"secret upstream detail"}`)), Header: http.Header{}}, nil
|
|
})
|
|
a := NewEvoLink(Config{BaseURL: "https://e.test", APIKey: "very-secret", Model: "m"}, client)
|
|
_, err := a.Submit(context.Background(), Request{Capability: "image.generate", Prompt: "p"})
|
|
if err == nil || strings.Contains(err.Error(), "secret") {
|
|
t.Fatalf("error=%v", err)
|
|
}
|
|
}
|
|
|
|
func TestBailianUsesThePreparedRequestModelForImageAndVideo(t *testing.T) {
|
|
models := []string{}
|
|
client := roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
|
var body map[string]any
|
|
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
models = append(models, body["model"].(string))
|
|
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"output":{"task_id":"task-1","task_status":"PENDING"}}`)), Header: http.Header{}}, nil
|
|
})
|
|
adapter := NewBailian(Config{BaseURL: "https://b.test", APIKey: "secret", Model: "fallback-model"}, client)
|
|
if _, err := adapter.Submit(context.Background(), Request{Capability: "image.generate", Model: "image-model", Prompt: "image"}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := adapter.Submit(context.Background(), Request{Capability: "video.generate", Model: "video-model", Prompt: "video"}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(models) != 2 || models[0] != "image-model" || models[1] != "video-model" {
|
|
t.Fatalf("models=%#v", models)
|
|
}
|
|
}
|
|
|
|
func TestBailianPayloadsMatchImageAndVideoProtocols(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
request Request
|
|
check func(*testing.T, map[string]any)
|
|
}{
|
|
{
|
|
name: "image messages and parameters",
|
|
request: Request{Capability: "image.generate", Model: "wan-image", Prompt: "draw", InputURLs: []string{"https://cdn.test/ref.png"}, Settings: map[string]any{
|
|
"width": 1024, "height": 768,
|
|
}},
|
|
check: func(t *testing.T, body map[string]any) {
|
|
input := body["input"].(map[string]any)
|
|
messages := input["messages"].([]any)
|
|
content := messages[0].(map[string]any)["content"].([]any)
|
|
parameters := body["parameters"].(map[string]any)
|
|
if len(content) != 2 || content[0].(map[string]any)["image"] != "https://cdn.test/ref.png" || content[1].(map[string]any)["text"] != "draw" {
|
|
t.Fatalf("content=%#v", content)
|
|
}
|
|
if parameters["size"] != "1024*768" || parameters["n"] != float64(1) || parameters["watermark"] != false {
|
|
t.Fatalf("parameters=%#v", parameters)
|
|
}
|
|
if _, exists := parameters["thinking_mode"]; exists {
|
|
t.Fatalf("editing parameters=%#v", parameters)
|
|
}
|
|
},
|
|
},
|
|
{
|
|
name: "video first and last frames",
|
|
request: Request{Capability: "video.generate", Model: "wan-video", Prompt: "move", InputURLs: []string{"https://cdn.test/first.png", "https://cdn.test/last.png"}, Settings: map[string]any{
|
|
"resolution": "1080p", "duration": 8,
|
|
}},
|
|
check: func(t *testing.T, body map[string]any) {
|
|
input := body["input"].(map[string]any)
|
|
media := input["media"].([]any)
|
|
parameters := body["parameters"].(map[string]any)
|
|
if len(media) != 2 || media[0].(map[string]any)["type"] != "first_frame" || media[1].(map[string]any)["type"] != "last_frame" {
|
|
t.Fatalf("media=%#v", media)
|
|
}
|
|
if parameters["resolution"] != "1080P" || parameters["duration"] != float64(8) || parameters["prompt_extend"] != true || parameters["watermark"] != false {
|
|
t.Fatalf("parameters=%#v", parameters)
|
|
}
|
|
},
|
|
},
|
|
}
|
|
for _, test := range tests {
|
|
t.Run(test.name, func(t *testing.T) {
|
|
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
|
var body map[string]any
|
|
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
test.check(t, body)
|
|
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"output":{"task_id":"task-1","task_status":"PENDING"}}`)), Header: http.Header{}}, nil
|
|
})
|
|
if _, err := NewBailian(Config{BaseURL: "https://b.test", APIKey: "secret"}, client).Submit(context.Background(), test.request); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestBailianDecodesCompatibleModeChoiceImages(t *testing.T) {
|
|
result := decodeBailian([]byte(`{"output":{"task_id":"task-1","task_status":"SUCCEEDED","choices":[{"message":{"content":[{"image":"https://cdn.test/choice.png"}]}}]}}`))
|
|
if result.Status != StatusSucceeded || len(result.OutputURLs) != 1 || result.OutputURLs[0] != "https://cdn.test/choice.png" {
|
|
t.Fatalf("result=%#v", result)
|
|
}
|
|
}
|
|
|
|
func TestSeedancePreservesTypedMultimodalMaterials(t *testing.T) {
|
|
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
|
var body map[string]any
|
|
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
content := body["content"].([]any)
|
|
if len(content) != 4 {
|
|
t.Fatalf("content=%#v", content)
|
|
}
|
|
checks := []struct{ materialType, urlKey, role string }{
|
|
{"image_url", "image_url", "reference_image"},
|
|
{"video_url", "video_url", "reference_video"},
|
|
{"audio_url", "audio_url", "reference_audio"},
|
|
}
|
|
for index, check := range checks {
|
|
item := content[index+1].(map[string]any)
|
|
if item["type"] != check.materialType || item["role"] != check.role || item["label"] != []string{"图片1", "视频1", "音频1"}[index] {
|
|
t.Fatalf("content[%d]=%#v", index+1, item)
|
|
}
|
|
if object := item[check.urlKey].(map[string]any); object["url"] == "" {
|
|
t.Fatalf("content[%d]=%#v", index+1, item)
|
|
}
|
|
}
|
|
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"id":"task-1","status":"queued"}`)), Header: http.Header{}}, nil
|
|
})
|
|
request := Request{
|
|
Capability: "video.generate", Prompt: "combine",
|
|
Materials: []Material{
|
|
{URL: "https://cdn.test/image.png", Type: MaterialImage, Label: "图片1"},
|
|
{URL: "https://cdn.test/video.mp4", Type: MaterialVideo, Label: "视频1"},
|
|
{URL: "https://cdn.test/audio.mp3", Type: MaterialAudio, Label: "音频1"},
|
|
},
|
|
}
|
|
if _, err := NewSeedance(Config{BaseURL: "https://s.test/api/v3", APIKey: "secret", Model: "seed"}, client).Submit(context.Background(), request); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestEvoLinkPayloadIncludesQualityAndSize(t *testing.T) {
|
|
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
|
var body map[string]any
|
|
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if body["quality"] != "high" || body["size"] != "16:9" || body["resolution"] != "1K" {
|
|
t.Fatalf("body=%#v", body)
|
|
}
|
|
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"id":"task-1","status":"queued"}`)), Header: http.Header{}}, nil
|
|
})
|
|
_, err := NewEvoLink(Config{BaseURL: "https://e.test", APIKey: "secret", Model: "gpt-image-2"}, client).Submit(context.Background(), Request{
|
|
Capability: "image.generate", Prompt: "draw", Settings: map[string]any{"quality": " high ", "width": 1920, "height": 1080},
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestMockIsDeterministic(t *testing.T) {
|
|
m := NewMock("fixture")
|
|
a, _ := m.Submit(context.Background(), Request{Capability: "video.generate", Prompt: "hello"})
|
|
b, _ := m.Submit(context.Background(), Request{Capability: "video.generate", Prompt: "hello"})
|
|
if a.TaskID != b.TaskID {
|
|
t.Fatalf("task IDs differ: %s %s", a.TaskID, b.TaskID)
|
|
}
|
|
got, _ := m.Query(context.Background(), a.TaskID)
|
|
if got.Status != "succeeded" || len(got.OutputURLs) != 1 {
|
|
t.Fatalf("Query=%#v", got)
|
|
}
|
|
}
|