381 lines
12 KiB
Go
381 lines
12 KiB
Go
// Package providers contains bounded protocol adapters for generation services.
|
|
package providers
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"log"
|
|
"net"
|
|
"net/http"
|
|
"net/url"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
type Status string
|
|
|
|
const (
|
|
StatusQueued Status = "queued"
|
|
StatusRunning Status = "running"
|
|
StatusSucceeded Status = "succeeded"
|
|
StatusFailed Status = "failed"
|
|
StatusCancelled Status = "cancelled"
|
|
StatusExpired Status = "expired"
|
|
)
|
|
|
|
type Request struct {
|
|
Capability string `json:"capability"`
|
|
Model string `json:"model,omitempty"`
|
|
Prompt string `json:"prompt"`
|
|
InputURLs []string `json:"inputUrls,omitempty"`
|
|
Materials []Material `json:"materials,omitempty"`
|
|
Settings map[string]any `json:"settings,omitempty"`
|
|
}
|
|
|
|
type MaterialType string
|
|
|
|
const (
|
|
MaterialImage MaterialType = "image"
|
|
MaterialVideo MaterialType = "video"
|
|
MaterialAudio MaterialType = "audio"
|
|
)
|
|
|
|
// Material retains the provider-facing type metadata that InputURLs cannot
|
|
// express. InputURLs remains supported for existing callers.
|
|
type Material struct {
|
|
URL string `json:"url"`
|
|
Type MaterialType `json:"type"`
|
|
Role string `json:"role,omitempty"`
|
|
Label string `json:"label,omitempty"`
|
|
}
|
|
|
|
func requestModel(request Request, fallback string) string {
|
|
if value := strings.TrimSpace(request.Model); value != "" {
|
|
return value
|
|
}
|
|
return fallback
|
|
}
|
|
|
|
type Result struct {
|
|
TaskID string
|
|
Status Status
|
|
OutputURLs []string
|
|
Raw json.RawMessage
|
|
ErrorMessage string
|
|
Usage map[string]int
|
|
}
|
|
|
|
// HTTPResult is the provider-neutral representation persisted with a Job.
|
|
// Output URLs and usage must survive a process restart after the external
|
|
// provider has already reached a terminal state.
|
|
type HTTPResult struct {
|
|
TaskID string `json:"taskId,omitempty"`
|
|
Status Status `json:"status"`
|
|
OutputURLs []string `json:"outputUrls"`
|
|
Raw json.RawMessage `json:"raw,omitempty"`
|
|
ErrorMessage string `json:"errorMessage,omitempty"`
|
|
Usage map[string]int `json:"usage,omitempty"`
|
|
}
|
|
|
|
func EncodeResult(result Result) (json.RawMessage, error) {
|
|
urls := result.OutputURLs
|
|
if urls == nil {
|
|
urls = []string{}
|
|
}
|
|
return json.Marshal(HTTPResult{TaskID: result.TaskID, Status: result.Status, OutputURLs: urls, Raw: result.Raw, ErrorMessage: result.ErrorMessage, Usage: result.Usage})
|
|
}
|
|
|
|
type Adapter interface {
|
|
Submit(context.Context, Request) (Result, error)
|
|
Query(context.Context, string) (Result, error)
|
|
}
|
|
|
|
// ModelQueryAdapter is implemented by providers whose query protocol requires
|
|
// the same model identifier that was used when the task was submitted. Callers
|
|
// can opt into it without widening the common Adapter contract.
|
|
type ModelQueryAdapter interface {
|
|
QueryModel(context.Context, string, string) (Result, error)
|
|
}
|
|
type HTTPClient interface {
|
|
Do(*http.Request) (*http.Response, error)
|
|
}
|
|
type Config struct {
|
|
BaseURL, APIKey, Model, AccessKeyID, SecretAccessKey, Region, Service string
|
|
MaxResponseBytes int64
|
|
}
|
|
|
|
type ProviderError struct {
|
|
Operation string
|
|
Status int
|
|
Code string
|
|
Cause error
|
|
}
|
|
|
|
func (e *ProviderError) Error() string {
|
|
if e.Status > 0 {
|
|
return fmt.Sprintf("provider %s failed with HTTP %d", e.Operation, e.Status)
|
|
}
|
|
return "provider " + e.Operation + " failed"
|
|
}
|
|
|
|
func (e *ProviderError) Unwrap() error { return e.Cause }
|
|
|
|
type httpAdapter struct {
|
|
name string
|
|
config Config
|
|
client HTTPClient
|
|
submitPath func(Request) string
|
|
queryPath func(string) string
|
|
payload func(Request) any
|
|
headers func(*http.Request)
|
|
decode func([]byte) Result
|
|
}
|
|
|
|
func (a *httpAdapter) submit(ctx context.Context, input Request) (Result, error) {
|
|
body, err := json.Marshal(a.payload(input))
|
|
if err != nil {
|
|
return Result{}, fmt.Errorf("encode provider request: %w", err)
|
|
}
|
|
return a.call(ctx, http.MethodPost, a.submitPath(input), body, "submit")
|
|
}
|
|
func (a *httpAdapter) query(ctx context.Context, id string) (Result, error) {
|
|
if strings.TrimSpace(id) == "" {
|
|
return Result{}, errors.New("provider task id is required")
|
|
}
|
|
return a.call(ctx, http.MethodGet, a.queryPath(url.PathEscape(id)), nil, "query")
|
|
}
|
|
func (a *httpAdapter) call(ctx context.Context, method, path string, body []byte, operation string) (Result, error) {
|
|
startedAt := time.Now()
|
|
base, err := url.Parse(strings.TrimRight(a.config.BaseURL, "/"))
|
|
if err != nil {
|
|
logHTTPProviderFailure(httpProviderDiagnostic{Provider: a.name, Operation: operation, ErrorClass: "config", ElapsedMS: elapsedMilliseconds(startedAt)})
|
|
return Result{}, errors.New("invalid provider base URL")
|
|
}
|
|
rel, err := url.Parse(path)
|
|
if err != nil {
|
|
logHTTPProviderFailure(httpProviderDiagnostic{Provider: a.name, Operation: operation, ErrorClass: "config", ElapsedMS: elapsedMilliseconds(startedAt)})
|
|
return Result{}, errors.New("invalid provider path")
|
|
}
|
|
base.Path = strings.TrimRight(base.Path, "/") + "/"
|
|
rel.Path = strings.TrimLeft(rel.Path, "/")
|
|
target := base.ResolveReference(rel)
|
|
req, err := http.NewRequestWithContext(ctx, method, target.String(), strings.NewReader(string(body)))
|
|
if err != nil {
|
|
logHTTPProviderFailure(httpProviderDiagnostic{Provider: a.name, Operation: operation, ErrorClass: "request", ElapsedMS: elapsedMilliseconds(startedAt)})
|
|
return Result{}, fmt.Errorf("build provider request: %w", err)
|
|
}
|
|
req.Header.Set("Content-Type", "application/json")
|
|
req.Header.Set("Authorization", "Bearer "+a.config.APIKey)
|
|
if a.headers != nil {
|
|
a.headers(req)
|
|
}
|
|
resp, err := a.client.Do(req)
|
|
if err != nil {
|
|
logHTTPProviderFailure(httpProviderDiagnostic{
|
|
Provider: a.name, Operation: operation, ErrorClass: classifyHTTPProviderTransportError(err), ElapsedMS: elapsedMilliseconds(startedAt),
|
|
})
|
|
return Result{}, &ProviderError{Operation: a.name + " " + operation, Cause: err}
|
|
}
|
|
if resp == nil {
|
|
logHTTPProviderFailure(httpProviderDiagnostic{
|
|
Provider: a.name, Operation: operation, ErrorClass: "invalid_response", ElapsedMS: elapsedMilliseconds(startedAt),
|
|
})
|
|
return Result{}, &ProviderError{Operation: a.name + " " + operation}
|
|
}
|
|
defer resp.Body.Close()
|
|
limit := a.config.MaxResponseBytes
|
|
if limit <= 0 {
|
|
limit = 2 << 20
|
|
}
|
|
raw, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
|
|
if err != nil {
|
|
logHTTPProviderFailure(httpProviderDiagnostic{
|
|
Provider: a.name, Operation: operation, Status: resp.StatusCode, ErrorClass: "response_read", ElapsedMS: elapsedMilliseconds(startedAt),
|
|
})
|
|
return Result{}, &ProviderError{Operation: a.name + " " + operation, Cause: err}
|
|
}
|
|
if int64(len(raw)) > limit {
|
|
logHTTPProviderFailure(httpProviderDiagnostic{
|
|
Provider: a.name, Operation: operation, Status: resp.StatusCode, ErrorClass: "response_too_large", ElapsedMS: elapsedMilliseconds(startedAt),
|
|
})
|
|
return Result{}, &ProviderError{Operation: a.name + " " + operation}
|
|
}
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
code, requestID, errorType := inspectHTTPProviderFailure(raw, resp.Header)
|
|
logHTTPProviderFailure(httpProviderDiagnostic{
|
|
Provider: a.name,
|
|
Operation: operation,
|
|
Status: resp.StatusCode,
|
|
Code: code,
|
|
RequestID: requestID,
|
|
ErrorType: errorType,
|
|
ErrorClass: "service",
|
|
ElapsedMS: elapsedMilliseconds(startedAt),
|
|
})
|
|
return Result{}, &ProviderError{Operation: a.name + " " + operation, Status: resp.StatusCode, Code: code}
|
|
}
|
|
if !json.Valid(raw) {
|
|
logHTTPProviderFailure(httpProviderDiagnostic{
|
|
Provider: a.name, Operation: operation, Status: resp.StatusCode, ErrorClass: "invalid_response", ElapsedMS: elapsedMilliseconds(startedAt),
|
|
})
|
|
return Result{}, &ProviderError{Operation: a.name + " " + operation}
|
|
}
|
|
result := a.decode(raw)
|
|
result.Raw = append(json.RawMessage(nil), raw...)
|
|
return result, nil
|
|
}
|
|
|
|
type httpProviderDiagnostic struct {
|
|
Provider string
|
|
Operation string
|
|
Status int
|
|
Code string
|
|
RequestID string
|
|
ErrorType string
|
|
ErrorClass string
|
|
ElapsedMS int64
|
|
}
|
|
|
|
func logHTTPProviderFailure(diagnostic httpProviderDiagnostic) {
|
|
log.Printf(
|
|
"zhinian-api provider operation failed provider=%s operation=%s status=%d code=%q requestId=%q errorType=%q errorClass=%s elapsedMs=%d",
|
|
diagnostic.Provider,
|
|
diagnostic.Operation,
|
|
diagnostic.Status,
|
|
diagnostic.Code,
|
|
diagnostic.RequestID,
|
|
diagnostic.ErrorType,
|
|
diagnostic.ErrorClass,
|
|
diagnostic.ElapsedMS,
|
|
)
|
|
}
|
|
|
|
func inspectHTTPProviderFailure(raw []byte, headers http.Header) (code, requestID, errorType string) {
|
|
root := map[string]any{}
|
|
_ = json.Unmarshal(raw, &root)
|
|
providerError := object(root["error"])
|
|
if len(providerError) == 0 {
|
|
providerError = object(root["Error"])
|
|
}
|
|
metadata := object(first(root["ResponseMetadata"], root["response_metadata"]))
|
|
metadataError := object(first(metadata["Error"], metadata["error"]))
|
|
|
|
code = safeHTTPProviderDiagnosticToken(first(
|
|
providerError["code"], providerError["Code"],
|
|
metadataError["code"], metadataError["Code"],
|
|
root["code"], root["Code"],
|
|
), 64)
|
|
requestID = safeHTTPProviderDiagnosticToken(first(
|
|
providerError["request_id"], providerError["requestId"], providerError["RequestId"], providerError["RequestID"],
|
|
root["request_id"], root["requestId"], root["RequestId"], root["RequestID"],
|
|
metadata["request_id"], metadata["requestId"], metadata["RequestId"], metadata["RequestID"],
|
|
headers.Get("X-Tt-Logid"), headers.Get("X-Request-Id"),
|
|
), 128)
|
|
errorType = safeHTTPProviderDiagnosticToken(first(
|
|
providerError["type"], providerError["Type"],
|
|
metadataError["type"], metadataError["Type"],
|
|
root["type"], root["Type"],
|
|
), 64)
|
|
return code, requestID, errorType
|
|
}
|
|
|
|
func safeHTTPProviderDiagnosticToken(value any, maxLength int) string {
|
|
var token string
|
|
switch typed := value.(type) {
|
|
case string:
|
|
token = strings.TrimSpace(typed)
|
|
case json.Number:
|
|
token = string(typed)
|
|
case float64:
|
|
token = strconv.FormatFloat(typed, 'f', -1, 64)
|
|
case float32:
|
|
token = strconv.FormatFloat(float64(typed), 'f', -1, 32)
|
|
case int:
|
|
token = strconv.Itoa(typed)
|
|
case int32:
|
|
token = strconv.FormatInt(int64(typed), 10)
|
|
case int64:
|
|
token = strconv.FormatInt(typed, 10)
|
|
}
|
|
if len(token) == 0 || len(token) > maxLength {
|
|
return ""
|
|
}
|
|
for _, character := range token {
|
|
if (character >= 'a' && character <= 'z') || (character >= 'A' && character <= 'Z') || (character >= '0' && character <= '9') || strings.ContainsRune("-_.:", character) {
|
|
continue
|
|
}
|
|
return ""
|
|
}
|
|
return token
|
|
}
|
|
|
|
func classifyHTTPProviderTransportError(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 record(raw []byte) map[string]any { var v map[string]any; _ = json.Unmarshal(raw, &v); return v }
|
|
func object(v any) map[string]any { x, _ := v.(map[string]any); return x }
|
|
func stringValue(values ...any) string {
|
|
for _, v := range values {
|
|
if s, ok := v.(string); ok && strings.TrimSpace(s) != "" {
|
|
return strings.TrimSpace(s)
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
func status(v any) Status {
|
|
s := strings.ToLower(stringValue(v))
|
|
switch s {
|
|
case "completed", "complete", "succeeded", "success", "done":
|
|
return StatusSucceeded
|
|
case "running", "processing", "generating", "in_progress":
|
|
return StatusRunning
|
|
case "failed", "error", "unknown":
|
|
return StatusFailed
|
|
case "cancelled", "canceled":
|
|
return StatusCancelled
|
|
case "expired", "not_found", "timeout":
|
|
return StatusExpired
|
|
default:
|
|
return StatusQueued
|
|
}
|
|
}
|
|
func collectURLs(v any, out *[]string) {
|
|
switch x := v.(type) {
|
|
case string:
|
|
if strings.HasPrefix(x, "http://") || strings.HasPrefix(x, "https://") {
|
|
*out = append(*out, x)
|
|
}
|
|
case []any:
|
|
for _, i := range x {
|
|
collectURLs(i, out)
|
|
}
|
|
case map[string]any:
|
|
for _, k := range []string{"url", "image_url", "imageUrl", "result_url", "resultUrl", "video_url", "file_url"} {
|
|
collectURLs(x[k], out)
|
|
}
|
|
}
|
|
}
|