package clients import ( "bytes" "context" "encoding/base64" "encoding/json" "fmt" "io" "mime/multipart" "net/http" "net/textproto" "net/url" "strings" "time" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" ) type OpenAIClient struct { HTTPClient *http.Client } func (c OpenAIClient) Run(ctx context.Context, request Request) (Response, error) { apiKey := credential(request.Candidate.Credentials, "apiKey", "api_key", "key", "token") if apiKey == "" { return Response{}, &ClientError{Code: "missing_credentials", Message: "openai api key is required", Retryable: false} } protocol := request.UpstreamProtocol if protocol == "" && request.Kind == "responses" { protocol = ProtocolOpenAIResponses } endpointKind := request.Kind if request.Kind == "responses" && protocol == ProtocolOpenAIChatCompletions { endpointKind = "chat.completions" } endpoint := openAIEndpoint(endpointKind) if endpoint == "" { return Response{}, &ClientError{Code: "unsupported_kind", Message: "unsupported openai request kind", Retryable: false} } body := cloneBody(request.Body) if request.Kind == "responses" && protocol == ProtocolOpenAIChatCompletions { var convertErr error body, convertErr = ResponsesRequestToChat(request.Body, request.PreviousResponseTurns) if convertErr != nil { return Response{}, convertErr } } if endpointKind == "chat.completions" { body = NormalizeChatCompletionRequestBody(body) applyOpenAIChatReasoningParams(body, request.Candidate) body = FilterOpenAIChatRequestBody(body) } else if request.Kind == "responses" { body = FilterOpenAIResponsesRequestBody(body) if _, hasInput := body["input"]; !hasInput { if messages, hasMessages := request.Body["messages"]; hasMessages { body["input"] = messages } } delete(body, "messages") if request.UpstreamPreviousResponseID != "" { body["previous_response_id"] = request.UpstreamPreviousResponseID } else { delete(body, "previous_response_id") } } body["model"] = upstreamModelName(request.Candidate) normalizeOpenAIImageRequestBody(endpointKind, body) stream := openAIEndpointSupportsStream(endpointKind) && (request.Stream || boolValue(body, "stream")) ensureOpenAIStreamUsage(body, endpointKind, stream) raw, contentType, err := openAIRequestPayload(endpointKind, body, request.Candidate) if err != nil { return Response{}, err } upstreamEndpoint := joinURL(openAIBaseURL(endpointKind, request.Candidate), endpoint) req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamEndpoint, bytes.NewReader(raw)) if err != nil { return Response{}, err } req.Header.Set("Content-Type", contentType) req.Header.Set("Authorization", "Bearer "+apiKey) responseStartedAt := time.Now() if err := notifySubmissionStarted(request); err != nil { return Response{}, err } resp, err := httpClient(request.HTTPClient, c.HTTPClient).Do(req) if err != nil { return Response{}, &ClientError{Code: "network", Message: err.Error(), Retryable: true} } if err := notifyResponseReceived(request); err != nil { return Response{}, err } requestID := requestIDFromHTTPResponse(resp) var result map[string]any var wire *WireResponse upstreamResponseID := "" nativeStreamDelta := openAIWireStreamDelta(request.StreamDelta, openAIWireProtocol(endpointKind), resp) if request.Kind == "responses" && protocol == ProtocolOpenAIResponses && stream { result, upstreamResponseID, err = decodeNativeResponsesStream(resp, nativeStreamDelta) wire = &WireResponse{Protocol: openAIWireProtocol(endpointKind), StatusCode: resp.StatusCode, Headers: compatibleResponseHeaders(resp.Header)} } else { var streamDelta StreamDelta = nativeStreamDelta var adapter *chatResponsesStreamAdapter if request.Kind == "responses" && protocol == ProtocolOpenAIChatCompletions && stream { adapter = newChatResponsesStreamAdapter(request.PublicResponseID, request.Model) streamDelta = func(event StreamDeltaEvent) error { return adapter.delta(event, request.StreamDelta) } } if stream { result, err = decodeOpenAIResponse(resp, true, streamDelta) wire = &WireResponse{Protocol: openAIWireProtocol(endpointKind), StatusCode: resp.StatusCode, Headers: compatibleResponseHeaders(resp.Header)} } else { result, wire, err = decodeHTTPResponseForProtocol(resp, openAIWireProtocol(endpointKind)) } if err == nil && endpointKind == "chat.completions" { result = NormalizeChatCompletionResult(result) } if err == nil && request.Kind == "responses" && protocol == ProtocolOpenAIChatCompletions { chatResult := result upstreamResponseID = requestIDFromResult(chatResult) result = ChatResultToResponse(chatResult, request.PublicResponseID, request.Model, request.Body) if adapter != nil { err = adapter.done(result, request.StreamDelta) } if err == nil { return Response{ Result: result, InternalResult: chatResult, RequestID: firstNonEmptyString(requestID, upstreamResponseID), Usage: usageFromOpenAI(chatResult), Progress: providerProgress(request), ResponseStartedAt: responseStartedAt, ResponseFinishedAt: time.Now(), UpstreamProtocol: protocol, UpstreamEndpoint: endpoint, UpstreamResponseID: upstreamResponseID, PublicResponseID: request.PublicResponseID, ResponseConverted: true, Wire: func() *WireResponse { if wire != nil { wire.Converted = true } return wire }(), }, nil } } } if err == nil && request.Kind == "chat.completions" { result = NormalizeChatCompletionResult(result) } responseFinishedAt := time.Now() if err != nil { return Response{}, annotateResponseError(err, requestID, responseStartedAt, responseFinishedAt) } if requestID == "" { requestID = requestIDFromResult(result) } if request.Kind == "responses" && protocol == ProtocolOpenAIResponses { if upstreamResponseID == "" { upstreamResponseID = requestIDFromResult(result) } } publicResponseID := request.PublicResponseID if request.Kind == "responses" && protocol == ProtocolOpenAIResponses { publicResponseID = upstreamResponseID } return Response{ Result: result, RequestID: requestID, Usage: usageFromOpenAI(result), Progress: providerProgress(request), ResponseStartedAt: responseStartedAt, ResponseFinishedAt: responseFinishedAt, ResponseDurationMS: responseDurationMS(responseStartedAt, responseFinishedAt), UpstreamProtocol: protocol, UpstreamEndpoint: endpoint, UpstreamResponseID: upstreamResponseID, PublicResponseID: publicResponseID, Wire: wire, }, nil } func normalizeOpenAIImageRequestBody(endpointKind string, body map[string]any) { if endpointKind != "images.generations" && endpointKind != "images.edits" { return } // OpenAI image endpoints express output geometry through size. The Gateway // keeps the generic fields for capability validation and billing, but they // are not valid OpenAI wire parameters. for _, key := range []string{ "aspect_ratio", "aspectRatio", "resolution", "width", "height", } { delete(body, key) } } func openAIRequestPayload(endpointKind string, body map[string]any, candidate store.RuntimeModelCandidate) ([]byte, string, error) { if endpointKind != "images.edits" { raw, err := json.Marshal(body) return raw, "application/json", err } var payload bytes.Buffer writer := multipart.NewWriter(&payload) imageFieldName := openAIImageEditFieldName(candidate) images := openAIImageEditValues(firstPresent(body["images"], body["image"])) if len(images) == 0 { return nil, "", &ClientError{ Code: "invalid_parameter", Message: "image is required", Param: "image", StatusCode: http.StatusBadRequest, Retryable: false, } } for index, value := range images { contentType, image, err := openAIImageEditPayload(value) if err != nil { return nil, "", err } if err := writeOpenAIImageEditFile(writer, imageFieldName, fmt.Sprintf("image-%d%s", index+1, openAIImageFileExtension(contentType)), contentType, image); err != nil { return nil, "", err } } if mask := firstPresent(body["mask"], body["mask_image"], body["maskImage"]); mask != nil { contentType, image, err := openAIImageEditPayload(mask) if err != nil { return nil, "", err } if err := writeOpenAIImageEditFile(writer, "mask", "mask"+openAIImageFileExtension(contentType), contentType, image); err != nil { return nil, "", err } } for key, value := range body { switch key { case "image", "images", "mask", "mask_image", "maskImage": continue } if strings.HasPrefix(key, "_") || value == nil { continue } fieldValue, err := openAIFormFieldValue(value) if err != nil { return nil, "", err } if fieldValue == "" { continue } if err := writer.WriteField(key, fieldValue); err != nil { return nil, "", err } } if err := writer.Close(); err != nil { return nil, "", err } return payload.Bytes(), writer.FormDataContentType(), nil } func openAIImageEditFieldName(candidate store.RuntimeModelCandidate) string { if configured := strings.TrimSpace(stringFromAny(candidate.PlatformConfig["imageEditMultipartFieldName"])); configured == "image" || configured == "image[]" { return configured } parsed, err := url.Parse(strings.TrimSpace(candidate.BaseURL)) if err == nil && strings.EqualFold(parsed.Hostname(), "api.openai.com") { return "image[]" } return "image" } func openAIImageEditValues(value any) []any { switch typed := value.(type) { case []any: return typed case []string: out := make([]any, 0, len(typed)) for _, item := range typed { out = append(out, item) } return out case nil: return nil default: return []any{typed} } } func openAIImageEditPayload(value any) (string, []byte, error) { switch typed := value.(type) { case map[string]any: for _, key := range []string{"data", "b64_json", "base64", "url"} { if nested := typed[key]; nested != nil { return openAIImageEditPayload(nested) } } case string: raw := strings.TrimSpace(typed) if raw == "" { break } contentType := "" encoded := raw if strings.HasPrefix(strings.ToLower(raw), "data:") { prefix, payload, ok := strings.Cut(raw, ",") if !ok || !strings.Contains(strings.ToLower(prefix), ";base64") { break } contentType = strings.TrimSpace(strings.Split(strings.TrimPrefix(prefix, "data:"), ";")[0]) encoded = payload } else if strings.Contains(raw, "://") { return "", nil, &ClientError{ Code: "invalid_parameter", Message: "OpenAI image edit input must be hydrated before multipart submission", Param: "image", StatusCode: http.StatusBadRequest, Retryable: false, } } image, err := decodeOpenAIImageEditBase64(encoded) if err != nil { break } if contentType == "" { contentType = strings.TrimSpace(strings.Split(http.DetectContentType(image), ";")[0]) } return contentType, image, nil case []byte: if len(typed) > 0 { return strings.TrimSpace(strings.Split(http.DetectContentType(typed), ";")[0]), typed, nil } } return "", nil, &ClientError{ Code: "invalid_parameter", Message: "image must be a base64 or data URL payload", Param: "image", StatusCode: http.StatusBadRequest, Retryable: false, } } func decodeOpenAIImageEditBase64(value string) ([]byte, error) { normalized := strings.Map(func(char rune) rune { switch char { case '\n', '\r', '\t', ' ': return -1 default: return char } }, value) var lastErr error for _, encoding := range []*base64.Encoding{ base64.StdEncoding, base64.RawStdEncoding, base64.URLEncoding, base64.RawURLEncoding, } { payload, err := encoding.DecodeString(normalized) if err == nil && len(payload) > 0 { return payload, nil } lastErr = err } return nil, lastErr } func writeOpenAIImageEditFile(writer *multipart.Writer, fieldName string, fileName string, contentType string, payload []byte) error { header := make(textproto.MIMEHeader) header.Set("Content-Disposition", fmt.Sprintf(`form-data; name="%s"; filename="%s"`, fieldName, escapeMultipartFilename(fileName))) if contentType != "" { header.Set("Content-Type", contentType) } part, err := writer.CreatePart(header) if err != nil { return err } _, err = io.Copy(part, bytes.NewReader(payload)) return err } func openAIImageFileExtension(contentType string) string { switch strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0])) { case "image/jpeg": return ".jpg" case "image/webp": return ".webp" case "image/gif": return ".gif" default: return ".png" } } func openAIFormFieldValue(value any) (string, error) { switch typed := value.(type) { case string: return typed, nil case bool: return fmt.Sprintf("%t", typed), nil case float64: return fmt.Sprintf("%v", typed), nil case float32: return fmt.Sprintf("%v", typed), nil case int: return fmt.Sprintf("%d", typed), nil case int64: return fmt.Sprintf("%d", typed), nil case json.Number: return typed.String(), nil default: encoded, err := json.Marshal(value) return string(encoded), err } } func openAIWireStreamDelta(next StreamDelta, protocol string, response *http.Response) StreamDelta { if next == nil { return nil } return func(event StreamDeltaEvent) error { event.WireProtocol = protocol if response != nil { event.WireStatusCode = response.StatusCode event.WireHeaders = compatibleResponseHeaders(response.Header) } return next(event) } } func openAIWireProtocol(kind string) string { switch kind { case "chat.completions": return ProtocolOpenAIChatCompletions case "responses": return ProtocolOpenAIResponses case "embeddings": return ProtocolOpenAIEmbeddings case "images.generations", "images.edits": return ProtocolOpenAIImages default: return "openai_" + strings.ReplaceAll(kind, ".", "_") } } func decodeOpenAIResponse(resp *http.Response, stream bool, onDelta StreamDelta) (map[string]any, error) { if stream { result, err := decodeOpenAIStreamResponse(resp, onDelta) if err == nil { return result, nil } return nil, err } return decodeHTTPResponse(resp) } func openAIEndpoint(kind string) string { switch kind { case "chat.completions": return "/chat/completions" case "responses": return "/responses" case "embeddings": return "/embeddings" case "reranks": return "/reranks" case "images.generations": return "/images/generations" case "images.edits": return "/images/edits" default: return "" } } func openAIEndpointSupportsStream(kind string) bool { return kind == "chat.completions" || kind == "responses" } func openAIBaseURL(kind string, candidate store.RuntimeModelCandidate) string { base := strings.TrimSpace(candidate.BaseURL) if kind != "reranks" { return base } if strings.Contains(base, "/compatible-mode/") && (strings.EqualFold(candidate.Provider, "aliyun-bailian-openai") || strings.Contains(base, "dashscope")) { return strings.Replace(base, "/compatible-mode/", "/compatible-api/", 1) } if base == "" && strings.EqualFold(candidate.Provider, "aliyun-bailian-openai") { return "https://dashscope.aliyuncs.com/compatible-api/v1" } return base } func cloneBody(body map[string]any) map[string]any { out := map[string]any{} for key, value := range body { out[key] = value } return out } func ensureOpenAIStreamUsage(body map[string]any, kind string, stream bool) { if !stream || kind != "chat.completions" { return } streamOptions := map[string]any{} if existing, ok := body["stream_options"].(map[string]any); ok { for key, value := range existing { streamOptions[key] = value } } streamOptions["include_usage"] = true body["stream_options"] = streamOptions } func joinURL(base string, path string) string { base = strings.TrimRight(strings.TrimSpace(base), "/") if base == "" { base = "https://api.openai.com/v1" } return base + path } func httpClient(clients ...*http.Client) *http.Client { for _, client := range clients { if client != nil { return client } } return http.DefaultClient } func providerProgress(request Request) []Progress { return []Progress{ {Phase: "submitting", Progress: 0.35, Message: "provider request submitted", Payload: map[string]any{"clientId": request.Candidate.ClientID}}, {Phase: "fetching_result", Progress: 0.8, Message: "provider response received", Payload: map[string]any{"provider": request.Candidate.Provider}}, } }