Files
easyai-ai-gateway/apps/api/internal/clients/openai.go
T
easyai e07a997aa9 feat(api): 统一官方兼容接口响应协议
兼容接口现在以入口协议作为最终响应协议,同协议保留官方 Wire 响应,跨协议统一转换成功、任务状态与错误结构。

同时修正异步提交状态边界,持久化兼容公开任务标识和官方提交响应,并新增迁移、流式响应及协议契约测试。

验证:go vet ./...;go test ./...;govulncheck ./...;pnpm lint;pnpm test;pnpm build;pnpm audit --audit-level high;pnpm openapi;全部 CI 脚本。
2026-07-22 15:34:59 +08:00

287 lines
9.5 KiB
Go

package clients
import (
"bytes"
"context"
"encoding/json"
"net/http"
"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, _ := json.Marshal(body)
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", "application/json")
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 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}},
}
}