90 lines
2.4 KiB
Go
90 lines
2.4 KiB
Go
package providers
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"net/http"
|
|
"strings"
|
|
)
|
|
|
|
const Seedream50ProModel = "doubao-seedream-5-0-pro-260628"
|
|
|
|
// Seedream implements the synchronous Ark image generation API. Unlike the
|
|
// Visual/Jimeng adapter, a successful submission already contains the output
|
|
// URLs and therefore never needs provider-side polling.
|
|
type Seedream struct{ *httpAdapter }
|
|
|
|
func NewSeedream(c Config, client HTTPClient) *Seedream {
|
|
if client == nil {
|
|
client = http.DefaultClient
|
|
}
|
|
a := &httpAdapter{
|
|
name: "seedream",
|
|
config: c,
|
|
client: client,
|
|
submitPath: func(Request) string { return "/images/generations" },
|
|
payload: func(r Request) any {
|
|
payload := map[string]any{
|
|
"model": requestModel(r, c.Model),
|
|
"size": seedreamSetting(r.Settings, "size", "1.5K"),
|
|
"output_format": seedreamSetting(r.Settings, "outputFormat", "png"),
|
|
"response_format": "url",
|
|
"watermark": false,
|
|
}
|
|
if prompt := strings.TrimSpace(r.Prompt); prompt != "" {
|
|
payload["prompt"] = prompt
|
|
}
|
|
layerDecomposition, _ := r.Settings["layerDecomposition"].(bool)
|
|
if layerDecomposition {
|
|
payload["layer_decomposition"] = true
|
|
} else {
|
|
payload["optimize_prompt_options"] = map[string]any{
|
|
"mode": seedreamSetting(r.Settings, "optimizeMode", "standard"),
|
|
}
|
|
}
|
|
switch len(r.InputURLs) {
|
|
case 0:
|
|
case 1:
|
|
payload["image"] = r.InputURLs[0]
|
|
default:
|
|
payload["image"] = r.InputURLs
|
|
}
|
|
return payload
|
|
},
|
|
decode: decodeSeedream,
|
|
}
|
|
return &Seedream{a}
|
|
}
|
|
|
|
func (a *Seedream) Submit(ctx context.Context, request Request) (Result, error) {
|
|
return a.submit(ctx, request)
|
|
}
|
|
|
|
func (a *Seedream) Query(context.Context, string) (Result, error) {
|
|
return Result{}, errors.New("seedream generation is synchronous")
|
|
}
|
|
|
|
func decodeSeedream(raw []byte) Result {
|
|
record := record(raw)
|
|
urls := []string{}
|
|
collectURLs(record["data"], &urls)
|
|
result := Result{OutputURLs: urls}
|
|
if len(urls) > 0 {
|
|
result.Status = StatusSucceeded
|
|
} else {
|
|
result.Status = StatusFailed
|
|
result.ErrorMessage = stringValue(object(record["error"])["message"], record["message"])
|
|
}
|
|
return result
|
|
}
|
|
|
|
func seedreamSetting(settings map[string]any, key, fallback string) string {
|
|
value := strings.TrimSpace(stringValue(settings[key]))
|
|
if value == "" {
|
|
return fallback
|
|
}
|
|
return value
|
|
}
|
|
|
|
var _ Adapter = (*Seedream)(nil)
|