初始化
This commit is contained in:
@@ -0,0 +1,103 @@
|
||||
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 ""
|
||||
}
|
||||
Reference in New Issue
Block a user