feat: complete remaining Go backend modules
This commit is contained in:
1 parent
cea2751dc5
commit
aef5a97165
145 files changed
+18376
-199
No files matched your search
@@ -0,0 +1,223 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type EvoLink struct{ *httpAdapter }
|
||||
|
||||
func NewEvoLink(c Config, client HTTPClient) *EvoLink {
|
||||
if client == nil {
|
||||
client = http.DefaultClient
|
||||
}
|
||||
a := &httpAdapter{name: "evolink", config: c, client: client, submitPath: func(Request) string { return "/v1/images/generations" }, queryPath: func(id string) string { return "/v1/tasks/" + id }, payload: func(r Request) any {
|
||||
payload := map[string]any{"model": requestModel(r, c.Model), "prompt": r.Prompt, "n": 1, "resolution": "1K"}
|
||||
if len(r.InputURLs) > 0 {
|
||||
payload["image_urls"] = r.InputURLs
|
||||
}
|
||||
if quality := strings.TrimSpace(stringValue(r.Settings["quality"])); quality != "" {
|
||||
payload["quality"] = quality
|
||||
}
|
||||
if size := strings.TrimSpace(stringValue(r.Settings["size"])); size != "" {
|
||||
payload["size"] = size
|
||||
} else if width, widthOK := positiveInteger(r.Settings["width"]); widthOK {
|
||||
if height, heightOK := positiveInteger(r.Settings["height"]); heightOK {
|
||||
payload["size"] = supportedRatio(width, height)
|
||||
}
|
||||
}
|
||||
return payload
|
||||
}, decode: decodeEvoLink}
|
||||
return &EvoLink{a}
|
||||
}
|
||||
func (a *EvoLink) Submit(c context.Context, r Request) (Result, error) { return a.submit(c, r) }
|
||||
func (a *EvoLink) Query(c context.Context, id string) (Result, error) { return a.query(c, id) }
|
||||
func decodeEvoLink(raw []byte) Result {
|
||||
r := record(raw)
|
||||
d := object(r["data"])
|
||||
out := []string{}
|
||||
for _, v := range []any{r["results"], d["results"], d["images"], d["image_urls"], d["output"], d["outputs"]} {
|
||||
collectURLs(v, &out)
|
||||
}
|
||||
return Result{TaskID: stringValue(r["id"], r["task_id"], d["id"], d["task_id"]), Status: status(first(r["status"], d["status"])), OutputURLs: out}
|
||||
}
|
||||
|
||||
type Bailian struct{ *httpAdapter }
|
||||
|
||||
func NewBailian(c Config, client HTTPClient) *Bailian {
|
||||
if client == nil {
|
||||
client = http.DefaultClient
|
||||
}
|
||||
a := &httpAdapter{name: "bailian", config: c, client: client, submitPath: func(r Request) string {
|
||||
if r.Capability == "video.generate" {
|
||||
return "/api/v1/services/aigc/video-generation/video-synthesis"
|
||||
}
|
||||
return "/api/v1/services/aigc/image-generation/generation"
|
||||
}, queryPath: func(id string) string { return "/api/v1/tasks/" + id }, payload: func(r Request) any {
|
||||
if r.Capability == "video.generate" {
|
||||
media := make([]any, 0, len(r.InputURLs))
|
||||
for index, inputURL := range r.InputURLs {
|
||||
frameType := "first_frame"
|
||||
if index > 0 {
|
||||
frameType = "last_frame"
|
||||
}
|
||||
media = append(media, map[string]any{"type": frameType, "url": inputURL})
|
||||
}
|
||||
parameters := map[string]any{
|
||||
"resolution": strings.ToUpper(stringValue(r.Settings["resolution"])),
|
||||
"duration": r.Settings["duration"],
|
||||
"prompt_extend": true,
|
||||
"watermark": false,
|
||||
}
|
||||
if parameters["resolution"] == "" {
|
||||
parameters["resolution"] = "720P"
|
||||
}
|
||||
if parameters["duration"] == nil {
|
||||
parameters["duration"] = 10
|
||||
}
|
||||
return map[string]any{"model": requestModel(r, c.Model), "input": map[string]any{"prompt": r.Prompt, "media": media}, "parameters": parameters}
|
||||
}
|
||||
content := make([]any, 0, len(r.InputURLs)+1)
|
||||
for _, inputURL := range r.InputURLs {
|
||||
content = append(content, map[string]any{"image": inputURL})
|
||||
}
|
||||
content = append(content, map[string]any{"text": r.Prompt})
|
||||
parameters := map[string]any{"size": "2K", "n": 1, "watermark": false}
|
||||
if width, widthOK := positiveInteger(r.Settings["width"]); widthOK {
|
||||
if height, heightOK := positiveInteger(r.Settings["height"]); heightOK {
|
||||
parameters["size"] = fmt.Sprintf("%d*%d", width, height)
|
||||
}
|
||||
}
|
||||
if len(r.InputURLs) == 0 {
|
||||
parameters["thinking_mode"] = true
|
||||
}
|
||||
return map[string]any{"model": requestModel(r, c.Model), "input": map[string]any{"messages": []any{map[string]any{"role": "user", "content": content}}}, "parameters": parameters}
|
||||
}, headers: func(r *http.Request) { r.Header.Set("X-DashScope-Async", "enable") }, decode: decodeBailian}
|
||||
return &Bailian{a}
|
||||
}
|
||||
func (a *Bailian) Submit(c context.Context, r Request) (Result, error) { return a.submit(c, r) }
|
||||
func (a *Bailian) Query(c context.Context, id string) (Result, error) { return a.query(c, id) }
|
||||
func decodeBailian(raw []byte) Result {
|
||||
r := record(raw)
|
||||
o := object(r["output"])
|
||||
out := []string{}
|
||||
collectURLs(o["results"], &out)
|
||||
if choices, ok := o["choices"].([]any); ok {
|
||||
for _, choice := range choices {
|
||||
message := object(object(choice)["message"])
|
||||
if content, ok := message["content"].([]any); ok {
|
||||
for _, item := range content {
|
||||
collectURLs(object(item)["image"], &out)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
collectURLs(o["video_url"], &out)
|
||||
return Result{TaskID: stringValue(o["task_id"], r["task_id"]), Status: status(o["task_status"]), OutputURLs: out, ErrorMessage: stringValue(r["message"])}
|
||||
}
|
||||
|
||||
type Seedance struct{ *httpAdapter }
|
||||
|
||||
func NewSeedance(c Config, client HTTPClient) *Seedance {
|
||||
if client == nil {
|
||||
client = http.DefaultClient
|
||||
}
|
||||
a := &httpAdapter{name: "seedance", config: c, client: client, submitPath: func(Request) string { return "/contents/generations/tasks" }, queryPath: func(id string) string { return "/contents/generations/tasks/" + id }, payload: func(r Request) any {
|
||||
content := []any{map[string]any{"type": "text", "text": r.Prompt}}
|
||||
materials := r.Materials
|
||||
if len(materials) == 0 {
|
||||
materials = make([]Material, 0, len(r.InputURLs))
|
||||
for _, inputURL := range r.InputURLs {
|
||||
materials = append(materials, Material{URL: inputURL, Type: MaterialImage})
|
||||
}
|
||||
}
|
||||
for _, material := range materials {
|
||||
materialType, urlKey, role := "image_url", "image_url", "reference_image"
|
||||
switch material.Type {
|
||||
case MaterialVideo:
|
||||
materialType, urlKey, role = "video_url", "video_url", "reference_video"
|
||||
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)
|
||||
}
|
||||
p := map[string]any{"model": requestModel(r, c.Model), "content": content, "generate_audio": true, "watermark": false}
|
||||
for k, v := range r.Settings {
|
||||
p[k] = v
|
||||
}
|
||||
return p
|
||||
}, decode: decodeSeedance}
|
||||
return &Seedance{a}
|
||||
}
|
||||
func (a *Seedance) Submit(c context.Context, r Request) (Result, error) { return a.submit(c, r) }
|
||||
func (a *Seedance) Query(c context.Context, id string) (Result, error) { return a.query(c, id) }
|
||||
func decodeSeedance(raw []byte) Result {
|
||||
r := record(raw)
|
||||
d := object(r["data"])
|
||||
content := object(r["content"])
|
||||
if len(content) == 0 {
|
||||
content = object(d["content"])
|
||||
}
|
||||
out := []string{}
|
||||
for _, v := range []any{content, r["video_url"], r["url"], r["output"], d} {
|
||||
collectURLs(v, &out)
|
||||
}
|
||||
usage := object(r["usage"])
|
||||
if len(usage) == 0 {
|
||||
usage = object(d["usage"])
|
||||
}
|
||||
u := map[string]int{}
|
||||
if n, ok := usage["completion_tokens"].(float64); ok && n > 0 {
|
||||
u["completionTokens"] = int(n)
|
||||
}
|
||||
return Result{TaskID: stringValue(r["id"], r["task_id"], d["id"], d["task_id"]), Status: status(first(r["status"], d["status"])), OutputURLs: out, ErrorMessage: stringValue(object(r["error"])["message"], object(d["error"])["message"]), Usage: u}
|
||||
}
|
||||
|
||||
func first(values ...any) any {
|
||||
for _, v := range values {
|
||||
if v != nil {
|
||||
return v
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func positiveInteger(value any) (int, bool) {
|
||||
switch number := value.(type) {
|
||||
case int:
|
||||
return number, number > 0
|
||||
case int64:
|
||||
return int(number), number > 0
|
||||
case float64:
|
||||
integer := int(number)
|
||||
return integer, number > 0 && float64(integer) == number
|
||||
default:
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
|
||||
func supportedRatio(width, height int) string {
|
||||
divisor := greatestCommonDivisor(width, height)
|
||||
ratio := fmt.Sprintf("%d:%d", width/divisor, height/divisor)
|
||||
switch ratio {
|
||||
case "1:1", "1:2", "2:1", "1:3", "3:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9", "9:21", "21:9":
|
||||
return ratio
|
||||
default:
|
||||
return fmt.Sprintf("%dx%d", width, height)
|
||||
}
|
||||
}
|
||||
|
||||
func greatestCommonDivisor(left, right int) int {
|
||||
for right != 0 {
|
||||
left, right = right, left%right
|
||||
}
|
||||
return left
|
||||
}
|
||||
|
||||
var _ = fmt.Sprint
|
||||
@@ -0,0 +1,18 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
)
|
||||
|
||||
type Mock struct{ seed string }
|
||||
|
||||
func NewMock(seed string) *Mock { return &Mock{seed: seed} }
|
||||
func (m *Mock) Submit(_ context.Context, r Request) (Result, error) {
|
||||
sum := sha256.Sum256([]byte(m.seed + "\x00" + r.Capability + "\x00" + r.Prompt))
|
||||
return Result{TaskID: "mock-" + hex.EncodeToString(sum[:8]), Status: StatusQueued}, nil
|
||||
}
|
||||
func (m *Mock) Query(_ context.Context, id string) (Result, error) {
|
||||
return Result{TaskID: id, Status: StatusSucceeded, OutputURLs: []string{"/generated-results/" + id}}, nil
|
||||
}
|
||||
@@ -0,0 +1,230 @@
|
||||
// Package providers contains bounded protocol adapters for generation services.
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type Status string
|
||||
|
||||
const (
|
||||
StatusQueued Status = "queued"
|
||||
StatusRunning Status = "running"
|
||||
StatusSucceeded Status = "succeeded"
|
||||
StatusFailed Status = "failed"
|
||||
StatusCancelled Status = "cancelled"
|
||||
StatusExpired Status = "expired"
|
||||
)
|
||||
|
||||
type Request struct {
|
||||
Capability string `json:"capability"`
|
||||
Model string `json:"model,omitempty"`
|
||||
Prompt string `json:"prompt"`
|
||||
InputURLs []string `json:"inputUrls,omitempty"`
|
||||
Materials []Material `json:"materials,omitempty"`
|
||||
Settings map[string]any `json:"settings,omitempty"`
|
||||
}
|
||||
|
||||
type MaterialType string
|
||||
|
||||
const (
|
||||
MaterialImage MaterialType = "image"
|
||||
MaterialVideo MaterialType = "video"
|
||||
MaterialAudio MaterialType = "audio"
|
||||
)
|
||||
|
||||
// Material retains the provider-facing type metadata that InputURLs cannot
|
||||
// express. InputURLs remains supported for existing callers.
|
||||
type Material struct {
|
||||
URL string `json:"url"`
|
||||
Type MaterialType `json:"type"`
|
||||
Role string `json:"role,omitempty"`
|
||||
Label string `json:"label,omitempty"`
|
||||
}
|
||||
|
||||
func requestModel(request Request, fallback string) string {
|
||||
if value := strings.TrimSpace(request.Model); value != "" {
|
||||
return value
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
type Result struct {
|
||||
TaskID string
|
||||
Status Status
|
||||
OutputURLs []string
|
||||
Raw json.RawMessage
|
||||
ErrorMessage string
|
||||
Usage map[string]int
|
||||
}
|
||||
|
||||
// HTTPResult is the provider-neutral representation persisted with a Job.
|
||||
// Output URLs and usage must survive a process restart after the external
|
||||
// provider has already reached a terminal state.
|
||||
type HTTPResult struct {
|
||||
TaskID string `json:"taskId,omitempty"`
|
||||
Status Status `json:"status"`
|
||||
OutputURLs []string `json:"outputUrls"`
|
||||
Raw json.RawMessage `json:"raw,omitempty"`
|
||||
ErrorMessage string `json:"errorMessage,omitempty"`
|
||||
Usage map[string]int `json:"usage,omitempty"`
|
||||
}
|
||||
|
||||
func EncodeResult(result Result) (json.RawMessage, error) {
|
||||
urls := result.OutputURLs
|
||||
if urls == nil {
|
||||
urls = []string{}
|
||||
}
|
||||
return json.Marshal(HTTPResult{TaskID: result.TaskID, Status: result.Status, OutputURLs: urls, Raw: result.Raw, ErrorMessage: result.ErrorMessage, Usage: result.Usage})
|
||||
}
|
||||
|
||||
type Adapter interface {
|
||||
Submit(context.Context, Request) (Result, error)
|
||||
Query(context.Context, string) (Result, error)
|
||||
}
|
||||
|
||||
// ModelQueryAdapter is implemented by providers whose query protocol requires
|
||||
// the same model identifier that was used when the task was submitted. Callers
|
||||
// can opt into it without widening the common Adapter contract.
|
||||
type ModelQueryAdapter interface {
|
||||
QueryModel(context.Context, string, string) (Result, error)
|
||||
}
|
||||
type HTTPClient interface {
|
||||
Do(*http.Request) (*http.Response, error)
|
||||
}
|
||||
type Config struct {
|
||||
BaseURL, APIKey, Model, AccessKeyID, SecretAccessKey, Region, Service string
|
||||
MaxResponseBytes int64
|
||||
}
|
||||
|
||||
type ProviderError struct {
|
||||
Operation string
|
||||
Status int
|
||||
}
|
||||
|
||||
func (e *ProviderError) Error() string {
|
||||
if e.Status > 0 {
|
||||
return fmt.Sprintf("provider %s failed with HTTP %d", e.Operation, e.Status)
|
||||
}
|
||||
return "provider " + e.Operation + " failed"
|
||||
}
|
||||
|
||||
type httpAdapter struct {
|
||||
name string
|
||||
config Config
|
||||
client HTTPClient
|
||||
submitPath func(Request) string
|
||||
queryPath func(string) string
|
||||
payload func(Request) any
|
||||
headers func(*http.Request)
|
||||
decode func([]byte) Result
|
||||
}
|
||||
|
||||
func (a *httpAdapter) submit(ctx context.Context, input Request) (Result, error) {
|
||||
body, err := json.Marshal(a.payload(input))
|
||||
if err != nil {
|
||||
return Result{}, fmt.Errorf("encode provider request: %w", err)
|
||||
}
|
||||
return a.call(ctx, http.MethodPost, a.submitPath(input), body, "submit")
|
||||
}
|
||||
func (a *httpAdapter) query(ctx context.Context, id string) (Result, error) {
|
||||
if strings.TrimSpace(id) == "" {
|
||||
return Result{}, errors.New("provider task id is required")
|
||||
}
|
||||
return a.call(ctx, http.MethodGet, a.queryPath(url.PathEscape(id)), nil, "query")
|
||||
}
|
||||
func (a *httpAdapter) call(ctx context.Context, method, path string, body []byte, operation string) (Result, error) {
|
||||
base, err := url.Parse(strings.TrimRight(a.config.BaseURL, "/"))
|
||||
if err != nil {
|
||||
return Result{}, errors.New("invalid provider base URL")
|
||||
}
|
||||
rel, err := url.Parse(path)
|
||||
if err != nil {
|
||||
return Result{}, errors.New("invalid provider path")
|
||||
}
|
||||
base.Path = strings.TrimRight(base.Path, "/") + "/"
|
||||
rel.Path = strings.TrimLeft(rel.Path, "/")
|
||||
target := base.ResolveReference(rel)
|
||||
req, err := http.NewRequestWithContext(ctx, method, target.String(), strings.NewReader(string(body)))
|
||||
if err != nil {
|
||||
return Result{}, fmt.Errorf("build provider request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+a.config.APIKey)
|
||||
if a.headers != nil {
|
||||
a.headers(req)
|
||||
}
|
||||
resp, err := a.client.Do(req)
|
||||
if err != nil {
|
||||
return Result{}, &ProviderError{Operation: a.name + " " + operation}
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
limit := a.config.MaxResponseBytes
|
||||
if limit <= 0 {
|
||||
limit = 2 << 20
|
||||
}
|
||||
raw, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
|
||||
if err != nil || int64(len(raw)) > limit {
|
||||
return Result{}, &ProviderError{Operation: a.name + " " + operation}
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return Result{}, &ProviderError{Operation: a.name + " " + operation, Status: resp.StatusCode}
|
||||
}
|
||||
if !json.Valid(raw) {
|
||||
return Result{}, &ProviderError{Operation: a.name + " " + operation}
|
||||
}
|
||||
result := a.decode(raw)
|
||||
result.Raw = append(json.RawMessage(nil), raw...)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func record(raw []byte) map[string]any { var v map[string]any; _ = json.Unmarshal(raw, &v); return v }
|
||||
func object(v any) map[string]any { x, _ := v.(map[string]any); return x }
|
||||
func stringValue(values ...any) string {
|
||||
for _, v := range values {
|
||||
if s, ok := v.(string); ok && strings.TrimSpace(s) != "" {
|
||||
return strings.TrimSpace(s)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
func status(v any) Status {
|
||||
s := strings.ToLower(stringValue(v))
|
||||
switch s {
|
||||
case "completed", "complete", "succeeded", "success", "done":
|
||||
return StatusSucceeded
|
||||
case "running", "processing", "generating", "in_progress":
|
||||
return StatusRunning
|
||||
case "failed", "error", "unknown":
|
||||
return StatusFailed
|
||||
case "cancelled", "canceled":
|
||||
return StatusCancelled
|
||||
case "expired", "not_found", "timeout":
|
||||
return StatusExpired
|
||||
default:
|
||||
return StatusQueued
|
||||
}
|
||||
}
|
||||
func collectURLs(v any, out *[]string) {
|
||||
switch x := v.(type) {
|
||||
case string:
|
||||
if strings.HasPrefix(x, "http://") || strings.HasPrefix(x, "https://") {
|
||||
*out = append(*out, x)
|
||||
}
|
||||
case []any:
|
||||
for _, i := range x {
|
||||
collectURLs(i, out)
|
||||
}
|
||||
case map[string]any:
|
||||
for _, k := range []string{"url", "image_url", "imageUrl", "result_url", "resultUrl", "video_url", "file_url"} {
|
||||
collectURLs(x[k], out)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,306 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
type roundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) Do(r *http.Request) (*http.Response, error) { return f(r) }
|
||||
|
||||
func TestHTTPAdaptersMapRequestsAndResponses(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
adapter Adapter
|
||||
wantSubmitPath, wantQueryPath, response string
|
||||
}{
|
||||
{"evolink", NewEvoLink(Config{BaseURL: "https://e.test", APIKey: "secret", Model: "gpt-image-2"}, nil), "/v1/images/generations", "/v1/tasks/task-1", `{"id":"task-1","status":"completed","results":[{"url":"https://cdn.test/a.png"}]}`},
|
||||
{"bailian", NewBailian(Config{BaseURL: "https://b.test", APIKey: "secret", Model: "wan"}, nil), "/api/v1/services/aigc/image-generation/generation", "/api/v1/tasks/task-1", `{"output":{"task_id":"task-1","task_status":"SUCCEEDED","results":[{"url":"https://cdn.test/a.png"}]}}`},
|
||||
{"seedance", NewSeedance(Config{BaseURL: "https://s.test/api/v3", APIKey: "secret", Model: "seed"}, nil), "/api/v3/contents/generations/tasks", "/api/v3/contents/generations/tasks/task-1", `{"id":"task-1","status":"succeeded","content":{"video_url":"https://cdn.test/a.mp4"}}`},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var calls int
|
||||
client := roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
||||
calls++
|
||||
want := tt.wantSubmitPath
|
||||
if calls == 2 {
|
||||
want = tt.wantQueryPath
|
||||
}
|
||||
if r.URL.Path != want {
|
||||
t.Fatalf("path=%s want=%s", r.URL.Path, want)
|
||||
}
|
||||
if !strings.HasPrefix(r.Header.Get("Authorization"), "Bearer ") {
|
||||
t.Fatal("missing bearer")
|
||||
}
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(tt.response)), Header: http.Header{}}, nil
|
||||
})
|
||||
switch a := tt.adapter.(type) {
|
||||
case *EvoLink:
|
||||
a.client = client
|
||||
case *Bailian:
|
||||
a.client = client
|
||||
case *Seedance:
|
||||
a.client = client
|
||||
}
|
||||
submitted, err := tt.adapter.Submit(context.Background(), Request{Capability: "image.generate", Prompt: "hello"})
|
||||
if err != nil || submitted.TaskID != "task-1" {
|
||||
t.Fatalf("Submit=%#v,%v", submitted, err)
|
||||
}
|
||||
queried, err := tt.adapter.Query(context.Background(), "task-1")
|
||||
if err != nil || queried.Status != "succeeded" || len(queried.OutputURLs) != 1 {
|
||||
t.Fatalf("Query=%#v,%v", queried, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestVolcengineSignsSubmitAndMapsResult(t *testing.T) {
|
||||
var request *http.Request
|
||||
client := roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
||||
request = r
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"data":{"task_id":"task-1","status":"done","image_urls":["https://cdn.test/a.png"]}}`)), Header: http.Header{}}, nil
|
||||
})
|
||||
a := NewVolcengine(Config{BaseURL: "https://visual.test", AccessKeyID: "ak", SecretAccessKey: "sk", Region: "cn-north-1", Service: "cv", Model: "jimeng"}, client, func() time.Time { return time.Date(2026, 8, 13, 0, 0, 0, 0, time.UTC) })
|
||||
got, err := a.Submit(context.Background(), Request{Capability: "image.generate", Prompt: "hello"})
|
||||
if err != nil || got.TaskID != "task-1" {
|
||||
t.Fatalf("Submit=%#v,%v", got, err)
|
||||
}
|
||||
if request.URL.Query().Get("Action") != "CVSync2AsyncSubmitTask" || !strings.Contains(request.Header.Get("Authorization"), "Credential=ak/") || request.Header.Get("X-Date") != "20260813T000000Z" {
|
||||
t.Fatalf("request=%#v headers=%#v", request.URL, request.Header)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVolcenginePayloadsMatchJimengSubmitAndQueryProtocols(t *testing.T) {
|
||||
var bodies []map[string]any
|
||||
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
var body map[string]any
|
||||
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bodies = append(bodies, body)
|
||||
response := `{"data":{"task_id":"task-1","status":"queued"}}`
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(response)), Header: http.Header{}}, nil
|
||||
})
|
||||
adapter := NewVolcengine(Config{BaseURL: "https://visual.test", AccessKeyID: "ak", SecretAccessKey: "sk", Model: "jimeng"}, client, func() time.Time { return time.Date(2026, 8, 13, 0, 0, 0, 0, time.UTC) })
|
||||
if _, err := adapter.Submit(context.Background(), Request{Prompt: "draw", InputURLs: []string{"https://cdn.test/ref.png"}, Settings: map[string]any{"width": 1024, "height": 768, "force_single": true, "ignored": "value"}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := adapter.Query(context.Background(), "task-1"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if bodies[0]["width"] != float64(1024) || bodies[0]["height"] != float64(768) || bodies[0]["force_single"] != true || bodies[0]["ignored"] != nil {
|
||||
t.Fatalf("submit body=%#v", bodies[0])
|
||||
}
|
||||
queryJSON, ok := bodies[1]["req_json"].(string)
|
||||
if !ok || queryJSON == "" {
|
||||
t.Fatalf("query body=%#v", bodies[1])
|
||||
}
|
||||
var queryOptions map[string]any
|
||||
if err := json.Unmarshal([]byte(queryJSON), &queryOptions); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
logo := queryOptions["logo_info"].(map[string]any)
|
||||
if queryOptions["return_url"] != true || logo["add_logo"] != false || logo["opacity"] != float64(1) {
|
||||
t.Fatalf("query options=%#v", queryOptions)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVolcengineQueryUsesPersistedTaskModelInsteadOfAdapterFallback(t *testing.T) {
|
||||
var body map[string]any
|
||||
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"data":{"task_id":"task-1","status":"queued"}}`)), Header: http.Header{}}, nil
|
||||
})
|
||||
adapter := NewVolcengine(Config{BaseURL: "https://visual.test", AccessKeyID: "ak", SecretAccessKey: "sk", Model: "fallback-model-b"}, client, func() time.Time {
|
||||
return time.Date(2026, 8, 13, 0, 0, 0, 0, time.UTC)
|
||||
})
|
||||
|
||||
if _, err := adapter.QueryModel(context.Background(), "task-1", "persisted-model-a"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if body["req_key"] != "persisted-model-a" || body["task_id"] != "task-1" {
|
||||
t.Fatalf("query body=%#v", body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderErrorsAreGenericAndDoNotLeakSecrets(t *testing.T) {
|
||||
client := roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: 401, Body: io.NopCloser(strings.NewReader(`{"message":"secret upstream detail"}`)), Header: http.Header{}}, nil
|
||||
})
|
||||
a := NewEvoLink(Config{BaseURL: "https://e.test", APIKey: "very-secret", Model: "m"}, client)
|
||||
_, err := a.Submit(context.Background(), Request{Capability: "image.generate", Prompt: "p"})
|
||||
if err == nil || strings.Contains(err.Error(), "secret") {
|
||||
t.Fatalf("error=%v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBailianUsesThePreparedRequestModelForImageAndVideo(t *testing.T) {
|
||||
models := []string{}
|
||||
client := roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
||||
var body map[string]any
|
||||
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
models = append(models, body["model"].(string))
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"output":{"task_id":"task-1","task_status":"PENDING"}}`)), Header: http.Header{}}, nil
|
||||
})
|
||||
adapter := NewBailian(Config{BaseURL: "https://b.test", APIKey: "secret", Model: "fallback-model"}, client)
|
||||
if _, err := adapter.Submit(context.Background(), Request{Capability: "image.generate", Model: "image-model", Prompt: "image"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := adapter.Submit(context.Background(), Request{Capability: "video.generate", Model: "video-model", Prompt: "video"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(models) != 2 || models[0] != "image-model" || models[1] != "video-model" {
|
||||
t.Fatalf("models=%#v", models)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBailianPayloadsMatchImageAndVideoProtocols(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
request Request
|
||||
check func(*testing.T, map[string]any)
|
||||
}{
|
||||
{
|
||||
name: "image messages and parameters",
|
||||
request: Request{Capability: "image.generate", Model: "wan-image", Prompt: "draw", InputURLs: []string{"https://cdn.test/ref.png"}, Settings: map[string]any{
|
||||
"width": 1024, "height": 768,
|
||||
}},
|
||||
check: func(t *testing.T, body map[string]any) {
|
||||
input := body["input"].(map[string]any)
|
||||
messages := input["messages"].([]any)
|
||||
content := messages[0].(map[string]any)["content"].([]any)
|
||||
parameters := body["parameters"].(map[string]any)
|
||||
if len(content) != 2 || content[0].(map[string]any)["image"] != "https://cdn.test/ref.png" || content[1].(map[string]any)["text"] != "draw" {
|
||||
t.Fatalf("content=%#v", content)
|
||||
}
|
||||
if parameters["size"] != "1024*768" || parameters["n"] != float64(1) || parameters["watermark"] != false {
|
||||
t.Fatalf("parameters=%#v", parameters)
|
||||
}
|
||||
if _, exists := parameters["thinking_mode"]; exists {
|
||||
t.Fatalf("editing parameters=%#v", parameters)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "video first and last frames",
|
||||
request: Request{Capability: "video.generate", Model: "wan-video", Prompt: "move", InputURLs: []string{"https://cdn.test/first.png", "https://cdn.test/last.png"}, Settings: map[string]any{
|
||||
"resolution": "1080p", "duration": 8,
|
||||
}},
|
||||
check: func(t *testing.T, body map[string]any) {
|
||||
input := body["input"].(map[string]any)
|
||||
media := input["media"].([]any)
|
||||
parameters := body["parameters"].(map[string]any)
|
||||
if len(media) != 2 || media[0].(map[string]any)["type"] != "first_frame" || media[1].(map[string]any)["type"] != "last_frame" {
|
||||
t.Fatalf("media=%#v", media)
|
||||
}
|
||||
if parameters["resolution"] != "1080P" || parameters["duration"] != float64(8) || parameters["prompt_extend"] != true || parameters["watermark"] != false {
|
||||
t.Fatalf("parameters=%#v", parameters)
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
var body map[string]any
|
||||
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
test.check(t, body)
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"output":{"task_id":"task-1","task_status":"PENDING"}}`)), Header: http.Header{}}, nil
|
||||
})
|
||||
if _, err := NewBailian(Config{BaseURL: "https://b.test", APIKey: "secret"}, client).Submit(context.Background(), test.request); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBailianDecodesCompatibleModeChoiceImages(t *testing.T) {
|
||||
result := decodeBailian([]byte(`{"output":{"task_id":"task-1","task_status":"SUCCEEDED","choices":[{"message":{"content":[{"image":"https://cdn.test/choice.png"}]}}]}}`))
|
||||
if result.Status != StatusSucceeded || len(result.OutputURLs) != 1 || result.OutputURLs[0] != "https://cdn.test/choice.png" {
|
||||
t.Fatalf("result=%#v", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSeedancePreservesTypedMultimodalMaterials(t *testing.T) {
|
||||
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
var body map[string]any
|
||||
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
content := body["content"].([]any)
|
||||
if len(content) != 4 {
|
||||
t.Fatalf("content=%#v", content)
|
||||
}
|
||||
checks := []struct{ materialType, urlKey, role string }{
|
||||
{"image_url", "image_url", "reference_image"},
|
||||
{"video_url", "video_url", "reference_video"},
|
||||
{"audio_url", "audio_url", "reference_audio"},
|
||||
}
|
||||
for index, check := range checks {
|
||||
item := content[index+1].(map[string]any)
|
||||
if item["type"] != check.materialType || item["role"] != check.role || item["label"] != []string{"图片1", "视频1", "音频1"}[index] {
|
||||
t.Fatalf("content[%d]=%#v", index+1, item)
|
||||
}
|
||||
if object := item[check.urlKey].(map[string]any); object["url"] == "" {
|
||||
t.Fatalf("content[%d]=%#v", index+1, item)
|
||||
}
|
||||
}
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"id":"task-1","status":"queued"}`)), Header: http.Header{}}, nil
|
||||
})
|
||||
request := Request{
|
||||
Capability: "video.generate", Prompt: "combine",
|
||||
Materials: []Material{
|
||||
{URL: "https://cdn.test/image.png", Type: MaterialImage, Label: "图片1"},
|
||||
{URL: "https://cdn.test/video.mp4", Type: MaterialVideo, Label: "视频1"},
|
||||
{URL: "https://cdn.test/audio.mp3", Type: MaterialAudio, Label: "音频1"},
|
||||
},
|
||||
}
|
||||
if _, err := NewSeedance(Config{BaseURL: "https://s.test/api/v3", APIKey: "secret", Model: "seed"}, client).Submit(context.Background(), request); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvoLinkPayloadIncludesQualityAndSize(t *testing.T) {
|
||||
client := roundTripFunc(func(request *http.Request) (*http.Response, error) {
|
||||
var body map[string]any
|
||||
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if body["quality"] != "high" || body["size"] != "16:9" || body["resolution"] != "1K" {
|
||||
t.Fatalf("body=%#v", body)
|
||||
}
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"id":"task-1","status":"queued"}`)), Header: http.Header{}}, nil
|
||||
})
|
||||
_, err := NewEvoLink(Config{BaseURL: "https://e.test", APIKey: "secret", Model: "gpt-image-2"}, client).Submit(context.Background(), Request{
|
||||
Capability: "image.generate", Prompt: "draw", Settings: map[string]any{"quality": " high ", "width": 1920, "height": 1080},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMockIsDeterministic(t *testing.T) {
|
||||
m := NewMock("fixture")
|
||||
a, _ := m.Submit(context.Background(), Request{Capability: "video.generate", Prompt: "hello"})
|
||||
b, _ := m.Submit(context.Background(), Request{Capability: "video.generate", Prompt: "hello"})
|
||||
if a.TaskID != b.TaskID {
|
||||
t.Fatalf("task IDs differ: %s %s", a.TaskID, b.TaskID)
|
||||
}
|
||||
got, _ := m.Query(context.Background(), a.TaskID)
|
||||
if got.Status != "succeeded" || len(got.OutputURLs) != 1 {
|
||||
t.Fatalf("Query=%#v", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Volcengine struct {
|
||||
config Config
|
||||
client HTTPClient
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
func NewVolcengine(c Config, client HTTPClient, now func() time.Time) *Volcengine {
|
||||
if client == nil {
|
||||
client = http.DefaultClient
|
||||
}
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
if c.Region == "" {
|
||||
c.Region = "cn-north-1"
|
||||
}
|
||||
if c.Service == "" {
|
||||
c.Service = "cv"
|
||||
}
|
||||
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}
|
||||
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)
|
||||
}
|
||||
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) {
|
||||
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)})
|
||||
}
|
||||
func (v *Volcengine) call(ctx context.Context, action string, payload any) (Result, error) {
|
||||
body, _ := json.Marshal(payload)
|
||||
endpoint, err := url.Parse(v.config.BaseURL)
|
||||
if err != nil {
|
||||
return Result{}, errors.New("invalid provider base URL")
|
||||
}
|
||||
q := endpoint.Query()
|
||||
q.Set("Action", action)
|
||||
q.Set("Version", "2022-08-31")
|
||||
endpoint.RawQuery = canonicalQuery(q)
|
||||
date := v.now().UTC()
|
||||
xdate := date.Format("20060102T150405Z")
|
||||
short := xdate[:8]
|
||||
hash := sha(body)
|
||||
headers := "content-type:application/json\nhost:" + endpoint.Host + "\nx-content-sha256:" + hash + "\nx-date:" + xdate + "\n"
|
||||
signed := "content-type;host;x-content-sha256;x-date"
|
||||
canonicalPath := endpoint.EscapedPath()
|
||||
if canonicalPath == "" {
|
||||
canonicalPath = "/"
|
||||
}
|
||||
canonical := "POST\n" + canonicalPath + "\n" + endpoint.RawQuery + "\n" + headers + "\n" + signed + "\n" + hash
|
||||
scope := short + "/" + v.config.Region + "/" + v.config.Service + "/request"
|
||||
stringToSign := "HMAC-SHA256\n" + xdate + "\n" + scope + "\n" + sha([]byte(canonical))
|
||||
key := hmacBytes([]byte(v.config.SecretAccessKey), short)
|
||||
key = hmacBytes(key, v.config.Region)
|
||||
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.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 {
|
||||
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) {
|
||||
return Result{}, &ProviderError{Operation: "volcengine request", Status: resp.StatusCode}
|
||||
}
|
||||
r := record(raw)
|
||||
d := object(r["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
|
||||
}
|
||||
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)
|
||||
_, _ = h.Write([]byte(s))
|
||||
return h.Sum(nil)
|
||||
}
|
||||
func canonicalQuery(q url.Values) string {
|
||||
keys := make([]string, 0, len(q))
|
||||
for k := range q {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
parts := []string{}
|
||||
for _, k := range keys {
|
||||
for _, v := range q[k] {
|
||||
parts = append(parts, url.QueryEscape(k)+"="+url.QueryEscape(v))
|
||||
}
|
||||
}
|
||||
return strings.Join(parts, "&")
|
||||
}
|
||||
Reference in new issue
Block a user