| package openai |
|
|
| import ( |
| "bytes" |
| "context" |
| "encoding/json" |
| "fmt" |
| "io" |
| "net/http" |
| "net/url" |
| "strconv" |
| "strings" |
| "time" |
|
|
| "github.com/QuantumNous/new-api/common" |
| "github.com/QuantumNous/new-api/logger" |
| relaycommon "github.com/QuantumNous/new-api/relay/common" |
| "github.com/QuantumNous/new-api/service" |
|
|
| "github.com/gin-gonic/gin" |
| ) |
|
|
| const ( |
| maxImageTaskResponseBytes int64 = 64 * 1024 * 1024 |
| defaultImageTaskPollTimeout = 120 * time.Second |
| defaultImageTaskPollDelay = 2 * time.Second |
| minImageTaskPollDelay = 250 * time.Millisecond |
| maxImageTaskPollDelay = 5 * time.Second |
| maxImageTaskPollErrors = 3 |
| ) |
|
|
| type openAIImageTaskEnvelope struct { |
| Object string `json:"object"` |
| Status string `json:"status"` |
| TaskID string `json:"task_id"` |
| ID string `json:"id"` |
| PollURL string `json:"poll_url"` |
| ResultURL string `json:"result_url"` |
| PollAfterMS int64 `json:"poll_after_ms"` |
| Error string `json:"error"` |
| ErrorCode string `json:"error_code"` |
| } |
|
|
| type imageTaskPollStats struct { |
| TaskID string |
| Attempts int |
| Duration time.Duration |
| } |
|
|
| type imageTaskWaitFunc func(context.Context, time.Duration) error |
|
|
| func resolveOpenAIImageTask(c *gin.Context, info *relaycommon.RelayInfo, response *http.Response) (*http.Response, error) { |
| if !common.GetEnvOrDefaultBool("CODEX_IMAGE_TASK_POLL", false) || response == nil || response.StatusCode != http.StatusAccepted { |
| return response, nil |
| } |
| if info == nil { |
| return nil, fmt.Errorf("image task relay info is missing") |
| } |
|
|
| client, err := service.GetHttpClientWithProxy(info.ChannelSetting.Proxy) |
| if err != nil { |
| return nil, fmt.Errorf("create image task poll client: %w", err) |
| } |
|
|
| timeoutSeconds := common.GetEnvOrDefault("CODEX_IMAGE_TASK_POLL_TIMEOUT_SECONDS", int(defaultImageTaskPollTimeout/time.Second)) |
| if timeoutSeconds < 1 { |
| timeoutSeconds = int(defaultImageTaskPollTimeout / time.Second) |
| } |
| if timeoutSeconds > 600 { |
| timeoutSeconds = 600 |
| } |
|
|
| startedAt := time.Now() |
| resolved, stats, err := pollOpenAIImageTask( |
| c.Request.Context(), |
| response, |
| client, |
| time.Duration(timeoutSeconds)*time.Second, |
| waitForOpenAIImageTask, |
| ) |
| stats.Duration = time.Since(startedAt) |
| if err != nil { |
| return nil, err |
| } |
| if stats.Attempts > 0 { |
| logger.LogInfo(c, fmt.Sprintf( |
| "image task compatibility completed task_id=%s attempts=%d total_ms=%.3f", |
| maskImageTaskID(stats.TaskID), |
| stats.Attempts, |
| float64(stats.Duration.Microseconds())/1000, |
| )) |
| resolved.Header.Set("X-New-API-Image-Task-Compat", "polled") |
| } |
| return resolved, nil |
| } |
|
|
| func pollOpenAIImageTask( |
| ctx context.Context, |
| initial *http.Response, |
| client *http.Client, |
| timeout time.Duration, |
| wait imageTaskWaitFunc, |
| ) (*http.Response, imageTaskPollStats, error) { |
| var stats imageTaskPollStats |
| if initial == nil || initial.Body == nil { |
| return nil, stats, fmt.Errorf("image task response is empty") |
| } |
| if initial.StatusCode != http.StatusAccepted { |
| return initial, stats, nil |
| } |
|
|
| initialBody, err := readImageTaskResponseBody(initial) |
| if err != nil { |
| return nil, stats, err |
| } |
| initial.Body = io.NopCloser(bytes.NewReader(initialBody)) |
| initial.ContentLength = int64(len(initialBody)) |
|
|
| task, isTask, err := parseOpenAIImageTask(initialBody) |
| if err != nil { |
| return nil, stats, err |
| } |
| if !isTask { |
| return initial, stats, nil |
| } |
|
|
| stats.TaskID = firstNonEmpty(task.TaskID, task.ID) |
| baseURL, err := imageTaskBaseURL(initial) |
| if err != nil { |
| return nil, stats, err |
| } |
| pollClient := *client |
| pollClient.CheckRedirect = func(request *http.Request, via []*http.Request) error { |
| if len(via) >= 10 { |
| return fmt.Errorf("image task polling stopped after 10 redirects") |
| } |
| if request == nil || request.URL == nil || |
| !strings.EqualFold(baseURL.Scheme, request.URL.Scheme) || |
| !strings.EqualFold(baseURL.Host, request.URL.Host) { |
| return fmt.Errorf("image task redirect must use the original upstream origin") |
| } |
| return nil |
| } |
| client = &pollClient |
| pollURL, err := resolveImageTaskPollURL(baseURL, firstNonEmpty(task.PollURL, task.ResultURL)) |
| if err != nil { |
| return nil, stats, err |
| } |
|
|
| deadline := time.Now().Add(timeout) |
| consecutiveErrors := 0 |
| for { |
| if remaining := time.Until(deadline); remaining <= 0 { |
| return nil, stats, fmt.Errorf("image task polling timed out after %s (task_id=%s)", timeout, maskImageTaskID(stats.TaskID)) |
| } |
|
|
| delay := imageTaskPollDelay(task.PollAfterMS) |
| if remaining := time.Until(deadline); delay > remaining { |
| delay = remaining |
| } |
| if err := wait(ctx, delay); err != nil { |
| return nil, stats, fmt.Errorf("image task polling stopped: %w", err) |
| } |
|
|
| request, err := http.NewRequestWithContext(ctx, http.MethodGet, pollURL.String(), nil) |
| if err != nil { |
| return nil, stats, fmt.Errorf("create image task poll request: %w", err) |
| } |
| copyImageTaskAuthHeaders(request.Header, initial.Request) |
| request.Header.Set("Accept", "application/json") |
|
|
| pollResponse, err := client.Do(request) |
| stats.Attempts++ |
| if err != nil { |
| consecutiveErrors++ |
| if consecutiveErrors >= maxImageTaskPollErrors { |
| return nil, stats, fmt.Errorf("poll image task after %d consecutive errors: %w", consecutiveErrors, err) |
| } |
| continue |
| } |
|
|
| pollBody, readErr := readImageTaskResponseBody(pollResponse) |
| if readErr != nil { |
| return nil, stats, readErr |
| } |
| if pollResponse.StatusCode >= http.StatusInternalServerError { |
| consecutiveErrors++ |
| if consecutiveErrors >= maxImageTaskPollErrors { |
| return nil, stats, fmt.Errorf("poll image task returned HTTP %d after %d attempts", pollResponse.StatusCode, consecutiveErrors) |
| } |
| continue |
| } |
| consecutiveErrors = 0 |
| if pollResponse.StatusCode < http.StatusOK || pollResponse.StatusCode >= http.StatusMultipleChoices { |
| return nil, stats, fmt.Errorf("poll image task returned HTTP %d", pollResponse.StatusCode) |
| } |
|
|
| if isFinalOpenAIImageResponse(pollBody) { |
| pollResponse.StatusCode = http.StatusOK |
| pollResponse.Status = fmt.Sprintf("%d %s", http.StatusOK, http.StatusText(http.StatusOK)) |
| pollResponse.Body = io.NopCloser(bytes.NewReader(pollBody)) |
| pollResponse.ContentLength = int64(len(pollBody)) |
| pollResponse.Header.Del("Content-Encoding") |
| pollResponse.Header.Set("Content-Length", strconv.FormatInt(int64(len(pollBody)), 10)) |
| return pollResponse, stats, nil |
| } |
|
|
| nextTask, nextIsTask, err := parseOpenAIImageTask(pollBody) |
| if err != nil { |
| return nil, stats, err |
| } |
| if !nextIsTask { |
| return nil, stats, fmt.Errorf("image task poll response is neither a task nor a final image response") |
| } |
| task = nextTask |
| if nextID := firstNonEmpty(task.TaskID, task.ID); nextID != "" { |
| stats.TaskID = nextID |
| } |
|
|
| switch strings.ToLower(strings.TrimSpace(task.Status)) { |
| case "failed", "cancelled", "canceled", "expired": |
| message := firstNonEmpty(task.Error, task.ErrorCode, task.Status) |
| return nil, stats, fmt.Errorf("image task %s: %s", maskImageTaskID(stats.TaskID), message) |
| case "completed", "succeeded", "success": |
| return nil, stats, fmt.Errorf("image task %s completed without image data", maskImageTaskID(stats.TaskID)) |
| } |
|
|
| if nextRawURL := firstNonEmpty(task.PollURL, task.ResultURL); nextRawURL != "" { |
| pollURL, err = resolveImageTaskPollURL(baseURL, nextRawURL) |
| if err != nil { |
| return nil, stats, err |
| } |
| } |
| } |
| } |
|
|
| func parseOpenAIImageTask(body []byte) (openAIImageTaskEnvelope, bool, error) { |
| var task openAIImageTaskEnvelope |
| if err := common.Unmarshal(body, &task); err != nil { |
| return task, false, fmt.Errorf("decode image task response: %w", err) |
| } |
| if !strings.EqualFold(strings.TrimSpace(task.Object), "image.task") { |
| return task, false, nil |
| } |
| if firstNonEmpty(task.PollURL, task.ResultURL) == "" { |
| return task, false, fmt.Errorf("image task response is missing poll_url") |
| } |
| return task, true, nil |
| } |
|
|
| func isFinalOpenAIImageResponse(body []byte) bool { |
| var payload map[string]json.RawMessage |
| if err := common.Unmarshal(body, &payload); err != nil { |
| return false |
| } |
| if _, ok := payload["created"]; !ok { |
| return false |
| } |
| var images []json.RawMessage |
| if err := common.Unmarshal(payload["data"], &images); err != nil { |
| return false |
| } |
| return len(images) > 0 |
| } |
|
|
| func imageTaskBaseURL(response *http.Response) (*url.URL, error) { |
| if response.Request == nil || response.Request.URL == nil { |
| return nil, fmt.Errorf("image task response is missing its upstream request URL") |
| } |
| base := response.Request.URL |
| if base.Scheme != "http" && base.Scheme != "https" { |
| return nil, fmt.Errorf("unsupported image task upstream scheme: %s", base.Scheme) |
| } |
| if base.User != nil || base.Hostname() == "" { |
| return nil, fmt.Errorf("invalid image task upstream URL") |
| } |
| return base, nil |
| } |
|
|
| func resolveImageTaskPollURL(baseURL *url.URL, rawURL string) (*url.URL, error) { |
| if baseURL == nil { |
| return nil, fmt.Errorf("image task upstream URL is missing") |
| } |
| reference, err := url.Parse(strings.TrimSpace(rawURL)) |
| if err != nil { |
| return nil, fmt.Errorf("parse image task poll URL: %w", err) |
| } |
| resolved := baseURL.ResolveReference(reference) |
| if resolved.Scheme != "http" && resolved.Scheme != "https" { |
| return nil, fmt.Errorf("unsupported image task poll scheme: %s", resolved.Scheme) |
| } |
| if resolved.User != nil || resolved.Hostname() == "" { |
| return nil, fmt.Errorf("invalid image task poll URL") |
| } |
| if !strings.EqualFold(baseURL.Scheme, resolved.Scheme) || !strings.EqualFold(baseURL.Host, resolved.Host) { |
| return nil, fmt.Errorf("image task poll URL must use the original upstream origin") |
| } |
| return resolved, nil |
| } |
|
|
| func copyImageTaskAuthHeaders(destination http.Header, source *http.Request) { |
| if source == nil { |
| return |
| } |
| for _, name := range []string{"Authorization", "Api-Key", "X-Api-Key"} { |
| for _, value := range source.Header.Values(name) { |
| destination.Add(name, value) |
| } |
| } |
| } |
|
|
| func readImageTaskResponseBody(response *http.Response) ([]byte, error) { |
| if response == nil || response.Body == nil { |
| return nil, fmt.Errorf("image task response is empty") |
| } |
| defer response.Body.Close() |
| if response.ContentLength > maxImageTaskResponseBytes { |
| return nil, fmt.Errorf("image task response exceeds %d bytes", maxImageTaskResponseBytes) |
| } |
| body, err := io.ReadAll(io.LimitReader(response.Body, maxImageTaskResponseBytes+1)) |
| if err != nil { |
| return nil, fmt.Errorf("read image task response: %w", err) |
| } |
| if int64(len(body)) > maxImageTaskResponseBytes { |
| return nil, fmt.Errorf("image task response exceeds %d bytes", maxImageTaskResponseBytes) |
| } |
| return body, nil |
| } |
|
|
| func imageTaskPollDelay(milliseconds int64) time.Duration { |
| if milliseconds <= 0 { |
| return defaultImageTaskPollDelay |
| } |
| delay := time.Duration(milliseconds) * time.Millisecond |
| if delay < minImageTaskPollDelay { |
| return minImageTaskPollDelay |
| } |
| if delay > maxImageTaskPollDelay { |
| return maxImageTaskPollDelay |
| } |
| return delay |
| } |
|
|
| func waitForOpenAIImageTask(ctx context.Context, delay time.Duration) error { |
| timer := time.NewTimer(delay) |
| defer timer.Stop() |
| select { |
| case <-ctx.Done(): |
| return ctx.Err() |
| case <-timer.C: |
| return nil |
| } |
| } |
|
|
| func firstNonEmpty(values ...string) string { |
| for _, value := range values { |
| if trimmed := strings.TrimSpace(value); trimmed != "" { |
| return trimmed |
| } |
| } |
| return "" |
| } |
|
|
| func maskImageTaskID(taskID string) string { |
| taskID = strings.TrimSpace(taskID) |
| if len(taskID) <= 12 { |
| return taskID |
| } |
| return taskID[:8] + "..." + taskID[len(taskID)-4:] |
| } |
|
|