Files
NianAIGC/backend/internal/jobs/provider_test.go

338 lines
15 KiB
Go

package jobs
import (
"context"
"encoding/json"
"errors"
"reflect"
"strings"
"testing"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/providers"
)
func TestProviderJobBuilderPreparesImageRequestAndEngineOverride(t *testing.T) {
b := testProviderBuilder()
body := map[string]any{
"engine": "evolink",
"promptAssembly": map[string]any{
"mode": "image", "manualPrompt": " assembled image ",
"materials": []any{map[string]any{"id": "asset-1", "url": "https://in.test/reference.png", "type": "image"}},
},
"scale": 0.7, "width": 1536.0, "height": 1024.0,
"min_ratio": 0.5, "max_ratio": 2.0, "force_single": true,
"quality": " HIGH ", "settings": map[string]any{"imageCount": 2.0},
}
cmd, err := b.Build(context.Background(), "owner", "client", "image.generate", "idem", body)
if err != nil {
t.Fatal(err)
}
wantRequest := providers.Request{
Capability: "image.generate", Model: "gpt-image-test", Prompt: "assembled image", InputURLs: []string{"https://in.test/reference.png"},
Materials: []providers.Material{{URL: "https://in.test/reference.png", Type: providers.MaterialImage, Label: "@图片1"}},
Settings: map[string]any{"scale": 0.7, "width": float64(1536), "height": float64(1024), "min_ratio": 0.5, "max_ratio": 2.0, "force_single": true, "quality": "high", "imageCount": float64(2)},
}
assertPreparedJob(t, cmd.Job, "evolink", "gpt-image-test", wantRequest)
if !reflect.DeepEqual(cmd.Job.InputAssetIDs, []string{"asset-1"}) {
t.Fatalf("InputAssetIDs = %#v", cmd.Job.InputAssetIDs)
}
}
func TestProviderJobBuilderUsesPublicInputURLsForImageMaterialCoverage(t *testing.T) {
b := testProviderBuilder()
cmd, err := b.Build(context.Background(), "owner", "client", "image.generate", "", map[string]any{
"prompt": "compose @图片2", "inputUrls": []any{"https://in.test/one.png", "https://in.test/two.png"},
})
if err != nil {
t.Fatal(err)
}
var request providers.Request
if err := json.Unmarshal(cmd.Job.RequestPayload, &request); err != nil {
t.Fatal(err)
}
if !reflect.DeepEqual(request.InputURLs, []string{"https://in.test/one.png", "https://in.test/two.png"}) {
t.Fatalf("InputURLs = %#v", request.InputURLs)
}
}
func TestProviderJobBuilderPreparesSeedanceVideoRequest(t *testing.T) {
b := testProviderBuilder()
body := map[string]any{
"engine": "seedance", "prompt": " launch video ",
"materials": []any{
map[string]any{"id": "image-1", "url": "https://in.test/first.png", "type": "image"},
map[string]any{"id": "video-1", "url": "https://in.test/reference.mp4", "type": "video"},
map[string]any{"id": "audio-1", "url": "https://in.test/music.mp3", "type": "audio"},
},
"settings": map[string]any{"ratio": "16:9", "duration": 8.0, "resolution": "1080p", "ignored": "value"},
}
cmd, err := b.Build(context.Background(), "owner", "", "video.generate", "", body)
if err != nil {
t.Fatal(err)
}
wantRequest := providers.Request{
Capability: "video.generate", Model: "seedance-test", Prompt: "launch video",
InputURLs: []string{"https://in.test/first.png", "https://in.test/reference.mp4", "https://in.test/music.mp3"},
Materials: []providers.Material{
{URL: "https://in.test/first.png", Type: providers.MaterialImage, Label: "@图片1"},
{URL: "https://in.test/reference.mp4", Type: providers.MaterialVideo, Label: "@视频1"},
{URL: "https://in.test/music.mp3", Type: providers.MaterialAudio, Label: "@音频1"},
},
Settings: map[string]any{"ratio": "16:9", "duration": float64(8), "resolution": "1080p"},
}
assertPreparedJob(t, cmd.Job, "seedance", "seedance-test", wantRequest)
}
func TestProviderJobBuilderAppliesInjectedVideoDefaults(t *testing.T) {
b := testProviderBuilder()
cmd, err := b.Build(context.Background(), "owner", "", "video.generate", "", map[string]any{
"engine": "seedance", "prompt": "video",
})
if err != nil {
t.Fatal(err)
}
var request providers.Request
if err := json.Unmarshal(cmd.Job.RequestPayload, &request); err != nil {
t.Fatal(err)
}
want := map[string]any{"ratio": "9:16", "duration": float64(5), "resolution": "720p"}
if !reflect.DeepEqual(request.Settings, want) {
t.Fatalf("Settings = %#v, want %#v", request.Settings, want)
}
}
func TestProviderJobBuilderPreparesBailianFirstAndLastFrameVideo(t *testing.T) {
b := testProviderBuilder()
body := map[string]any{
"engine": "bailian",
"promptAssembly": map[string]any{
"mode": "video", "manualPrompt": "animate frames",
"materials": []any{
map[string]any{"id": "first", "url": "https://in.test/first.png", "type": "image"},
map[string]any{"id": "last", "url": "https://in.test/last.png", "type": "image"},
},
},
"settings": map[string]any{"duration": 12.0, "resolution": "1080p"},
}
cmd, err := b.Build(context.Background(), "owner", "", "video.generate", "", body)
if err != nil {
t.Fatal(err)
}
wantRequest := providers.Request{
Capability: "video.generate", Model: "bailian-video-test", Prompt: "animate frames",
InputURLs: []string{"https://in.test/first.png", "https://in.test/last.png"},
Materials: []providers.Material{
{URL: "https://in.test/first.png", Type: providers.MaterialImage, Label: "@图片1"},
{URL: "https://in.test/last.png", Type: providers.MaterialImage, Label: "@图片2"},
},
Settings: map[string]any{"duration": float64(12), "resolution": "1080P"},
}
assertPreparedJob(t, cmd.Job, "bailian", "bailian-video-test", wantRequest)
}
func TestProviderJobBuilderRejectsInvalidPreparation(t *testing.T) {
tests := []struct {
name, capability, message string
body map[string]any
}{
{name: "missing image prompt", capability: "image.generate", body: map[string]any{}, message: "prompt is required"},
{name: "arbitrary provider", capability: "image.generate", body: map[string]any{"prompt": "p", "engine": "attacker-provider"}, message: "unsupported image engine"},
{name: "bailian too many references", capability: "image.generate", body: map[string]any{"prompt": "p", "engine": "bailian", "imageUrls": tenURLs()}, message: "at most 9 reference images"},
{name: "bailian invalid image pixels", capability: "image.generate", body: map[string]any{"prompt": "p", "engine": "bailian", "width": 100.0, "height": 100.0}, message: "image dimensions"},
{name: "seedance missing materials referenced by prompt", capability: "video.generate", body: map[string]any{"engine": "seedance", "prompt": "use @图片2"}, message: "requires @图片2"},
{name: "bailian requires frame", capability: "video.generate", body: map[string]any{"engine": "bailian", "prompt": "p", "materials": []any{}}, message: "1 or 2 image materials"},
{name: "bailian rejects video material", capability: "video.generate", body: map[string]any{"engine": "bailian", "prompt": "p", "materials": []any{map[string]any{"url": "https://in.test/a.mp4", "type": "video"}}}, message: "1 or 2 image materials"},
{name: "bad seedance settings", capability: "video.generate", body: map[string]any{"engine": "seedance", "prompt": "p", "settings": map[string]any{"duration": 99.0}}, message: "video duration"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
_, err := testProviderBuilder().Build(context.Background(), "owner", "", test.capability, "", test.body)
if err == nil || !strings.Contains(strings.ToLower(err.Error()), strings.ToLower(test.message)) {
t.Fatalf("error = %v, want containing %q", err, test.message)
}
var jobErr *Error
if !strings.Contains(err.Error(), "configured") && (!errorsAs(err, &jobErr) || jobErr.Status != 400 || jobErr.Kind != ErrorInvalid) {
t.Fatalf("error = %#v, want generic invalid job error", err)
}
})
}
}
func TestProviderBuilderAndProcessor(t *testing.T) {
b := testProviderBuilder()
cmd, err := b.Build(context.Background(), "owner", "client", "image.generate", "idem", map[string]any{"prompt": "hello", "inputUrls": []any{"https://in.test/a.png"}})
if err != nil || cmd.Job.OwnerID != "owner" || cmd.Job.ExternalClientID != "client" {
t.Fatalf("Build=%#v,%v", cmd, err)
}
store := newMemoryJobStore()
store.jobs[cmd.Job.ID] = cmd.Job
p := ProviderProcessor{Providers: ProviderRegistry{"image-default": providers.NewMock("fixture")}, Store: store}
submitted, err := p.Advance(context.Background(), cmd.Job)
if err != nil || submitted.ProviderTaskID == "" || submitted.Status != StatusQueued {
t.Fatalf("submit=%#v,%v", submitted, err)
}
done, err := p.Advance(context.Background(), submitted)
if err != nil || done.Status != StatusSucceeded {
t.Fatalf("query=%#v,%v", done, err)
}
if store.jobs[cmd.Job.ID].Status != StatusSucceeded {
t.Fatal("provider result was not persisted")
}
var persisted providers.HTTPResult
if err := json.Unmarshal(done.ResponsePayload, &persisted); err != nil || len(persisted.OutputURLs) != 1 || persisted.OutputURLs[0] == "" {
t.Fatalf("persisted result = %#v, %v raw=%s", persisted, err, done.ResponsePayload)
}
}
func TestProviderProcessorNeverResubmitsAfterPersistedDispatchIntent(t *testing.T) {
store := newMemoryJobStore()
started := time.Date(2026, 8, 13, 8, 0, 0, 0, time.UTC)
job := Job{ID: "job-unknown", OwnerID: "owner", Provider: "mock", Capability: "image.generate", Status: StatusRunning, LockedBy: "worker", ProviderDispatchStartedAt: &started, RequestPayload: json.RawMessage(`{"capability":"image.generate","model":"fixture","prompt":"hello"}`)}
store.jobs[job.ID] = job
adapter := &countingProvider{result: providers.Result{TaskID: "duplicate", Status: providers.StatusQueued}}
processor := ProviderProcessor{Providers: ProviderRegistry{"mock": adapter}, Store: store}
got, err := processor.Advance(context.Background(), job)
if err != nil || adapter.submits != 0 || got.Status != StatusFailed || got.Error == nil || got.Error.Retryable {
t.Fatalf("got=%#v err=%v submits=%d", got, err, adapter.submits)
}
}
func TestProviderProcessorClearsTransientErrorAfterSuccessfulPoll(t *testing.T) {
store := newMemoryJobStore()
job := Job{ID: "job-recovered", OwnerID: "owner", Provider: "fixture", ReqKey: "model-a", Capability: "image.generate", Status: StatusQueued, LockedBy: "worker", ProviderTaskID: "provider-task", Error: &JobError{Message: "temporary timeout", Retryable: true}, RequestPayload: json.RawMessage(`{"capability":"image.generate","model":"model-a","prompt":"hello"}`)}
store.jobs[job.ID] = job
adapter := &countingProvider{result: providers.Result{TaskID: "provider-task", Status: providers.StatusSucceeded, OutputURLs: []string{"https://cdn.test/result.png"}}}
processor := ProviderProcessor{Providers: ProviderRegistry{"fixture": adapter}, Store: store}
got, err := processor.Advance(context.Background(), job)
if err != nil || got.Status != StatusSucceeded || got.Error != nil {
t.Fatalf("got=%#v err=%v", got, err)
}
}
func TestProviderProcessorQueriesWithPersistedRequestModel(t *testing.T) {
store := newMemoryJobStore()
job := Job{ID: "job-model", OwnerID: "owner", Provider: "fixture", ReqKey: "persisted-model-a", Capability: "image.generate", Status: StatusQueued, LockedBy: "worker", ProviderTaskID: "provider-task", RequestPayload: json.RawMessage(`{"capability":"image.generate","model":"persisted-model-a","prompt":"hello"}`)}
store.jobs[job.ID] = job
adapter := &modelQueryProvider{result: providers.Result{TaskID: "provider-task", Status: providers.StatusRunning}}
processor := ProviderProcessor{Providers: ProviderRegistry{"fixture": adapter}, Store: store}
if _, err := processor.Advance(context.Background(), job); err != nil {
t.Fatal(err)
}
if adapter.model != "persisted-model-a" {
t.Fatalf("query model = %q", adapter.model)
}
}
func TestProviderJobBuilderRejectsInvalidPublicWebhookURL(t *testing.T) {
for _, value := range []string{"/internal/callback", "javascript:alert(1)", "https://user:pass@example.test/hook"} {
_, err := testProviderBuilder().Build(context.Background(), "api:client", "client", "image.generate", "", map[string]any{"prompt": "hello", "webhookUrl": value})
var invalid *Error
if !errorsAs(err, &invalid) || invalid.Status != 400 {
t.Fatalf("webhook %q error = %#v", value, err)
}
}
cmd, err := testProviderBuilder().Build(context.Background(), "api:client", "client", "image.generate", "", map[string]any{"prompt": "hello", "webhookUrl": "https://hooks.example.test/done"})
if err != nil || cmd.Job.WebhookURL != "https://hooks.example.test/done" {
t.Fatalf("valid webhook = %#v, %v", cmd.Job, err)
}
}
func TestProviderJobBuilderClampsPriorityToPublicContract(t *testing.T) {
for _, test := range []struct {
raw any
want int
}{{raw: 999.0, want: 100}, {raw: -999.0, want: -100}, {raw: 42.0, want: 42}} {
cmd, err := testProviderBuilder().Build(context.Background(), "owner", "client", "image.generate", "", map[string]any{"prompt": "hello", "priority": test.raw})
if err != nil || cmd.Job.Priority != test.want {
t.Fatalf("priority(%v)=%d err=%v, want %d", test.raw, cmd.Job.Priority, err, test.want)
}
}
}
type countingProvider struct {
submits int
result providers.Result
}
type modelQueryProvider struct {
result providers.Result
model string
}
func (provider *modelQueryProvider) Submit(context.Context, providers.Request) (providers.Result, error) {
return provider.result, nil
}
func (provider *modelQueryProvider) Query(context.Context, string) (providers.Result, error) {
return providers.Result{}, errors.New("fallback query must not be used")
}
func (provider *modelQueryProvider) QueryModel(_ context.Context, _ string, model string) (providers.Result, error) {
provider.model = model
return provider.result, nil
}
func (p *countingProvider) Submit(context.Context, providers.Request) (providers.Result, error) {
p.submits++
return p.result, nil
}
func (p *countingProvider) Query(context.Context, string) (providers.Result, error) {
return p.result, nil
}
func testProviderBuilder() ProviderJobBuilder {
return ProviderJobBuilder{
ImageProvider: "image-default", ImageModel: "image-default-model",
VideoProvider: "video-default", VideoModel: "video-default-model",
ImageEngine: "jimeng", VideoEngine: "seedance",
ImageEngines: map[string]ProviderTarget{
"jimeng": {Provider: "volcengine-visual", Model: "jimeng-test"},
"evolink": {Provider: "evolink", Model: "gpt-image-test"},
"bailian": {Provider: "bailian", Model: "bailian-image-test"},
},
VideoEngines: map[string]ProviderTarget{
"seedance": {Provider: "seedance", Model: "seedance-test", Settings: map[string]any{"ratio": "9:16", "duration": 5, "resolution": "720p"}},
"bailian": {Provider: "bailian", Model: "bailian-video-test"},
},
NewID: func() string { return "job-1" },
}
}
func assertPreparedJob(t *testing.T, job Job, provider, model string, want providers.Request) {
t.Helper()
if job.Provider != provider || job.ReqKey != model {
t.Fatalf("provider/model = %q/%q, want %q/%q", job.Provider, job.ReqKey, provider, model)
}
var got providers.Request
if err := json.Unmarshal(job.RequestPayload, &got); err != nil {
t.Fatal(err)
}
if !reflect.DeepEqual(got, want) {
t.Fatalf("request = %#v, want %#v; raw=%s", got, want, job.RequestPayload)
}
if job.Prompt != want.Prompt || !reflect.DeepEqual(job.InputURLs, want.InputURLs) {
t.Fatalf("job preparation = prompt %q URLs %#v", job.Prompt, job.InputURLs)
}
}
func tenURLs() []any {
urls := make([]any, 10)
for i := range urls {
urls[i] = "https://in.test/reference.png"
}
return urls
}
// Kept local so the test remains compatible with the package's supported Go version.
func errorsAs(err error, target any) bool {
jobErr, ok := err.(*Error)
pointer, targetOK := target.(**Error)
if ok && targetOK {
*pointer = jobErr
}
return ok && targetOK
}