From d754be0d8ec8b85eb477b4b813af3045a1e9445d Mon Sep 17 00:00:00 2001 From: shaw Date: Tue, 7 Jul 2026 22:14:46 +0800 Subject: [PATCH] =?UTF-8?q?refactor(gateway):=20=E6=8A=BD=E5=8F=96=20CC=20?= =?UTF-8?q?forwarder=20=E5=85=AC=E5=85=B1=E7=AE=A1=E7=BA=BF=E5=B9=B6?= =?UTF-8?q?=E6=8B=86=E5=88=86=E4=B8=A4=E5=A4=A7=20service=20=E6=96=87?= =?UTF-8?q?=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit PR #3802 遗留项:三个 CC forwarder(raw 直转 / responses 回退 / messages 回退)间约 85% 重复的 HTTP 管线与 SSE 循环骨架收敛到新文件 openai_gateway_cc_pipeline.go,messages / chat_completions 两条主路径中 逐字相同的错误处理块一并接入。各路径有意保留的行为差异(GLM effort 归一化、fast policy 适用范围、ClientDisconnect 语义、Grok 分支等)留在 调用方,未强行统一。 同时对两个最臃肿的网关文件做纯移动拆分(零语义变化,逐字节校验): - openai_gateway_service.go 7821→4872 行: 调度 → openai_gateway_scheduling.go passthrough → openai_gateway_passthrough.go 用量/计费/codex 快照 → openai_gateway_usage.go - gateway_service.go 10912→7294 行: 调度 → gateway_scheduling.go Anthropic APIKey 直通 → gateway_anthropic_passthrough.go Bedrock → gateway_bedrock.go 回归保障:全部搬迁块与 HEAD 逐字节 diff 一致;留存文件经"HEAD 减去搬迁 范围"重构比对,差异仅为 goimports 移除的孤儿 import;定向单测 (fallback/raw/cyber/grok/transport/调度/用量/广域 Forward-Handle 扫描) 全绿;另经独立对抗审计确认零行为差异。 --- .../service/gateway_anthropic_passthrough.go | 810 ++++ backend/internal/service/gateway_bedrock.go | 412 ++ .../internal/service/gateway_scheduling.go | 2459 +++++++++++ backend/internal/service/gateway_service.go | 3618 ----------------- .../service/openai_gateway_cc_pipeline.go | 324 ++ .../openai_gateway_chat_completions.go | 61 +- .../openai_gateway_chat_completions_raw.go | 120 +- .../service/openai_gateway_messages.go | 67 +- .../openai_gateway_messages_chat_fallback.go | 195 +- .../service/openai_gateway_passthrough.go | 1088 +++++ .../openai_gateway_responses_chat_fallback.go | 229 +- .../service/openai_gateway_scheduling.go | 1266 ++++++ .../service/openai_gateway_service.go | 2949 -------------- .../internal/service/openai_gateway_usage.go | 658 +++ 14 files changed, 7090 insertions(+), 7166 deletions(-) create mode 100644 backend/internal/service/gateway_anthropic_passthrough.go create mode 100644 backend/internal/service/gateway_bedrock.go create mode 100644 backend/internal/service/gateway_scheduling.go create mode 100644 backend/internal/service/openai_gateway_cc_pipeline.go create mode 100644 backend/internal/service/openai_gateway_passthrough.go create mode 100644 backend/internal/service/openai_gateway_scheduling.go create mode 100644 backend/internal/service/openai_gateway_usage.go diff --git a/backend/internal/service/gateway_anthropic_passthrough.go b/backend/internal/service/gateway_anthropic_passthrough.go new file mode 100644 index 0000000000..7309e141aa --- /dev/null +++ b/backend/internal/service/gateway_anthropic_passthrough.go @@ -0,0 +1,810 @@ +package service + +// 本文件由 gateway_service.go 纯移动拆分而来:Anthropic APIKey 直通 +// (passthrough)转发路径及其流式/非流式响应与 usage 解析。仅做代码搬迁, +// 无任何行为变更。 + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "sync/atomic" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" + "github.com/tidwall/gjson" + + "github.com/gin-gonic/gin" +) + +type anthropicPassthroughForwardInput struct { + Body []byte + Parsed *ParsedRequest + RequestModel string + OriginalModel string + RequestStream bool + StartTime time.Time +} + +func (s *GatewayService) forwardAnthropicAPIKeyPassthrough( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + reqModel string, + originalModel string, + reqStream bool, + startTime time.Time, +) (*ForwardResult, error) { + return s.forwardAnthropicAPIKeyPassthroughWithInput(ctx, c, account, anthropicPassthroughForwardInput{ + Body: body, + RequestModel: reqModel, + OriginalModel: originalModel, + RequestStream: reqStream, + StartTime: startTime, + }) +} + +func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput( + ctx context.Context, + c *gin.Context, + account *Account, + input anthropicPassthroughForwardInput, +) (*ForwardResult, error) { + token, tokenType, err := s.GetAccessToken(ctx, account) + if err != nil { + return nil, err + } + if tokenType != "apikey" { + return nil, fmt.Errorf("anthropic api key passthrough requires apikey token, got: %s", tokenType) + } + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + logger.LegacyPrintf("service.gateway", "[Anthropic 自动透传] 命中 API Key 透传分支: account=%d name=%s model=%s stream=%v", + account.ID, account.Name, input.RequestModel, input.RequestStream) + + if c != nil { + c.Set("anthropic_passthrough", true) + } + // Pre-filter: strip empty text blocks (including nested in tool_result) to prevent upstream 400. + input.Body = StripEmptyTextBlocks(input.Body) + // Pre-filter: strip web-search history blocks the upstream cannot accept + // (emulation-synthesized ones always; genuine ones additionally for + // passback-required third-party upstreams such as GLM/Kimi/DeepSeek, + // which reject server_tool_use with 400). input.RequestModel 已是映射后的模型 ID。 + input.Body = FilterWebSearchHistoryBlocks(input.Body, input.RequestModel) + if input.Parsed != nil { + // 透传分支也会改写实际 wire body,成功 usage hash 依赖这里同步当前 body。 + if err := input.Parsed.ReplaceBody(input.Body); err != nil { + return nil, err + } + } + + var resp *http.Response + retryStart := time.Now() + for attempt := 1; attempt <= maxRetryAttempts; attempt++ { + upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, input.RequestStream) + upstreamReq, wireBody, err := s.buildUpstreamRequestAnthropicAPIKeyPassthrough(upstreamCtx, c, account, input.Body, token) + releaseUpstreamCtx() + if err != nil { + return nil, err + } + if input.Parsed != nil && !bytes.Equal(wireBody, input.Body) { + // build 阶段会按 beta 能力清理 body,发送前同步到 ParsedRequest 当前视图。 + if err := input.Parsed.ReplaceBody(wireBody); err != nil { + return nil, err + } + input.Body = input.Parsed.Body.Bytes() + } + + resp, err = s.httpUpstream.DoWithTLS(upstreamReq, proxyURL, account.ID, account.Concurrency, s.tlsFPProfileService.ResolveTLSProfile(account)) + if err != nil { + if resp != nil && resp.Body != nil { + _ = resp.Body.Close() + } + safeErr := sanitizeUpstreamErrorMessage(err.Error()) + setOpsUpstreamError(c, 0, safeErr, "") + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: 0, + UpstreamURL: safeUpstreamURL(upstreamReq.URL.String()), + Passthrough: true, + Kind: "request_error", + Message: safeErr, + }) + c.JSON(http.StatusBadGateway, gin.H{ + "type": "error", + "error": gin.H{ + "type": "upstream_error", + "message": "Upstream request failed", + }, + }) + return nil, fmt.Errorf("upstream request failed: %s", safeErr) + } + + // 透传分支禁止 400 请求体降级重试(该重试会改写请求体) + if resp.StatusCode >= 400 && resp.StatusCode != 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) { + if attempt < maxRetryAttempts { + elapsed := time.Since(retryStart) + if elapsed >= maxRetryElapsed { + break + } + + delay := retryBackoffDelay(attempt) + remaining := maxRetryElapsed - elapsed + if delay > remaining { + delay = remaining + } + if delay <= 0 { + break + } + + respBody, _ := s.readUpstreamErrorBody(resp) + _ = resp.Body.Close() + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + UpstreamURL: safeUpstreamURL(upstreamReq.URL.String()), + Passthrough: true, + Kind: "retry", + Message: extractUpstreamErrorMessage(respBody), + Detail: func() string { + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + return truncateString(string(respBody), s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes) + } + return "" + }(), + }) + logger.LegacyPrintf("service.gateway", "Anthropic passthrough account %d: upstream error %d, retry %d/%d after %v (elapsed=%v/%v)", + account.ID, resp.StatusCode, attempt, maxRetryAttempts, delay, elapsed, maxRetryElapsed) + if err := sleepWithContext(ctx, delay); err != nil { + return nil, err + } + continue + } + break + } + + break + } + if resp == nil || resp.Body == nil { + return nil, errors.New("upstream request failed: empty response") + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) { + if s.shouldFailoverUpstreamError(resp.StatusCode) { + respBody, _ := s.readUpstreamErrorBody(resp) + _ = resp.Body.Close() + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + + logger.LegacyPrintf("service.gateway", "[Anthropic Passthrough] Upstream error (retry exhausted, failover): Account=%d(%s) Status=%d RequestID=%s Body=%s", + account.ID, account.Name, resp.StatusCode, resp.Header.Get("x-request-id"), truncateString(string(respBody), 1000)) + + s.handleRetryExhaustedSideEffects(ctx, resp, account) + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + Passthrough: true, + Kind: "retry_exhausted_failover", + Message: extractUpstreamErrorMessage(respBody), + Detail: func() string { + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + return truncateString(string(respBody), s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes) + } + return "" + }(), + }) + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + return s.handleRetryExhaustedError(ctx, resp, c, account) + } + + if resp.StatusCode >= 400 && s.shouldFailoverUpstreamError(resp.StatusCode) { + respBody, _ := s.readUpstreamErrorBody(resp) + _ = resp.Body.Close() + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + + logger.LegacyPrintf("service.gateway", "[Anthropic Passthrough] Upstream error (failover): Account=%d(%s) Status=%d RequestID=%s Body=%s", + account.ID, account.Name, resp.StatusCode, resp.Header.Get("x-request-id"), truncateString(string(respBody), 1000)) + + s.handleFailoverSideEffects(ctx, resp, account, input.RequestModel) + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + Passthrough: true, + Kind: "failover", + Message: extractUpstreamErrorMessage(respBody), + Detail: func() string { + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + return truncateString(string(respBody), s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes) + } + return "" + }(), + }) + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + + if resp.StatusCode >= 400 { + return s.handleErrorResponse(ctx, resp, c, account, input.RequestModel) + } + + var usage *ClaudeUsage + var firstTokenMs *int + var clientDisconnect bool + if input.RequestStream { + streamResult, err := s.handleStreamingResponseAnthropicAPIKeyPassthrough(ctx, resp, c, account, input.StartTime, input.RequestModel) + if err != nil { + return nil, err + } + usage = streamResult.usage + firstTokenMs = streamResult.firstTokenMs + clientDisconnect = streamResult.clientDisconnect + } else { + usage, err = s.handleNonStreamingResponseAnthropicAPIKeyPassthrough(ctx, resp, c, account) + if err != nil { + return nil, err + } + } + if usage == nil { + usage = &ClaudeUsage{} + } + + return &ForwardResult{ + RequestID: resp.Header.Get("x-request-id"), + Usage: *usage, + Model: input.OriginalModel, + UpstreamModel: input.RequestModel, + Stream: input.RequestStream, + Duration: time.Since(input.StartTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnect, + }, nil +} + +func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + token string, +) (*http.Request, []byte, error) { + targetURL := claudeAPIURL + baseURL := account.GetBaseURL() + if baseURL != "" { + validatedURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return nil, nil, err + } + targetURL = validatedURL + "/v1/messages?beta=true" + } + + // 能力维度 body sanitize:透传路径上 anthropic-beta header 原样透传客户端值, + // 依此决定是否保留 body 中的 context_management。避免“客户端 body 带字段但 + // header 忘记带 beta token”的客户端 bug 在透传场景下让上游 400。 + clientBeta := "" + if c != nil && c.Request != nil { + clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") + } + // 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准 + if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { + clientBeta = beta + } + if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { + body = sanitized + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, nil, err + } + + if c != nil && c.Request != nil { + for key, values := range c.Request.Header { + lowerKey := strings.ToLower(strings.TrimSpace(key)) + if !allowedHeaders[lowerKey] { + continue + } + wireKey := resolveWireCasing(key) + for _, v := range values { + addHeaderRaw(req.Header, wireKey, v) + } + } + } + + // 覆盖入站鉴权残留,并注入上游认证 + req.Header.Del("authorization") + req.Header.Del("x-api-key") + req.Header.Del("x-goog-api-key") + req.Header.Del("cookie") + setAnthropicAPIKeyAuthHeader(req.Header, account, token) + + if getHeaderRaw(req.Header, "content-type") == "" { + setHeaderRaw(req.Header, "content-type", "application/json") + } + if getHeaderRaw(req.Header, "anthropic-version") == "" { + setHeaderRaw(req.Header, "anthropic-version", "2023-06-01") + } + + // 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头) + account.ApplyHeaderOverrides(req.Header) + + return req, body, nil +} + +func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, + startTime time.Time, + model string, +) (*streamingResult, error) { + if s.rateLimitService != nil { + s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header) + } + + writeAnthropicPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + + contentType := strings.TrimSpace(resp.Header.Get("Content-Type")) + if contentType == "" { + contentType = "text/event-stream" + } + c.Header("Content-Type", contentType) + if c.Writer.Header().Get("Cache-Control") == "" { + c.Header("Cache-Control", "no-cache") + } + if c.Writer.Header().Get("Connection") == "" { + c.Header("Connection", "keep-alive") + } + c.Header("X-Accel-Buffering", "no") + if v := resp.Header.Get("x-request-id"); v != "" { + c.Header("x-request-id", v) + } + + w := c.Writer + flusher, ok := w.(http.Flusher) + if !ok { + return nil, errors.New("streaming not supported") + } + + usage := &ClaudeUsage{} + var firstTokenMs *int + clientDisconnected := false + sawTerminalEvent := false + + scanner := bufio.NewScanner(resp.Body) + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + scanBuf := getSSEScannerBuf64K() + scanner.Buffer(scanBuf[:0], maxLineSize) + + type scanEvent struct { + line string + err error + } + events := make(chan scanEvent, 16) + done := make(chan struct{}) + sendEvent := func(ev scanEvent) bool { + select { + case events <- ev: + return true + case <-done: + return false + } + } + var lastReadAt int64 + atomic.StoreInt64(&lastReadAt, time.Now().UnixNano()) + go func(scanBuf *sseScannerBuf64K) { + defer putSSEScannerBuf64K(scanBuf) + defer close(events) + for scanner.Scan() { + atomic.StoreInt64(&lastReadAt, time.Now().UnixNano()) + if !sendEvent(scanEvent{line: scanner.Text()}) { + return + } + } + if err := scanner.Err(); err != nil { + _ = sendEvent(scanEvent{err: err}) + } + }(scanBuf) + defer close(done) + + streamInterval := time.Duration(0) + if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 { + streamInterval = time.Duration(s.cfg.Gateway.StreamDataIntervalTimeout) * time.Second + } + var intervalTicker *time.Ticker + if streamInterval > 0 { + intervalTicker = time.NewTicker(streamInterval) + defer intervalTicker.Stop() + } + var intervalCh <-chan time.Time + if intervalTicker != nil { + intervalCh = intervalTicker.C + } + + keepaliveInterval := time.Duration(0) + if s.cfg != nil && s.cfg.Gateway.StreamKeepaliveInterval > 0 { + keepaliveInterval = time.Duration(s.cfg.Gateway.StreamKeepaliveInterval) * time.Second + } + var keepaliveTimer *time.Timer + if keepaliveInterval > 0 { + keepaliveTimer = time.NewTimer(keepaliveInterval) + defer keepaliveTimer.Stop() + } + var keepaliveCh <-chan time.Time + if keepaliveTimer != nil { + keepaliveCh = keepaliveTimer.C + } + lastDataAt := time.Now() + resetKeepaliveTimer := func() { + if keepaliveTimer == nil { + return + } + if !keepaliveTimer.Stop() { + select { + case <-keepaliveTimer.C: + default: + } + } + keepaliveTimer.Reset(keepaliveInterval) + } + inPartialEvent := false + + for { + select { + case ev, ok := <-events: + if !ok { + if !clientDisconnected { + // 兜底补刷,确保最后一个未以空行结尾的事件也能及时送达客户端。 + flusher.Flush() + } + if !sawTerminalEvent { + if clientDisconnected && streamInterval > 0 { + lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) + if time.Since(lastRead) >= streamInterval { + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, fmt.Errorf("stream usage incomplete after timeout") + } + } + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: clientDisconnected}, fmt.Errorf("stream usage incomplete: missing terminal event") + } + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: clientDisconnected}, nil + } + if ev.err != nil { + if sawTerminalEvent { + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: clientDisconnected}, nil + } + if clientDisconnected { + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, fmt.Errorf("stream usage incomplete after disconnect: %w", ev.err) + } + if errors.Is(ev.err, context.Canceled) || errors.Is(ev.err, context.DeadlineExceeded) { + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, fmt.Errorf("stream usage incomplete: %w", ev.err) + } + if errors.Is(ev.err, bufio.ErrTooLong) { + logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] SSE line too long: account=%d max_size=%d error=%v", account.ID, maxLineSize, ev.err) + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs}, ev.err + } + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs}, fmt.Errorf("stream read error: %w", ev.err) + } + + line := ev.line + if data, ok := extractAnthropicSSEDataLine(line); ok { + trimmed := strings.TrimSpace(data) + if anthropicStreamEventIsTerminal("", trimmed) { + sawTerminalEvent = true + } + if firstTokenMs == nil && trimmed != "" && trimmed != "[DONE]" { + ms := int(time.Since(startTime).Milliseconds()) + firstTokenMs = &ms + } + s.parseSSEUsagePassthrough(data, usage) + } else { + trimmed := strings.TrimSpace(line) + if strings.HasPrefix(trimmed, "event:") && anthropicStreamEventIsTerminal(strings.TrimSpace(strings.TrimPrefix(trimmed, "event:")), "") { + sawTerminalEvent = true + } + } + + if !clientDisconnected { + restored := string(reverseToolNamesIfPresent(c, []byte(line))) + if _, err := io.WriteString(w, restored); err != nil { + clientDisconnected = true + logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) + } else if _, err := io.WriteString(w, "\n"); err != nil { + clientDisconnected = true + logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) + } else if line == "" { + // 按 SSE 事件边界刷出,减少每行 flush 带来的 syscall 开销。 + flusher.Flush() + lastDataAt = time.Now() + resetKeepaliveTimer() + inPartialEvent = false + } else { + inPartialEvent = true + } + } + + case <-intervalCh: + lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) + if time.Since(lastRead) < streamInterval { + continue + } + if clientDisconnected { + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, fmt.Errorf("stream usage incomplete after timeout") + } + logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] Stream data interval timeout: account=%d model=%s interval=%s", account.ID, model, streamInterval) + if s.rateLimitService != nil { + s.rateLimitService.HandleStreamTimeout(ctx, account, model) + } + return &streamingResult{usage: usage, firstTokenMs: firstTokenMs}, fmt.Errorf("stream data interval timeout") + + case <-keepaliveCh: + if clientDisconnected { + continue + } + if inPartialEvent { + resetKeepaliveTimer() + continue + } + if time.Since(lastDataAt) < keepaliveInterval { + resetKeepaliveTimer() + continue + } + if _, err := fmt.Fprint(w, "event: ping\ndata: {\"type\": \"ping\"}\n\n"); err != nil { + clientDisconnected = true + logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] Client disconnected during keepalive ping, continue draining upstream for usage: account=%d", account.ID) + continue + } + flusher.Flush() + lastDataAt = time.Now() + resetKeepaliveTimer() + } + } +} + +func extractAnthropicSSEDataLine(line string) (string, bool) { + if !strings.HasPrefix(line, "data:") { + return "", false + } + start := len("data:") + for start < len(line) { + if line[start] != ' ' && line[start] != '\t' { + break + } + start++ + } + return line[start:], true +} + +func (s *GatewayService) parseSSEUsagePassthrough(data string, usage *ClaudeUsage) { + if usage == nil || data == "" || data == "[DONE]" { + return + } + + parsed := gjson.Parse(data) + switch parsed.Get("type").String() { + case "message_start": + msgUsage := parsed.Get("message.usage") + if msgUsage.Exists() { + usage.InputTokens = int(msgUsage.Get("input_tokens").Int()) + usage.CacheCreationInputTokens = int(msgUsage.Get("cache_creation_input_tokens").Int()) + usage.CacheReadInputTokens = int(msgUsage.Get("cache_read_input_tokens").Int()) + + // 保持与通用解析一致:message_start 允许覆盖 5m/1h 明细(包括 0)。 + cc5m := msgUsage.Get("cache_creation.ephemeral_5m_input_tokens") + cc1h := msgUsage.Get("cache_creation.ephemeral_1h_input_tokens") + if cc5m.Exists() || cc1h.Exists() { + usage.CacheCreation5mTokens = int(cc5m.Int()) + usage.CacheCreation1hTokens = int(cc1h.Int()) + } + } + case "message_delta": + deltaUsage := parsed.Get("usage") + if deltaUsage.Exists() { + if v := deltaUsage.Get("input_tokens").Int(); v > 0 { + usage.InputTokens = int(v) + } + if v := deltaUsage.Get("output_tokens").Int(); v > 0 { + usage.OutputTokens = int(v) + } + if v := deltaUsage.Get("cache_creation_input_tokens").Int(); v > 0 { + usage.CacheCreationInputTokens = int(v) + } + if v := deltaUsage.Get("cache_read_input_tokens").Int(); v > 0 { + usage.CacheReadInputTokens = int(v) + } + + cc5m := deltaUsage.Get("cache_creation.ephemeral_5m_input_tokens") + cc1h := deltaUsage.Get("cache_creation.ephemeral_1h_input_tokens") + if cc5m.Exists() && cc5m.Int() > 0 { + usage.CacheCreation5mTokens = int(cc5m.Int()) + } + if cc1h.Exists() && cc1h.Int() > 0 { + usage.CacheCreation1hTokens = int(cc1h.Int()) + } + } + } + + if usage.CacheReadInputTokens == 0 { + if cached := parsed.Get("message.usage.cached_tokens").Int(); cached > 0 { + usage.CacheReadInputTokens = int(cached) + } + if cached := parsed.Get("usage.cached_tokens").Int(); usage.CacheReadInputTokens == 0 && cached > 0 { + usage.CacheReadInputTokens = int(cached) + } + } + if usage.CacheCreationInputTokens == 0 { + cc5m := parsed.Get("message.usage.cache_creation.ephemeral_5m_input_tokens").Int() + cc1h := parsed.Get("message.usage.cache_creation.ephemeral_1h_input_tokens").Int() + if cc5m == 0 && cc1h == 0 { + cc5m = parsed.Get("usage.cache_creation.ephemeral_5m_input_tokens").Int() + cc1h = parsed.Get("usage.cache_creation.ephemeral_1h_input_tokens").Int() + } + total := cc5m + cc1h + if total > 0 { + usage.CacheCreationInputTokens = int(total) + } + } +} + +func parseClaudeUsageFromResponseBody(body []byte) *ClaudeUsage { + usage := &ClaudeUsage{} + if len(body) == 0 { + return usage + } + + parsed := gjson.ParseBytes(body) + usageNode := parsed.Get("usage") + if !usageNode.Exists() { + return usage + } + + usage.InputTokens = int(usageNode.Get("input_tokens").Int()) + usage.OutputTokens = int(usageNode.Get("output_tokens").Int()) + usage.CacheCreationInputTokens = int(usageNode.Get("cache_creation_input_tokens").Int()) + usage.CacheReadInputTokens = int(usageNode.Get("cache_read_input_tokens").Int()) + + cc5m := usageNode.Get("cache_creation.ephemeral_5m_input_tokens").Int() + cc1h := usageNode.Get("cache_creation.ephemeral_1h_input_tokens").Int() + if cc5m > 0 || cc1h > 0 { + usage.CacheCreation5mTokens = int(cc5m) + usage.CacheCreation1hTokens = int(cc1h) + } + if usage.CacheCreationInputTokens == 0 && (cc5m > 0 || cc1h > 0) { + usage.CacheCreationInputTokens = int(cc5m + cc1h) + } + if usage.CacheReadInputTokens == 0 { + if cached := usageNode.Get("cached_tokens").Int(); cached > 0 { + usage.CacheReadInputTokens = int(cached) + } + } + return usage +} + +func (s *GatewayService) invalidNonStreamingJSONFailoverError( + ctx context.Context, + resp *http.Response, + account *Account, + body []byte, + parseErr error, + requestedModel ...string, +) error { + const statusCode = http.StatusBadGateway + + accountID := int64(0) + accountName := "" + retryableOnSameAccount := false + if account != nil { + accountID = account.ID + accountName = account.Name + retryableOnSameAccount = account.IsPoolMode() && account.IsPoolModeRetryableStatus(statusCode) + } + + logger.LegacyPrintf( + "service.gateway", + "Account %d(%s): upstream returned non-JSON 2xx response, attempting failover: status=%d request_id=%s error=%v", + accountID, + accountName, + resp.StatusCode, + resp.Header.Get("x-request-id"), + parseErr, + ) + + if s.rateLimitService != nil && account != nil { + if len(requestedModel) > 0 { + s.rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body, requestedModel[0]) + } else { + s.rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body) + } + } + + return &UpstreamFailoverError{ + StatusCode: statusCode, + ResponseBody: body, + ResponseHeaders: resp.Header, + RetryableOnSameAccount: retryableOnSameAccount, + } +} + +func (s *GatewayService) handleNonStreamingResponseAnthropicAPIKeyPassthrough( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, +) (*ClaudeUsage, error) { + if s.rateLimitService != nil { + s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header) + } + + body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, anthropicTooLargeError) + if err != nil { + return nil, err + } + + if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices { + var raw json.RawMessage + if err := json.Unmarshal(body, &raw); err != nil { + return nil, s.invalidNonStreamingJSONFailoverError(ctx, resp, account, body, err) + } + } + + usage := parseClaudeUsageFromResponseBody(body) + + writeAnthropicPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + contentType := strings.TrimSpace(resp.Header.Get("Content-Type")) + if contentType == "" { + contentType = "application/json" + } + body = reverseToolNamesIfPresent(c, body) + c.Data(resp.StatusCode, contentType, body) + return usage, nil +} + +func writeAnthropicPassthroughResponseHeaders(dst http.Header, src http.Header, filter *responseheaders.CompiledHeaderFilter) { + if dst == nil || src == nil { + return + } + if filter != nil { + responseheaders.WriteFilteredHeaders(dst, src, filter) + return + } + if v := strings.TrimSpace(src.Get("Content-Type")); v != "" { + dst.Set("Content-Type", v) + } + if v := strings.TrimSpace(src.Get("x-request-id")); v != "" { + dst.Set("x-request-id", v) + } +} diff --git a/backend/internal/service/gateway_bedrock.go b/backend/internal/service/gateway_bedrock.go new file mode 100644 index 0000000000..8c6042e2bf --- /dev/null +++ b/backend/internal/service/gateway_bedrock.go @@ -0,0 +1,412 @@ +package service + +// 本文件由 gateway_service.go 纯移动拆分而来:Bedrock 上游转发(CC 兼容转换、 +// 请求构建、错误处理与非流式响应)。仅做代码搬迁,无任何行为变更。 + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + + "github.com/gin-gonic/gin" +) + +// ApplyBedrockCCCompat 应用 Bedrock CC 兼容转换(渠道级模型映射后调用) +// 清理 body 中 Anthropic API 专有字段、修复 thinking/tool_use ID、过滤 beta token, +// 同时过滤 HTTP header 中的 anthropic-beta(防止 Passthrough 路径透传不支持的 token)。 +func (s *GatewayService) ApplyBedrockCCCompat(c *gin.Context, body []byte, model string, account *Account, groupID *int64) []byte { + if !s.isBedrockCCCompatEnabled(c.Request.Context(), account, groupID) { + return body + } + body = sanitizeBedrockCCFields(body) + body = sanitizeBedrockThinking(body, model) + body = sanitizeBedrockToolUseIDs(body) + body = sanitizeBedrockCCBetaTokens(body, model) + // 过滤 HTTP header 中的 anthropic-beta,只保留 Bedrock 支持的 token + if betaHeader := c.GetHeader("anthropic-beta"); betaHeader != "" { + if filtered := ResolveBedrockBetaTokens(betaHeader, body, model); len(filtered) > 0 { + c.Request.Header.Set("anthropic-beta", strings.Join(filtered, ", ")) + } else { + c.Request.Header.Del("anthropic-beta") + } + } + return body +} + +// isBedrockCCCompatEnabled 检查渠道是否启用了 Bedrock CC 兼容模式 +func (s *GatewayService) isBedrockCCCompatEnabled(ctx context.Context, account *Account, groupID *int64) bool { + if groupID == nil || s.channelService == nil { + return false + } + ch, err := s.channelService.GetChannelForGroup(ctx, *groupID) + if err != nil || ch == nil { + return false + } + return ch.IsBedrockCCCompatEnabled(account.Platform) +} + +// forwardBedrock 转发请求到 AWS Bedrock +func (s *GatewayService) forwardBedrock( + ctx context.Context, + c *gin.Context, + account *Account, + parsed *ParsedRequest, + startTime time.Time, +) (*ForwardResult, error) { + reqModel := parsed.Model + reqStream := parsed.Stream + body := parsed.Body.Bytes() + + region := bedrockRuntimeRegion(account) + mappedModel, ok := ResolveBedrockModelID(account, reqModel) + if !ok { + return nil, fmt.Errorf("unsupported bedrock model: %s", reqModel) + } + if mappedModel != reqModel { + logger.LegacyPrintf("service.gateway", "[Bedrock] Model mapping: %s -> %s (account: %s)", reqModel, mappedModel, account.Name) + } + + betaHeader := "" + if c != nil && c.Request != nil { + betaHeader = c.GetHeader("anthropic-beta") + } + + // 准备请求体(注入 anthropic_version/anthropic_beta,移除 Bedrock 不支持的字段,清理 cache_control) + betaTokens, err := s.resolveBedrockBetaTokensForRequest(ctx, account, betaHeader, body, mappedModel) + if err != nil { + return nil, err + } + + bedrockBody, err := PrepareBedrockRequestBodyWithTokens(body, mappedModel, betaTokens, false) + if err != nil { + return nil, fmt.Errorf("prepare bedrock request body: %w", err) + } + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + logger.LegacyPrintf("service.gateway", "[Bedrock] 命中 Bedrock 分支: account=%d name=%s model=%s->%s stream=%v", + account.ID, account.Name, reqModel, mappedModel, reqStream) + + // 根据账号类型选择认证方式 + var signer *BedrockSigner + var bedrockAPIKey string + if account.IsBedrockAPIKey() { + bedrockAPIKey = account.GetCredential("api_key") + if bedrockAPIKey == "" { + return nil, fmt.Errorf("api_key not found in bedrock credentials") + } + } else { + signer, err = NewBedrockSignerFromAccount(account) + if err != nil { + return nil, fmt.Errorf("create bedrock signer: %w", err) + } + } + + // 执行上游请求(含重试) + resp, err := s.executeBedrockUpstream(ctx, c, account, bedrockBody, mappedModel, region, reqStream, signer, bedrockAPIKey, proxyURL) + if err != nil { + return nil, err + } + defer func() { _ = resp.Body.Close() }() + + // 将 Bedrock 的 x-amzn-requestid 映射到 x-request-id, + // 使通用错误处理函数(handleErrorResponse、handleRetryExhaustedError)能正确提取 AWS request ID。 + if awsReqID := resp.Header.Get("x-amzn-requestid"); awsReqID != "" && resp.Header.Get("x-request-id") == "" { + resp.Header.Set("x-request-id", awsReqID) + } + + // 错误/failover 处理 + if resp.StatusCode >= 400 { + return s.handleBedrockUpstreamErrors(ctx, resp, c, account) + } + + // Bedrock 分支绕过通用 Forward 成功路径,这里保持上游接受回调语义一致。 + if parsed.OnUpstreamAccepted != nil { + parsed.OnUpstreamAccepted() + } + + // 响应处理 + var usage *ClaudeUsage + var firstTokenMs *int + var clientDisconnect bool + if reqStream { + streamResult, err := s.handleBedrockStreamingResponse(ctx, resp, c, account, startTime, reqModel) + if err != nil { + return nil, err + } + usage = streamResult.usage + firstTokenMs = streamResult.firstTokenMs + clientDisconnect = streamResult.clientDisconnect + } else { + usage, err = s.handleBedrockNonStreamingResponse(ctx, resp, c, account) + if err != nil { + return nil, err + } + } + if usage == nil { + usage = &ClaudeUsage{} + } + + return &ForwardResult{ + RequestID: resp.Header.Get("x-amzn-requestid"), + Usage: *usage, + Model: reqModel, + UpstreamModel: mappedModel, + Stream: reqStream, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnect, + }, nil +} + +// executeBedrockUpstream 执行 Bedrock 上游请求(含重试逻辑) +func (s *GatewayService) executeBedrockUpstream( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + modelID string, + region string, + stream bool, + signer *BedrockSigner, + apiKey string, + proxyURL string, +) (*http.Response, error) { + var resp *http.Response + var err error + retryStart := time.Now() + for attempt := 1; attempt <= maxRetryAttempts; attempt++ { + var upstreamReq *http.Request + if account.IsBedrockAPIKey() { + upstreamReq, err = s.buildUpstreamRequestBedrockAPIKey(ctx, body, modelID, region, stream, apiKey) + } else { + upstreamReq, err = s.buildUpstreamRequestBedrock(ctx, body, modelID, region, stream, signer) + } + if err != nil { + return nil, err + } + + resp, err = s.httpUpstream.DoWithTLS(upstreamReq, proxyURL, account.ID, account.Concurrency, nil) + if err != nil { + if resp != nil && resp.Body != nil { + _ = resp.Body.Close() + } + safeErr := sanitizeUpstreamErrorMessage(err.Error()) + setOpsUpstreamError(c, 0, safeErr, "") + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: 0, + UpstreamURL: safeUpstreamURL(upstreamReq.URL.String()), + Kind: "request_error", + Message: safeErr, + }) + c.JSON(http.StatusBadGateway, gin.H{ + "type": "error", + "error": gin.H{ + "type": "upstream_error", + "message": "Upstream request failed", + }, + }) + return nil, fmt.Errorf("upstream request failed: %s", safeErr) + } + + if resp.StatusCode >= 400 && resp.StatusCode != 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) { + if attempt < maxRetryAttempts { + elapsed := time.Since(retryStart) + if elapsed >= maxRetryElapsed { + break + } + + delay := retryBackoffDelay(attempt) + remaining := maxRetryElapsed - elapsed + if delay > remaining { + delay = remaining + } + if delay <= 0 { + break + } + + respBody, _ := s.readUpstreamErrorBody(resp) + _ = resp.Body.Close() + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamURL: safeUpstreamURL(upstreamReq.URL.String()), + Kind: "retry", + Message: extractUpstreamErrorMessage(respBody), + Detail: func() string { + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + return truncateString(string(respBody), s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes) + } + return "" + }(), + }) + logger.LegacyPrintf("service.gateway", "[Bedrock] account %d: upstream error %d, retry %d/%d after %v", + account.ID, resp.StatusCode, attempt, maxRetryAttempts, delay) + if err := sleepWithContext(ctx, delay); err != nil { + return nil, err + } + continue + } + break + } + + break + } + if resp == nil || resp.Body == nil { + return nil, errors.New("upstream request failed: empty response") + } + return resp, nil +} + +// handleBedrockUpstreamErrors 处理 Bedrock 上游 4xx/5xx 错误(failover + 错误响应) +func (s *GatewayService) handleBedrockUpstreamErrors( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, +) (*ForwardResult, error) { + // retry exhausted + failover + if s.shouldRetryUpstreamError(account, resp.StatusCode) { + if s.shouldFailoverUpstreamError(resp.StatusCode) { + respBody, _ := s.readUpstreamErrorBody(resp) + _ = resp.Body.Close() + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + + logger.LegacyPrintf("service.gateway", "[Bedrock] Upstream error (retry exhausted, failover): Account=%d(%s) Status=%d Body=%s", + account.ID, account.Name, resp.StatusCode, truncateString(string(respBody), 1000)) + + s.handleRetryExhaustedSideEffects(ctx, resp, account) + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + Kind: "retry_exhausted_failover", + Message: extractUpstreamErrorMessage(respBody), + }) + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + return s.handleRetryExhaustedError(ctx, resp, c, account) + } + + // non-retryable failover + if s.shouldFailoverUpstreamError(resp.StatusCode) { + respBody, _ := s.readUpstreamErrorBody(resp) + _ = resp.Body.Close() + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + + s.handleFailoverSideEffects(ctx, resp, account) + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + Kind: "failover", + Message: extractUpstreamErrorMessage(respBody), + }) + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + + // other errors + return s.handleErrorResponse(ctx, resp, c, account) +} + +// buildUpstreamRequestBedrock 构建 Bedrock 上游请求 +func (s *GatewayService) buildUpstreamRequestBedrock( + ctx context.Context, + body []byte, + modelID string, + region string, + stream bool, + signer *BedrockSigner, +) (*http.Request, error) { + targetURL := BuildBedrockURL(region, modelID, stream) + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, err + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + + // SigV4 签名 + if err := signer.SignRequest(ctx, req, body); err != nil { + return nil, fmt.Errorf("sign bedrock request: %w", err) + } + + return req, nil +} + +// buildUpstreamRequestBedrockAPIKey 构建 Bedrock API Key (Bearer Token) 上游请求 +func (s *GatewayService) buildUpstreamRequestBedrockAPIKey( + ctx context.Context, + body []byte, + modelID string, + region string, + stream bool, + apiKey string, +) (*http.Request, error) { + targetURL := BuildBedrockURL(region, modelID, stream) + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, err + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + req.Header.Set("Authorization", "Bearer "+apiKey) + + return req, nil +} + +// handleBedrockNonStreamingResponse 处理 Bedrock 非流式响应 +// Bedrock InvokeModel 非流式响应的 body 格式与 Claude API 兼容 +func (s *GatewayService) handleBedrockNonStreamingResponse( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, +) (*ClaudeUsage, error) { + body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, anthropicTooLargeError) + if err != nil { + return nil, err + } + + // 转换 Bedrock 特有的 amazon-bedrock-invocationMetrics 为标准 Anthropic usage 格式 + // 并移除该字段避免透传给客户端 + body = transformBedrockInvocationMetrics(body) + + usage := parseClaudeUsageFromResponseBody(body) + + c.Header("Content-Type", "application/json") + if v := resp.Header.Get("x-amzn-requestid"); v != "" { + c.Header("x-request-id", v) + } + c.Data(resp.StatusCode, "application/json", body) + return usage, nil +} diff --git a/backend/internal/service/gateway_scheduling.go b/backend/internal/service/gateway_scheduling.go new file mode 100644 index 0000000000..9a640d9fe8 --- /dev/null +++ b/backend/internal/service/gateway_scheduling.go @@ -0,0 +1,2459 @@ +package service + +// 本文件由 gateway_service.go 纯移动拆分而来:账号选择与负载感知调度、窗口费用 +// 与 RPM 预取、候选排序/过滤、混合平台调度与选择失败诊断。仅做代码搬迁, +// 无任何行为变更。 + +import ( + "context" + "fmt" + "log/slog" + mathrand "math/rand" + "sort" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/claude" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" +) + +// SelectAccount 选择账号(粘性会话+优先级) +func (s *GatewayService) SelectAccount(ctx context.Context, groupID *int64, sessionHash string) (*Account, error) { + return s.SelectAccountForModel(ctx, groupID, sessionHash, "") +} + +// SelectAccountForModel 选择支持指定模型的账号(粘性会话+优先级+模型映射) +func (s *GatewayService) SelectAccountForModel(ctx context.Context, groupID *int64, sessionHash string, requestedModel string) (*Account, error) { + return s.SelectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, nil) +} + +// SelectAccountForModelWithExclusions selects an account supporting the requested model while excluding specified accounts. +func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*Account, error) { + // 优先检查 context 中的强制平台(/antigravity 路由) + var platform string + forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string) + if hasForcePlatform && forcePlatform != "" { + platform = forcePlatform + } else if groupID != nil { + group, resolvedGroupID, err := s.resolveGatewayGroup(ctx, groupID) + if err != nil { + return nil, err + } + groupID = resolvedGroupID + ctx = s.withGroupContext(ctx, group) + platform = group.Platform + } else { + // 无分组时只使用原生 anthropic 平台 + platform = PlatformAnthropic + } + + // Claude Code 限制可能已将 groupID 解析为 fallback group, + // 渠道限制预检查必须使用解析后的分组。 + if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { + slog.Warn("channel pricing restriction blocked request", + "group_id", derefGroupID(groupID), + "model", requestedModel) + return nil, fmt.Errorf("%w supporting model: %s (channel pricing restriction)", ErrNoAvailableAccounts, requestedModel) + } + + // anthropic/gemini 分组支持混合调度(包含启用了 mixed_scheduling 的 antigravity 账户) + // 注意:强制平台模式不走混合调度 + if (platform == PlatformAnthropic || platform == PlatformGemini) && !hasForcePlatform { + account, err := s.selectAccountWithMixedScheduling(ctx, groupID, sessionHash, requestedModel, excludedIDs, platform) + if err != nil { + return nil, err + } + return s.hydrateSelectedAccount(ctx, account) + } + + // antigravity 分组、强制平台模式或无分组使用单平台选择 + // 注意:强制平台模式也必须遵守分组限制,不再回退到全平台查询 + account, err := s.selectAccountForModelWithPlatform(ctx, groupID, sessionHash, requestedModel, excludedIDs, platform) + if err != nil { + return nil, err + } + return s.hydrateSelectedAccount(ctx, account) +} + +// SelectAccountWithLoadAwareness selects account with load-awareness and wait plan. +// metadataUserID: 用于客户端亲和调度,从中提取客户端 ID +// sub2apiUserID: 系统用户 ID,用于二维亲和调度 +func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string, sub2apiUserID int64) (*AccountSelectionResult, error) { + // 调试日志:记录调度入口参数 + excludedIDsList := make([]int64, 0, len(excludedIDs)) + for id := range excludedIDs { + excludedIDsList = append(excludedIDsList, id) + } + slog.Debug("account_scheduling_starting", + "group_id", derefGroupID(groupID), + "model", requestedModel, + "session", shortSessionHash(sessionHash), + "excluded_ids", excludedIDsList) + + cfg := s.schedulingConfig() + + // 检查 Claude Code 客户端限制(可能会替换 groupID 为降级分组) + group, groupID, err := s.checkClaudeCodeRestriction(ctx, groupID) + if err != nil { + return nil, err + } + ctx = s.withGroupContext(ctx, group) + + // Claude Code 限制可能已将 groupID 解析为 fallback group, + // 渠道限制预检查必须使用解析后的分组。 + if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { + slog.Warn("channel pricing restriction blocked request", + "group_id", derefGroupID(groupID), + "model", requestedModel) + return nil, fmt.Errorf("%w supporting model: %s (channel pricing restriction)", ErrNoAvailableAccounts, requestedModel) + } + + var stickyAccountID int64 + var stickySource string + if prefetch := prefetchedStickyAccountIDFromContext(ctx, groupID); prefetch > 0 { + stickyAccountID = prefetch + stickySource = "prefetch" + } else if sessionHash != "" && s.cache != nil { + if accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash); err == nil { + stickyAccountID = accountID + stickySource = "cache" + } + } + + // [DEBUG-STICKY] 调度器入口日志 + slog.Info("sticky.scheduler_entry", + "group_id", derefGroupID(groupID), + "session_hash", shortSessionHash(sessionHash), + "sticky_account_id", stickyAccountID, + "sticky_source", stickySource, + "model", requestedModel, + "load_batch", cfg.LoadBatchEnabled, + "has_concurrency_svc", s.concurrencyService != nil, + "excluded_count", len(excludedIDs), + ) + + if s.debugModelRoutingEnabled() && requestedModel != "" { + groupPlatform := "" + if group != nil { + groupPlatform = group.Platform + } + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] select entry: group_id=%v group_platform=%s model=%s session=%s sticky_account=%d load_batch=%v concurrency=%v", + derefGroupID(groupID), groupPlatform, requestedModel, shortSessionHash(sessionHash), stickyAccountID, cfg.LoadBatchEnabled, s.concurrencyService != nil) + } + + if s.concurrencyService == nil || !cfg.LoadBatchEnabled { + // 复制排除列表,用于会话限制拒绝时的重试 + localExcluded := make(map[int64]struct{}) + for k, v := range excludedIDs { + localExcluded[k] = v + } + + for { + account, err := s.SelectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, localExcluded) + if err != nil { + return nil, err + } + + result, err := s.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) + if err == nil && result.Acquired { + // 获取槽位后检查会话限制(使用 sessionHash 作为会话标识符) + if !s.checkAndRegisterSession(ctx, account, sessionHash) { + result.ReleaseFunc() // 释放槽位 + localExcluded[account.ID] = struct{}{} // 排除此账号 + continue // 重新选择 + } + return s.newSelectionResult(ctx, account, true, result.ReleaseFunc, nil) + } + + // 对于等待计划的情况,也需要先检查会话限制 + if !s.checkAndRegisterSession(ctx, account, sessionHash) { + localExcluded[account.ID] = struct{}{} + continue + } + + if stickyAccountID > 0 && stickyAccountID == account.ID && s.concurrencyService != nil { + waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, account.ID) + if waitingCount < cfg.StickySessionMaxWaiting { + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) + } + } + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.FallbackWaitTimeout, + MaxWaiting: cfg.FallbackMaxWaiting, + }) + } + } + + platform, hasForcePlatform, err := s.resolvePlatform(ctx, groupID, group) + if err != nil { + return nil, err + } + preferOAuth := platform == PlatformGemini + if s.debugModelRoutingEnabled() && platform == PlatformAnthropic && requestedModel != "" { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] load-aware enabled: group_id=%v model=%s session=%s platform=%s", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), platform) + } + + accounts, useMixed, err := s.listSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) + if err != nil { + return nil, err + } + if len(accounts) == 0 { + return nil, ErrNoAvailableAccounts + } + ctx = s.withWindowCostPrefetch(ctx, accounts) + ctx = s.withRPMPrefetch(ctx, accounts) + + // 提前构建 accountByID(供 Layer 1 和 Layer 1.5 使用) + accountByID := make(map[int64]*Account, len(accounts)) + for i := range accounts { + accountByID[accounts[i].ID] = &accounts[i] + } + isExcluded := func(accountID int64) bool { + if excludedIDs == nil { + return false + } + _, excluded := excludedIDs[accountID] + return excluded + } + + // 获取模型路由配置(仅 anthropic 平台) + var routingAccountIDs []int64 + if group != nil && requestedModel != "" && group.Platform == PlatformAnthropic { + routingAccountIDs = group.GetRoutingAccountIDs(requestedModel) + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] context group routing: group_id=%d model=%s enabled=%v rules=%d matched_ids=%v session=%s sticky_account=%d", + group.ID, requestedModel, group.ModelRoutingEnabled, len(group.ModelRouting), routingAccountIDs, shortSessionHash(sessionHash), stickyAccountID) + if len(routingAccountIDs) == 0 && group.ModelRoutingEnabled && len(group.ModelRouting) > 0 { + keys := make([]string, 0, len(group.ModelRouting)) + for k := range group.ModelRouting { + keys = append(keys, k) + } + sort.Strings(keys) + const maxKeys = 20 + if len(keys) > maxKeys { + keys = keys[:maxKeys] + } + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] context group routing miss: group_id=%d model=%s patterns(sample)=%v", group.ID, requestedModel, keys) + } + } + } + + // ============ Layer 1: 模型路由优先选择(优先级高于粘性会话) ============ + if len(routingAccountIDs) > 0 && s.concurrencyService != nil { + // 1. 过滤出路由列表中可调度的账号 + var routingCandidates []*Account + var filteredExcluded, filteredMissing, filteredUnsched, filteredPlatform, filteredModelScope, filteredModelMapping, filteredWindowCost int + var modelScopeSkippedIDs []int64 // 记录因模型限流被跳过的账号 ID + for _, routingAccountID := range routingAccountIDs { + if isExcluded(routingAccountID) { + filteredExcluded++ + continue + } + account, ok := accountByID[routingAccountID] + if !ok || !s.isAccountSchedulableForSelection(account) { + if !ok { + filteredMissing++ + } else { + filteredUnsched++ + } + continue + } + if !s.isAccountAllowedForPlatform(account, platform, useMixed) { + filteredPlatform++ + continue + } + if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, account, requestedModel) { + filteredModelMapping++ + continue + } + if !s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) { + filteredModelScope++ + modelScopeSkippedIDs = append(modelScopeSkippedIDs, account.ID) + continue + } + // 配额检查 + if !s.isAccountSchedulableForQuota(account) { + continue + } + // 窗口费用检查(非粘性会话路径) + if !s.isAccountSchedulableForWindowCost(ctx, account, false) { + filteredWindowCost++ + continue + } + // RPM 检查(非粘性会话路径) + if !s.isAccountSchedulableForRPM(ctx, account, false) { + continue + } + routingCandidates = append(routingCandidates, account) + } + + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed candidates: group_id=%v model=%s routed=%d candidates=%d filtered(excluded=%d missing=%d unsched=%d platform=%d model_scope=%d model_mapping=%d window_cost=%d)", + derefGroupID(groupID), requestedModel, len(routingAccountIDs), len(routingCandidates), + filteredExcluded, filteredMissing, filteredUnsched, filteredPlatform, filteredModelScope, filteredModelMapping, filteredWindowCost) + if len(modelScopeSkippedIDs) > 0 { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] model_rate_limited accounts skipped: group_id=%v model=%s account_ids=%v", + derefGroupID(groupID), requestedModel, modelScopeSkippedIDs) + } + } + + if len(routingCandidates) > 0 { + // 1.5. 在路由账号范围内检查粘性会话 + if sessionHash != "" && stickyAccountID > 0 { + slog.Debug("sticky.layer1_5_checking", + "sticky_account_id", stickyAccountID, + "in_routing_list", containsInt64(routingAccountIDs, stickyAccountID), + "is_excluded", isExcluded(stickyAccountID), + "in_account_map", func() bool { _, ok := accountByID[stickyAccountID]; return ok }(), + "session", shortSessionHash(sessionHash), + ) + if containsInt64(routingAccountIDs, stickyAccountID) && !isExcluded(stickyAccountID) { + // 粘性账号在路由列表中,优先使用 + if stickyAccount, ok := accountByID[stickyAccountID]; ok { + var stickyCacheMissReason string + + gatePass := s.isAccountSchedulableForSelection(stickyAccount) && + s.isAccountAllowedForPlatform(stickyAccount, platform, useMixed) && + (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, stickyAccount, requestedModel)) && + s.isAccountSchedulableForModelSelection(ctx, stickyAccount, requestedModel) && + s.isAccountSchedulableForQuota(stickyAccount) && + s.isAccountSchedulableForWindowCost(ctx, stickyAccount, true) + + rpmPass := gatePass && s.isAccountSchedulableForRPM(ctx, stickyAccount, true) + + if rpmPass { // 粘性会话窗口费用+RPM 检查 + result, err := s.tryAcquireAccountSlot(ctx, stickyAccountID, stickyAccount.Concurrency) + if err == nil && result.Acquired { + // 会话数量限制检查 + if !s.checkAndRegisterSession(ctx, stickyAccount, sessionHash) { + result.ReleaseFunc() // 释放槽位 + stickyCacheMissReason = "session_limit" + // 继续到负载感知选择 + } else { + slog.Debug("sticky.layer1_5_hit", + "account_id", stickyAccountID, + "session", shortSessionHash(sessionHash), + "result", "slot_acquired", + ) + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), stickyAccountID) + } + return s.newSelectionResult(ctx, stickyAccount, true, result.ReleaseFunc, nil) + } + } + + if stickyCacheMissReason == "" { + waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, stickyAccountID) + if waitingCount < cfg.StickySessionMaxWaiting { + // 会话数量限制检查(等待计划也需要占用会话配额) + if !s.checkAndRegisterSession(ctx, stickyAccount, sessionHash) { + stickyCacheMissReason = "session_limit" + // 会话限制已满,继续到负载感知选择 + } else { + return &AccountSelectionResult{ + Account: stickyAccount, + WaitPlan: &AccountWaitPlan{ + AccountID: stickyAccountID, + MaxConcurrency: stickyAccount.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }, + }, nil + } + } else { + stickyCacheMissReason = "wait_queue_full" + } + } + // 粘性账号槽位满且等待队列已满,继续使用负载感知选择 + } else if !gatePass { + stickyCacheMissReason = "gate_check" + } else { + stickyCacheMissReason = "rpm_red" + } + + // 记录粘性缓存未命中的结构化日志 + if stickyCacheMissReason != "" { + baseRPM := stickyAccount.GetBaseRPM() + var currentRPM int + if count, ok := rpmFromPrefetchContext(ctx, stickyAccount.ID); ok { + currentRPM = count + } + logger.LegacyPrintf("service.gateway", "[StickyCacheMiss] reason=%s account_id=%d session=%s current_rpm=%d base_rpm=%d", + stickyCacheMissReason, stickyAccountID, shortSessionHash(sessionHash), currentRPM, baseRPM) + } + } else { + _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + logger.LegacyPrintf("service.gateway", "[StickyCacheMiss] reason=account_cleared account_id=%d session=%s current_rpm=0 base_rpm=0", + stickyAccountID, shortSessionHash(sessionHash)) + } + } + } + + // 2. 批量获取负载信息 + routingLoads := make([]AccountWithConcurrency, 0, len(routingCandidates)) + for _, acc := range routingCandidates { + routingLoads = append(routingLoads, AccountWithConcurrency{ + ID: acc.ID, + MaxConcurrency: acc.EffectiveLoadFactor(), + }) + } + routingLoadMap, _ := s.concurrencyService.GetAccountsLoadBatch(ctx, routingLoads) + + // 3. 按负载感知排序 + var routingAvailable []accountWithLoad + for _, acc := range routingCandidates { + loadInfo := routingLoadMap[acc.ID] + if loadInfo == nil { + loadInfo = &AccountLoadInfo{AccountID: acc.ID} + } + if loadInfo.LoadRate < 100 { + routingAvailable = append(routingAvailable, accountWithLoad{account: acc, loadInfo: loadInfo}) + } + } + + if len(routingAvailable) > 0 { + // 排序:优先级 > 负载率 > 最后使用时间 + sort.SliceStable(routingAvailable, func(i, j int) bool { + a, b := routingAvailable[i], routingAvailable[j] + if a.account.Priority != b.account.Priority { + return a.account.Priority < b.account.Priority + } + if a.loadInfo.LoadRate != b.loadInfo.LoadRate { + return a.loadInfo.LoadRate < b.loadInfo.LoadRate + } + switch { + case a.account.LastUsedAt == nil && b.account.LastUsedAt != nil: + return true + case a.account.LastUsedAt != nil && b.account.LastUsedAt == nil: + return false + case a.account.LastUsedAt == nil && b.account.LastUsedAt == nil: + return false + default: + return a.account.LastUsedAt.Before(*b.account.LastUsedAt) + } + }) + shuffleWithinSortGroups(routingAvailable) + + // 4. 尝试获取槽位 + for _, item := range routingAvailable { + result, err := s.tryAcquireAccountSlot(ctx, item.account.ID, item.account.Concurrency) + if err == nil && result.Acquired { + // 会话数量限制检查 + if !s.checkAndRegisterSession(ctx, item.account, sessionHash) { + result.ReleaseFunc() // 释放槽位,继续尝试下一个账号 + continue + } + if sessionHash != "" && s.cache != nil { + _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, item.account.ID, stickySessionTTL) + } + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) + } + return s.newSelectionResult(ctx, item.account, true, result.ReleaseFunc, nil) + } + } + + // 5. 所有路由账号槽位满,尝试返回等待计划(选择负载最低的) + // 遍历找到第一个满足会话限制的账号 + for _, item := range routingAvailable { + if !s.checkAndRegisterSession(ctx, item.account, sessionHash) { + continue // 会话限制已满,尝试下一个 + } + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed wait: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) + } + return s.newSelectionResult(ctx, item.account, false, nil, &AccountWaitPlan{ + AccountID: item.account.ID, + MaxConcurrency: item.account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) + } + // 所有路由账号会话限制都已满,继续到 Layer 2 回退 + } + // 路由列表中的账号都不可用(负载率 >= 100),继续到 Layer 2 回退 + logger.LegacyPrintf("service.gateway", "[ModelRouting] All routed accounts unavailable for model=%s, falling back to normal selection", requestedModel) + } + } + + // ============ Layer 1.5: 粘性会话(仅在无模型路由配置时生效) ============ + if len(routingAccountIDs) == 0 && sessionHash != "" && stickyAccountID > 0 && !isExcluded(stickyAccountID) { + accountID := stickyAccountID + if accountID > 0 && !isExcluded(accountID) { + account, ok := accountByID[accountID] + if ok { + // 检查账户是否需要清理粘性会话绑定 + clearSticky := shouldClearStickySession(account, requestedModel) + if clearSticky { + slog.Debug("sticky.layer1_5_no_routing_clear", + "account_id", accountID, + "reason", "should_clear_sticky_session", + "session", shortSessionHash(sessionHash), + ) + _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + } + + // 注意:不再检查 isAccountInGroup,因为 accountByID 已经从按分组过滤的 + // accounts 列表构建,账号一定在分组内。而 scheduler snapshot 缓存 + // 反序列化后 AccountGroups 字段为空,导致 isAccountInGroup 永远返回 false。 + platformOK := s.isAccountAllowedForPlatform(account, platform, useMixed) + modelSupported := requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel) + modelSchedulable := s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) + quotaOK := s.isAccountSchedulableForQuota(account) + windowCostOK := s.isAccountSchedulableForWindowCost(ctx, account, true) + rpmOK := s.isAccountSchedulableForRPM(ctx, account, true) + schedulable := s.isAccountSchedulableForSelection(account) + + slog.Debug("sticky.layer1_5_no_routing_checks", + "account_id", accountID, + "session", shortSessionHash(sessionHash), + "clear_sticky", clearSticky, + "schedulable", schedulable, + "platform_ok", platformOK, + "model_supported", modelSupported, + "model_schedulable", modelSchedulable, + "quota_ok", quotaOK, + "window_cost_ok", windowCostOK, + "rpm_ok", rpmOK, + ) + + if !clearSticky && platformOK && modelSupported && modelSchedulable && quotaOK && windowCostOK && rpmOK && schedulable { + result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency) + if err == nil && result.Acquired { + // 会话数量限制检查 + if !s.checkAndRegisterSession(ctx, account, sessionHash) { + result.ReleaseFunc() // 释放槽位,继续到 Layer 2 + slog.Debug("sticky.layer1_5_no_routing_miss", + "account_id", accountID, + "reason", "session_limit", + "session", shortSessionHash(sessionHash), + ) + } else { + slog.Debug("sticky.layer1_5_no_routing_hit", + "account_id", accountID, + "session", shortSessionHash(sessionHash), + "result", "slot_acquired", + ) + if s.cache != nil { + _ = s.cache.RefreshSessionTTL(ctx, derefGroupID(groupID), sessionHash, stickySessionTTL) + } + return s.newSelectionResult(ctx, account, true, result.ReleaseFunc, nil) + } + } else { + slog.Debug("sticky.layer1_5_no_routing_slot_busy", + "account_id", accountID, + "session", shortSessionHash(sessionHash), + ) + } + + waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, accountID) + if waitingCount < cfg.StickySessionMaxWaiting { + // 会话数量限制检查(等待计划也需要占用会话配额) + if !s.checkAndRegisterSession(ctx, account, sessionHash) { + // 会话限制已满,继续到 Layer 2 + } else { + slog.Debug("sticky.layer1_5_no_routing_hit", + "account_id", accountID, + "session", shortSessionHash(sessionHash), + "result", "wait_plan", + ) + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: accountID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) + } + } + } else if !clearSticky { + slog.Debug("sticky.layer1_5_no_routing_miss", + "account_id", accountID, + "reason", "gate_check_failed", + "session", shortSessionHash(sessionHash), + ) + } + } else { + slog.Debug("sticky.layer1_5_no_routing_miss", + "account_id", accountID, + "reason", "account_not_in_map", + "session", shortSessionHash(sessionHash), + ) + } + } + } else if len(routingAccountIDs) == 0 && sessionHash != "" { + slog.Debug("sticky.layer1_5_no_routing_skip", + "sticky_account_id", stickyAccountID, + "is_excluded", func() bool { return stickyAccountID > 0 && isExcluded(stickyAccountID) }(), + "session", shortSessionHash(sessionHash), + "reason", func() string { + if stickyAccountID == 0 { + return "no_sticky_binding" + } + return "sticky_account_excluded" + }(), + ) + } + + // ============ Layer 2: 负载感知选择 ============ + slog.Debug("sticky.layer2_fallback", + "session", shortSessionHash(sessionHash), + "sticky_account_id", stickyAccountID, + "reason", "sticky_not_used_falling_back_to_load_balance", + "total_accounts", len(accounts), + ) + candidates := make([]*Account, 0, len(accounts)) + for i := range accounts { + acc := &accounts[i] + if isExcluded(acc.ID) { + continue + } + // Scheduler snapshots can be temporarily stale (bucket rebuild is throttled); + // re-check schedulability here so recently rate-limited/overloaded accounts + // are not selected again before the bucket is rebuilt. + if !s.isAccountSchedulableForSelection(acc) { + continue + } + if !s.isAccountAllowedForPlatform(acc, platform, useMixed) { + continue + } + if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { + continue + } + // 配额检查 + if !s.isAccountSchedulableForQuota(acc) { + continue + } + // 窗口费用检查(非粘性会话路径) + if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { + continue + } + // RPM 检查(非粘性会话路径) + if !s.isAccountSchedulableForRPM(ctx, acc, false) { + continue + } + candidates = append(candidates, acc) + } + + if len(candidates) == 0 { + return nil, ErrNoAvailableAccounts + } + + accountLoads := make([]AccountWithConcurrency, 0, len(candidates)) + for _, acc := range candidates { + accountLoads = append(accountLoads, AccountWithConcurrency{ + ID: acc.ID, + MaxConcurrency: acc.EffectiveLoadFactor(), + }) + } + + loadMap, err := s.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) + if err != nil { + if result, ok, legacyErr := s.tryAcquireByLegacyOrder(ctx, candidates, groupID, sessionHash, preferOAuth); legacyErr != nil { + return nil, legacyErr + } else if ok { + return result, nil + } + } else { + var available []accountWithLoad + for _, acc := range candidates { + loadInfo := loadMap[acc.ID] + if loadInfo == nil { + loadInfo = &AccountLoadInfo{AccountID: acc.ID} + } + if loadInfo.LoadRate < 100 { + available = append(available, accountWithLoad{ + account: acc, + loadInfo: loadInfo, + }) + } + } + + // 分层过滤选择:优先级 →(可选)最早重置 → 负载率 → LRU + for len(available) > 0 { + // 1. 取优先级最小的集合 + candidates := filterByMinPriority(available) + // 2. (可选)use-it-or-lose-it:优先选用会话窗口最早重置的账号 + if cfg.PreferSoonestReset { + candidates = filterBySoonestReset(candidates) + } + // 3. 取负载率最低的集合 + candidates = filterByMinLoadRate(candidates) + // 4. LRU 选择最久未用的账号 + selected := selectByLRU(candidates, preferOAuth) + if selected == nil { + break + } + + result, err := s.tryAcquireAccountSlot(ctx, selected.account.ID, selected.account.Concurrency) + if err == nil && result.Acquired { + // 会话数量限制检查 + if !s.checkAndRegisterSession(ctx, selected.account, sessionHash) { + result.ReleaseFunc() // 释放槽位,继续尝试下一个账号 + } else { + if sessionHash != "" && s.cache != nil { + _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.account.ID, stickySessionTTL) + } + return s.newSelectionResult(ctx, selected.account, true, result.ReleaseFunc, nil) + } + } + + // 移除已尝试的账号,重新进行分层过滤 + selectedID := selected.account.ID + newAvailable := make([]accountWithLoad, 0, len(available)-1) + for _, acc := range available { + if acc.account.ID != selectedID { + newAvailable = append(newAvailable, acc) + } + } + available = newAvailable + } + } + + // ============ Layer 3: 兜底排队 ============ + s.sortCandidatesForFallback(candidates, preferOAuth, cfg.FallbackSelectionMode) + for _, acc := range candidates { + // 会话数量限制检查(等待计划也需要占用会话配额) + if !s.checkAndRegisterSession(ctx, acc, sessionHash) { + continue // 会话限制已满,尝试下一个账号 + } + return s.newSelectionResult(ctx, acc, false, nil, &AccountWaitPlan{ + AccountID: acc.ID, + MaxConcurrency: acc.Concurrency, + Timeout: cfg.FallbackWaitTimeout, + MaxWaiting: cfg.FallbackMaxWaiting, + }) + } + return nil, ErrNoAvailableAccounts +} + +func (s *GatewayService) tryAcquireByLegacyOrder(ctx context.Context, candidates []*Account, groupID *int64, sessionHash string, preferOAuth bool) (*AccountSelectionResult, bool, error) { + ordered := append([]*Account(nil), candidates...) + sortAccountsByPriorityAndLastUsed(ordered, preferOAuth) + + for _, acc := range ordered { + result, err := s.tryAcquireAccountSlot(ctx, acc.ID, acc.Concurrency) + if err == nil && result.Acquired { + // 会话数量限制检查 + if !s.checkAndRegisterSession(ctx, acc, sessionHash) { + result.ReleaseFunc() // 释放槽位,继续尝试下一个账号 + continue + } + if sessionHash != "" && s.cache != nil { + _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, acc.ID, stickySessionTTL) + } + selection, err := s.newSelectionResult(ctx, acc, true, result.ReleaseFunc, nil) + if err != nil { + return nil, false, err + } + return selection, true, nil + } + } + + return nil, false, nil +} + +func (s *GatewayService) schedulingConfig() config.GatewaySchedulingConfig { + if s.cfg != nil { + return s.cfg.Gateway.Scheduling + } + return config.GatewaySchedulingConfig{ + StickySessionMaxWaiting: 3, + StickySessionWaitTimeout: 45 * time.Second, + FallbackWaitTimeout: 30 * time.Second, + FallbackMaxWaiting: 100, + LoadBatchEnabled: true, + SlotCleanupInterval: 30 * time.Second, + } +} + +func (s *GatewayService) withGroupContext(ctx context.Context, group *Group) context.Context { + if !IsGroupContextValid(group) { + return ctx + } + if existing, ok := ctx.Value(ctxkey.Group).(*Group); ok && existing != nil && existing.ID == group.ID && IsGroupContextValid(existing) { + return ctx + } + return context.WithValue(ctx, ctxkey.Group, group) +} + +func (s *GatewayService) groupFromContext(ctx context.Context, groupID int64) *Group { + if group, ok := ctx.Value(ctxkey.Group).(*Group); ok && IsGroupContextValid(group) && group.ID == groupID { + return group + } + return nil +} + +func (s *GatewayService) resolveGroupByID(ctx context.Context, groupID int64) (*Group, error) { + if group := s.groupFromContext(ctx, groupID); group != nil { + return group, nil + } + group, err := s.groupRepo.GetByIDLite(ctx, groupID) + if err != nil { + return nil, fmt.Errorf("get group failed: %w", err) + } + return group, nil +} + +func (s *GatewayService) ResolveGroupByID(ctx context.Context, groupID int64) (*Group, error) { + return s.resolveGroupByID(ctx, groupID) +} + +func (s *GatewayService) routingAccountIDsForRequest(ctx context.Context, groupID *int64, requestedModel string, platform string) []int64 { + if groupID == nil || requestedModel == "" || platform != PlatformAnthropic { + return nil + } + group, err := s.resolveGroupByID(ctx, *groupID) + if err != nil || group == nil { + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] resolve group failed: group_id=%v model=%s platform=%s err=%v", derefGroupID(groupID), requestedModel, platform, err) + } + return nil + } + // Preserve existing behavior: model routing only applies to anthropic groups. + if group.Platform != PlatformAnthropic { + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] skip: non-anthropic group platform: group_id=%d group_platform=%s model=%s", group.ID, group.Platform, requestedModel) + } + return nil + } + ids := group.GetRoutingAccountIDs(requestedModel) + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routing lookup: group_id=%d model=%s enabled=%v rules=%d matched_ids=%v", + group.ID, requestedModel, group.ModelRoutingEnabled, len(group.ModelRouting), ids) + } + return ids +} + +func (s *GatewayService) resolveGatewayGroup(ctx context.Context, groupID *int64) (*Group, *int64, error) { + if groupID == nil { + return nil, nil, nil + } + + currentID := *groupID + visited := map[int64]struct{}{} + for { + if _, seen := visited[currentID]; seen { + return nil, nil, fmt.Errorf("fallback group cycle detected") + } + visited[currentID] = struct{}{} + + group, err := s.resolveGroupByID(ctx, currentID) + if err != nil { + return nil, nil, err + } + + if !group.ClaudeCodeOnly || IsClaudeCodeClient(ctx) { + return group, ¤tID, nil + } + + if group.FallbackGroupID == nil { + return nil, nil, ErrClaudeCodeOnly + } + currentID = *group.FallbackGroupID + } +} + +// checkClaudeCodeRestriction 检查分组的 Claude Code 客户端限制 +// 如果分组启用了 claude_code_only 且请求不是来自 Claude Code 客户端: +// - 有降级分组:返回降级分组的 ID +// - 无降级分组:返回 ErrClaudeCodeOnly 错误 +func (s *GatewayService) checkClaudeCodeRestriction(ctx context.Context, groupID *int64) (*Group, *int64, error) { + if groupID == nil { + return nil, groupID, nil + } + + // 强制平台模式不检查 Claude Code 限制 + if forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string); hasForcePlatform && forcePlatform != "" { + return nil, groupID, nil + } + + group, resolvedID, err := s.resolveGatewayGroup(ctx, groupID) + if err != nil { + return nil, nil, err + } + + return group, resolvedID, nil +} + +func (s *GatewayService) resolvePlatform(ctx context.Context, groupID *int64, group *Group) (string, bool, error) { + forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string) + if hasForcePlatform && forcePlatform != "" { + return forcePlatform, true, nil + } + if group != nil { + return group.Platform, false, nil + } + if groupID != nil { + group, err := s.resolveGroupByID(ctx, *groupID) + if err != nil { + return "", false, err + } + return group.Platform, false, nil + } + return PlatformAnthropic, false, nil +} + +func (s *GatewayService) listSchedulableAccounts(ctx context.Context, groupID *int64, platform string, hasForcePlatform bool) ([]Account, bool, error) { + if s.schedulerSnapshot != nil { + accounts, useMixed, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) + if err == nil { + slog.Debug("account_scheduling_list_snapshot", + "group_id", derefGroupID(groupID), + "platform", platform, + "use_mixed", useMixed, + "count", len(accounts)) + if slog.Default().Enabled(ctx, slog.LevelDebug) { + for _, acc := range accounts { + slog.Debug("account_scheduling_account_detail", + "account_id", acc.ID, + "name", acc.Name, + "platform", acc.Platform, + "type", acc.Type, + "status", acc.Status, + "tls_fingerprint", acc.IsTLSFingerprintEnabled()) + } + } + } + return accounts, useMixed, err + } + useMixed := (platform == PlatformAnthropic || platform == PlatformGemini) && !hasForcePlatform + if useMixed { + platforms := []string{platform, PlatformAntigravity} + var accounts []Account + var err error + if groupID != nil { + accounts, err = s.accountRepo.ListSchedulableByGroupIDAndPlatforms(ctx, *groupID, platforms) + } else if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { + accounts, err = s.accountRepo.ListSchedulableByPlatforms(ctx, platforms) + } else { + accounts, err = s.accountRepo.ListSchedulableUngroupedByPlatforms(ctx, platforms) + } + if err != nil { + slog.Debug("account_scheduling_list_failed", + "group_id", derefGroupID(groupID), + "platform", platform, + "error", err) + return nil, useMixed, err + } + filtered := make([]Account, 0, len(accounts)) + for _, acc := range accounts { + if acc.Platform == PlatformAntigravity && !acc.IsMixedSchedulingEnabled() { + continue + } + filtered = append(filtered, acc) + } + slog.Debug("account_scheduling_list_mixed", + "group_id", derefGroupID(groupID), + "platform", platform, + "raw_count", len(accounts), + "filtered_count", len(filtered)) + if slog.Default().Enabled(ctx, slog.LevelDebug) { + for _, acc := range filtered { + slog.Debug("account_scheduling_account_detail", + "account_id", acc.ID, + "name", acc.Name, + "platform", acc.Platform, + "type", acc.Type, + "status", acc.Status, + "tls_fingerprint", acc.IsTLSFingerprintEnabled()) + } + } + return filtered, useMixed, nil + } + + var accounts []Account + var err error + if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { + accounts, err = s.accountRepo.ListSchedulableByPlatform(ctx, platform) + } else if groupID != nil { + accounts, err = s.accountRepo.ListSchedulableByGroupIDAndPlatform(ctx, *groupID, platform) + // 分组内无账号则返回空列表,由上层处理错误,不再回退到全平台查询 + } else { + accounts, err = s.accountRepo.ListSchedulableUngroupedByPlatform(ctx, platform) + } + if err != nil { + slog.Debug("account_scheduling_list_failed", + "group_id", derefGroupID(groupID), + "platform", platform, + "error", err) + return nil, useMixed, err + } + slog.Debug("account_scheduling_list_single", + "group_id", derefGroupID(groupID), + "platform", platform, + "count", len(accounts)) + if slog.Default().Enabled(ctx, slog.LevelDebug) { + for _, acc := range accounts { + slog.Debug("account_scheduling_account_detail", + "account_id", acc.ID, + "name", acc.Name, + "platform", acc.Platform, + "type", acc.Type, + "status", acc.Status, + "tls_fingerprint", acc.IsTLSFingerprintEnabled()) + } + } + return accounts, useMixed, nil +} + +// IsSingleAntigravityAccountGroup 检查指定分组是否只有一个 antigravity 平台的可调度账号。 +// 用于 Handler 层在首次请求时提前设置 SingleAccountRetry context, +// 避免单账号分组收到 503 时错误地设置模型限流标记导致后续请求连续快速失败。 +func (s *GatewayService) IsSingleAntigravityAccountGroup(ctx context.Context, groupID *int64) bool { + accounts, _, err := s.listSchedulableAccounts(ctx, groupID, PlatformAntigravity, true) + if err != nil { + return false + } + return len(accounts) == 1 +} + +func (s *GatewayService) isAccountAllowedForPlatform(account *Account, platform string, useMixed bool) bool { + if account == nil { + return false + } + if useMixed { + if account.Platform == platform { + return true + } + return account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled() + } + return account.Platform == platform +} + +func (s *GatewayService) isAccountSchedulableForSelection(account *Account) bool { + if account == nil { + return false + } + return account.IsSchedulable() +} + +func (s *GatewayService) isAccountSchedulableForModelSelection(ctx context.Context, account *Account, requestedModel string) bool { + if account == nil { + return false + } + return account.IsSchedulableForModelWithContext(ctx, requestedModel) +} + +// isAccountInGroup checks if the account belongs to the specified group. +// When groupID is nil, returns true only for ungrouped accounts (no group assignments). +func (s *GatewayService) isAccountInGroup(account *Account, groupID *int64) bool { + if account == nil { + return false + } + if groupID == nil { + // 无分组的 API Key 只能使用未分组的账号 + return len(account.AccountGroups) == 0 + } + for _, ag := range account.AccountGroups { + if ag.GroupID == *groupID { + return true + } + } + return false +} + +func (s *GatewayService) tryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (*AcquireResult, error) { + if s.concurrencyService == nil { + return &AcquireResult{Acquired: true, ReleaseFunc: func() {}}, nil + } + return s.concurrencyService.AcquireAccountSlot(ctx, accountID, maxConcurrency) +} + +type usageLogWindowStatsBatchProvider interface { + GetAccountWindowStatsBatch(ctx context.Context, accountIDs []int64, startTime time.Time) (map[int64]*usagestats.AccountStats, error) +} + +type windowCostPrefetchContextKeyType struct{} + +var windowCostPrefetchContextKey = windowCostPrefetchContextKeyType{} + +func windowCostFromPrefetchContext(ctx context.Context, accountID int64) (float64, bool) { + if ctx == nil || accountID <= 0 { + return 0, false + } + m, ok := ctx.Value(windowCostPrefetchContextKey).(map[int64]float64) + if !ok || len(m) == 0 { + return 0, false + } + v, exists := m[accountID] + return v, exists +} + +func (s *GatewayService) withWindowCostPrefetch(ctx context.Context, accounts []Account) context.Context { + if ctx == nil || len(accounts) == 0 || s.sessionLimitCache == nil || s.usageLogRepo == nil { + return ctx + } + + accountByID := make(map[int64]*Account) + accountIDs := make([]int64, 0, len(accounts)) + for i := range accounts { + account := &accounts[i] + if account == nil || !account.IsAnthropicOAuthOrSetupToken() { + continue + } + if account.GetWindowCostLimit() <= 0 { + continue + } + accountByID[account.ID] = account + accountIDs = append(accountIDs, account.ID) + } + if len(accountIDs) == 0 { + return ctx + } + + costs := make(map[int64]float64, len(accountIDs)) + cacheValues, err := s.sessionLimitCache.GetWindowCostBatch(ctx, accountIDs) + if err == nil { + for accountID, cost := range cacheValues { + costs[accountID] = cost + } + windowCostPrefetchCacheHitTotal.Add(int64(len(cacheValues))) + } else { + windowCostPrefetchErrorTotal.Add(1) + logger.LegacyPrintf("service.gateway", "window_cost batch cache read failed: %v", err) + } + cacheMissCount := len(accountIDs) - len(costs) + if cacheMissCount < 0 { + cacheMissCount = 0 + } + windowCostPrefetchCacheMissTotal.Add(int64(cacheMissCount)) + + missingByStart := make(map[int64][]int64) + startTimes := make(map[int64]time.Time) + for _, accountID := range accountIDs { + if _, ok := costs[accountID]; ok { + continue + } + account := accountByID[accountID] + if account == nil { + continue + } + startTime := account.GetCurrentWindowStartTime() + startKey := startTime.Unix() + missingByStart[startKey] = append(missingByStart[startKey], accountID) + startTimes[startKey] = startTime + } + if len(missingByStart) == 0 { + return context.WithValue(ctx, windowCostPrefetchContextKey, costs) + } + + batchReader, hasBatch := s.usageLogRepo.(usageLogWindowStatsBatchProvider) + for startKey, ids := range missingByStart { + startTime := startTimes[startKey] + + if hasBatch { + windowCostPrefetchBatchSQLTotal.Add(1) + queryStart := time.Now() + statsByAccount, err := batchReader.GetAccountWindowStatsBatch(ctx, ids, startTime) + if err == nil { + slog.Debug("window_cost_batch_query_ok", + "accounts", len(ids), + "window_start", startTime.Format(time.RFC3339), + "duration_ms", time.Since(queryStart).Milliseconds()) + for _, accountID := range ids { + stats := statsByAccount[accountID] + cost := 0.0 + if stats != nil { + cost = stats.StandardCost + } + costs[accountID] = cost + _ = s.sessionLimitCache.SetWindowCost(ctx, accountID, cost) + } + continue + } + windowCostPrefetchErrorTotal.Add(1) + logger.LegacyPrintf("service.gateway", "window_cost batch db query failed: start=%s err=%v", startTime.Format(time.RFC3339), err) + } + + // 回退路径:缺少批量仓储能力或批量查询失败时,按账号单查(失败开放)。 + windowCostPrefetchFallbackTotal.Add(int64(len(ids))) + for _, accountID := range ids { + stats, err := s.usageLogRepo.GetAccountWindowStats(ctx, accountID, startTime) + if err != nil { + windowCostPrefetchErrorTotal.Add(1) + continue + } + cost := stats.StandardCost + costs[accountID] = cost + _ = s.sessionLimitCache.SetWindowCost(ctx, accountID, cost) + } + } + + return context.WithValue(ctx, windowCostPrefetchContextKey, costs) +} + +// isAccountSchedulableForQuota 检查账号是否在配额限制内 +// 适用于配置了 quota_limit 的 apikey 和 bedrock 类型账号 +func (s *GatewayService) isAccountSchedulableForQuota(account *Account) bool { + if !account.IsAPIKeyOrBedrock() { + return true + } + return !account.IsQuotaExceeded() +} + +// isAccountSchedulableForWindowCost 检查账号是否可根据窗口费用进行调度 +// 仅适用于 Anthropic OAuth/SetupToken 账号 +// 返回 true 表示可调度,false 表示不可调度 +func (s *GatewayService) isAccountSchedulableForWindowCost(ctx context.Context, account *Account, isSticky bool) bool { + // 只检查 Anthropic OAuth/SetupToken 账号 + if !account.IsAnthropicOAuthOrSetupToken() { + return true + } + + limit := account.GetWindowCostLimit() + if limit <= 0 { + return true // 未启用窗口费用限制 + } + + // 尝试从缓存获取窗口费用 + var currentCost float64 + if cost, ok := windowCostFromPrefetchContext(ctx, account.ID); ok { + currentCost = cost + goto checkSchedulability + } + if s.sessionLimitCache != nil { + if cost, hit, err := s.sessionLimitCache.GetWindowCost(ctx, account.ID); err == nil && hit { + currentCost = cost + goto checkSchedulability + } + } + + // 缓存未命中,从数据库查询 + { + // 使用统一的窗口开始时间计算逻辑(考虑窗口过期情况) + startTime := account.GetCurrentWindowStartTime() + + stats, err := s.usageLogRepo.GetAccountWindowStats(ctx, account.ID, startTime) + if err != nil { + // 失败开放:查询失败时允许调度 + return true + } + + // 使用标准费用(不含账号倍率) + currentCost = stats.StandardCost + + // 设置缓存(忽略错误) + if s.sessionLimitCache != nil { + _ = s.sessionLimitCache.SetWindowCost(ctx, account.ID, currentCost) + } + } + +checkSchedulability: + schedulability := account.CheckWindowCostSchedulability(currentCost) + + switch schedulability { + case WindowCostSchedulable: + return true + case WindowCostStickyOnly: + return isSticky + case WindowCostNotSchedulable: + return false + } + return true +} + +// rpmPrefetchContextKey is the context key for prefetched RPM counts. +type rpmPrefetchContextKeyType struct{} + +var rpmPrefetchContextKey = rpmPrefetchContextKeyType{} + +func rpmFromPrefetchContext(ctx context.Context, accountID int64) (int, bool) { + if v, ok := ctx.Value(rpmPrefetchContextKey).(map[int64]int); ok { + count, found := v[accountID] + return count, found + } + return 0, false +} + +// withRPMPrefetch 批量预取所有候选账号的 RPM 计数 +func (s *GatewayService) withRPMPrefetch(ctx context.Context, accounts []Account) context.Context { + if s.rpmCache == nil { + return ctx + } + + var ids []int64 + for i := range accounts { + if accounts[i].IsAnthropicOAuthOrSetupToken() && accounts[i].GetBaseRPM() > 0 { + ids = append(ids, accounts[i].ID) + } + } + if len(ids) == 0 { + return ctx + } + + counts, err := s.rpmCache.GetRPMBatch(ctx, ids) + if err != nil { + return ctx // 失败开放 + } + return context.WithValue(ctx, rpmPrefetchContextKey, counts) +} + +// isAccountSchedulableForRPM 检查账号是否可根据 RPM 进行调度 +// 仅适用于 Anthropic OAuth/SetupToken 账号 +func (s *GatewayService) isAccountSchedulableForRPM(ctx context.Context, account *Account, isSticky bool) bool { + if !account.IsAnthropicOAuthOrSetupToken() { + return true + } + baseRPM := account.GetBaseRPM() + if baseRPM <= 0 { + return true + } + + // 尝试从预取缓存获取 + var currentRPM int + if count, ok := rpmFromPrefetchContext(ctx, account.ID); ok { + currentRPM = count + } else if s.rpmCache != nil { + if count, err := s.rpmCache.GetRPM(ctx, account.ID); err == nil { + currentRPM = count + } + // 失败开放:GetRPM 错误时允许调度 + } + + schedulability := account.CheckRPMSchedulability(currentRPM) + switch schedulability { + case WindowCostSchedulable: + return true + case WindowCostStickyOnly: + return isSticky + case WindowCostNotSchedulable: + return false + } + return true +} + +// IncrementAccountRPM increments the RPM counter for the given account. +// 已知 TOCTOU 竞态:调度时读取 RPM 计数与此处递增之间存在时间窗口, +// 高并发下可能短暂超出 RPM 限制。这是与 WindowCost 一致的 soft-limit +// 设计权衡——可接受的少量超额优于加锁带来的延迟和复杂度。 +func (s *GatewayService) IncrementAccountRPM(ctx context.Context, accountID int64) error { + if s.rpmCache == nil { + return nil + } + _, err := s.rpmCache.IncrementRPM(ctx, accountID) + return err +} + +// checkAndRegisterSession 检查并注册会话,用于会话数量限制 +// 仅适用于 Anthropic OAuth/SetupToken 账号 +// sessionID: 会话标识符(使用粘性会话的 hash) +// 返回 true 表示允许(在限制内或会话已存在),false 表示拒绝(超出限制且是新会话) +func (s *GatewayService) checkAndRegisterSession(ctx context.Context, account *Account, sessionID string) bool { + // 只检查 Anthropic OAuth/SetupToken 账号 + if !account.IsAnthropicOAuthOrSetupToken() { + return true + } + + maxSessions := account.GetMaxSessions() + if maxSessions <= 0 || sessionID == "" { + return true // 未启用会话限制或无会话ID + } + + if s.sessionLimitCache == nil { + return true // 缓存不可用时允许通过 + } + + idleTimeout := time.Duration(account.GetSessionIdleTimeoutMinutes()) * time.Minute + + allowed, err := s.sessionLimitCache.RegisterSession(ctx, account.ID, sessionID, maxSessions, idleTimeout) + if err != nil { + // 失败开放:缓存错误时允许通过 + return true + } + return allowed +} + +func (s *GatewayService) getSchedulableAccount(ctx context.Context, accountID int64) (*Account, error) { + if s.schedulerSnapshot != nil { + return s.schedulerSnapshot.GetAccount(ctx, accountID) + } + return s.accountRepo.GetByID(ctx, accountID) +} + +func (s *GatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { + if account == nil || s.schedulerSnapshot == nil { + return account, nil + } + hydrated, err := s.schedulerSnapshot.GetAccount(ctx, account.ID) + if err != nil { + return nil, err + } + if hydrated == nil { + return nil, fmt.Errorf("selected gateway account %d not found during hydration", account.ID) + } + return hydrated, nil +} + +func (s *GatewayService) newSelectionResult(ctx context.Context, account *Account, acquired bool, release func(), waitPlan *AccountWaitPlan) (*AccountSelectionResult, error) { + hydrated, err := s.hydrateSelectedAccount(ctx, account) + if err != nil { + return nil, err + } + return &AccountSelectionResult{ + Account: hydrated, + Acquired: acquired, + ReleaseFunc: release, + WaitPlan: waitPlan, + }, nil +} + +// filterByMinPriority 过滤出优先级最小的账号集合 +func filterByMinPriority(accounts []accountWithLoad) []accountWithLoad { + if len(accounts) == 0 { + return accounts + } + minPriority := accounts[0].account.Priority + for _, acc := range accounts[1:] { + if acc.account.Priority < minPriority { + minPriority = acc.account.Priority + } + } + result := make([]accountWithLoad, 0, len(accounts)) + for _, acc := range accounts { + if acc.account.Priority == minPriority { + result = append(result, acc) + } + } + return result +} + +// filterByMinLoadRate 过滤出负载率最低的账号集合 +func filterByMinLoadRate(accounts []accountWithLoad) []accountWithLoad { + if len(accounts) == 0 { + return accounts + } + minLoadRate := accounts[0].loadInfo.LoadRate + for _, acc := range accounts[1:] { + if acc.loadInfo.LoadRate < minLoadRate { + minLoadRate = acc.loadInfo.LoadRate + } + } + result := make([]accountWithLoad, 0, len(accounts)) + for _, acc := range accounts { + if acc.loadInfo.LoadRate == minLoadRate { + result = append(result, acc) + } + } + return result +} + +// filterBySoonestReset 过滤出「会话窗口最早重置」的账号集合(use-it-or-lose-it)。 +// 仅保留拥有未来重置时间(SessionWindowEnd 在当前时间之后)且最早的账号; +// 窗口为空或已过期的账号视为无活跃窗口、优先级最低。 +// 当所有账号都没有活跃窗口时,返回原集合(不改变后续 LRU 选择)。 +func filterBySoonestReset(accounts []accountWithLoad) []accountWithLoad { + if len(accounts) <= 1 { + return accounts + } + now := time.Now() + var minEnd *time.Time + for _, acc := range accounts { + end := acc.account.SessionWindowEnd + if end == nil || !now.Before(*end) { + continue + } + if minEnd == nil || end.Before(*minEnd) { + minEnd = end + } + } + if minEnd == nil { + // 没有任何账号拥有活跃窗口,保持原集合 + return accounts + } + result := make([]accountWithLoad, 0, len(accounts)) + for _, acc := range accounts { + end := acc.account.SessionWindowEnd + if end != nil && now.Before(*end) && end.Equal(*minEnd) { + result = append(result, acc) + } + } + return result +} + +// selectByLRU 从集合中选择最久未用的账号 +// 如果有多个账号具有相同的最小 LastUsedAt,则随机选择一个 +func selectByLRU(accounts []accountWithLoad, preferOAuth bool) *accountWithLoad { + if len(accounts) == 0 { + return nil + } + if len(accounts) == 1 { + return &accounts[0] + } + + // 1. 找到最小的 LastUsedAt(nil 被视为最小) + var minTime *time.Time + hasNil := false + for _, acc := range accounts { + if acc.account.LastUsedAt == nil { + hasNil = true + break + } + if minTime == nil || acc.account.LastUsedAt.Before(*minTime) { + minTime = acc.account.LastUsedAt + } + } + + // 2. 收集所有具有最小 LastUsedAt 的账号索引 + var candidateIdxs []int + for i, acc := range accounts { + if hasNil { + if acc.account.LastUsedAt == nil { + candidateIdxs = append(candidateIdxs, i) + } + } else { + if acc.account.LastUsedAt != nil && acc.account.LastUsedAt.Equal(*minTime) { + candidateIdxs = append(candidateIdxs, i) + } + } + } + + // 3. 如果只有一个候选,直接返回 + if len(candidateIdxs) == 1 { + return &accounts[candidateIdxs[0]] + } + + // 4. 如果有多个候选且 preferOAuth,优先选择 OAuth 类型 + if preferOAuth { + var oauthIdxs []int + for _, idx := range candidateIdxs { + if accounts[idx].account.Type == AccountTypeOAuth { + oauthIdxs = append(oauthIdxs, idx) + } + } + if len(oauthIdxs) > 0 { + candidateIdxs = oauthIdxs + } + } + + // 5. 随机选择一个 + selectedIdx := candidateIdxs[mathrand.Intn(len(candidateIdxs))] + return &accounts[selectedIdx] +} + +func sortAccountsByPriorityAndLastUsed(accounts []*Account, preferOAuth bool) { + sort.SliceStable(accounts, func(i, j int) bool { + a, b := accounts[i], accounts[j] + if a.Priority != b.Priority { + return a.Priority < b.Priority + } + switch { + case a.LastUsedAt == nil && b.LastUsedAt != nil: + return true + case a.LastUsedAt != nil && b.LastUsedAt == nil: + return false + case a.LastUsedAt == nil && b.LastUsedAt == nil: + if preferOAuth && a.Type != b.Type { + return a.Type == AccountTypeOAuth + } + return false + default: + return a.LastUsedAt.Before(*b.LastUsedAt) + } + }) + shuffleWithinPriorityAndLastUsed(accounts, preferOAuth) +} + +// shuffleWithinSortGroups 对排序后的 accountWithLoad 切片,按 (Priority, LoadRate, LastUsedAt) 分组后组内随机打乱。 +// 防止并发请求读取同一快照时,确定性排序导致所有请求命中相同账号。 +func shuffleWithinSortGroups(accounts []accountWithLoad) { + if len(accounts) <= 1 { + return + } + i := 0 + for i < len(accounts) { + j := i + 1 + for j < len(accounts) && sameAccountWithLoadGroup(accounts[i], accounts[j]) { + j++ + } + if j-i > 1 { + mathrand.Shuffle(j-i, func(a, b int) { + accounts[i+a], accounts[i+b] = accounts[i+b], accounts[i+a] + }) + } + i = j + } +} + +// sameAccountWithLoadGroup 判断两个 accountWithLoad 是否属于同一排序组 +func sameAccountWithLoadGroup(a, b accountWithLoad) bool { + if a.account.Priority != b.account.Priority { + return false + } + if a.loadInfo.LoadRate != b.loadInfo.LoadRate { + return false + } + return sameLastUsedAt(a.account.LastUsedAt, b.account.LastUsedAt) +} + +// shuffleWithinPriorityAndLastUsed 对排序后的 []*Account 切片,按 (Priority, LastUsedAt) 分组后组内随机打乱。 +// +// 注意:当 preferOAuth=true 时,需要保证 OAuth 账号在同组内仍然优先,否则会把排序时的偏好打散掉。 +// 因此这里采用"组内分区 + 分区内 shuffle"的方式: +// - 先把同组账号按 (OAuth / 非 OAuth) 拆成两段,保持 OAuth 段在前; +// - 再分别在各段内随机打散,避免热点。 +func shuffleWithinPriorityAndLastUsed(accounts []*Account, preferOAuth bool) { + if len(accounts) <= 1 { + return + } + i := 0 + for i < len(accounts) { + j := i + 1 + for j < len(accounts) && sameAccountGroup(accounts[i], accounts[j]) { + j++ + } + if j-i > 1 { + if preferOAuth { + oauth := make([]*Account, 0, j-i) + others := make([]*Account, 0, j-i) + for _, acc := range accounts[i:j] { + if acc.Type == AccountTypeOAuth { + oauth = append(oauth, acc) + } else { + others = append(others, acc) + } + } + if len(oauth) > 1 { + mathrand.Shuffle(len(oauth), func(a, b int) { oauth[a], oauth[b] = oauth[b], oauth[a] }) + } + if len(others) > 1 { + mathrand.Shuffle(len(others), func(a, b int) { others[a], others[b] = others[b], others[a] }) + } + copy(accounts[i:], oauth) + copy(accounts[i+len(oauth):], others) + } else { + mathrand.Shuffle(j-i, func(a, b int) { + accounts[i+a], accounts[i+b] = accounts[i+b], accounts[i+a] + }) + } + } + i = j + } +} + +// sameAccountGroup 判断两个 Account 是否属于同一排序组(Priority + LastUsedAt) +func sameAccountGroup(a, b *Account) bool { + if a.Priority != b.Priority { + return false + } + return sameLastUsedAt(a.LastUsedAt, b.LastUsedAt) +} + +// sameLastUsedAt 判断两个 LastUsedAt 是否相同(精度到秒) +func sameLastUsedAt(a, b *time.Time) bool { + switch { + case a == nil && b == nil: + return true + case a == nil || b == nil: + return false + default: + return a.Unix() == b.Unix() + } +} + +// sortCandidatesForFallback 根据配置选择排序策略 +// mode: "last_used"(按最后使用时间) 或 "random"(随机) +func (s *GatewayService) sortCandidatesForFallback(accounts []*Account, preferOAuth bool, mode string) { + if mode == "random" { + // 先按优先级排序,然后在同优先级内随机打乱 + sortAccountsByPriorityOnly(accounts, preferOAuth) + shuffleWithinPriority(accounts) + } else { + // 默认按最后使用时间排序 + sortAccountsByPriorityAndLastUsed(accounts, preferOAuth) + } +} + +// sortAccountsByPriorityOnly 仅按优先级排序 +func sortAccountsByPriorityOnly(accounts []*Account, preferOAuth bool) { + sort.SliceStable(accounts, func(i, j int) bool { + a, b := accounts[i], accounts[j] + if a.Priority != b.Priority { + return a.Priority < b.Priority + } + if preferOAuth && a.Type != b.Type { + return a.Type == AccountTypeOAuth + } + return false + }) +} + +// shuffleWithinPriority 在同优先级内随机打乱顺序 +func shuffleWithinPriority(accounts []*Account) { + if len(accounts) <= 1 { + return + } + r := mathrand.New(mathrand.NewSource(time.Now().UnixNano())) + start := 0 + for start < len(accounts) { + priority := accounts[start].Priority + end := start + 1 + for end < len(accounts) && accounts[end].Priority == priority { + end++ + } + // 对 [start, end) 范围内的账户随机打乱 + if end-start > 1 { + r.Shuffle(end-start, func(i, j int) { + accounts[start+i], accounts[start+j] = accounts[start+j], accounts[start+i] + }) + } + start = end + } +} + +// selectAccountForModelWithPlatform 选择单平台账户(完全隔离) +func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, platform string) (*Account, error) { + preferOAuth := platform == PlatformGemini + routingAccountIDs := s.routingAccountIDsForRequest(ctx, groupID, requestedModel, platform) + + // require_privacy_set: 获取分组信息 + var schedGroup *Group + if groupID != nil && s.groupRepo != nil { + schedGroup, _ = s.groupRepo.GetByID(ctx, *groupID) + } + + var accounts []Account + accountsLoaded := false + + // ============ Model Routing (legacy path): apply before sticky session ============ + // When load-awareness is disabled (e.g. concurrency service not configured), we still honor model routing + // so switching model can switch upstream account within the same sticky session. + if len(routingAccountIDs) > 0 { + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy routed begin: group_id=%v model=%s platform=%s session=%s routed_ids=%v", + derefGroupID(groupID), requestedModel, platform, shortSessionHash(sessionHash), routingAccountIDs) + } + // 1) Sticky session only applies if the bound account is within the routing set. + if sessionHash != "" && s.cache != nil { + accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + if err == nil && accountID > 0 && containsInt64(routingAccountIDs, accountID) { + if _, excluded := excludedIDs[accountID]; !excluded { + account, err := s.getSchedulableAccount(ctx, accountID) + // 检查账号分组归属和平台匹配(确保粘性会话不会跨分组或跨平台) + if err == nil { + clearSticky := shouldClearStickySession(account, requestedModel) + if clearSticky { + _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + } + if !clearSticky && s.isAccountInGroup(account, groupID) && account.Platform == platform && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) && !s.isStickyAccountUpstreamRestricted(ctx, groupID, account, requestedModel) { + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), accountID) + } + return account, nil + } + } + } + } + } + + // 2) Select an account from the routed candidates. + forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string) + if hasForcePlatform && forcePlatform == "" { + hasForcePlatform = false + } + var err error + accounts, _, err = s.listSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) + if err != nil { + return nil, fmt.Errorf("query accounts failed: %w", err) + } + accountsLoaded = true + + // 提前预取窗口费用+RPM 计数,确保 routing 段内的调度检查调用能命中缓存 + ctx = s.withWindowCostPrefetch(ctx, accounts) + ctx = s.withRPMPrefetch(ctx, accounts) + + routingSet := make(map[int64]struct{}, len(routingAccountIDs)) + for _, id := range routingAccountIDs { + if id > 0 { + routingSet[id] = struct{}{} + } + } + + var selected *Account + for i := range accounts { + acc := &accounts[i] + if _, ok := routingSet[acc.ID]; !ok { + continue + } + if _, excluded := excludedIDs[acc.ID]; excluded { + continue + } + // Scheduler snapshots can be temporarily stale; re-check schedulability here to + // avoid selecting accounts that were recently rate-limited/overloaded. + if !s.isAccountSchedulableForSelection(acc) { + continue + } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { + _ = s.accountRepo.SetError(ctx, acc.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } + if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForQuota(acc) { + continue + } + if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { + continue + } + if !s.isAccountSchedulableForRPM(ctx, acc, false) { + continue + } + if selected == nil { + selected = acc + continue + } + if acc.Priority < selected.Priority { + selected = acc + } else if acc.Priority == selected.Priority { + switch { + case acc.LastUsedAt == nil && selected.LastUsedAt != nil: + selected = acc + case acc.LastUsedAt != nil && selected.LastUsedAt == nil: + // keep selected (never used is preferred) + case acc.LastUsedAt == nil && selected.LastUsedAt == nil: + if preferOAuth && acc.Type != selected.Type && acc.Type == AccountTypeOAuth { + selected = acc + } + default: + if acc.LastUsedAt.Before(*selected.LastUsedAt) { + selected = acc + } + } + } + } + + if selected != nil { + if sessionHash != "" && s.cache != nil { + if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { + logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) + } + } + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), selected.ID) + } + return selected, nil + } + logger.LegacyPrintf("service.gateway", "[ModelRouting] No routed accounts available for model=%s, falling back to normal selection", requestedModel) + } + + // 1. 查询粘性会话 + if sessionHash != "" && s.cache != nil { + accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + if err == nil && accountID > 0 { + if _, excluded := excludedIDs[accountID]; !excluded { + account, err := s.getSchedulableAccount(ctx, accountID) + // 检查账号分组归属和平台匹配(确保粘性会话不会跨分组或跨平台) + if err == nil { + clearSticky := shouldClearStickySession(account, requestedModel) + if clearSticky { + _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + } + if !clearSticky && s.isAccountInGroup(account, groupID) && account.Platform == platform && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) { + return account, nil + } + } + } + } + } + + // 2. 获取可调度账号列表(单平台) + if !accountsLoaded { + forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string) + if hasForcePlatform && forcePlatform == "" { + hasForcePlatform = false + } + var err error + accounts, _, err = s.listSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) + if err != nil { + return nil, fmt.Errorf("query accounts failed: %w", err) + } + } + + // 批量预取窗口费用+RPM 计数,避免逐个账号查询(N+1) + ctx = s.withWindowCostPrefetch(ctx, accounts) + ctx = s.withRPMPrefetch(ctx, accounts) + + // 3. 按优先级+最久未用选择(考虑模型支持) + // needsUpstreamCheck 仅在主选择循环中使用;粘性会话命中时跳过此检查, + // 因为粘性会话优先保持连接一致性,且 upstream 计费基准极少使用。 + needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) + var selected *Account + for i := range accounts { + acc := &accounts[i] + if _, excluded := excludedIDs[acc.ID]; excluded { + continue + } + // Scheduler snapshots can be temporarily stale; re-check schedulability here to + // avoid selecting accounts that were recently rate-limited/overloaded. + if !s.isAccountSchedulableForSelection(acc) { + continue + } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { + _ = s.accountRepo.SetError(ctx, acc.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } + if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { + continue + } + if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForQuota(acc) { + continue + } + if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { + continue + } + if !s.isAccountSchedulableForRPM(ctx, acc, false) { + continue + } + if selected == nil { + selected = acc + continue + } + if acc.Priority < selected.Priority { + selected = acc + } else if acc.Priority == selected.Priority { + switch { + case acc.LastUsedAt == nil && selected.LastUsedAt != nil: + selected = acc + case acc.LastUsedAt != nil && selected.LastUsedAt == nil: + // keep selected (never used is preferred) + case acc.LastUsedAt == nil && selected.LastUsedAt == nil: + if preferOAuth && acc.Type != selected.Type && acc.Type == AccountTypeOAuth { + selected = acc + } + default: + if acc.LastUsedAt.Before(*selected.LastUsedAt) { + selected = acc + } + } + } + } + + if selected == nil { + stats := s.logDetailedSelectionFailure(ctx, groupID, sessionHash, requestedModel, platform, accounts, excludedIDs, false) + if requestedModel != "" { + return nil, fmt.Errorf("%w supporting model: %s (%s)", ErrNoAvailableAccounts, requestedModel, summarizeSelectionFailureStats(stats)) + } + return nil, ErrNoAvailableAccounts + } + + // 4. 建立粘性绑定 + if sessionHash != "" && s.cache != nil { + if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { + logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) + } + } + + return selected, nil +} + +// selectAccountWithMixedScheduling 选择账户(支持混合调度) +// 查询原生平台账户 + 启用 mixed_scheduling 的 antigravity 账户 +func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, nativePlatform string) (*Account, error) { + preferOAuth := nativePlatform == PlatformGemini + routingAccountIDs := s.routingAccountIDsForRequest(ctx, groupID, requestedModel, nativePlatform) + + // require_privacy_set: 获取分组信息 + var schedGroup *Group + if groupID != nil && s.groupRepo != nil { + schedGroup, _ = s.groupRepo.GetByID(ctx, *groupID) + } + + var accounts []Account + accountsLoaded := false + + // ============ Model Routing (legacy path): apply before sticky session ============ + if len(routingAccountIDs) > 0 { + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy mixed routed begin: group_id=%v model=%s platform=%s session=%s routed_ids=%v", + derefGroupID(groupID), requestedModel, nativePlatform, shortSessionHash(sessionHash), routingAccountIDs) + } + // 1) Sticky session only applies if the bound account is within the routing set. + if sessionHash != "" && s.cache != nil { + accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + if err == nil && accountID > 0 && containsInt64(routingAccountIDs, accountID) { + if _, excluded := excludedIDs[accountID]; !excluded { + account, err := s.getSchedulableAccount(ctx, accountID) + // 检查账号分组归属和有效性:原生平台直接匹配,antigravity 需要启用混合调度 + if err == nil { + clearSticky := shouldClearStickySession(account, requestedModel) + if clearSticky { + _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + } + if !clearSticky && s.isAccountInGroup(account, groupID) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) { + if account.Platform == nativePlatform || (account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled()) { + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy mixed routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), accountID) + } + return account, nil + } + } + } + } + } + } + + // 2) Select an account from the routed candidates. + var err error + accounts, _, err = s.listSchedulableAccounts(ctx, groupID, nativePlatform, false) + if err != nil { + return nil, fmt.Errorf("query accounts failed: %w", err) + } + accountsLoaded = true + + // 提前预取窗口费用+RPM 计数,确保 routing 段内的调度检查调用能命中缓存 + ctx = s.withWindowCostPrefetch(ctx, accounts) + ctx = s.withRPMPrefetch(ctx, accounts) + + routingSet := make(map[int64]struct{}, len(routingAccountIDs)) + for _, id := range routingAccountIDs { + if id > 0 { + routingSet[id] = struct{}{} + } + } + + var selected *Account + for i := range accounts { + acc := &accounts[i] + if _, ok := routingSet[acc.ID]; !ok { + continue + } + if _, excluded := excludedIDs[acc.ID]; excluded { + continue + } + // Scheduler snapshots can be temporarily stale; re-check schedulability here to + // avoid selecting accounts that were recently rate-limited/overloaded. + if !s.isAccountSchedulableForSelection(acc) { + continue + } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { + _ = s.accountRepo.SetError(ctx, acc.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } + // 过滤:原生平台直接通过,antigravity 需要启用混合调度 + if acc.Platform == PlatformAntigravity && !acc.IsMixedSchedulingEnabled() { + continue + } + if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForQuota(acc) { + continue + } + if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { + continue + } + if !s.isAccountSchedulableForRPM(ctx, acc, false) { + continue + } + if selected == nil { + selected = acc + continue + } + if acc.Priority < selected.Priority { + selected = acc + } else if acc.Priority == selected.Priority { + switch { + case acc.LastUsedAt == nil && selected.LastUsedAt != nil: + selected = acc + case acc.LastUsedAt != nil && selected.LastUsedAt == nil: + // keep selected (never used is preferred) + case acc.LastUsedAt == nil && selected.LastUsedAt == nil: + if preferOAuth && acc.Platform == PlatformGemini && selected.Platform == PlatformGemini && acc.Type != selected.Type && acc.Type == AccountTypeOAuth { + selected = acc + } + default: + if acc.LastUsedAt.Before(*selected.LastUsedAt) { + selected = acc + } + } + } + } + + if selected != nil { + if sessionHash != "" && s.cache != nil { + if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { + logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) + } + } + if s.debugModelRoutingEnabled() { + logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy mixed routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), selected.ID) + } + return selected, nil + } + logger.LegacyPrintf("service.gateway", "[ModelRouting] No routed accounts available for model=%s, falling back to normal selection", requestedModel) + } + + // 1. 查询粘性会话 + if sessionHash != "" && s.cache != nil { + accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + if err == nil && accountID > 0 { + if _, excluded := excludedIDs[accountID]; !excluded { + account, err := s.getSchedulableAccount(ctx, accountID) + // 检查账号分组归属和有效性:原生平台直接匹配,antigravity 需要启用混合调度 + if err == nil { + clearSticky := shouldClearStickySession(account, requestedModel) + if clearSticky { + _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + } + if !clearSticky && s.isAccountInGroup(account, groupID) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) && !s.isStickyAccountUpstreamRestricted(ctx, groupID, account, requestedModel) { + if account.Platform == nativePlatform || (account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled()) { + return account, nil + } + } + } + } + } + } + + // 2. 获取可调度账号列表 + if !accountsLoaded { + var err error + accounts, _, err = s.listSchedulableAccounts(ctx, groupID, nativePlatform, false) + if err != nil { + return nil, fmt.Errorf("query accounts failed: %w", err) + } + } + + // 批量预取窗口费用+RPM 计数,避免逐个账号查询(N+1) + ctx = s.withWindowCostPrefetch(ctx, accounts) + ctx = s.withRPMPrefetch(ctx, accounts) + + // 3. 按优先级+最久未用选择(考虑模型支持和混合调度) + // needsUpstreamCheck 仅在主选择循环中使用;粘性会话命中时跳过此检查。 + needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) + var selected *Account + for i := range accounts { + acc := &accounts[i] + if _, excluded := excludedIDs[acc.ID]; excluded { + continue + } + // Scheduler snapshots can be temporarily stale; re-check schedulability here to + // avoid selecting accounts that were recently rate-limited/overloaded. + if !s.isAccountSchedulableForSelection(acc) { + continue + } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { + _ = s.accountRepo.SetError(ctx, acc.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } + // 过滤:原生平台直接通过,antigravity 需要启用混合调度 + if acc.Platform == PlatformAntigravity && !acc.IsMixedSchedulingEnabled() { + continue + } + if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { + continue + } + if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { + continue + } + if !s.isAccountSchedulableForQuota(acc) { + continue + } + if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { + continue + } + if !s.isAccountSchedulableForRPM(ctx, acc, false) { + continue + } + if selected == nil { + selected = acc + continue + } + if acc.Priority < selected.Priority { + selected = acc + } else if acc.Priority == selected.Priority { + switch { + case acc.LastUsedAt == nil && selected.LastUsedAt != nil: + selected = acc + case acc.LastUsedAt != nil && selected.LastUsedAt == nil: + // keep selected (never used is preferred) + case acc.LastUsedAt == nil && selected.LastUsedAt == nil: + if preferOAuth && acc.Platform == PlatformGemini && selected.Platform == PlatformGemini && acc.Type != selected.Type && acc.Type == AccountTypeOAuth { + selected = acc + } + default: + if acc.LastUsedAt.Before(*selected.LastUsedAt) { + selected = acc + } + } + } + } + + if selected == nil { + stats := s.logDetailedSelectionFailure(ctx, groupID, sessionHash, requestedModel, nativePlatform, accounts, excludedIDs, true) + if requestedModel != "" { + return nil, fmt.Errorf("%w supporting model: %s (%s)", ErrNoAvailableAccounts, requestedModel, summarizeSelectionFailureStats(stats)) + } + return nil, ErrNoAvailableAccounts + } + + // 4. 建立粘性绑定 + if sessionHash != "" && s.cache != nil { + if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { + logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) + } + } + + return selected, nil +} + +type selectionFailureStats struct { + Total int + Eligible int + Excluded int + Unschedulable int + PlatformFiltered int + ModelUnsupported int + ModelRateLimited int + SamplePlatformIDs []int64 + SampleMappingIDs []int64 + SampleRateLimitIDs []string +} + +type selectionFailureDiagnosis struct { + Category string + Detail string +} + +func (s *GatewayService) logDetailedSelectionFailure( + ctx context.Context, + groupID *int64, + sessionHash string, + requestedModel string, + platform string, + accounts []Account, + excludedIDs map[int64]struct{}, + allowMixedScheduling bool, +) selectionFailureStats { + stats := s.collectSelectionFailureStats(ctx, accounts, requestedModel, platform, excludedIDs, allowMixedScheduling) + logger.LegacyPrintf( + "service.gateway", + "[SelectAccountDetailed] group_id=%v model=%s platform=%s session=%s total=%d eligible=%d excluded=%d unschedulable=%d platform_filtered=%d model_unsupported=%d model_rate_limited=%d sample_platform_filtered=%v sample_model_unsupported=%v sample_model_rate_limited=%v", + derefGroupID(groupID), + requestedModel, + platform, + shortSessionHash(sessionHash), + stats.Total, + stats.Eligible, + stats.Excluded, + stats.Unschedulable, + stats.PlatformFiltered, + stats.ModelUnsupported, + stats.ModelRateLimited, + stats.SamplePlatformIDs, + stats.SampleMappingIDs, + stats.SampleRateLimitIDs, + ) + return stats +} + +func (s *GatewayService) collectSelectionFailureStats( + ctx context.Context, + accounts []Account, + requestedModel string, + platform string, + excludedIDs map[int64]struct{}, + allowMixedScheduling bool, +) selectionFailureStats { + stats := selectionFailureStats{ + Total: len(accounts), + } + + for i := range accounts { + acc := &accounts[i] + diagnosis := s.diagnoseSelectionFailure(ctx, acc, requestedModel, platform, excludedIDs, allowMixedScheduling) + switch diagnosis.Category { + case "excluded": + stats.Excluded++ + case "unschedulable": + stats.Unschedulable++ + case "platform_filtered": + stats.PlatformFiltered++ + stats.SamplePlatformIDs = appendSelectionFailureSampleID(stats.SamplePlatformIDs, acc.ID) + case "model_unsupported": + stats.ModelUnsupported++ + stats.SampleMappingIDs = appendSelectionFailureSampleID(stats.SampleMappingIDs, acc.ID) + case "model_rate_limited": + stats.ModelRateLimited++ + remaining := acc.GetRateLimitRemainingTimeWithContext(ctx, requestedModel).Truncate(time.Second) + stats.SampleRateLimitIDs = appendSelectionFailureRateSample(stats.SampleRateLimitIDs, acc.ID, remaining) + default: + stats.Eligible++ + } + } + + return stats +} + +func (s *GatewayService) diagnoseSelectionFailure( + ctx context.Context, + acc *Account, + requestedModel string, + platform string, + excludedIDs map[int64]struct{}, + allowMixedScheduling bool, +) selectionFailureDiagnosis { + if acc == nil { + return selectionFailureDiagnosis{Category: "unschedulable", Detail: "account_nil"} + } + if _, excluded := excludedIDs[acc.ID]; excluded { + return selectionFailureDiagnosis{Category: "excluded"} + } + if !s.isAccountSchedulableForSelection(acc) { + return selectionFailureDiagnosis{Category: "unschedulable", Detail: "generic_unschedulable"} + } + if isPlatformFilteredForSelection(acc, platform, allowMixedScheduling) { + return selectionFailureDiagnosis{ + Category: "platform_filtered", + Detail: fmt.Sprintf("account_platform=%s requested_platform=%s", acc.Platform, strings.TrimSpace(platform)), + } + } + if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { + return selectionFailureDiagnosis{ + Category: "model_unsupported", + Detail: fmt.Sprintf("model=%s", requestedModel), + } + } + if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { + remaining := acc.GetRateLimitRemainingTimeWithContext(ctx, requestedModel).Truncate(time.Second) + return selectionFailureDiagnosis{ + Category: "model_rate_limited", + Detail: fmt.Sprintf("remaining=%s", remaining), + } + } + return selectionFailureDiagnosis{Category: "eligible"} +} + +func isPlatformFilteredForSelection(acc *Account, platform string, allowMixedScheduling bool) bool { + if acc == nil { + return true + } + if allowMixedScheduling { + if acc.Platform == PlatformAntigravity { + return !acc.IsMixedSchedulingEnabled() + } + return acc.Platform != platform + } + if strings.TrimSpace(platform) == "" { + return false + } + return acc.Platform != platform +} + +func appendSelectionFailureSampleID(samples []int64, id int64) []int64 { + const limit = 5 + if len(samples) >= limit { + return samples + } + return append(samples, id) +} + +func appendSelectionFailureRateSample(samples []string, accountID int64, remaining time.Duration) []string { + const limit = 5 + if len(samples) >= limit { + return samples + } + return append(samples, fmt.Sprintf("%d(%s)", accountID, remaining)) +} + +func summarizeSelectionFailureStats(stats selectionFailureStats) string { + return fmt.Sprintf( + "total=%d eligible=%d excluded=%d unschedulable=%d platform_filtered=%d model_unsupported=%d model_rate_limited=%d", + stats.Total, + stats.Eligible, + stats.Excluded, + stats.Unschedulable, + stats.PlatformFiltered, + stats.ModelUnsupported, + stats.ModelRateLimited, + ) +} + +// isModelSupportedByAccountWithContext 根据账户平台检查模型支持(带 context) +// 对于 Antigravity 平台,会先获取映射后的最终模型名(包括 thinking 后缀)再检查支持 +func (s *GatewayService) isModelSupportedByAccountWithContext(ctx context.Context, account *Account, requestedModel string) bool { + if account.Platform == PlatformAntigravity { + if strings.TrimSpace(requestedModel) == "" { + return true + } + // 使用与转发阶段一致的映射逻辑:自定义映射优先 → 默认映射兜底 + mapped := mapAntigravityModel(account, requestedModel) + if mapped == "" { + return false + } + // 应用 thinking 后缀后检查最终模型是否在账号映射中 + if enabled, ok := ThinkingEnabledFromContext(ctx); ok { + finalModel := applyThinkingModelSuffix(mapped, enabled) + if finalModel == mapped { + return true // thinking 后缀未改变模型名,映射已通过 + } + return account.IsModelSupported(finalModel) + } + return true + } + return s.isModelSupportedByAccount(account, requestedModel) +} + +// isModelSupportedByAccount 根据账户平台检查模型支持(无 context,用于非 Antigravity 平台) +func (s *GatewayService) isModelSupportedByAccount(account *Account, requestedModel string) bool { + if account.Platform == PlatformAntigravity { + if strings.TrimSpace(requestedModel) == "" { + return true + } + return mapAntigravityModel(account, requestedModel) != "" + } + if account.IsBedrock() { + _, ok := ResolveBedrockModelID(account, requestedModel) + return ok + } + // OpenAI 透传模式:仅替换认证,允许所有模型 + if account.Platform == PlatformOpenAI && account.IsOpenAIPassthroughEnabled() { + return true + } + // OAuth/SetupToken 账号使用 Anthropic 标准映射(短ID → 长ID) + if account.Platform == PlatformAnthropic && account.Type != AccountTypeAPIKey { + if account.Type == AccountTypeServiceAccount { + requestedModel = normalizeVertexAnthropicModelID(claude.NormalizeModelID(requestedModel)) + } else { + requestedModel = claude.NormalizeModelID(requestedModel) + } + } + // 其他平台使用账户的模型支持检查 + return account.IsModelSupported(requestedModel) +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 2ca554edb3..79150f5598 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -10,7 +10,6 @@ import ( "fmt" "io" "log/slog" - mathrand "math/rand" "net" "net/http" "net/url" @@ -32,7 +31,6 @@ import ( infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" - "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" "github.com/cespare/xxhash/v2" @@ -1560,2444 +1558,6 @@ func generateSessionUUID(seed string) string { bytes[0:4], bytes[4:6], bytes[6:8], bytes[8:10], bytes[10:16]) } -// SelectAccount 选择账号(粘性会话+优先级) -func (s *GatewayService) SelectAccount(ctx context.Context, groupID *int64, sessionHash string) (*Account, error) { - return s.SelectAccountForModel(ctx, groupID, sessionHash, "") -} - -// SelectAccountForModel 选择支持指定模型的账号(粘性会话+优先级+模型映射) -func (s *GatewayService) SelectAccountForModel(ctx context.Context, groupID *int64, sessionHash string, requestedModel string) (*Account, error) { - return s.SelectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, nil) -} - -// SelectAccountForModelWithExclusions selects an account supporting the requested model while excluding specified accounts. -func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*Account, error) { - // 优先检查 context 中的强制平台(/antigravity 路由) - var platform string - forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string) - if hasForcePlatform && forcePlatform != "" { - platform = forcePlatform - } else if groupID != nil { - group, resolvedGroupID, err := s.resolveGatewayGroup(ctx, groupID) - if err != nil { - return nil, err - } - groupID = resolvedGroupID - ctx = s.withGroupContext(ctx, group) - platform = group.Platform - } else { - // 无分组时只使用原生 anthropic 平台 - platform = PlatformAnthropic - } - - // Claude Code 限制可能已将 groupID 解析为 fallback group, - // 渠道限制预检查必须使用解析后的分组。 - if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { - slog.Warn("channel pricing restriction blocked request", - "group_id", derefGroupID(groupID), - "model", requestedModel) - return nil, fmt.Errorf("%w supporting model: %s (channel pricing restriction)", ErrNoAvailableAccounts, requestedModel) - } - - // anthropic/gemini 分组支持混合调度(包含启用了 mixed_scheduling 的 antigravity 账户) - // 注意:强制平台模式不走混合调度 - if (platform == PlatformAnthropic || platform == PlatformGemini) && !hasForcePlatform { - account, err := s.selectAccountWithMixedScheduling(ctx, groupID, sessionHash, requestedModel, excludedIDs, platform) - if err != nil { - return nil, err - } - return s.hydrateSelectedAccount(ctx, account) - } - - // antigravity 分组、强制平台模式或无分组使用单平台选择 - // 注意:强制平台模式也必须遵守分组限制,不再回退到全平台查询 - account, err := s.selectAccountForModelWithPlatform(ctx, groupID, sessionHash, requestedModel, excludedIDs, platform) - if err != nil { - return nil, err - } - return s.hydrateSelectedAccount(ctx, account) -} - -// SelectAccountWithLoadAwareness selects account with load-awareness and wait plan. -// metadataUserID: 用于客户端亲和调度,从中提取客户端 ID -// sub2apiUserID: 系统用户 ID,用于二维亲和调度 -func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string, sub2apiUserID int64) (*AccountSelectionResult, error) { - // 调试日志:记录调度入口参数 - excludedIDsList := make([]int64, 0, len(excludedIDs)) - for id := range excludedIDs { - excludedIDsList = append(excludedIDsList, id) - } - slog.Debug("account_scheduling_starting", - "group_id", derefGroupID(groupID), - "model", requestedModel, - "session", shortSessionHash(sessionHash), - "excluded_ids", excludedIDsList) - - cfg := s.schedulingConfig() - - // 检查 Claude Code 客户端限制(可能会替换 groupID 为降级分组) - group, groupID, err := s.checkClaudeCodeRestriction(ctx, groupID) - if err != nil { - return nil, err - } - ctx = s.withGroupContext(ctx, group) - - // Claude Code 限制可能已将 groupID 解析为 fallback group, - // 渠道限制预检查必须使用解析后的分组。 - if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { - slog.Warn("channel pricing restriction blocked request", - "group_id", derefGroupID(groupID), - "model", requestedModel) - return nil, fmt.Errorf("%w supporting model: %s (channel pricing restriction)", ErrNoAvailableAccounts, requestedModel) - } - - var stickyAccountID int64 - var stickySource string - if prefetch := prefetchedStickyAccountIDFromContext(ctx, groupID); prefetch > 0 { - stickyAccountID = prefetch - stickySource = "prefetch" - } else if sessionHash != "" && s.cache != nil { - if accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash); err == nil { - stickyAccountID = accountID - stickySource = "cache" - } - } - - // [DEBUG-STICKY] 调度器入口日志 - slog.Info("sticky.scheduler_entry", - "group_id", derefGroupID(groupID), - "session_hash", shortSessionHash(sessionHash), - "sticky_account_id", stickyAccountID, - "sticky_source", stickySource, - "model", requestedModel, - "load_batch", cfg.LoadBatchEnabled, - "has_concurrency_svc", s.concurrencyService != nil, - "excluded_count", len(excludedIDs), - ) - - if s.debugModelRoutingEnabled() && requestedModel != "" { - groupPlatform := "" - if group != nil { - groupPlatform = group.Platform - } - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] select entry: group_id=%v group_platform=%s model=%s session=%s sticky_account=%d load_batch=%v concurrency=%v", - derefGroupID(groupID), groupPlatform, requestedModel, shortSessionHash(sessionHash), stickyAccountID, cfg.LoadBatchEnabled, s.concurrencyService != nil) - } - - if s.concurrencyService == nil || !cfg.LoadBatchEnabled { - // 复制排除列表,用于会话限制拒绝时的重试 - localExcluded := make(map[int64]struct{}) - for k, v := range excludedIDs { - localExcluded[k] = v - } - - for { - account, err := s.SelectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, localExcluded) - if err != nil { - return nil, err - } - - result, err := s.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) - if err == nil && result.Acquired { - // 获取槽位后检查会话限制(使用 sessionHash 作为会话标识符) - if !s.checkAndRegisterSession(ctx, account, sessionHash) { - result.ReleaseFunc() // 释放槽位 - localExcluded[account.ID] = struct{}{} // 排除此账号 - continue // 重新选择 - } - return s.newSelectionResult(ctx, account, true, result.ReleaseFunc, nil) - } - - // 对于等待计划的情况,也需要先检查会话限制 - if !s.checkAndRegisterSession(ctx, account, sessionHash) { - localExcluded[account.ID] = struct{}{} - continue - } - - if stickyAccountID > 0 && stickyAccountID == account.ID && s.concurrencyService != nil { - waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, account.ID) - if waitingCount < cfg.StickySessionMaxWaiting { - return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ - AccountID: account.ID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }) - } - } - return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ - AccountID: account.ID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.FallbackWaitTimeout, - MaxWaiting: cfg.FallbackMaxWaiting, - }) - } - } - - platform, hasForcePlatform, err := s.resolvePlatform(ctx, groupID, group) - if err != nil { - return nil, err - } - preferOAuth := platform == PlatformGemini - if s.debugModelRoutingEnabled() && platform == PlatformAnthropic && requestedModel != "" { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] load-aware enabled: group_id=%v model=%s session=%s platform=%s", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), platform) - } - - accounts, useMixed, err := s.listSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) - if err != nil { - return nil, err - } - if len(accounts) == 0 { - return nil, ErrNoAvailableAccounts - } - ctx = s.withWindowCostPrefetch(ctx, accounts) - ctx = s.withRPMPrefetch(ctx, accounts) - - // 提前构建 accountByID(供 Layer 1 和 Layer 1.5 使用) - accountByID := make(map[int64]*Account, len(accounts)) - for i := range accounts { - accountByID[accounts[i].ID] = &accounts[i] - } - isExcluded := func(accountID int64) bool { - if excludedIDs == nil { - return false - } - _, excluded := excludedIDs[accountID] - return excluded - } - - // 获取模型路由配置(仅 anthropic 平台) - var routingAccountIDs []int64 - if group != nil && requestedModel != "" && group.Platform == PlatformAnthropic { - routingAccountIDs = group.GetRoutingAccountIDs(requestedModel) - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] context group routing: group_id=%d model=%s enabled=%v rules=%d matched_ids=%v session=%s sticky_account=%d", - group.ID, requestedModel, group.ModelRoutingEnabled, len(group.ModelRouting), routingAccountIDs, shortSessionHash(sessionHash), stickyAccountID) - if len(routingAccountIDs) == 0 && group.ModelRoutingEnabled && len(group.ModelRouting) > 0 { - keys := make([]string, 0, len(group.ModelRouting)) - for k := range group.ModelRouting { - keys = append(keys, k) - } - sort.Strings(keys) - const maxKeys = 20 - if len(keys) > maxKeys { - keys = keys[:maxKeys] - } - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] context group routing miss: group_id=%d model=%s patterns(sample)=%v", group.ID, requestedModel, keys) - } - } - } - - // ============ Layer 1: 模型路由优先选择(优先级高于粘性会话) ============ - if len(routingAccountIDs) > 0 && s.concurrencyService != nil { - // 1. 过滤出路由列表中可调度的账号 - var routingCandidates []*Account - var filteredExcluded, filteredMissing, filteredUnsched, filteredPlatform, filteredModelScope, filteredModelMapping, filteredWindowCost int - var modelScopeSkippedIDs []int64 // 记录因模型限流被跳过的账号 ID - for _, routingAccountID := range routingAccountIDs { - if isExcluded(routingAccountID) { - filteredExcluded++ - continue - } - account, ok := accountByID[routingAccountID] - if !ok || !s.isAccountSchedulableForSelection(account) { - if !ok { - filteredMissing++ - } else { - filteredUnsched++ - } - continue - } - if !s.isAccountAllowedForPlatform(account, platform, useMixed) { - filteredPlatform++ - continue - } - if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, account, requestedModel) { - filteredModelMapping++ - continue - } - if !s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) { - filteredModelScope++ - modelScopeSkippedIDs = append(modelScopeSkippedIDs, account.ID) - continue - } - // 配额检查 - if !s.isAccountSchedulableForQuota(account) { - continue - } - // 窗口费用检查(非粘性会话路径) - if !s.isAccountSchedulableForWindowCost(ctx, account, false) { - filteredWindowCost++ - continue - } - // RPM 检查(非粘性会话路径) - if !s.isAccountSchedulableForRPM(ctx, account, false) { - continue - } - routingCandidates = append(routingCandidates, account) - } - - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed candidates: group_id=%v model=%s routed=%d candidates=%d filtered(excluded=%d missing=%d unsched=%d platform=%d model_scope=%d model_mapping=%d window_cost=%d)", - derefGroupID(groupID), requestedModel, len(routingAccountIDs), len(routingCandidates), - filteredExcluded, filteredMissing, filteredUnsched, filteredPlatform, filteredModelScope, filteredModelMapping, filteredWindowCost) - if len(modelScopeSkippedIDs) > 0 { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] model_rate_limited accounts skipped: group_id=%v model=%s account_ids=%v", - derefGroupID(groupID), requestedModel, modelScopeSkippedIDs) - } - } - - if len(routingCandidates) > 0 { - // 1.5. 在路由账号范围内检查粘性会话 - if sessionHash != "" && stickyAccountID > 0 { - slog.Debug("sticky.layer1_5_checking", - "sticky_account_id", stickyAccountID, - "in_routing_list", containsInt64(routingAccountIDs, stickyAccountID), - "is_excluded", isExcluded(stickyAccountID), - "in_account_map", func() bool { _, ok := accountByID[stickyAccountID]; return ok }(), - "session", shortSessionHash(sessionHash), - ) - if containsInt64(routingAccountIDs, stickyAccountID) && !isExcluded(stickyAccountID) { - // 粘性账号在路由列表中,优先使用 - if stickyAccount, ok := accountByID[stickyAccountID]; ok { - var stickyCacheMissReason string - - gatePass := s.isAccountSchedulableForSelection(stickyAccount) && - s.isAccountAllowedForPlatform(stickyAccount, platform, useMixed) && - (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, stickyAccount, requestedModel)) && - s.isAccountSchedulableForModelSelection(ctx, stickyAccount, requestedModel) && - s.isAccountSchedulableForQuota(stickyAccount) && - s.isAccountSchedulableForWindowCost(ctx, stickyAccount, true) - - rpmPass := gatePass && s.isAccountSchedulableForRPM(ctx, stickyAccount, true) - - if rpmPass { // 粘性会话窗口费用+RPM 检查 - result, err := s.tryAcquireAccountSlot(ctx, stickyAccountID, stickyAccount.Concurrency) - if err == nil && result.Acquired { - // 会话数量限制检查 - if !s.checkAndRegisterSession(ctx, stickyAccount, sessionHash) { - result.ReleaseFunc() // 释放槽位 - stickyCacheMissReason = "session_limit" - // 继续到负载感知选择 - } else { - slog.Debug("sticky.layer1_5_hit", - "account_id", stickyAccountID, - "session", shortSessionHash(sessionHash), - "result", "slot_acquired", - ) - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), stickyAccountID) - } - return s.newSelectionResult(ctx, stickyAccount, true, result.ReleaseFunc, nil) - } - } - - if stickyCacheMissReason == "" { - waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, stickyAccountID) - if waitingCount < cfg.StickySessionMaxWaiting { - // 会话数量限制检查(等待计划也需要占用会话配额) - if !s.checkAndRegisterSession(ctx, stickyAccount, sessionHash) { - stickyCacheMissReason = "session_limit" - // 会话限制已满,继续到负载感知选择 - } else { - return &AccountSelectionResult{ - Account: stickyAccount, - WaitPlan: &AccountWaitPlan{ - AccountID: stickyAccountID, - MaxConcurrency: stickyAccount.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }, - }, nil - } - } else { - stickyCacheMissReason = "wait_queue_full" - } - } - // 粘性账号槽位满且等待队列已满,继续使用负载感知选择 - } else if !gatePass { - stickyCacheMissReason = "gate_check" - } else { - stickyCacheMissReason = "rpm_red" - } - - // 记录粘性缓存未命中的结构化日志 - if stickyCacheMissReason != "" { - baseRPM := stickyAccount.GetBaseRPM() - var currentRPM int - if count, ok := rpmFromPrefetchContext(ctx, stickyAccount.ID); ok { - currentRPM = count - } - logger.LegacyPrintf("service.gateway", "[StickyCacheMiss] reason=%s account_id=%d session=%s current_rpm=%d base_rpm=%d", - stickyCacheMissReason, stickyAccountID, shortSessionHash(sessionHash), currentRPM, baseRPM) - } - } else { - _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - logger.LegacyPrintf("service.gateway", "[StickyCacheMiss] reason=account_cleared account_id=%d session=%s current_rpm=0 base_rpm=0", - stickyAccountID, shortSessionHash(sessionHash)) - } - } - } - - // 2. 批量获取负载信息 - routingLoads := make([]AccountWithConcurrency, 0, len(routingCandidates)) - for _, acc := range routingCandidates { - routingLoads = append(routingLoads, AccountWithConcurrency{ - ID: acc.ID, - MaxConcurrency: acc.EffectiveLoadFactor(), - }) - } - routingLoadMap, _ := s.concurrencyService.GetAccountsLoadBatch(ctx, routingLoads) - - // 3. 按负载感知排序 - var routingAvailable []accountWithLoad - for _, acc := range routingCandidates { - loadInfo := routingLoadMap[acc.ID] - if loadInfo == nil { - loadInfo = &AccountLoadInfo{AccountID: acc.ID} - } - if loadInfo.LoadRate < 100 { - routingAvailable = append(routingAvailable, accountWithLoad{account: acc, loadInfo: loadInfo}) - } - } - - if len(routingAvailable) > 0 { - // 排序:优先级 > 负载率 > 最后使用时间 - sort.SliceStable(routingAvailable, func(i, j int) bool { - a, b := routingAvailable[i], routingAvailable[j] - if a.account.Priority != b.account.Priority { - return a.account.Priority < b.account.Priority - } - if a.loadInfo.LoadRate != b.loadInfo.LoadRate { - return a.loadInfo.LoadRate < b.loadInfo.LoadRate - } - switch { - case a.account.LastUsedAt == nil && b.account.LastUsedAt != nil: - return true - case a.account.LastUsedAt != nil && b.account.LastUsedAt == nil: - return false - case a.account.LastUsedAt == nil && b.account.LastUsedAt == nil: - return false - default: - return a.account.LastUsedAt.Before(*b.account.LastUsedAt) - } - }) - shuffleWithinSortGroups(routingAvailable) - - // 4. 尝试获取槽位 - for _, item := range routingAvailable { - result, err := s.tryAcquireAccountSlot(ctx, item.account.ID, item.account.Concurrency) - if err == nil && result.Acquired { - // 会话数量限制检查 - if !s.checkAndRegisterSession(ctx, item.account, sessionHash) { - result.ReleaseFunc() // 释放槽位,继续尝试下一个账号 - continue - } - if sessionHash != "" && s.cache != nil { - _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, item.account.ID, stickySessionTTL) - } - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) - } - return s.newSelectionResult(ctx, item.account, true, result.ReleaseFunc, nil) - } - } - - // 5. 所有路由账号槽位满,尝试返回等待计划(选择负载最低的) - // 遍历找到第一个满足会话限制的账号 - for _, item := range routingAvailable { - if !s.checkAndRegisterSession(ctx, item.account, sessionHash) { - continue // 会话限制已满,尝试下一个 - } - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed wait: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) - } - return s.newSelectionResult(ctx, item.account, false, nil, &AccountWaitPlan{ - AccountID: item.account.ID, - MaxConcurrency: item.account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }) - } - // 所有路由账号会话限制都已满,继续到 Layer 2 回退 - } - // 路由列表中的账号都不可用(负载率 >= 100),继续到 Layer 2 回退 - logger.LegacyPrintf("service.gateway", "[ModelRouting] All routed accounts unavailable for model=%s, falling back to normal selection", requestedModel) - } - } - - // ============ Layer 1.5: 粘性会话(仅在无模型路由配置时生效) ============ - if len(routingAccountIDs) == 0 && sessionHash != "" && stickyAccountID > 0 && !isExcluded(stickyAccountID) { - accountID := stickyAccountID - if accountID > 0 && !isExcluded(accountID) { - account, ok := accountByID[accountID] - if ok { - // 检查账户是否需要清理粘性会话绑定 - clearSticky := shouldClearStickySession(account, requestedModel) - if clearSticky { - slog.Debug("sticky.layer1_5_no_routing_clear", - "account_id", accountID, - "reason", "should_clear_sticky_session", - "session", shortSessionHash(sessionHash), - ) - _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - } - - // 注意:不再检查 isAccountInGroup,因为 accountByID 已经从按分组过滤的 - // accounts 列表构建,账号一定在分组内。而 scheduler snapshot 缓存 - // 反序列化后 AccountGroups 字段为空,导致 isAccountInGroup 永远返回 false。 - platformOK := s.isAccountAllowedForPlatform(account, platform, useMixed) - modelSupported := requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel) - modelSchedulable := s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) - quotaOK := s.isAccountSchedulableForQuota(account) - windowCostOK := s.isAccountSchedulableForWindowCost(ctx, account, true) - rpmOK := s.isAccountSchedulableForRPM(ctx, account, true) - schedulable := s.isAccountSchedulableForSelection(account) - - slog.Debug("sticky.layer1_5_no_routing_checks", - "account_id", accountID, - "session", shortSessionHash(sessionHash), - "clear_sticky", clearSticky, - "schedulable", schedulable, - "platform_ok", platformOK, - "model_supported", modelSupported, - "model_schedulable", modelSchedulable, - "quota_ok", quotaOK, - "window_cost_ok", windowCostOK, - "rpm_ok", rpmOK, - ) - - if !clearSticky && platformOK && modelSupported && modelSchedulable && quotaOK && windowCostOK && rpmOK && schedulable { - result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency) - if err == nil && result.Acquired { - // 会话数量限制检查 - if !s.checkAndRegisterSession(ctx, account, sessionHash) { - result.ReleaseFunc() // 释放槽位,继续到 Layer 2 - slog.Debug("sticky.layer1_5_no_routing_miss", - "account_id", accountID, - "reason", "session_limit", - "session", shortSessionHash(sessionHash), - ) - } else { - slog.Debug("sticky.layer1_5_no_routing_hit", - "account_id", accountID, - "session", shortSessionHash(sessionHash), - "result", "slot_acquired", - ) - if s.cache != nil { - _ = s.cache.RefreshSessionTTL(ctx, derefGroupID(groupID), sessionHash, stickySessionTTL) - } - return s.newSelectionResult(ctx, account, true, result.ReleaseFunc, nil) - } - } else { - slog.Debug("sticky.layer1_5_no_routing_slot_busy", - "account_id", accountID, - "session", shortSessionHash(sessionHash), - ) - } - - waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, accountID) - if waitingCount < cfg.StickySessionMaxWaiting { - // 会话数量限制检查(等待计划也需要占用会话配额) - if !s.checkAndRegisterSession(ctx, account, sessionHash) { - // 会话限制已满,继续到 Layer 2 - } else { - slog.Debug("sticky.layer1_5_no_routing_hit", - "account_id", accountID, - "session", shortSessionHash(sessionHash), - "result", "wait_plan", - ) - return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ - AccountID: accountID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }) - } - } - } else if !clearSticky { - slog.Debug("sticky.layer1_5_no_routing_miss", - "account_id", accountID, - "reason", "gate_check_failed", - "session", shortSessionHash(sessionHash), - ) - } - } else { - slog.Debug("sticky.layer1_5_no_routing_miss", - "account_id", accountID, - "reason", "account_not_in_map", - "session", shortSessionHash(sessionHash), - ) - } - } - } else if len(routingAccountIDs) == 0 && sessionHash != "" { - slog.Debug("sticky.layer1_5_no_routing_skip", - "sticky_account_id", stickyAccountID, - "is_excluded", func() bool { return stickyAccountID > 0 && isExcluded(stickyAccountID) }(), - "session", shortSessionHash(sessionHash), - "reason", func() string { - if stickyAccountID == 0 { - return "no_sticky_binding" - } - return "sticky_account_excluded" - }(), - ) - } - - // ============ Layer 2: 负载感知选择 ============ - slog.Debug("sticky.layer2_fallback", - "session", shortSessionHash(sessionHash), - "sticky_account_id", stickyAccountID, - "reason", "sticky_not_used_falling_back_to_load_balance", - "total_accounts", len(accounts), - ) - candidates := make([]*Account, 0, len(accounts)) - for i := range accounts { - acc := &accounts[i] - if isExcluded(acc.ID) { - continue - } - // Scheduler snapshots can be temporarily stale (bucket rebuild is throttled); - // re-check schedulability here so recently rate-limited/overloaded accounts - // are not selected again before the bucket is rebuilt. - if !s.isAccountSchedulableForSelection(acc) { - continue - } - if !s.isAccountAllowedForPlatform(acc, platform, useMixed) { - continue - } - if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { - continue - } - // 配额检查 - if !s.isAccountSchedulableForQuota(acc) { - continue - } - // 窗口费用检查(非粘性会话路径) - if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { - continue - } - // RPM 检查(非粘性会话路径) - if !s.isAccountSchedulableForRPM(ctx, acc, false) { - continue - } - candidates = append(candidates, acc) - } - - if len(candidates) == 0 { - return nil, ErrNoAvailableAccounts - } - - accountLoads := make([]AccountWithConcurrency, 0, len(candidates)) - for _, acc := range candidates { - accountLoads = append(accountLoads, AccountWithConcurrency{ - ID: acc.ID, - MaxConcurrency: acc.EffectiveLoadFactor(), - }) - } - - loadMap, err := s.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) - if err != nil { - if result, ok, legacyErr := s.tryAcquireByLegacyOrder(ctx, candidates, groupID, sessionHash, preferOAuth); legacyErr != nil { - return nil, legacyErr - } else if ok { - return result, nil - } - } else { - var available []accountWithLoad - for _, acc := range candidates { - loadInfo := loadMap[acc.ID] - if loadInfo == nil { - loadInfo = &AccountLoadInfo{AccountID: acc.ID} - } - if loadInfo.LoadRate < 100 { - available = append(available, accountWithLoad{ - account: acc, - loadInfo: loadInfo, - }) - } - } - - // 分层过滤选择:优先级 →(可选)最早重置 → 负载率 → LRU - for len(available) > 0 { - // 1. 取优先级最小的集合 - candidates := filterByMinPriority(available) - // 2. (可选)use-it-or-lose-it:优先选用会话窗口最早重置的账号 - if cfg.PreferSoonestReset { - candidates = filterBySoonestReset(candidates) - } - // 3. 取负载率最低的集合 - candidates = filterByMinLoadRate(candidates) - // 4. LRU 选择最久未用的账号 - selected := selectByLRU(candidates, preferOAuth) - if selected == nil { - break - } - - result, err := s.tryAcquireAccountSlot(ctx, selected.account.ID, selected.account.Concurrency) - if err == nil && result.Acquired { - // 会话数量限制检查 - if !s.checkAndRegisterSession(ctx, selected.account, sessionHash) { - result.ReleaseFunc() // 释放槽位,继续尝试下一个账号 - } else { - if sessionHash != "" && s.cache != nil { - _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.account.ID, stickySessionTTL) - } - return s.newSelectionResult(ctx, selected.account, true, result.ReleaseFunc, nil) - } - } - - // 移除已尝试的账号,重新进行分层过滤 - selectedID := selected.account.ID - newAvailable := make([]accountWithLoad, 0, len(available)-1) - for _, acc := range available { - if acc.account.ID != selectedID { - newAvailable = append(newAvailable, acc) - } - } - available = newAvailable - } - } - - // ============ Layer 3: 兜底排队 ============ - s.sortCandidatesForFallback(candidates, preferOAuth, cfg.FallbackSelectionMode) - for _, acc := range candidates { - // 会话数量限制检查(等待计划也需要占用会话配额) - if !s.checkAndRegisterSession(ctx, acc, sessionHash) { - continue // 会话限制已满,尝试下一个账号 - } - return s.newSelectionResult(ctx, acc, false, nil, &AccountWaitPlan{ - AccountID: acc.ID, - MaxConcurrency: acc.Concurrency, - Timeout: cfg.FallbackWaitTimeout, - MaxWaiting: cfg.FallbackMaxWaiting, - }) - } - return nil, ErrNoAvailableAccounts -} - -func (s *GatewayService) tryAcquireByLegacyOrder(ctx context.Context, candidates []*Account, groupID *int64, sessionHash string, preferOAuth bool) (*AccountSelectionResult, bool, error) { - ordered := append([]*Account(nil), candidates...) - sortAccountsByPriorityAndLastUsed(ordered, preferOAuth) - - for _, acc := range ordered { - result, err := s.tryAcquireAccountSlot(ctx, acc.ID, acc.Concurrency) - if err == nil && result.Acquired { - // 会话数量限制检查 - if !s.checkAndRegisterSession(ctx, acc, sessionHash) { - result.ReleaseFunc() // 释放槽位,继续尝试下一个账号 - continue - } - if sessionHash != "" && s.cache != nil { - _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, acc.ID, stickySessionTTL) - } - selection, err := s.newSelectionResult(ctx, acc, true, result.ReleaseFunc, nil) - if err != nil { - return nil, false, err - } - return selection, true, nil - } - } - - return nil, false, nil -} - -func (s *GatewayService) schedulingConfig() config.GatewaySchedulingConfig { - if s.cfg != nil { - return s.cfg.Gateway.Scheduling - } - return config.GatewaySchedulingConfig{ - StickySessionMaxWaiting: 3, - StickySessionWaitTimeout: 45 * time.Second, - FallbackWaitTimeout: 30 * time.Second, - FallbackMaxWaiting: 100, - LoadBatchEnabled: true, - SlotCleanupInterval: 30 * time.Second, - } -} - -func (s *GatewayService) withGroupContext(ctx context.Context, group *Group) context.Context { - if !IsGroupContextValid(group) { - return ctx - } - if existing, ok := ctx.Value(ctxkey.Group).(*Group); ok && existing != nil && existing.ID == group.ID && IsGroupContextValid(existing) { - return ctx - } - return context.WithValue(ctx, ctxkey.Group, group) -} - -func (s *GatewayService) groupFromContext(ctx context.Context, groupID int64) *Group { - if group, ok := ctx.Value(ctxkey.Group).(*Group); ok && IsGroupContextValid(group) && group.ID == groupID { - return group - } - return nil -} - -func (s *GatewayService) resolveGroupByID(ctx context.Context, groupID int64) (*Group, error) { - if group := s.groupFromContext(ctx, groupID); group != nil { - return group, nil - } - group, err := s.groupRepo.GetByIDLite(ctx, groupID) - if err != nil { - return nil, fmt.Errorf("get group failed: %w", err) - } - return group, nil -} - -func (s *GatewayService) ResolveGroupByID(ctx context.Context, groupID int64) (*Group, error) { - return s.resolveGroupByID(ctx, groupID) -} - -func (s *GatewayService) routingAccountIDsForRequest(ctx context.Context, groupID *int64, requestedModel string, platform string) []int64 { - if groupID == nil || requestedModel == "" || platform != PlatformAnthropic { - return nil - } - group, err := s.resolveGroupByID(ctx, *groupID) - if err != nil || group == nil { - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] resolve group failed: group_id=%v model=%s platform=%s err=%v", derefGroupID(groupID), requestedModel, platform, err) - } - return nil - } - // Preserve existing behavior: model routing only applies to anthropic groups. - if group.Platform != PlatformAnthropic { - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] skip: non-anthropic group platform: group_id=%d group_platform=%s model=%s", group.ID, group.Platform, requestedModel) - } - return nil - } - ids := group.GetRoutingAccountIDs(requestedModel) - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routing lookup: group_id=%d model=%s enabled=%v rules=%d matched_ids=%v", - group.ID, requestedModel, group.ModelRoutingEnabled, len(group.ModelRouting), ids) - } - return ids -} - -func (s *GatewayService) resolveGatewayGroup(ctx context.Context, groupID *int64) (*Group, *int64, error) { - if groupID == nil { - return nil, nil, nil - } - - currentID := *groupID - visited := map[int64]struct{}{} - for { - if _, seen := visited[currentID]; seen { - return nil, nil, fmt.Errorf("fallback group cycle detected") - } - visited[currentID] = struct{}{} - - group, err := s.resolveGroupByID(ctx, currentID) - if err != nil { - return nil, nil, err - } - - if !group.ClaudeCodeOnly || IsClaudeCodeClient(ctx) { - return group, ¤tID, nil - } - - if group.FallbackGroupID == nil { - return nil, nil, ErrClaudeCodeOnly - } - currentID = *group.FallbackGroupID - } -} - -// checkClaudeCodeRestriction 检查分组的 Claude Code 客户端限制 -// 如果分组启用了 claude_code_only 且请求不是来自 Claude Code 客户端: -// - 有降级分组:返回降级分组的 ID -// - 无降级分组:返回 ErrClaudeCodeOnly 错误 -func (s *GatewayService) checkClaudeCodeRestriction(ctx context.Context, groupID *int64) (*Group, *int64, error) { - if groupID == nil { - return nil, groupID, nil - } - - // 强制平台模式不检查 Claude Code 限制 - if forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string); hasForcePlatform && forcePlatform != "" { - return nil, groupID, nil - } - - group, resolvedID, err := s.resolveGatewayGroup(ctx, groupID) - if err != nil { - return nil, nil, err - } - - return group, resolvedID, nil -} - -func (s *GatewayService) resolvePlatform(ctx context.Context, groupID *int64, group *Group) (string, bool, error) { - forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string) - if hasForcePlatform && forcePlatform != "" { - return forcePlatform, true, nil - } - if group != nil { - return group.Platform, false, nil - } - if groupID != nil { - group, err := s.resolveGroupByID(ctx, *groupID) - if err != nil { - return "", false, err - } - return group.Platform, false, nil - } - return PlatformAnthropic, false, nil -} - -func (s *GatewayService) listSchedulableAccounts(ctx context.Context, groupID *int64, platform string, hasForcePlatform bool) ([]Account, bool, error) { - if s.schedulerSnapshot != nil { - accounts, useMixed, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) - if err == nil { - slog.Debug("account_scheduling_list_snapshot", - "group_id", derefGroupID(groupID), - "platform", platform, - "use_mixed", useMixed, - "count", len(accounts)) - if slog.Default().Enabled(ctx, slog.LevelDebug) { - for _, acc := range accounts { - slog.Debug("account_scheduling_account_detail", - "account_id", acc.ID, - "name", acc.Name, - "platform", acc.Platform, - "type", acc.Type, - "status", acc.Status, - "tls_fingerprint", acc.IsTLSFingerprintEnabled()) - } - } - } - return accounts, useMixed, err - } - useMixed := (platform == PlatformAnthropic || platform == PlatformGemini) && !hasForcePlatform - if useMixed { - platforms := []string{platform, PlatformAntigravity} - var accounts []Account - var err error - if groupID != nil { - accounts, err = s.accountRepo.ListSchedulableByGroupIDAndPlatforms(ctx, *groupID, platforms) - } else if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { - accounts, err = s.accountRepo.ListSchedulableByPlatforms(ctx, platforms) - } else { - accounts, err = s.accountRepo.ListSchedulableUngroupedByPlatforms(ctx, platforms) - } - if err != nil { - slog.Debug("account_scheduling_list_failed", - "group_id", derefGroupID(groupID), - "platform", platform, - "error", err) - return nil, useMixed, err - } - filtered := make([]Account, 0, len(accounts)) - for _, acc := range accounts { - if acc.Platform == PlatformAntigravity && !acc.IsMixedSchedulingEnabled() { - continue - } - filtered = append(filtered, acc) - } - slog.Debug("account_scheduling_list_mixed", - "group_id", derefGroupID(groupID), - "platform", platform, - "raw_count", len(accounts), - "filtered_count", len(filtered)) - if slog.Default().Enabled(ctx, slog.LevelDebug) { - for _, acc := range filtered { - slog.Debug("account_scheduling_account_detail", - "account_id", acc.ID, - "name", acc.Name, - "platform", acc.Platform, - "type", acc.Type, - "status", acc.Status, - "tls_fingerprint", acc.IsTLSFingerprintEnabled()) - } - } - return filtered, useMixed, nil - } - - var accounts []Account - var err error - if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { - accounts, err = s.accountRepo.ListSchedulableByPlatform(ctx, platform) - } else if groupID != nil { - accounts, err = s.accountRepo.ListSchedulableByGroupIDAndPlatform(ctx, *groupID, platform) - // 分组内无账号则返回空列表,由上层处理错误,不再回退到全平台查询 - } else { - accounts, err = s.accountRepo.ListSchedulableUngroupedByPlatform(ctx, platform) - } - if err != nil { - slog.Debug("account_scheduling_list_failed", - "group_id", derefGroupID(groupID), - "platform", platform, - "error", err) - return nil, useMixed, err - } - slog.Debug("account_scheduling_list_single", - "group_id", derefGroupID(groupID), - "platform", platform, - "count", len(accounts)) - if slog.Default().Enabled(ctx, slog.LevelDebug) { - for _, acc := range accounts { - slog.Debug("account_scheduling_account_detail", - "account_id", acc.ID, - "name", acc.Name, - "platform", acc.Platform, - "type", acc.Type, - "status", acc.Status, - "tls_fingerprint", acc.IsTLSFingerprintEnabled()) - } - } - return accounts, useMixed, nil -} - -// IsSingleAntigravityAccountGroup 检查指定分组是否只有一个 antigravity 平台的可调度账号。 -// 用于 Handler 层在首次请求时提前设置 SingleAccountRetry context, -// 避免单账号分组收到 503 时错误地设置模型限流标记导致后续请求连续快速失败。 -func (s *GatewayService) IsSingleAntigravityAccountGroup(ctx context.Context, groupID *int64) bool { - accounts, _, err := s.listSchedulableAccounts(ctx, groupID, PlatformAntigravity, true) - if err != nil { - return false - } - return len(accounts) == 1 -} - -func (s *GatewayService) isAccountAllowedForPlatform(account *Account, platform string, useMixed bool) bool { - if account == nil { - return false - } - if useMixed { - if account.Platform == platform { - return true - } - return account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled() - } - return account.Platform == platform -} - -func (s *GatewayService) isAccountSchedulableForSelection(account *Account) bool { - if account == nil { - return false - } - return account.IsSchedulable() -} - -func (s *GatewayService) isAccountSchedulableForModelSelection(ctx context.Context, account *Account, requestedModel string) bool { - if account == nil { - return false - } - return account.IsSchedulableForModelWithContext(ctx, requestedModel) -} - -// isAccountInGroup checks if the account belongs to the specified group. -// When groupID is nil, returns true only for ungrouped accounts (no group assignments). -func (s *GatewayService) isAccountInGroup(account *Account, groupID *int64) bool { - if account == nil { - return false - } - if groupID == nil { - // 无分组的 API Key 只能使用未分组的账号 - return len(account.AccountGroups) == 0 - } - for _, ag := range account.AccountGroups { - if ag.GroupID == *groupID { - return true - } - } - return false -} - -func (s *GatewayService) tryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (*AcquireResult, error) { - if s.concurrencyService == nil { - return &AcquireResult{Acquired: true, ReleaseFunc: func() {}}, nil - } - return s.concurrencyService.AcquireAccountSlot(ctx, accountID, maxConcurrency) -} - -type usageLogWindowStatsBatchProvider interface { - GetAccountWindowStatsBatch(ctx context.Context, accountIDs []int64, startTime time.Time) (map[int64]*usagestats.AccountStats, error) -} - -type windowCostPrefetchContextKeyType struct{} - -var windowCostPrefetchContextKey = windowCostPrefetchContextKeyType{} - -func windowCostFromPrefetchContext(ctx context.Context, accountID int64) (float64, bool) { - if ctx == nil || accountID <= 0 { - return 0, false - } - m, ok := ctx.Value(windowCostPrefetchContextKey).(map[int64]float64) - if !ok || len(m) == 0 { - return 0, false - } - v, exists := m[accountID] - return v, exists -} - -func (s *GatewayService) withWindowCostPrefetch(ctx context.Context, accounts []Account) context.Context { - if ctx == nil || len(accounts) == 0 || s.sessionLimitCache == nil || s.usageLogRepo == nil { - return ctx - } - - accountByID := make(map[int64]*Account) - accountIDs := make([]int64, 0, len(accounts)) - for i := range accounts { - account := &accounts[i] - if account == nil || !account.IsAnthropicOAuthOrSetupToken() { - continue - } - if account.GetWindowCostLimit() <= 0 { - continue - } - accountByID[account.ID] = account - accountIDs = append(accountIDs, account.ID) - } - if len(accountIDs) == 0 { - return ctx - } - - costs := make(map[int64]float64, len(accountIDs)) - cacheValues, err := s.sessionLimitCache.GetWindowCostBatch(ctx, accountIDs) - if err == nil { - for accountID, cost := range cacheValues { - costs[accountID] = cost - } - windowCostPrefetchCacheHitTotal.Add(int64(len(cacheValues))) - } else { - windowCostPrefetchErrorTotal.Add(1) - logger.LegacyPrintf("service.gateway", "window_cost batch cache read failed: %v", err) - } - cacheMissCount := len(accountIDs) - len(costs) - if cacheMissCount < 0 { - cacheMissCount = 0 - } - windowCostPrefetchCacheMissTotal.Add(int64(cacheMissCount)) - - missingByStart := make(map[int64][]int64) - startTimes := make(map[int64]time.Time) - for _, accountID := range accountIDs { - if _, ok := costs[accountID]; ok { - continue - } - account := accountByID[accountID] - if account == nil { - continue - } - startTime := account.GetCurrentWindowStartTime() - startKey := startTime.Unix() - missingByStart[startKey] = append(missingByStart[startKey], accountID) - startTimes[startKey] = startTime - } - if len(missingByStart) == 0 { - return context.WithValue(ctx, windowCostPrefetchContextKey, costs) - } - - batchReader, hasBatch := s.usageLogRepo.(usageLogWindowStatsBatchProvider) - for startKey, ids := range missingByStart { - startTime := startTimes[startKey] - - if hasBatch { - windowCostPrefetchBatchSQLTotal.Add(1) - queryStart := time.Now() - statsByAccount, err := batchReader.GetAccountWindowStatsBatch(ctx, ids, startTime) - if err == nil { - slog.Debug("window_cost_batch_query_ok", - "accounts", len(ids), - "window_start", startTime.Format(time.RFC3339), - "duration_ms", time.Since(queryStart).Milliseconds()) - for _, accountID := range ids { - stats := statsByAccount[accountID] - cost := 0.0 - if stats != nil { - cost = stats.StandardCost - } - costs[accountID] = cost - _ = s.sessionLimitCache.SetWindowCost(ctx, accountID, cost) - } - continue - } - windowCostPrefetchErrorTotal.Add(1) - logger.LegacyPrintf("service.gateway", "window_cost batch db query failed: start=%s err=%v", startTime.Format(time.RFC3339), err) - } - - // 回退路径:缺少批量仓储能力或批量查询失败时,按账号单查(失败开放)。 - windowCostPrefetchFallbackTotal.Add(int64(len(ids))) - for _, accountID := range ids { - stats, err := s.usageLogRepo.GetAccountWindowStats(ctx, accountID, startTime) - if err != nil { - windowCostPrefetchErrorTotal.Add(1) - continue - } - cost := stats.StandardCost - costs[accountID] = cost - _ = s.sessionLimitCache.SetWindowCost(ctx, accountID, cost) - } - } - - return context.WithValue(ctx, windowCostPrefetchContextKey, costs) -} - -// isAccountSchedulableForQuota 检查账号是否在配额限制内 -// 适用于配置了 quota_limit 的 apikey 和 bedrock 类型账号 -func (s *GatewayService) isAccountSchedulableForQuota(account *Account) bool { - if !account.IsAPIKeyOrBedrock() { - return true - } - return !account.IsQuotaExceeded() -} - -// isAccountSchedulableForWindowCost 检查账号是否可根据窗口费用进行调度 -// 仅适用于 Anthropic OAuth/SetupToken 账号 -// 返回 true 表示可调度,false 表示不可调度 -func (s *GatewayService) isAccountSchedulableForWindowCost(ctx context.Context, account *Account, isSticky bool) bool { - // 只检查 Anthropic OAuth/SetupToken 账号 - if !account.IsAnthropicOAuthOrSetupToken() { - return true - } - - limit := account.GetWindowCostLimit() - if limit <= 0 { - return true // 未启用窗口费用限制 - } - - // 尝试从缓存获取窗口费用 - var currentCost float64 - if cost, ok := windowCostFromPrefetchContext(ctx, account.ID); ok { - currentCost = cost - goto checkSchedulability - } - if s.sessionLimitCache != nil { - if cost, hit, err := s.sessionLimitCache.GetWindowCost(ctx, account.ID); err == nil && hit { - currentCost = cost - goto checkSchedulability - } - } - - // 缓存未命中,从数据库查询 - { - // 使用统一的窗口开始时间计算逻辑(考虑窗口过期情况) - startTime := account.GetCurrentWindowStartTime() - - stats, err := s.usageLogRepo.GetAccountWindowStats(ctx, account.ID, startTime) - if err != nil { - // 失败开放:查询失败时允许调度 - return true - } - - // 使用标准费用(不含账号倍率) - currentCost = stats.StandardCost - - // 设置缓存(忽略错误) - if s.sessionLimitCache != nil { - _ = s.sessionLimitCache.SetWindowCost(ctx, account.ID, currentCost) - } - } - -checkSchedulability: - schedulability := account.CheckWindowCostSchedulability(currentCost) - - switch schedulability { - case WindowCostSchedulable: - return true - case WindowCostStickyOnly: - return isSticky - case WindowCostNotSchedulable: - return false - } - return true -} - -// rpmPrefetchContextKey is the context key for prefetched RPM counts. -type rpmPrefetchContextKeyType struct{} - -var rpmPrefetchContextKey = rpmPrefetchContextKeyType{} - -func rpmFromPrefetchContext(ctx context.Context, accountID int64) (int, bool) { - if v, ok := ctx.Value(rpmPrefetchContextKey).(map[int64]int); ok { - count, found := v[accountID] - return count, found - } - return 0, false -} - -// withRPMPrefetch 批量预取所有候选账号的 RPM 计数 -func (s *GatewayService) withRPMPrefetch(ctx context.Context, accounts []Account) context.Context { - if s.rpmCache == nil { - return ctx - } - - var ids []int64 - for i := range accounts { - if accounts[i].IsAnthropicOAuthOrSetupToken() && accounts[i].GetBaseRPM() > 0 { - ids = append(ids, accounts[i].ID) - } - } - if len(ids) == 0 { - return ctx - } - - counts, err := s.rpmCache.GetRPMBatch(ctx, ids) - if err != nil { - return ctx // 失败开放 - } - return context.WithValue(ctx, rpmPrefetchContextKey, counts) -} - -// isAccountSchedulableForRPM 检查账号是否可根据 RPM 进行调度 -// 仅适用于 Anthropic OAuth/SetupToken 账号 -func (s *GatewayService) isAccountSchedulableForRPM(ctx context.Context, account *Account, isSticky bool) bool { - if !account.IsAnthropicOAuthOrSetupToken() { - return true - } - baseRPM := account.GetBaseRPM() - if baseRPM <= 0 { - return true - } - - // 尝试从预取缓存获取 - var currentRPM int - if count, ok := rpmFromPrefetchContext(ctx, account.ID); ok { - currentRPM = count - } else if s.rpmCache != nil { - if count, err := s.rpmCache.GetRPM(ctx, account.ID); err == nil { - currentRPM = count - } - // 失败开放:GetRPM 错误时允许调度 - } - - schedulability := account.CheckRPMSchedulability(currentRPM) - switch schedulability { - case WindowCostSchedulable: - return true - case WindowCostStickyOnly: - return isSticky - case WindowCostNotSchedulable: - return false - } - return true -} - -// IncrementAccountRPM increments the RPM counter for the given account. -// 已知 TOCTOU 竞态:调度时读取 RPM 计数与此处递增之间存在时间窗口, -// 高并发下可能短暂超出 RPM 限制。这是与 WindowCost 一致的 soft-limit -// 设计权衡——可接受的少量超额优于加锁带来的延迟和复杂度。 -func (s *GatewayService) IncrementAccountRPM(ctx context.Context, accountID int64) error { - if s.rpmCache == nil { - return nil - } - _, err := s.rpmCache.IncrementRPM(ctx, accountID) - return err -} - -// checkAndRegisterSession 检查并注册会话,用于会话数量限制 -// 仅适用于 Anthropic OAuth/SetupToken 账号 -// sessionID: 会话标识符(使用粘性会话的 hash) -// 返回 true 表示允许(在限制内或会话已存在),false 表示拒绝(超出限制且是新会话) -func (s *GatewayService) checkAndRegisterSession(ctx context.Context, account *Account, sessionID string) bool { - // 只检查 Anthropic OAuth/SetupToken 账号 - if !account.IsAnthropicOAuthOrSetupToken() { - return true - } - - maxSessions := account.GetMaxSessions() - if maxSessions <= 0 || sessionID == "" { - return true // 未启用会话限制或无会话ID - } - - if s.sessionLimitCache == nil { - return true // 缓存不可用时允许通过 - } - - idleTimeout := time.Duration(account.GetSessionIdleTimeoutMinutes()) * time.Minute - - allowed, err := s.sessionLimitCache.RegisterSession(ctx, account.ID, sessionID, maxSessions, idleTimeout) - if err != nil { - // 失败开放:缓存错误时允许通过 - return true - } - return allowed -} - -func (s *GatewayService) getSchedulableAccount(ctx context.Context, accountID int64) (*Account, error) { - if s.schedulerSnapshot != nil { - return s.schedulerSnapshot.GetAccount(ctx, accountID) - } - return s.accountRepo.GetByID(ctx, accountID) -} - -func (s *GatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { - if account == nil || s.schedulerSnapshot == nil { - return account, nil - } - hydrated, err := s.schedulerSnapshot.GetAccount(ctx, account.ID) - if err != nil { - return nil, err - } - if hydrated == nil { - return nil, fmt.Errorf("selected gateway account %d not found during hydration", account.ID) - } - return hydrated, nil -} - -func (s *GatewayService) newSelectionResult(ctx context.Context, account *Account, acquired bool, release func(), waitPlan *AccountWaitPlan) (*AccountSelectionResult, error) { - hydrated, err := s.hydrateSelectedAccount(ctx, account) - if err != nil { - return nil, err - } - return &AccountSelectionResult{ - Account: hydrated, - Acquired: acquired, - ReleaseFunc: release, - WaitPlan: waitPlan, - }, nil -} - -// filterByMinPriority 过滤出优先级最小的账号集合 -func filterByMinPriority(accounts []accountWithLoad) []accountWithLoad { - if len(accounts) == 0 { - return accounts - } - minPriority := accounts[0].account.Priority - for _, acc := range accounts[1:] { - if acc.account.Priority < minPriority { - minPriority = acc.account.Priority - } - } - result := make([]accountWithLoad, 0, len(accounts)) - for _, acc := range accounts { - if acc.account.Priority == minPriority { - result = append(result, acc) - } - } - return result -} - -// filterByMinLoadRate 过滤出负载率最低的账号集合 -func filterByMinLoadRate(accounts []accountWithLoad) []accountWithLoad { - if len(accounts) == 0 { - return accounts - } - minLoadRate := accounts[0].loadInfo.LoadRate - for _, acc := range accounts[1:] { - if acc.loadInfo.LoadRate < minLoadRate { - minLoadRate = acc.loadInfo.LoadRate - } - } - result := make([]accountWithLoad, 0, len(accounts)) - for _, acc := range accounts { - if acc.loadInfo.LoadRate == minLoadRate { - result = append(result, acc) - } - } - return result -} - -// filterBySoonestReset 过滤出「会话窗口最早重置」的账号集合(use-it-or-lose-it)。 -// 仅保留拥有未来重置时间(SessionWindowEnd 在当前时间之后)且最早的账号; -// 窗口为空或已过期的账号视为无活跃窗口、优先级最低。 -// 当所有账号都没有活跃窗口时,返回原集合(不改变后续 LRU 选择)。 -func filterBySoonestReset(accounts []accountWithLoad) []accountWithLoad { - if len(accounts) <= 1 { - return accounts - } - now := time.Now() - var minEnd *time.Time - for _, acc := range accounts { - end := acc.account.SessionWindowEnd - if end == nil || !now.Before(*end) { - continue - } - if minEnd == nil || end.Before(*minEnd) { - minEnd = end - } - } - if minEnd == nil { - // 没有任何账号拥有活跃窗口,保持原集合 - return accounts - } - result := make([]accountWithLoad, 0, len(accounts)) - for _, acc := range accounts { - end := acc.account.SessionWindowEnd - if end != nil && now.Before(*end) && end.Equal(*minEnd) { - result = append(result, acc) - } - } - return result -} - -// selectByLRU 从集合中选择最久未用的账号 -// 如果有多个账号具有相同的最小 LastUsedAt,则随机选择一个 -func selectByLRU(accounts []accountWithLoad, preferOAuth bool) *accountWithLoad { - if len(accounts) == 0 { - return nil - } - if len(accounts) == 1 { - return &accounts[0] - } - - // 1. 找到最小的 LastUsedAt(nil 被视为最小) - var minTime *time.Time - hasNil := false - for _, acc := range accounts { - if acc.account.LastUsedAt == nil { - hasNil = true - break - } - if minTime == nil || acc.account.LastUsedAt.Before(*minTime) { - minTime = acc.account.LastUsedAt - } - } - - // 2. 收集所有具有最小 LastUsedAt 的账号索引 - var candidateIdxs []int - for i, acc := range accounts { - if hasNil { - if acc.account.LastUsedAt == nil { - candidateIdxs = append(candidateIdxs, i) - } - } else { - if acc.account.LastUsedAt != nil && acc.account.LastUsedAt.Equal(*minTime) { - candidateIdxs = append(candidateIdxs, i) - } - } - } - - // 3. 如果只有一个候选,直接返回 - if len(candidateIdxs) == 1 { - return &accounts[candidateIdxs[0]] - } - - // 4. 如果有多个候选且 preferOAuth,优先选择 OAuth 类型 - if preferOAuth { - var oauthIdxs []int - for _, idx := range candidateIdxs { - if accounts[idx].account.Type == AccountTypeOAuth { - oauthIdxs = append(oauthIdxs, idx) - } - } - if len(oauthIdxs) > 0 { - candidateIdxs = oauthIdxs - } - } - - // 5. 随机选择一个 - selectedIdx := candidateIdxs[mathrand.Intn(len(candidateIdxs))] - return &accounts[selectedIdx] -} - -func sortAccountsByPriorityAndLastUsed(accounts []*Account, preferOAuth bool) { - sort.SliceStable(accounts, func(i, j int) bool { - a, b := accounts[i], accounts[j] - if a.Priority != b.Priority { - return a.Priority < b.Priority - } - switch { - case a.LastUsedAt == nil && b.LastUsedAt != nil: - return true - case a.LastUsedAt != nil && b.LastUsedAt == nil: - return false - case a.LastUsedAt == nil && b.LastUsedAt == nil: - if preferOAuth && a.Type != b.Type { - return a.Type == AccountTypeOAuth - } - return false - default: - return a.LastUsedAt.Before(*b.LastUsedAt) - } - }) - shuffleWithinPriorityAndLastUsed(accounts, preferOAuth) -} - -// shuffleWithinSortGroups 对排序后的 accountWithLoad 切片,按 (Priority, LoadRate, LastUsedAt) 分组后组内随机打乱。 -// 防止并发请求读取同一快照时,确定性排序导致所有请求命中相同账号。 -func shuffleWithinSortGroups(accounts []accountWithLoad) { - if len(accounts) <= 1 { - return - } - i := 0 - for i < len(accounts) { - j := i + 1 - for j < len(accounts) && sameAccountWithLoadGroup(accounts[i], accounts[j]) { - j++ - } - if j-i > 1 { - mathrand.Shuffle(j-i, func(a, b int) { - accounts[i+a], accounts[i+b] = accounts[i+b], accounts[i+a] - }) - } - i = j - } -} - -// sameAccountWithLoadGroup 判断两个 accountWithLoad 是否属于同一排序组 -func sameAccountWithLoadGroup(a, b accountWithLoad) bool { - if a.account.Priority != b.account.Priority { - return false - } - if a.loadInfo.LoadRate != b.loadInfo.LoadRate { - return false - } - return sameLastUsedAt(a.account.LastUsedAt, b.account.LastUsedAt) -} - -// shuffleWithinPriorityAndLastUsed 对排序后的 []*Account 切片,按 (Priority, LastUsedAt) 分组后组内随机打乱。 -// -// 注意:当 preferOAuth=true 时,需要保证 OAuth 账号在同组内仍然优先,否则会把排序时的偏好打散掉。 -// 因此这里采用"组内分区 + 分区内 shuffle"的方式: -// - 先把同组账号按 (OAuth / 非 OAuth) 拆成两段,保持 OAuth 段在前; -// - 再分别在各段内随机打散,避免热点。 -func shuffleWithinPriorityAndLastUsed(accounts []*Account, preferOAuth bool) { - if len(accounts) <= 1 { - return - } - i := 0 - for i < len(accounts) { - j := i + 1 - for j < len(accounts) && sameAccountGroup(accounts[i], accounts[j]) { - j++ - } - if j-i > 1 { - if preferOAuth { - oauth := make([]*Account, 0, j-i) - others := make([]*Account, 0, j-i) - for _, acc := range accounts[i:j] { - if acc.Type == AccountTypeOAuth { - oauth = append(oauth, acc) - } else { - others = append(others, acc) - } - } - if len(oauth) > 1 { - mathrand.Shuffle(len(oauth), func(a, b int) { oauth[a], oauth[b] = oauth[b], oauth[a] }) - } - if len(others) > 1 { - mathrand.Shuffle(len(others), func(a, b int) { others[a], others[b] = others[b], others[a] }) - } - copy(accounts[i:], oauth) - copy(accounts[i+len(oauth):], others) - } else { - mathrand.Shuffle(j-i, func(a, b int) { - accounts[i+a], accounts[i+b] = accounts[i+b], accounts[i+a] - }) - } - } - i = j - } -} - -// sameAccountGroup 判断两个 Account 是否属于同一排序组(Priority + LastUsedAt) -func sameAccountGroup(a, b *Account) bool { - if a.Priority != b.Priority { - return false - } - return sameLastUsedAt(a.LastUsedAt, b.LastUsedAt) -} - -// sameLastUsedAt 判断两个 LastUsedAt 是否相同(精度到秒) -func sameLastUsedAt(a, b *time.Time) bool { - switch { - case a == nil && b == nil: - return true - case a == nil || b == nil: - return false - default: - return a.Unix() == b.Unix() - } -} - -// sortCandidatesForFallback 根据配置选择排序策略 -// mode: "last_used"(按最后使用时间) 或 "random"(随机) -func (s *GatewayService) sortCandidatesForFallback(accounts []*Account, preferOAuth bool, mode string) { - if mode == "random" { - // 先按优先级排序,然后在同优先级内随机打乱 - sortAccountsByPriorityOnly(accounts, preferOAuth) - shuffleWithinPriority(accounts) - } else { - // 默认按最后使用时间排序 - sortAccountsByPriorityAndLastUsed(accounts, preferOAuth) - } -} - -// sortAccountsByPriorityOnly 仅按优先级排序 -func sortAccountsByPriorityOnly(accounts []*Account, preferOAuth bool) { - sort.SliceStable(accounts, func(i, j int) bool { - a, b := accounts[i], accounts[j] - if a.Priority != b.Priority { - return a.Priority < b.Priority - } - if preferOAuth && a.Type != b.Type { - return a.Type == AccountTypeOAuth - } - return false - }) -} - -// shuffleWithinPriority 在同优先级内随机打乱顺序 -func shuffleWithinPriority(accounts []*Account) { - if len(accounts) <= 1 { - return - } - r := mathrand.New(mathrand.NewSource(time.Now().UnixNano())) - start := 0 - for start < len(accounts) { - priority := accounts[start].Priority - end := start + 1 - for end < len(accounts) && accounts[end].Priority == priority { - end++ - } - // 对 [start, end) 范围内的账户随机打乱 - if end-start > 1 { - r.Shuffle(end-start, func(i, j int) { - accounts[start+i], accounts[start+j] = accounts[start+j], accounts[start+i] - }) - } - start = end - } -} - -// selectAccountForModelWithPlatform 选择单平台账户(完全隔离) -func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, platform string) (*Account, error) { - preferOAuth := platform == PlatformGemini - routingAccountIDs := s.routingAccountIDsForRequest(ctx, groupID, requestedModel, platform) - - // require_privacy_set: 获取分组信息 - var schedGroup *Group - if groupID != nil && s.groupRepo != nil { - schedGroup, _ = s.groupRepo.GetByID(ctx, *groupID) - } - - var accounts []Account - accountsLoaded := false - - // ============ Model Routing (legacy path): apply before sticky session ============ - // When load-awareness is disabled (e.g. concurrency service not configured), we still honor model routing - // so switching model can switch upstream account within the same sticky session. - if len(routingAccountIDs) > 0 { - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy routed begin: group_id=%v model=%s platform=%s session=%s routed_ids=%v", - derefGroupID(groupID), requestedModel, platform, shortSessionHash(sessionHash), routingAccountIDs) - } - // 1) Sticky session only applies if the bound account is within the routing set. - if sessionHash != "" && s.cache != nil { - accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - if err == nil && accountID > 0 && containsInt64(routingAccountIDs, accountID) { - if _, excluded := excludedIDs[accountID]; !excluded { - account, err := s.getSchedulableAccount(ctx, accountID) - // 检查账号分组归属和平台匹配(确保粘性会话不会跨分组或跨平台) - if err == nil { - clearSticky := shouldClearStickySession(account, requestedModel) - if clearSticky { - _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - } - if !clearSticky && s.isAccountInGroup(account, groupID) && account.Platform == platform && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) && !s.isStickyAccountUpstreamRestricted(ctx, groupID, account, requestedModel) { - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), accountID) - } - return account, nil - } - } - } - } - } - - // 2) Select an account from the routed candidates. - forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string) - if hasForcePlatform && forcePlatform == "" { - hasForcePlatform = false - } - var err error - accounts, _, err = s.listSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) - if err != nil { - return nil, fmt.Errorf("query accounts failed: %w", err) - } - accountsLoaded = true - - // 提前预取窗口费用+RPM 计数,确保 routing 段内的调度检查调用能命中缓存 - ctx = s.withWindowCostPrefetch(ctx, accounts) - ctx = s.withRPMPrefetch(ctx, accounts) - - routingSet := make(map[int64]struct{}, len(routingAccountIDs)) - for _, id := range routingAccountIDs { - if id > 0 { - routingSet[id] = struct{}{} - } - } - - var selected *Account - for i := range accounts { - acc := &accounts[i] - if _, ok := routingSet[acc.ID]; !ok { - continue - } - if _, excluded := excludedIDs[acc.ID]; excluded { - continue - } - // Scheduler snapshots can be temporarily stale; re-check schedulability here to - // avoid selecting accounts that were recently rate-limited/overloaded. - if !s.isAccountSchedulableForSelection(acc) { - continue - } - // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 - if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { - _ = s.accountRepo.SetError(ctx, acc.ID, - fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) - continue - } - if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForQuota(acc) { - continue - } - if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { - continue - } - if !s.isAccountSchedulableForRPM(ctx, acc, false) { - continue - } - if selected == nil { - selected = acc - continue - } - if acc.Priority < selected.Priority { - selected = acc - } else if acc.Priority == selected.Priority { - switch { - case acc.LastUsedAt == nil && selected.LastUsedAt != nil: - selected = acc - case acc.LastUsedAt != nil && selected.LastUsedAt == nil: - // keep selected (never used is preferred) - case acc.LastUsedAt == nil && selected.LastUsedAt == nil: - if preferOAuth && acc.Type != selected.Type && acc.Type == AccountTypeOAuth { - selected = acc - } - default: - if acc.LastUsedAt.Before(*selected.LastUsedAt) { - selected = acc - } - } - } - } - - if selected != nil { - if sessionHash != "" && s.cache != nil { - if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { - logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) - } - } - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), selected.ID) - } - return selected, nil - } - logger.LegacyPrintf("service.gateway", "[ModelRouting] No routed accounts available for model=%s, falling back to normal selection", requestedModel) - } - - // 1. 查询粘性会话 - if sessionHash != "" && s.cache != nil { - accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - if err == nil && accountID > 0 { - if _, excluded := excludedIDs[accountID]; !excluded { - account, err := s.getSchedulableAccount(ctx, accountID) - // 检查账号分组归属和平台匹配(确保粘性会话不会跨分组或跨平台) - if err == nil { - clearSticky := shouldClearStickySession(account, requestedModel) - if clearSticky { - _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - } - if !clearSticky && s.isAccountInGroup(account, groupID) && account.Platform == platform && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) { - return account, nil - } - } - } - } - } - - // 2. 获取可调度账号列表(单平台) - if !accountsLoaded { - forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string) - if hasForcePlatform && forcePlatform == "" { - hasForcePlatform = false - } - var err error - accounts, _, err = s.listSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) - if err != nil { - return nil, fmt.Errorf("query accounts failed: %w", err) - } - } - - // 批量预取窗口费用+RPM 计数,避免逐个账号查询(N+1) - ctx = s.withWindowCostPrefetch(ctx, accounts) - ctx = s.withRPMPrefetch(ctx, accounts) - - // 3. 按优先级+最久未用选择(考虑模型支持) - // needsUpstreamCheck 仅在主选择循环中使用;粘性会话命中时跳过此检查, - // 因为粘性会话优先保持连接一致性,且 upstream 计费基准极少使用。 - needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) - var selected *Account - for i := range accounts { - acc := &accounts[i] - if _, excluded := excludedIDs[acc.ID]; excluded { - continue - } - // Scheduler snapshots can be temporarily stale; re-check schedulability here to - // avoid selecting accounts that were recently rate-limited/overloaded. - if !s.isAccountSchedulableForSelection(acc) { - continue - } - // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 - if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { - _ = s.accountRepo.SetError(ctx, acc.ID, - fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) - continue - } - if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { - continue - } - if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForQuota(acc) { - continue - } - if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { - continue - } - if !s.isAccountSchedulableForRPM(ctx, acc, false) { - continue - } - if selected == nil { - selected = acc - continue - } - if acc.Priority < selected.Priority { - selected = acc - } else if acc.Priority == selected.Priority { - switch { - case acc.LastUsedAt == nil && selected.LastUsedAt != nil: - selected = acc - case acc.LastUsedAt != nil && selected.LastUsedAt == nil: - // keep selected (never used is preferred) - case acc.LastUsedAt == nil && selected.LastUsedAt == nil: - if preferOAuth && acc.Type != selected.Type && acc.Type == AccountTypeOAuth { - selected = acc - } - default: - if acc.LastUsedAt.Before(*selected.LastUsedAt) { - selected = acc - } - } - } - } - - if selected == nil { - stats := s.logDetailedSelectionFailure(ctx, groupID, sessionHash, requestedModel, platform, accounts, excludedIDs, false) - if requestedModel != "" { - return nil, fmt.Errorf("%w supporting model: %s (%s)", ErrNoAvailableAccounts, requestedModel, summarizeSelectionFailureStats(stats)) - } - return nil, ErrNoAvailableAccounts - } - - // 4. 建立粘性绑定 - if sessionHash != "" && s.cache != nil { - if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { - logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) - } - } - - return selected, nil -} - -// selectAccountWithMixedScheduling 选择账户(支持混合调度) -// 查询原生平台账户 + 启用 mixed_scheduling 的 antigravity 账户 -func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, nativePlatform string) (*Account, error) { - preferOAuth := nativePlatform == PlatformGemini - routingAccountIDs := s.routingAccountIDsForRequest(ctx, groupID, requestedModel, nativePlatform) - - // require_privacy_set: 获取分组信息 - var schedGroup *Group - if groupID != nil && s.groupRepo != nil { - schedGroup, _ = s.groupRepo.GetByID(ctx, *groupID) - } - - var accounts []Account - accountsLoaded := false - - // ============ Model Routing (legacy path): apply before sticky session ============ - if len(routingAccountIDs) > 0 { - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy mixed routed begin: group_id=%v model=%s platform=%s session=%s routed_ids=%v", - derefGroupID(groupID), requestedModel, nativePlatform, shortSessionHash(sessionHash), routingAccountIDs) - } - // 1) Sticky session only applies if the bound account is within the routing set. - if sessionHash != "" && s.cache != nil { - accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - if err == nil && accountID > 0 && containsInt64(routingAccountIDs, accountID) { - if _, excluded := excludedIDs[accountID]; !excluded { - account, err := s.getSchedulableAccount(ctx, accountID) - // 检查账号分组归属和有效性:原生平台直接匹配,antigravity 需要启用混合调度 - if err == nil { - clearSticky := shouldClearStickySession(account, requestedModel) - if clearSticky { - _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - } - if !clearSticky && s.isAccountInGroup(account, groupID) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) { - if account.Platform == nativePlatform || (account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled()) { - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy mixed routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), accountID) - } - return account, nil - } - } - } - } - } - } - - // 2) Select an account from the routed candidates. - var err error - accounts, _, err = s.listSchedulableAccounts(ctx, groupID, nativePlatform, false) - if err != nil { - return nil, fmt.Errorf("query accounts failed: %w", err) - } - accountsLoaded = true - - // 提前预取窗口费用+RPM 计数,确保 routing 段内的调度检查调用能命中缓存 - ctx = s.withWindowCostPrefetch(ctx, accounts) - ctx = s.withRPMPrefetch(ctx, accounts) - - routingSet := make(map[int64]struct{}, len(routingAccountIDs)) - for _, id := range routingAccountIDs { - if id > 0 { - routingSet[id] = struct{}{} - } - } - - var selected *Account - for i := range accounts { - acc := &accounts[i] - if _, ok := routingSet[acc.ID]; !ok { - continue - } - if _, excluded := excludedIDs[acc.ID]; excluded { - continue - } - // Scheduler snapshots can be temporarily stale; re-check schedulability here to - // avoid selecting accounts that were recently rate-limited/overloaded. - if !s.isAccountSchedulableForSelection(acc) { - continue - } - // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 - if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { - _ = s.accountRepo.SetError(ctx, acc.ID, - fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) - continue - } - // 过滤:原生平台直接通过,antigravity 需要启用混合调度 - if acc.Platform == PlatformAntigravity && !acc.IsMixedSchedulingEnabled() { - continue - } - if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForQuota(acc) { - continue - } - if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { - continue - } - if !s.isAccountSchedulableForRPM(ctx, acc, false) { - continue - } - if selected == nil { - selected = acc - continue - } - if acc.Priority < selected.Priority { - selected = acc - } else if acc.Priority == selected.Priority { - switch { - case acc.LastUsedAt == nil && selected.LastUsedAt != nil: - selected = acc - case acc.LastUsedAt != nil && selected.LastUsedAt == nil: - // keep selected (never used is preferred) - case acc.LastUsedAt == nil && selected.LastUsedAt == nil: - if preferOAuth && acc.Platform == PlatformGemini && selected.Platform == PlatformGemini && acc.Type != selected.Type && acc.Type == AccountTypeOAuth { - selected = acc - } - default: - if acc.LastUsedAt.Before(*selected.LastUsedAt) { - selected = acc - } - } - } - } - - if selected != nil { - if sessionHash != "" && s.cache != nil { - if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { - logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) - } - } - if s.debugModelRoutingEnabled() { - logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy mixed routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), selected.ID) - } - return selected, nil - } - logger.LegacyPrintf("service.gateway", "[ModelRouting] No routed accounts available for model=%s, falling back to normal selection", requestedModel) - } - - // 1. 查询粘性会话 - if sessionHash != "" && s.cache != nil { - accountID, err := s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - if err == nil && accountID > 0 { - if _, excluded := excludedIDs[accountID]; !excluded { - account, err := s.getSchedulableAccount(ctx, accountID) - // 检查账号分组归属和有效性:原生平台直接匹配,antigravity 需要启用混合调度 - if err == nil { - clearSticky := shouldClearStickySession(account, requestedModel) - if clearSticky { - _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) - } - if !clearSticky && s.isAccountInGroup(account, groupID) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) && !s.isStickyAccountUpstreamRestricted(ctx, groupID, account, requestedModel) { - if account.Platform == nativePlatform || (account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled()) { - return account, nil - } - } - } - } - } - } - - // 2. 获取可调度账号列表 - if !accountsLoaded { - var err error - accounts, _, err = s.listSchedulableAccounts(ctx, groupID, nativePlatform, false) - if err != nil { - return nil, fmt.Errorf("query accounts failed: %w", err) - } - } - - // 批量预取窗口费用+RPM 计数,避免逐个账号查询(N+1) - ctx = s.withWindowCostPrefetch(ctx, accounts) - ctx = s.withRPMPrefetch(ctx, accounts) - - // 3. 按优先级+最久未用选择(考虑模型支持和混合调度) - // needsUpstreamCheck 仅在主选择循环中使用;粘性会话命中时跳过此检查。 - needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) - var selected *Account - for i := range accounts { - acc := &accounts[i] - if _, excluded := excludedIDs[acc.ID]; excluded { - continue - } - // Scheduler snapshots can be temporarily stale; re-check schedulability here to - // avoid selecting accounts that were recently rate-limited/overloaded. - if !s.isAccountSchedulableForSelection(acc) { - continue - } - // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 - if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { - _ = s.accountRepo.SetError(ctx, acc.ID, - fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) - continue - } - // 过滤:原生平台直接通过,antigravity 需要启用混合调度 - if acc.Platform == PlatformAntigravity && !acc.IsMixedSchedulingEnabled() { - continue - } - if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { - continue - } - if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { - continue - } - if !s.isAccountSchedulableForQuota(acc) { - continue - } - if !s.isAccountSchedulableForWindowCost(ctx, acc, false) { - continue - } - if !s.isAccountSchedulableForRPM(ctx, acc, false) { - continue - } - if selected == nil { - selected = acc - continue - } - if acc.Priority < selected.Priority { - selected = acc - } else if acc.Priority == selected.Priority { - switch { - case acc.LastUsedAt == nil && selected.LastUsedAt != nil: - selected = acc - case acc.LastUsedAt != nil && selected.LastUsedAt == nil: - // keep selected (never used is preferred) - case acc.LastUsedAt == nil && selected.LastUsedAt == nil: - if preferOAuth && acc.Platform == PlatformGemini && selected.Platform == PlatformGemini && acc.Type != selected.Type && acc.Type == AccountTypeOAuth { - selected = acc - } - default: - if acc.LastUsedAt.Before(*selected.LastUsedAt) { - selected = acc - } - } - } - } - - if selected == nil { - stats := s.logDetailedSelectionFailure(ctx, groupID, sessionHash, requestedModel, nativePlatform, accounts, excludedIDs, true) - if requestedModel != "" { - return nil, fmt.Errorf("%w supporting model: %s (%s)", ErrNoAvailableAccounts, requestedModel, summarizeSelectionFailureStats(stats)) - } - return nil, ErrNoAvailableAccounts - } - - // 4. 建立粘性绑定 - if sessionHash != "" && s.cache != nil { - if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { - logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) - } - } - - return selected, nil -} - -type selectionFailureStats struct { - Total int - Eligible int - Excluded int - Unschedulable int - PlatformFiltered int - ModelUnsupported int - ModelRateLimited int - SamplePlatformIDs []int64 - SampleMappingIDs []int64 - SampleRateLimitIDs []string -} - -type selectionFailureDiagnosis struct { - Category string - Detail string -} - -func (s *GatewayService) logDetailedSelectionFailure( - ctx context.Context, - groupID *int64, - sessionHash string, - requestedModel string, - platform string, - accounts []Account, - excludedIDs map[int64]struct{}, - allowMixedScheduling bool, -) selectionFailureStats { - stats := s.collectSelectionFailureStats(ctx, accounts, requestedModel, platform, excludedIDs, allowMixedScheduling) - logger.LegacyPrintf( - "service.gateway", - "[SelectAccountDetailed] group_id=%v model=%s platform=%s session=%s total=%d eligible=%d excluded=%d unschedulable=%d platform_filtered=%d model_unsupported=%d model_rate_limited=%d sample_platform_filtered=%v sample_model_unsupported=%v sample_model_rate_limited=%v", - derefGroupID(groupID), - requestedModel, - platform, - shortSessionHash(sessionHash), - stats.Total, - stats.Eligible, - stats.Excluded, - stats.Unschedulable, - stats.PlatformFiltered, - stats.ModelUnsupported, - stats.ModelRateLimited, - stats.SamplePlatformIDs, - stats.SampleMappingIDs, - stats.SampleRateLimitIDs, - ) - return stats -} - -func (s *GatewayService) collectSelectionFailureStats( - ctx context.Context, - accounts []Account, - requestedModel string, - platform string, - excludedIDs map[int64]struct{}, - allowMixedScheduling bool, -) selectionFailureStats { - stats := selectionFailureStats{ - Total: len(accounts), - } - - for i := range accounts { - acc := &accounts[i] - diagnosis := s.diagnoseSelectionFailure(ctx, acc, requestedModel, platform, excludedIDs, allowMixedScheduling) - switch diagnosis.Category { - case "excluded": - stats.Excluded++ - case "unschedulable": - stats.Unschedulable++ - case "platform_filtered": - stats.PlatformFiltered++ - stats.SamplePlatformIDs = appendSelectionFailureSampleID(stats.SamplePlatformIDs, acc.ID) - case "model_unsupported": - stats.ModelUnsupported++ - stats.SampleMappingIDs = appendSelectionFailureSampleID(stats.SampleMappingIDs, acc.ID) - case "model_rate_limited": - stats.ModelRateLimited++ - remaining := acc.GetRateLimitRemainingTimeWithContext(ctx, requestedModel).Truncate(time.Second) - stats.SampleRateLimitIDs = appendSelectionFailureRateSample(stats.SampleRateLimitIDs, acc.ID, remaining) - default: - stats.Eligible++ - } - } - - return stats -} - -func (s *GatewayService) diagnoseSelectionFailure( - ctx context.Context, - acc *Account, - requestedModel string, - platform string, - excludedIDs map[int64]struct{}, - allowMixedScheduling bool, -) selectionFailureDiagnosis { - if acc == nil { - return selectionFailureDiagnosis{Category: "unschedulable", Detail: "account_nil"} - } - if _, excluded := excludedIDs[acc.ID]; excluded { - return selectionFailureDiagnosis{Category: "excluded"} - } - if !s.isAccountSchedulableForSelection(acc) { - return selectionFailureDiagnosis{Category: "unschedulable", Detail: "generic_unschedulable"} - } - if isPlatformFilteredForSelection(acc, platform, allowMixedScheduling) { - return selectionFailureDiagnosis{ - Category: "platform_filtered", - Detail: fmt.Sprintf("account_platform=%s requested_platform=%s", acc.Platform, strings.TrimSpace(platform)), - } - } - if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { - return selectionFailureDiagnosis{ - Category: "model_unsupported", - Detail: fmt.Sprintf("model=%s", requestedModel), - } - } - if !s.isAccountSchedulableForModelSelection(ctx, acc, requestedModel) { - remaining := acc.GetRateLimitRemainingTimeWithContext(ctx, requestedModel).Truncate(time.Second) - return selectionFailureDiagnosis{ - Category: "model_rate_limited", - Detail: fmt.Sprintf("remaining=%s", remaining), - } - } - return selectionFailureDiagnosis{Category: "eligible"} -} - -func isPlatformFilteredForSelection(acc *Account, platform string, allowMixedScheduling bool) bool { - if acc == nil { - return true - } - if allowMixedScheduling { - if acc.Platform == PlatformAntigravity { - return !acc.IsMixedSchedulingEnabled() - } - return acc.Platform != platform - } - if strings.TrimSpace(platform) == "" { - return false - } - return acc.Platform != platform -} - -func appendSelectionFailureSampleID(samples []int64, id int64) []int64 { - const limit = 5 - if len(samples) >= limit { - return samples - } - return append(samples, id) -} - -func appendSelectionFailureRateSample(samples []string, accountID int64, remaining time.Duration) []string { - const limit = 5 - if len(samples) >= limit { - return samples - } - return append(samples, fmt.Sprintf("%d(%s)", accountID, remaining)) -} - -func summarizeSelectionFailureStats(stats selectionFailureStats) string { - return fmt.Sprintf( - "total=%d eligible=%d excluded=%d unschedulable=%d platform_filtered=%d model_unsupported=%d model_rate_limited=%d", - stats.Total, - stats.Eligible, - stats.Excluded, - stats.Unschedulable, - stats.PlatformFiltered, - stats.ModelUnsupported, - stats.ModelRateLimited, - ) -} - -// isModelSupportedByAccountWithContext 根据账户平台检查模型支持(带 context) -// 对于 Antigravity 平台,会先获取映射后的最终模型名(包括 thinking 后缀)再检查支持 -func (s *GatewayService) isModelSupportedByAccountWithContext(ctx context.Context, account *Account, requestedModel string) bool { - if account.Platform == PlatformAntigravity { - if strings.TrimSpace(requestedModel) == "" { - return true - } - // 使用与转发阶段一致的映射逻辑:自定义映射优先 → 默认映射兜底 - mapped := mapAntigravityModel(account, requestedModel) - if mapped == "" { - return false - } - // 应用 thinking 后缀后检查最终模型是否在账号映射中 - if enabled, ok := ThinkingEnabledFromContext(ctx); ok { - finalModel := applyThinkingModelSuffix(mapped, enabled) - if finalModel == mapped { - return true // thinking 后缀未改变模型名,映射已通过 - } - return account.IsModelSupported(finalModel) - } - return true - } - return s.isModelSupportedByAccount(account, requestedModel) -} - -// isModelSupportedByAccount 根据账户平台检查模型支持(无 context,用于非 Antigravity 平台) -func (s *GatewayService) isModelSupportedByAccount(account *Account, requestedModel string) bool { - if account.Platform == PlatformAntigravity { - if strings.TrimSpace(requestedModel) == "" { - return true - } - return mapAntigravityModel(account, requestedModel) != "" - } - if account.IsBedrock() { - _, ok := ResolveBedrockModelID(account, requestedModel) - return ok - } - // OpenAI 透传模式:仅替换认证,允许所有模型 - if account.Platform == PlatformOpenAI && account.IsOpenAIPassthroughEnabled() { - return true - } - // OAuth/SetupToken 账号使用 Anthropic 标准映射(短ID → 长ID) - if account.Platform == PlatformAnthropic && account.Type != AccountTypeAPIKey { - if account.Type == AccountTypeServiceAccount { - requestedModel = normalizeVertexAnthropicModelID(claude.NormalizeModelID(requestedModel)) - } else { - requestedModel = claude.NormalizeModelID(requestedModel) - } - } - // 其他平台使用账户的模型支持检查 - return account.IsModelSupported(requestedModel) -} - // GetAccessToken 获取账号凭证 func (s *GatewayService) GetAccessToken(ctx context.Context, account *Account) (string, string, error) { switch account.Type { @@ -5640,1184 +3200,6 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A }, nil } -type anthropicPassthroughForwardInput struct { - Body []byte - Parsed *ParsedRequest - RequestModel string - OriginalModel string - RequestStream bool - StartTime time.Time -} - -func (s *GatewayService) forwardAnthropicAPIKeyPassthrough( - ctx context.Context, - c *gin.Context, - account *Account, - body []byte, - reqModel string, - originalModel string, - reqStream bool, - startTime time.Time, -) (*ForwardResult, error) { - return s.forwardAnthropicAPIKeyPassthroughWithInput(ctx, c, account, anthropicPassthroughForwardInput{ - Body: body, - RequestModel: reqModel, - OriginalModel: originalModel, - RequestStream: reqStream, - StartTime: startTime, - }) -} - -func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput( - ctx context.Context, - c *gin.Context, - account *Account, - input anthropicPassthroughForwardInput, -) (*ForwardResult, error) { - token, tokenType, err := s.GetAccessToken(ctx, account) - if err != nil { - return nil, err - } - if tokenType != "apikey" { - return nil, fmt.Errorf("anthropic api key passthrough requires apikey token, got: %s", tokenType) - } - - proxyURL := "" - if account.ProxyID != nil && account.Proxy != nil { - proxyURL = account.Proxy.URL() - } - - logger.LegacyPrintf("service.gateway", "[Anthropic 自动透传] 命中 API Key 透传分支: account=%d name=%s model=%s stream=%v", - account.ID, account.Name, input.RequestModel, input.RequestStream) - - if c != nil { - c.Set("anthropic_passthrough", true) - } - // Pre-filter: strip empty text blocks (including nested in tool_result) to prevent upstream 400. - input.Body = StripEmptyTextBlocks(input.Body) - // Pre-filter: strip web-search history blocks the upstream cannot accept - // (emulation-synthesized ones always; genuine ones additionally for - // passback-required third-party upstreams such as GLM/Kimi/DeepSeek, - // which reject server_tool_use with 400). input.RequestModel 已是映射后的模型 ID。 - input.Body = FilterWebSearchHistoryBlocks(input.Body, input.RequestModel) - if input.Parsed != nil { - // 透传分支也会改写实际 wire body,成功 usage hash 依赖这里同步当前 body。 - if err := input.Parsed.ReplaceBody(input.Body); err != nil { - return nil, err - } - } - - var resp *http.Response - retryStart := time.Now() - for attempt := 1; attempt <= maxRetryAttempts; attempt++ { - upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, input.RequestStream) - upstreamReq, wireBody, err := s.buildUpstreamRequestAnthropicAPIKeyPassthrough(upstreamCtx, c, account, input.Body, token) - releaseUpstreamCtx() - if err != nil { - return nil, err - } - if input.Parsed != nil && !bytes.Equal(wireBody, input.Body) { - // build 阶段会按 beta 能力清理 body,发送前同步到 ParsedRequest 当前视图。 - if err := input.Parsed.ReplaceBody(wireBody); err != nil { - return nil, err - } - input.Body = input.Parsed.Body.Bytes() - } - - resp, err = s.httpUpstream.DoWithTLS(upstreamReq, proxyURL, account.ID, account.Concurrency, s.tlsFPProfileService.ResolveTLSProfile(account)) - if err != nil { - if resp != nil && resp.Body != nil { - _ = resp.Body.Close() - } - safeErr := sanitizeUpstreamErrorMessage(err.Error()) - setOpsUpstreamError(c, 0, safeErr, "") - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: 0, - UpstreamURL: safeUpstreamURL(upstreamReq.URL.String()), - Passthrough: true, - Kind: "request_error", - Message: safeErr, - }) - c.JSON(http.StatusBadGateway, gin.H{ - "type": "error", - "error": gin.H{ - "type": "upstream_error", - "message": "Upstream request failed", - }, - }) - return nil, fmt.Errorf("upstream request failed: %s", safeErr) - } - - // 透传分支禁止 400 请求体降级重试(该重试会改写请求体) - if resp.StatusCode >= 400 && resp.StatusCode != 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) { - if attempt < maxRetryAttempts { - elapsed := time.Since(retryStart) - if elapsed >= maxRetryElapsed { - break - } - - delay := retryBackoffDelay(attempt) - remaining := maxRetryElapsed - elapsed - if delay > remaining { - delay = remaining - } - if delay <= 0 { - break - } - - respBody, _ := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - UpstreamURL: safeUpstreamURL(upstreamReq.URL.String()), - Passthrough: true, - Kind: "retry", - Message: extractUpstreamErrorMessage(respBody), - Detail: func() string { - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - return truncateString(string(respBody), s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes) - } - return "" - }(), - }) - logger.LegacyPrintf("service.gateway", "Anthropic passthrough account %d: upstream error %d, retry %d/%d after %v (elapsed=%v/%v)", - account.ID, resp.StatusCode, attempt, maxRetryAttempts, delay, elapsed, maxRetryElapsed) - if err := sleepWithContext(ctx, delay); err != nil { - return nil, err - } - continue - } - break - } - - break - } - if resp == nil || resp.Body == nil { - return nil, errors.New("upstream request failed: empty response") - } - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode >= 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) { - if s.shouldFailoverUpstreamError(resp.StatusCode) { - respBody, _ := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) - - logger.LegacyPrintf("service.gateway", "[Anthropic Passthrough] Upstream error (retry exhausted, failover): Account=%d(%s) Status=%d RequestID=%s Body=%s", - account.ID, account.Name, resp.StatusCode, resp.Header.Get("x-request-id"), truncateString(string(respBody), 1000)) - - s.handleRetryExhaustedSideEffects(ctx, resp, account) - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Passthrough: true, - Kind: "retry_exhausted_failover", - Message: extractUpstreamErrorMessage(respBody), - Detail: func() string { - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - return truncateString(string(respBody), s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes) - } - return "" - }(), - }) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), - } - } - return s.handleRetryExhaustedError(ctx, resp, c, account) - } - - if resp.StatusCode >= 400 && s.shouldFailoverUpstreamError(resp.StatusCode) { - respBody, _ := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) - - logger.LegacyPrintf("service.gateway", "[Anthropic Passthrough] Upstream error (failover): Account=%d(%s) Status=%d RequestID=%s Body=%s", - account.ID, account.Name, resp.StatusCode, resp.Header.Get("x-request-id"), truncateString(string(respBody), 1000)) - - s.handleFailoverSideEffects(ctx, resp, account, input.RequestModel) - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Passthrough: true, - Kind: "failover", - Message: extractUpstreamErrorMessage(respBody), - Detail: func() string { - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - return truncateString(string(respBody), s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes) - } - return "" - }(), - }) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), - } - } - - if resp.StatusCode >= 400 { - return s.handleErrorResponse(ctx, resp, c, account, input.RequestModel) - } - - var usage *ClaudeUsage - var firstTokenMs *int - var clientDisconnect bool - if input.RequestStream { - streamResult, err := s.handleStreamingResponseAnthropicAPIKeyPassthrough(ctx, resp, c, account, input.StartTime, input.RequestModel) - if err != nil { - return nil, err - } - usage = streamResult.usage - firstTokenMs = streamResult.firstTokenMs - clientDisconnect = streamResult.clientDisconnect - } else { - usage, err = s.handleNonStreamingResponseAnthropicAPIKeyPassthrough(ctx, resp, c, account) - if err != nil { - return nil, err - } - } - if usage == nil { - usage = &ClaudeUsage{} - } - - return &ForwardResult{ - RequestID: resp.Header.Get("x-request-id"), - Usage: *usage, - Model: input.OriginalModel, - UpstreamModel: input.RequestModel, - Stream: input.RequestStream, - Duration: time.Since(input.StartTime), - FirstTokenMs: firstTokenMs, - ClientDisconnect: clientDisconnect, - }, nil -} - -func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough( - ctx context.Context, - c *gin.Context, - account *Account, - body []byte, - token string, -) (*http.Request, []byte, error) { - targetURL := claudeAPIURL - baseURL := account.GetBaseURL() - if baseURL != "" { - validatedURL, err := s.validateUpstreamBaseURL(baseURL) - if err != nil { - return nil, nil, err - } - targetURL = validatedURL + "/v1/messages?beta=true" - } - - // 能力维度 body sanitize:透传路径上 anthropic-beta header 原样透传客户端值, - // 依此决定是否保留 body 中的 context_management。避免“客户端 body 带字段但 - // header 忘记带 beta token”的客户端 bug 在透传场景下让上游 400。 - clientBeta := "" - if c != nil && c.Request != nil { - clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") - } - // 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准 - if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok { - clientBeta = beta - } - if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { - body = sanitized - } - - req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) - if err != nil { - return nil, nil, err - } - - if c != nil && c.Request != nil { - for key, values := range c.Request.Header { - lowerKey := strings.ToLower(strings.TrimSpace(key)) - if !allowedHeaders[lowerKey] { - continue - } - wireKey := resolveWireCasing(key) - for _, v := range values { - addHeaderRaw(req.Header, wireKey, v) - } - } - } - - // 覆盖入站鉴权残留,并注入上游认证 - req.Header.Del("authorization") - req.Header.Del("x-api-key") - req.Header.Del("x-goog-api-key") - req.Header.Del("cookie") - setAnthropicAPIKeyAuthHeader(req.Header, account, token) - - if getHeaderRaw(req.Header, "content-type") == "" { - setHeaderRaw(req.Header, "content-type", "application/json") - } - if getHeaderRaw(req.Header, "anthropic-version") == "" { - setHeaderRaw(req.Header, "anthropic-version", "2023-06-01") - } - - // 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头) - account.ApplyHeaderOverrides(req.Header) - - return req, body, nil -} - -func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough( - ctx context.Context, - resp *http.Response, - c *gin.Context, - account *Account, - startTime time.Time, - model string, -) (*streamingResult, error) { - if s.rateLimitService != nil { - s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header) - } - - writeAnthropicPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - - contentType := strings.TrimSpace(resp.Header.Get("Content-Type")) - if contentType == "" { - contentType = "text/event-stream" - } - c.Header("Content-Type", contentType) - if c.Writer.Header().Get("Cache-Control") == "" { - c.Header("Cache-Control", "no-cache") - } - if c.Writer.Header().Get("Connection") == "" { - c.Header("Connection", "keep-alive") - } - c.Header("X-Accel-Buffering", "no") - if v := resp.Header.Get("x-request-id"); v != "" { - c.Header("x-request-id", v) - } - - w := c.Writer - flusher, ok := w.(http.Flusher) - if !ok { - return nil, errors.New("streaming not supported") - } - - usage := &ClaudeUsage{} - var firstTokenMs *int - clientDisconnected := false - sawTerminalEvent := false - - scanner := bufio.NewScanner(resp.Body) - maxLineSize := defaultMaxLineSize - if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { - maxLineSize = s.cfg.Gateway.MaxLineSize - } - scanBuf := getSSEScannerBuf64K() - scanner.Buffer(scanBuf[:0], maxLineSize) - - type scanEvent struct { - line string - err error - } - events := make(chan scanEvent, 16) - done := make(chan struct{}) - sendEvent := func(ev scanEvent) bool { - select { - case events <- ev: - return true - case <-done: - return false - } - } - var lastReadAt int64 - atomic.StoreInt64(&lastReadAt, time.Now().UnixNano()) - go func(scanBuf *sseScannerBuf64K) { - defer putSSEScannerBuf64K(scanBuf) - defer close(events) - for scanner.Scan() { - atomic.StoreInt64(&lastReadAt, time.Now().UnixNano()) - if !sendEvent(scanEvent{line: scanner.Text()}) { - return - } - } - if err := scanner.Err(); err != nil { - _ = sendEvent(scanEvent{err: err}) - } - }(scanBuf) - defer close(done) - - streamInterval := time.Duration(0) - if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 { - streamInterval = time.Duration(s.cfg.Gateway.StreamDataIntervalTimeout) * time.Second - } - var intervalTicker *time.Ticker - if streamInterval > 0 { - intervalTicker = time.NewTicker(streamInterval) - defer intervalTicker.Stop() - } - var intervalCh <-chan time.Time - if intervalTicker != nil { - intervalCh = intervalTicker.C - } - - keepaliveInterval := time.Duration(0) - if s.cfg != nil && s.cfg.Gateway.StreamKeepaliveInterval > 0 { - keepaliveInterval = time.Duration(s.cfg.Gateway.StreamKeepaliveInterval) * time.Second - } - var keepaliveTimer *time.Timer - if keepaliveInterval > 0 { - keepaliveTimer = time.NewTimer(keepaliveInterval) - defer keepaliveTimer.Stop() - } - var keepaliveCh <-chan time.Time - if keepaliveTimer != nil { - keepaliveCh = keepaliveTimer.C - } - lastDataAt := time.Now() - resetKeepaliveTimer := func() { - if keepaliveTimer == nil { - return - } - if !keepaliveTimer.Stop() { - select { - case <-keepaliveTimer.C: - default: - } - } - keepaliveTimer.Reset(keepaliveInterval) - } - inPartialEvent := false - - for { - select { - case ev, ok := <-events: - if !ok { - if !clientDisconnected { - // 兜底补刷,确保最后一个未以空行结尾的事件也能及时送达客户端。 - flusher.Flush() - } - if !sawTerminalEvent { - if clientDisconnected && streamInterval > 0 { - lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) - if time.Since(lastRead) >= streamInterval { - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, fmt.Errorf("stream usage incomplete after timeout") - } - } - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: clientDisconnected}, fmt.Errorf("stream usage incomplete: missing terminal event") - } - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: clientDisconnected}, nil - } - if ev.err != nil { - if sawTerminalEvent { - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: clientDisconnected}, nil - } - if clientDisconnected { - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, fmt.Errorf("stream usage incomplete after disconnect: %w", ev.err) - } - if errors.Is(ev.err, context.Canceled) || errors.Is(ev.err, context.DeadlineExceeded) { - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, fmt.Errorf("stream usage incomplete: %w", ev.err) - } - if errors.Is(ev.err, bufio.ErrTooLong) { - logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] SSE line too long: account=%d max_size=%d error=%v", account.ID, maxLineSize, ev.err) - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs}, ev.err - } - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs}, fmt.Errorf("stream read error: %w", ev.err) - } - - line := ev.line - if data, ok := extractAnthropicSSEDataLine(line); ok { - trimmed := strings.TrimSpace(data) - if anthropicStreamEventIsTerminal("", trimmed) { - sawTerminalEvent = true - } - if firstTokenMs == nil && trimmed != "" && trimmed != "[DONE]" { - ms := int(time.Since(startTime).Milliseconds()) - firstTokenMs = &ms - } - s.parseSSEUsagePassthrough(data, usage) - } else { - trimmed := strings.TrimSpace(line) - if strings.HasPrefix(trimmed, "event:") && anthropicStreamEventIsTerminal(strings.TrimSpace(strings.TrimPrefix(trimmed, "event:")), "") { - sawTerminalEvent = true - } - } - - if !clientDisconnected { - restored := string(reverseToolNamesIfPresent(c, []byte(line))) - if _, err := io.WriteString(w, restored); err != nil { - clientDisconnected = true - logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) - } else if _, err := io.WriteString(w, "\n"); err != nil { - clientDisconnected = true - logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) - } else if line == "" { - // 按 SSE 事件边界刷出,减少每行 flush 带来的 syscall 开销。 - flusher.Flush() - lastDataAt = time.Now() - resetKeepaliveTimer() - inPartialEvent = false - } else { - inPartialEvent = true - } - } - - case <-intervalCh: - lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) - if time.Since(lastRead) < streamInterval { - continue - } - if clientDisconnected { - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, fmt.Errorf("stream usage incomplete after timeout") - } - logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] Stream data interval timeout: account=%d model=%s interval=%s", account.ID, model, streamInterval) - if s.rateLimitService != nil { - s.rateLimitService.HandleStreamTimeout(ctx, account, model) - } - return &streamingResult{usage: usage, firstTokenMs: firstTokenMs}, fmt.Errorf("stream data interval timeout") - - case <-keepaliveCh: - if clientDisconnected { - continue - } - if inPartialEvent { - resetKeepaliveTimer() - continue - } - if time.Since(lastDataAt) < keepaliveInterval { - resetKeepaliveTimer() - continue - } - if _, err := fmt.Fprint(w, "event: ping\ndata: {\"type\": \"ping\"}\n\n"); err != nil { - clientDisconnected = true - logger.LegacyPrintf("service.gateway", "[Anthropic passthrough] Client disconnected during keepalive ping, continue draining upstream for usage: account=%d", account.ID) - continue - } - flusher.Flush() - lastDataAt = time.Now() - resetKeepaliveTimer() - } - } -} - -func extractAnthropicSSEDataLine(line string) (string, bool) { - if !strings.HasPrefix(line, "data:") { - return "", false - } - start := len("data:") - for start < len(line) { - if line[start] != ' ' && line[start] != '\t' { - break - } - start++ - } - return line[start:], true -} - -func (s *GatewayService) parseSSEUsagePassthrough(data string, usage *ClaudeUsage) { - if usage == nil || data == "" || data == "[DONE]" { - return - } - - parsed := gjson.Parse(data) - switch parsed.Get("type").String() { - case "message_start": - msgUsage := parsed.Get("message.usage") - if msgUsage.Exists() { - usage.InputTokens = int(msgUsage.Get("input_tokens").Int()) - usage.CacheCreationInputTokens = int(msgUsage.Get("cache_creation_input_tokens").Int()) - usage.CacheReadInputTokens = int(msgUsage.Get("cache_read_input_tokens").Int()) - - // 保持与通用解析一致:message_start 允许覆盖 5m/1h 明细(包括 0)。 - cc5m := msgUsage.Get("cache_creation.ephemeral_5m_input_tokens") - cc1h := msgUsage.Get("cache_creation.ephemeral_1h_input_tokens") - if cc5m.Exists() || cc1h.Exists() { - usage.CacheCreation5mTokens = int(cc5m.Int()) - usage.CacheCreation1hTokens = int(cc1h.Int()) - } - } - case "message_delta": - deltaUsage := parsed.Get("usage") - if deltaUsage.Exists() { - if v := deltaUsage.Get("input_tokens").Int(); v > 0 { - usage.InputTokens = int(v) - } - if v := deltaUsage.Get("output_tokens").Int(); v > 0 { - usage.OutputTokens = int(v) - } - if v := deltaUsage.Get("cache_creation_input_tokens").Int(); v > 0 { - usage.CacheCreationInputTokens = int(v) - } - if v := deltaUsage.Get("cache_read_input_tokens").Int(); v > 0 { - usage.CacheReadInputTokens = int(v) - } - - cc5m := deltaUsage.Get("cache_creation.ephemeral_5m_input_tokens") - cc1h := deltaUsage.Get("cache_creation.ephemeral_1h_input_tokens") - if cc5m.Exists() && cc5m.Int() > 0 { - usage.CacheCreation5mTokens = int(cc5m.Int()) - } - if cc1h.Exists() && cc1h.Int() > 0 { - usage.CacheCreation1hTokens = int(cc1h.Int()) - } - } - } - - if usage.CacheReadInputTokens == 0 { - if cached := parsed.Get("message.usage.cached_tokens").Int(); cached > 0 { - usage.CacheReadInputTokens = int(cached) - } - if cached := parsed.Get("usage.cached_tokens").Int(); usage.CacheReadInputTokens == 0 && cached > 0 { - usage.CacheReadInputTokens = int(cached) - } - } - if usage.CacheCreationInputTokens == 0 { - cc5m := parsed.Get("message.usage.cache_creation.ephemeral_5m_input_tokens").Int() - cc1h := parsed.Get("message.usage.cache_creation.ephemeral_1h_input_tokens").Int() - if cc5m == 0 && cc1h == 0 { - cc5m = parsed.Get("usage.cache_creation.ephemeral_5m_input_tokens").Int() - cc1h = parsed.Get("usage.cache_creation.ephemeral_1h_input_tokens").Int() - } - total := cc5m + cc1h - if total > 0 { - usage.CacheCreationInputTokens = int(total) - } - } -} - -func parseClaudeUsageFromResponseBody(body []byte) *ClaudeUsage { - usage := &ClaudeUsage{} - if len(body) == 0 { - return usage - } - - parsed := gjson.ParseBytes(body) - usageNode := parsed.Get("usage") - if !usageNode.Exists() { - return usage - } - - usage.InputTokens = int(usageNode.Get("input_tokens").Int()) - usage.OutputTokens = int(usageNode.Get("output_tokens").Int()) - usage.CacheCreationInputTokens = int(usageNode.Get("cache_creation_input_tokens").Int()) - usage.CacheReadInputTokens = int(usageNode.Get("cache_read_input_tokens").Int()) - - cc5m := usageNode.Get("cache_creation.ephemeral_5m_input_tokens").Int() - cc1h := usageNode.Get("cache_creation.ephemeral_1h_input_tokens").Int() - if cc5m > 0 || cc1h > 0 { - usage.CacheCreation5mTokens = int(cc5m) - usage.CacheCreation1hTokens = int(cc1h) - } - if usage.CacheCreationInputTokens == 0 && (cc5m > 0 || cc1h > 0) { - usage.CacheCreationInputTokens = int(cc5m + cc1h) - } - if usage.CacheReadInputTokens == 0 { - if cached := usageNode.Get("cached_tokens").Int(); cached > 0 { - usage.CacheReadInputTokens = int(cached) - } - } - return usage -} - -func (s *GatewayService) invalidNonStreamingJSONFailoverError( - ctx context.Context, - resp *http.Response, - account *Account, - body []byte, - parseErr error, - requestedModel ...string, -) error { - const statusCode = http.StatusBadGateway - - accountID := int64(0) - accountName := "" - retryableOnSameAccount := false - if account != nil { - accountID = account.ID - accountName = account.Name - retryableOnSameAccount = account.IsPoolMode() && account.IsPoolModeRetryableStatus(statusCode) - } - - logger.LegacyPrintf( - "service.gateway", - "Account %d(%s): upstream returned non-JSON 2xx response, attempting failover: status=%d request_id=%s error=%v", - accountID, - accountName, - resp.StatusCode, - resp.Header.Get("x-request-id"), - parseErr, - ) - - if s.rateLimitService != nil && account != nil { - if len(requestedModel) > 0 { - s.rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body, requestedModel[0]) - } else { - s.rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body) - } - } - - return &UpstreamFailoverError{ - StatusCode: statusCode, - ResponseBody: body, - ResponseHeaders: resp.Header, - RetryableOnSameAccount: retryableOnSameAccount, - } -} - -func (s *GatewayService) handleNonStreamingResponseAnthropicAPIKeyPassthrough( - ctx context.Context, - resp *http.Response, - c *gin.Context, - account *Account, -) (*ClaudeUsage, error) { - if s.rateLimitService != nil { - s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header) - } - - body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, anthropicTooLargeError) - if err != nil { - return nil, err - } - - if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices { - var raw json.RawMessage - if err := json.Unmarshal(body, &raw); err != nil { - return nil, s.invalidNonStreamingJSONFailoverError(ctx, resp, account, body, err) - } - } - - usage := parseClaudeUsageFromResponseBody(body) - - writeAnthropicPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - contentType := strings.TrimSpace(resp.Header.Get("Content-Type")) - if contentType == "" { - contentType = "application/json" - } - body = reverseToolNamesIfPresent(c, body) - c.Data(resp.StatusCode, contentType, body) - return usage, nil -} - -func writeAnthropicPassthroughResponseHeaders(dst http.Header, src http.Header, filter *responseheaders.CompiledHeaderFilter) { - if dst == nil || src == nil { - return - } - if filter != nil { - responseheaders.WriteFilteredHeaders(dst, src, filter) - return - } - if v := strings.TrimSpace(src.Get("Content-Type")); v != "" { - dst.Set("Content-Type", v) - } - if v := strings.TrimSpace(src.Get("x-request-id")); v != "" { - dst.Set("x-request-id", v) - } -} - -// ApplyBedrockCCCompat 应用 Bedrock CC 兼容转换(渠道级模型映射后调用) -// 清理 body 中 Anthropic API 专有字段、修复 thinking/tool_use ID、过滤 beta token, -// 同时过滤 HTTP header 中的 anthropic-beta(防止 Passthrough 路径透传不支持的 token)。 -func (s *GatewayService) ApplyBedrockCCCompat(c *gin.Context, body []byte, model string, account *Account, groupID *int64) []byte { - if !s.isBedrockCCCompatEnabled(c.Request.Context(), account, groupID) { - return body - } - body = sanitizeBedrockCCFields(body) - body = sanitizeBedrockThinking(body, model) - body = sanitizeBedrockToolUseIDs(body) - body = sanitizeBedrockCCBetaTokens(body, model) - // 过滤 HTTP header 中的 anthropic-beta,只保留 Bedrock 支持的 token - if betaHeader := c.GetHeader("anthropic-beta"); betaHeader != "" { - if filtered := ResolveBedrockBetaTokens(betaHeader, body, model); len(filtered) > 0 { - c.Request.Header.Set("anthropic-beta", strings.Join(filtered, ", ")) - } else { - c.Request.Header.Del("anthropic-beta") - } - } - return body -} - -// isBedrockCCCompatEnabled 检查渠道是否启用了 Bedrock CC 兼容模式 -func (s *GatewayService) isBedrockCCCompatEnabled(ctx context.Context, account *Account, groupID *int64) bool { - if groupID == nil || s.channelService == nil { - return false - } - ch, err := s.channelService.GetChannelForGroup(ctx, *groupID) - if err != nil || ch == nil { - return false - } - return ch.IsBedrockCCCompatEnabled(account.Platform) -} - -// forwardBedrock 转发请求到 AWS Bedrock -func (s *GatewayService) forwardBedrock( - ctx context.Context, - c *gin.Context, - account *Account, - parsed *ParsedRequest, - startTime time.Time, -) (*ForwardResult, error) { - reqModel := parsed.Model - reqStream := parsed.Stream - body := parsed.Body.Bytes() - - region := bedrockRuntimeRegion(account) - mappedModel, ok := ResolveBedrockModelID(account, reqModel) - if !ok { - return nil, fmt.Errorf("unsupported bedrock model: %s", reqModel) - } - if mappedModel != reqModel { - logger.LegacyPrintf("service.gateway", "[Bedrock] Model mapping: %s -> %s (account: %s)", reqModel, mappedModel, account.Name) - } - - betaHeader := "" - if c != nil && c.Request != nil { - betaHeader = c.GetHeader("anthropic-beta") - } - - // 准备请求体(注入 anthropic_version/anthropic_beta,移除 Bedrock 不支持的字段,清理 cache_control) - betaTokens, err := s.resolveBedrockBetaTokensForRequest(ctx, account, betaHeader, body, mappedModel) - if err != nil { - return nil, err - } - - bedrockBody, err := PrepareBedrockRequestBodyWithTokens(body, mappedModel, betaTokens, false) - if err != nil { - return nil, fmt.Errorf("prepare bedrock request body: %w", err) - } - - proxyURL := "" - if account.ProxyID != nil && account.Proxy != nil { - proxyURL = account.Proxy.URL() - } - - logger.LegacyPrintf("service.gateway", "[Bedrock] 命中 Bedrock 分支: account=%d name=%s model=%s->%s stream=%v", - account.ID, account.Name, reqModel, mappedModel, reqStream) - - // 根据账号类型选择认证方式 - var signer *BedrockSigner - var bedrockAPIKey string - if account.IsBedrockAPIKey() { - bedrockAPIKey = account.GetCredential("api_key") - if bedrockAPIKey == "" { - return nil, fmt.Errorf("api_key not found in bedrock credentials") - } - } else { - signer, err = NewBedrockSignerFromAccount(account) - if err != nil { - return nil, fmt.Errorf("create bedrock signer: %w", err) - } - } - - // 执行上游请求(含重试) - resp, err := s.executeBedrockUpstream(ctx, c, account, bedrockBody, mappedModel, region, reqStream, signer, bedrockAPIKey, proxyURL) - if err != nil { - return nil, err - } - defer func() { _ = resp.Body.Close() }() - - // 将 Bedrock 的 x-amzn-requestid 映射到 x-request-id, - // 使通用错误处理函数(handleErrorResponse、handleRetryExhaustedError)能正确提取 AWS request ID。 - if awsReqID := resp.Header.Get("x-amzn-requestid"); awsReqID != "" && resp.Header.Get("x-request-id") == "" { - resp.Header.Set("x-request-id", awsReqID) - } - - // 错误/failover 处理 - if resp.StatusCode >= 400 { - return s.handleBedrockUpstreamErrors(ctx, resp, c, account) - } - - // Bedrock 分支绕过通用 Forward 成功路径,这里保持上游接受回调语义一致。 - if parsed.OnUpstreamAccepted != nil { - parsed.OnUpstreamAccepted() - } - - // 响应处理 - var usage *ClaudeUsage - var firstTokenMs *int - var clientDisconnect bool - if reqStream { - streamResult, err := s.handleBedrockStreamingResponse(ctx, resp, c, account, startTime, reqModel) - if err != nil { - return nil, err - } - usage = streamResult.usage - firstTokenMs = streamResult.firstTokenMs - clientDisconnect = streamResult.clientDisconnect - } else { - usage, err = s.handleBedrockNonStreamingResponse(ctx, resp, c, account) - if err != nil { - return nil, err - } - } - if usage == nil { - usage = &ClaudeUsage{} - } - - return &ForwardResult{ - RequestID: resp.Header.Get("x-amzn-requestid"), - Usage: *usage, - Model: reqModel, - UpstreamModel: mappedModel, - Stream: reqStream, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - ClientDisconnect: clientDisconnect, - }, nil -} - -// executeBedrockUpstream 执行 Bedrock 上游请求(含重试逻辑) -func (s *GatewayService) executeBedrockUpstream( - ctx context.Context, - c *gin.Context, - account *Account, - body []byte, - modelID string, - region string, - stream bool, - signer *BedrockSigner, - apiKey string, - proxyURL string, -) (*http.Response, error) { - var resp *http.Response - var err error - retryStart := time.Now() - for attempt := 1; attempt <= maxRetryAttempts; attempt++ { - var upstreamReq *http.Request - if account.IsBedrockAPIKey() { - upstreamReq, err = s.buildUpstreamRequestBedrockAPIKey(ctx, body, modelID, region, stream, apiKey) - } else { - upstreamReq, err = s.buildUpstreamRequestBedrock(ctx, body, modelID, region, stream, signer) - } - if err != nil { - return nil, err - } - - resp, err = s.httpUpstream.DoWithTLS(upstreamReq, proxyURL, account.ID, account.Concurrency, nil) - if err != nil { - if resp != nil && resp.Body != nil { - _ = resp.Body.Close() - } - safeErr := sanitizeUpstreamErrorMessage(err.Error()) - setOpsUpstreamError(c, 0, safeErr, "") - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: 0, - UpstreamURL: safeUpstreamURL(upstreamReq.URL.String()), - Kind: "request_error", - Message: safeErr, - }) - c.JSON(http.StatusBadGateway, gin.H{ - "type": "error", - "error": gin.H{ - "type": "upstream_error", - "message": "Upstream request failed", - }, - }) - return nil, fmt.Errorf("upstream request failed: %s", safeErr) - } - - if resp.StatusCode >= 400 && resp.StatusCode != 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) { - if attempt < maxRetryAttempts { - elapsed := time.Since(retryStart) - if elapsed >= maxRetryElapsed { - break - } - - delay := retryBackoffDelay(attempt) - remaining := maxRetryElapsed - elapsed - if delay > remaining { - delay = remaining - } - if delay <= 0 { - break - } - - respBody, _ := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamURL: safeUpstreamURL(upstreamReq.URL.String()), - Kind: "retry", - Message: extractUpstreamErrorMessage(respBody), - Detail: func() string { - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - return truncateString(string(respBody), s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes) - } - return "" - }(), - }) - logger.LegacyPrintf("service.gateway", "[Bedrock] account %d: upstream error %d, retry %d/%d after %v", - account.ID, resp.StatusCode, attempt, maxRetryAttempts, delay) - if err := sleepWithContext(ctx, delay); err != nil { - return nil, err - } - continue - } - break - } - - break - } - if resp == nil || resp.Body == nil { - return nil, errors.New("upstream request failed: empty response") - } - return resp, nil -} - -// handleBedrockUpstreamErrors 处理 Bedrock 上游 4xx/5xx 错误(failover + 错误响应) -func (s *GatewayService) handleBedrockUpstreamErrors( - ctx context.Context, - resp *http.Response, - c *gin.Context, - account *Account, -) (*ForwardResult, error) { - // retry exhausted + failover - if s.shouldRetryUpstreamError(account, resp.StatusCode) { - if s.shouldFailoverUpstreamError(resp.StatusCode) { - respBody, _ := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) - - logger.LegacyPrintf("service.gateway", "[Bedrock] Upstream error (retry exhausted, failover): Account=%d(%s) Status=%d Body=%s", - account.ID, account.Name, resp.StatusCode, truncateString(string(respBody), 1000)) - - s.handleRetryExhaustedSideEffects(ctx, resp, account) - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - Kind: "retry_exhausted_failover", - Message: extractUpstreamErrorMessage(respBody), - }) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), - } - } - return s.handleRetryExhaustedError(ctx, resp, c, account) - } - - // non-retryable failover - if s.shouldFailoverUpstreamError(resp.StatusCode) { - respBody, _ := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) - - s.handleFailoverSideEffects(ctx, resp, account) - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - Kind: "failover", - Message: extractUpstreamErrorMessage(respBody), - }) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), - } - } - - // other errors - return s.handleErrorResponse(ctx, resp, c, account) -} - -// buildUpstreamRequestBedrock 构建 Bedrock 上游请求 -func (s *GatewayService) buildUpstreamRequestBedrock( - ctx context.Context, - body []byte, - modelID string, - region string, - stream bool, - signer *BedrockSigner, -) (*http.Request, error) { - targetURL := BuildBedrockURL(region, modelID, stream) - - req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) - if err != nil { - return nil, err - } - - req.Header.Set("Content-Type", "application/json") - req.Header.Set("Accept", "application/json") - - // SigV4 签名 - if err := signer.SignRequest(ctx, req, body); err != nil { - return nil, fmt.Errorf("sign bedrock request: %w", err) - } - - return req, nil -} - -// buildUpstreamRequestBedrockAPIKey 构建 Bedrock API Key (Bearer Token) 上游请求 -func (s *GatewayService) buildUpstreamRequestBedrockAPIKey( - ctx context.Context, - body []byte, - modelID string, - region string, - stream bool, - apiKey string, -) (*http.Request, error) { - targetURL := BuildBedrockURL(region, modelID, stream) - - req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) - if err != nil { - return nil, err - } - - req.Header.Set("Content-Type", "application/json") - req.Header.Set("Accept", "application/json") - req.Header.Set("Authorization", "Bearer "+apiKey) - - return req, nil -} - -// handleBedrockNonStreamingResponse 处理 Bedrock 非流式响应 -// Bedrock InvokeModel 非流式响应的 body 格式与 Claude API 兼容 -func (s *GatewayService) handleBedrockNonStreamingResponse( - ctx context.Context, - resp *http.Response, - c *gin.Context, - account *Account, -) (*ClaudeUsage, error) { - body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, anthropicTooLargeError) - if err != nil { - return nil, err - } - - // 转换 Bedrock 特有的 amazon-bedrock-invocationMetrics 为标准 Anthropic usage 格式 - // 并移除该字段避免透传给客户端 - body = transformBedrockInvocationMetrics(body) - - usage := parseClaudeUsageFromResponseBody(body) - - c.Header("Content-Type", "application/json") - if v := resp.Header.Get("x-amzn-requestid"); v != "" { - c.Header("x-request-id", v) - } - c.Data(resp.StatusCode, "application/json", body) - return usage, nil -} - func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, reqStream bool, mimicClaudeCode bool) (*http.Request, []byte, error) { if account.Platform == PlatformAnthropic && account.Type == AccountTypeServiceAccount { req, err := s.buildUpstreamRequestAnthropicVertex(ctx, c, account, body, token, modelID, reqStream) diff --git a/backend/internal/service/openai_gateway_cc_pipeline.go b/backend/internal/service/openai_gateway_cc_pipeline.go new file mode 100644 index 0000000000..816f5a26e4 --- /dev/null +++ b/backend/internal/service/openai_gateway_cc_pipeline.go @@ -0,0 +1,324 @@ +package service + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" + "github.com/gin-gonic/gin" + "go.uber.org/zap" +) + +// 本文件收敛三个 CC(Chat Completions)forwarder 之间重复的 HTTP 管线与 SSE +// 循环骨架(PR #3802 遗留项): +// +// - forwardAsRawChatCompletions (原生 CC 直转) +// - forwardResponsesViaRawChatCompletions(/v1/responses → CC 回退) +// - forwardAnthropicViaRawChatCompletions(/v1/messages → CC 回退) +// +// 以及 messages / chat_completions 两条 Responses 主路径中逐字相同的错误处理块。 +// 所有 helper 都是对既有内联代码的等价提取,不改变任何行为;各路径的差异 +// (GLM effort 归一化、fast policy、Grok 分支、ClientDisconnect 语义等)仍留在 +// 调用方,属于有意保留的行为差异,不在此强行统一。 + +// newUpstreamSSEScanner 构造读取上游 SSE 流的行扫描器,按配置放大单行上限。 +func (s *OpenAIGatewayService) newUpstreamSSEScanner(r io.Reader) *bufio.Scanner { + scanner := bufio.NewScanner(r) + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + return scanner +} + +// newStreamHeaderWriter 返回幂等的 SSE 响应头写入闭包:首次调用时透传过滤后的 +// 上游响应头并写入标准 SSE 头 + 200 状态码,后续调用为 no-op。延迟到首个事件 +// 写出前才提交响应头,使上游早期失败仍可改走 failover 或非流式错误响应。 +func (s *OpenAIGatewayService) newStreamHeaderWriter(c *gin.Context, upstream http.Header) func() { + headersWritten := false + return func() { + if headersWritten { + return + } + headersWritten = true + if s.responseHeaderFilter != nil { + responseheaders.WriteFilteredHeaders(c.Writer.Header(), upstream, s.responseHeaderFilter) + } + c.Writer.Header().Set("Content-Type", "text/event-stream") + c.Writer.Header().Set("Cache-Control", "no-cache") + c.Writer.Header().Set("Connection", "keep-alive") + c.Writer.Header().Set("X-Accel-Buffering", "no") + c.Writer.WriteHeader(http.StatusOK) + } +} + +// readOpenAIUpstreamError 读取上游错误体并把 resp.Body 回卷为可重读的副本 +// (下游 handleXxxErrorResponse 需要再次读取),返回原始错误体与脱敏后的 +// 上游错误消息。 +func (s *OpenAIGatewayService) readOpenAIUpstreamError(resp *http.Response) ([]byte, string) { + respBody := s.readUpstreamErrorBody(resp) + _ = resp.Body.Close() + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + + upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) + upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) + return respBody, upstreamMsg +} + +// failoverOpenAIUpstreamHTTPError 对 >=400 的上游响应做 failover 判定:命中时 +// 记录 ops 事件、执行账号级错误处置并返回 *UpstreamFailoverError;未命中返回 +// nil,调用方继续走各自端点格式的非 failover 错误处理链。 +func (s *OpenAIGatewayService) failoverOpenAIUpstreamHTTPError( + ctx context.Context, + c *gin.Context, + account *Account, + resp *http.Response, + respBody []byte, + upstreamMsg string, + upstreamModel string, +) *UpstreamFailoverError { + if !s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { + return nil + } + upstreamDetail := "" + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes + if maxBytes <= 0 { + maxBytes = 2048 + } + upstreamDetail = truncateString(string(respBody), maxBytes) + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + Kind: "failover", + Message: upstreamMsg, + Detail: upstreamDetail, + }) + s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) + return &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)), + } +} + +// openAIChatCompletionsTargetURL 解析账号的(非 Grok)Chat Completions 上游端点。 +func (s *OpenAIGatewayService) openAIChatCompletionsTargetURL(account *Account) (string, error) { + baseURL := account.GetOpenAIBaseURL() + if baseURL == "" { + baseURL = "https://api.openai.com" + } + validatedURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return "", fmt.Errorf("invalid base_url: %w", err) + } + return buildOpenAIChatCompletionsURL(validatedURL), nil +} + +// resolveCCFallbackTarget 解析两条 CC 回退路径共用的账号凭证与上游端点 +// (回退路径仅面向 APIKey 账号,凭证恒为 openai api_key)。 +func (s *OpenAIGatewayService) resolveCCFallbackTarget(account *Account) (apiKey string, targetURL string, err error) { + apiKey = account.GetOpenAIApiKey() + if apiKey == "" { + return "", "", fmt.Errorf("account %d missing api_key", account.ID) + } + targetURL, err = s.openAIChatCompletionsTargetURL(account) + if err != nil { + return "", "", err + } + return apiKey, targetURL, nil +} + +// sendCCUpstreamRequest 构建并发送 CC 上游请求:分离的上游 context、OpenAI HTTP +// profile、标准头(含流式 Accept 切换)、客户端 header 白名单透传、自定义 UA 与 +// 账号级 header 覆写,最后经代理发出。传输层失败(DNS/TCP/TLS,无 HTTP 响应) +// 统一由 handleOpenAIUpstreamTransportError 归一为 failover。 +// +// userAgent 为空时保留默认 UA;Grok 的默认 UA 兜底由调用方解析后传入。 +func (s *OpenAIGatewayService) sendCCUpstreamRequest( + ctx context.Context, + c *gin.Context, + account *Account, + targetURL string, + body []byte, + stream bool, + bearerToken string, + userAgent string, +) (*http.Response, error) { + upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) + upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(body)) + releaseUpstreamCtx() + if err != nil { + return nil, fmt.Errorf("build upstream request: %w", err) + } + upstreamReq = upstreamReq.WithContext(WithHTTPUpstreamProfile(upstreamReq.Context(), HTTPUpstreamProfileOpenAI)) + upstreamReq.Header.Set("Content-Type", "application/json") + upstreamReq.Header.Set("Authorization", "Bearer "+bearerToken) + if stream { + upstreamReq.Header.Set("Accept", "text/event-stream") + } else { + upstreamReq.Header.Set("Accept", "application/json") + } + + // 透传白名单中的客户端 header。详见 openaiCCRawAllowedHeaders 的设计说明。 + for key, values := range c.Request.Header { + lowerKey := strings.ToLower(key) + if openaiCCRawAllowedHeaders[lowerKey] { + for _, v := range values { + upstreamReq.Header.Add(key, v) + } + } + } + if userAgent != "" { + upstreamReq.Header.Set("user-agent", userAgent) + } + + // 账号级请求头覆写(仅 openai api_key 账号启用时生效) + account.ApplyHeaderOverrides(upstreamReq.Header) + + proxyURL := "" + if account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false) + } + return resp, nil +} + +// ccStreamScanState 是 scanCCStream 返回的读取状态快照。 +type ccStreamScanState struct { + // Usage 为 include_usage chunk 中最近一次出现的用量(上游可能重复发送, + // 总是保留最新值);终态事件中的用量由调用方在 finalize 阶段自行覆盖。 + Usage OpenAIUsage + // FirstTokenMs 为首个实际输出 chunk(排除 usage-only chunk)的到达时延。 + FirstTokenMs *int + // SawDone 表示上游发出了 [DONE] 哨兵。 + SawDone bool + // Err 为 scanner 读错误(客户端 context 取消不属于此类,会原样带出)。 + // 非 nil 时调用方必须跳过 finalize 并返回 usage-incomplete 错误,避免 + // 把上游截断伪装成正常收尾。 + Err error +} + +// scanCCStream 驱动两条 CC 回退路径共享的 SSE 读循环:提取 data 行、在 [DONE] +// 哨兵处停止、保留最新 usage、记录首 token 时延,并把每个解析成功的 chunk 交给 +// emit 回调做各自的协议转换与写出。读错误按既有约定过滤 context 取消类噪声后 +// 记入 Warn 日志。 +func (s *OpenAIGatewayService) scanCCStream( + resp *http.Response, + logPrefix string, + requestID string, + startTime time.Time, + emit func(*apicompat.ChatCompletionsChunk), +) ccStreamScanState { + var st ccStreamScanState + + scanner := s.newUpstreamSSEScanner(resp.Body) + for scanner.Scan() { + line := scanner.Text() + payload, ok := extractOpenAISSEDataLine(line) + if !ok { + continue + } + payload = strings.TrimSpace(payload) + if payload == "" { + continue + } + if payload == "[DONE]" { + st.SawDone = true + break + } + + if u := extractCCStreamUsage(payload); u != nil { + st.Usage = *u + } + + var chunk apicompat.ChatCompletionsChunk + if err := json.Unmarshal([]byte(payload), &chunk); err != nil { + logger.L().Warn(logPrefix+": failed to parse chat stream chunk", + zap.Error(err), + zap.String("request_id", requestID), + ) + continue + } + if st.FirstTokenMs == nil && !isOpenAIChatUsageOnlyStreamChunk(payload) && chatChunkStartsResponsesOutput(&chunk) { + ms := int(time.Since(startTime).Milliseconds()) + st.FirstTokenMs = &ms + } + emit(&chunk) + } + + if err := scanner.Err(); err != nil { + if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + logger.L().Warn(logPrefix+": stream read error", + zap.Error(err), + zap.String("request_id", requestID), + ) + } + st.Err = err + } + return st +} + +// logCCStreamMissingDoneSentinel 记录"上游未发 [DONE] 哨兵即结束"的 debug 日志。 +func logCCStreamMissingDoneSentinel(logPrefix, requestID string) { + logger.L().Debug(logPrefix+": upstream stream ended without done sentinel", + zap.String("request_id", requestID), + ) +} + +// readCCUpstreamJSONResponse 读取并解析 CC 非流式 JSON 响应,失败时以调用方 +// 端点格式回写错误;成功时顺带提取 usage。 +func (s *OpenAIGatewayService) readCCUpstreamJSONResponse( + c *gin.Context, + resp *http.Response, + writeError compatErrorWriter, +) (*apicompat.ChatCompletionsResponse, OpenAIUsage, error) { + respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) + if err != nil { + if !errors.Is(err, ErrUpstreamResponseBodyTooLarge) { + writeError(c, http.StatusBadGateway, "api_error", "Failed to read upstream response") + } + return nil, OpenAIUsage{}, fmt.Errorf("read upstream body: %w", err) + } + + var ccResp apicompat.ChatCompletionsResponse + if err := json.Unmarshal(respBody, &ccResp); err != nil { + writeError(c, http.StatusBadGateway, "api_error", "Failed to parse upstream response") + return nil, OpenAIUsage{}, fmt.Errorf("parse chat completions response: %w", err) + } + + usage := OpenAIUsage{} + if parsed, ok := extractOpenAIUsageFromJSONBytes(respBody); ok { + usage = parsed + } + return &ccResp, usage, nil +} + +// writeOpenAIResponsesFallbackError 以 /v1/responses 回退路径的既有错误格式回写 +// (裸 error 对象;不调用 MarkResponseCommitted,与原内联写法保持一致)。 +func writeOpenAIResponsesFallbackError(c *gin.Context, statusCode int, errType, message string) { + c.JSON(statusCode, gin.H{ + "error": gin.H{ + "type": errType, + "message": message, + }, + }) +} diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 20361a6952..3c086d5ee4 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -1,13 +1,10 @@ package service import ( - "bufio" - "bytes" "context" "encoding/json" "errors" "fmt" - "io" "net/http" "strings" "sync/atomic" @@ -269,12 +266,7 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( // 8. Handle error response with failover if resp.StatusCode >= 400 { - respBody := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) - - upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) - upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) if account.Type == AccountTypeAPIKey && openai_compat.ResolveResponsesSupport(account.Extra) == openai_compat.ResponsesSupportUnknown && !isResponsesEndpointSupportedByStatus(resp.StatusCode) { @@ -285,31 +277,8 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( ) return s.forwardAsRawChatCompletions(ctx, c, account, body, defaultMappedModel) } - if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { - upstreamDetail := "" - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes - if maxBytes <= 0 { - maxBytes = 2048 - } - upstreamDetail = truncateString(string(respBody), maxBytes) - } - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Kind: "failover", - Message: upstreamMsg, - Detail: upstreamDetail, - }) - s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)), - } + if foErr := s.failoverOpenAIUpstreamHTTPError(ctx, c, account, resp, respBody, upstreamMsg, upstreamModel); foErr != nil { + return nil, foErr } return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel) } @@ -500,22 +469,7 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( requestBodyLen int, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - - headersWritten := false - writeStreamHeaders := func() { - if headersWritten { - return - } - headersWritten = true - if s.responseHeaderFilter != nil { - responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - } - c.Writer.Header().Set("Content-Type", "text/event-stream") - c.Writer.Header().Set("Cache-Control", "no-cache") - c.Writer.Header().Set("Connection", "keep-alive") - c.Writer.Header().Set("X-Accel-Buffering", "no") - c.Writer.WriteHeader(http.StatusOK) - } + writeStreamHeaders := s.newStreamHeaderWriter(c, resp.Header) state := apicompat.NewResponsesEventToChatState() state.Model = originalModel @@ -533,12 +487,7 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( var streamFailoverErr *UpstreamFailoverError var streamNonFailoverErr error - scanner := bufio.NewScanner(resp.Body) - maxLineSize := defaultMaxLineSize - if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { - maxLineSize = s.cfg.Gateway.MaxLineSize - } - scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + scanner := s.newUpstreamSSEScanner(resp.Body) streamInterval := time.Duration(0) if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 { diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 023440f94f..9b31b803d2 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -1,13 +1,10 @@ package service import ( - "bufio" - "bytes" "context" "encoding/json" "errors" "fmt" - "io" "net/http" "strings" "time" @@ -148,69 +145,26 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( zap.Bool("stream", clientStream), ) - // 5. Build upstream request + // 5. Build and send upstream request via the shared CC pipeline targetURL, err := s.rawChatCompletionsURL(account) if err != nil { return nil, err } - - upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) - upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(upstreamBody)) - releaseUpstreamCtx() - if err != nil { - return nil, fmt.Errorf("build upstream request: %w", err) - } - upstreamReq = upstreamReq.WithContext(WithHTTPUpstreamProfile(upstreamReq.Context(), HTTPUpstreamProfileOpenAI)) - upstreamReq.Header.Set("Content-Type", "application/json") - upstreamReq.Header.Set("Authorization", "Bearer "+token) - if clientStream { - upstreamReq.Header.Set("Accept", "text/event-stream") - } else { - upstreamReq.Header.Set("Accept", "application/json") - } - - // 透传白名单中的客户端 header。详见 openaiCCRawAllowedHeaders 的设计说明。 - for key, values := range c.Request.Header { - lowerKey := strings.ToLower(key) - if openaiCCRawAllowedHeaders[lowerKey] { - for _, v := range values { - upstreamReq.Header.Add(key, v) - } - } - } customUA := account.GetOpenAIUserAgent() - if customUA != "" { - upstreamReq.Header.Set("user-agent", customUA) - } else if account.Platform == PlatformGrok { - upstreamReq.Header.Set("user-agent", "sub2api-grok/1.0") + if customUA == "" && account.Platform == PlatformGrok { + customUA = "sub2api-grok/1.0" } - - // 账号级请求头覆写(仅 openai api_key 账号启用时生效) - account.ApplyHeaderOverrides(upstreamReq.Header) - - // 6. Send request - proxyURL := "" - if account.Proxy != nil { - proxyURL = account.Proxy.URL() - } - resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, upstreamBody, clientStream, token, customUA) if err != nil { - return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false) + return nil, err } defer func() { _ = resp.Body.Close() }() // 7. Handle error response with failover if resp.StatusCode >= 400 { - respBody := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) if account.Platform == PlatformGrok { s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) - } - - upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) - upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) - if account.Platform == PlatformGrok { appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ Platform: account.Platform, AccountID: account.ID, @@ -230,31 +184,8 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel) } - if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { - upstreamDetail := "" - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes - if maxBytes <= 0 { - maxBytes = 2048 - } - upstreamDetail = truncateString(string(respBody), maxBytes) - } - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Kind: "failover", - Message: upstreamMsg, - Detail: upstreamDetail, - }) - s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)), - } + if foErr := s.failoverOpenAIUpstreamHTTPError(ctx, c, account, resp, respBody, upstreamMsg, upstreamModel); foErr != nil { + return nil, foErr } return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel) } @@ -286,15 +217,7 @@ func (s *OpenAIGatewayService) rawChatCompletionsURL(account *Account) (string, return targetURL, nil } - baseURL := account.GetOpenAIBaseURL() - if baseURL == "" { - baseURL = "https://api.openai.com" - } - validatedURL, err := s.validateUpstreamBaseURL(baseURL) - if err != nil { - return "", fmt.Errorf("invalid base_url: %w", err) - } - return buildOpenAIChatCompletionsURL(validatedURL), nil + return s.openAIChatCompletionsTargetURL(account) } // streamRawChatCompletions 透传上游 CC SSE 流到客户端,并提取 usage(包括 @@ -316,29 +239,8 @@ func (s *OpenAIGatewayService) streamRawChatCompletions( requestBodyLen int, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - - headersWritten := false - writeStreamHeaders := func() { - if headersWritten { - return - } - headersWritten = true - if s.responseHeaderFilter != nil { - responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - } - c.Writer.Header().Set("Content-Type", "text/event-stream") - c.Writer.Header().Set("Cache-Control", "no-cache") - c.Writer.Header().Set("Connection", "keep-alive") - c.Writer.Header().Set("X-Accel-Buffering", "no") - c.Writer.WriteHeader(http.StatusOK) - } - - scanner := bufio.NewScanner(resp.Body) - maxLineSize := defaultMaxLineSize - if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { - maxLineSize = s.cfg.Gateway.MaxLineSize - } - scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + writeStreamHeaders := s.newStreamHeaderWriter(c, resp.Header) + scanner := s.newUpstreamSSEScanner(resp.Body) var usage OpenAIUsage var firstTokenMs *int diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 2afd155646..8c5c801944 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -1,13 +1,10 @@ package service import ( - "bufio" - "bytes" "context" "encoding/json" "errors" "fmt" - "io" "net/http" "strings" "sync/atomic" @@ -322,16 +319,12 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( // 8. Handle error response with failover if resp.StatusCode >= 400 { - respBody := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) if account.Platform == PlatformGrok { s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) } - upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) - upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) if previousResponseID != "" && (isOpenAICompatPreviousResponseNotFound(resp.StatusCode, upstreamMsg, respBody) || isOpenAICompatPreviousResponseUnsupported(resp.StatusCode, upstreamMsg, respBody)) { if isOpenAICompatPreviousResponseUnsupported(resp.StatusCode, upstreamMsg, respBody) { s.disableOpenAICompatSessionContinuation(ctx, c, account, promptCacheKey) @@ -345,31 +338,8 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( ) return s.ForwardAsAnthropic(ctx, c, account, body, promptCacheKey, defaultMappedModel) } - if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { - upstreamDetail := "" - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes - if maxBytes <= 0 { - maxBytes = 2048 - } - upstreamDetail = truncateString(string(respBody), maxBytes) - } - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Kind: "failover", - Message: upstreamMsg, - Detail: upstreamDetail, - }) - s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)), - } + if foErr := s.failoverOpenAIUpstreamHTTPError(ctx, c, account, resp, respBody, upstreamMsg, upstreamModel); foErr != nil { + return nil, foErr } // Non-failover error: return Anthropic-formatted error to client return s.handleAnthropicErrorResponse(resp, c, account, billingModel) @@ -573,12 +543,7 @@ func (s *OpenAIGatewayService) readOpenAICompatBufferedTerminal( return nil, usage, acc, errors.New("upstream response body is nil") } - scanner := bufio.NewScanner(resp.Body) - maxLineSize := defaultMaxLineSize - if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { - maxLineSize = s.cfg.Gateway.MaxLineSize - } - scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + scanner := s.newUpstreamSSEScanner(resp.Body) streamInterval := time.Duration(0) if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 { @@ -737,22 +702,7 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( startTime time.Time, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - - headersWritten := false - writeStreamHeaders := func() { - if headersWritten { - return - } - headersWritten = true - if s.responseHeaderFilter != nil { - responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - } - c.Writer.Header().Set("Content-Type", "text/event-stream") - c.Writer.Header().Set("Cache-Control", "no-cache") - c.Writer.Header().Set("Connection", "keep-alive") - c.Writer.Header().Set("X-Accel-Buffering", "no") - c.Writer.WriteHeader(http.StatusOK) - } + writeStreamHeaders := s.newStreamHeaderWriter(c, resp.Header) state := apicompat.NewResponsesEventToAnthropicState() state.Model = originalModel @@ -763,12 +713,7 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( clientDisconnected := false clientOutputStarted := false - scanner := bufio.NewScanner(resp.Body) - maxLineSize := defaultMaxLineSize - if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { - maxLineSize = s.cfg.Gateway.MaxLineSize - } - scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + scanner := s.newUpstreamSSEScanner(resp.Body) streamInterval := time.Duration(0) if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 { diff --git a/backend/internal/service/openai_gateway_messages_chat_fallback.go b/backend/internal/service/openai_gateway_messages_chat_fallback.go index c33aa911ad..4596177468 100644 --- a/backend/internal/service/openai_gateway_messages_chat_fallback.go +++ b/backend/internal/service/openai_gateway_messages_chat_fallback.go @@ -1,13 +1,9 @@ package service import ( - "bufio" - "bytes" "context" "encoding/json" - "errors" "fmt" - "io" "net/http" "strings" "time" @@ -100,91 +96,22 @@ func (s *OpenAIGatewayService) forwardAnthropicViaRawChatCompletions( zap.Bool("stream", clientStream), ) - // 3. Build upstream request - apiKey := account.GetOpenAIApiKey() - if apiKey == "" { - return nil, fmt.Errorf("account %d missing api_key", account.ID) - } - baseURL := account.GetOpenAIBaseURL() - if baseURL == "" { - baseURL = "https://api.openai.com" - } - validatedURL, err := s.validateUpstreamBaseURL(baseURL) + // 3. Build and send upstream request via the shared CC pipeline + apiKey, targetURL, err := s.resolveCCFallbackTarget(account) if err != nil { - return nil, fmt.Errorf("invalid base_url: %w", err) + return nil, err } - targetURL := buildOpenAIChatCompletionsURL(validatedURL) - - upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) - upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(chatBody)) - releaseUpstreamCtx() + resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent()) if err != nil { - return nil, fmt.Errorf("build upstream request: %w", err) - } - upstreamReq = upstreamReq.WithContext(WithHTTPUpstreamProfile(upstreamReq.Context(), HTTPUpstreamProfileOpenAI)) - upstreamReq.Header.Set("Content-Type", "application/json") - upstreamReq.Header.Set("Authorization", "Bearer "+apiKey) - if clientStream { - upstreamReq.Header.Set("Accept", "text/event-stream") - } else { - upstreamReq.Header.Set("Accept", "application/json") - } - for key, values := range c.Request.Header { - lowerKey := strings.ToLower(key) - if openaiCCRawAllowedHeaders[lowerKey] { - for _, v := range values { - upstreamReq.Header.Add(key, v) - } - } - } - if customUA := account.GetOpenAIUserAgent(); customUA != "" { - upstreamReq.Header.Set("user-agent", customUA) - } - account.ApplyHeaderOverrides(upstreamReq.Header) - - proxyURL := "" - if account.Proxy != nil { - proxyURL = account.Proxy.URL() - } - resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) - if err != nil { - return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false) + return nil, err } defer func() { _ = resp.Body.Close() }() // 4. Handle error responses if resp.StatusCode >= 400 { - respBody := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) - - upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) - upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) - if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { - upstreamDetail := "" - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes - if maxBytes <= 0 { - maxBytes = 2048 - } - upstreamDetail = truncateString(string(respBody), maxBytes) - } - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Kind: "failover", - Message: upstreamMsg, - Detail: upstreamDetail, - }) - s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)), - } + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) + if foErr := s.failoverOpenAIUpstreamHTTPError(ctx, c, account, resp, respBody, upstreamMsg, upstreamModel); foErr != nil { + return nil, foErr } // Non-failover error: return Anthropic-formatted error to client via the // shared compat handler (passthrough rules, ops recording, cyber_policy). @@ -209,28 +136,14 @@ func (s *OpenAIGatewayService) bufferChatCompletionsAsAnthropic( startTime time.Time, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) + ccResp, usage, err := s.readCCUpstreamJSONResponse(c, resp, writeAnthropicError) if err != nil { - if !errors.Is(err, ErrUpstreamResponseBodyTooLarge) { - writeAnthropicError(c, http.StatusBadGateway, "api_error", "Failed to read upstream response") - } - return nil, fmt.Errorf("read upstream body: %w", err) + return nil, err } - - var ccResp apicompat.ChatCompletionsResponse - if err := json.Unmarshal(respBody, &ccResp); err != nil { - writeAnthropicError(c, http.StatusBadGateway, "api_error", "Failed to parse upstream response") - return nil, fmt.Errorf("parse chat completions response: %w", err) - } - responsesResp := apicompat.ChatCompletionsResponseToResponses(&ccResp, originalModel) + responsesResp := apicompat.ChatCompletionsResponseToResponses(ccResp, originalModel) anthropicResp := apicompat.ResponsesToAnthropic(responsesResp, originalModel) - usage := OpenAIUsage{} - if parsed, ok := extractOpenAIUsageFromJSONBytes(respBody); ok { - usage = parsed - } - if s.responseHeaderFilter != nil { responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) } @@ -260,71 +173,18 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsAnthropic( startTime time.Time, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - headersWritten := false - writeStreamHeaders := func() { - if headersWritten { - return - } - headersWritten = true - if s.responseHeaderFilter != nil { - responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - } - c.Writer.Header().Set("Content-Type", "text/event-stream") - c.Writer.Header().Set("Cache-Control", "no-cache") - c.Writer.Header().Set("Connection", "keep-alive") - c.Writer.Header().Set("X-Accel-Buffering", "no") - c.Writer.WriteHeader(http.StatusOK) - } + writeStreamHeaders := s.newStreamHeaderWriter(c, resp.Header) ccState := apicompat.NewChatCompletionsToResponsesStreamState(originalModel) anthropicState := apicompat.NewResponsesEventToAnthropicState() anthropicState.Model = originalModel - var usage OpenAIUsage - var firstTokenMs *int clientDisconnected := false - sawDone := false - - scanner := bufio.NewScanner(resp.Body) - maxLineSize := defaultMaxLineSize - if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { - maxLineSize = s.cfg.Gateway.MaxLineSize - } - scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) - - for scanner.Scan() { - line := scanner.Text() - payload, ok := extractOpenAISSEDataLine(line) - if !ok { - continue - } - payload = strings.TrimSpace(payload) - if payload == "" { - continue - } - if payload == "[DONE]" { - sawDone = true - break - } - - if u := extractCCStreamUsage(payload); u != nil { - usage = *u - } - - var chunk apicompat.ChatCompletionsChunk - if err := json.Unmarshal([]byte(payload), &chunk); err != nil { - logger.L().Warn("openai messages chat fallback: failed to parse chat stream chunk", - zap.Error(err), - zap.String("request_id", requestID), - ) - continue - } - if firstTokenMs == nil && !isOpenAIChatUsageOnlyStreamChunk(payload) && chatChunkStartsResponsesOutput(&chunk) { - ms := int(time.Since(startTime).Milliseconds()) - firstTokenMs = &ms - } + // 与 responses 兄弟不同:客户端断开后仍继续做事件转换(喂 anthropicState), + // 仅跳过写出,保证 finalize 阶段的 usage 汇总不受断开影响。 + emitChunk := func(chunk *apicompat.ChatCompletionsChunk) { // CC chunk → Responses events → Anthropic events - responsesEvents := apicompat.ChatCompletionsChunkToResponsesEvents(&chunk, ccState) + responsesEvents := apicompat.ChatCompletionsChunkToResponsesEvents(chunk, ccState) for _, rEvent := range responsesEvents { anthropicEvents := apicompat.ResponsesEventToAnthropicEvents(&rEvent, anthropicState) if clientDisconnected { @@ -347,13 +207,10 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsAnthropic( } } - if err := scanner.Err(); err != nil { - if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { - logger.L().Warn("openai messages chat fallback: stream read error", - zap.Error(err), - zap.String("request_id", requestID), - ) - } + scan := s.scanCCStream(resp, "openai messages chat fallback", requestID, startTime, emitChunk) + usage := scan.Usage + + if scan.Err != nil { // Broken upstream read: skip finalization so no synthetic message_stop // masks the truncation, and surface the error to flag usage incomplete // (mirrors forwardResponsesViaRawChatCompletions). @@ -367,9 +224,9 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsAnthropic( ServiceTier: serviceTier, Stream: true, Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, + FirstTokenMs: scan.FirstTokenMs, ClientDisconnect: clientDisconnected, - }, fmt.Errorf("stream usage incomplete: %w", err) + }, fmt.Errorf("stream usage incomplete: %w", scan.Err) } // Finalize CC→Responses stream (emit response.completed) @@ -397,10 +254,8 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsAnthropic( if !clientDisconnected { c.Writer.Flush() } - if !sawDone { - logger.L().Debug("openai messages chat fallback: upstream stream ended without done sentinel", - zap.String("request_id", requestID), - ) + if !scan.SawDone { + logCCStreamMissingDoneSentinel("openai messages chat fallback", requestID) } return &OpenAIForwardResult{ @@ -413,7 +268,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsAnthropic( ServiceTier: serviceTier, Stream: true, Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, + FirstTokenMs: scan.FirstTokenMs, ClientDisconnect: clientDisconnected, }, nil } diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go new file mode 100644 index 0000000000..c4e0702eb0 --- /dev/null +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -0,0 +1,1088 @@ +package service + +// 本文件由 openai_gateway_service.go 纯移动拆分而来:/v1/responses 直通 +// (passthrough)转发路径及其流式/非流式响应处理与错误处理。仅做代码搬迁, +// 无任何行为变更。 + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "sort" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/openai" + "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" + "go.uber.org/zap" +) + +func (s *OpenAIGatewayService) forwardOpenAIPassthrough( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + reqModel string, + reasoningEffort *string, + reqStream bool, + startTime time.Time, +) (*OpenAIForwardResult, error) { + upstreamPassthroughModel := "" + if isOpenAIResponsesCompactPath(c) { + compactMappedModel := resolveOpenAICompactForwardModel(account, reqModel) + if compactMappedModel != "" && compactMappedModel != reqModel { + nextBody, setErr := sjson.SetBytes(body, "model", compactMappedModel) + if setErr != nil { + return nil, fmt.Errorf("set compact passthrough model: %w", setErr) + } + body = nextBody + upstreamPassthroughModel = compactMappedModel + } + } + + if account != nil && account.Type == AccountTypeOAuth { + if rejectReason := detectOpenAIPassthroughInstructionsRejectReason(reqModel, body); rejectReason != "" { + rejectMsg := "OpenAI codex passthrough requires a non-empty instructions field" + MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalPolicyDenied) + logOpenAIPassthroughInstructionsRejected(ctx, c, account, reqModel, rejectReason, body) + c.JSON(http.StatusForbidden, gin.H{ + "error": gin.H{ + "type": "forbidden_error", + "message": rejectMsg, + }, + }) + return nil, fmt.Errorf("openai passthrough rejected before upstream: %s", rejectReason) + } + + normalizedBody, normalized, err := normalizeOpenAIPassthroughOAuthBody(body, isOpenAIResponsesCompactPath(c)) + if err != nil { + return nil, err + } + if normalized { + body = normalizedBody + } + reqStream = gjson.GetBytes(body, "stream").Bool() + } + + sanitizedBody, sanitized, err := sanitizeEmptyBase64InputImagesInOpenAIBody(body) + if err != nil { + return nil, err + } + if sanitized { + body = sanitizedBody + } + + // Apply OpenAI fast policy to the passthrough body (filter/block by service_tier). + // 统一使用 upstream 视角的 model:透传路径下 body 已经过 compact 映射 + + // OAuth normalize,body 中的 model 字段即上游真正会看到的 slug。 + // 这样可以与 chat-completions / messages / native /responses 入口的 + // upstreamModel 保持一致,避免 whitelist 命中差异。当 body 中没有 + // model 字段时退回 reqModel。 + policyModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) + if policyModel == "" { + policyModel = reqModel + } + updatedBody, policyErr := s.applyOpenAIFastPolicyToBody(ctx, account, policyModel, body) + if policyErr != nil { + var blocked *OpenAIFastBlockedError + if errors.As(policyErr, &blocked) { + writeOpenAIFastPolicyBlockedResponse(c, blocked) + } + return nil, policyErr + } + body = updatedBody + + apiKey := getAPIKeyFromContext(c) + if IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body) && !GroupAllowsImageGeneration(apiKeyGroup(apiKey)) { + MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusForbidden, gin.H{ + "error": gin.H{ + "type": "permission_error", + "message": ImageGenerationPermissionMessage(), + }, + }) + return nil, errors.New("image generation disabled for group") + } + imageBillingModel := "" + imageSizeTier := "" + imageInputSize := "" + if IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body) { + var imageCfgErr error + imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, reqModel) + if imageCfgErr != nil { + setOpsUpstreamError(c, http.StatusBadRequest, imageCfgErr.Error(), "") + c.JSON(http.StatusBadRequest, gin.H{ + "error": gin.H{ + "type": "invalid_request_error", + "message": imageCfgErr.Error(), + "param": "size", + }, + }) + return nil, imageCfgErr + } + imageBillingModel = imageCfg.Model + imageSizeTier = imageCfg.SizeTier + imageInputSize = imageCfg.InputSize + } + + logger.LegacyPrintf("service.openai_gateway", + "[OpenAI 自动透传] 命中自动透传分支: account=%d name=%s type=%s model=%s stream=%v", + account.ID, + account.Name, + account.Type, + reqModel, + reqStream, + ) + if reqStream && c != nil && c.Request != nil { + if timeoutHeaders := collectOpenAIPassthroughTimeoutHeaders(c.Request.Header); len(timeoutHeaders) > 0 { + streamWarnLogger := logger.FromContext(ctx).With( + zap.String("component", "service.openai_gateway"), + zap.Int64("account_id", account.ID), + zap.Strings("timeout_headers", timeoutHeaders), + ) + if s.isOpenAIPassthroughTimeoutHeadersAllowed() { + streamWarnLogger.Warn("OpenAI passthrough 透传请求包含超时相关请求头,且当前配置为放行,可能导致上游提前断流") + } else { + streamWarnLogger.Warn("OpenAI passthrough 检测到超时相关请求头,将按配置过滤以降低断流风险") + } + } + } + + // Get access token + token, _, err := s.GetAccessToken(ctx, account) + if err != nil { + return nil, err + } + + upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) + upstreamReq, err := s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token) + releaseUpstreamCtx() + if err != nil { + return nil, err + } + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + if c != nil { + c.Set("openai_passthrough", true) + } + + upstreamStart := time.Now() + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds()) + if err != nil { + // Transport-level failure (proxy/DNS/TCP/TLS — no HTTP response). Convert to + // a failover so the handler switches to a healthy account, and temporarily + // unschedule the account on durable faults (e.g. rejected proxy credentials). + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 400 { + // 透传模式默认保持原样代理;但 429/529 属于网关必须兜底的 + // 上游容量类错误,应先触发多账号 failover 以维持基础 SLA。 + if shouldFailoverOpenAIPassthroughResponse(resp.StatusCode) { + return nil, s.handleFailoverErrorResponsePassthrough(ctx, resp, c, account, body) + } + return nil, s.handleErrorResponsePassthrough(ctx, resp, c, account, body) + } + + serviceTier := extractOpenAIServiceTierFromBody(body) + + var usage *OpenAIUsage + var firstTokenMs *int + responseID := "" + imageCount := 0 + var imageOutputSizes []string + if reqStream { + result, err := s.handleStreamingResponsePassthrough(ctx, resp, c, account, startTime, reqModel, upstreamPassthroughModel) + if err != nil { + return nil, err + } + usage = result.usage + firstTokenMs = result.firstTokenMs + responseID = strings.TrimSpace(result.responseID) + imageCount = result.imageCount + imageOutputSizes = result.imageOutputSizes + } else { + result, err := s.handleNonStreamingResponsePassthrough(ctx, resp, c, reqModel, upstreamPassthroughModel) + if err != nil { + return nil, err + } + usage = result.usage + responseID = strings.TrimSpace(result.responseID) + imageCount = result.imageCount + imageOutputSizes = result.imageOutputSizes + } + s.bindHTTPResponseAccount(ctx, c, account, responseID) + + // 排除 spark 影子:其 codex_* 仅由 QueryUsage(/wham/usage bengalfox)更新(外审第7轮 P1)。 + if !account.IsShadow() { + if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil { + s.updateCodexUsageSnapshot(ctx, account.ID, snapshot) + } + } + + if usage == nil { + usage = &OpenAIUsage{} + } + + forwardResult := &OpenAIForwardResult{ + RequestID: resp.Header.Get("x-request-id"), + ResponseID: responseID, + Usage: *usage, + Model: reqModel, + UpstreamModel: upstreamPassthroughModel, + ServiceTier: serviceTier, + ReasoningEffort: reasoningEffort, + Stream: reqStream, + OpenAIWSMode: false, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + } + if imageCount > 0 { + forwardResult.ImageCount = imageCount + forwardResult.ImageSize = imageSizeTier + forwardResult.ImageInputSize = imageInputSize + forwardResult.ImageOutputSizes = imageOutputSizes + forwardResult.BillingModel = imageBillingModel + } + return forwardResult, nil +} + +func logOpenAIPassthroughInstructionsRejected( + ctx context.Context, + c *gin.Context, + account *Account, + reqModel string, + rejectReason string, + body []byte, +) { + if ctx == nil { + ctx = context.Background() + } + accountID := int64(0) + accountName := "" + accountType := "" + if account != nil { + accountID = account.ID + accountName = strings.TrimSpace(account.Name) + accountType = strings.TrimSpace(string(account.Type)) + } + fields := []zap.Field{ + zap.String("component", "service.openai_gateway"), + zap.Int64("account_id", accountID), + zap.String("account_name", accountName), + zap.String("account_type", accountType), + zap.String("request_model", strings.TrimSpace(reqModel)), + zap.String("reject_reason", strings.TrimSpace(rejectReason)), + } + fields = appendCodexCLIOnlyRejectedRequestFields(fields, c, body) + logger.FromContext(ctx).With(fields...).Warn("OpenAI passthrough 本地拦截:Codex 请求缺少有效 instructions") +} + +func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + token string, +) (*http.Request, error) { + targetURL := openaiPlatformAPIURL + switch account.Type { + case AccountTypeOAuth: + targetURL = chatgptCodexURL + case AccountTypeAPIKey: + baseURL := account.GetOpenAIBaseURL() + if baseURL != "" { + validatedURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return nil, err + } + targetURL = buildOpenAIResponsesURL(validatedURL) + } + } + targetURL = appendOpenAIResponsesRequestPathSuffix(targetURL, openAIResponsesRequestPathSuffix(c)) + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, err + } + req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) + + // 透传客户端请求头(安全白名单)。 + allowTimeoutHeaders := s.isOpenAIPassthroughTimeoutHeadersAllowed() + if c != nil && c.Request != nil { + for key, values := range c.Request.Header { + lower := strings.ToLower(strings.TrimSpace(key)) + if !isOpenAIPassthroughAllowedRequestHeader(lower, allowTimeoutHeaders) { + continue + } + for _, v := range values { + req.Header.Add(key, v) + } + } + } + + // 覆盖入站鉴权残留,并注入上游认证 + req.Header.Del("authorization") + req.Header.Del("x-api-key") + req.Header.Del("x-goog-api-key") + req.Header.Set("authorization", "Bearer "+token) + + // OAuth 透传到 ChatGPT internal API 时补齐必要头。 + if account.Type == AccountTypeOAuth { + promptCacheKey := strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) + req.Host = "chatgpt.com" + if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, req.Header, account); err != nil { + return nil, fmt.Errorf("resolve chatgpt account headers: %w", err) + } + apiKeyID := getAPIKeyIDFromContext(c) + // 先保存客户端原始值,再做 compact 补充,避免后续统一隔离时读到已处理的值。 + clientSessionID := strings.TrimSpace(req.Header.Get("session_id")) + clientConversationID := strings.TrimSpace(req.Header.Get("conversation_id")) + if isOpenAIResponsesCompactPath(c) { + req.Header.Set("accept", "application/json") + if req.Header.Get("version") == "" { + req.Header.Set("version", codexCLIVersion) + } + if clientSessionID == "" { + clientSessionID = resolveOpenAICompactSessionID(c) + } + } else if req.Header.Get("accept") == "" { + req.Header.Set("accept", "text/event-stream") + } + if req.Header.Get("OpenAI-Beta") == "" { + req.Header.Set("OpenAI-Beta", "responses=experimental") + } + if req.Header.Get("originator") == "" { + req.Header.Set("originator", "codex_cli_rs") + } + // 用隔离后的 session 标识符覆盖客户端透传值,防止跨用户会话碰撞。 + if clientSessionID == "" { + clientSessionID = promptCacheKey + } + if clientConversationID == "" { + clientConversationID = promptCacheKey + } + if clientSessionID != "" { + req.Header.Set("session_id", isolateOpenAISessionID(apiKeyID, clientSessionID)) + } + if clientConversationID != "" { + req.Header.Set("conversation_id", isolateOpenAISessionID(apiKeyID, clientConversationID)) + } + } + + // 透传模式也支持账户自定义 User-Agent 与 ForceCodexCLI 兜底。 + customUA := account.GetOpenAIUserAgent() + if customUA != "" { + req.Header.Set("user-agent", customUA) + } + if s.cfg != nil && s.cfg.Gateway.ForceCodexCLI { + req.Header.Set("user-agent", codexCLIUserAgent) + } + // OAuth 安全透传:对非 Codex UA 统一兜底,降低被上游风控拦截概率。 + if account.Type == AccountTypeOAuth && !openai.IsCodexCLIRequest(req.Header.Get("user-agent")) { + req.Header.Set("user-agent", codexCLIUserAgent) + } + + // 浏览器型 UA 兜底:仅 OAuth(ChatGPT 内部接口)账号生效,若最终 user-agent 仍为浏览器 + // (Chrome/Firefox/Safari/Edge 等),替换为后台配置的 Codex UA,避免 Cloudflare 触发 JS 质询。 + s.overrideBrowserUserAgent(ctx, account, req) + + if req.Header.Get("content-type") == "" { + req.Header.Set("content-type", "application/json") + } + + // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) + account.ApplyHeaderOverrides(req.Header) + + return req, nil +} + +func shouldFailoverOpenAIPassthroughResponse(statusCode int) bool { + switch statusCode { + case http.StatusTooManyRequests, 529: + return true + default: + return false + } +} + +func (s *OpenAIGatewayService) handleFailoverErrorResponsePassthrough( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, + requestBody []byte, +) error { + body := s.readUpstreamErrorBody(resp) + + upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body)) + upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) + upstreamDetail := "" + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes + if maxBytes <= 0 { + maxBytes = 2048 + } + upstreamDetail = truncateString(string(body), maxBytes) + } + setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail) + logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body) + reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody) + _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel) + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + Passthrough: true, + Kind: "failover", + Message: upstreamMsg, + Detail: upstreamDetail, + UpstreamResponseBody: upstreamDetail, + }) + return &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: body, + ResponseHeaders: resp.Header.Clone(), + } +} + +func (s *OpenAIGatewayService) handleErrorResponsePassthrough( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, + requestBody []byte, +) error { + MarkResponseCommitted(c) + body := s.readUpstreamErrorBody(resp) + + // cyber_policy:透传账号本就把原始 body 回给客户端(下方 c.Data),此处仅打标记, + // 供 handler 事后写风控/邮件。cyber 是上游网络安全策略拦截,不冷却账号, + // 故下方跳过 handleOpenAIAccountUpstreamError(避免自定义 temp-unschedulable 规则误冷却)。 + cyberHit, cyberCode, cyberMsg := detectOpenAICyberPolicy(body) + if cyberHit { + MarkOpsCyberPolicy(c, CyberPolicyMark{ + Code: cyberCode, + Message: cyberMsg, + Body: truncateString(string(body), 4096), + UpstreamStatus: resp.StatusCode, + }) + } + + upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body)) + upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) + upstreamDetail := "" + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes + if maxBytes <= 0 { + maxBytes = 2048 + } + upstreamDetail = truncateString(string(body), maxBytes) + } + setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail) + logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body) + // 透传模式保留原始上游错误响应,但运行态账号状态仍需更新, + // 避免粘性路由继续复用刚被限流的账号。cyber 例外:不冷却账号。 + if !cyberHit { + reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody) + _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel) + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + Passthrough: true, + Kind: "http_error", + Message: upstreamMsg, + Detail: upstreamDetail, + UpstreamResponseBody: upstreamDetail, + }) + + writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + contentType := resp.Header.Get("Content-Type") + if contentType == "" { + contentType = "application/json" + } + c.Data(resp.StatusCode, contentType, body) + + if upstreamMsg == "" { + return fmt.Errorf("upstream error: %d", resp.StatusCode) + } + return fmt.Errorf("upstream error: %d message=%s", resp.StatusCode, upstreamMsg) +} + +func isOpenAIPassthroughAllowedRequestHeader(lowerKey string, allowTimeoutHeaders bool) bool { + if lowerKey == "" { + return false + } + if isOpenAIPassthroughTimeoutHeader(lowerKey) { + return allowTimeoutHeaders + } + return openaiPassthroughAllowedHeaders[lowerKey] +} + +func isOpenAIPassthroughTimeoutHeader(lowerKey string) bool { + switch lowerKey { + case "x-stainless-timeout", "x-stainless-read-timeout", "x-stainless-connect-timeout", "x-request-timeout", "request-timeout", "grpc-timeout": + return true + default: + return false + } +} + +func (s *OpenAIGatewayService) isOpenAIPassthroughTimeoutHeadersAllowed() bool { + return s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIPassthroughAllowTimeoutHeaders +} + +func collectOpenAIPassthroughTimeoutHeaders(h http.Header) []string { + if h == nil { + return nil + } + var matched []string + for key, values := range h { + lowerKey := strings.ToLower(strings.TrimSpace(key)) + if isOpenAIPassthroughTimeoutHeader(lowerKey) { + entry := lowerKey + if len(values) > 0 { + entry = fmt.Sprintf("%s=%s", lowerKey, strings.Join(values, "|")) + } + matched = append(matched, entry) + } + } + sort.Strings(matched) + return matched +} + +type openaiStreamingResultPassthrough struct { + usage *OpenAIUsage + firstTokenMs *int + responseID string + imageCount int + imageOutputSizes []string +} + +type openaiNonStreamingResultPassthrough struct { + *OpenAIUsage + usage *OpenAIUsage + responseID string + imageCount int + imageOutputSizes []string +} + +func openAIStreamClientOutputStarted(c *gin.Context, localStarted bool) bool { + if localStarted { + return true + } + return c != nil && c.Writer != nil && c.Writer.Written() +} + +func openAIStreamEventIsPreamble(eventType string) bool { + switch strings.TrimSpace(eventType) { + case "response.created", "response.in_progress": + return true + default: + return false + } +} + +func openAIStreamDataStartsClientOutput(data, eventType string) bool { + trimmed := strings.TrimSpace(data) + if trimmed == "" { + return false + } + if strings.TrimSpace(eventType) == "response.failed" { + return false + } + return !openAIStreamEventIsPreamble(eventType) +} + +func openAIStreamFailedEventShouldFailover(payload []byte, message string) bool { + if isOpenAIContextWindowError(message, payload) { + return false + } + if isOpenAITransientProcessingError(http.StatusBadRequest, message, payload) { + return true + } + code := strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "response.error.code").String())) + if code == "" { + code = strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "error.code").String())) + } + errType := strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "response.error.type").String())) + if errType == "" { + errType = strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "error.type").String())) + } + combined := strings.ToLower(strings.TrimSpace(message + " " + code + " " + errType)) + if combined == "" { + return true + } + nonRetryableMarkers := []string{ + "invalid_request", + "content_policy", + "policy", + "safety", + "high-risk cyber", + "not allowed", + "violat", + } + for _, marker := range nonRetryableMarkers { + if strings.Contains(combined, marker) { + return false + } + } + return true +} + +func (s *OpenAIGatewayService) recordOpenAIStreamUpstreamError( + c *gin.Context, + account *Account, + passthrough bool, + upstreamRequestID string, + kind string, + payload []byte, + message string, +) string { + message = sanitizeUpstreamErrorMessage(strings.TrimSpace(message)) + if message == "" { + message = "OpenAI upstream response failed" + } + detail := "" + if len(payload) > 0 && s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes + if maxBytes <= 0 { + maxBytes = 2048 + } + detail = truncateString(string(payload), maxBytes) + } + if c != nil { + setOpsUpstreamError(c, http.StatusBadGateway, message, detail) + event := OpsUpstreamErrorEvent{ + Platform: PlatformOpenAI, + UpstreamStatusCode: http.StatusBadGateway, + UpstreamRequestID: strings.TrimSpace(upstreamRequestID), + Passthrough: passthrough, + Kind: kind, + Message: message, + Detail: detail, + } + if account != nil { + event.Platform = account.Platform + event.AccountID = account.ID + event.AccountName = account.Name + } + appendOpsUpstreamError(c, event) + } + return message +} + +func (s *OpenAIGatewayService) newOpenAIStreamFailoverError( + c *gin.Context, + account *Account, + passthrough bool, + upstreamRequestID string, + payload []byte, + message string, +) *UpstreamFailoverError { + message = sanitizeUpstreamErrorMessage(strings.TrimSpace(message)) + if message == "" { + message = "OpenAI stream disconnected before completion" + } + message = s.recordOpenAIStreamUpstreamError(c, account, passthrough, upstreamRequestID, "failover", payload, message) + body, _ := json.Marshal(gin.H{ + "error": gin.H{ + "type": "upstream_error", + "message": message, + }, + }) + return &UpstreamFailoverError{ + StatusCode: http.StatusBadGateway, + ResponseBody: body, + } +} + +func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, + startTime time.Time, + originalModel string, + mappedModel string, +) (*openaiStreamingResultPassthrough, error) { + writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + + // SSE headers + c.Header("Content-Type", "text/event-stream") + c.Header("Cache-Control", "no-cache") + c.Header("Connection", "keep-alive") + c.Header("X-Accel-Buffering", "no") + if v := resp.Header.Get("x-request-id"); v != "" { + c.Header("x-request-id", v) + } + + w := c.Writer + flusher, ok := w.(http.Flusher) + if !ok { + return nil, errors.New("streaming not supported") + } + + usage := &OpenAIUsage{} + imageCounter := newOpenAIImageOutputCounter() + var firstTokenMs *int + responseID := "" + clientDisconnected := false + sawDone := false + sawTerminalEvent := false + sawFailedEvent := false + failedMessage := "" + clientOutputStarted := false + upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id")) + pendingLines := make([]string, 0, 8) + writePendingLines := func() bool { + for _, pending := range pendingLines { + if _, err := fmt.Fprintln(w, pending); err != nil { + clientDisconnected = true + logger.LegacyPrintf("service.openai_gateway", "[OpenAI passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) + return false + } + } + pendingLines = pendingLines[:0] + return true + } + + scanner := bufio.NewScanner(resp.Body) + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + scanBuf := getSSEScannerBuf64K() + scanner.Buffer(scanBuf[:0], maxLineSize) + defer putSSEScannerBuf64K(scanBuf) + + needModelReplace := strings.TrimSpace(originalModel) != "" && strings.TrimSpace(mappedModel) != "" && strings.TrimSpace(originalModel) != strings.TrimSpace(mappedModel) + resultWithUsage := func() *openaiStreamingResultPassthrough { + return &openaiStreamingResultPassthrough{ + usage: usage, + firstTokenMs: firstTokenMs, + responseID: responseID, + imageCount: imageCounter.Count(), + imageOutputSizes: imageCounter.Sizes(), + } + } + + for scanner.Scan() { + line := scanner.Text() + lineStartsClientOutput := false + forceFlushFailedEvent := false + if data, ok := extractOpenAISSEDataLine(line); ok { + dataBytes := []byte(data) + trimmedData := strings.TrimSpace(data) + if needModelReplace && strings.Contains(data, mappedModel) { + line = s.replaceModelInSSELine(line, mappedModel, originalModel) + if replacedData, replaced := extractOpenAISSEDataLine(line); replaced { + dataBytes = []byte(replacedData) + trimmedData = strings.TrimSpace(replacedData) + } + } + if normalizedData, normalized := normalizeOpenAIResponsesFunctionCallArguments(dataBytes); normalized { + dataBytes = normalizedData + trimmedData = strings.TrimSpace(string(normalizedData)) + line = "data: " + string(normalizedData) + } + eventType := strings.TrimSpace(gjson.Get(trimmedData, "type").String()) + if eventType == "response.failed" { + failedMessage = extractOpenAISSEErrorMessage(dataBytes) + // response.failed 自带上游已消耗的 usage(input token 通常已扣);必须先解析 + // 再打 cyber 标记,否则 mark 记到的是解析前的 0,导致流式 cyber 按 0 token 计费 + // 而漏记真实用量。对齐 WS V2 / Chat 流式路径(均先解析 usage 再 Mark)。 + s.parseSSEUsageBytes(dataBytes, usage) + if hit, code, msg := detectOpenAICyberPolicy(dataBytes); hit { + MarkOpsCyberPolicy(c, CyberPolicyMark{ + Code: code, + Message: msg, + Body: truncateString(string(dataBytes), 4096), + UpstreamStatus: http.StatusOK, + UpstreamInTok: usage.InputTokens, + UpstreamOutTok: usage.OutputTokens, + }) + } else if !openAIStreamClientOutputStarted(c, clientOutputStarted) && openAIStreamFailedEventShouldFailover(dataBytes, failedMessage) { + return resultWithUsage(), + s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, dataBytes, failedMessage) + } + forceFlushFailedEvent = true + sawFailedEvent = true + } + if trimmedData == "[DONE]" { + sawDone = true + } + if openAIStreamEventIsTerminal(trimmedData) { + sawTerminalEvent = true + } + if responseID == "" { + responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes) + } + imageCounter.AddSSEData(dataBytes) + if sanitizedData, sanitized := sanitizeOpenAIResponseFailedEventForClient(dataBytes, eventType); sanitized { + dataBytes = sanitizedData + trimmedData = strings.TrimSpace(string(sanitizedData)) + line = "data: " + string(sanitizedData) + } + lineStartsClientOutput = forceFlushFailedEvent || openAIStreamDataStartsClientOutput(trimmedData, eventType) + if firstTokenMs == nil && lineStartsClientOutput && trimmedData != "[DONE]" { + ms := int(time.Since(startTime).Milliseconds()) + firstTokenMs = &ms + } + s.parseSSEUsageBytes(dataBytes, usage) + } + + if !clientDisconnected { + if !clientOutputStarted && !lineStartsClientOutput { + pendingLines = append(pendingLines, line) + continue + } + if !clientOutputStarted && len(pendingLines) > 0 { + if !writePendingLines() { + continue + } + } + if _, err := fmt.Fprintln(w, line); err != nil { + clientDisconnected = true + logger.LegacyPrintf("service.openai_gateway", "[OpenAI passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) + } else { + clientOutputStarted = true + flusher.Flush() + } + } + } + if err := scanner.Err(); err != nil { + if sawTerminalEvent && !sawFailedEvent { + return resultWithUsage(), nil + } + if sawFailedEvent { + return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage) + } + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return resultWithUsage(), fmt.Errorf("stream usage incomplete: %w", err) + } + if errors.Is(err, bufio.ErrTooLong) { + logger.LegacyPrintf("service.openai_gateway", "[OpenAI passthrough] SSE line too long: account=%d max_size=%d error=%v", account.ID, maxLineSize, err) + return resultWithUsage(), err + } + if !openAIStreamClientOutputStarted(c, clientOutputStarted) { + msg := "OpenAI stream disconnected before completion" + if errText := strings.TrimSpace(err.Error()); errText != "" { + msg += ": " + errText + } + return resultWithUsage(), + s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, nil, msg) + } + if clientDisconnected { + return resultWithUsage(), fmt.Errorf("stream usage incomplete after disconnect: %w", err) + } + logger.LegacyPrintf("service.openai_gateway", + "[OpenAI passthrough] 流读取异常中断: account=%d request_id=%s err=%v", + account.ID, + upstreamRequestID, + err, + ) + return resultWithUsage(), fmt.Errorf("stream read error: %w", err) + } + if sawFailedEvent { + return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage) + } + if !clientDisconnected && !sawDone && !sawTerminalEvent && ctx.Err() == nil { + logger.FromContext(ctx).With( + zap.String("component", "service.openai_gateway"), + zap.Int64("account_id", account.ID), + zap.String("upstream_request_id", upstreamRequestID), + ).Info("OpenAI passthrough 上游流在未收到 [DONE] 时结束,疑似断流") + if !openAIStreamClientOutputStarted(c, clientOutputStarted) { + return resultWithUsage(), + s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, nil, "OpenAI stream ended before a terminal event") + } + return resultWithUsage(), errors.New("stream usage incomplete: missing terminal event") + } + + return resultWithUsage(), nil +} + +func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( + ctx context.Context, + resp *http.Response, + c *gin.Context, + originalModel string, + mappedModel string, +) (*openaiNonStreamingResultPassthrough, error) { + body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) + if err != nil { + return nil, err + } + + // Detect SSE responses from upstream and convert to JSON. + // Some upstreams (e.g. other sub2api instances) may return SSE even when + // stream=false was requested. Without this conversion the client would + // receive raw SSE text or a terminal event with empty output. + if isEventStreamResponse(resp.Header) { + return s.handlePassthroughSSEToJSON(resp, c, body, originalModel, mappedModel) + } + + usage := &OpenAIUsage{} + usageParsed := false + if len(body) > 0 { + if parsedUsage, ok := extractOpenAIUsageFromJSONBytes(body); ok { + *usage = parsedUsage + usageParsed = true + } + } + if !usageParsed { + // 兜底:尝试从 SSE 文本中解析 usage + usage = s.parseSSEUsageFromBody(string(body)) + } + + writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + + contentType := resp.Header.Get("Content-Type") + if contentType == "" { + contentType = "application/json" + } + if originalModel != "" && mappedModel != "" && originalModel != mappedModel { + body = s.replaceModelInResponseBody(body, mappedModel, originalModel) + } + c.Data(resp.StatusCode, contentType, body) + return &openaiNonStreamingResultPassthrough{ + OpenAIUsage: usage, + usage: usage, + responseID: extractOpenAIResponseIDFromJSONBytes(body), + imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body), + imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body), + }, nil +} + +// handlePassthroughSSEToJSON converts an SSE response body into a JSON +// response for the passthrough path. It mirrors handleSSEToJSON while +// preserving passthrough payloads, except compact-only model remapping may +// rewrite model fields back to the original requested model. +func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c *gin.Context, body []byte, originalModel string, mappedModel string) (*openaiNonStreamingResultPassthrough, error) { + bodyText := string(body) + finalResponse, ok := extractCodexFinalResponse(bodyText) + + usage := &OpenAIUsage{} + if ok { + if parsedUsage, parsed := extractOpenAIUsageFromJSONBytes(finalResponse); parsed { + *usage = parsedUsage + } + // When the terminal event has an empty output array, reconstruct + // output from accumulated delta events so the client gets full content. + if len(gjson.GetBytes(finalResponse, "output").Array()) == 0 { + if outputJSON, reconstructed := reconstructResponseOutputFromSSE(bodyText); reconstructed { + if patched, err := sjson.SetRawBytes(finalResponse, "output", outputJSON); err == nil { + finalResponse = patched + } + } + } + body = finalResponse + if originalModel != "" && mappedModel != "" && originalModel != mappedModel { + body = s.replaceModelInResponseBody(body, mappedModel, originalModel) + } + // Correct tool calls in final response + body = s.correctToolCallsInResponseBody(body) + } else { + terminalType, terminalPayload, terminalOK := extractOpenAISSETerminalEvent(bodyText) + if terminalOK && terminalType == "response.failed" { + msg := extractOpenAISSEErrorMessage(terminalPayload) + if msg == "" { + msg = "Upstream compact response failed" + } + return nil, s.writeOpenAINonStreamingProtocolError(resp, c, msg) + } + usage = s.parseSSEUsageFromBody(bodyText) + if originalModel != "" && mappedModel != "" && originalModel != mappedModel { + bodyText = s.replaceModelInSSEBody(bodyText, mappedModel, originalModel) + } + body = []byte(bodyText) + } + + writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + + contentType := "application/json; charset=utf-8" + if !ok { + contentType = resp.Header.Get("Content-Type") + if contentType == "" { + contentType = "text/event-stream" + } + } + c.Data(resp.StatusCode, contentType, body) + + return &openaiNonStreamingResultPassthrough{ + OpenAIUsage: usage, + usage: usage, + responseID: extractOpenAIResponseIDFromJSONBytes(body), + imageCount: countOpenAIImageOutputsFromSSEBody(bodyText), + imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText), + }, nil +} + +func writeOpenAIPassthroughResponseHeaders(dst http.Header, src http.Header, filter *responseheaders.CompiledHeaderFilter) { + if dst == nil || src == nil { + return + } + if filter != nil { + responseheaders.WriteFilteredHeaders(dst, src, filter) + } else { + // 兜底:尽量保留最基础的 content-type + if v := strings.TrimSpace(src.Get("Content-Type")); v != "" { + dst.Set("Content-Type", v) + } + } + // 透传模式强制放行 x-codex-* 响应头(若上游返回)。 + // 注意:真实 http.Response.Header 的 key 一般会被 canonicalize;但为了兼容测试/自建响应, + // 这里用 EqualFold 做一次大小写不敏感的查找。 + getCaseInsensitiveValues := func(h http.Header, want string) []string { + if h == nil { + return nil + } + for k, vals := range h { + if strings.EqualFold(k, want) { + return vals + } + } + return nil + } + + for _, rawKey := range []string{ + "x-codex-primary-used-percent", + "x-codex-primary-reset-after-seconds", + "x-codex-primary-window-minutes", + "x-codex-secondary-used-percent", + "x-codex-secondary-reset-after-seconds", + "x-codex-secondary-window-minutes", + "x-codex-primary-over-secondary-limit-percent", + } { + vals := getCaseInsensitiveValues(src, rawKey) + if len(vals) == 0 { + continue + } + key := http.CanonicalHeaderKey(rawKey) + dst.Del(key) + for _, v := range vals { + dst.Add(key, v) + } + } +} diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go index c499bec778..16e552f21f 100644 --- a/backend/internal/service/openai_gateway_responses_chat_fallback.go +++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go @@ -1,13 +1,10 @@ package service import ( - "bufio" - "bytes" "context" "encoding/json" "errors" "fmt" - "io" "net/http" "strings" "time" @@ -31,22 +28,12 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( var responsesReq apicompat.ResponsesRequest if err := json.Unmarshal(body, &responsesReq); err != nil { - c.JSON(http.StatusBadRequest, gin.H{ - "error": gin.H{ - "type": "invalid_request_error", - "message": "Failed to parse request body", - }, - }) + writeOpenAIResponsesFallbackError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return nil, fmt.Errorf("parse responses request: %w", err) } originalModel := strings.TrimSpace(responsesReq.Model) if originalModel == "" { - c.JSON(http.StatusBadRequest, gin.H{ - "error": gin.H{ - "type": "invalid_request_error", - "message": "model is required", - }, - }) + writeOpenAIResponsesFallbackError(c, http.StatusBadRequest, "invalid_request_error", "model is required") return nil, fmt.Errorf("missing model in request") } @@ -56,12 +43,7 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( chatReq, err := apicompat.ResponsesToChatCompletionsRequest(&responsesReq) if err != nil { - c.JSON(http.StatusBadRequest, gin.H{ - "error": gin.H{ - "type": "invalid_request_error", - "message": err.Error(), - }, - }) + writeOpenAIResponsesFallbackError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) return nil, fmt.Errorf("convert responses to chat completions: %w", err) } @@ -98,94 +80,21 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions( zap.Bool("stream", clientStream), ) - apiKey := account.GetOpenAIApiKey() - if apiKey == "" { - return nil, fmt.Errorf("account %d missing api_key", account.ID) - } - baseURL := account.GetOpenAIBaseURL() - if baseURL == "" { - baseURL = "https://api.openai.com" - } - validatedURL, err := s.validateUpstreamBaseURL(baseURL) + // Build and send upstream request via the shared CC pipeline + apiKey, targetURL, err := s.resolveCCFallbackTarget(account) if err != nil { - return nil, fmt.Errorf("invalid base_url: %w", err) + return nil, err } - targetURL := buildOpenAIChatCompletionsURL(validatedURL) - - upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) - upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(chatBody)) - releaseUpstreamCtx() + resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent()) if err != nil { - return nil, fmt.Errorf("build upstream request: %w", err) - } - upstreamReq = upstreamReq.WithContext(WithHTTPUpstreamProfile(upstreamReq.Context(), HTTPUpstreamProfileOpenAI)) - upstreamReq.Header.Set("Content-Type", "application/json") - upstreamReq.Header.Set("Authorization", "Bearer "+apiKey) - if clientStream { - upstreamReq.Header.Set("Accept", "text/event-stream") - } else { - upstreamReq.Header.Set("Accept", "application/json") - } - for key, values := range c.Request.Header { - lowerKey := strings.ToLower(key) - if openaiCCRawAllowedHeaders[lowerKey] { - for _, v := range values { - upstreamReq.Header.Add(key, v) - } - } - } - if customUA := account.GetOpenAIUserAgent(); customUA != "" { - upstreamReq.Header.Set("user-agent", customUA) - } - - // 账号级请求头覆写(仅 openai api_key 账号启用时生效) - account.ApplyHeaderOverrides(upstreamReq.Header) - - proxyURL := "" - if account.Proxy != nil { - proxyURL = account.Proxy.URL() - } - resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) - if err != nil { - // Transport-level failure (proxy/DNS/TCP/TLS — no HTTP response). Convert to - // a failover so the handler switches to a healthy account, and temporarily - // unschedule the account on durable faults (e.g. rejected proxy credentials). - return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false) + return nil, err } defer func() { _ = resp.Body.Close() }() if resp.StatusCode >= 400 { - respBody := s.readUpstreamErrorBody(resp) - _ = resp.Body.Close() - resp.Body = io.NopCloser(bytes.NewReader(respBody)) - - upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) - upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) - if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { - upstreamDetail := "" - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes - if maxBytes <= 0 { - maxBytes = 2048 - } - upstreamDetail = truncateString(string(respBody), maxBytes) - } - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Kind: "failover", - Message: upstreamMsg, - Detail: upstreamDetail, - }) - s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) - return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, - RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)), - } + respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) + if foErr := s.failoverOpenAIUpstreamHTTPError(ctx, c, account, resp, respBody, upstreamMsg, upstreamModel); foErr != nil { + return nil, foErr } return s.handleErrorResponse(ctx, resp, c, account, chatBody, billingModel) } @@ -207,35 +116,11 @@ func (s *OpenAIGatewayService) bufferChatCompletionsAsResponses( startTime time.Time, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) + ccResp, usage, err := s.readCCUpstreamJSONResponse(c, resp, writeOpenAIResponsesFallbackError) if err != nil { - if !errors.Is(err, ErrUpstreamResponseBodyTooLarge) { - c.JSON(http.StatusBadGateway, gin.H{ - "error": gin.H{ - "type": "api_error", - "message": "Failed to read upstream response", - }, - }) - } - return nil, fmt.Errorf("read upstream body: %w", err) - } - - var ccResp apicompat.ChatCompletionsResponse - if err := json.Unmarshal(respBody, &ccResp); err != nil { - c.JSON(http.StatusBadGateway, gin.H{ - "error": gin.H{ - "type": "api_error", - "message": "Failed to parse upstream response", - }, - }) - return nil, fmt.Errorf("parse chat completions response: %w", err) - } - responsesResp := apicompat.ChatCompletionsResponseToResponses(&ccResp, originalModel) - - usage := OpenAIUsage{} - if parsed, ok := extractOpenAIUsageFromJSONBytes(respBody); ok { - usage = parsed + return nil, err } + responsesResp := apicompat.ChatCompletionsResponseToResponses(ccResp, originalModel) if s.responseHeaderFilter != nil { responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) @@ -266,27 +151,10 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses( startTime time.Time, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - headersWritten := false - writeStreamHeaders := func() { - if headersWritten { - return - } - headersWritten = true - if s.responseHeaderFilter != nil { - responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - } - c.Writer.Header().Set("Content-Type", "text/event-stream") - c.Writer.Header().Set("Cache-Control", "no-cache") - c.Writer.Header().Set("Connection", "keep-alive") - c.Writer.Header().Set("X-Accel-Buffering", "no") - c.Writer.WriteHeader(http.StatusOK) - } + writeStreamHeaders := s.newStreamHeaderWriter(c, resp.Header) state := apicompat.NewChatCompletionsToResponsesStreamState(originalModel) - var usage OpenAIUsage - var firstTokenMs *int clientDisconnected := false - sawDone := false writeEvents := func(events []apicompat.ResponsesStreamEvent) { if clientDisconnected || len(events) == 0 { @@ -314,57 +182,14 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses( c.Writer.Flush() } - scanner := bufio.NewScanner(resp.Body) - maxLineSize := defaultMaxLineSize - if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { - maxLineSize = s.cfg.Gateway.MaxLineSize - } - scanner.Buffer(make([]byte, 0, 64*1024), maxLineSize) + scan := s.scanCCStream(resp, "openai responses chat fallback", requestID, startTime, func(chunk *apicompat.ChatCompletionsChunk) { + writeEvents(apicompat.ChatCompletionsChunkToResponsesEvents(chunk, state)) + }) - for scanner.Scan() { - line := scanner.Text() - payload, ok := extractOpenAISSEDataLine(line) - if !ok { - continue - } - payload = strings.TrimSpace(payload) - if payload == "" { - continue - } - if payload == "[DONE]" { - sawDone = true - break - } - - if u := extractCCStreamUsage(payload); u != nil { - usage = *u - } - - var chunk apicompat.ChatCompletionsChunk - if err := json.Unmarshal([]byte(payload), &chunk); err != nil { - logger.L().Warn("openai responses chat fallback: failed to parse chat stream chunk", - zap.Error(err), - zap.String("request_id", requestID), - ) - continue - } - if firstTokenMs == nil && !isOpenAIChatUsageOnlyStreamChunk(payload) && chatChunkStartsResponsesOutput(&chunk) { - ms := int(time.Since(startTime).Milliseconds()) - firstTokenMs = &ms - } - writeEvents(apicompat.ChatCompletionsChunkToResponsesEvents(&chunk, state)) - } - - if err := scanner.Err(); err != nil { - if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { - logger.L().Warn("openai responses chat fallback: stream read error", - zap.Error(err), - zap.String("request_id", requestID), - ) - } + if scan.Err != nil { return &OpenAIForwardResult{ RequestID: requestID, - Usage: usage, + Usage: scan.Usage, Model: originalModel, BillingModel: billingModel, UpstreamModel: upstreamModel, @@ -372,8 +197,8 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses( ServiceTier: serviceTier, Stream: true, Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - }, fmt.Errorf("stream usage incomplete: %w", err) + FirstTokenMs: scan.FirstTokenMs, + }, fmt.Errorf("stream usage incomplete: %w", scan.Err) } writeEvents(apicompat.FinalizeChatCompletionsResponsesStream(state)) @@ -386,15 +211,13 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses( c.Writer.Flush() } } - if !sawDone { - logger.L().Debug("openai responses chat fallback: upstream stream ended without done sentinel", - zap.String("request_id", requestID), - ) + if !scan.SawDone { + logCCStreamMissingDoneSentinel("openai responses chat fallback", requestID) } return &OpenAIForwardResult{ RequestID: requestID, - Usage: usage, + Usage: scan.Usage, Model: originalModel, BillingModel: billingModel, UpstreamModel: upstreamModel, @@ -402,7 +225,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses( ServiceTier: serviceTier, Stream: true, Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, + FirstTokenMs: scan.FirstTokenMs, }, nil } diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go new file mode 100644 index 0000000000..6ac0de47b4 --- /dev/null +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -0,0 +1,1266 @@ +package service + +// 本文件由 openai_gateway_service.go 纯移动拆分而来:粘性会话哈希、账号选择与 +// 负载感知调度、配额自动暂停判定、并发槽位获取。仅做代码搬迁,无任何行为变更。 + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log/slog" + "sort" + "strconv" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" +) + +// ExtractSessionID extracts the raw session ID from headers or body without hashing. +// Used by ForwardAsAnthropic to pass as prompt_cache_key for upstream cache. +func (s *OpenAIGatewayService) ExtractSessionID(c *gin.Context, body []byte) string { + if c == nil { + return "" + } + sessionID := strings.TrimSpace(c.GetHeader("session_id")) + if sessionID == "" { + sessionID = strings.TrimSpace(c.GetHeader("conversation_id")) + } + if sessionID == "" && len(body) > 0 { + sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) + } + return sessionID +} + +func explicitOpenAISessionID(c *gin.Context, body []byte) string { + if c == nil { + return "" + } + + sessionID := strings.TrimSpace(c.GetHeader("session_id")) + if sessionID == "" { + sessionID = strings.TrimSpace(c.GetHeader("conversation_id")) + } + if sessionID == "" && len(body) > 0 { + sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) + } + return sessionID +} + +// GenerateExplicitSessionHash generates a sticky-session hash only from explicit +// client session signals. It intentionally skips content-derived fallback and is +// used by stateless endpoints such as /v1/images. +func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body []byte) string { + sessionID := explicitOpenAISessionID(c, body) + if sessionID == "" { + return "" + } + + currentHash, legacyHash := deriveOpenAISessionHashes(sessionID) + attachOpenAILegacySessionHashToGin(c, legacyHash) + return currentHash +} + +// GenerateSessionHash generates a sticky-session hash for OpenAI requests. +// +// Priority: +// 1. Header: session_id +// 2. Header: conversation_id +// 3. Body: prompt_cache_key (opencode) +// 4. Body: content-based fallback (model + system + tools + first user message) +func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) string { + if c == nil { + return "" + } + + sessionID := explicitOpenAISessionID(c, body) + if sessionID == "" && len(body) > 0 { + sessionID = deriveOpenAIContentSessionSeed(body) + } + if sessionID == "" { + return "" + } + + currentHash, legacyHash := deriveOpenAISessionHashes(sessionID) + attachOpenAILegacySessionHashToGin(c, legacyHash) + return currentHash +} + +// GenerateSessionHashWithFallback 先按常规信号生成会话哈希; +// 当未携带 session_id/conversation_id/prompt_cache_key 时,使用 fallbackSeed 生成稳定哈希。 +// 该方法用于 WS ingress,避免会话信号缺失时发生跨账号漂移。 +func (s *OpenAIGatewayService) GenerateSessionHashWithFallback(c *gin.Context, body []byte, fallbackSeed string) string { + sessionHash := s.GenerateSessionHash(c, body) + if sessionHash != "" { + return sessionHash + } + + seed := strings.TrimSpace(fallbackSeed) + if seed == "" { + return "" + } + + currentHash, legacyHash := deriveOpenAISessionHashes(seed) + attachOpenAILegacySessionHashToGin(c, legacyHash) + return currentHash +} + +func resolveOpenAIUpstreamOriginator(c *gin.Context, isOfficialClient bool) string { + if c != nil { + if originator := strings.TrimSpace(c.GetHeader("originator")); originator != "" { + return originator + } + } + if isOfficialClient { + return "codex_cli_rs" + } + return "opencode" +} + +// BindStickySession sets session -> account binding with standard TTL. +func (s *OpenAIGatewayService) BindStickySession(ctx context.Context, groupID *int64, sessionHash string, accountID int64) error { + if sessionHash == "" || accountID <= 0 { + return nil + } + ttl := openaiStickySessionTTL + if s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.StickySessionTTLSeconds > 0 { + ttl = time.Duration(s.cfg.Gateway.OpenAIWS.StickySessionTTLSeconds) * time.Second + } + return s.setStickySessionAccountID(ctx, groupID, sessionHash, accountID, ttl) +} + +// SelectAccount selects an OpenAI account with sticky session support +func (s *OpenAIGatewayService) SelectAccount(ctx context.Context, groupID *int64, sessionHash string) (*Account, error) { + return s.SelectAccountForModel(ctx, groupID, sessionHash, "") +} + +// SelectAccountForModel selects an account supporting the requested model +func (s *OpenAIGatewayService) SelectAccountForModel(ctx context.Context, groupID *int64, sessionHash string, requestedModel string) (*Account, error) { + return s.SelectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, nil) +} + +// SelectAccountForModelWithExclusions selects an account supporting the requested model while excluding specified accounts. +// SelectAccountForModelWithExclusions 选择支持指定模型的账号,同时排除指定的账号。 +func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*Account, error) { + return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, 0, "") +} + +// noAvailableOpenAISelectionError builds the standard "no account available" error +// while preserving the compact-specific error when applicable. +func normalizeOpenAICompatiblePlatform(platform string) string { + if platform == PlatformGrok { + return PlatformGrok + } + return PlatformOpenAI +} + +func noAvailableOpenAISelectionError(requestedModel string, compactBlocked bool) error { + if compactBlocked { + return ErrNoAvailableCompactAccounts + } + if requestedModel != "" { + return fmt.Errorf("no available OpenAI accounts supporting model: %s", requestedModel) + } + return errors.New("no available OpenAI accounts") +} + +// openAICompactSupportTier classifies an OpenAI account by compact capability. +// 0 = explicitly unsupported, 1 = unknown / not yet probed, 2 = explicitly supported. +func openAICompactSupportTier(account *Account) int { + if account == nil || !account.IsOpenAI() { + return 0 + } + supported, known := account.OpenAICompactSupportKnown() + if !known { + return 1 + } + if supported { + return 2 + } + return 0 +} + +// isOpenAICompatibleAccountEligibleForRequest 判断 OpenAI 兼容账号是否满足本次请求的调度条件。 +// 检查内容包括:平台匹配、账号可用性、quota 自动暂停、spark 路由限制、模型支持及端点能力。 +// +// 注意:对 spark 影子账号,调用方还须额外调用 parentHealthyForShadow(account, lookup) +// 检查母账号凭据可用性;该检查未内置于本函数,以避免注入 DB 依赖。 +func isOpenAICompatibleAccountEligibleForRequest(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) bool { + platform = normalizeOpenAICompatiblePlatform(platform) + if account == nil || account.Platform != platform || !account.IsOpenAICompatible() || !account.IsSchedulableForModelWithContext(ctx, requestedModel) { + return false + } + if account.IsOpenAI() { + if paused, reason := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { + // Debug level: this fires per-candidate on the scheduling hot path, so Info + // would amplify into log spam once several accounts cross the threshold. + slog.Debug("account_auto_paused_by_quota", + "account_id", account.ID, + "window", reason.window, + "threshold", reason.threshold, + "utilization", reason.utilization, + ) + return false + } + } + if account.IsGrok() { + if paused, reason := shouldAutoPauseGrokAccountByQuota(account); paused { + slog.Debug("grok_account_auto_paused_by_quota", + "account_id", account.ID, + "window", reason.window, + "threshold", reason.threshold, + "utilization", reason.utilization, + ) + return false + } + } + if requestedModel != "" && !account.IsModelSupported(requestedModel) { + return false + } + if !account.SupportsOpenAIEndpointCapability(requiredCapability) { + return false + } + if requireCompact && (!account.IsOpenAI() || openAICompactSupportTier(account) == 0) { + return false + } + return true +} + +type openAIQuotaAutoPauseDecision struct { + window string + threshold float64 + utilization float64 +} + +func shouldAutoPauseGrokAccountByQuota(account *Account) (bool, openAIQuotaAutoPauseDecision) { + if account == nil || !account.IsGrok() || account.Type != AccountTypeOAuth { + return false, openAIQuotaAutoPauseDecision{} + } + snapshot, err := grokQuotaSnapshotFromExtra(account.Extra) + if err != nil || snapshot == nil { + return false, openAIQuotaAutoPauseDecision{} + } + now := time.Now() + if grokQuotaSnapshotStaleForPause(snapshot, now) { + return false, openAIQuotaAutoPauseDecision{} + } + if grokQuotaRetryAfterActive(snapshot, now) { + return true, openAIQuotaAutoPauseDecision{window: "retry_after", threshold: 1, utilization: 1} + } + if paused, decision := shouldAutoPauseGrokQuotaWindow("requests", snapshot.Requests, now); paused { + return true, decision + } + if paused, decision := shouldAutoPauseGrokQuotaWindow("tokens", snapshot.Tokens, now); paused { + return true, decision + } + return false, openAIQuotaAutoPauseDecision{} +} + +func grokQuotaRetryAfterActive(snapshot *xai.QuotaSnapshot, now time.Time) bool { + if snapshot == nil || snapshot.RetryAfterSeconds == nil || *snapshot.RetryAfterSeconds <= 0 { + return false + } + if strings.TrimSpace(snapshot.UpdatedAt) == "" { + return true + } + updatedAt, err := parseTime(snapshot.UpdatedAt) + if err != nil { + return true + } + retryAfterUntil := updatedAt.Add(time.Duration(*snapshot.RetryAfterSeconds) * time.Second) + return now.Before(retryAfterUntil) +} + +func shouldAutoPauseGrokQuotaWindow(name string, window *xai.QuotaWindow, now time.Time) (bool, openAIQuotaAutoPauseDecision) { + if window == nil || window.Limit == nil || window.Remaining == nil || *window.Limit <= 0 { + return false, openAIQuotaAutoPauseDecision{} + } + if window.ResetUnix != nil && *window.ResetUnix > 0 && !now.Before(time.Unix(*window.ResetUnix, 0)) { + return false, openAIQuotaAutoPauseDecision{} + } + utilization := float64(*window.Limit-*window.Remaining) / float64(*window.Limit) + if *window.Remaining <= 0 || utilization >= 1 { + return true, openAIQuotaAutoPauseDecision{window: name, threshold: 1, utilization: utilization} + } + return false, openAIQuotaAutoPauseDecision{} +} + +func grokQuotaSnapshotStaleForPause(snapshot *xai.QuotaSnapshot, now time.Time) bool { + if snapshot == nil || strings.TrimSpace(snapshot.UpdatedAt) == "" { + return false + } + updatedAt, err := parseTime(snapshot.UpdatedAt) + if err != nil { + return false + } + return now.Sub(updatedAt) >= openAICodexAutoPauseStaleAfter +} + +func shouldAutoPauseOpenAIAccountByQuota(ctx context.Context, account *Account) (bool, openAIQuotaAutoPauseDecision) { + if account == nil || !account.IsOpenAI() { + return false, openAIQuotaAutoPauseDecision{} + } + // Per-account explicit-disable flags must take precedence over the global default. + // Without these, leaving the account threshold blank means "use global default", + // so an admin has no way to exempt a single account from auto-pause once a global + // default exists. The disable flag is per-window so an account can opt out of + // only 5h or only 7d auto-pause. + disabled5h := resolveAccountExtraBool(account.Extra, "auto_pause_5h_disabled") + disabled7d := resolveAccountExtraBool(account.Extra, "auto_pause_7d_disabled") + threshold5h, threshold7d := resolveOpenAIQuotaAutoPauseThresholds(ctx, account) + now := time.Now() + if !disabled5h && threshold5h > 0 { + if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "5h", now); ok && utilization >= threshold5h { + return true, openAIQuotaAutoPauseDecision{window: "5h", threshold: threshold5h, utilization: utilization} + } + } + if !disabled7d && threshold7d > 0 { + if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "7d", now); ok && utilization >= threshold7d { + return true, openAIQuotaAutoPauseDecision{window: "7d", threshold: threshold7d, utilization: utilization} + } + } + return false, openAIQuotaAutoPauseDecision{} +} + +// resolveAccountExtraBool reads a bool-like value from account extra, tolerating +// the few shapes JSON unmarshalling may produce (real bool, "true"/"false" +// strings, 0/1 numbers). +func resolveAccountExtraBool(extra map[string]any, key string) bool { + if len(extra) == 0 { + return false + } + value, ok := extra[key] + if !ok || value == nil { + return false + } + switch v := value.(type) { + case bool: + return v + case string: + parsed, err := strconv.ParseBool(strings.TrimSpace(v)) + return err == nil && parsed + case float64: + return v != 0 + case float32: + return v != 0 + case int: + return v != 0 + case int64: + return v != 0 + case json.Number: + if i, err := v.Int64(); err == nil { + return i != 0 + } + } + return false +} + +func resolveOpenAIQuotaAutoPauseThresholds(ctx context.Context, account *Account) (float64, float64) { + threshold5h, _ := resolveAccountExtraNumber(account.Extra, "auto_pause_5h_threshold") + threshold7d, _ := resolveAccountExtraNumber(account.Extra, "auto_pause_7d_threshold") + threshold5h = clamp01(threshold5h) + threshold7d = clamp01(threshold7d) + if threshold5h > 0 && threshold7d > 0 { + return threshold5h, threshold7d + } + settings := openAIQuotaAutoPauseSettingsFromContext(ctx) + if threshold5h <= 0 { + threshold5h = clamp01(settings.DefaultThreshold5h) + } + if threshold7d <= 0 { + threshold7d = clamp01(settings.DefaultThreshold7d) + } + return threshold5h, threshold7d +} + +func resolveAccountExtraNumber(extra map[string]any, keys ...string) (float64, bool) { + if len(extra) == 0 { + return 0, false + } + for _, key := range keys { + value, ok := extra[key] + if !ok || value == nil { + continue + } + switch v := value.(type) { + case float64: + return v, true + case float32: + return float64(v), true + case int: + return float64(v), true + case int64: + return float64(v), true + case json.Number: + parsed, err := v.Float64() + if err == nil { + return parsed, true + } + case string: + parsed, err := strconv.ParseFloat(strings.TrimSpace(v), 64) + if err == nil { + return parsed, true + } + } + } + return 0, false +} + +// resolveOpenAIQuotaUtilization returns the current utilization ratio (0..1) for the +// given Codex usage window. ok=false means there is no usable signal to pause on: +// either no snapshot exists, or the window has already rolled over so the cached +// percentage is stale. The stale guard matters because a paused account stops +// receiving requests, so its snapshot is never refreshed from upstream headers — +// without this check an old used_percent would keep the account paused forever even +// after the real window reset. +func resolveOpenAIQuotaUtilization(extra map[string]any, window string, now time.Time) (float64, bool) { + usedPercent := readOpenAIQuotaUsedPercent(extra, window) + if usedPercent <= 0 { + return 0, false + } + if openAIQuotaWindowReset(extra, window, now) { + return 0, false + } + // 快照过于陈旧(账号长期未收到流量刷新)时,不再据此暂停。放行后下一次响应头 + // 会刷新快照实现自愈,避免账号在错误/过期的 used% 上被永久跳过(issue #2994)。 + if openAICodexSnapshotStaleForPause(extra, now) { + return 0, false + } + return usedPercent / 100, true +} + +// openAICodexSnapshotStaleForPause reports whether the Codex usage snapshot is stale +// enough that it should no longer keep an account auto-paused. It anchors on +// codex_usage_updated_at (always written by buildCodexUsageExtraUpdates). A missing or +// unparseable timestamp returns false (treated as fresh, so the account stays paused) — +// this is deliberate: it prevents any snapshot without a write time from silently escaping +// auto-pause, and a genuinely-exhausted account that is actively served refreshes the +// timestamp on every response so it never crosses the staleness bound. +func openAICodexSnapshotStaleForPause(extra map[string]any, now time.Time) bool { + if len(extra) == 0 { + return false + } + updatedRaw, ok := extra["codex_usage_updated_at"] + if !ok { + return false + } + updatedAt, err := parseTime(fmt.Sprint(updatedRaw)) + if err != nil { + return false + } + return now.Sub(updatedAt) >= openAICodexAutoPauseStaleAfter +} + +// openAIQuotaWindowReset reports whether the Codex usage window's reset time has +// already passed relative to now. It prefers the absolute codex__reset_at +// timestamp and falls back to codex__reset_after_seconds anchored at +// codex_usage_updated_at, mirroring AccountUsageService's window-progress logic. +func openAIQuotaWindowReset(extra map[string]any, window string, now time.Time) bool { + if len(extra) == 0 { + return false + } + if resetAtRaw, ok := extra["codex_"+window+"_reset_at"]; ok { + if resetAt, err := parseTime(fmt.Sprint(resetAtRaw)); err == nil { + return !now.Before(resetAt) + } + } + resetAfter := parseExtraInt(extra["codex_"+window+"_reset_after_seconds"]) + if resetAfter <= 0 { + return false + } + base := now + if updatedRaw, ok := extra["codex_usage_updated_at"]; ok { + if updatedAt, err := parseTime(fmt.Sprint(updatedRaw)); err == nil { + base = updatedAt + } + } + resetAt := base.Add(time.Duration(resetAfter) * time.Second) + return !now.Before(resetAt) +} + +func readOpenAIQuotaUsedPercent(extra map[string]any, window string) float64 { + if len(extra) == 0 { + return 0 + } + if value, ok := resolveAccountExtraNumber(extra, "codex_"+window+"_used_percent"); ok { + return value + } + return 0 +} + +type openAIQuotaAutoPauseCtxKey struct{} + +func withOpenAIQuotaAutoPauseSettings(ctx context.Context, settings OpsOpenAIAccountQuotaAutoPauseSettings) context.Context { + if ctx == nil { + ctx = context.Background() + } + return context.WithValue(ctx, openAIQuotaAutoPauseCtxKey{}, settings) +} + +func openAIQuotaAutoPauseSettingsFromContext(ctx context.Context) OpsOpenAIAccountQuotaAutoPauseSettings { + if ctx == nil { + return OpsOpenAIAccountQuotaAutoPauseSettings{} + } + settings, _ := ctx.Value(openAIQuotaAutoPauseCtxKey{}).(OpsOpenAIAccountQuotaAutoPauseSettings) + return settings +} + +func (s *OpenAIGatewayService) withOpenAIQuotaAutoPauseContext(ctx context.Context) context.Context { + if s == nil || s.settingService == nil { + return ctx + } + return withOpenAIQuotaAutoPauseSettings(ctx, s.settingService.GetOpenAIQuotaAutoPauseSettings(ctx)) +} + +// prioritizeOpenAICompactAccounts re-orders a slice so that accounts with known +// compact support are tried first, followed by unknown, then explicitly unsupported. +// The relative order within each tier is preserved. +func prioritizeOpenAICompactAccounts(accounts []*Account) []*Account { + if len(accounts) == 0 { + return nil + } + supported := make([]*Account, 0, len(accounts)) + unknown := make([]*Account, 0, len(accounts)) + unsupported := make([]*Account, 0, len(accounts)) + for _, account := range accounts { + switch openAICompactSupportTier(account) { + case 2: + supported = append(supported, account) + case 1: + unknown = append(unknown, account) + default: + unsupported = append(unsupported, account) + } + } + out := make([]*Account, 0, len(accounts)) + out = append(out, supported...) + out = append(out, unknown...) + out = append(out, unsupported...) + return out +} + +// resolveOpenAIAccountUpstreamModelForRequest resolves the upstream model that +// would be sent for a given request, honouring compact-only mappings when the +// caller is on the /responses/compact path. +func resolveOpenAIAccountUpstreamModelForRequest(account *Account, requestedModel string, requireCompact bool) string { + upstreamModel := resolveOpenAIForwardModel(account, requestedModel, "") + if upstreamModel == "" { + return "" + } + if requireCompact { + return resolveOpenAICompactForwardModel(account, upstreamModel) + } + return upstreamModel +} + +func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) (*Account, error) { + platform = normalizeOpenAICompatiblePlatform(platform) + if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { + slog.Warn("channel pricing restriction blocked request", + "group_id", derefGroupID(groupID), + "model", requestedModel) + return nil, fmt.Errorf("%w supporting model: %s (channel pricing restriction)", ErrNoAvailableAccounts, requestedModel) + } + + // 1. 尝试粘性会话命中 + // Try sticky session hit + if account := s.tryStickySessionHit(ctx, groupID, platform, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability); account != nil { + return account, nil + } + + // 2. 获取可调度的 OpenAI 账号 + // Get schedulable OpenAI accounts + accounts, err := s.listSchedulableAccounts(ctx, groupID, platform) + if err != nil { + return nil, fmt.Errorf("query accounts failed: %w", err) + } + + // 3. 按优先级 + LRU 选择最佳账号 + // Select by priority + LRU + selected, compactBlocked := s.selectBestAccount(ctx, groupID, platform, accounts, requestedModel, excludedIDs, requireCompact, requiredCapability) + + if selected == nil { + return nil, noAvailableOpenAISelectionError(requestedModel, compactBlocked) + } + + hydrated, err := s.hydrateSelectedAccount(ctx, selected) + if err != nil { + return nil, err + } + + // 4. 设置粘性会话绑定 + // Set sticky session binding + if sessionHash != "" { + _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, selected.ID, openaiStickySessionTTL) + } + + return hydrated, nil +} + +// tryStickySessionHit 尝试从粘性会话获取账号。 +// 如果命中且账号可用则返回账号;如果账号不可用则清理会话并返回 nil。 +// +// tryStickySessionHit attempts to get account from sticky session. +// Returns account if hit and usable; clears session and returns nil if account is unavailable. +func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID *int64, platform string, sessionHash, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) *Account { + if sessionHash == "" { + return nil + } + platform = normalizeOpenAICompatiblePlatform(platform) + + accountID := stickyAccountID + if accountID <= 0 { + var err error + accountID, err = s.getStickySessionAccountID(ctx, groupID, sessionHash) + if err != nil || accountID <= 0 { + return nil + } + } + + if _, excluded := excludedIDs[accountID]; excluded { + return nil + } + + account, err := s.getSchedulableAccount(ctx, accountID) + if err != nil { + return nil + } + + // 检查账号是否需要清理粘性会话 + // Check if sticky session should be cleared + if shouldClearStickySession(account, requestedModel) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + return nil + } + + // 验证账号是否可用于当前请求 + // Verify account is usable for current request + if !isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, false, requiredCapability) { + return nil + } + if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + return nil + } + if s.isOpenAIAccountRuntimeBlocked(account) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + return nil + } + account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, platform, requestedModel, requireCompact, requiredCapability) + if account == nil || !openAIStickyAccountMatchesGroup(account, groupID) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + return nil + } + if groupID != nil && s.needsUpstreamChannelRestrictionCheck(ctx, groupID) && + s.isUpstreamModelRestrictedByChannel(ctx, *groupID, account, requestedModel, requireCompact) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + return nil + } + + // 刷新会话 TTL 并返回账号 + // Refresh session TTL and return account + _ = s.refreshStickySessionTTL(ctx, groupID, sessionHash, openaiStickySessionTTL) + return account +} + +// selectBestAccount 从候选账号中选择最佳账号(优先级 + LRU)。 +// 返回 nil 表示无可用账号。 +// +// selectBestAccount selects the best account from candidates (priority + LRU). +// Returns nil if no available account. The second return reports whether at +// least one candidate was filtered out solely because it lacks compact support +// (only meaningful when requireCompact=true). +func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, platform string, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*Account, bool) { + platform = normalizeOpenAICompatiblePlatform(platform) + var selected *Account + selectedCompactTier := -1 + compactBlocked := false + needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) + + for i := range accounts { + acc := &accounts[i] + + // 跳过被排除的账号 + // Skip excluded accounts + if _, excluded := excludedIDs[acc.ID]; excluded { + continue + } + + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) + if fresh == nil { + continue + } + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, false, requiredCapability) + if fresh == nil { + continue + } + if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel, requireCompact) { + continue + } + compactTier := 0 + if requireCompact { + compactTier = openAICompactSupportTier(fresh) + if compactTier == 0 { + compactBlocked = true + continue + } + } + + // 选择优先级最高且最久未使用的账号 + // Select highest priority and least recently used + if selected == nil { + selected = fresh + selectedCompactTier = compactTier + continue + } + + // compact 模式下高 tier 优先;同 tier 内才比较 priority/LRU。 + if requireCompact && compactTier != selectedCompactTier { + if compactTier > selectedCompactTier { + selected = fresh + selectedCompactTier = compactTier + } + continue + } + + if s.isBetterAccount(fresh, selected) { + selected = fresh + selectedCompactTier = compactTier + } + } + + return selected, compactBlocked +} + +// isBetterAccount 判断 candidate 是否比 current 更优。 +// 规则:优先级更高(数值更小)优先;同优先级时,未使用过的优先,其次是最久未使用的。 +// +// isBetterAccount checks if candidate is better than current. +// Rules: higher priority (lower value) wins; same priority: never used > least recently used. +func (s *OpenAIGatewayService) isBetterAccount(candidate, current *Account) bool { + // 优先级更高(数值更小) + // Higher priority (lower value) + if candidate.Priority < current.Priority { + return true + } + if candidate.Priority > current.Priority { + return false + } + + // 同优先级,比较最后使用时间 + // Same priority, compare last used time + switch { + case candidate.LastUsedAt == nil && current.LastUsedAt != nil: + // candidate 从未使用,优先 + return true + case candidate.LastUsedAt != nil && current.LastUsedAt == nil: + // current 从未使用,保持 + return false + case candidate.LastUsedAt == nil && current.LastUsedAt == nil: + // 都未使用,保持 + return false + default: + // 都使用过,选择最久未使用的 + return candidate.LastUsedAt.Before(*current.LastUsedAt) + } +} + +// SelectAccountWithLoadAwareness selects an account with load-awareness and wait plan. +func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*AccountSelectionResult, error) { + return s.selectAccountWithLoadAwareness(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, "") +} + +func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*AccountSelectionResult, error) { + platform = normalizeOpenAICompatiblePlatform(platform) + if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { + slog.Warn("channel pricing restriction blocked request", + "group_id", derefGroupID(groupID), + "model", requestedModel) + return nil, fmt.Errorf("%w supporting model: %s (channel pricing restriction)", ErrNoAvailableAccounts, requestedModel) + } + + cfg := s.schedulingConfig() + needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) + var stickyAccountID int64 + if sessionHash != "" && s.cache != nil { + if accountID, err := s.getStickySessionAccountID(ctx, groupID, sessionHash); err == nil { + stickyAccountID = accountID + } + } + if s.concurrencyService == nil || !cfg.LoadBatchEnabled { + account, err := s.selectAccountForModelWithExclusions(ctx, groupID, platform, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability) + if err != nil { + return nil, err + } + result, err := s.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) + if err == nil && result != nil && result.Acquired { + return s.newAcquiredSelectionResult(ctx, account, result.ReleaseFunc) + } + if stickyAccountID > 0 && stickyAccountID == account.ID && s.concurrencyService != nil { + waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, account.ID) + if waitingCount < cfg.StickySessionMaxWaiting { + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) + } + } + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.FallbackWaitTimeout, + MaxWaiting: cfg.FallbackMaxWaiting, + }) + } + + accounts, err := s.listSchedulableAccounts(ctx, groupID, platform) + if err != nil { + return nil, err + } + if len(accounts) == 0 { + return nil, ErrNoAvailableAccounts + } + + isExcluded := func(accountID int64) bool { + if excludedIDs == nil { + return false + } + _, excluded := excludedIDs[accountID] + return excluded + } + + // ============ Layer 1: Sticky session ============ + if sessionHash != "" { + accountID := stickyAccountID + if accountID > 0 && !isExcluded(accountID) { + account, err := s.getSchedulableAccount(ctx, accountID) + if err == nil { + clearSticky := shouldClearStickySession(account, requestedModel) + if clearSticky { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + } + if !clearSticky && isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, false, requiredCapability) { + account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, platform, requestedModel, requireCompact, requiredCapability) + if account == nil { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + } else if !openAIStickyAccountMatchesGroup(account, groupID) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + } else if s.isOpenAIAccountRuntimeBlocked(account) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + } else if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, account, requestedModel, requireCompact) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + } else if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { + _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) + } else { + result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency) + if err == nil && result != nil && result.Acquired { + selection, selectErr := s.newAcquiredSelectionResult(ctx, account, result.ReleaseFunc) + if selectErr != nil { + return nil, selectErr + } + _ = s.refreshStickySessionTTL(ctx, groupID, sessionHash, openaiStickySessionTTL) + return selection, nil + } + + waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, accountID) + if waitingCount < cfg.StickySessionMaxWaiting { + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: accountID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) + } + } + } + } + } + } + + // ============ Layer 2: Load-aware selection ============ + // Per-pass parent-health cache to avoid repeated DB calls when multiple shadow + // accounts share the same parent. + parentCacheL2 := make(map[int64]*Account) + parentLookupL2 := func(id int64) *Account { + if a, ok := parentCacheL2[id]; ok { + return a + } + if s.accountRepo == nil { + return nil + } + a, _ := s.accountRepo.GetByID(ctx, id) + parentCacheL2[id] = a + return a + } + baseCandidateCount := 0 + candidates := make([]*Account, 0, len(accounts)) + for i := range accounts { + acc := &accounts[i] + if isExcluded(acc.ID) { + continue + } + // Scheduler snapshots can be temporarily stale (bucket rebuild is throttled); + // re-check schedulability here so recently rate-limited/overloaded accounts + // are not selected again before the bucket is rebuilt. + if !isOpenAICompatibleAccountEligibleForRequest(ctx, acc, platform, requestedModel, false, requiredCapability) { + continue + } + if !parentHealthyForShadow(acc, parentLookupL2) { + continue + } + if s.isOpenAIAccountRuntimeBlocked(acc) { + continue + } + if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel, requireCompact) { + continue + } + baseCandidateCount++ + candidates = append(candidates, acc) + } + + if len(candidates) == 0 { + return nil, ErrNoAvailableAccounts + } + + accountLoads := make([]AccountWithConcurrency, 0, len(candidates)) + for _, acc := range candidates { + accountLoads = append(accountLoads, AccountWithConcurrency{ + ID: acc.ID, + MaxConcurrency: acc.EffectiveLoadFactor(), + }) + } + + tryAcquireFromLoadMap := func(loadMap map[int64]*AccountLoadInfo) (*AccountSelectionResult, bool, error) { + var available []accountWithLoad + for _, acc := range candidates { + loadInfo := loadMap[acc.ID] + if loadInfo == nil { + loadInfo = &AccountLoadInfo{AccountID: acc.ID} + } + if loadInfo.LoadRate < 100 { + available = append(available, accountWithLoad{ + account: acc, + loadInfo: loadInfo, + }) + } + } + + if len(available) == 0 { + return nil, false, nil + } + + sort.SliceStable(available, func(i, j int) bool { + a, b := available[i], available[j] + if a.account.Priority != b.account.Priority { + return a.account.Priority < b.account.Priority + } + if a.loadInfo.LoadRate != b.loadInfo.LoadRate { + return a.loadInfo.LoadRate < b.loadInfo.LoadRate + } + switch { + case a.account.LastUsedAt == nil && b.account.LastUsedAt != nil: + return true + case a.account.LastUsedAt != nil && b.account.LastUsedAt == nil: + return false + case a.account.LastUsedAt == nil && b.account.LastUsedAt == nil: + return false + default: + return a.account.LastUsedAt.Before(*b.account.LastUsedAt) + } + }) + shuffleWithinSortGroups(available) + + selectionOrder := make([]accountWithLoad, 0, len(available)) + if requireCompact { + appendTier := func(out []accountWithLoad, tier int) []accountWithLoad { + for _, item := range available { + if openAICompactSupportTier(item.account) == tier { + out = append(out, item) + } + } + return out + } + selectionOrder = appendTier(selectionOrder, 2) + selectionOrder = appendTier(selectionOrder, 1) + // tier 0 候选作为兜底追加:DB recheck 时若发现 cache tier 0 实际 + // 已升级为 1/2(探测刚跑完,cache 尚未刷新),仍可正常命中。 + selectionOrder = appendTier(selectionOrder, 0) + } else { + selectionOrder = append(selectionOrder, available...) + } + + for _, item := range selectionOrder { + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, item.account, platform, requestedModel, false, requiredCapability) + if fresh == nil { + continue + } + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) + if fresh == nil { + continue + } + if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel, requireCompact) { + continue + } + result, err := s.tryAcquireAccountSlot(ctx, fresh.ID, fresh.Concurrency) + if err == nil && result != nil && result.Acquired { + selection, selectErr := s.newAcquiredSelectionResult(ctx, fresh, result.ReleaseFunc) + if selectErr != nil { + return nil, true, selectErr + } + if sessionHash != "" { + _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, fresh.ID, openaiStickySessionTTL) + } + return selection, true, nil + } + } + return nil, true, nil + } + + loadMap, err := s.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) + if err != nil { + ordered := append([]*Account(nil), candidates...) + sortAccountsByPriorityAndLastUsed(ordered, false) + if requireCompact { + ordered = prioritizeOpenAICompactAccounts(ordered) + } + for _, acc := range ordered { + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) + if fresh == nil { + continue + } + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) + if fresh == nil { + continue + } + if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel, requireCompact) { + continue + } + result, err := s.tryAcquireAccountSlot(ctx, fresh.ID, fresh.Concurrency) + if err == nil && result != nil && result.Acquired { + selection, selectErr := s.newAcquiredSelectionResult(ctx, fresh, result.ReleaseFunc) + if selectErr != nil { + return nil, selectErr + } + if sessionHash != "" { + _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, fresh.ID, openaiStickySessionTTL) + } + return selection, nil + } + } + } else { + if selection, attempted, selectErr := tryAcquireFromLoadMap(loadMap); selectErr != nil { + return nil, selectErr + } else if selection != nil { + return selection, nil + } else if attempted { + if freshLoadMap, loadErr := s.concurrencyService.GetAccountsLoadBatchFresh(ctx, accountLoads); loadErr == nil { + if selection, _, selectErr := tryAcquireFromLoadMap(freshLoadMap); selectErr != nil { + return nil, selectErr + } else if selection != nil { + return selection, nil + } + } + } + } + + // ============ Layer 3: Fallback wait ============ + sortAccountsByPriorityAndLastUsed(candidates, false) + if requireCompact { + candidates = prioritizeOpenAICompactAccounts(candidates) + } + for _, acc := range candidates { + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) + if fresh == nil { + continue + } + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) + if fresh == nil { + continue + } + if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel, requireCompact) { + continue + } + return s.newSelectionResult(ctx, fresh, false, nil, &AccountWaitPlan{ + AccountID: fresh.ID, + MaxConcurrency: fresh.Concurrency, + Timeout: cfg.FallbackWaitTimeout, + MaxWaiting: cfg.FallbackMaxWaiting, + }) + } + + if requireCompact && baseCandidateCount > 0 { + return nil, ErrNoAvailableCompactAccounts + } + return nil, ErrNoAvailableAccounts +} + +func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, groupID *int64, platform string) ([]Account, error) { + platform = normalizeOpenAICompatiblePlatform(platform) + if s.schedulerSnapshot != nil { + accounts, _, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, false) + return accounts, err + } + var accounts []Account + var err error + if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { + accounts, err = s.accountRepo.ListSchedulableByPlatform(ctx, platform) + } else if groupID != nil { + accounts, err = s.accountRepo.ListSchedulableByGroupIDAndPlatform(ctx, *groupID, platform) + } else { + accounts, err = s.accountRepo.ListSchedulableUngroupedByPlatform(ctx, platform) + } + if err != nil { + return nil, fmt.Errorf("query accounts failed: %w", err) + } + return accounts, nil +} + +func (s *OpenAIGatewayService) tryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (*AcquireResult, error) { + if s.concurrencyService == nil { + return &AcquireResult{Acquired: true, ReleaseFunc: func() {}}, nil + } + return s.concurrencyService.AcquireAccountSlot(ctx, accountID, maxConcurrency) +} + +func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { + if account == nil { + return nil + } + platform = normalizeOpenAICompatiblePlatform(platform) + + fresh := account + if s.schedulerSnapshot != nil { + current, err := s.getSchedulableAccount(ctx, account.ID) + if err != nil || current == nil { + return nil + } + fresh = current + } + + if !isOpenAICompatibleAccountEligibleForRequest(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) { + return nil + } + if !parentHealthyForShadow(fresh, s.parentAccountLookup(ctx)) { + return nil + } + if s.isOpenAIAccountRuntimeBlocked(fresh) { + return nil + } + return fresh +} + +// parentAccountLookup 返回供 parentHealthyForShadow 使用的母账号解析闭包:经 accountRepo +// 按 ID 取当前 Account(repo 为空时 fail-closed 返回 nil)。统一调度/粘连各路径的母账号解析, +// 取代各调用点重复内联的同一闭包(历史上 recheck 等路径还漏写过 accountRepo==nil 守卫)。 +// L2 候选循环改用带 per-pass 缓存的 parentLookupL2,不走此方法。 +func (s *OpenAIGatewayService) parentAccountLookup(ctx context.Context) func(int64) *Account { + return func(id int64) *Account { + if s.accountRepo == nil { + return nil + } + a, _ := s.accountRepo.GetByID(ctx, id) + return a + } +} + +func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { + if account == nil { + return nil + } + platform = normalizeOpenAICompatiblePlatform(platform) + if s.schedulerSnapshot == nil || s.accountRepo == nil { + if !isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, requireCompact, requiredCapability) { + return nil + } + if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { + return nil + } + return account + } + + latest, err := s.accountRepo.GetByID(ctx, account.ID) + if err != nil || latest == nil { + return nil + } + if !isOpenAICompatibleAccountEligibleForRequest(ctx, latest, platform, requestedModel, requireCompact, requiredCapability) { + return nil + } + if !parentHealthyForShadow(latest, s.parentAccountLookup(ctx)) { + return nil + } + if s.isOpenAIAccountRuntimeBlocked(latest) { + return nil + } + return latest +} + +func (s *OpenAIGatewayService) getSchedulableAccount(ctx context.Context, accountID int64) (*Account, error) { + var ( + account *Account + err error + ) + if s.schedulerSnapshot != nil { + account, err = s.schedulerSnapshot.GetAccount(ctx, accountID) + } else { + account, err = s.accountRepo.GetByID(ctx, accountID) + } + if err != nil || account == nil { + return account, err + } + return account, nil +} + +func (s *OpenAIGatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { + if account == nil || s.schedulerSnapshot == nil { + return account, nil + } + hydrated, err := s.schedulerSnapshot.GetAccount(ctx, account.ID) + if err != nil { + return nil, err + } + if hydrated == nil { + return nil, fmt.Errorf("selected openai account %d not found during hydration", account.ID) + } + return hydrated, nil +} + +func (s *OpenAIGatewayService) newSelectionResult(ctx context.Context, account *Account, acquired bool, release func(), waitPlan *AccountWaitPlan) (*AccountSelectionResult, error) { + hydrated, err := s.hydrateSelectedAccount(ctx, account) + if err != nil { + return nil, err + } + return &AccountSelectionResult{ + Account: hydrated, + Acquired: acquired, + ReleaseFunc: release, + WaitPlan: waitPlan, + }, nil +} + +func (s *OpenAIGatewayService) newAcquiredSelectionResult(ctx context.Context, account *Account, release func()) (*AccountSelectionResult, error) { + selection, err := s.newSelectionResult(ctx, account, true, release, nil) + if err != nil && release != nil { + release() + } + return selection, err +} + +func (s *OpenAIGatewayService) schedulingConfig() config.GatewaySchedulingConfig { + if s.cfg != nil { + return s.cfg.Gateway.Scheduling + } + return config.GatewaySchedulingConfig{ + StickySessionMaxWaiting: 3, + StickySessionWaitTimeout: 45 * time.Second, + FallbackWaitTimeout: 30 * time.Second, + FallbackMaxWaiting: 100, + LoadBatchEnabled: true, + SlotCleanupInterval: 30 * time.Second, + } +} diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index de3192ae54..935f4d9c45 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -13,7 +13,6 @@ import ( "log/slog" "math/rand" "net/http" - "sort" "strconv" "strings" "sync" @@ -26,8 +25,6 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat" - "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" - "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" "github.com/cespare/xxhash/v2" @@ -1242,1251 +1239,6 @@ func isOpenAIContextWindowError(upstreamMsg string, upstreamBody []byte) bool { return match(string(upstreamBody)) } -// ExtractSessionID extracts the raw session ID from headers or body without hashing. -// Used by ForwardAsAnthropic to pass as prompt_cache_key for upstream cache. -func (s *OpenAIGatewayService) ExtractSessionID(c *gin.Context, body []byte) string { - if c == nil { - return "" - } - sessionID := strings.TrimSpace(c.GetHeader("session_id")) - if sessionID == "" { - sessionID = strings.TrimSpace(c.GetHeader("conversation_id")) - } - if sessionID == "" && len(body) > 0 { - sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) - } - return sessionID -} - -func explicitOpenAISessionID(c *gin.Context, body []byte) string { - if c == nil { - return "" - } - - sessionID := strings.TrimSpace(c.GetHeader("session_id")) - if sessionID == "" { - sessionID = strings.TrimSpace(c.GetHeader("conversation_id")) - } - if sessionID == "" && len(body) > 0 { - sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) - } - return sessionID -} - -// GenerateExplicitSessionHash generates a sticky-session hash only from explicit -// client session signals. It intentionally skips content-derived fallback and is -// used by stateless endpoints such as /v1/images. -func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body []byte) string { - sessionID := explicitOpenAISessionID(c, body) - if sessionID == "" { - return "" - } - - currentHash, legacyHash := deriveOpenAISessionHashes(sessionID) - attachOpenAILegacySessionHashToGin(c, legacyHash) - return currentHash -} - -// GenerateSessionHash generates a sticky-session hash for OpenAI requests. -// -// Priority: -// 1. Header: session_id -// 2. Header: conversation_id -// 3. Body: prompt_cache_key (opencode) -// 4. Body: content-based fallback (model + system + tools + first user message) -func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) string { - if c == nil { - return "" - } - - sessionID := explicitOpenAISessionID(c, body) - if sessionID == "" && len(body) > 0 { - sessionID = deriveOpenAIContentSessionSeed(body) - } - if sessionID == "" { - return "" - } - - currentHash, legacyHash := deriveOpenAISessionHashes(sessionID) - attachOpenAILegacySessionHashToGin(c, legacyHash) - return currentHash -} - -// GenerateSessionHashWithFallback 先按常规信号生成会话哈希; -// 当未携带 session_id/conversation_id/prompt_cache_key 时,使用 fallbackSeed 生成稳定哈希。 -// 该方法用于 WS ingress,避免会话信号缺失时发生跨账号漂移。 -func (s *OpenAIGatewayService) GenerateSessionHashWithFallback(c *gin.Context, body []byte, fallbackSeed string) string { - sessionHash := s.GenerateSessionHash(c, body) - if sessionHash != "" { - return sessionHash - } - - seed := strings.TrimSpace(fallbackSeed) - if seed == "" { - return "" - } - - currentHash, legacyHash := deriveOpenAISessionHashes(seed) - attachOpenAILegacySessionHashToGin(c, legacyHash) - return currentHash -} - -func resolveOpenAIUpstreamOriginator(c *gin.Context, isOfficialClient bool) string { - if c != nil { - if originator := strings.TrimSpace(c.GetHeader("originator")); originator != "" { - return originator - } - } - if isOfficialClient { - return "codex_cli_rs" - } - return "opencode" -} - -// BindStickySession sets session -> account binding with standard TTL. -func (s *OpenAIGatewayService) BindStickySession(ctx context.Context, groupID *int64, sessionHash string, accountID int64) error { - if sessionHash == "" || accountID <= 0 { - return nil - } - ttl := openaiStickySessionTTL - if s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.StickySessionTTLSeconds > 0 { - ttl = time.Duration(s.cfg.Gateway.OpenAIWS.StickySessionTTLSeconds) * time.Second - } - return s.setStickySessionAccountID(ctx, groupID, sessionHash, accountID, ttl) -} - -// SelectAccount selects an OpenAI account with sticky session support -func (s *OpenAIGatewayService) SelectAccount(ctx context.Context, groupID *int64, sessionHash string) (*Account, error) { - return s.SelectAccountForModel(ctx, groupID, sessionHash, "") -} - -// SelectAccountForModel selects an account supporting the requested model -func (s *OpenAIGatewayService) SelectAccountForModel(ctx context.Context, groupID *int64, sessionHash string, requestedModel string) (*Account, error) { - return s.SelectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, nil) -} - -// SelectAccountForModelWithExclusions selects an account supporting the requested model while excluding specified accounts. -// SelectAccountForModelWithExclusions 选择支持指定模型的账号,同时排除指定的账号。 -func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*Account, error) { - return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, 0, "") -} - -// noAvailableOpenAISelectionError builds the standard "no account available" error -// while preserving the compact-specific error when applicable. -func normalizeOpenAICompatiblePlatform(platform string) string { - if platform == PlatformGrok { - return PlatformGrok - } - return PlatformOpenAI -} - -func noAvailableOpenAISelectionError(requestedModel string, compactBlocked bool) error { - if compactBlocked { - return ErrNoAvailableCompactAccounts - } - if requestedModel != "" { - return fmt.Errorf("no available OpenAI accounts supporting model: %s", requestedModel) - } - return errors.New("no available OpenAI accounts") -} - -// openAICompactSupportTier classifies an OpenAI account by compact capability. -// 0 = explicitly unsupported, 1 = unknown / not yet probed, 2 = explicitly supported. -func openAICompactSupportTier(account *Account) int { - if account == nil || !account.IsOpenAI() { - return 0 - } - supported, known := account.OpenAICompactSupportKnown() - if !known { - return 1 - } - if supported { - return 2 - } - return 0 -} - -// isOpenAICompatibleAccountEligibleForRequest 判断 OpenAI 兼容账号是否满足本次请求的调度条件。 -// 检查内容包括:平台匹配、账号可用性、quota 自动暂停、spark 路由限制、模型支持及端点能力。 -// -// 注意:对 spark 影子账号,调用方还须额外调用 parentHealthyForShadow(account, lookup) -// 检查母账号凭据可用性;该检查未内置于本函数,以避免注入 DB 依赖。 -func isOpenAICompatibleAccountEligibleForRequest(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) bool { - platform = normalizeOpenAICompatiblePlatform(platform) - if account == nil || account.Platform != platform || !account.IsOpenAICompatible() || !account.IsSchedulableForModelWithContext(ctx, requestedModel) { - return false - } - if account.IsOpenAI() { - if paused, reason := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { - // Debug level: this fires per-candidate on the scheduling hot path, so Info - // would amplify into log spam once several accounts cross the threshold. - slog.Debug("account_auto_paused_by_quota", - "account_id", account.ID, - "window", reason.window, - "threshold", reason.threshold, - "utilization", reason.utilization, - ) - return false - } - } - if account.IsGrok() { - if paused, reason := shouldAutoPauseGrokAccountByQuota(account); paused { - slog.Debug("grok_account_auto_paused_by_quota", - "account_id", account.ID, - "window", reason.window, - "threshold", reason.threshold, - "utilization", reason.utilization, - ) - return false - } - } - if requestedModel != "" && !account.IsModelSupported(requestedModel) { - return false - } - if !account.SupportsOpenAIEndpointCapability(requiredCapability) { - return false - } - if requireCompact && (!account.IsOpenAI() || openAICompactSupportTier(account) == 0) { - return false - } - return true -} - -type openAIQuotaAutoPauseDecision struct { - window string - threshold float64 - utilization float64 -} - -func shouldAutoPauseGrokAccountByQuota(account *Account) (bool, openAIQuotaAutoPauseDecision) { - if account == nil || !account.IsGrok() || account.Type != AccountTypeOAuth { - return false, openAIQuotaAutoPauseDecision{} - } - snapshot, err := grokQuotaSnapshotFromExtra(account.Extra) - if err != nil || snapshot == nil { - return false, openAIQuotaAutoPauseDecision{} - } - now := time.Now() - if grokQuotaSnapshotStaleForPause(snapshot, now) { - return false, openAIQuotaAutoPauseDecision{} - } - if grokQuotaRetryAfterActive(snapshot, now) { - return true, openAIQuotaAutoPauseDecision{window: "retry_after", threshold: 1, utilization: 1} - } - if paused, decision := shouldAutoPauseGrokQuotaWindow("requests", snapshot.Requests, now); paused { - return true, decision - } - if paused, decision := shouldAutoPauseGrokQuotaWindow("tokens", snapshot.Tokens, now); paused { - return true, decision - } - return false, openAIQuotaAutoPauseDecision{} -} - -func grokQuotaRetryAfterActive(snapshot *xai.QuotaSnapshot, now time.Time) bool { - if snapshot == nil || snapshot.RetryAfterSeconds == nil || *snapshot.RetryAfterSeconds <= 0 { - return false - } - if strings.TrimSpace(snapshot.UpdatedAt) == "" { - return true - } - updatedAt, err := parseTime(snapshot.UpdatedAt) - if err != nil { - return true - } - retryAfterUntil := updatedAt.Add(time.Duration(*snapshot.RetryAfterSeconds) * time.Second) - return now.Before(retryAfterUntil) -} - -func shouldAutoPauseGrokQuotaWindow(name string, window *xai.QuotaWindow, now time.Time) (bool, openAIQuotaAutoPauseDecision) { - if window == nil || window.Limit == nil || window.Remaining == nil || *window.Limit <= 0 { - return false, openAIQuotaAutoPauseDecision{} - } - if window.ResetUnix != nil && *window.ResetUnix > 0 && !now.Before(time.Unix(*window.ResetUnix, 0)) { - return false, openAIQuotaAutoPauseDecision{} - } - utilization := float64(*window.Limit-*window.Remaining) / float64(*window.Limit) - if *window.Remaining <= 0 || utilization >= 1 { - return true, openAIQuotaAutoPauseDecision{window: name, threshold: 1, utilization: utilization} - } - return false, openAIQuotaAutoPauseDecision{} -} - -func grokQuotaSnapshotStaleForPause(snapshot *xai.QuotaSnapshot, now time.Time) bool { - if snapshot == nil || strings.TrimSpace(snapshot.UpdatedAt) == "" { - return false - } - updatedAt, err := parseTime(snapshot.UpdatedAt) - if err != nil { - return false - } - return now.Sub(updatedAt) >= openAICodexAutoPauseStaleAfter -} - -func shouldAutoPauseOpenAIAccountByQuota(ctx context.Context, account *Account) (bool, openAIQuotaAutoPauseDecision) { - if account == nil || !account.IsOpenAI() { - return false, openAIQuotaAutoPauseDecision{} - } - // Per-account explicit-disable flags must take precedence over the global default. - // Without these, leaving the account threshold blank means "use global default", - // so an admin has no way to exempt a single account from auto-pause once a global - // default exists. The disable flag is per-window so an account can opt out of - // only 5h or only 7d auto-pause. - disabled5h := resolveAccountExtraBool(account.Extra, "auto_pause_5h_disabled") - disabled7d := resolveAccountExtraBool(account.Extra, "auto_pause_7d_disabled") - threshold5h, threshold7d := resolveOpenAIQuotaAutoPauseThresholds(ctx, account) - now := time.Now() - if !disabled5h && threshold5h > 0 { - if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "5h", now); ok && utilization >= threshold5h { - return true, openAIQuotaAutoPauseDecision{window: "5h", threshold: threshold5h, utilization: utilization} - } - } - if !disabled7d && threshold7d > 0 { - if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "7d", now); ok && utilization >= threshold7d { - return true, openAIQuotaAutoPauseDecision{window: "7d", threshold: threshold7d, utilization: utilization} - } - } - return false, openAIQuotaAutoPauseDecision{} -} - -// resolveAccountExtraBool reads a bool-like value from account extra, tolerating -// the few shapes JSON unmarshalling may produce (real bool, "true"/"false" -// strings, 0/1 numbers). -func resolveAccountExtraBool(extra map[string]any, key string) bool { - if len(extra) == 0 { - return false - } - value, ok := extra[key] - if !ok || value == nil { - return false - } - switch v := value.(type) { - case bool: - return v - case string: - parsed, err := strconv.ParseBool(strings.TrimSpace(v)) - return err == nil && parsed - case float64: - return v != 0 - case float32: - return v != 0 - case int: - return v != 0 - case int64: - return v != 0 - case json.Number: - if i, err := v.Int64(); err == nil { - return i != 0 - } - } - return false -} - -func resolveOpenAIQuotaAutoPauseThresholds(ctx context.Context, account *Account) (float64, float64) { - threshold5h, _ := resolveAccountExtraNumber(account.Extra, "auto_pause_5h_threshold") - threshold7d, _ := resolveAccountExtraNumber(account.Extra, "auto_pause_7d_threshold") - threshold5h = clamp01(threshold5h) - threshold7d = clamp01(threshold7d) - if threshold5h > 0 && threshold7d > 0 { - return threshold5h, threshold7d - } - settings := openAIQuotaAutoPauseSettingsFromContext(ctx) - if threshold5h <= 0 { - threshold5h = clamp01(settings.DefaultThreshold5h) - } - if threshold7d <= 0 { - threshold7d = clamp01(settings.DefaultThreshold7d) - } - return threshold5h, threshold7d -} - -func resolveAccountExtraNumber(extra map[string]any, keys ...string) (float64, bool) { - if len(extra) == 0 { - return 0, false - } - for _, key := range keys { - value, ok := extra[key] - if !ok || value == nil { - continue - } - switch v := value.(type) { - case float64: - return v, true - case float32: - return float64(v), true - case int: - return float64(v), true - case int64: - return float64(v), true - case json.Number: - parsed, err := v.Float64() - if err == nil { - return parsed, true - } - case string: - parsed, err := strconv.ParseFloat(strings.TrimSpace(v), 64) - if err == nil { - return parsed, true - } - } - } - return 0, false -} - -// resolveOpenAIQuotaUtilization returns the current utilization ratio (0..1) for the -// given Codex usage window. ok=false means there is no usable signal to pause on: -// either no snapshot exists, or the window has already rolled over so the cached -// percentage is stale. The stale guard matters because a paused account stops -// receiving requests, so its snapshot is never refreshed from upstream headers — -// without this check an old used_percent would keep the account paused forever even -// after the real window reset. -func resolveOpenAIQuotaUtilization(extra map[string]any, window string, now time.Time) (float64, bool) { - usedPercent := readOpenAIQuotaUsedPercent(extra, window) - if usedPercent <= 0 { - return 0, false - } - if openAIQuotaWindowReset(extra, window, now) { - return 0, false - } - // 快照过于陈旧(账号长期未收到流量刷新)时,不再据此暂停。放行后下一次响应头 - // 会刷新快照实现自愈,避免账号在错误/过期的 used% 上被永久跳过(issue #2994)。 - if openAICodexSnapshotStaleForPause(extra, now) { - return 0, false - } - return usedPercent / 100, true -} - -// openAICodexSnapshotStaleForPause reports whether the Codex usage snapshot is stale -// enough that it should no longer keep an account auto-paused. It anchors on -// codex_usage_updated_at (always written by buildCodexUsageExtraUpdates). A missing or -// unparseable timestamp returns false (treated as fresh, so the account stays paused) — -// this is deliberate: it prevents any snapshot without a write time from silently escaping -// auto-pause, and a genuinely-exhausted account that is actively served refreshes the -// timestamp on every response so it never crosses the staleness bound. -func openAICodexSnapshotStaleForPause(extra map[string]any, now time.Time) bool { - if len(extra) == 0 { - return false - } - updatedRaw, ok := extra["codex_usage_updated_at"] - if !ok { - return false - } - updatedAt, err := parseTime(fmt.Sprint(updatedRaw)) - if err != nil { - return false - } - return now.Sub(updatedAt) >= openAICodexAutoPauseStaleAfter -} - -// openAIQuotaWindowReset reports whether the Codex usage window's reset time has -// already passed relative to now. It prefers the absolute codex__reset_at -// timestamp and falls back to codex__reset_after_seconds anchored at -// codex_usage_updated_at, mirroring AccountUsageService's window-progress logic. -func openAIQuotaWindowReset(extra map[string]any, window string, now time.Time) bool { - if len(extra) == 0 { - return false - } - if resetAtRaw, ok := extra["codex_"+window+"_reset_at"]; ok { - if resetAt, err := parseTime(fmt.Sprint(resetAtRaw)); err == nil { - return !now.Before(resetAt) - } - } - resetAfter := parseExtraInt(extra["codex_"+window+"_reset_after_seconds"]) - if resetAfter <= 0 { - return false - } - base := now - if updatedRaw, ok := extra["codex_usage_updated_at"]; ok { - if updatedAt, err := parseTime(fmt.Sprint(updatedRaw)); err == nil { - base = updatedAt - } - } - resetAt := base.Add(time.Duration(resetAfter) * time.Second) - return !now.Before(resetAt) -} - -func readOpenAIQuotaUsedPercent(extra map[string]any, window string) float64 { - if len(extra) == 0 { - return 0 - } - if value, ok := resolveAccountExtraNumber(extra, "codex_"+window+"_used_percent"); ok { - return value - } - return 0 -} - -type openAIQuotaAutoPauseCtxKey struct{} - -func withOpenAIQuotaAutoPauseSettings(ctx context.Context, settings OpsOpenAIAccountQuotaAutoPauseSettings) context.Context { - if ctx == nil { - ctx = context.Background() - } - return context.WithValue(ctx, openAIQuotaAutoPauseCtxKey{}, settings) -} - -func openAIQuotaAutoPauseSettingsFromContext(ctx context.Context) OpsOpenAIAccountQuotaAutoPauseSettings { - if ctx == nil { - return OpsOpenAIAccountQuotaAutoPauseSettings{} - } - settings, _ := ctx.Value(openAIQuotaAutoPauseCtxKey{}).(OpsOpenAIAccountQuotaAutoPauseSettings) - return settings -} - -func (s *OpenAIGatewayService) withOpenAIQuotaAutoPauseContext(ctx context.Context) context.Context { - if s == nil || s.settingService == nil { - return ctx - } - return withOpenAIQuotaAutoPauseSettings(ctx, s.settingService.GetOpenAIQuotaAutoPauseSettings(ctx)) -} - -// prioritizeOpenAICompactAccounts re-orders a slice so that accounts with known -// compact support are tried first, followed by unknown, then explicitly unsupported. -// The relative order within each tier is preserved. -func prioritizeOpenAICompactAccounts(accounts []*Account) []*Account { - if len(accounts) == 0 { - return nil - } - supported := make([]*Account, 0, len(accounts)) - unknown := make([]*Account, 0, len(accounts)) - unsupported := make([]*Account, 0, len(accounts)) - for _, account := range accounts { - switch openAICompactSupportTier(account) { - case 2: - supported = append(supported, account) - case 1: - unknown = append(unknown, account) - default: - unsupported = append(unsupported, account) - } - } - out := make([]*Account, 0, len(accounts)) - out = append(out, supported...) - out = append(out, unknown...) - out = append(out, unsupported...) - return out -} - -// resolveOpenAIAccountUpstreamModelForRequest resolves the upstream model that -// would be sent for a given request, honouring compact-only mappings when the -// caller is on the /responses/compact path. -func resolveOpenAIAccountUpstreamModelForRequest(account *Account, requestedModel string, requireCompact bool) string { - upstreamModel := resolveOpenAIForwardModel(account, requestedModel, "") - if upstreamModel == "" { - return "" - } - if requireCompact { - return resolveOpenAICompactForwardModel(account, upstreamModel) - } - return upstreamModel -} - -func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) (*Account, error) { - platform = normalizeOpenAICompatiblePlatform(platform) - if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { - slog.Warn("channel pricing restriction blocked request", - "group_id", derefGroupID(groupID), - "model", requestedModel) - return nil, fmt.Errorf("%w supporting model: %s (channel pricing restriction)", ErrNoAvailableAccounts, requestedModel) - } - - // 1. 尝试粘性会话命中 - // Try sticky session hit - if account := s.tryStickySessionHit(ctx, groupID, platform, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability); account != nil { - return account, nil - } - - // 2. 获取可调度的 OpenAI 账号 - // Get schedulable OpenAI accounts - accounts, err := s.listSchedulableAccounts(ctx, groupID, platform) - if err != nil { - return nil, fmt.Errorf("query accounts failed: %w", err) - } - - // 3. 按优先级 + LRU 选择最佳账号 - // Select by priority + LRU - selected, compactBlocked := s.selectBestAccount(ctx, groupID, platform, accounts, requestedModel, excludedIDs, requireCompact, requiredCapability) - - if selected == nil { - return nil, noAvailableOpenAISelectionError(requestedModel, compactBlocked) - } - - hydrated, err := s.hydrateSelectedAccount(ctx, selected) - if err != nil { - return nil, err - } - - // 4. 设置粘性会话绑定 - // Set sticky session binding - if sessionHash != "" { - _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, selected.ID, openaiStickySessionTTL) - } - - return hydrated, nil -} - -// tryStickySessionHit 尝试从粘性会话获取账号。 -// 如果命中且账号可用则返回账号;如果账号不可用则清理会话并返回 nil。 -// -// tryStickySessionHit attempts to get account from sticky session. -// Returns account if hit and usable; clears session and returns nil if account is unavailable. -func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID *int64, platform string, sessionHash, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) *Account { - if sessionHash == "" { - return nil - } - platform = normalizeOpenAICompatiblePlatform(platform) - - accountID := stickyAccountID - if accountID <= 0 { - var err error - accountID, err = s.getStickySessionAccountID(ctx, groupID, sessionHash) - if err != nil || accountID <= 0 { - return nil - } - } - - if _, excluded := excludedIDs[accountID]; excluded { - return nil - } - - account, err := s.getSchedulableAccount(ctx, accountID) - if err != nil { - return nil - } - - // 检查账号是否需要清理粘性会话 - // Check if sticky session should be cleared - if shouldClearStickySession(account, requestedModel) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - return nil - } - - // 验证账号是否可用于当前请求 - // Verify account is usable for current request - if !isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, false, requiredCapability) { - return nil - } - if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - return nil - } - if s.isOpenAIAccountRuntimeBlocked(account) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - return nil - } - account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, platform, requestedModel, requireCompact, requiredCapability) - if account == nil || !openAIStickyAccountMatchesGroup(account, groupID) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - return nil - } - if groupID != nil && s.needsUpstreamChannelRestrictionCheck(ctx, groupID) && - s.isUpstreamModelRestrictedByChannel(ctx, *groupID, account, requestedModel, requireCompact) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - return nil - } - - // 刷新会话 TTL 并返回账号 - // Refresh session TTL and return account - _ = s.refreshStickySessionTTL(ctx, groupID, sessionHash, openaiStickySessionTTL) - return account -} - -// selectBestAccount 从候选账号中选择最佳账号(优先级 + LRU)。 -// 返回 nil 表示无可用账号。 -// -// selectBestAccount selects the best account from candidates (priority + LRU). -// Returns nil if no available account. The second return reports whether at -// least one candidate was filtered out solely because it lacks compact support -// (only meaningful when requireCompact=true). -func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, platform string, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*Account, bool) { - platform = normalizeOpenAICompatiblePlatform(platform) - var selected *Account - selectedCompactTier := -1 - compactBlocked := false - needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) - - for i := range accounts { - acc := &accounts[i] - - // 跳过被排除的账号 - // Skip excluded accounts - if _, excluded := excludedIDs[acc.ID]; excluded { - continue - } - - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) - if fresh == nil { - continue - } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, false, requiredCapability) - if fresh == nil { - continue - } - if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel, requireCompact) { - continue - } - compactTier := 0 - if requireCompact { - compactTier = openAICompactSupportTier(fresh) - if compactTier == 0 { - compactBlocked = true - continue - } - } - - // 选择优先级最高且最久未使用的账号 - // Select highest priority and least recently used - if selected == nil { - selected = fresh - selectedCompactTier = compactTier - continue - } - - // compact 模式下高 tier 优先;同 tier 内才比较 priority/LRU。 - if requireCompact && compactTier != selectedCompactTier { - if compactTier > selectedCompactTier { - selected = fresh - selectedCompactTier = compactTier - } - continue - } - - if s.isBetterAccount(fresh, selected) { - selected = fresh - selectedCompactTier = compactTier - } - } - - return selected, compactBlocked -} - -// isBetterAccount 判断 candidate 是否比 current 更优。 -// 规则:优先级更高(数值更小)优先;同优先级时,未使用过的优先,其次是最久未使用的。 -// -// isBetterAccount checks if candidate is better than current. -// Rules: higher priority (lower value) wins; same priority: never used > least recently used. -func (s *OpenAIGatewayService) isBetterAccount(candidate, current *Account) bool { - // 优先级更高(数值更小) - // Higher priority (lower value) - if candidate.Priority < current.Priority { - return true - } - if candidate.Priority > current.Priority { - return false - } - - // 同优先级,比较最后使用时间 - // Same priority, compare last used time - switch { - case candidate.LastUsedAt == nil && current.LastUsedAt != nil: - // candidate 从未使用,优先 - return true - case candidate.LastUsedAt != nil && current.LastUsedAt == nil: - // current 从未使用,保持 - return false - case candidate.LastUsedAt == nil && current.LastUsedAt == nil: - // 都未使用,保持 - return false - default: - // 都使用过,选择最久未使用的 - return candidate.LastUsedAt.Before(*current.LastUsedAt) - } -} - -// SelectAccountWithLoadAwareness selects an account with load-awareness and wait plan. -func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*AccountSelectionResult, error) { - return s.selectAccountWithLoadAwareness(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, "") -} - -func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*AccountSelectionResult, error) { - platform = normalizeOpenAICompatiblePlatform(platform) - if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { - slog.Warn("channel pricing restriction blocked request", - "group_id", derefGroupID(groupID), - "model", requestedModel) - return nil, fmt.Errorf("%w supporting model: %s (channel pricing restriction)", ErrNoAvailableAccounts, requestedModel) - } - - cfg := s.schedulingConfig() - needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) - var stickyAccountID int64 - if sessionHash != "" && s.cache != nil { - if accountID, err := s.getStickySessionAccountID(ctx, groupID, sessionHash); err == nil { - stickyAccountID = accountID - } - } - if s.concurrencyService == nil || !cfg.LoadBatchEnabled { - account, err := s.selectAccountForModelWithExclusions(ctx, groupID, platform, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability) - if err != nil { - return nil, err - } - result, err := s.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) - if err == nil && result != nil && result.Acquired { - return s.newAcquiredSelectionResult(ctx, account, result.ReleaseFunc) - } - if stickyAccountID > 0 && stickyAccountID == account.ID && s.concurrencyService != nil { - waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, account.ID) - if waitingCount < cfg.StickySessionMaxWaiting { - return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ - AccountID: account.ID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }) - } - } - return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ - AccountID: account.ID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.FallbackWaitTimeout, - MaxWaiting: cfg.FallbackMaxWaiting, - }) - } - - accounts, err := s.listSchedulableAccounts(ctx, groupID, platform) - if err != nil { - return nil, err - } - if len(accounts) == 0 { - return nil, ErrNoAvailableAccounts - } - - isExcluded := func(accountID int64) bool { - if excludedIDs == nil { - return false - } - _, excluded := excludedIDs[accountID] - return excluded - } - - // ============ Layer 1: Sticky session ============ - if sessionHash != "" { - accountID := stickyAccountID - if accountID > 0 && !isExcluded(accountID) { - account, err := s.getSchedulableAccount(ctx, accountID) - if err == nil { - clearSticky := shouldClearStickySession(account, requestedModel) - if clearSticky { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - } - if !clearSticky && isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, false, requiredCapability) { - account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, platform, requestedModel, requireCompact, requiredCapability) - if account == nil { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - } else if !openAIStickyAccountMatchesGroup(account, groupID) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - } else if s.isOpenAIAccountRuntimeBlocked(account) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - } else if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, account, requestedModel, requireCompact) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - } else if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { - _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) - } else { - result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency) - if err == nil && result != nil && result.Acquired { - selection, selectErr := s.newAcquiredSelectionResult(ctx, account, result.ReleaseFunc) - if selectErr != nil { - return nil, selectErr - } - _ = s.refreshStickySessionTTL(ctx, groupID, sessionHash, openaiStickySessionTTL) - return selection, nil - } - - waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, accountID) - if waitingCount < cfg.StickySessionMaxWaiting { - return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ - AccountID: accountID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }) - } - } - } - } - } - } - - // ============ Layer 2: Load-aware selection ============ - // Per-pass parent-health cache to avoid repeated DB calls when multiple shadow - // accounts share the same parent. - parentCacheL2 := make(map[int64]*Account) - parentLookupL2 := func(id int64) *Account { - if a, ok := parentCacheL2[id]; ok { - return a - } - if s.accountRepo == nil { - return nil - } - a, _ := s.accountRepo.GetByID(ctx, id) - parentCacheL2[id] = a - return a - } - baseCandidateCount := 0 - candidates := make([]*Account, 0, len(accounts)) - for i := range accounts { - acc := &accounts[i] - if isExcluded(acc.ID) { - continue - } - // Scheduler snapshots can be temporarily stale (bucket rebuild is throttled); - // re-check schedulability here so recently rate-limited/overloaded accounts - // are not selected again before the bucket is rebuilt. - if !isOpenAICompatibleAccountEligibleForRequest(ctx, acc, platform, requestedModel, false, requiredCapability) { - continue - } - if !parentHealthyForShadow(acc, parentLookupL2) { - continue - } - if s.isOpenAIAccountRuntimeBlocked(acc) { - continue - } - if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel, requireCompact) { - continue - } - baseCandidateCount++ - candidates = append(candidates, acc) - } - - if len(candidates) == 0 { - return nil, ErrNoAvailableAccounts - } - - accountLoads := make([]AccountWithConcurrency, 0, len(candidates)) - for _, acc := range candidates { - accountLoads = append(accountLoads, AccountWithConcurrency{ - ID: acc.ID, - MaxConcurrency: acc.EffectiveLoadFactor(), - }) - } - - tryAcquireFromLoadMap := func(loadMap map[int64]*AccountLoadInfo) (*AccountSelectionResult, bool, error) { - var available []accountWithLoad - for _, acc := range candidates { - loadInfo := loadMap[acc.ID] - if loadInfo == nil { - loadInfo = &AccountLoadInfo{AccountID: acc.ID} - } - if loadInfo.LoadRate < 100 { - available = append(available, accountWithLoad{ - account: acc, - loadInfo: loadInfo, - }) - } - } - - if len(available) == 0 { - return nil, false, nil - } - - sort.SliceStable(available, func(i, j int) bool { - a, b := available[i], available[j] - if a.account.Priority != b.account.Priority { - return a.account.Priority < b.account.Priority - } - if a.loadInfo.LoadRate != b.loadInfo.LoadRate { - return a.loadInfo.LoadRate < b.loadInfo.LoadRate - } - switch { - case a.account.LastUsedAt == nil && b.account.LastUsedAt != nil: - return true - case a.account.LastUsedAt != nil && b.account.LastUsedAt == nil: - return false - case a.account.LastUsedAt == nil && b.account.LastUsedAt == nil: - return false - default: - return a.account.LastUsedAt.Before(*b.account.LastUsedAt) - } - }) - shuffleWithinSortGroups(available) - - selectionOrder := make([]accountWithLoad, 0, len(available)) - if requireCompact { - appendTier := func(out []accountWithLoad, tier int) []accountWithLoad { - for _, item := range available { - if openAICompactSupportTier(item.account) == tier { - out = append(out, item) - } - } - return out - } - selectionOrder = appendTier(selectionOrder, 2) - selectionOrder = appendTier(selectionOrder, 1) - // tier 0 候选作为兜底追加:DB recheck 时若发现 cache tier 0 实际 - // 已升级为 1/2(探测刚跑完,cache 尚未刷新),仍可正常命中。 - selectionOrder = appendTier(selectionOrder, 0) - } else { - selectionOrder = append(selectionOrder, available...) - } - - for _, item := range selectionOrder { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, item.account, platform, requestedModel, false, requiredCapability) - if fresh == nil { - continue - } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) - if fresh == nil { - continue - } - if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel, requireCompact) { - continue - } - result, err := s.tryAcquireAccountSlot(ctx, fresh.ID, fresh.Concurrency) - if err == nil && result != nil && result.Acquired { - selection, selectErr := s.newAcquiredSelectionResult(ctx, fresh, result.ReleaseFunc) - if selectErr != nil { - return nil, true, selectErr - } - if sessionHash != "" { - _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, fresh.ID, openaiStickySessionTTL) - } - return selection, true, nil - } - } - return nil, true, nil - } - - loadMap, err := s.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) - if err != nil { - ordered := append([]*Account(nil), candidates...) - sortAccountsByPriorityAndLastUsed(ordered, false) - if requireCompact { - ordered = prioritizeOpenAICompactAccounts(ordered) - } - for _, acc := range ordered { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) - if fresh == nil { - continue - } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) - if fresh == nil { - continue - } - if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel, requireCompact) { - continue - } - result, err := s.tryAcquireAccountSlot(ctx, fresh.ID, fresh.Concurrency) - if err == nil && result != nil && result.Acquired { - selection, selectErr := s.newAcquiredSelectionResult(ctx, fresh, result.ReleaseFunc) - if selectErr != nil { - return nil, selectErr - } - if sessionHash != "" { - _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, fresh.ID, openaiStickySessionTTL) - } - return selection, nil - } - } - } else { - if selection, attempted, selectErr := tryAcquireFromLoadMap(loadMap); selectErr != nil { - return nil, selectErr - } else if selection != nil { - return selection, nil - } else if attempted { - if freshLoadMap, loadErr := s.concurrencyService.GetAccountsLoadBatchFresh(ctx, accountLoads); loadErr == nil { - if selection, _, selectErr := tryAcquireFromLoadMap(freshLoadMap); selectErr != nil { - return nil, selectErr - } else if selection != nil { - return selection, nil - } - } - } - } - - // ============ Layer 3: Fallback wait ============ - sortAccountsByPriorityAndLastUsed(candidates, false) - if requireCompact { - candidates = prioritizeOpenAICompactAccounts(candidates) - } - for _, acc := range candidates { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) - if fresh == nil { - continue - } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) - if fresh == nil { - continue - } - if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel, requireCompact) { - continue - } - return s.newSelectionResult(ctx, fresh, false, nil, &AccountWaitPlan{ - AccountID: fresh.ID, - MaxConcurrency: fresh.Concurrency, - Timeout: cfg.FallbackWaitTimeout, - MaxWaiting: cfg.FallbackMaxWaiting, - }) - } - - if requireCompact && baseCandidateCount > 0 { - return nil, ErrNoAvailableCompactAccounts - } - return nil, ErrNoAvailableAccounts -} - -func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, groupID *int64, platform string) ([]Account, error) { - platform = normalizeOpenAICompatiblePlatform(platform) - if s.schedulerSnapshot != nil { - accounts, _, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, false) - return accounts, err - } - var accounts []Account - var err error - if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { - accounts, err = s.accountRepo.ListSchedulableByPlatform(ctx, platform) - } else if groupID != nil { - accounts, err = s.accountRepo.ListSchedulableByGroupIDAndPlatform(ctx, *groupID, platform) - } else { - accounts, err = s.accountRepo.ListSchedulableUngroupedByPlatform(ctx, platform) - } - if err != nil { - return nil, fmt.Errorf("query accounts failed: %w", err) - } - return accounts, nil -} - -func (s *OpenAIGatewayService) tryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (*AcquireResult, error) { - if s.concurrencyService == nil { - return &AcquireResult{Acquired: true, ReleaseFunc: func() {}}, nil - } - return s.concurrencyService.AcquireAccountSlot(ctx, accountID, maxConcurrency) -} - -func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { - if account == nil { - return nil - } - platform = normalizeOpenAICompatiblePlatform(platform) - - fresh := account - if s.schedulerSnapshot != nil { - current, err := s.getSchedulableAccount(ctx, account.ID) - if err != nil || current == nil { - return nil - } - fresh = current - } - - if !isOpenAICompatibleAccountEligibleForRequest(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) { - return nil - } - if !parentHealthyForShadow(fresh, s.parentAccountLookup(ctx)) { - return nil - } - if s.isOpenAIAccountRuntimeBlocked(fresh) { - return nil - } - return fresh -} - -// parentAccountLookup 返回供 parentHealthyForShadow 使用的母账号解析闭包:经 accountRepo -// 按 ID 取当前 Account(repo 为空时 fail-closed 返回 nil)。统一调度/粘连各路径的母账号解析, -// 取代各调用点重复内联的同一闭包(历史上 recheck 等路径还漏写过 accountRepo==nil 守卫)。 -// L2 候选循环改用带 per-pass 缓存的 parentLookupL2,不走此方法。 -func (s *OpenAIGatewayService) parentAccountLookup(ctx context.Context) func(int64) *Account { - return func(id int64) *Account { - if s.accountRepo == nil { - return nil - } - a, _ := s.accountRepo.GetByID(ctx, id) - return a - } -} - -func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { - if account == nil { - return nil - } - platform = normalizeOpenAICompatiblePlatform(platform) - if s.schedulerSnapshot == nil || s.accountRepo == nil { - if !isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, requireCompact, requiredCapability) { - return nil - } - if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { - return nil - } - return account - } - - latest, err := s.accountRepo.GetByID(ctx, account.ID) - if err != nil || latest == nil { - return nil - } - if !isOpenAICompatibleAccountEligibleForRequest(ctx, latest, platform, requestedModel, requireCompact, requiredCapability) { - return nil - } - if !parentHealthyForShadow(latest, s.parentAccountLookup(ctx)) { - return nil - } - if s.isOpenAIAccountRuntimeBlocked(latest) { - return nil - } - return latest -} - -func (s *OpenAIGatewayService) getSchedulableAccount(ctx context.Context, accountID int64) (*Account, error) { - var ( - account *Account - err error - ) - if s.schedulerSnapshot != nil { - account, err = s.schedulerSnapshot.GetAccount(ctx, accountID) - } else { - account, err = s.accountRepo.GetByID(ctx, accountID) - } - if err != nil || account == nil { - return account, err - } - return account, nil -} - -func (s *OpenAIGatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { - if account == nil || s.schedulerSnapshot == nil { - return account, nil - } - hydrated, err := s.schedulerSnapshot.GetAccount(ctx, account.ID) - if err != nil { - return nil, err - } - if hydrated == nil { - return nil, fmt.Errorf("selected openai account %d not found during hydration", account.ID) - } - return hydrated, nil -} - -func (s *OpenAIGatewayService) newSelectionResult(ctx context.Context, account *Account, acquired bool, release func(), waitPlan *AccountWaitPlan) (*AccountSelectionResult, error) { - hydrated, err := s.hydrateSelectedAccount(ctx, account) - if err != nil { - return nil, err - } - return &AccountSelectionResult{ - Account: hydrated, - Acquired: acquired, - ReleaseFunc: release, - WaitPlan: waitPlan, - }, nil -} - -func (s *OpenAIGatewayService) newAcquiredSelectionResult(ctx context.Context, account *Account, release func()) (*AccountSelectionResult, error) { - selection, err := s.newSelectionResult(ctx, account, true, release, nil) - if err != nil && release != nil { - release() - } - return selection, err -} - -func (s *OpenAIGatewayService) schedulingConfig() config.GatewaySchedulingConfig { - if s.cfg != nil { - return s.cfg.Gateway.Scheduling - } - return config.GatewaySchedulingConfig{ - StickySessionMaxWaiting: 3, - StickySessionWaitTimeout: 45 * time.Second, - FallbackWaitTimeout: 30 * time.Second, - FallbackMaxWaiting: 100, - LoadBatchEnabled: true, - SlotCleanupInterval: 30 * time.Second, - } -} - // GetAccessToken gets the access token for an OpenAI account func (s *OpenAIGatewayService) GetAccessToken(ctx context.Context, account *Account) (string, string, error) { if account.IsShadow() { @@ -3407,1068 +2159,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } } -func (s *OpenAIGatewayService) forwardOpenAIPassthrough( - ctx context.Context, - c *gin.Context, - account *Account, - body []byte, - reqModel string, - reasoningEffort *string, - reqStream bool, - startTime time.Time, -) (*OpenAIForwardResult, error) { - upstreamPassthroughModel := "" - if isOpenAIResponsesCompactPath(c) { - compactMappedModel := resolveOpenAICompactForwardModel(account, reqModel) - if compactMappedModel != "" && compactMappedModel != reqModel { - nextBody, setErr := sjson.SetBytes(body, "model", compactMappedModel) - if setErr != nil { - return nil, fmt.Errorf("set compact passthrough model: %w", setErr) - } - body = nextBody - upstreamPassthroughModel = compactMappedModel - } - } - - if account != nil && account.Type == AccountTypeOAuth { - if rejectReason := detectOpenAIPassthroughInstructionsRejectReason(reqModel, body); rejectReason != "" { - rejectMsg := "OpenAI codex passthrough requires a non-empty instructions field" - MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalPolicyDenied) - logOpenAIPassthroughInstructionsRejected(ctx, c, account, reqModel, rejectReason, body) - c.JSON(http.StatusForbidden, gin.H{ - "error": gin.H{ - "type": "forbidden_error", - "message": rejectMsg, - }, - }) - return nil, fmt.Errorf("openai passthrough rejected before upstream: %s", rejectReason) - } - - normalizedBody, normalized, err := normalizeOpenAIPassthroughOAuthBody(body, isOpenAIResponsesCompactPath(c)) - if err != nil { - return nil, err - } - if normalized { - body = normalizedBody - } - reqStream = gjson.GetBytes(body, "stream").Bool() - } - - sanitizedBody, sanitized, err := sanitizeEmptyBase64InputImagesInOpenAIBody(body) - if err != nil { - return nil, err - } - if sanitized { - body = sanitizedBody - } - - // Apply OpenAI fast policy to the passthrough body (filter/block by service_tier). - // 统一使用 upstream 视角的 model:透传路径下 body 已经过 compact 映射 + - // OAuth normalize,body 中的 model 字段即上游真正会看到的 slug。 - // 这样可以与 chat-completions / messages / native /responses 入口的 - // upstreamModel 保持一致,避免 whitelist 命中差异。当 body 中没有 - // model 字段时退回 reqModel。 - policyModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) - if policyModel == "" { - policyModel = reqModel - } - updatedBody, policyErr := s.applyOpenAIFastPolicyToBody(ctx, account, policyModel, body) - if policyErr != nil { - var blocked *OpenAIFastBlockedError - if errors.As(policyErr, &blocked) { - writeOpenAIFastPolicyBlockedResponse(c, blocked) - } - return nil, policyErr - } - body = updatedBody - - apiKey := getAPIKeyFromContext(c) - if IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body) && !GroupAllowsImageGeneration(apiKeyGroup(apiKey)) { - MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate) - c.JSON(http.StatusForbidden, gin.H{ - "error": gin.H{ - "type": "permission_error", - "message": ImageGenerationPermissionMessage(), - }, - }) - return nil, errors.New("image generation disabled for group") - } - imageBillingModel := "" - imageSizeTier := "" - imageInputSize := "" - if IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body) { - var imageCfgErr error - imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, reqModel) - if imageCfgErr != nil { - setOpsUpstreamError(c, http.StatusBadRequest, imageCfgErr.Error(), "") - c.JSON(http.StatusBadRequest, gin.H{ - "error": gin.H{ - "type": "invalid_request_error", - "message": imageCfgErr.Error(), - "param": "size", - }, - }) - return nil, imageCfgErr - } - imageBillingModel = imageCfg.Model - imageSizeTier = imageCfg.SizeTier - imageInputSize = imageCfg.InputSize - } - - logger.LegacyPrintf("service.openai_gateway", - "[OpenAI 自动透传] 命中自动透传分支: account=%d name=%s type=%s model=%s stream=%v", - account.ID, - account.Name, - account.Type, - reqModel, - reqStream, - ) - if reqStream && c != nil && c.Request != nil { - if timeoutHeaders := collectOpenAIPassthroughTimeoutHeaders(c.Request.Header); len(timeoutHeaders) > 0 { - streamWarnLogger := logger.FromContext(ctx).With( - zap.String("component", "service.openai_gateway"), - zap.Int64("account_id", account.ID), - zap.Strings("timeout_headers", timeoutHeaders), - ) - if s.isOpenAIPassthroughTimeoutHeadersAllowed() { - streamWarnLogger.Warn("OpenAI passthrough 透传请求包含超时相关请求头,且当前配置为放行,可能导致上游提前断流") - } else { - streamWarnLogger.Warn("OpenAI passthrough 检测到超时相关请求头,将按配置过滤以降低断流风险") - } - } - } - - // Get access token - token, _, err := s.GetAccessToken(ctx, account) - if err != nil { - return nil, err - } - - upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) - upstreamReq, err := s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token) - releaseUpstreamCtx() - if err != nil { - return nil, err - } - - proxyURL := "" - if account.ProxyID != nil && account.Proxy != nil { - proxyURL = account.Proxy.URL() - } - - if c != nil { - c.Set("openai_passthrough", true) - } - - upstreamStart := time.Now() - resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) - SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds()) - if err != nil { - // Transport-level failure (proxy/DNS/TCP/TLS — no HTTP response). Convert to - // a failover so the handler switches to a healthy account, and temporarily - // unschedule the account on durable faults (e.g. rejected proxy credentials). - return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) - } - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode >= 400 { - // 透传模式默认保持原样代理;但 429/529 属于网关必须兜底的 - // 上游容量类错误,应先触发多账号 failover 以维持基础 SLA。 - if shouldFailoverOpenAIPassthroughResponse(resp.StatusCode) { - return nil, s.handleFailoverErrorResponsePassthrough(ctx, resp, c, account, body) - } - return nil, s.handleErrorResponsePassthrough(ctx, resp, c, account, body) - } - - serviceTier := extractOpenAIServiceTierFromBody(body) - - var usage *OpenAIUsage - var firstTokenMs *int - responseID := "" - imageCount := 0 - var imageOutputSizes []string - if reqStream { - result, err := s.handleStreamingResponsePassthrough(ctx, resp, c, account, startTime, reqModel, upstreamPassthroughModel) - if err != nil { - return nil, err - } - usage = result.usage - firstTokenMs = result.firstTokenMs - responseID = strings.TrimSpace(result.responseID) - imageCount = result.imageCount - imageOutputSizes = result.imageOutputSizes - } else { - result, err := s.handleNonStreamingResponsePassthrough(ctx, resp, c, reqModel, upstreamPassthroughModel) - if err != nil { - return nil, err - } - usage = result.usage - responseID = strings.TrimSpace(result.responseID) - imageCount = result.imageCount - imageOutputSizes = result.imageOutputSizes - } - s.bindHTTPResponseAccount(ctx, c, account, responseID) - - // 排除 spark 影子:其 codex_* 仅由 QueryUsage(/wham/usage bengalfox)更新(外审第7轮 P1)。 - if !account.IsShadow() { - if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil { - s.updateCodexUsageSnapshot(ctx, account.ID, snapshot) - } - } - - if usage == nil { - usage = &OpenAIUsage{} - } - - forwardResult := &OpenAIForwardResult{ - RequestID: resp.Header.Get("x-request-id"), - ResponseID: responseID, - Usage: *usage, - Model: reqModel, - UpstreamModel: upstreamPassthroughModel, - ServiceTier: serviceTier, - ReasoningEffort: reasoningEffort, - Stream: reqStream, - OpenAIWSMode: false, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - } - if imageCount > 0 { - forwardResult.ImageCount = imageCount - forwardResult.ImageSize = imageSizeTier - forwardResult.ImageInputSize = imageInputSize - forwardResult.ImageOutputSizes = imageOutputSizes - forwardResult.BillingModel = imageBillingModel - } - return forwardResult, nil -} - -func logOpenAIPassthroughInstructionsRejected( - ctx context.Context, - c *gin.Context, - account *Account, - reqModel string, - rejectReason string, - body []byte, -) { - if ctx == nil { - ctx = context.Background() - } - accountID := int64(0) - accountName := "" - accountType := "" - if account != nil { - accountID = account.ID - accountName = strings.TrimSpace(account.Name) - accountType = strings.TrimSpace(string(account.Type)) - } - fields := []zap.Field{ - zap.String("component", "service.openai_gateway"), - zap.Int64("account_id", accountID), - zap.String("account_name", accountName), - zap.String("account_type", accountType), - zap.String("request_model", strings.TrimSpace(reqModel)), - zap.String("reject_reason", strings.TrimSpace(rejectReason)), - } - fields = appendCodexCLIOnlyRejectedRequestFields(fields, c, body) - logger.FromContext(ctx).With(fields...).Warn("OpenAI passthrough 本地拦截:Codex 请求缺少有效 instructions") -} - -func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( - ctx context.Context, - c *gin.Context, - account *Account, - body []byte, - token string, -) (*http.Request, error) { - targetURL := openaiPlatformAPIURL - switch account.Type { - case AccountTypeOAuth: - targetURL = chatgptCodexURL - case AccountTypeAPIKey: - baseURL := account.GetOpenAIBaseURL() - if baseURL != "" { - validatedURL, err := s.validateUpstreamBaseURL(baseURL) - if err != nil { - return nil, err - } - targetURL = buildOpenAIResponsesURL(validatedURL) - } - } - targetURL = appendOpenAIResponsesRequestPathSuffix(targetURL, openAIResponsesRequestPathSuffix(c)) - - req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) - if err != nil { - return nil, err - } - req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) - - // 透传客户端请求头(安全白名单)。 - allowTimeoutHeaders := s.isOpenAIPassthroughTimeoutHeadersAllowed() - if c != nil && c.Request != nil { - for key, values := range c.Request.Header { - lower := strings.ToLower(strings.TrimSpace(key)) - if !isOpenAIPassthroughAllowedRequestHeader(lower, allowTimeoutHeaders) { - continue - } - for _, v := range values { - req.Header.Add(key, v) - } - } - } - - // 覆盖入站鉴权残留,并注入上游认证 - req.Header.Del("authorization") - req.Header.Del("x-api-key") - req.Header.Del("x-goog-api-key") - req.Header.Set("authorization", "Bearer "+token) - - // OAuth 透传到 ChatGPT internal API 时补齐必要头。 - if account.Type == AccountTypeOAuth { - promptCacheKey := strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) - req.Host = "chatgpt.com" - if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, req.Header, account); err != nil { - return nil, fmt.Errorf("resolve chatgpt account headers: %w", err) - } - apiKeyID := getAPIKeyIDFromContext(c) - // 先保存客户端原始值,再做 compact 补充,避免后续统一隔离时读到已处理的值。 - clientSessionID := strings.TrimSpace(req.Header.Get("session_id")) - clientConversationID := strings.TrimSpace(req.Header.Get("conversation_id")) - if isOpenAIResponsesCompactPath(c) { - req.Header.Set("accept", "application/json") - if req.Header.Get("version") == "" { - req.Header.Set("version", codexCLIVersion) - } - if clientSessionID == "" { - clientSessionID = resolveOpenAICompactSessionID(c) - } - } else if req.Header.Get("accept") == "" { - req.Header.Set("accept", "text/event-stream") - } - if req.Header.Get("OpenAI-Beta") == "" { - req.Header.Set("OpenAI-Beta", "responses=experimental") - } - if req.Header.Get("originator") == "" { - req.Header.Set("originator", "codex_cli_rs") - } - // 用隔离后的 session 标识符覆盖客户端透传值,防止跨用户会话碰撞。 - if clientSessionID == "" { - clientSessionID = promptCacheKey - } - if clientConversationID == "" { - clientConversationID = promptCacheKey - } - if clientSessionID != "" { - req.Header.Set("session_id", isolateOpenAISessionID(apiKeyID, clientSessionID)) - } - if clientConversationID != "" { - req.Header.Set("conversation_id", isolateOpenAISessionID(apiKeyID, clientConversationID)) - } - } - - // 透传模式也支持账户自定义 User-Agent 与 ForceCodexCLI 兜底。 - customUA := account.GetOpenAIUserAgent() - if customUA != "" { - req.Header.Set("user-agent", customUA) - } - if s.cfg != nil && s.cfg.Gateway.ForceCodexCLI { - req.Header.Set("user-agent", codexCLIUserAgent) - } - // OAuth 安全透传:对非 Codex UA 统一兜底,降低被上游风控拦截概率。 - if account.Type == AccountTypeOAuth && !openai.IsCodexCLIRequest(req.Header.Get("user-agent")) { - req.Header.Set("user-agent", codexCLIUserAgent) - } - - // 浏览器型 UA 兜底:仅 OAuth(ChatGPT 内部接口)账号生效,若最终 user-agent 仍为浏览器 - // (Chrome/Firefox/Safari/Edge 等),替换为后台配置的 Codex UA,避免 Cloudflare 触发 JS 质询。 - s.overrideBrowserUserAgent(ctx, account, req) - - if req.Header.Get("content-type") == "" { - req.Header.Set("content-type", "application/json") - } - - // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) - account.ApplyHeaderOverrides(req.Header) - - return req, nil -} - -func shouldFailoverOpenAIPassthroughResponse(statusCode int) bool { - switch statusCode { - case http.StatusTooManyRequests, 529: - return true - default: - return false - } -} - -func (s *OpenAIGatewayService) handleFailoverErrorResponsePassthrough( - ctx context.Context, - resp *http.Response, - c *gin.Context, - account *Account, - requestBody []byte, -) error { - body := s.readUpstreamErrorBody(resp) - - upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body)) - upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) - upstreamDetail := "" - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes - if maxBytes <= 0 { - maxBytes = 2048 - } - upstreamDetail = truncateString(string(body), maxBytes) - } - setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail) - logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body) - reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody) - _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel) - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Passthrough: true, - Kind: "failover", - Message: upstreamMsg, - Detail: upstreamDetail, - UpstreamResponseBody: upstreamDetail, - }) - return &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: body, - ResponseHeaders: resp.Header.Clone(), - } -} - -func (s *OpenAIGatewayService) handleErrorResponsePassthrough( - ctx context.Context, - resp *http.Response, - c *gin.Context, - account *Account, - requestBody []byte, -) error { - MarkResponseCommitted(c) - body := s.readUpstreamErrorBody(resp) - - // cyber_policy:透传账号本就把原始 body 回给客户端(下方 c.Data),此处仅打标记, - // 供 handler 事后写风控/邮件。cyber 是上游网络安全策略拦截,不冷却账号, - // 故下方跳过 handleOpenAIAccountUpstreamError(避免自定义 temp-unschedulable 规则误冷却)。 - cyberHit, cyberCode, cyberMsg := detectOpenAICyberPolicy(body) - if cyberHit { - MarkOpsCyberPolicy(c, CyberPolicyMark{ - Code: cyberCode, - Message: cyberMsg, - Body: truncateString(string(body), 4096), - UpstreamStatus: resp.StatusCode, - }) - } - - upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body)) - upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) - upstreamDetail := "" - if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes - if maxBytes <= 0 { - maxBytes = 2048 - } - upstreamDetail = truncateString(string(body), maxBytes) - } - setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail) - logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body) - // 透传模式保留原始上游错误响应,但运行态账号状态仍需更新, - // 避免粘性路由继续复用刚被限流的账号。cyber 例外:不冷却账号。 - if !cyberHit { - reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody) - _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel) - } - appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ - Platform: account.Platform, - AccountID: account.ID, - AccountName: account.Name, - UpstreamStatusCode: resp.StatusCode, - UpstreamRequestID: resp.Header.Get("x-request-id"), - Passthrough: true, - Kind: "http_error", - Message: upstreamMsg, - Detail: upstreamDetail, - UpstreamResponseBody: upstreamDetail, - }) - - writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - contentType := resp.Header.Get("Content-Type") - if contentType == "" { - contentType = "application/json" - } - c.Data(resp.StatusCode, contentType, body) - - if upstreamMsg == "" { - return fmt.Errorf("upstream error: %d", resp.StatusCode) - } - return fmt.Errorf("upstream error: %d message=%s", resp.StatusCode, upstreamMsg) -} - -func isOpenAIPassthroughAllowedRequestHeader(lowerKey string, allowTimeoutHeaders bool) bool { - if lowerKey == "" { - return false - } - if isOpenAIPassthroughTimeoutHeader(lowerKey) { - return allowTimeoutHeaders - } - return openaiPassthroughAllowedHeaders[lowerKey] -} - -func isOpenAIPassthroughTimeoutHeader(lowerKey string) bool { - switch lowerKey { - case "x-stainless-timeout", "x-stainless-read-timeout", "x-stainless-connect-timeout", "x-request-timeout", "request-timeout", "grpc-timeout": - return true - default: - return false - } -} - -func (s *OpenAIGatewayService) isOpenAIPassthroughTimeoutHeadersAllowed() bool { - return s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIPassthroughAllowTimeoutHeaders -} - -func collectOpenAIPassthroughTimeoutHeaders(h http.Header) []string { - if h == nil { - return nil - } - var matched []string - for key, values := range h { - lowerKey := strings.ToLower(strings.TrimSpace(key)) - if isOpenAIPassthroughTimeoutHeader(lowerKey) { - entry := lowerKey - if len(values) > 0 { - entry = fmt.Sprintf("%s=%s", lowerKey, strings.Join(values, "|")) - } - matched = append(matched, entry) - } - } - sort.Strings(matched) - return matched -} - -type openaiStreamingResultPassthrough struct { - usage *OpenAIUsage - firstTokenMs *int - responseID string - imageCount int - imageOutputSizes []string -} - -type openaiNonStreamingResultPassthrough struct { - *OpenAIUsage - usage *OpenAIUsage - responseID string - imageCount int - imageOutputSizes []string -} - -func openAIStreamClientOutputStarted(c *gin.Context, localStarted bool) bool { - if localStarted { - return true - } - return c != nil && c.Writer != nil && c.Writer.Written() -} - -func openAIStreamEventIsPreamble(eventType string) bool { - switch strings.TrimSpace(eventType) { - case "response.created", "response.in_progress": - return true - default: - return false - } -} - -func openAIStreamDataStartsClientOutput(data, eventType string) bool { - trimmed := strings.TrimSpace(data) - if trimmed == "" { - return false - } - if strings.TrimSpace(eventType) == "response.failed" { - return false - } - return !openAIStreamEventIsPreamble(eventType) -} - -func openAIStreamFailedEventShouldFailover(payload []byte, message string) bool { - if isOpenAIContextWindowError(message, payload) { - return false - } - if isOpenAITransientProcessingError(http.StatusBadRequest, message, payload) { - return true - } - code := strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "response.error.code").String())) - if code == "" { - code = strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "error.code").String())) - } - errType := strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "response.error.type").String())) - if errType == "" { - errType = strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "error.type").String())) - } - combined := strings.ToLower(strings.TrimSpace(message + " " + code + " " + errType)) - if combined == "" { - return true - } - nonRetryableMarkers := []string{ - "invalid_request", - "content_policy", - "policy", - "safety", - "high-risk cyber", - "not allowed", - "violat", - } - for _, marker := range nonRetryableMarkers { - if strings.Contains(combined, marker) { - return false - } - } - return true -} - -func (s *OpenAIGatewayService) recordOpenAIStreamUpstreamError( - c *gin.Context, - account *Account, - passthrough bool, - upstreamRequestID string, - kind string, - payload []byte, - message string, -) string { - message = sanitizeUpstreamErrorMessage(strings.TrimSpace(message)) - if message == "" { - message = "OpenAI upstream response failed" - } - detail := "" - if len(payload) > 0 && s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { - maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes - if maxBytes <= 0 { - maxBytes = 2048 - } - detail = truncateString(string(payload), maxBytes) - } - if c != nil { - setOpsUpstreamError(c, http.StatusBadGateway, message, detail) - event := OpsUpstreamErrorEvent{ - Platform: PlatformOpenAI, - UpstreamStatusCode: http.StatusBadGateway, - UpstreamRequestID: strings.TrimSpace(upstreamRequestID), - Passthrough: passthrough, - Kind: kind, - Message: message, - Detail: detail, - } - if account != nil { - event.Platform = account.Platform - event.AccountID = account.ID - event.AccountName = account.Name - } - appendOpsUpstreamError(c, event) - } - return message -} - -func (s *OpenAIGatewayService) newOpenAIStreamFailoverError( - c *gin.Context, - account *Account, - passthrough bool, - upstreamRequestID string, - payload []byte, - message string, -) *UpstreamFailoverError { - message = sanitizeUpstreamErrorMessage(strings.TrimSpace(message)) - if message == "" { - message = "OpenAI stream disconnected before completion" - } - message = s.recordOpenAIStreamUpstreamError(c, account, passthrough, upstreamRequestID, "failover", payload, message) - body, _ := json.Marshal(gin.H{ - "error": gin.H{ - "type": "upstream_error", - "message": message, - }, - }) - return &UpstreamFailoverError{ - StatusCode: http.StatusBadGateway, - ResponseBody: body, - } -} - -func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( - ctx context.Context, - resp *http.Response, - c *gin.Context, - account *Account, - startTime time.Time, - originalModel string, - mappedModel string, -) (*openaiStreamingResultPassthrough, error) { - writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - - // SSE headers - c.Header("Content-Type", "text/event-stream") - c.Header("Cache-Control", "no-cache") - c.Header("Connection", "keep-alive") - c.Header("X-Accel-Buffering", "no") - if v := resp.Header.Get("x-request-id"); v != "" { - c.Header("x-request-id", v) - } - - w := c.Writer - flusher, ok := w.(http.Flusher) - if !ok { - return nil, errors.New("streaming not supported") - } - - usage := &OpenAIUsage{} - imageCounter := newOpenAIImageOutputCounter() - var firstTokenMs *int - responseID := "" - clientDisconnected := false - sawDone := false - sawTerminalEvent := false - sawFailedEvent := false - failedMessage := "" - clientOutputStarted := false - upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id")) - pendingLines := make([]string, 0, 8) - writePendingLines := func() bool { - for _, pending := range pendingLines { - if _, err := fmt.Fprintln(w, pending); err != nil { - clientDisconnected = true - logger.LegacyPrintf("service.openai_gateway", "[OpenAI passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) - return false - } - } - pendingLines = pendingLines[:0] - return true - } - - scanner := bufio.NewScanner(resp.Body) - maxLineSize := defaultMaxLineSize - if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { - maxLineSize = s.cfg.Gateway.MaxLineSize - } - scanBuf := getSSEScannerBuf64K() - scanner.Buffer(scanBuf[:0], maxLineSize) - defer putSSEScannerBuf64K(scanBuf) - - needModelReplace := strings.TrimSpace(originalModel) != "" && strings.TrimSpace(mappedModel) != "" && strings.TrimSpace(originalModel) != strings.TrimSpace(mappedModel) - resultWithUsage := func() *openaiStreamingResultPassthrough { - return &openaiStreamingResultPassthrough{ - usage: usage, - firstTokenMs: firstTokenMs, - responseID: responseID, - imageCount: imageCounter.Count(), - imageOutputSizes: imageCounter.Sizes(), - } - } - - for scanner.Scan() { - line := scanner.Text() - lineStartsClientOutput := false - forceFlushFailedEvent := false - if data, ok := extractOpenAISSEDataLine(line); ok { - dataBytes := []byte(data) - trimmedData := strings.TrimSpace(data) - if needModelReplace && strings.Contains(data, mappedModel) { - line = s.replaceModelInSSELine(line, mappedModel, originalModel) - if replacedData, replaced := extractOpenAISSEDataLine(line); replaced { - dataBytes = []byte(replacedData) - trimmedData = strings.TrimSpace(replacedData) - } - } - if normalizedData, normalized := normalizeOpenAIResponsesFunctionCallArguments(dataBytes); normalized { - dataBytes = normalizedData - trimmedData = strings.TrimSpace(string(normalizedData)) - line = "data: " + string(normalizedData) - } - eventType := strings.TrimSpace(gjson.Get(trimmedData, "type").String()) - if eventType == "response.failed" { - failedMessage = extractOpenAISSEErrorMessage(dataBytes) - // response.failed 自带上游已消耗的 usage(input token 通常已扣);必须先解析 - // 再打 cyber 标记,否则 mark 记到的是解析前的 0,导致流式 cyber 按 0 token 计费 - // 而漏记真实用量。对齐 WS V2 / Chat 流式路径(均先解析 usage 再 Mark)。 - s.parseSSEUsageBytes(dataBytes, usage) - if hit, code, msg := detectOpenAICyberPolicy(dataBytes); hit { - MarkOpsCyberPolicy(c, CyberPolicyMark{ - Code: code, - Message: msg, - Body: truncateString(string(dataBytes), 4096), - UpstreamStatus: http.StatusOK, - UpstreamInTok: usage.InputTokens, - UpstreamOutTok: usage.OutputTokens, - }) - } else if !openAIStreamClientOutputStarted(c, clientOutputStarted) && openAIStreamFailedEventShouldFailover(dataBytes, failedMessage) { - return resultWithUsage(), - s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, dataBytes, failedMessage) - } - forceFlushFailedEvent = true - sawFailedEvent = true - } - if trimmedData == "[DONE]" { - sawDone = true - } - if openAIStreamEventIsTerminal(trimmedData) { - sawTerminalEvent = true - } - if responseID == "" { - responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes) - } - imageCounter.AddSSEData(dataBytes) - if sanitizedData, sanitized := sanitizeOpenAIResponseFailedEventForClient(dataBytes, eventType); sanitized { - dataBytes = sanitizedData - trimmedData = strings.TrimSpace(string(sanitizedData)) - line = "data: " + string(sanitizedData) - } - lineStartsClientOutput = forceFlushFailedEvent || openAIStreamDataStartsClientOutput(trimmedData, eventType) - if firstTokenMs == nil && lineStartsClientOutput && trimmedData != "[DONE]" { - ms := int(time.Since(startTime).Milliseconds()) - firstTokenMs = &ms - } - s.parseSSEUsageBytes(dataBytes, usage) - } - - if !clientDisconnected { - if !clientOutputStarted && !lineStartsClientOutput { - pendingLines = append(pendingLines, line) - continue - } - if !clientOutputStarted && len(pendingLines) > 0 { - if !writePendingLines() { - continue - } - } - if _, err := fmt.Fprintln(w, line); err != nil { - clientDisconnected = true - logger.LegacyPrintf("service.openai_gateway", "[OpenAI passthrough] Client disconnected during streaming, continue draining upstream for usage: account=%d", account.ID) - } else { - clientOutputStarted = true - flusher.Flush() - } - } - } - if err := scanner.Err(); err != nil { - if sawTerminalEvent && !sawFailedEvent { - return resultWithUsage(), nil - } - if sawFailedEvent { - return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage) - } - if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { - return resultWithUsage(), fmt.Errorf("stream usage incomplete: %w", err) - } - if errors.Is(err, bufio.ErrTooLong) { - logger.LegacyPrintf("service.openai_gateway", "[OpenAI passthrough] SSE line too long: account=%d max_size=%d error=%v", account.ID, maxLineSize, err) - return resultWithUsage(), err - } - if !openAIStreamClientOutputStarted(c, clientOutputStarted) { - msg := "OpenAI stream disconnected before completion" - if errText := strings.TrimSpace(err.Error()); errText != "" { - msg += ": " + errText - } - return resultWithUsage(), - s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, nil, msg) - } - if clientDisconnected { - return resultWithUsage(), fmt.Errorf("stream usage incomplete after disconnect: %w", err) - } - logger.LegacyPrintf("service.openai_gateway", - "[OpenAI passthrough] 流读取异常中断: account=%d request_id=%s err=%v", - account.ID, - upstreamRequestID, - err, - ) - return resultWithUsage(), fmt.Errorf("stream read error: %w", err) - } - if sawFailedEvent { - return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage) - } - if !clientDisconnected && !sawDone && !sawTerminalEvent && ctx.Err() == nil { - logger.FromContext(ctx).With( - zap.String("component", "service.openai_gateway"), - zap.Int64("account_id", account.ID), - zap.String("upstream_request_id", upstreamRequestID), - ).Info("OpenAI passthrough 上游流在未收到 [DONE] 时结束,疑似断流") - if !openAIStreamClientOutputStarted(c, clientOutputStarted) { - return resultWithUsage(), - s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, nil, "OpenAI stream ended before a terminal event") - } - return resultWithUsage(), errors.New("stream usage incomplete: missing terminal event") - } - - return resultWithUsage(), nil -} - -func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( - ctx context.Context, - resp *http.Response, - c *gin.Context, - originalModel string, - mappedModel string, -) (*openaiNonStreamingResultPassthrough, error) { - body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) - if err != nil { - return nil, err - } - - // Detect SSE responses from upstream and convert to JSON. - // Some upstreams (e.g. other sub2api instances) may return SSE even when - // stream=false was requested. Without this conversion the client would - // receive raw SSE text or a terminal event with empty output. - if isEventStreamResponse(resp.Header) { - return s.handlePassthroughSSEToJSON(resp, c, body, originalModel, mappedModel) - } - - usage := &OpenAIUsage{} - usageParsed := false - if len(body) > 0 { - if parsedUsage, ok := extractOpenAIUsageFromJSONBytes(body); ok { - *usage = parsedUsage - usageParsed = true - } - } - if !usageParsed { - // 兜底:尝试从 SSE 文本中解析 usage - usage = s.parseSSEUsageFromBody(string(body)) - } - - writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - - contentType := resp.Header.Get("Content-Type") - if contentType == "" { - contentType = "application/json" - } - if originalModel != "" && mappedModel != "" && originalModel != mappedModel { - body = s.replaceModelInResponseBody(body, mappedModel, originalModel) - } - c.Data(resp.StatusCode, contentType, body) - return &openaiNonStreamingResultPassthrough{ - OpenAIUsage: usage, - usage: usage, - responseID: extractOpenAIResponseIDFromJSONBytes(body), - imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body), - imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body), - }, nil -} - -// handlePassthroughSSEToJSON converts an SSE response body into a JSON -// response for the passthrough path. It mirrors handleSSEToJSON while -// preserving passthrough payloads, except compact-only model remapping may -// rewrite model fields back to the original requested model. -func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c *gin.Context, body []byte, originalModel string, mappedModel string) (*openaiNonStreamingResultPassthrough, error) { - bodyText := string(body) - finalResponse, ok := extractCodexFinalResponse(bodyText) - - usage := &OpenAIUsage{} - if ok { - if parsedUsage, parsed := extractOpenAIUsageFromJSONBytes(finalResponse); parsed { - *usage = parsedUsage - } - // When the terminal event has an empty output array, reconstruct - // output from accumulated delta events so the client gets full content. - if len(gjson.GetBytes(finalResponse, "output").Array()) == 0 { - if outputJSON, reconstructed := reconstructResponseOutputFromSSE(bodyText); reconstructed { - if patched, err := sjson.SetRawBytes(finalResponse, "output", outputJSON); err == nil { - finalResponse = patched - } - } - } - body = finalResponse - if originalModel != "" && mappedModel != "" && originalModel != mappedModel { - body = s.replaceModelInResponseBody(body, mappedModel, originalModel) - } - // Correct tool calls in final response - body = s.correctToolCallsInResponseBody(body) - } else { - terminalType, terminalPayload, terminalOK := extractOpenAISSETerminalEvent(bodyText) - if terminalOK && terminalType == "response.failed" { - msg := extractOpenAISSEErrorMessage(terminalPayload) - if msg == "" { - msg = "Upstream compact response failed" - } - return nil, s.writeOpenAINonStreamingProtocolError(resp, c, msg) - } - usage = s.parseSSEUsageFromBody(bodyText) - if originalModel != "" && mappedModel != "" && originalModel != mappedModel { - bodyText = s.replaceModelInSSEBody(bodyText, mappedModel, originalModel) - } - body = []byte(bodyText) - } - - writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) - - contentType := "application/json; charset=utf-8" - if !ok { - contentType = resp.Header.Get("Content-Type") - if contentType == "" { - contentType = "text/event-stream" - } - } - c.Data(resp.StatusCode, contentType, body) - - return &openaiNonStreamingResultPassthrough{ - OpenAIUsage: usage, - usage: usage, - responseID: extractOpenAIResponseIDFromJSONBytes(body), - imageCount: countOpenAIImageOutputsFromSSEBody(bodyText), - imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText), - }, nil -} - -func writeOpenAIPassthroughResponseHeaders(dst http.Header, src http.Header, filter *responseheaders.CompiledHeaderFilter) { - if dst == nil || src == nil { - return - } - if filter != nil { - responseheaders.WriteFilteredHeaders(dst, src, filter) - } else { - // 兜底:尽量保留最基础的 content-type - if v := strings.TrimSpace(src.Get("Content-Type")); v != "" { - dst.Set("Content-Type", v) - } - } - // 透传模式强制放行 x-codex-* 响应头(若上游返回)。 - // 注意:真实 http.Response.Header 的 key 一般会被 canonicalize;但为了兼容测试/自建响应, - // 这里用 EqualFold 做一次大小写不敏感的查找。 - getCaseInsensitiveValues := func(h http.Header, want string) []string { - if h == nil { - return nil - } - for k, vals := range h { - if strings.EqualFold(k, want) { - return vals - } - } - return nil - } - - for _, rawKey := range []string{ - "x-codex-primary-used-percent", - "x-codex-primary-reset-after-seconds", - "x-codex-primary-window-minutes", - "x-codex-secondary-used-percent", - "x-codex-secondary-reset-after-seconds", - "x-codex-secondary-window-minutes", - "x-codex-primary-over-secondary-limit-percent", - } { - vals := getCaseInsensitiveValues(src, rawKey) - if len(vals) == 0 { - continue - } - key := http.CanonicalHeaderKey(rawKey) - dst.Del(key) - for _, v := range vals { - dst.Add(key, v) - } - } -} - func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string, isStream bool, promptCacheKey string, isCodexCLI bool) (*http.Request, error) { // Determine target URL based on account type var targetURL string @@ -6277,645 +3967,6 @@ func (s *OpenAIGatewayService) replaceModelInResponseBody(body []byte, fromModel return body } -// OpenAIRecordUsageInput input for recording usage -type OpenAIRecordUsageInput struct { - Result *OpenAIForwardResult - APIKey *APIKey - User *User - Account *Account - Subscription *UserSubscription - InboundEndpoint string - UpstreamEndpoint string - UserAgent string // 请求的 User-Agent - IPAddress string // 请求的客户端 IP 地址 - RequestPayloadHash string - APIKeyService APIKeyQuotaUpdater - QuotaPlatform string // user×platform quota platform resolved by the handler before async billing. - // CyberBlocked 为 true 时把该用量行标记为 cyber(request_type=cyber),计费逻辑不变。 - CyberBlocked bool - ChannelUsageFields -} - -// CyberPolicyUsageInput 是 cyber 拒绝、未走正常 RecordUsage 的请求记录用量的入参。 -// 用量按上游真实 token 计费,与 WS cyber 及正常请求口径一致(InputTokens/OutputTokens -// 取自上游 response.failed 报告的 usage,即 mark.UpstreamInTok/OutTok)。 -type CyberPolicyUsageInput struct { - APIKey *APIKey - Account *Account - Subscription *UserSubscription - RequestID string - Model string - Stream bool - InputTokens int - OutputTokens int - // 渠道归因与请求级 meta,使 cyber 计费行与正常 RecordUsage 行口径一致 - // (否则 cyber 行 channel_id 等为空,渠道维度统计会遗漏 cyber 命中)。 - InboundEndpoint string - UpstreamEndpoint string - UserAgent string - IPAddress string - RequestPayloadHash string - APIKeyService APIKeyQuotaUpdater - ChannelUsageFields -} - -// RecordCyberPolicyUsageLog 为被上游 cyber_policy 拒绝、未走正常 RecordUsage 的请求 -// (HTTP forward 返回错误路径)记录用量并按上游真实 token 计费,使其与 WS cyber 路径、 -// 与正常请求的计费口径统一(不再是 tokens=0 免费行)。token 取自上游 response.failed -// 报告的 usage(非流式直接拒通常为 0,cost 随之为 0)。复用 RecordUsage 完成成本计算、 -// 扣费与用量行写入(request_type=cyber 由 CyberBlocked 置位)。仅 forward 返回错误的 -// 路径由 handler 调用,避免与成功路径的正常 RecordUsage 重复。 -func (s *OpenAIGatewayService) RecordCyberPolicyUsageLog(ctx context.Context, in CyberPolicyUsageInput) { - if s == nil || in.APIKey == nil || in.APIKey.User == nil || in.Account == nil || strings.TrimSpace(in.Model) == "" { - return - } - result := &OpenAIForwardResult{ - RequestID: in.RequestID, - Model: in.Model, - Stream: in.Stream, - Usage: OpenAIUsage{ - InputTokens: in.InputTokens, - OutputTokens: in.OutputTokens, - }, - } - if err := s.RecordUsage(ctx, &OpenAIRecordUsageInput{ - Result: result, - APIKey: in.APIKey, - User: in.APIKey.User, - Account: in.Account, - Subscription: in.Subscription, - InboundEndpoint: in.InboundEndpoint, - UpstreamEndpoint: in.UpstreamEndpoint, - UserAgent: in.UserAgent, - IPAddress: in.IPAddress, - RequestPayloadHash: in.RequestPayloadHash, - APIKeyService: in.APIKeyService, - ChannelUsageFields: in.ChannelUsageFields, - CyberBlocked: true, - }); err != nil { - logger.LegacyPrintf("service.openai_gateway", "cyber usage record failed: request_id=%s err=%v", in.RequestID, err) - } -} - -// RecordUsage records usage and deducts balance -func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRecordUsageInput) error { - if input == nil { - return errors.New("openai usage input is nil") - } - result := input.Result - if result == nil { - return errors.New("openai usage result is nil") - } - if s.rateLimitService != nil && input.Account != nil && input.Account.Platform == PlatformOpenAI { - s.rateLimitService.ResetOpenAI403Counter(ctx, input.Account.ID) - } - - apiKey := input.APIKey - user := input.User - account := input.Account - subscription := input.Subscription - ApplyOpenAIImageBillingResolution(result) - - // 计算实际的新输入token(减去缓存读取的token) - // 因为 input_tokens 包含了 cache_read_tokens,而缓存读取的token不应按输入价格计费 - actualInputTokens := result.Usage.InputTokens - result.Usage.CacheReadInputTokens - if actualInputTokens < 0 { - actualInputTokens = 0 - } - - // Calculate cost - tokens := UsageTokens{ - InputTokens: actualInputTokens, - ImageInputTokens: result.Usage.ImageInputTokens, - OutputTokens: result.Usage.OutputTokens, - CacheCreationTokens: result.Usage.CacheCreationInputTokens, - CacheReadTokens: result.Usage.CacheReadInputTokens, - ImageOutputTokens: result.Usage.ImageOutputTokens, - } - - // Get rate multiplier - multiplier := 1.0 - if s.cfg != nil { - multiplier = s.cfg.Default.RateMultiplier - } - if apiKey.GroupID != nil && apiKey.Group != nil { - resolver := s.userGroupRateResolver - if resolver == nil { - resolver = newUserGroupRateResolver(nil, nil, resolveUserGroupRateCacheTTL(s.cfg), nil, "service.openai_gateway") - } - multiplier = resolver.Resolve(ctx, user.ID, *apiKey.GroupID, apiKey.Group.RateMultiplier) - } - // token 倍率叠加高峰因子(token 计费含图片 token,图片按次倍率不受影响)。高峰因子按请求时刻现算, - // 不并入上面的 Resolve,以免污染 user:group 倍率缓存。 - multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, multiplier, timezone.Now()) - - var cost *CostBreakdown - var err error - billingModel := forwardResultBillingModel(result.Model, result.UpstreamModel) - if result.BillingModel != "" { - billingModel = strings.TrimSpace(result.BillingModel) - } - if input.BillingModelSource == BillingModelSourceChannelMapped && input.ChannelMappedModel != "" && input.ChannelMappedModel != input.OriginalModel { - billingModel = input.ChannelMappedModel - } - if input.BillingModelSource == BillingModelSourceRequested && input.OriginalModel != "" { - billingModel = input.OriginalModel - } - billingModels := usageBillingModelCandidates( - billingModel, - result.BillingModel, - input.ChannelMappedModel, - input.OriginalModel, - result.UpstreamModel, - result.Model, - ) - serviceTier := "" - if result.ServiceTier != nil { - serviceTier = strings.TrimSpace(*result.ServiceTier) - } - cost, err = s.calculateOpenAIRecordUsageCost(ctx, result, apiKey, billingModels, multiplier, imageMultiplier, tokens, serviceTier) - if err != nil { - if !isUsagePricingUnavailableError(err) { - return err - } - logger.L().With( - zap.String("component", "service.openai_gateway"), - zap.Strings("billing_models", billingModels), - zap.String("requested_model", input.OriginalModel), - zap.String("mapped_model", input.ChannelMappedModel), - zap.String("upstream_model", result.UpstreamModel), - zap.Int64("api_key_id", apiKey.ID), - zap.Int64("account_id", account.ID), - ).Warn("openai_usage.pricing_missing_record_zero_cost", zap.Error(err)) - cost = &CostBreakdown{BillingMode: string(BillingModeToken)} - } - - // Determine billing type - isSubscriptionBilling := subscription != nil && apiKey.Group != nil && apiKey.Group.IsSubscriptionType() - billingType := BillingTypeBalance - if isSubscriptionBilling { - billingType = BillingTypeSubscription - } - - // Create usage log - durationMs := int(result.Duration.Milliseconds()) - accountRateMultiplier := account.BillingRateMultiplier() - requestID := resolveUsageBillingRequestID(ctx, result.RequestID) - if result.OpenAIWSMode { - if upstreamRequestID := strings.TrimSpace(result.RequestID); upstreamRequestID != "" { - requestID = upstreamRequestID - } - } - - // 确定 RequestedModel(渠道映射前的原始模型) - requestedModel := result.Model - if input.OriginalModel != "" { - requestedModel = input.OriginalModel - } - - usageLog := &UsageLog{ - UserID: user.ID, - APIKeyID: apiKey.ID, - AccountID: account.ID, - RequestID: requestID, - Model: result.Model, - RequestedModel: requestedModel, - UpstreamModel: optionalNonEqualStringPtr(result.UpstreamModel, result.Model), - ServiceTier: result.ServiceTier, - ReasoningEffort: result.ReasoningEffort, - InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint), - UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint), - InputTokens: actualInputTokens, - OutputTokens: result.Usage.OutputTokens, - CacheCreationTokens: result.Usage.CacheCreationInputTokens, - CacheReadTokens: result.Usage.CacheReadInputTokens, - ImageOutputTokens: result.Usage.ImageOutputTokens, - ImageCount: result.ImageCount, - ImageSize: optionalTrimmedStringPtr(result.ImageSize), - ImageInputSize: optionalTrimmedStringPtr(result.ImageInputSize), - ImageOutputSize: optionalTrimmedStringPtr(result.ImageOutputSize), - ImageSizeSource: optionalTrimmedStringPtr(result.ImageSizeSource), - ImageSizeBreakdown: result.ImageSizeBreakdown, - } - if cost != nil { - usageLog.InputCost = cost.InputCost - usageLog.OutputCost = cost.OutputCost - usageLog.ImageOutputCost = cost.ImageOutputCost - usageLog.CacheCreationCost = cost.CacheCreationCost - usageLog.CacheReadCost = cost.CacheReadCost - usageLog.TotalCost = cost.TotalCost - usageLog.ActualCost = cost.ActualCost - } - if result.ImageCount > 0 && (cost == nil || cost.BillingMode != string(BillingModeToken)) { - usageLog.RateMultiplier = imageMultiplier - } else { - usageLog.RateMultiplier = multiplier - } - usageLog.AccountRateMultiplier = &accountRateMultiplier - usageLog.BillingType = billingType - usageLog.Stream = result.Stream - if input.CyberBlocked { - usageLog.RequestType = RequestTypeCyberBlocked - } - usageLog.OpenAIWSMode = result.OpenAIWSMode - usageLog.DurationMs = &durationMs - usageLog.FirstTokenMs = result.FirstTokenMs - usageLog.CreatedAt = time.Now() - // 设置渠道信息 - usageLog.ChannelID = optionalInt64Ptr(input.ChannelID) - usageLog.ModelMappingChain = optionalTrimmedStringPtr(input.ModelMappingChain) - // 设置计费模式 - if cost != nil && cost.BillingMode != "" { - billingMode := cost.BillingMode - usageLog.BillingMode = &billingMode - } else if result.ImageCount > 0 { - billingMode := string(BillingModeImage) - usageLog.BillingMode = &billingMode - } else { - billingMode := string(BillingModeToken) - usageLog.BillingMode = &billingMode - } - // 添加 UserAgent - if input.UserAgent != "" { - usageLog.UserAgent = &input.UserAgent - } - - // 添加 IPAddress - if input.IPAddress != "" { - usageLog.IPAddress = &input.IPAddress - } - - if apiKey.GroupID != nil { - usageLog.GroupID = apiKey.GroupID - } - if subscription != nil { - usageLog.SubscriptionID = &subscription.ID - } - - // 计算账号统计定价费用(使用最终上游模型匹配自定义规则) - if apiKey.GroupID != nil { - applyAccountStatsCost(ctx, usageLog, s.channelService, s.billingService, - account.ID, *apiKey.GroupID, result.UpstreamModel, result.Model, - tokens, cost.TotalCost, - ) - } - - if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { - writeUsageLogBestEffort(ctx, s.usageLogRepo, usageLog, "service.openai_gateway") - logger.LegacyPrintf("service.openai_gateway", "[SIMPLE MODE] Usage recorded (not billed): user=%d, tokens=%d", usageLog.UserID, usageLog.TotalTokens()) - s.deferredService.ScheduleLastUsedUpdate(account.ID) - return nil - } - - // Async usage billing runs outside the original request context, so it - // cannot recover ForcePlatform there. Fall back for internal/test callers. - quotaPlatform := input.QuotaPlatform - if quotaPlatform == "" { - quotaPlatform = PlatformFromAPIKey(apiKey) - } - - billingErr := func() error { - _, err := applyUsageBilling(ctx, requestID, usageLog, &postUsageBillingParams{ - Cost: cost, - User: user, - APIKey: apiKey, - Account: account, - Subscription: subscription, - RequestPayloadHash: resolveUsageBillingPayloadFingerprint(ctx, input.RequestPayloadHash), - IsSubscriptionBill: isSubscriptionBilling, - AccountRateMultiplier: accountRateMultiplier, - APIKeyService: input.APIKeyService, - Platform: quotaPlatform, - }, s.billingDeps(), s.usageBillingRepo) - return err - }() - - if billingErr != nil { - return billingErr - } - writeUsageLogBestEffort(ctx, s.usageLogRepo, usageLog, "service.openai_gateway") - - return nil -} - -func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( - ctx context.Context, - result *OpenAIForwardResult, - apiKey *APIKey, - billingModels []string, - multiplier float64, - imageMultiplier float64, - tokens UsageTokens, - serviceTier string, -) (*CostBreakdown, error) { - billingModel := firstUsageBillingModel(billingModels) - if result != nil && result.ImageCount > 0 { - // 渠道定价为 token 计费时走 token 路径,否则走图片计费 - if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved == nil || resolved.Mode != BillingModeToken { - return s.calculateOpenAIImageCost(ctx, billingModel, apiKey, result, imageMultiplier), nil - } - } - if len(billingModels) == 0 || billingModel == "" { - return nil, errors.New("openai usage billing model is empty") - } - var lastErr error - for _, candidate := range billingModels { - candidate = strings.TrimSpace(candidate) - if candidate == "" { - continue - } - cost, err := s.calculateOpenAIRecordUsageTokenCost(ctx, apiKey, candidate, multiplier, tokens, serviceTier) - if err == nil { - return cost, nil - } - lastErr = err - } - if lastErr == nil { - lastErr = errors.New("no non-empty billing model candidates") - } - return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr) -} - -func isUsagePricingUnavailableError(err error) bool { - if err == nil { - return false - } - if errors.Is(err, ErrModelPricingUnavailable) { - return true - } - msg := strings.ToLower(err.Error()) - return strings.Contains(msg, "no pricing available") || strings.Contains(msg, "pricing not found") -} - -func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost( - ctx context.Context, - apiKey *APIKey, - billingModel string, - multiplier float64, - tokens UsageTokens, - serviceTier string, -) (*CostBreakdown, error) { - if s.resolver != nil && apiKey.Group != nil { - gid := apiKey.Group.ID - return s.billingService.CalculateCostUnified(CostInput{ - Ctx: ctx, - Model: billingModel, - GroupID: &gid, - Tokens: tokens, - RequestCount: 1, - RateMultiplier: multiplier, - ServiceTier: serviceTier, - Resolver: s.resolver, - }) - } - return s.billingService.CalculateCostWithServiceTier(billingModel, tokens, multiplier, serviceTier) -} - -func (s *OpenAIGatewayService) calculateOpenAIImageCost( - ctx context.Context, - billingModel string, - apiKey *APIKey, - result *OpenAIForwardResult, - multiplier float64, -) *CostBreakdown { - sizeTier := NormalizeImageBillingTierOrDefault(result.ImageSize) - if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved != nil && - (resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage) { - gid := apiKey.Group.ID - cost, err := s.billingService.CalculateCostUnified(CostInput{ - Ctx: ctx, - Model: billingModel, - GroupID: &gid, - RequestCount: result.ImageCount, - SizeTier: sizeTier, - RateMultiplier: multiplier, - Resolver: s.resolver, - Resolved: resolved, - }) - if err == nil { - return cost - } - logger.LegacyPrintf("service.openai_gateway", "Calculate image channel cost failed: %v", err) - } - - var groupConfig *ImagePriceConfig - if apiKey != nil && apiKey.Group != nil { - groupConfig = &ImagePriceConfig{ - Price1K: apiKey.Group.ImagePrice1K, - Price2K: apiKey.Group.ImagePrice2K, - Price4K: apiKey.Group.ImagePrice4K, - } - } - return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier) -} - -func (s *OpenAIGatewayService) resolveOpenAIChannelPricing(ctx context.Context, billingModel string, apiKey *APIKey) *ResolvedPricing { - if s.resolver == nil || apiKey == nil || apiKey.Group == nil { - return nil - } - gid := apiKey.Group.ID - resolved := s.resolver.Resolve(ctx, PricingInput{Model: billingModel, GroupID: &gid}) - if resolved.Source == PricingSourceChannel { - return resolved - } - return nil -} - -// ParseCodexRateLimitHeaders extracts Codex usage limits from response headers. -// Exported for use in ratelimit_service when handling OpenAI 429 responses. -func ParseCodexRateLimitHeaders(headers http.Header) *OpenAICodexUsageSnapshot { - snapshot := &OpenAICodexUsageSnapshot{} - hasData := false - - // Helper to parse float64 from header - parseFloat := func(key string) *float64 { - if v := headers.Get(key); v != "" { - if f, err := strconv.ParseFloat(v, 64); err == nil { - return &f - } - } - return nil - } - - // Helper to parse int from header - parseInt := func(key string) *int { - if v := headers.Get(key); v != "" { - if i, err := strconv.Atoi(v); err == nil { - return &i - } - } - return nil - } - - // Primary (weekly) limits - if v := parseFloat("x-codex-primary-used-percent"); v != nil { - snapshot.PrimaryUsedPercent = v - hasData = true - } - if v := parseInt("x-codex-primary-reset-after-seconds"); v != nil { - snapshot.PrimaryResetAfterSeconds = v - hasData = true - } - if v := parseInt("x-codex-primary-window-minutes"); v != nil { - snapshot.PrimaryWindowMinutes = v - hasData = true - } - - // Secondary (5h) limits - if v := parseFloat("x-codex-secondary-used-percent"); v != nil { - snapshot.SecondaryUsedPercent = v - hasData = true - } - if v := parseInt("x-codex-secondary-reset-after-seconds"); v != nil { - snapshot.SecondaryResetAfterSeconds = v - hasData = true - } - if v := parseInt("x-codex-secondary-window-minutes"); v != nil { - snapshot.SecondaryWindowMinutes = v - hasData = true - } - - // Overflow ratio - if v := parseFloat("x-codex-primary-over-secondary-limit-percent"); v != nil { - snapshot.PrimaryOverSecondaryPercent = v - hasData = true - } - - if !hasData { - return nil - } - - snapshot.UpdatedAt = time.Now().Format(time.RFC3339) - return snapshot -} - -func codexSnapshotBaseTime(snapshot *OpenAICodexUsageSnapshot, fallback time.Time) time.Time { - if snapshot == nil { - return fallback - } - if snapshot.UpdatedAt == "" { - return fallback - } - base, err := time.Parse(time.RFC3339, snapshot.UpdatedAt) - if err != nil { - return fallback - } - return base -} - -func codexResetAtRFC3339(base time.Time, resetAfterSeconds *int) *string { - if resetAfterSeconds == nil { - return nil - } - sec := *resetAfterSeconds - if sec < 0 { - sec = 0 - } - resetAt := base.Add(time.Duration(sec) * time.Second).Format(time.RFC3339) - return &resetAt -} - -func buildCodexUsageExtraUpdates(snapshot *OpenAICodexUsageSnapshot, fallbackNow time.Time) map[string]any { - if snapshot == nil { - return nil - } - - baseTime := codexSnapshotBaseTime(snapshot, fallbackNow) - updates := make(map[string]any) - - // 保存原始 primary/secondary 字段,便于排查问题 - if snapshot.PrimaryUsedPercent != nil { - updates["codex_primary_used_percent"] = *snapshot.PrimaryUsedPercent - } - if snapshot.PrimaryResetAfterSeconds != nil { - updates["codex_primary_reset_after_seconds"] = *snapshot.PrimaryResetAfterSeconds - } - if snapshot.PrimaryWindowMinutes != nil { - updates["codex_primary_window_minutes"] = *snapshot.PrimaryWindowMinutes - } - if snapshot.SecondaryUsedPercent != nil { - updates["codex_secondary_used_percent"] = *snapshot.SecondaryUsedPercent - } - if snapshot.SecondaryResetAfterSeconds != nil { - updates["codex_secondary_reset_after_seconds"] = *snapshot.SecondaryResetAfterSeconds - } - if snapshot.SecondaryWindowMinutes != nil { - updates["codex_secondary_window_minutes"] = *snapshot.SecondaryWindowMinutes - } - if snapshot.PrimaryOverSecondaryPercent != nil { - updates["codex_primary_over_secondary_percent"] = *snapshot.PrimaryOverSecondaryPercent - } - updates["codex_usage_updated_at"] = baseTime.Format(time.RFC3339) - - // 归一化到 5h/7d 规范字段 - if normalized := snapshot.Normalize(); normalized != nil { - if normalized.Used5hPercent != nil { - updates["codex_5h_used_percent"] = *normalized.Used5hPercent - } - if normalized.Reset5hSeconds != nil { - updates["codex_5h_reset_after_seconds"] = *normalized.Reset5hSeconds - } - if normalized.Window5hMinutes != nil { - updates["codex_5h_window_minutes"] = *normalized.Window5hMinutes - } - if normalized.Used7dPercent != nil { - updates["codex_7d_used_percent"] = *normalized.Used7dPercent - } - if normalized.Reset7dSeconds != nil { - updates["codex_7d_reset_after_seconds"] = *normalized.Reset7dSeconds - } - if normalized.Window7dMinutes != nil { - updates["codex_7d_window_minutes"] = *normalized.Window7dMinutes - } - if reset5hAt := codexResetAtRFC3339(baseTime, normalized.Reset5hSeconds); reset5hAt != nil { - updates["codex_5h_reset_at"] = *reset5hAt - } - if reset7dAt := codexResetAtRFC3339(baseTime, normalized.Reset7dSeconds); reset7dAt != nil { - updates["codex_7d_reset_at"] = *reset7dAt - } - } - - return updates -} - -// updateCodexUsageSnapshot saves the Codex usage snapshot to account's Extra field -// updateCodexUsageSnapshot 把 /responses 的 x-codex-* 全局头快照写入账号 codex_* Extra。 -// ⚠️ 调用方必须排除 spark 影子账号(account.IsShadow()):影子的 codex_* 仅由 QueryUsage -// (/wham/usage bengalfox 道)更新,不能被全局头口径污染(外审第7轮 P1)。本函数仅持 accountID, -// 无法在此自检影子,故守卫前置到各调用点。 -func (s *OpenAIGatewayService) updateCodexUsageSnapshot(ctx context.Context, accountID int64, snapshot *OpenAICodexUsageSnapshot) { - if snapshot == nil { - return - } - if s == nil || s.accountRepo == nil { - return - } - - now := time.Now() - updates := buildCodexUsageExtraUpdates(snapshot, now) - if len(updates) == 0 { - return - } - if !s.getCodexSnapshotThrottle().Allow(accountID, now) { - return - } - - go func() { - updateCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - _ = s.accountRepo.UpdateExtra(updateCtx, accountID, updates) - }() -} - -func (s *OpenAIGatewayService) UpdateCodexUsageSnapshotFromHeaders(ctx context.Context, accountID int64, headers http.Header) { - if accountID <= 0 || headers == nil { - return - } - if snapshot := ParseCodexRateLimitHeaders(headers); snapshot != nil { - s.updateCodexUsageSnapshot(ctx, accountID, snapshot) - } -} - func getOpenAIReasoningEffortFromReqBody(reqBody map[string]any) (value string, present bool) { if reqBody == nil { return "", false diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go new file mode 100644 index 0000000000..58b6029ffa --- /dev/null +++ b/backend/internal/service/openai_gateway_usage.go @@ -0,0 +1,658 @@ +package service + +// 本文件由 openai_gateway_service.go 纯移动拆分而来:用量记录、计费成本计算与 +// Codex 用量快照。仅做代码搬迁,无任何行为变更。 + +import ( + "context" + "errors" + "fmt" + "net/http" + "strconv" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "go.uber.org/zap" +) + +// OpenAIRecordUsageInput input for recording usage +type OpenAIRecordUsageInput struct { + Result *OpenAIForwardResult + APIKey *APIKey + User *User + Account *Account + Subscription *UserSubscription + InboundEndpoint string + UpstreamEndpoint string + UserAgent string // 请求的 User-Agent + IPAddress string // 请求的客户端 IP 地址 + RequestPayloadHash string + APIKeyService APIKeyQuotaUpdater + QuotaPlatform string // user×platform quota platform resolved by the handler before async billing. + // CyberBlocked 为 true 时把该用量行标记为 cyber(request_type=cyber),计费逻辑不变。 + CyberBlocked bool + ChannelUsageFields +} + +// CyberPolicyUsageInput 是 cyber 拒绝、未走正常 RecordUsage 的请求记录用量的入参。 +// 用量按上游真实 token 计费,与 WS cyber 及正常请求口径一致(InputTokens/OutputTokens +// 取自上游 response.failed 报告的 usage,即 mark.UpstreamInTok/OutTok)。 +type CyberPolicyUsageInput struct { + APIKey *APIKey + Account *Account + Subscription *UserSubscription + RequestID string + Model string + Stream bool + InputTokens int + OutputTokens int + // 渠道归因与请求级 meta,使 cyber 计费行与正常 RecordUsage 行口径一致 + // (否则 cyber 行 channel_id 等为空,渠道维度统计会遗漏 cyber 命中)。 + InboundEndpoint string + UpstreamEndpoint string + UserAgent string + IPAddress string + RequestPayloadHash string + APIKeyService APIKeyQuotaUpdater + ChannelUsageFields +} + +// RecordCyberPolicyUsageLog 为被上游 cyber_policy 拒绝、未走正常 RecordUsage 的请求 +// (HTTP forward 返回错误路径)记录用量并按上游真实 token 计费,使其与 WS cyber 路径、 +// 与正常请求的计费口径统一(不再是 tokens=0 免费行)。token 取自上游 response.failed +// 报告的 usage(非流式直接拒通常为 0,cost 随之为 0)。复用 RecordUsage 完成成本计算、 +// 扣费与用量行写入(request_type=cyber 由 CyberBlocked 置位)。仅 forward 返回错误的 +// 路径由 handler 调用,避免与成功路径的正常 RecordUsage 重复。 +func (s *OpenAIGatewayService) RecordCyberPolicyUsageLog(ctx context.Context, in CyberPolicyUsageInput) { + if s == nil || in.APIKey == nil || in.APIKey.User == nil || in.Account == nil || strings.TrimSpace(in.Model) == "" { + return + } + result := &OpenAIForwardResult{ + RequestID: in.RequestID, + Model: in.Model, + Stream: in.Stream, + Usage: OpenAIUsage{ + InputTokens: in.InputTokens, + OutputTokens: in.OutputTokens, + }, + } + if err := s.RecordUsage(ctx, &OpenAIRecordUsageInput{ + Result: result, + APIKey: in.APIKey, + User: in.APIKey.User, + Account: in.Account, + Subscription: in.Subscription, + InboundEndpoint: in.InboundEndpoint, + UpstreamEndpoint: in.UpstreamEndpoint, + UserAgent: in.UserAgent, + IPAddress: in.IPAddress, + RequestPayloadHash: in.RequestPayloadHash, + APIKeyService: in.APIKeyService, + ChannelUsageFields: in.ChannelUsageFields, + CyberBlocked: true, + }); err != nil { + logger.LegacyPrintf("service.openai_gateway", "cyber usage record failed: request_id=%s err=%v", in.RequestID, err) + } +} + +// RecordUsage records usage and deducts balance +func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRecordUsageInput) error { + if input == nil { + return errors.New("openai usage input is nil") + } + result := input.Result + if result == nil { + return errors.New("openai usage result is nil") + } + if s.rateLimitService != nil && input.Account != nil && input.Account.Platform == PlatformOpenAI { + s.rateLimitService.ResetOpenAI403Counter(ctx, input.Account.ID) + } + + apiKey := input.APIKey + user := input.User + account := input.Account + subscription := input.Subscription + ApplyOpenAIImageBillingResolution(result) + + // 计算实际的新输入token(减去缓存读取的token) + // 因为 input_tokens 包含了 cache_read_tokens,而缓存读取的token不应按输入价格计费 + actualInputTokens := result.Usage.InputTokens - result.Usage.CacheReadInputTokens + if actualInputTokens < 0 { + actualInputTokens = 0 + } + + // Calculate cost + tokens := UsageTokens{ + InputTokens: actualInputTokens, + ImageInputTokens: result.Usage.ImageInputTokens, + OutputTokens: result.Usage.OutputTokens, + CacheCreationTokens: result.Usage.CacheCreationInputTokens, + CacheReadTokens: result.Usage.CacheReadInputTokens, + ImageOutputTokens: result.Usage.ImageOutputTokens, + } + + // Get rate multiplier + multiplier := 1.0 + if s.cfg != nil { + multiplier = s.cfg.Default.RateMultiplier + } + if apiKey.GroupID != nil && apiKey.Group != nil { + resolver := s.userGroupRateResolver + if resolver == nil { + resolver = newUserGroupRateResolver(nil, nil, resolveUserGroupRateCacheTTL(s.cfg), nil, "service.openai_gateway") + } + multiplier = resolver.Resolve(ctx, user.ID, *apiKey.GroupID, apiKey.Group.RateMultiplier) + } + // token 倍率叠加高峰因子(token 计费含图片 token,图片按次倍率不受影响)。高峰因子按请求时刻现算, + // 不并入上面的 Resolve,以免污染 user:group 倍率缓存。 + multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, multiplier, timezone.Now()) + + var cost *CostBreakdown + var err error + billingModel := forwardResultBillingModel(result.Model, result.UpstreamModel) + if result.BillingModel != "" { + billingModel = strings.TrimSpace(result.BillingModel) + } + if input.BillingModelSource == BillingModelSourceChannelMapped && input.ChannelMappedModel != "" && input.ChannelMappedModel != input.OriginalModel { + billingModel = input.ChannelMappedModel + } + if input.BillingModelSource == BillingModelSourceRequested && input.OriginalModel != "" { + billingModel = input.OriginalModel + } + billingModels := usageBillingModelCandidates( + billingModel, + result.BillingModel, + input.ChannelMappedModel, + input.OriginalModel, + result.UpstreamModel, + result.Model, + ) + serviceTier := "" + if result.ServiceTier != nil { + serviceTier = strings.TrimSpace(*result.ServiceTier) + } + cost, err = s.calculateOpenAIRecordUsageCost(ctx, result, apiKey, billingModels, multiplier, imageMultiplier, tokens, serviceTier) + if err != nil { + if !isUsagePricingUnavailableError(err) { + return err + } + logger.L().With( + zap.String("component", "service.openai_gateway"), + zap.Strings("billing_models", billingModels), + zap.String("requested_model", input.OriginalModel), + zap.String("mapped_model", input.ChannelMappedModel), + zap.String("upstream_model", result.UpstreamModel), + zap.Int64("api_key_id", apiKey.ID), + zap.Int64("account_id", account.ID), + ).Warn("openai_usage.pricing_missing_record_zero_cost", zap.Error(err)) + cost = &CostBreakdown{BillingMode: string(BillingModeToken)} + } + + // Determine billing type + isSubscriptionBilling := subscription != nil && apiKey.Group != nil && apiKey.Group.IsSubscriptionType() + billingType := BillingTypeBalance + if isSubscriptionBilling { + billingType = BillingTypeSubscription + } + + // Create usage log + durationMs := int(result.Duration.Milliseconds()) + accountRateMultiplier := account.BillingRateMultiplier() + requestID := resolveUsageBillingRequestID(ctx, result.RequestID) + if result.OpenAIWSMode { + if upstreamRequestID := strings.TrimSpace(result.RequestID); upstreamRequestID != "" { + requestID = upstreamRequestID + } + } + + // 确定 RequestedModel(渠道映射前的原始模型) + requestedModel := result.Model + if input.OriginalModel != "" { + requestedModel = input.OriginalModel + } + + usageLog := &UsageLog{ + UserID: user.ID, + APIKeyID: apiKey.ID, + AccountID: account.ID, + RequestID: requestID, + Model: result.Model, + RequestedModel: requestedModel, + UpstreamModel: optionalNonEqualStringPtr(result.UpstreamModel, result.Model), + ServiceTier: result.ServiceTier, + ReasoningEffort: result.ReasoningEffort, + InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint), + UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint), + InputTokens: actualInputTokens, + OutputTokens: result.Usage.OutputTokens, + CacheCreationTokens: result.Usage.CacheCreationInputTokens, + CacheReadTokens: result.Usage.CacheReadInputTokens, + ImageOutputTokens: result.Usage.ImageOutputTokens, + ImageCount: result.ImageCount, + ImageSize: optionalTrimmedStringPtr(result.ImageSize), + ImageInputSize: optionalTrimmedStringPtr(result.ImageInputSize), + ImageOutputSize: optionalTrimmedStringPtr(result.ImageOutputSize), + ImageSizeSource: optionalTrimmedStringPtr(result.ImageSizeSource), + ImageSizeBreakdown: result.ImageSizeBreakdown, + } + if cost != nil { + usageLog.InputCost = cost.InputCost + usageLog.OutputCost = cost.OutputCost + usageLog.ImageOutputCost = cost.ImageOutputCost + usageLog.CacheCreationCost = cost.CacheCreationCost + usageLog.CacheReadCost = cost.CacheReadCost + usageLog.TotalCost = cost.TotalCost + usageLog.ActualCost = cost.ActualCost + } + if result.ImageCount > 0 && (cost == nil || cost.BillingMode != string(BillingModeToken)) { + usageLog.RateMultiplier = imageMultiplier + } else { + usageLog.RateMultiplier = multiplier + } + usageLog.AccountRateMultiplier = &accountRateMultiplier + usageLog.BillingType = billingType + usageLog.Stream = result.Stream + if input.CyberBlocked { + usageLog.RequestType = RequestTypeCyberBlocked + } + usageLog.OpenAIWSMode = result.OpenAIWSMode + usageLog.DurationMs = &durationMs + usageLog.FirstTokenMs = result.FirstTokenMs + usageLog.CreatedAt = time.Now() + // 设置渠道信息 + usageLog.ChannelID = optionalInt64Ptr(input.ChannelID) + usageLog.ModelMappingChain = optionalTrimmedStringPtr(input.ModelMappingChain) + // 设置计费模式 + if cost != nil && cost.BillingMode != "" { + billingMode := cost.BillingMode + usageLog.BillingMode = &billingMode + } else if result.ImageCount > 0 { + billingMode := string(BillingModeImage) + usageLog.BillingMode = &billingMode + } else { + billingMode := string(BillingModeToken) + usageLog.BillingMode = &billingMode + } + // 添加 UserAgent + if input.UserAgent != "" { + usageLog.UserAgent = &input.UserAgent + } + + // 添加 IPAddress + if input.IPAddress != "" { + usageLog.IPAddress = &input.IPAddress + } + + if apiKey.GroupID != nil { + usageLog.GroupID = apiKey.GroupID + } + if subscription != nil { + usageLog.SubscriptionID = &subscription.ID + } + + // 计算账号统计定价费用(使用最终上游模型匹配自定义规则) + if apiKey.GroupID != nil { + applyAccountStatsCost(ctx, usageLog, s.channelService, s.billingService, + account.ID, *apiKey.GroupID, result.UpstreamModel, result.Model, + tokens, cost.TotalCost, + ) + } + + if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { + writeUsageLogBestEffort(ctx, s.usageLogRepo, usageLog, "service.openai_gateway") + logger.LegacyPrintf("service.openai_gateway", "[SIMPLE MODE] Usage recorded (not billed): user=%d, tokens=%d", usageLog.UserID, usageLog.TotalTokens()) + s.deferredService.ScheduleLastUsedUpdate(account.ID) + return nil + } + + // Async usage billing runs outside the original request context, so it + // cannot recover ForcePlatform there. Fall back for internal/test callers. + quotaPlatform := input.QuotaPlatform + if quotaPlatform == "" { + quotaPlatform = PlatformFromAPIKey(apiKey) + } + + billingErr := func() error { + _, err := applyUsageBilling(ctx, requestID, usageLog, &postUsageBillingParams{ + Cost: cost, + User: user, + APIKey: apiKey, + Account: account, + Subscription: subscription, + RequestPayloadHash: resolveUsageBillingPayloadFingerprint(ctx, input.RequestPayloadHash), + IsSubscriptionBill: isSubscriptionBilling, + AccountRateMultiplier: accountRateMultiplier, + APIKeyService: input.APIKeyService, + Platform: quotaPlatform, + }, s.billingDeps(), s.usageBillingRepo) + return err + }() + + if billingErr != nil { + return billingErr + } + writeUsageLogBestEffort(ctx, s.usageLogRepo, usageLog, "service.openai_gateway") + + return nil +} + +func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( + ctx context.Context, + result *OpenAIForwardResult, + apiKey *APIKey, + billingModels []string, + multiplier float64, + imageMultiplier float64, + tokens UsageTokens, + serviceTier string, +) (*CostBreakdown, error) { + billingModel := firstUsageBillingModel(billingModels) + if result != nil && result.ImageCount > 0 { + // 渠道定价为 token 计费时走 token 路径,否则走图片计费 + if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved == nil || resolved.Mode != BillingModeToken { + return s.calculateOpenAIImageCost(ctx, billingModel, apiKey, result, imageMultiplier), nil + } + } + if len(billingModels) == 0 || billingModel == "" { + return nil, errors.New("openai usage billing model is empty") + } + var lastErr error + for _, candidate := range billingModels { + candidate = strings.TrimSpace(candidate) + if candidate == "" { + continue + } + cost, err := s.calculateOpenAIRecordUsageTokenCost(ctx, apiKey, candidate, multiplier, tokens, serviceTier) + if err == nil { + return cost, nil + } + lastErr = err + } + if lastErr == nil { + lastErr = errors.New("no non-empty billing model candidates") + } + return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr) +} + +func isUsagePricingUnavailableError(err error) bool { + if err == nil { + return false + } + if errors.Is(err, ErrModelPricingUnavailable) { + return true + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "no pricing available") || strings.Contains(msg, "pricing not found") +} + +func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost( + ctx context.Context, + apiKey *APIKey, + billingModel string, + multiplier float64, + tokens UsageTokens, + serviceTier string, +) (*CostBreakdown, error) { + if s.resolver != nil && apiKey.Group != nil { + gid := apiKey.Group.ID + return s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, + Model: billingModel, + GroupID: &gid, + Tokens: tokens, + RequestCount: 1, + RateMultiplier: multiplier, + ServiceTier: serviceTier, + Resolver: s.resolver, + }) + } + return s.billingService.CalculateCostWithServiceTier(billingModel, tokens, multiplier, serviceTier) +} + +func (s *OpenAIGatewayService) calculateOpenAIImageCost( + ctx context.Context, + billingModel string, + apiKey *APIKey, + result *OpenAIForwardResult, + multiplier float64, +) *CostBreakdown { + sizeTier := NormalizeImageBillingTierOrDefault(result.ImageSize) + if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved != nil && + (resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage) { + gid := apiKey.Group.ID + cost, err := s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, + Model: billingModel, + GroupID: &gid, + RequestCount: result.ImageCount, + SizeTier: sizeTier, + RateMultiplier: multiplier, + Resolver: s.resolver, + Resolved: resolved, + }) + if err == nil { + return cost + } + logger.LegacyPrintf("service.openai_gateway", "Calculate image channel cost failed: %v", err) + } + + var groupConfig *ImagePriceConfig + if apiKey != nil && apiKey.Group != nil { + groupConfig = &ImagePriceConfig{ + Price1K: apiKey.Group.ImagePrice1K, + Price2K: apiKey.Group.ImagePrice2K, + Price4K: apiKey.Group.ImagePrice4K, + } + } + return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier) +} + +func (s *OpenAIGatewayService) resolveOpenAIChannelPricing(ctx context.Context, billingModel string, apiKey *APIKey) *ResolvedPricing { + if s.resolver == nil || apiKey == nil || apiKey.Group == nil { + return nil + } + gid := apiKey.Group.ID + resolved := s.resolver.Resolve(ctx, PricingInput{Model: billingModel, GroupID: &gid}) + if resolved.Source == PricingSourceChannel { + return resolved + } + return nil +} + +// ParseCodexRateLimitHeaders extracts Codex usage limits from response headers. +// Exported for use in ratelimit_service when handling OpenAI 429 responses. +func ParseCodexRateLimitHeaders(headers http.Header) *OpenAICodexUsageSnapshot { + snapshot := &OpenAICodexUsageSnapshot{} + hasData := false + + // Helper to parse float64 from header + parseFloat := func(key string) *float64 { + if v := headers.Get(key); v != "" { + if f, err := strconv.ParseFloat(v, 64); err == nil { + return &f + } + } + return nil + } + + // Helper to parse int from header + parseInt := func(key string) *int { + if v := headers.Get(key); v != "" { + if i, err := strconv.Atoi(v); err == nil { + return &i + } + } + return nil + } + + // Primary (weekly) limits + if v := parseFloat("x-codex-primary-used-percent"); v != nil { + snapshot.PrimaryUsedPercent = v + hasData = true + } + if v := parseInt("x-codex-primary-reset-after-seconds"); v != nil { + snapshot.PrimaryResetAfterSeconds = v + hasData = true + } + if v := parseInt("x-codex-primary-window-minutes"); v != nil { + snapshot.PrimaryWindowMinutes = v + hasData = true + } + + // Secondary (5h) limits + if v := parseFloat("x-codex-secondary-used-percent"); v != nil { + snapshot.SecondaryUsedPercent = v + hasData = true + } + if v := parseInt("x-codex-secondary-reset-after-seconds"); v != nil { + snapshot.SecondaryResetAfterSeconds = v + hasData = true + } + if v := parseInt("x-codex-secondary-window-minutes"); v != nil { + snapshot.SecondaryWindowMinutes = v + hasData = true + } + + // Overflow ratio + if v := parseFloat("x-codex-primary-over-secondary-limit-percent"); v != nil { + snapshot.PrimaryOverSecondaryPercent = v + hasData = true + } + + if !hasData { + return nil + } + + snapshot.UpdatedAt = time.Now().Format(time.RFC3339) + return snapshot +} + +func codexSnapshotBaseTime(snapshot *OpenAICodexUsageSnapshot, fallback time.Time) time.Time { + if snapshot == nil { + return fallback + } + if snapshot.UpdatedAt == "" { + return fallback + } + base, err := time.Parse(time.RFC3339, snapshot.UpdatedAt) + if err != nil { + return fallback + } + return base +} + +func codexResetAtRFC3339(base time.Time, resetAfterSeconds *int) *string { + if resetAfterSeconds == nil { + return nil + } + sec := *resetAfterSeconds + if sec < 0 { + sec = 0 + } + resetAt := base.Add(time.Duration(sec) * time.Second).Format(time.RFC3339) + return &resetAt +} + +func buildCodexUsageExtraUpdates(snapshot *OpenAICodexUsageSnapshot, fallbackNow time.Time) map[string]any { + if snapshot == nil { + return nil + } + + baseTime := codexSnapshotBaseTime(snapshot, fallbackNow) + updates := make(map[string]any) + + // 保存原始 primary/secondary 字段,便于排查问题 + if snapshot.PrimaryUsedPercent != nil { + updates["codex_primary_used_percent"] = *snapshot.PrimaryUsedPercent + } + if snapshot.PrimaryResetAfterSeconds != nil { + updates["codex_primary_reset_after_seconds"] = *snapshot.PrimaryResetAfterSeconds + } + if snapshot.PrimaryWindowMinutes != nil { + updates["codex_primary_window_minutes"] = *snapshot.PrimaryWindowMinutes + } + if snapshot.SecondaryUsedPercent != nil { + updates["codex_secondary_used_percent"] = *snapshot.SecondaryUsedPercent + } + if snapshot.SecondaryResetAfterSeconds != nil { + updates["codex_secondary_reset_after_seconds"] = *snapshot.SecondaryResetAfterSeconds + } + if snapshot.SecondaryWindowMinutes != nil { + updates["codex_secondary_window_minutes"] = *snapshot.SecondaryWindowMinutes + } + if snapshot.PrimaryOverSecondaryPercent != nil { + updates["codex_primary_over_secondary_percent"] = *snapshot.PrimaryOverSecondaryPercent + } + updates["codex_usage_updated_at"] = baseTime.Format(time.RFC3339) + + // 归一化到 5h/7d 规范字段 + if normalized := snapshot.Normalize(); normalized != nil { + if normalized.Used5hPercent != nil { + updates["codex_5h_used_percent"] = *normalized.Used5hPercent + } + if normalized.Reset5hSeconds != nil { + updates["codex_5h_reset_after_seconds"] = *normalized.Reset5hSeconds + } + if normalized.Window5hMinutes != nil { + updates["codex_5h_window_minutes"] = *normalized.Window5hMinutes + } + if normalized.Used7dPercent != nil { + updates["codex_7d_used_percent"] = *normalized.Used7dPercent + } + if normalized.Reset7dSeconds != nil { + updates["codex_7d_reset_after_seconds"] = *normalized.Reset7dSeconds + } + if normalized.Window7dMinutes != nil { + updates["codex_7d_window_minutes"] = *normalized.Window7dMinutes + } + if reset5hAt := codexResetAtRFC3339(baseTime, normalized.Reset5hSeconds); reset5hAt != nil { + updates["codex_5h_reset_at"] = *reset5hAt + } + if reset7dAt := codexResetAtRFC3339(baseTime, normalized.Reset7dSeconds); reset7dAt != nil { + updates["codex_7d_reset_at"] = *reset7dAt + } + } + + return updates +} + +// updateCodexUsageSnapshot saves the Codex usage snapshot to account's Extra field +// updateCodexUsageSnapshot 把 /responses 的 x-codex-* 全局头快照写入账号 codex_* Extra。 +// ⚠️ 调用方必须排除 spark 影子账号(account.IsShadow()):影子的 codex_* 仅由 QueryUsage +// (/wham/usage bengalfox 道)更新,不能被全局头口径污染(外审第7轮 P1)。本函数仅持 accountID, +// 无法在此自检影子,故守卫前置到各调用点。 +func (s *OpenAIGatewayService) updateCodexUsageSnapshot(ctx context.Context, accountID int64, snapshot *OpenAICodexUsageSnapshot) { + if snapshot == nil { + return + } + if s == nil || s.accountRepo == nil { + return + } + + now := time.Now() + updates := buildCodexUsageExtraUpdates(snapshot, now) + if len(updates) == 0 { + return + } + if !s.getCodexSnapshotThrottle().Allow(accountID, now) { + return + } + + go func() { + updateCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _ = s.accountRepo.UpdateExtra(updateCtx, accountID, updates) + }() +} + +func (s *OpenAIGatewayService) UpdateCodexUsageSnapshotFromHeaders(ctx context.Context, accountID int64, headers http.Header) { + if accountID <= 0 || headers == nil { + return + } + if snapshot := ParseCodexRateLimitHeaders(headers); snapshot != nil { + s.updateCodexUsageSnapshot(ctx, accountID, snapshot) + } +}