Files
sub2api/backend/internal/service/openai_gateway_cc_pipeline.go
T

350 lines
13 KiB
Go
Raw Normal View History

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,
},
})
}