Files
JuYou/API/internal/provider/apimart/image.go
T
2026-08-25 17:59:42 +08:00

104 lines
3.2 KiB
Go

package apimart
import (
"context"
"encoding/json"
"errors"
"net/http"
"strings"
"juhe-factory/api/internal/provider"
)
func (c Client) SubmitImage(ctx context.Context, baseURL, apiKey, requestID string, payload map[string]any) (provider.SubmitResult, error) {
data, err := c.requestJSON(ctx, http.MethodPost, baseURL, "/images/generations", apiKey, requestID, payload)
if err != nil {
return provider.SubmitResult{}, err
}
result := provider.SubmitResult{Raw: append(json.RawMessage(nil), data...)}
var decoded any
if err := json.Unmarshal(data, &decoded); err != nil {
return provider.SubmitResult{}, errors.New("中转站提交响应不是有效 JSON")
}
result.TaskID = findString(decoded, "task_id", "id", "request_id")
result.URL = findURL(decoded)
if result.TaskID == "" && result.URL == "" {
return provider.SubmitResult{}, errors.New("中转站提交响应缺少任务 ID 或结果 URL")
}
return result, nil
}
// PollImage 通过 APIMart 统一任务接口查询图片生成状态和结果。
func (c Client) PollImage(ctx context.Context, baseURL, apiKey, taskID string) (provider.PollResult, error) {
data, err := c.requestJSON(ctx, http.MethodGet, baseURL, "/tasks/"+taskID, apiKey, "", nil)
if err != nil {
return provider.PollResult{}, err
}
var decoded any
if err := json.Unmarshal(data, &decoded); err != nil {
return provider.PollResult{}, errors.New("中转站轮询响应不是有效 JSON")
}
status := strings.ToLower(findString(decoded, "task_status", "status", "state"))
if status == "" {
status = "processing"
}
return provider.PollResult{Status: status, URL: findURL(decoded), Error: findString(decoded, "error_message", "error", "message", "detail"), Raw: append(json.RawMessage(nil), data...)}, nil
}
func findURL(value any) string {
var walk func(any, string) string
walk = func(node any, parentKey string) string {
switch data := node.(type) {
case string:
if (strings.Contains(strings.ToLower(parentKey), "url") || strings.HasPrefix(data, "http://") || strings.HasPrefix(data, "https://")) && (strings.HasPrefix(data, "https://") || strings.HasPrefix(data, "http://")) {
return data
}
case map[string]any:
for key, child := range data {
if found := walk(child, key); found != "" {
return found
}
}
case []any:
for _, child := range data {
if found := walk(child, parentKey); found != "" {
return found
}
}
}
return ""
}
return walk(value, "")
}
func findString(value any, keys ...string) string {
wanted := map[string]bool{}
for _, key := range keys {
wanted[strings.ToLower(key)] = true
}
return findMatchingString(value, func(key, _ string) bool { return wanted[strings.ToLower(key)] })
}
func findMatchingString(value any, match func(key, value string) bool) string {
switch data := value.(type) {
case map[string]any:
for key, child := range data {
if text, ok := child.(string); ok && match(key, text) {
return text
}
}
for _, child := range data {
if found := findMatchingString(child, match); found != "" {
return found
}
}
case []any:
for _, child := range data {
if found := findMatchingString(child, match); found != "" {
return found
}
}
}
return ""
}