// 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 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} } 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) } } }