增加即梦的生图和生视频功能
This commit is contained in:
1 parent
a9a4cdf125
commit
542246e4ce
20 files changed
+1008
-53
No files matched your search
@@ -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
|
||||
|
||||
@@ -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 != "" {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in new issue
Block a user