Files
sub2api/backend/internal/service/openai_gateway_cc_pipeline.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

350 lines
13 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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"
)
// 本文件收敛三个 CCChat Completionsforwarder 之间重复的 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 {
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody)
tempUnscheduled := false
if c != nil && account != nil && account.Platform != PlatformGrok && !shouldFailover && !IsResponseCommitted(c) && s.rateLimitService != nil {
tempUnscheduled = s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, respBody, upstreamModel) == ErrorPolicyTempUnscheduled
shouldFailover = tempUnscheduled
}
if account != nil && account.Platform == PlatformGrok {
shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody)
}
if account != nil && account.Platform == PlatformGrok {
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
}
if !shouldFailover {
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,
})
shouldDisable := tempUnscheduled
if account.Platform != PlatformGrok && !tempUnscheduled {
shouldDisable = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel)
}
return newOpenAIUpstreamFailoverError(
resp.StatusCode,
resp.Header,
respBody,
upstreamMsg,
!shouldDisable && account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
)
}
// openAIChatCompletionsTargetURL 解析账号的(非 GrokChat 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 = strings.TrimSpace(account.GetOpenAIProtocolAPIKey())
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,
grokCacheIdentity 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)
}
if account.Platform == PlatformGrok {
if account.IsGrokOAuth() {
applyGrokCLIHeaders(upstreamReq.Header)
}
applyGrokCacheHeaders(upstreamReq.Header, grokCacheIdentity)
}
// 账号级请求头覆写:放在所有内置默认头(含 Grok CLI 身份头)之后应用,
// 使配置值获得除共享传输层强制头之外的最高优先级。
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,
},
})
}