Files
NianAIGC/backend/internal/providers/volcengine.go
2026-09-03 17:23:00 +08:00

392 lines
12 KiB
Go

package providers
import (
"context"
"crypto/hmac"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"io"
"log"
"net"
"net/http"
"net/url"
"sort"
"strconv"
"strings"
"time"
)
type Volcengine struct {
config Config
client HTTPClient
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
}
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) {
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, 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},
})
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, 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", version)
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, 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()
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}
}
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, Code: diagnostic.Code}
}
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", Code: diagnostic.Code}
}
d := object(first(response.business["data"], response.business["Data"]))
out := []string{}
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)
_, _ = 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, "&")
}