Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
This commit is contained in:
@@ -0,0 +1,705 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
const (
|
||||
openAIWSClientReadLimitBytesDefault int64 = 64 * 1024 * 1024
|
||||
openAIWSHTTPBridgeThresholdBytesDefault int64 = 15 * 1024 * 1024
|
||||
openAIWSHTTPBridgeErrorBodyLimitBytes = 64 * 1024
|
||||
)
|
||||
|
||||
const openAIWSHTTPBridgeToolStateContextKey = "openai_ws_http_bridge_tool_state"
|
||||
|
||||
type openAIWSHTTPBridgeToolState struct {
|
||||
ClientMapping apicompat.ResponsesClientToolMapping
|
||||
LoweredTools json.RawMessage
|
||||
}
|
||||
|
||||
func openAIWSHTTPBridgeToolStateFromContext(c *gin.Context) (openAIWSHTTPBridgeToolState, bool) {
|
||||
if c == nil {
|
||||
return openAIWSHTTPBridgeToolState{}, false
|
||||
}
|
||||
value, ok := c.Get(openAIWSHTTPBridgeToolStateContextKey)
|
||||
state, typed := value.(openAIWSHTTPBridgeToolState)
|
||||
return state, ok && typed
|
||||
}
|
||||
|
||||
func setOpenAIWSHTTPBridgeToolState(c *gin.Context, state openAIWSHTTPBridgeToolState) {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
state.LoweredTools = append(json.RawMessage(nil), state.LoweredTools...)
|
||||
c.Set(openAIWSHTTPBridgeToolStateContextKey, state)
|
||||
}
|
||||
|
||||
func decodeOpenAIWSHTTPBridgeLoweredTools(raw json.RawMessage) []any {
|
||||
if len(raw) == 0 {
|
||||
return nil
|
||||
}
|
||||
var tools []any
|
||||
if err := json.Unmarshal(raw, &tools); err != nil {
|
||||
return nil
|
||||
}
|
||||
return tools
|
||||
}
|
||||
|
||||
func openAIWSHTTPBridgeRawField(body []byte, name string) (json.RawMessage, bool) {
|
||||
var fields map[string]json.RawMessage
|
||||
if err := json.Unmarshal(body, &fields); err != nil {
|
||||
return nil, false
|
||||
}
|
||||
raw, present := fields[name]
|
||||
return append(json.RawMessage(nil), raw...), present
|
||||
}
|
||||
|
||||
func openAIWSHTTPBridgeToolUpstreamName(account *Account) string {
|
||||
if account != nil && account.Platform == PlatformGrok {
|
||||
return "Grok WS HTTP bridge"
|
||||
}
|
||||
return "OpenAI WS HTTP bridge"
|
||||
}
|
||||
|
||||
// ResolveOpenAIWSClientFirstMessageTimeout returns the effective client ingress deadline.
|
||||
func ResolveOpenAIWSClientFirstMessageTimeout(cfg *config.Config) time.Duration {
|
||||
seconds := config.DefaultOpenAIWSClientFirstMessageTimeoutSeconds
|
||||
if cfg != nil && cfg.Gateway.OpenAIWS.ClientFirstMessageTimeoutSeconds > 0 {
|
||||
seconds = cfg.Gateway.OpenAIWS.ClientFirstMessageTimeoutSeconds
|
||||
}
|
||||
return time.Duration(seconds) * time.Second
|
||||
}
|
||||
|
||||
func ResolveOpenAIWSClientReadLimitBytes(cfg *config.Config) int64 {
|
||||
if cfg == nil || cfg.Gateway.OpenAIWS.ClientReadLimitBytes <= 0 {
|
||||
return openAIWSClientReadLimitBytesDefault
|
||||
}
|
||||
return cfg.Gateway.OpenAIWS.ClientReadLimitBytes
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) openAIWSHTTPBridgeEnabled() bool {
|
||||
return s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.HTTPBridgeEnabled
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) openAIWSHTTPBridgeThresholdBytes() int64 {
|
||||
if s == nil || s.cfg == nil || s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes <= 0 {
|
||||
return openAIWSHTTPBridgeThresholdBytesDefault
|
||||
}
|
||||
return s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) shouldBridgeOpenAIWSHTTP(account *Account, payloadBytes int, previousResponseID string) bool {
|
||||
if account != nil && account.Platform == PlatformGrok {
|
||||
return true
|
||||
}
|
||||
if !s.openAIWSHTTPBridgeEnabled() {
|
||||
return false
|
||||
}
|
||||
if strings.TrimSpace(previousResponseID) != "" {
|
||||
return false
|
||||
}
|
||||
threshold := s.openAIWSHTTPBridgeThresholdBytes()
|
||||
return threshold > 0 && int64(payloadBytes) >= threshold
|
||||
}
|
||||
|
||||
func prepareOpenAIWSHTTPBridgeBody(payload []byte) ([]byte, error) {
|
||||
var body map[string]any
|
||||
if err := json.Unmarshal(payload, &body); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if body == nil {
|
||||
return nil, errors.New("response.create payload must be a JSON object")
|
||||
}
|
||||
delete(body, "type")
|
||||
delete(body, "generate")
|
||||
delete(body, "previous_response_id")
|
||||
body["stream"] = true
|
||||
return json.Marshal(body)
|
||||
}
|
||||
|
||||
type openAIWSToolCallReplayCollector struct {
|
||||
items []json.RawMessage
|
||||
seen map[string]struct{}
|
||||
allItems []json.RawMessage
|
||||
allSeen map[string]struct{}
|
||||
}
|
||||
|
||||
func (c *openAIWSToolCallReplayCollector) AddEvent(eventType string, message []byte) {
|
||||
switch strings.TrimSpace(eventType) {
|
||||
case "response.output_item.done":
|
||||
item := gjson.GetBytes(message, "item")
|
||||
c.addAllItem(item)
|
||||
c.addItem(item)
|
||||
case "response.completed", "response.done":
|
||||
output := gjson.GetBytes(message, "response.output")
|
||||
if !output.IsArray() {
|
||||
return
|
||||
}
|
||||
for _, item := range output.Array() {
|
||||
c.addAllItem(item)
|
||||
c.addItem(item)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *openAIWSToolCallReplayCollector) Items() []json.RawMessage {
|
||||
return cloneOpenAIWSRawMessages(c.items)
|
||||
}
|
||||
|
||||
func (c *openAIWSToolCallReplayCollector) AllItems() []json.RawMessage {
|
||||
return cloneOpenAIWSRawMessages(c.allItems)
|
||||
}
|
||||
|
||||
func (c *openAIWSToolCallReplayCollector) addAllItem(item gjson.Result) {
|
||||
if !item.Exists() || item.Type != gjson.JSON {
|
||||
return
|
||||
}
|
||||
raw := strings.TrimSpace(item.Raw)
|
||||
if raw == "" || !strings.HasPrefix(raw, "{") || strings.TrimSpace(item.Get("type").String()) == "" {
|
||||
return
|
||||
}
|
||||
key := strings.TrimSpace(item.Get("id").String())
|
||||
if key == "" {
|
||||
key = strings.TrimSpace(item.Get("call_id").String())
|
||||
}
|
||||
if key == "" {
|
||||
key = raw
|
||||
}
|
||||
if c.allSeen == nil {
|
||||
c.allSeen = make(map[string]struct{})
|
||||
}
|
||||
if _, ok := c.allSeen[key]; ok {
|
||||
return
|
||||
}
|
||||
c.allSeen[key] = struct{}{}
|
||||
c.allItems = append(c.allItems, json.RawMessage(raw))
|
||||
}
|
||||
|
||||
func (c *openAIWSToolCallReplayCollector) addItem(item gjson.Result) {
|
||||
if !item.Exists() || item.Type != gjson.JSON {
|
||||
return
|
||||
}
|
||||
raw := strings.TrimSpace(item.Raw)
|
||||
if raw == "" || !strings.HasPrefix(raw, "{") {
|
||||
return
|
||||
}
|
||||
if !isCodexToolCallContextItemType(item.Get("type").String()) {
|
||||
return
|
||||
}
|
||||
key := strings.TrimSpace(item.Get("id").String())
|
||||
if key == "" {
|
||||
key = strings.TrimSpace(item.Get("call_id").String())
|
||||
}
|
||||
if key == "" {
|
||||
key = raw
|
||||
}
|
||||
if c.seen == nil {
|
||||
c.seen = make(map[string]struct{})
|
||||
}
|
||||
if _, ok := c.seen[key]; ok {
|
||||
return
|
||||
}
|
||||
c.seen[key] = struct{}{}
|
||||
c.items = append(c.items, json.RawMessage(raw))
|
||||
}
|
||||
|
||||
func buildOpenAIWSHTTPBridgeErrorEvent(statusCode int, message string) []byte {
|
||||
message = strings.TrimSpace(message)
|
||||
if message == "" {
|
||||
message = http.StatusText(statusCode)
|
||||
}
|
||||
if message == "" {
|
||||
message = "upstream request failed"
|
||||
}
|
||||
event := map[string]any{
|
||||
"type": "error",
|
||||
"status": statusCode,
|
||||
"error": map[string]any{
|
||||
"type": "upstream_error",
|
||||
"message": message,
|
||||
},
|
||||
}
|
||||
body, err := json.Marshal(event)
|
||||
if err != nil {
|
||||
return []byte(`{"type":"error","error":{"type":"upstream_error","message":"upstream request failed"}}`)
|
||||
}
|
||||
return body
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
ctx context.Context,
|
||||
c *gin.Context,
|
||||
account *Account,
|
||||
token string,
|
||||
payload []byte,
|
||||
payloadBytes int,
|
||||
originalModel string,
|
||||
imageBillingModel string,
|
||||
imageSizeTier string,
|
||||
imageInputSize string,
|
||||
grokCacheIdentity string,
|
||||
turn int,
|
||||
writeClientMessage func([]byte) error,
|
||||
) (*OpenAIForwardResult, error) {
|
||||
if s == nil {
|
||||
return nil, errors.New("service is nil")
|
||||
}
|
||||
if s.httpUpstream == nil {
|
||||
return nil, errors.New("openai http upstream is nil")
|
||||
}
|
||||
if account == nil {
|
||||
return nil, errors.New("account is nil")
|
||||
}
|
||||
if writeClientMessage == nil {
|
||||
return nil, errors.New("client websocket writer is nil")
|
||||
}
|
||||
responseModelObserver := &upstreamResponseModelObserver{}
|
||||
|
||||
body, err := prepareOpenAIWSHTTPBridgeBody(payload)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prepare http bridge body: %w", err)
|
||||
}
|
||||
grokIntentSourceBody := append([]byte(nil), body...)
|
||||
_, grokExplicitToolsField := openAIWSHTTPBridgeRawField(grokIntentSourceBody, "tools")
|
||||
grokExplicitToolIntent := account.Platform == PlatformGrok && hasGrokResponsesToolIntent(grokIntentSourceBody)
|
||||
var clientToolMapping apicompat.ResponsesClientToolMapping
|
||||
functionToolUpstream := (account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey) || account.Platform == PlatformGrok
|
||||
if functionToolUpstream {
|
||||
if account.Platform == PlatformGrok {
|
||||
body, err = sanitizeGrokResponsesInput(body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("sanitize Grok WS HTTP bridge input: %w", err)
|
||||
}
|
||||
}
|
||||
inheritedState, _ := openAIWSHTTPBridgeToolStateFromContext(c)
|
||||
inheritedLoweredTools := decodeOpenAIWSHTTPBridgeLoweredTools(inheritedState.LoweredTools)
|
||||
body, clientToolMapping, err = adaptResponsesClientToolsForFunctionUpstreamWithMapping(
|
||||
body,
|
||||
openAIWSHTTPBridgeToolUpstreamName(account),
|
||||
inheritedState.ClientMapping,
|
||||
inheritedLoweredTools,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("adapt %s client tools: %w", openAIWSHTTPBridgeToolUpstreamName(account), err)
|
||||
}
|
||||
if account.Platform == PlatformGrok && !grokExplicitToolsField && !grokExplicitToolIntent && len(inheritedLoweredTools) > 0 && hasGrokResponsesToolIntent(body) {
|
||||
// This continuation omitted tools, so the pre-adapter source cannot
|
||||
// represent the effective inherited declarations. Cache routing must
|
||||
// see the rehydrated tool intent or it will replace client functions
|
||||
// with the native-search tool-free route. Explicit current-turn tool
|
||||
// intent still uses the original pre-sanitization source above.
|
||||
grokIntentSourceBody = append(grokIntentSourceBody[:0], body...)
|
||||
}
|
||||
loweredTools := inheritedState.LoweredTools
|
||||
if currentTools, present := openAIWSHTTPBridgeRawField(body, "tools"); present {
|
||||
loweredTools = currentTools
|
||||
}
|
||||
setOpenAIWSHTTPBridgeToolState(c, openAIWSHTTPBridgeToolState{
|
||||
ClientMapping: clientToolMapping,
|
||||
LoweredTools: loweredTools,
|
||||
})
|
||||
}
|
||||
|
||||
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
|
||||
var upstreamReq *http.Request
|
||||
if account.Platform == PlatformGrok {
|
||||
upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel)
|
||||
body, err = patchGrokResponsesBody(body, upstreamModel)
|
||||
if err != nil {
|
||||
releaseUpstreamCtx()
|
||||
return nil, err
|
||||
}
|
||||
grokMixedCacheIntentBody := append([]byte(nil), body...)
|
||||
body, err = applyGrokResponsesCacheIdentity(body, grokIntentSourceBody, grokCacheIdentity, account.IsGrokOAuth())
|
||||
if err != nil {
|
||||
releaseUpstreamCtx()
|
||||
return nil, fmt.Errorf("apply grok prompt cache identity: %w", err)
|
||||
}
|
||||
body, err = applyGrokFreeRequestToolCacheRoute(c, body, grokMixedCacheIntentBody, account, grokCacheIdentity)
|
||||
if err != nil {
|
||||
releaseUpstreamCtx()
|
||||
return nil, fmt.Errorf("apply grok Free function-tool cache route: %w", err)
|
||||
}
|
||||
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity, s.cfg, s.settingService)
|
||||
} else {
|
||||
upstreamReq, err = s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token)
|
||||
}
|
||||
releaseUpstreamCtx()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if account.Platform != PlatformGrok && isOpenAIResponsesLiteWebSocketPayload(payload) {
|
||||
upstreamReq.Header.Set(responsesLiteHeader, "true")
|
||||
}
|
||||
|
||||
proxyURL := ""
|
||||
if account.ProxyID != nil && account.Proxy != nil {
|
||||
proxyURL = account.Proxy.URL()
|
||||
}
|
||||
if c != nil {
|
||||
c.Set("openai_passthrough", true)
|
||||
c.Set("openai_ws_http_bridge", true)
|
||||
}
|
||||
|
||||
turnStart := time.Now()
|
||||
resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency)
|
||||
if err != nil {
|
||||
if turn == 1 {
|
||||
return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true)
|
||||
}
|
||||
safeErr := sanitizeUpstreamErrorMessage(err.Error())
|
||||
_ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(http.StatusBadGateway, "Upstream request failed"))
|
||||
return nil, fmt.Errorf("upstream http bridge request failed: %s", safeErr)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode >= 400 {
|
||||
respBody, _ := io.ReadAll(io.LimitReader(resp.Body, openAIWSHTTPBridgeErrorBodyLimitBytes))
|
||||
upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
|
||||
if upstreamMsg == "" {
|
||||
upstreamMsg = http.StatusText(resp.StatusCode)
|
||||
}
|
||||
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody)
|
||||
if account.Platform == PlatformGrok {
|
||||
shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody)
|
||||
s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, resolveGrokWSUpstreamModel(account, body, originalModel)), account, resp.StatusCode, resp.Header, respBody)
|
||||
if shouldFailover && (turn == 1 || resp.StatusCode == http.StatusTooManyRequests) {
|
||||
return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, false)
|
||||
}
|
||||
} else if shouldFailover && (turn == 1 || resp.StatusCode == http.StatusTooManyRequests) {
|
||||
return nil, s.handleFailoverErrorResponsePassthrough(ctx, resp, c, account, body, respBody)
|
||||
}
|
||||
if account.Platform != PlatformGrok && (shouldFailover || shouldCooldownOpenAITransientUpstreamError(resp.StatusCode, respBody)) {
|
||||
canonicalModel := canonicalOpenAIAccountSchedulingModel(account, originalModel)
|
||||
s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, canonicalModel)
|
||||
}
|
||||
_ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(resp.StatusCode, upstreamMsg))
|
||||
return nil, fmt.Errorf("upstream http bridge error: status=%d message=%s", resp.StatusCode, upstreamMsg)
|
||||
}
|
||||
if account.Platform == PlatformGrok {
|
||||
s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, resolveGrokWSUpstreamModel(account, body, originalModel)), account, resp.Header, resp.StatusCode)
|
||||
}
|
||||
|
||||
responseID := ""
|
||||
usage := OpenAIUsage{}
|
||||
imageCounter := newOpenAIImageOutputCounter()
|
||||
var firstTokenMs *int
|
||||
reqStream := openAIWSPayloadBoolFromRaw(body, "stream", true)
|
||||
eventCount := 0
|
||||
tokenEventCount := 0
|
||||
terminalEventCount := 0
|
||||
replayCollector := &openAIWSToolCallReplayCollector{}
|
||||
firstEventType := ""
|
||||
lastEventType := ""
|
||||
upstreamTerminalEvent := ""
|
||||
sawDone := false
|
||||
wroteDownstream := false
|
||||
pendingClientMessages := make([][]byte, 0, 4)
|
||||
pendingClientMessageBytes := int64(0)
|
||||
capacityFailoverSuppressedLogged := false
|
||||
clientDisconnected := false
|
||||
mappedModel := ""
|
||||
needModelReplace := false
|
||||
var mappedModelBytes []byte
|
||||
if originalModel != "" {
|
||||
mappedModel = strings.TrimSpace(gjson.GetBytes(body, "model").String())
|
||||
if mappedModel == "" {
|
||||
mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
|
||||
}
|
||||
needModelReplace = mappedModel != "" && mappedModel != originalModel
|
||||
if needModelReplace {
|
||||
mappedModelBytes = []byte(mappedModel)
|
||||
}
|
||||
}
|
||||
|
||||
resultWithUsage := func() *OpenAIForwardResult {
|
||||
imageCount := imageCounter.Count()
|
||||
result := &OpenAIForwardResult{
|
||||
RequestID: responseID,
|
||||
Usage: usage,
|
||||
Model: originalModel,
|
||||
UpstreamModel: mappedModel,
|
||||
UpstreamResponseModel: responseModelObserver.Model(),
|
||||
UpstreamResponseModelConflict: responseModelObserver.Conflict(),
|
||||
ServiceTier: extractOpenAIServiceTierFromBody(body),
|
||||
ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, mappedModel, originalModel), body, mappedModel),
|
||||
Stream: reqStream,
|
||||
OpenAIWSMode: true,
|
||||
UpstreamTerminalEvent: upstreamTerminalEvent,
|
||||
ResponseHeaders: cloneHeader(resp.Header),
|
||||
Duration: time.Since(turnStart),
|
||||
FirstTokenMs: firstTokenMs,
|
||||
}
|
||||
if replayInput := replayCollector.Items(); len(replayInput) > 0 {
|
||||
result.wsReplayInput = replayInput
|
||||
result.wsReplayInputExists = true
|
||||
}
|
||||
result.wsAccountFailoverReplayInput = replayCollector.AllItems()
|
||||
if imageCount > 0 {
|
||||
result.ImageCount = imageCount
|
||||
result.ImageSize = imageSizeTier
|
||||
result.ImageInputSize = imageInputSize
|
||||
result.ImageOutputSizes = imageCounter.Sizes()
|
||||
result.BillingModel = imageBillingModel
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
maxLineSize := defaultMaxLineSize
|
||||
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
|
||||
maxLineSize = s.cfg.Gateway.MaxLineSize
|
||||
}
|
||||
if hasResponsesClientToolMapping(clientToolMapping) {
|
||||
resp.Body = newResponsesClientToolStreamBody(resp.Body, clientToolMapping, maxLineSize)
|
||||
}
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
scanBuf := getSSEScannerBuf64K()
|
||||
scanner.Buffer(scanBuf[:0], maxLineSize)
|
||||
defer putSSEScannerBuf64K(scanBuf)
|
||||
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
data, ok := extractOpenAISSEDataLine(line)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
trimmedData := strings.TrimSpace(data)
|
||||
if trimmedData == "" {
|
||||
continue
|
||||
}
|
||||
if trimmedData == "[DONE]" {
|
||||
sawDone = true
|
||||
continue
|
||||
}
|
||||
|
||||
upstreamMessage := []byte(trimmedData)
|
||||
if normalized, changed := normalizeCompletedImageGenerationStatus(upstreamMessage); changed {
|
||||
upstreamMessage = normalized
|
||||
}
|
||||
eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(upstreamMessage)
|
||||
responseModelObserver.ObserveOpenAI(upstreamMessage, eventType)
|
||||
if responseID == "" && eventResponseID != "" {
|
||||
responseID = eventResponseID
|
||||
}
|
||||
if eventType != "" {
|
||||
eventCount++
|
||||
if firstEventType == "" {
|
||||
firstEventType = eventType
|
||||
}
|
||||
lastEventType = eventType
|
||||
}
|
||||
if isOpenAIWSTokenEvent(eventType) {
|
||||
tokenEventCount++
|
||||
if firstTokenMs == nil {
|
||||
ms := int(time.Since(turnStart).Milliseconds())
|
||||
firstTokenMs = &ms
|
||||
}
|
||||
}
|
||||
if openAIWSEventShouldParseUsage(eventType) {
|
||||
parseOpenAIWSResponseUsageFromCompletedEvent(upstreamMessage, &usage)
|
||||
}
|
||||
imageCounter.AddSSEData(upstreamMessage)
|
||||
|
||||
if needModelReplace && len(mappedModelBytes) > 0 && openAIWSEventMayContainModel(eventType) && strings.Contains(trimmedData, mappedModel) {
|
||||
upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel)
|
||||
}
|
||||
if s.toolCorrector != nil && openAIWSEventMayContainToolCalls(eventType) && openAIWSMessageLikelyContainsToolCalls(upstreamMessage) {
|
||||
if corrected, changed := s.toolCorrector.CorrectToolCallsInSSEBytes(upstreamMessage); changed {
|
||||
upstreamMessage = corrected
|
||||
}
|
||||
}
|
||||
replayCollector.AddEvent(eventType, upstreamMessage)
|
||||
|
||||
var upstreamEventErr error
|
||||
if eventType == "error" || eventType == "response.failed" {
|
||||
errMessage := extractOpenAISSEErrorMessage(upstreamMessage)
|
||||
if errMessage == "" {
|
||||
errMessage = "upstream error event"
|
||||
}
|
||||
statusCode := openAIStreamFailureStatus(upstreamMessage, errMessage)
|
||||
shouldFailover := openAIStreamFailedEventShouldFailover(upstreamMessage, errMessage)
|
||||
if eventType == "error" {
|
||||
errCodeRaw, errTypeRaw, _ := parseOpenAIWSErrorEventFields(upstreamMessage)
|
||||
statusCode = openAIWSErrorHTTPStatusFromRaw(errCodeRaw, errTypeRaw)
|
||||
shouldFailover = s.shouldFailoverOpenAIUpstreamResponse(statusCode, errMessage, upstreamMessage)
|
||||
}
|
||||
requestScopedCapacity := isOpenAIUpstreamCapacityShedEvent(upstreamMessage)
|
||||
if account.Platform == PlatformGrok && eventType == "error" {
|
||||
// SSE error events do not carry an HTTP status. The local status
|
||||
// mapper therefore defaults unknown xAI codes (for example
|
||||
// new_sensitive) to 502; classify the body as a request-scoped
|
||||
// 403 before applying status-based failover or account state.
|
||||
if isGrokContentPolicyRejection(http.StatusForbidden, upstreamMessage) {
|
||||
shouldFailover = false
|
||||
} else {
|
||||
shouldFailover = s.shouldFailoverGrokUpstreamError(statusCode, upstreamMessage)
|
||||
s.handleGrokAccountUpstreamError(ctx, account, statusCode, resp.Header, upstreamMessage)
|
||||
}
|
||||
} else if eventType == "error" && shouldFailover && !requestScopedCapacity {
|
||||
accountStatus := statusCode
|
||||
if transientStatus := openAIWSPayloadTransientStatus(upstreamMessage); transientStatus != 0 {
|
||||
accountStatus = transientStatus
|
||||
}
|
||||
canonicalModel := canonicalOpenAIAccountSchedulingModel(account, originalModel)
|
||||
s.handleOpenAIAccountUpstreamError(ctx, account, accountStatus, resp.Header, upstreamMessage, canonicalModel)
|
||||
}
|
||||
if !wroteDownstream && shouldFailover && (turn == 1 || statusCode == http.StatusTooManyRequests) {
|
||||
if account.Platform == PlatformGrok {
|
||||
return nil, newOpenAIUpstreamFailoverError(statusCode, resp.Header, upstreamMessage, errMessage, false)
|
||||
}
|
||||
return nil, s.newOpenAIStreamFailoverError(c, account, true, resp.Header.Get("x-request-id"), upstreamMessage, errMessage, resp.Header)
|
||||
}
|
||||
if wroteDownstream && requestScopedCapacity && !capacityFailoverSuppressedLogged {
|
||||
logOpenAICapacityFailoverSuppressed(ctx, account, "ws_http_bridge", resp.Header.Get("x-request-id"), eventType)
|
||||
capacityFailoverSuppressedLogged = true
|
||||
}
|
||||
if eventType == "error" {
|
||||
upstreamEventErr = errors.New(errMessage)
|
||||
}
|
||||
}
|
||||
|
||||
// 客户端写出副本改写容量降载码:Codex 对 error/response.failed 中的
|
||||
// server_is_overloaded / slow_down 判致命并终止会话,改写后走客户端内置
|
||||
// 重试。账号状态与终止事件判定(下方 handleOpenAIWSTerminalTransientFailure)
|
||||
// 仍使用未改写的 upstreamMessage。
|
||||
clientMessage := upstreamMessage
|
||||
if eventType == "error" || eventType == "response.failed" {
|
||||
if rewritten, changed := sanitizeOpenAICapacityShedErrorCodeForClient(clientMessage); changed {
|
||||
clientMessage = rewritten
|
||||
}
|
||||
}
|
||||
if !clientDisconnected {
|
||||
stageBeforeSemanticOutput := turn == 1 && account.Platform == PlatformOpenAI && !wroteDownstream
|
||||
commitStagedMessages := !stageBeforeSemanticOutput ||
|
||||
openAIStreamDataStartsClientOutput(string(clientMessage), eventType) ||
|
||||
isOpenAIWSTerminalEvent(eventType)
|
||||
if stageBeforeSemanticOutput && !commitStagedMessages {
|
||||
if pendingClientMessageBytes+int64(len(clientMessage)) > openAIFirstOutputStageMaxBytes {
|
||||
return nil, s.newOpenAIStreamFailoverError(
|
||||
c,
|
||||
account,
|
||||
true,
|
||||
resp.Header.Get("x-request-id"),
|
||||
nil,
|
||||
"OpenAI WS HTTP bridge first-output staging limit exceeded",
|
||||
resp.Header,
|
||||
)
|
||||
}
|
||||
pendingClientMessages = append(pendingClientMessages, append([]byte(nil), clientMessage...))
|
||||
pendingClientMessageBytes += int64(len(clientMessage))
|
||||
} else {
|
||||
messages := append(pendingClientMessages, clientMessage)
|
||||
pendingClientMessages = nil
|
||||
pendingClientMessageBytes = 0
|
||||
for _, message := range messages {
|
||||
if err := writeClientMessage(message); err != nil {
|
||||
if isOpenAIWSClientDisconnectError(err) {
|
||||
clientDisconnected = true
|
||||
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err)
|
||||
logOpenAIWSModeInfo(
|
||||
"ingress_ws_http_bridge_client_disconnected_drain account_id=%d turn=%d close_status=%s close_reason=%s",
|
||||
account.ID,
|
||||
turn,
|
||||
closeStatus,
|
||||
truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
|
||||
)
|
||||
break
|
||||
}
|
||||
return nil, wrapOpenAIWSIngressTurnError(
|
||||
"write_client",
|
||||
fmt.Errorf("write client websocket event: %w", err),
|
||||
wroteDownstream,
|
||||
)
|
||||
}
|
||||
wroteDownstream = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if upstreamEventErr != nil {
|
||||
return resultWithUsage(), upstreamEventErr
|
||||
}
|
||||
if isOpenAIWSTerminalEvent(eventType) {
|
||||
upstreamTerminalEvent = s.handleOpenAIWSTerminalTransientFailure(ctx, account, canonicalOpenAIAccountSchedulingModel(account, originalModel), resp.Header, upstreamMessage)
|
||||
terminalEventCount++
|
||||
firstTokenMsValue := -1
|
||||
if firstTokenMs != nil {
|
||||
firstTokenMsValue = *firstTokenMs
|
||||
}
|
||||
logOpenAIWSModeInfo(
|
||||
"ingress_ws_http_bridge_turn_completed account_id=%d turn=%d response_id=%s payload_bytes=%d duration_ms=%d events=%d token_events=%d terminal_events=%d first_event=%s last_event=%s first_token_ms=%d client_disconnected=%v",
|
||||
account.ID,
|
||||
turn,
|
||||
truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen),
|
||||
payloadBytes,
|
||||
time.Since(turnStart).Milliseconds(),
|
||||
eventCount,
|
||||
tokenEventCount,
|
||||
terminalEventCount,
|
||||
truncateOpenAIWSLogValue(firstEventType, openAIWSLogValueMaxLen),
|
||||
truncateOpenAIWSLogValue(lastEventType, openAIWSLogValueMaxLen),
|
||||
firstTokenMsValue,
|
||||
clientDisconnected,
|
||||
)
|
||||
return resultWithUsage(), nil
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
streamErr := fmt.Errorf("read upstream http bridge stream: %w", err)
|
||||
if turn == 1 && !wroteDownstream {
|
||||
return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, streamErr, true)
|
||||
}
|
||||
return resultWithUsage(), streamErr
|
||||
}
|
||||
terminalErr := errors.New("upstream http bridge stream ended before terminal event")
|
||||
if sawDone {
|
||||
terminalErr = errors.New("upstream http bridge stream sent [DONE] before terminal event")
|
||||
}
|
||||
if turn == 1 && !wroteDownstream {
|
||||
return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, terminalErr, true)
|
||||
}
|
||||
return resultWithUsage(), terminalErr
|
||||
}
|
||||
|
||||
func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, seedPayload, currentPayload []byte, originalModel string) (string, error) {
|
||||
body, err := prepareOpenAIWSHTTPBridgeBody(seedPayload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
upstreamModel := resolveGrokWSUpstreamModel(account, currentPayload, originalModel)
|
||||
body, err = patchGrokResponsesBody(body, upstreamModel)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return resolveGrokCacheIdentity(c, body, "", upstreamModel), nil
|
||||
}
|
||||
|
||||
func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel string) string {
|
||||
upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
|
||||
originalModel = strings.TrimSpace(originalModel)
|
||||
// Shared ingress has already applied channel and account mappings when the
|
||||
// body model differs from the client-facing model. Only resolve from the
|
||||
// original model when the body still carries that original value.
|
||||
if account != nil && originalModel != "" && (upstreamModel == "" || upstreamModel == originalModel) {
|
||||
if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" {
|
||||
upstreamModel = mappedModel
|
||||
}
|
||||
}
|
||||
if upstreamModel == "" {
|
||||
upstreamModel = grokDefaultResponsesModel
|
||||
}
|
||||
return upstreamModel
|
||||
}
|
||||
Reference in New Issue
Block a user