- 将 OpenAI images/edits 请求构造成 multipart/form-data - 在运行时安全转存 URL 图片为 data URL 后提交二进制文件 - 覆盖多图字段、真实上游模型名和内部元数据过滤测试
512 lines
16 KiB
Go
512 lines
16 KiB
Go
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)
|
|
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 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}},
|
|
}
|
|
}
|