104 lines
3.2 KiB
Go
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 ""
|
|
}
|