增加即梦的生图和生视频功能

This commit is contained in:
andy committed 2026-08-19 17:58:22 +08:00
1 parent a9a4cdf125
commit 542246e4ce
20 files changed
+1008 -53

No files matched your search

+8 -4
View File
@@ -39,7 +39,11 @@ func firstNonEmpty(values ...string) string {
}
func positiveInt64Env(getenv postgres.Getenv, name string, fallback int64) int64 {
value, err := strconv.ParseInt(strings.TrimSpace(getenv(name)), 10, 64)
return positiveInt64Value(getenv(name), fallback)
}
func positiveInt64Value(raw string, fallback int64) int64 {
value, err := strconv.ParseInt(strings.TrimSpace(raw), 10, 64)
if err != nil || value <= 0 {
return fallback
}
@@ -180,9 +184,9 @@ func providerVideoTargets(getenv postgres.Getenv) map[string]jobs.ProviderTarget
Provider: "seedance",
Model: firstNonEmpty(getenv("SEEDANCE_MODEL"), "doubao-seedance-2-0-260128"),
Settings: map[string]any{
"ratio": firstNonEmpty(getenv("SEEDANCE_DEFAULT_RATIO"), "9:16"),
"duration": float64(positiveInt64Env(getenv, "SEEDANCE_DEFAULT_DURATION", 5)),
"resolution": firstNonEmpty(getenv("SEEDANCE_DEFAULT_RESOLUTION"), "720p"),
"ratio": firstNonEmpty(getenv("SEEDANCE_RATIO"), getenv("SEEDANCE_DEFAULT_RATIO"), "9:16"),
"duration": float64(positiveInt64Value(firstNonEmpty(getenv("SEEDANCE_DURATION"), getenv("SEEDANCE_DEFAULT_DURATION")), 5)),
"resolution": firstNonEmpty(getenv("SEEDANCE_RESOLUTION"), getenv("SEEDANCE_DEFAULT_RESOLUTION"), "720p"),
},
},
"bailian": {
@@ -273,6 +273,21 @@ func TestProviderTargetsNeverSelectRemovedProvider(t *testing.T) {
}
}
func TestProviderVideoTargetsUseDocumentedSeedanceDefaults(t *testing.T) {
values := map[string]string{
"SEEDANCE_RATIO": "16:9",
"SEEDANCE_DURATION": "12",
"SEEDANCE_RESOLUTION": "1080p",
}
target := providerVideoTargets(func(name string) string { return values[name] })["seedance"]
if target.Model != "doubao-seedance-2-0-260128" {
t.Fatalf("model=%q", target.Model)
}
if target.Settings["ratio"] != "16:9" || target.Settings["duration"] != float64(12) || target.Settings["resolution"] != "1080p" {
t.Fatalf("settings=%#v", target.Settings)
}
}
func TestCapabilitySummaryMatchesConfiguredDefaultVideoEngine(t *testing.T) {
for _, test := range []struct {
name, engine, wantEngine, wantProvider, wantModel string
+79 -1
View File
@@ -5,6 +5,7 @@ import (
"encoding/json"
"errors"
"fmt"
"log"
"math"
"net/url"
"strings"
@@ -62,9 +63,16 @@ func (p ProviderProcessor) Advance(ctx context.Context, job Job) (Job, error) {
}
}
var result providers.Result
phase := "submit"
expectedStatus := job.Status
if job.ProviderTaskID == "" {
if job.ProviderDispatchStartedAt != nil {
logProviderFailureDiagnostic(providerFailureDiagnostic{
JobID: safeProviderLogToken(job.ID),
Provider: safeProviderLogToken(job.Provider),
Phase: "submit",
ErrorClass: "unknown_outcome",
})
failed := StatusFailed
failure := &JobError{Message: "provider submission outcome is unknown; refusing duplicate submission", Retryable: false}
if p.Store == nil {
@@ -83,12 +91,16 @@ func (p ProviderProcessor) Advance(ctx context.Context, job Job) (Job, error) {
}
result, err = adapter.Submit(ctx, request)
} else {
phase = "query"
if modeled, ok := adapter.(providers.ModelQueryAdapter); ok {
result, err = modeled.QueryModel(ctx, job.ProviderTaskID, job.ReqKey)
} else {
result, err = adapter.Query(ctx, job.ProviderTaskID)
}
}
if err != nil {
logProviderFailure(job, phase, err)
}
if err != nil && job.ProviderTaskID == "" && p.Store != nil {
failed := StatusFailed
failure := &JobError{Message: "provider submission outcome is unknown; refusing duplicate submission", Retryable: false}
@@ -140,6 +152,69 @@ func (p ProviderProcessor) Advance(ctx context.Context, job Job) (Job, error) {
return job, nil
}
type providerFailureDiagnostic struct {
JobID string
Provider string
Phase string
Status int
ErrorClass string
}
func logProviderFailure(job Job, phase string, err error) {
diagnostic := providerFailureDiagnostic{
JobID: safeProviderLogToken(job.ID),
Provider: safeProviderLogToken(job.Provider),
Phase: safeProviderPhase(phase),
ErrorClass: "provider",
}
if errors.Is(err, context.Canceled) {
diagnostic.ErrorClass = "canceled"
} else if errors.Is(err, context.DeadlineExceeded) {
diagnostic.ErrorClass = "timeout"
} else {
var providerError *providers.ProviderError
if errors.As(err, &providerError) {
diagnostic.Status = providerError.Status
if providerError.Status > 0 {
diagnostic.ErrorClass = "service"
}
}
}
logProviderFailureDiagnostic(diagnostic)
}
func logProviderFailureDiagnostic(diagnostic providerFailureDiagnostic) {
log.Printf(
"zhinian-api generation provider failed jobId=%s provider=%s phase=%s status=%d errorClass=%s",
diagnostic.JobID,
diagnostic.Provider,
diagnostic.Phase,
diagnostic.Status,
diagnostic.ErrorClass,
)
}
func safeProviderPhase(phase string) string {
if phase == "query" {
return "query"
}
return "submit"
}
func safeProviderLogToken(value string) string {
value = strings.TrimSpace(value)
if len(value) == 0 || len(value) > 96 {
return "invalid"
}
for _, character := range value {
if (character >= 'a' && character <= 'z') || (character >= 'A' && character <= 'Z') || (character >= '0' && character <= '9') || strings.ContainsRune("-_.:", character) {
continue
}
return "invalid"
}
return value
}
func (p ProviderProcessor) refreshAssetURLs(ctx context.Context, job Job, request providers.Request) (providers.Request, error) {
if p.AssetURLs == nil || len(job.InputAssetIDs) == 0 {
return request, nil
@@ -561,6 +636,9 @@ func videoSettings(engine string, defaults map[string]any, raw any, materials []
if engine != "" && engine != "seedance" {
return nil, invalidPreparation("unsupported video engine")
}
if len(materials) > 4 {
return nil, invalidPreparation("seedance supports at most 4 materials")
}
ratio := stringValue(input["ratio"])
if ratio != "" && !oneOf(ratio, "9:16", "16:9", "1:1", "4:3", "3:4", "21:9", "adaptive") {
return nil, invalidPreparation("unsupported video ratio")
@@ -576,7 +654,7 @@ func videoSettings(engine string, defaults map[string]any, raw any, materials []
settings["duration"] = duration
}
resolution := stringValue(input["resolution"])
if resolution != "" && !oneOf(resolution, "480p", "720p", "1080p", "4k") {
if resolution != "" && !oneOf(resolution, "480p", "720p", "1080p") {
return nil, invalidPreparation("unsupported video resolution")
}
if resolution != "" {
+130 -2
View File
@@ -1,9 +1,13 @@
package jobs
import (
"bytes"
"context"
"encoding/json"
"errors"
"log"
"net/http"
"net/http/httptest"
"reflect"
"strings"
"testing"
@@ -162,6 +166,16 @@ func TestProviderJobBuilderRejectsInvalidPreparation(t *testing.T) {
{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"},
{name: "seedance unsupported 4k output", capability: "video.generate", body: map[string]any{"engine": "seedance", "prompt": "p", "settings": map[string]any{"resolution": "4k"}}, message: "unsupported video resolution"},
{name: "seedance too many materials", capability: "video.generate", body: map[string]any{
"engine": "seedance", "prompt": "p", "materials": []any{
map[string]any{"url": "https://in.test/1.png", "type": "image"},
map[string]any{"url": "https://in.test/2.png", "type": "image"},
map[string]any{"url": "https://in.test/3.png", "type": "image"},
map[string]any{"url": "https://in.test/4.png", "type": "image"},
map[string]any{"url": "https://in.test/5.png", "type": "image"},
},
}, message: "at most 4 materials"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
@@ -246,6 +260,79 @@ func TestProviderProcessorNeverResubmitsAfterPersistedDispatchIntent(t *testing.
}
}
func TestProviderProcessorLogsSafeUnknownSubmissionOutcome(t *testing.T) {
var output bytes.Buffer
previousOutput, previousFlags := log.Writer(), log.Flags()
log.SetOutput(&output)
log.SetFlags(0)
t.Cleanup(func() {
log.SetOutput(previousOutput)
log.SetFlags(previousFlags)
})
started := time.Now().UTC()
job := Job{
ID: "job-unknown-outcome", OwnerID: "private-owner", Provider: "volcengine-visual", ReqKey: "private-model",
Capability: "image.generate", Status: StatusRunning, ProviderDispatchStartedAt: &started,
RequestPayload: json.RawMessage(`{"capability":"image.generate","model":"private-model","prompt":"private prompt"}`),
}
adapter := &countingProvider{}
processor := ProviderProcessor{Providers: ProviderRegistry{"volcengine-visual": adapter}}
got, err := processor.Advance(context.Background(), job)
if err != nil || got.Status != StatusFailed || adapter.submits != 0 {
t.Fatalf("got=%#v err=%v submits=%d", got, err, adapter.submits)
}
logged := output.String()
for _, expected := range []string{"jobId=job-unknown-outcome", "provider=volcengine-visual", "phase=submit", "status=0", "errorClass=unknown_outcome"} {
if !strings.Contains(logged, expected) {
t.Fatalf("log %q does not contain %q", logged, expected)
}
}
for _, secret := range []string{"private-owner", "private-model", "private prompt"} {
if strings.Contains(logged, secret) {
t.Fatalf("log leaks %q: %s", secret, logged)
}
}
}
func TestProviderProcessorLogsSafeJobCorrelationWhenSubmitFails(t *testing.T) {
var output bytes.Buffer
previousOutput, previousFlags := log.Writer(), log.Flags()
log.SetOutput(&output)
log.SetFlags(0)
t.Cleanup(func() {
log.SetOutput(previousOutput)
log.SetFlags(previousFlags)
})
store := newMemoryJobStore()
job := Job{
ID: "job-safe-log", OwnerID: "owner", Provider: "volcengine-visual", ReqKey: "jimeng_seedream46_cvtob",
Capability: "image.generate", Status: StatusRunning, LockedBy: "worker",
RequestPayload: json.RawMessage(`{"capability":"image.generate","model":"jimeng_seedream46_cvtob","prompt":"private prompt"}`),
}
store.jobs[job.ID] = job
adapter := &countingProvider{err: &providers.ProviderError{Operation: "volcengine request", Status: http.StatusForbidden}}
processor := ProviderProcessor{Providers: ProviderRegistry{"volcengine-visual": adapter}, Store: store}
got, err := processor.Advance(context.Background(), job)
if err != nil || got.Status != StatusFailed {
t.Fatalf("got=%#v err=%v", got, err)
}
logged := output.String()
for _, expected := range []string{"jobId=job-safe-log", "provider=volcengine-visual", "phase=submit", "status=403", "errorClass=service"} {
if !strings.Contains(logged, expected) {
t.Fatalf("log %q does not contain %q", logged, expected)
}
}
for _, secret := range []string{"private prompt", "jimeng_seedream46_cvtob", "owner", "volcengine request"} {
if strings.Contains(logged, secret) {
t.Fatalf("log leaks %q: %s", secret, logged)
}
}
}
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"}`)}
@@ -290,6 +377,46 @@ func TestProviderProcessorQueriesWithPersistedRequestModel(t *testing.T) {
}
}
func TestProviderProcessorCompletesOfficialVolcengineQueryResponse(t *testing.T) {
var action string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
action = request.URL.Query().Get("Action")
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{
"code":10000,
"data":{"binary_data_base64":null,"image_urls":["https://cdn.test/generated.png"],"status":"done"},
"message":"Success",
"request_id":"request-1",
"status":10000
}`))
}))
defer server.Close()
store := newMemoryJobStore()
job := Job{
ID: "job-jimeng", OwnerID: "owner", Provider: "volcengine-visual", ReqKey: "jimeng_seedream46_cvtob",
Capability: "image.generate", Status: StatusRunning, LockedBy: "worker", ProviderTaskID: "task-1",
RequestPayload: json.RawMessage(`{"capability":"image.generate","model":"jimeng_seedream46_cvtob","prompt":"hello"}`),
}
store.jobs[job.ID] = job
adapter := providers.NewVolcengine(providers.Config{
BaseURL: server.URL, AccessKeyID: "ak", SecretAccessKey: "sk", Region: "cn-north-1", Service: "cv", Model: "jimeng_seedream46_cvtob",
}, server.Client(), func() time.Time { return time.Date(2026, 8, 19, 0, 0, 0, 0, time.UTC) })
processor := ProviderProcessor{Providers: ProviderRegistry{"volcengine-visual": adapter}, Store: store}
got, err := processor.Advance(context.Background(), job)
if err != nil {
t.Fatal(err)
}
var persisted providers.HTTPResult
if err := json.Unmarshal(got.ResponsePayload, &persisted); err != nil {
t.Fatal(err)
}
if action != "JimengSeedream46CVToBGetResult" || got.Status != StatusSucceeded || got.ProviderTaskID != "task-1" || len(persisted.OutputURLs) != 1 || persisted.OutputURLs[0] != "https://cdn.test/generated.png" {
t.Fatalf("job=%#v response=%#v", got, persisted)
}
}
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})
@@ -320,6 +447,7 @@ type countingProvider struct {
submits int
result providers.Result
request providers.Request
err error
}
type recordingProviderAssetURLResolver struct {
@@ -352,10 +480,10 @@ func (provider *modelQueryProvider) QueryModel(_ context.Context, _ string, mode
func (p *countingProvider) Submit(_ context.Context, request providers.Request) (providers.Result, error) {
p.submits++
p.request = request
return p.result, nil
return p.result, p.err
}
func (p *countingProvider) Query(context.Context, string) (providers.Result, error) {
return p.result, nil
return p.result, p.err
}
func testProviderBuilder() ProviderJobBuilder {
+1 -5
View File
@@ -145,11 +145,7 @@ func NewSeedance(c Config, client HTTPClient) *Seedance {
case MaterialAudio:
materialType, urlKey, role = "audio_url", "audio_url", "reference_audio"
}
item := map[string]any{"type": materialType, urlKey: map[string]any{"url": material.URL}, "role": role}
if material.Label != "" {
item["label"] = material.Label
}
content = append(content, item)
content = append(content, map[string]any{"type": materialType, urlKey: map[string]any{"url": material.URL}, "role": role})
}
p := map[string]any{"model": requestModel(r, c.Model), "content": content, "generate_audio": true, "watermark": false}
for k, v := range r.Settings {
+301 -3
View File
@@ -1,9 +1,11 @@
package providers
import (
"bytes"
"context"
"encoding/json"
"io"
"log"
"net/http"
"strings"
"testing"
@@ -128,6 +130,10 @@ func TestVolcenginePayloadsMatchJimengSubmitAndQueryProtocols(t *testing.T) {
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])
}
imageURLs, ok := bodies[0]["image_urls"].([]any)
if !ok || len(imageURLs) != 1 || imageURLs[0] != "https://cdn.test/ref.png" {
t.Fatalf("submit image_urls=%#v", bodies[0]["image_urls"])
}
queryJSON, ok := bodies[1]["req_json"].(string)
if !ok || queryJSON == "" {
t.Fatalf("query body=%#v", bodies[1])
@@ -142,6 +148,58 @@ func TestVolcenginePayloadsMatchJimengSubmitAndQueryProtocols(t *testing.T) {
}
}
func TestVolcengineSeedream46UsesDedicatedActionsAndVersion(t *testing.T) {
var requests []*http.Request
var bodies []map[string]any
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
requests = append(requests, request.Clone(request.Context()))
var body map[string]any
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
t.Fatal(err)
}
bodies = append(bodies, body)
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(`{"code":10000,"data":{"task_id":"task-1","status":"queued"}}`)),
Header: http.Header{},
}, nil
})
adapter := NewVolcengine(Config{
BaseURL: "https://visual.test", AccessKeyID: "ak", SecretAccessKey: "sk", Model: "jimeng_seedream46_cvtob",
}, client, func() time.Time { return time.Date(2026, 8, 19, 0, 0, 0, 0, time.UTC) })
if _, err := adapter.Submit(context.Background(), Request{
Prompt: "draw", Settings: map[string]any{"scale": 50, "width": 2048, "height": 2048, "force_single": true},
}); err != nil {
t.Fatal(err)
}
if _, err := adapter.QueryModel(context.Background(), "task-1", "jimeng_seedream46_cvtob"); err != nil {
t.Fatal(err)
}
wantActions := []string{"JimengSeedream46CVToBSubmitTask", "JimengSeedream46CVToBGetResult"}
if len(requests) != len(wantActions) {
t.Fatalf("requests=%d want=%d", len(requests), len(wantActions))
}
for index, wantAction := range wantActions {
if got := requests[index].URL.Query().Get("Action"); got != wantAction {
t.Fatalf("request[%d] action=%q want=%q", index, got, wantAction)
}
if got := requests[index].URL.Query().Get("Version"); got != "2024-06-06" {
t.Fatalf("request[%d] version=%q want=2024-06-06", index, got)
}
}
if bodies[0]["req_key"] != "jimeng_seedream46_cvtob" || bodies[0]["scale"] != float64(50) {
t.Fatalf("submit body=%#v", bodies[0])
}
if _, exists := bodies[0]["image_urls"]; exists {
t.Fatalf("text-to-image submit contains empty image_urls: %#v", bodies[0])
}
if bodies[1]["req_key"] != "jimeng_seedream46_cvtob" || bodies[1]["task_id"] != "task-1" {
t.Fatalf("query body=%#v", bodies[1])
}
}
func TestVolcengineQueryUsesPersistedTaskModelInsteadOfAdapterFallback(t *testing.T) {
var body map[string]any
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
@@ -162,6 +220,240 @@ func TestVolcengineQueryUsesPersistedTaskModelInsteadOfAdapterFallback(t *testin
}
}
func TestVolcengineDecodesOfficialQueryResponse(t *testing.T) {
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
if request.URL.Query().Get("Action") != seedream46QueryAction || request.URL.Query().Get("Version") != seedream46Version {
t.Fatalf("action=%q", request.URL.Query().Get("Action"))
}
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(`{
"code":10000,
"data":{"binary_data_base64":null,"image_urls":["https://cdn.test/generated.png"],"status":"done"},
"message":"Success",
"request_id":"request-1",
"status":10000,
"time_elapsed":"508.312154ms"
}`)),
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, 19, 0, 0, 0, 0, time.UTC)
})
result, err := adapter.QueryModel(context.Background(), "task-1", "jimeng_seedream46_cvtob")
if err != nil {
t.Fatal(err)
}
if result.TaskID != "task-1" || result.Status != StatusSucceeded || len(result.OutputURLs) != 1 || result.OutputURLs[0] != "https://cdn.test/generated.png" {
t.Fatalf("result=%#v", result)
}
}
func TestVolcengineDecodesGatewayWrappedBusinessResponse(t *testing.T) {
client := roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(`{
"ResponseMetadata":{"RequestId":"request-wrapper-1"},
"Result":{
"code":10000,
"data":{"task_id":"task-wrapper-1","status":"done","image_urls":["https://cdn.test/wrapped.png"]},
"message":"Success",
"request_id":"request-business-1"
}
}`)),
Header: http.Header{},
}, nil
})
adapter := NewVolcengine(Config{BaseURL: "https://visual.test", AccessKeyID: "ak", SecretAccessKey: "sk", Model: seedream46Model}, client, func() time.Time {
return time.Date(2026, 8, 19, 0, 0, 0, 0, time.UTC)
})
result, err := adapter.Submit(context.Background(), Request{Prompt: "draw"})
if err != nil {
t.Fatal(err)
}
if result.TaskID != "task-wrapper-1" || result.Status != StatusSucceeded || len(result.OutputURLs) != 1 || result.OutputURLs[0] != "https://cdn.test/wrapped.png" {
t.Fatalf("result=%#v", result)
}
}
func TestVolcengineRejectsOfficialBusinessErrorOnHTTP200(t *testing.T) {
client := roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(`{
"code":50413,
"data":null,
"message":"Post Text Risk Not Pass",
"request_id":"request-2",
"status":50413
}`)),
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, 19, 0, 0, 0, 0, time.UTC)
})
if _, err := adapter.Submit(context.Background(), Request{Prompt: "blocked prompt"}); err == nil {
t.Fatal("expected Volcengine business error")
}
}
func TestVolcengineServiceFailureLogsOnlySafeDiagnosticFields(t *testing.T) {
var output bytes.Buffer
previousOutput, previousFlags := log.Writer(), log.Flags()
log.SetOutput(&output)
log.SetFlags(0)
t.Cleanup(func() {
log.SetOutput(previousOutput)
log.SetFlags(previousFlags)
})
client := roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusForbidden,
Body: io.NopCloser(strings.NewReader(`{
"code":50400,
"message":"Authorization failed for private prompt",
"request_id":"request-safe-1",
"data":null
}`)),
Header: http.Header{},
}, nil
})
adapter := NewVolcengine(Config{
BaseURL: "https://visual.test", AccessKeyID: "private-access-key", SecretAccessKey: "private-secret-key", Model: "jimeng",
}, client, func() time.Time { return time.Date(2026, 8, 19, 0, 0, 0, 0, time.UTC) })
if _, err := adapter.Submit(context.Background(), Request{Prompt: "private prompt"}); err == nil {
t.Fatal("expected Volcengine service error")
}
got := output.String()
for _, expected := range []string{"operation=submit", "status=403", `code="50400"`, `requestId="request-safe-1"`, "errorClass=service", "elapsedMs="} {
if !strings.Contains(got, expected) {
t.Fatalf("log %q does not contain %q", got, expected)
}
}
for _, secret := range []string{"private-access-key", "private-secret-key", "Authorization failed", "private prompt", "Signature=", "Credential="} {
if strings.Contains(got, secret) {
t.Fatalf("log leaks %q: %s", secret, got)
}
}
}
func TestVolcengineServiceFailureLogsNestedGatewayDiagnosticFields(t *testing.T) {
var output bytes.Buffer
previousOutput, previousFlags := log.Writer(), log.Flags()
log.SetOutput(&output)
log.SetFlags(0)
t.Cleanup(func() {
log.SetOutput(previousOutput)
log.SetFlags(previousFlags)
})
client := roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusBadRequest,
Body: io.NopCloser(strings.NewReader(`{
"ResponseMetadata": {
"RequestId": "request-nested-1",
"Error": {
"CodeN": 100025,
"Code": "InvalidCredential",
"Message": "Authorization contains private-secret-key and private prompt"
}
}
}`)),
Header: http.Header{"X-Tt-Logid": []string{"request-header-ignored"}},
}, nil
})
adapter := NewVolcengine(Config{
BaseURL: "https://visual.test", AccessKeyID: "private-access-key", SecretAccessKey: "private-secret-key", Model: seedream46Model,
}, client, func() time.Time { return time.Date(2026, 8, 19, 0, 0, 0, 0, time.UTC) })
if _, err := adapter.Submit(context.Background(), Request{Prompt: "private prompt"}); err == nil {
t.Fatal("expected Volcengine gateway error")
}
got := output.String()
for _, expected := range []string{"operation=submit", "status=400", `code="InvalidCredential"`, `codeN="100025"`, `requestId="request-nested-1"`, "errorClass=service"} {
if !strings.Contains(got, expected) {
t.Fatalf("log %q does not contain %q", got, expected)
}
}
for _, secret := range []string{"private-access-key", "private-secret-key", "Authorization contains", "private prompt", "Signature=", "Credential=", "request-header-ignored"} {
if strings.Contains(got, secret) {
t.Fatalf("log leaks %q: %s", secret, got)
}
}
}
func TestVolcengineServiceFailureFallsBackToSafeResponseHeaderRequestID(t *testing.T) {
var output bytes.Buffer
previousOutput, previousFlags := log.Writer(), log.Flags()
log.SetOutput(&output)
log.SetFlags(0)
t.Cleanup(func() {
log.SetOutput(previousOutput)
log.SetFlags(previousFlags)
})
client := roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusBadGateway,
Body: io.NopCloser(strings.NewReader(`{"ResponseMetadata":{"Error":{"Code":"InternalServiceError","CodeN":"100023"}}}`)),
Header: http.Header{"X-Tt-Logid": []string{"request-header-1"}},
}, nil
})
adapter := NewVolcengine(Config{
BaseURL: "https://visual.test", AccessKeyID: "private-access-key", SecretAccessKey: "private-secret-key", Model: seedream46Model,
}, client, func() time.Time { return time.Date(2026, 8, 19, 0, 0, 0, 0, time.UTC) })
if _, err := adapter.Submit(context.Background(), Request{Prompt: "private prompt"}); err == nil {
t.Fatal("expected Volcengine gateway error")
}
got := output.String()
for _, expected := range []string{`code="InternalServiceError"`, `codeN="100023"`, `requestId="request-header-1"`} {
if !strings.Contains(got, expected) {
t.Fatalf("log %q does not contain %q", got, expected)
}
}
}
func TestVolcengineTimeoutLogsSafeTransportClassification(t *testing.T) {
var output bytes.Buffer
previousOutput, previousFlags := log.Writer(), log.Flags()
log.SetOutput(&output)
log.SetFlags(0)
t.Cleanup(func() {
log.SetOutput(previousOutput)
log.SetFlags(previousFlags)
})
client := roundTripFunc(func(*http.Request) (*http.Response, error) {
return nil, context.DeadlineExceeded
})
adapter := NewVolcengine(Config{BaseURL: "https://visual.test", AccessKeyID: "private-access-key", SecretAccessKey: "private-secret-key", Model: "jimeng"}, client, nil)
if _, err := adapter.Submit(context.Background(), Request{Prompt: "private prompt"}); err == nil {
t.Fatal("expected Volcengine timeout")
}
got := output.String()
for _, expected := range []string{"operation=submit", "status=0", `code=""`, `requestId=""`, "errorClass=timeout", "elapsedMs="} {
if !strings.Contains(got, expected) {
t.Fatalf("log %q does not contain %q", got, expected)
}
}
for _, secret := range []string{"private-access-key", "private-secret-key", "private prompt", "deadline exceeded"} {
if strings.Contains(got, secret) {
t.Fatalf("log leaks %q: %s", secret, got)
}
}
}
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
@@ -264,7 +556,7 @@ func TestBailianDecodesCompatibleModeChoiceImages(t *testing.T) {
}
}
func TestSeedancePreservesTypedMultimodalMaterials(t *testing.T) {
func TestSeedance20PayloadMatchesOfficialMultimodalContract(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 {
@@ -281,13 +573,19 @@ func TestSeedancePreservesTypedMultimodalMaterials(t *testing.T) {
}
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] {
if item["type"] != check.materialType || item["role"] != check.role {
t.Fatalf("content[%d]=%#v", index+1, item)
}
if _, exists := item["label"]; exists {
t.Fatalf("content[%d] contains unsupported label: %#v", index+1, item)
}
if object := item[check.urlKey].(map[string]any); object["url"] == "" {
t.Fatalf("content[%d]=%#v", index+1, item)
}
}
if body["model"] != "doubao-seedance-2-0-260128" || body["generate_audio"] != true || body["watermark"] != false {
t.Fatalf("body=%#v", body)
}
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"id":"task-1","status":"queued"}`)), Header: http.Header{}}, nil
})
request := Request{
@@ -298,7 +596,7 @@ func TestSeedancePreservesTypedMultimodalMaterials(t *testing.T) {
{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 {
if _, err := NewSeedance(Config{BaseURL: "https://s.test/api/v3", APIKey: "secret", Model: "doubao-seedance-2-0-260128"}, client).Submit(context.Background(), request); err != nil {
t.Fatal(err)
}
}
+279 -13
View File
@@ -8,9 +8,12 @@ import (
"encoding/json"
"errors"
"io"
"log"
"net"
"net/http"
"net/url"
"sort"
"strconv"
"strings"
"time"
)
@@ -21,6 +24,20 @@ type Volcengine struct {
now func() time.Time
}
const (
seedream46Model = "jimeng_seedream46_cvtob"
seedream46SubmitAction = "JimengSeedream46CVToBSubmitTask"
seedream46QueryAction = "JimengSeedream46CVToBGetResult"
seedream46Version = "2024-06-06"
visualLegacyVersion = "2022-08-31"
)
type volcengineProtocol struct {
submitAction string
queryAction string
version string
}
func NewVolcengine(c Config, client HTTPClient, now func() time.Time) *Volcengine {
if client == nil {
client = http.DefaultClient
@@ -37,33 +54,57 @@ func NewVolcengine(c Config, client HTTPClient, now func() time.Time) *Volcengin
return &Volcengine{c, client, now}
}
func (v *Volcengine) Submit(ctx context.Context, r Request) (Result, error) {
p := map[string]any{"req_key": requestModel(r, v.config.Model), "prompt": r.Prompt, "image_urls": r.InputURLs}
model := requestModel(r, v.config.Model)
protocol := volcengineProtocolForModel(model)
p := map[string]any{"req_key": model, "prompt": r.Prompt}
if len(r.InputURLs) > 0 {
p["image_urls"] = r.InputURLs
}
for _, key := range []string{"scale", "width", "height", "min_ratio", "max_ratio", "force_single"} {
if value, exists := r.Settings[key]; exists && value != nil && value != "" {
p[key] = value
}
}
return v.call(ctx, "CVSync2AsyncSubmitTask", p)
return v.call(ctx, protocol.submitAction, protocol.version, p)
}
func (v *Volcengine) Query(ctx context.Context, id string) (Result, error) {
return v.QueryModel(ctx, id, v.config.Model)
}
func (v *Volcengine) QueryModel(ctx context.Context, id, model string) (Result, error) {
if strings.TrimSpace(id) == "" {
return Result{}, errors.New("provider task id is required")
}
resolvedModel := requestModel(Request{Model: model}, v.config.Model)
protocol := volcengineProtocolForModel(resolvedModel)
queryOptions, _ := json.Marshal(map[string]any{
"return_url": true,
"logo_info": map[string]any{"add_logo": false, "position": 0, "language": 0, "opacity": 1},
})
return v.call(ctx, "CVSync2AsyncGetResult", map[string]any{"req_key": requestModel(Request{Model: model}, v.config.Model), "task_id": id, "req_json": string(queryOptions)})
result, err := v.call(ctx, protocol.queryAction, protocol.version, map[string]any{"req_key": resolvedModel, "task_id": id, "req_json": string(queryOptions)})
if err != nil {
return Result{}, err
}
if result.TaskID == "" {
result.TaskID = id
}
return result, nil
}
func (v *Volcengine) call(ctx context.Context, action string, payload any) (Result, error) {
body, _ := json.Marshal(payload)
func (v *Volcengine) call(ctx context.Context, action, version string, payload any) (Result, error) {
startedAt := time.Now()
operation := volcengineOperation(action)
body, err := json.Marshal(payload)
if err != nil {
logVolcengineFailure(volcengineDiagnostic{Operation: operation, ErrorClass: "encode", ElapsedMS: elapsedMilliseconds(startedAt)})
return Result{}, &ProviderError{Operation: "volcengine request"}
}
endpoint, err := url.Parse(v.config.BaseURL)
if err != nil {
logVolcengineFailure(volcengineDiagnostic{Operation: operation, ErrorClass: "config", ElapsedMS: elapsedMilliseconds(startedAt)})
return Result{}, errors.New("invalid provider base URL")
}
q := endpoint.Query()
q.Set("Action", action)
q.Set("Version", "2022-08-31")
q.Set("Version", version)
endpoint.RawQuery = canonicalQuery(q)
date := v.now().UTC()
xdate := date.Format("20060102T150405Z")
@@ -83,26 +124,251 @@ func (v *Volcengine) call(ctx context.Context, action string, payload any) (Resu
key = hmacBytes(key, v.config.Service)
key = hmacBytes(key, "request")
signature := hex.EncodeToString(hmacBytes(key, stringToSign))
req, _ := http.NewRequestWithContext(ctx, http.MethodPost, endpoint.String(), strings.NewReader(string(body)))
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint.String(), strings.NewReader(string(body)))
if err != nil {
logVolcengineFailure(volcengineDiagnostic{Operation: operation, ErrorClass: "config", ElapsedMS: elapsedMilliseconds(startedAt)})
return Result{}, &ProviderError{Operation: "volcengine request"}
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Content-Sha256", hash)
req.Header.Set("X-Date", xdate)
req.Header.Set("Authorization", "HMAC-SHA256 Credential="+v.config.AccessKeyID+"/"+scope+", SignedHeaders="+signed+", Signature="+signature)
resp, err := v.client.Do(req)
if err != nil {
logVolcengineFailure(volcengineDiagnostic{Operation: operation, ErrorClass: classifyVolcengineTransportError(err), ElapsedMS: elapsedMilliseconds(startedAt)})
return Result{}, &ProviderError{Operation: "volcengine request"}
}
defer resp.Body.Close()
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
if resp.StatusCode < 200 || resp.StatusCode >= 300 || !json.Valid(raw) {
limit := v.config.MaxResponseBytes
if limit <= 0 {
limit = 2 << 20
}
raw, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
if err != nil {
logVolcengineFailure(volcengineDiagnostic{Operation: operation, Status: resp.StatusCode, ErrorClass: "response_read", ElapsedMS: elapsedMilliseconds(startedAt)})
return Result{}, &ProviderError{Operation: "volcengine request", Status: resp.StatusCode}
}
r := record(raw)
d := object(r["data"])
if int64(len(raw)) > limit {
logVolcengineFailure(volcengineDiagnostic{Operation: operation, Status: resp.StatusCode, ErrorClass: "response_too_large", ElapsedMS: elapsedMilliseconds(startedAt)})
return Result{}, &ProviderError{Operation: "volcengine request", Status: resp.StatusCode}
}
validJSON := json.Valid(raw)
r := map[string]any{}
if validJSON {
r = record(raw)
}
response := inspectVolcengineResponse(r, resp.Header)
diagnostic := volcengineDiagnostic{
Operation: operation,
Status: resp.StatusCode,
Code: volcengineDiagnosticCode(response.code),
CodeN: volcengineDiagnosticNumericCode(response.codeN),
RequestID: safeVolcengineRequestID(response.requestID),
ElapsedMS: elapsedMilliseconds(startedAt),
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
diagnostic.ErrorClass = "service"
logVolcengineFailure(diagnostic)
return Result{}, &ProviderError{Operation: "volcengine request", Status: resp.StatusCode}
}
if !validJSON {
diagnostic.ErrorClass = "invalid_json"
logVolcengineFailure(diagnostic)
return Result{}, &ProviderError{Operation: "volcengine request"}
}
if response.code != nil && !volcengineRequestSucceeded(response.code) {
diagnostic.ErrorClass = "service"
logVolcengineFailure(diagnostic)
return Result{}, &ProviderError{Operation: "volcengine request"}
}
d := object(first(response.business["data"], response.business["Data"]))
out := []string{}
collectURLs(d, &out)
return Result{TaskID: stringValue(r["task_id"], d["task_id"]), Status: status(first(r["status"], d["status"])), OutputURLs: out, Raw: raw}, nil
for _, value := range []any{d["image_urls"], d["image_url"], d["url"], d["result_url"], d["output"], d["outputs"]} {
collectURLs(value, &out)
}
return Result{TaskID: stringValue(response.business["task_id"], response.business["TaskId"], d["task_id"], d["TaskId"]), Status: status(first(d["status"], d["Status"], response.business["status"], response.business["Status"])), OutputURLs: out, Raw: raw}, nil
}
type volcengineResponse struct {
business map[string]any
code any
codeN any
requestID any
}
func inspectVolcengineResponse(root map[string]any, headers http.Header) volcengineResponse {
business := root
if wrapped := object(first(root["Result"], root["result"])); len(wrapped) > 0 {
business = wrapped
}
metadata := object(first(root["ResponseMetadata"], root["response_metadata"]))
gatewayError := object(first(metadata["Error"], metadata["error"]))
return volcengineResponse{
business: business,
code: first(
business["code"], business["Code"],
root["code"], root["Code"],
gatewayError["Code"], gatewayError["code"],
),
codeN: first(
gatewayError["CodeN"], gatewayError["codeN"], gatewayError["code_n"],
business["code_n"], business["codeN"], business["CodeN"],
root["code_n"], root["codeN"], root["CodeN"],
),
requestID: first(
business["request_id"], business["requestId"], business["RequestId"], business["RequestID"],
root["request_id"], root["requestId"], root["RequestId"], root["RequestID"],
metadata["RequestId"], metadata["RequestID"], metadata["request_id"], metadata["requestId"],
headers.Get("X-Tt-Logid"), headers.Get("X-Request-Id"),
),
}
}
func volcengineProtocolForModel(model string) volcengineProtocol {
if strings.TrimSpace(model) == seedream46Model {
return volcengineProtocol{
submitAction: seedream46SubmitAction,
queryAction: seedream46QueryAction,
version: seedream46Version,
}
}
return volcengineProtocol{
submitAction: "CVSync2AsyncSubmitTask",
queryAction: "CVSync2AsyncGetResult",
version: visualLegacyVersion,
}
}
func volcengineRequestSucceeded(value any) bool {
return volcengineDiagnosticCode(value) == "10000"
}
type volcengineDiagnostic struct {
Operation string
Status int
Code string
CodeN string
RequestID string
ErrorClass string
ElapsedMS int64
}
func logVolcengineFailure(diagnostic volcengineDiagnostic) {
log.Printf(
"zhinian-api Volcengine operation failed operation=%s status=%d code=%q codeN=%q requestId=%q errorClass=%s elapsedMs=%d",
diagnostic.Operation,
diagnostic.Status,
diagnostic.Code,
diagnostic.CodeN,
diagnostic.RequestID,
diagnostic.ErrorClass,
diagnostic.ElapsedMS,
)
}
func volcengineOperation(action string) string {
switch action {
case "CVSync2AsyncSubmitTask", seedream46SubmitAction:
return "submit"
case "CVSync2AsyncGetResult", seedream46QueryAction:
return "query"
default:
return "unknown"
}
}
func volcengineDiagnosticCode(value any) string {
code := volcengineDiagnosticValue(value)
if len(code) == 0 || len(code) > 64 {
return ""
}
for _, character := range code {
if (character >= 'a' && character <= 'z') || (character >= 'A' && character <= 'Z') || (character >= '0' && character <= '9') || strings.ContainsRune("-_.:", character) {
continue
}
return ""
}
return code
}
func volcengineDiagnosticNumericCode(value any) string {
code := volcengineDiagnosticValue(value)
if len(code) == 0 || len(code) > 32 {
return ""
}
for _, character := range code {
if character < '0' || character > '9' {
return ""
}
}
return code
}
func volcengineDiagnosticValue(value any) string {
switch typed := value.(type) {
case float64:
return strconv.FormatFloat(typed, 'f', -1, 64)
case float32:
return strconv.FormatFloat(float64(typed), 'f', -1, 32)
case int:
return strconv.Itoa(typed)
case int32:
return strconv.FormatInt(int64(typed), 10)
case int64:
return strconv.FormatInt(typed, 10)
case json.Number:
return string(typed)
case string:
return strings.TrimSpace(typed)
default:
return ""
}
}
func safeVolcengineRequestID(value any) string {
requestID := stringValue(value)
if len(requestID) == 0 || len(requestID) > 128 {
return ""
}
for _, character := range requestID {
if (character >= 'a' && character <= 'z') || (character >= 'A' && character <= 'Z') || (character >= '0' && character <= '9') || strings.ContainsRune("-_.:", character) {
continue
}
return ""
}
return requestID
}
func classifyVolcengineTransportError(err error) string {
if errors.Is(err, context.Canceled) {
return "canceled"
}
if errors.Is(err, context.DeadlineExceeded) {
return "timeout"
}
var networkError net.Error
if errors.As(err, &networkError) && networkError.Timeout() {
return "timeout"
}
var dnsError *net.DNSError
if errors.As(err, &dnsError) {
return "dns"
}
var operationError *net.OpError
if errors.As(err, &operationError) && operationError.Op == "dial" {
return "connect"
}
return "transport"
}
func elapsedMilliseconds(startedAt time.Time) int64 {
elapsed := time.Since(startedAt).Milliseconds()
if elapsed < 0 {
return 0
}
return elapsed
}
func sha(b []byte) string { x := sha256.Sum256(b); return hex.EncodeToString(x[:]) }
func hmacBytes(k []byte, s string) []byte {
h := hmac.New(sha256.New, k)