package service import ( "bufio" "context" "encoding/json" "errors" "fmt" "net/http" "strings" "sync/atomic" "time" "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/gin-gonic/gin" ) type antigravityStreamResult struct { usage *ClaudeUsage firstTokenMs *int clientDisconnect bool // 客户端是否在流式传输过程中断开 } func (s *AntigravityGatewayService) observeAntigravityGeminiSSELine(c *gin.Context, line string) { observer := upstreamResponseModelObserverFromContext(c) if observer == nil { observer = beginUpstreamResponseModelObservation(c) } trimmed := strings.TrimSpace(line) if !strings.HasPrefix(trimmed, "data:") { return } payload := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) if payload == "" || payload == "[DONE]" { return } // Observe the original payload: ObserveGemini supports both the v1internal // wrapper and direct Gemini response shapes. The main stream handler will // unwrap the same line for business processing, so unwrapping here would be // duplicate work on every SSE event. observer.ObserveGemini([]byte(payload)) } // antigravityClientWriter 封装流式响应的客户端写入,自动检测断开并标记。 // 断开后所有写入操作变为 no-op,调用方通过 Disconnected() 判断是否继续 drain 上游。 type antigravityClientWriter struct { w gin.ResponseWriter flusher http.Flusher disconnected bool prefix string // 日志前缀,标识来源方法 beforeFirstWrite func() } func newAntigravityClientWriter(w gin.ResponseWriter, flusher http.Flusher, prefix string) *antigravityClientWriter { return &antigravityClientWriter{w: w, flusher: flusher, prefix: prefix} } // Write 写入数据到客户端,写入失败时标记断开并返回 false func (cw *antigravityClientWriter) Write(p []byte) bool { if cw.disconnected { return false } cw.prepareFirstWrite() if _, err := cw.w.Write(p); err != nil { cw.markDisconnected() return false } cw.flusher.Flush() return true } // Fprintf 格式化写入数据到客户端,写入失败时标记断开并返回 false func (cw *antigravityClientWriter) Fprintf(format string, args ...any) bool { if cw.disconnected { return false } cw.prepareFirstWrite() if _, err := fmt.Fprintf(cw.w, format, args...); err != nil { cw.markDisconnected() return false } cw.flusher.Flush() return true } func (cw *antigravityClientWriter) Disconnected() bool { return cw.disconnected } func (cw *antigravityClientWriter) prepareFirstWrite() { if cw.beforeFirstWrite == nil { return } prepare := cw.beforeFirstWrite cw.beforeFirstWrite = nil prepare() } func (cw *antigravityClientWriter) markDisconnected() { cw.disconnected = true logger.LegacyPrintf("service.antigravity_gateway", "Client disconnected during streaming (%s), continuing to drain upstream for billing", cw.prefix) } // handleStreamReadError 处理上游读取错误的通用逻辑。 // 返回 (clientDisconnect, handled):handled=true 表示错误已处理,调用方应返回已收集的 usage。 func handleStreamReadError(err error, clientDisconnected bool, prefix string) (disconnect bool, handled bool) { if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { logger.LegacyPrintf("service.antigravity_gateway", "Context canceled during streaming (%s), returning collected usage", prefix) return true, true } if clientDisconnected { logger.LegacyPrintf("service.antigravity_gateway", "Upstream read error after client disconnect (%s): %v, returning collected usage", prefix, err) return true, true } return false, false } func (s *AntigravityGatewayService) handleGeminiStreamingResponse(c *gin.Context, resp *http.Response, startTime time.Time) (*antigravityStreamResult, error) { if upstreamResponseModelObserverFromContext(c) == nil { beginUpstreamResponseModelObservation(c) } c.Status(resp.StatusCode) c.Header("Cache-Control", "no-cache") c.Header("Connection", "keep-alive") c.Header("X-Accel-Buffering", "no") contentType := resp.Header.Get("Content-Type") if contentType == "" { contentType = "text/event-stream; charset=utf-8" } c.Header("Content-Type", contentType) flusher, ok := c.Writer.(http.Flusher) if !ok { return nil, errors.New("streaming not supported") } // 使用 Scanner 并限制单行大小,避免 ReadString 无上限导致 OOM scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { maxLineSize = s.settingService.cfg.Gateway.MaxLineSize } scanBuf := getSSEScannerBuf64K() scanner.Buffer(scanBuf[:0], maxLineSize) usage := &ClaudeUsage{} var firstTokenMs *int type scanEvent struct { line string err error } // 独立 goroutine 读取上游,避免读取阻塞影响超时处理 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.settingService.cfg != nil && s.settingService.cfg.Gateway.StreamDataIntervalTimeout > 0 { streamInterval = time.Duration(s.settingService.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 } // 下游 keepalive:防止代理/Cloudflare Tunnel 因连接空闲而断开 keepaliveInterval := time.Duration(0) if s.settingService.cfg != nil && s.settingService.cfg.Gateway.StreamKeepaliveInterval > 0 { keepaliveInterval = time.Duration(s.settingService.cfg.Gateway.StreamKeepaliveInterval) * time.Second } var keepaliveTicker *time.Ticker if keepaliveInterval > 0 { keepaliveTicker = time.NewTicker(keepaliveInterval) defer keepaliveTicker.Stop() } var keepaliveCh <-chan time.Time if keepaliveTicker != nil { keepaliveCh = keepaliveTicker.C } lastDataAt := time.Now() cw := newAntigravityClientWriter(c.Writer, flusher, "antigravity gemini") // 仅发送一次错误事件,避免多次写入导致协议混乱 errorEventSent := false sendErrorEvent := func(reason string) { if errorEventSent || cw.Disconnected() { return } errorEventSent = true _, _ = fmt.Fprintf(c.Writer, "event: error\ndata: {\"error\":\"%s\"}\n\n", reason) flusher.Flush() } for { select { case ev, ok := <-events: if !ok { return &antigravityStreamResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: cw.Disconnected()}, nil } if ev.err != nil { if disconnect, handled := handleStreamReadError(ev.err, cw.Disconnected(), "antigravity gemini"); handled { return &antigravityStreamResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: disconnect}, nil } if errors.Is(ev.err, bufio.ErrTooLong) { logger.LegacyPrintf("service.antigravity_gateway", "SSE line too long (antigravity): max_size=%d error=%v", maxLineSize, ev.err) sendErrorEvent("response_too_large") return &antigravityStreamResult{usage: usage, firstTokenMs: firstTokenMs}, ev.err } sendErrorEvent("stream_read_error") return nil, ev.err } lastDataAt = time.Now() line := ev.line s.observeAntigravityGeminiSSELine(c, line) trimmed := strings.TrimRight(line, "\r\n") if strings.HasPrefix(trimmed, "data:") { payload := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) if payload == "" || payload == "[DONE]" { cw.Fprintf("%s\n", line) continue } // 解包 v1internal 响应 inner, parseErr := s.unwrapV1InternalResponse([]byte(payload)) if parseErr == nil && inner != nil { payload = string(inner) } // 解析 usage if u := extractGeminiUsage(inner); u != nil { usage = u } var parsed map[string]any if json.Unmarshal(inner, &parsed) == nil { // Check for MALFORMED_FUNCTION_CALL if candidates, ok := parsed["candidates"].([]any); ok && len(candidates) > 0 { if cand, ok := candidates[0].(map[string]any); ok { if fr, ok := cand["finishReason"].(string); ok && fr == "MALFORMED_FUNCTION_CALL" { logger.LegacyPrintf("service.antigravity_gateway", "[Antigravity] MALFORMED_FUNCTION_CALL detected in forward stream") if content, ok := cand["content"]; ok { if b, err := json.Marshal(content); err == nil { logger.LegacyPrintf("service.antigravity_gateway", "[Antigravity] Malformed content: %s", string(b)) } } } } } } if firstTokenMs == nil { ms := int(time.Since(startTime).Milliseconds()) firstTokenMs = &ms } cw.Fprintf("data: %s\n\n", payload) continue } cw.Fprintf("%s\n", line) case <-intervalCh: lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) if time.Since(lastRead) < streamInterval { continue } if cw.Disconnected() { logger.LegacyPrintf("service.antigravity_gateway", "Upstream timeout after client disconnect (antigravity gemini), returning collected usage") return &antigravityStreamResult{usage: usage, firstTokenMs: firstTokenMs, clientDisconnect: true}, nil } logger.LegacyPrintf("service.antigravity_gateway", "Stream data interval timeout (antigravity)") sendErrorEvent("stream_timeout") return &antigravityStreamResult{usage: usage, firstTokenMs: firstTokenMs}, fmt.Errorf("stream data interval timeout") case <-keepaliveCh: if cw.Disconnected() { continue } if time.Since(lastDataAt) < keepaliveInterval { continue } // SSE ping/keepalive:保持连接活跃防止 Cloudflare Tunnel 等代理断开 if !cw.Fprintf(":\n\n") { logger.LegacyPrintf("service.antigravity_gateway", "Client disconnected during keepalive ping (antigravity gemini), continuing to drain upstream for billing") continue } } } } // handleGeminiStreamToNonStreaming 读取上游流式响应,合并为非流式响应返回给客户端 // Gemini 流式响应是增量的,需要累积所有 chunk 的内容 func (s *AntigravityGatewayService) handleGeminiStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time) (*antigravityStreamResult, error) { if upstreamResponseModelObserverFromContext(c) == nil { beginUpstreamResponseModelObservation(c) } scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { maxLineSize = s.settingService.cfg.Gateway.MaxLineSize } scanBuf := getSSEScannerBuf64K() scanner.Buffer(scanBuf[:0], maxLineSize) usage := &ClaudeUsage{} var firstTokenMs *int var last map[string]any var lastWithParts map[string]any var collectedImageParts []map[string]any // 收集所有包含图片的 parts var collectedTextParts []string // 收集所有文本片段 type scanEvent struct { line string err error } // 独立 goroutine 读取上游,避免读取阻塞影响超时处理 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.settingService.cfg != nil && s.settingService.cfg.Gateway.StreamDataIntervalTimeout > 0 { streamInterval = time.Duration(s.settingService.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 } for { select { case ev, ok := <-events: if !ok { // 流结束,返回收集的响应 goto returnResponse } if ev.err != nil { if errors.Is(ev.err, bufio.ErrTooLong) { logger.LegacyPrintf("service.antigravity_gateway", "SSE line too long (antigravity non-stream): max_size=%d error=%v", maxLineSize, ev.err) } return nil, ev.err } line := ev.line s.observeAntigravityGeminiSSELine(c, line) trimmed := strings.TrimRight(line, "\r\n") if !strings.HasPrefix(trimmed, "data:") { continue } payload := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) if payload == "" || payload == "[DONE]" { continue } // 解包 v1internal 响应 inner, parseErr := s.unwrapV1InternalResponse([]byte(payload)) if parseErr != nil { continue } var parsed map[string]any if err := json.Unmarshal(inner, &parsed); err != nil { continue } // 记录首 token 时间 if firstTokenMs == nil { ms := int(time.Since(startTime).Milliseconds()) firstTokenMs = &ms } last = parsed // 提取 usage if u := extractGeminiUsage(inner); u != nil { usage = u } // Check for MALFORMED_FUNCTION_CALL if candidates, ok := parsed["candidates"].([]any); ok && len(candidates) > 0 { if cand, ok := candidates[0].(map[string]any); ok { if fr, ok := cand["finishReason"].(string); ok && fr == "MALFORMED_FUNCTION_CALL" { logger.LegacyPrintf("service.antigravity_gateway", "[Antigravity] MALFORMED_FUNCTION_CALL detected in forward non-stream collect") if content, ok := cand["content"]; ok { if b, err := json.Marshal(content); err == nil { logger.LegacyPrintf("service.antigravity_gateway", "[Antigravity] Malformed content: %s", string(b)) } } } } } // 保留最后一个有 parts 的响应 if parts := extractGeminiParts(parsed); len(parts) > 0 { lastWithParts = parsed // 收集包含图片和文本的 parts for _, part := range parts { if inlineData, ok := part["inlineData"].(map[string]any); ok { collectedImageParts = append(collectedImageParts, part) _ = inlineData // 避免 unused 警告 } if text, ok := part["text"].(string); ok && text != "" { collectedTextParts = append(collectedTextParts, text) } } } case <-intervalCh: lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) if time.Since(lastRead) < streamInterval { continue } logger.LegacyPrintf("service.antigravity_gateway", "Stream data interval timeout (antigravity non-stream)") return nil, fmt.Errorf("stream data interval timeout") } } returnResponse: // 选择最后一个有效响应 finalResponse := pickGeminiCollectResult(last, lastWithParts) // 处理空响应情况 — 触发同账号重试 + failover 切换账号 if last == nil && lastWithParts == nil { logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] warning: empty stream response (gemini non-stream), triggering failover") return nil, &UpstreamFailoverError{ StatusCode: http.StatusBadGateway, ResponseBody: []byte(`{"error":"empty stream response from upstream"}`), RetryableOnSameAccount: true, } } // 如果收集到了图片 parts,需要合并到最终响应中 if len(collectedImageParts) > 0 { finalResponse = mergeImagePartsToResponse(finalResponse, collectedImageParts) } // 如果收集到了文本,需要合并到最终响应中 if len(collectedTextParts) > 0 { finalResponse = mergeTextPartsToResponse(finalResponse, collectedTextParts) } respBody, err := json.Marshal(finalResponse) if err != nil { return nil, fmt.Errorf("failed to marshal response: %w", err) } c.Data(http.StatusOK, "application/json", respBody) return &antigravityStreamResult{usage: usage, firstTokenMs: firstTokenMs}, nil } // getOrCreateGeminiParts 获取 Gemini 响应的 parts 结构,返回深拷贝和更新回调 func getOrCreateGeminiParts(response map[string]any) (result map[string]any, existingParts []any, setParts func([]any)) { // 深拷贝 response result = make(map[string]any) for k, v := range response { result[k] = v } // 获取或创建 candidates candidates, ok := result["candidates"].([]any) if !ok || len(candidates) == 0 { candidates = []any{map[string]any{}} } // 获取第一个 candidate candidate, ok := candidates[0].(map[string]any) if !ok { candidate = make(map[string]any) candidates[0] = candidate } // 获取或创建 content content, ok := candidate["content"].(map[string]any) if !ok { content = map[string]any{"role": "model"} candidate["content"] = content } // 获取现有 parts existingParts, ok = content["parts"].([]any) if !ok { existingParts = []any{} } // 返回更新回调 setParts = func(newParts []any) { content["parts"] = newParts result["candidates"] = candidates } return result, existingParts, setParts } // mergeCollectedPartsToResponse 将收集的所有 parts 合并到 Gemini 响应中 // 这个函数会合并所有类型的 parts:text、thinking、functionCall、inlineData 等 // 保持原始顺序,只合并连续的普通 text parts func mergeCollectedPartsToResponse(response map[string]any, collectedParts []map[string]any) map[string]any { if len(collectedParts) == 0 { return response } result, _, setParts := getOrCreateGeminiParts(response) // 合并策略: // 1. 保持原始顺序 // 2. 连续的普通 text parts 合并为一个 // 3. thinking、functionCall、inlineData 等保持原样 var mergedParts []any var textBuffer strings.Builder flushTextBuffer := func() { if textBuffer.Len() > 0 { mergedParts = append(mergedParts, map[string]any{ "text": textBuffer.String(), }) textBuffer.Reset() } } for _, part := range collectedParts { // 检查是否是普通 text part if text, ok := part["text"].(string); ok { // 检查是否有 thought 标记 if thought, _ := part["thought"].(bool); thought { // thinking part,先刷新 text buffer,然后保留原样 flushTextBuffer() mergedParts = append(mergedParts, part) } else { // 普通 text,累积到 buffer _, _ = textBuffer.WriteString(text) } } else { // 非 text part(functionCall、inlineData 等),先刷新 text buffer,然后保留原样 flushTextBuffer() mergedParts = append(mergedParts, part) } } // 刷新剩余的 text flushTextBuffer() setParts(mergedParts) return result } // mergeImagePartsToResponse 将收集到的图片 parts 合并到 Gemini 响应中 func mergeImagePartsToResponse(response map[string]any, imageParts []map[string]any) map[string]any { if len(imageParts) == 0 { return response } result, existingParts, setParts := getOrCreateGeminiParts(response) // 检查现有 parts 中是否已经有图片 for _, p := range existingParts { if pm, ok := p.(map[string]any); ok { if _, hasInline := pm["inlineData"]; hasInline { return result // 已有图片,不重复添加 } } } // 添加收集到的图片 parts for _, imgPart := range imageParts { existingParts = append(existingParts, imgPart) } setParts(existingParts) return result } // mergeTextPartsToResponse 将收集到的文本合并到 Gemini 响应中 func mergeTextPartsToResponse(response map[string]any, textParts []string) map[string]any { if len(textParts) == 0 { return response } mergedText := strings.Join(textParts, "") result, existingParts, setParts := getOrCreateGeminiParts(response) // 查找并更新第一个 text part,或创建新的 newParts := make([]any, 0, len(existingParts)+1) textUpdated := false for _, p := range existingParts { pm, ok := p.(map[string]any) if !ok { newParts = append(newParts, p) continue } if _, hasText := pm["text"]; hasText && !textUpdated { // 用累积的文本替换 newPart := make(map[string]any) for k, v := range pm { newPart[k] = v } newPart["text"] = mergedText newParts = append(newParts, newPart) textUpdated = true } else { newParts = append(newParts, pm) } } if !textUpdated { newParts = append([]any{map[string]any{"text": mergedText}}, newParts...) } setParts(newParts) return result } func (s *AntigravityGatewayService) writeClaudeError(c *gin.Context, status int, errType, message string) error { MarkResponseCommitted(c) c.JSON(status, gin.H{ "type": "error", "error": gin.H{"type": errType, "message": message}, }) return fmt.Errorf("%s", message) } // WriteMappedClaudeError 导出版本,供 handler 层使用(如 fallback 错误处理) func (s *AntigravityGatewayService) WriteMappedClaudeError(c *gin.Context, account *Account, upstreamStatus int, upstreamRequestID string, body []byte) error { return s.writeMappedClaudeError(c, account, upstreamStatus, upstreamRequestID, body) } func (s *AntigravityGatewayService) writeMappedClaudeError(c *gin.Context, account *Account, upstreamStatus int, upstreamRequestID string, body []byte) error { MarkResponseCommitted(c) upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body)) upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) logBody, maxBytes := s.getLogConfig() upstreamDetail := s.getUpstreamErrorDetail(body) setOpsUpstreamError(c, upstreamStatus, upstreamMsg, upstreamDetail) appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ Platform: account.Platform, AccountID: account.ID, AccountName: account.Name, UpstreamStatusCode: upstreamStatus, UpstreamRequestID: upstreamRequestID, Kind: "http_error", Message: upstreamMsg, Detail: upstreamDetail, }) // 记录上游错误详情便于排障(可选:由配置控制;不回显到客户端) if logBody { logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] upstream_error status=%d body=%s", upstreamStatus, truncateForLog(body, maxBytes)) } // 检查错误透传规则 if ptStatus, ptErrType, ptErrMsg, matched := applyErrorPassthroughRule( c, account.Platform, upstreamStatus, body, 0, "", "", ); matched { c.JSON(ptStatus, gin.H{ "type": "error", "error": gin.H{"type": ptErrType, "message": ptErrMsg}, }) if upstreamMsg == "" { return fmt.Errorf("upstream error: %d", upstreamStatus) } return fmt.Errorf("upstream error: %d message=%s", upstreamStatus, upstreamMsg) } var statusCode int var errType, errMsg string switch upstreamStatus { case 400: statusCode = http.StatusBadRequest errType = "invalid_request_error" errMsg = getPassthroughOrDefault(upstreamMsg, "Invalid request") case 401: statusCode = http.StatusBadGateway errType = "authentication_error" errMsg = "Upstream authentication failed" case 403: statusCode = http.StatusBadGateway errType = "permission_error" errMsg = "Upstream access forbidden" case 429: statusCode = http.StatusTooManyRequests errType = "rate_limit_error" errMsg = "Upstream rate limit exceeded" case 529: statusCode = http.StatusServiceUnavailable errType = "overloaded_error" errMsg = "Upstream service overloaded" default: statusCode = http.StatusBadGateway errType = "upstream_error" errMsg = "Upstream request failed" } c.JSON(statusCode, gin.H{ "type": "error", "error": gin.H{"type": errType, "message": errMsg}, }) if upstreamMsg == "" { return fmt.Errorf("upstream error: %d", upstreamStatus) } return fmt.Errorf("upstream error: %d message=%s", upstreamStatus, upstreamMsg) } func (s *AntigravityGatewayService) writeGoogleError(c *gin.Context, status int, message string) error { MarkResponseCommitted(c) statusStr := "UNKNOWN" switch status { case 400: statusStr = "INVALID_ARGUMENT" case 404: statusStr = "NOT_FOUND" case 429: statusStr = "RESOURCE_EXHAUSTED" case 500: statusStr = "INTERNAL" case 502, 503: statusStr = "UNAVAILABLE" } c.JSON(status, gin.H{ "error": gin.H{ "code": status, "message": message, "status": statusStr, }, }) return fmt.Errorf("%s", message) } // collectClaudeStreamResponse 收集上游流式响应,转换为 Claude 非流式格式返回 // 用于处理客户端非流式请求但上游只支持流式的情况 func (s *AntigravityGatewayService) collectClaudeStreamResponse(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) ([]byte, *antigravityStreamResult, error) { if upstreamResponseModelObserverFromContext(c) == nil { beginUpstreamResponseModelObservation(c) } scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { maxLineSize = s.settingService.cfg.Gateway.MaxLineSize } scanBuf := getSSEScannerBuf64K() scanner.Buffer(scanBuf[:0], maxLineSize) var firstTokenMs *int var last map[string]any var lastWithParts map[string]any var collectedParts []map[string]any // 收集所有 parts(包括 text、thinking、functionCall、inlineData 等) var meaningfulResponse bool type scanEvent struct { line string err error } // 独立 goroutine 读取上游,避免读取阻塞影响超时处理 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.settingService.cfg != nil && s.settingService.cfg.Gateway.StreamDataIntervalTimeout > 0 { streamInterval = time.Duration(s.settingService.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 } for { select { case ev, ok := <-events: if !ok { // 流结束,转换并返回响应 goto returnResponse } if ev.err != nil { if errors.Is(ev.err, bufio.ErrTooLong) { logger.LegacyPrintf("service.antigravity_gateway", "SSE line too long (antigravity claude non-stream): max_size=%d error=%v", maxLineSize, ev.err) } return nil, nil, ev.err } line := ev.line trimmed := strings.TrimRight(line, "\r\n") if !strings.HasPrefix(trimmed, "data:") { continue } payload := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) if payload == "" || payload == "[DONE]" { continue } // 解包 v1internal 响应 inner, parseErr := s.unwrapV1InternalResponse([]byte(payload)) if parseErr != nil { continue } upstreamResponseModelObserverFromContext(c).ObserveGemini(inner) var parsed map[string]any if err := json.Unmarshal(inner, &parsed); err != nil { continue } last = parsed // 保留最后一个有 parts 的响应,并收集所有 parts parts := extractGeminiParts(parsed) if len(parts) > 0 { lastWithParts = parsed // 收集所有 parts(text、thinking、functionCall、inlineData 等) collectedParts = append(collectedParts, parts...) } if len(parts) > 0 || strings.TrimSpace(extractGeminiFinishReason(parsed)) != "" { meaningfulResponse = true if firstTokenMs == nil { ms := int(time.Since(startTime).Milliseconds()) firstTokenMs = &ms } } case <-intervalCh: lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) if time.Since(lastRead) < streamInterval { continue } logger.LegacyPrintf("service.antigravity_gateway", "Stream data interval timeout (antigravity claude non-stream)") return nil, nil, fmt.Errorf("stream data interval timeout") } } returnResponse: // 处理空响应情况 — 触发同账号重试 + failover 切换账号 if !meaningfulResponse { logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] warning: empty stream response (claude non-stream), triggering failover") return nil, nil, &UpstreamFailoverError{ StatusCode: http.StatusBadGateway, ResponseBody: []byte(`{"error":"empty stream response from upstream"}`), RetryableOnSameAccount: true, } } // 选择最后一个有效响应 finalResponse := pickGeminiCollectResult(last, lastWithParts) // 将收集的所有 parts 合并到最终响应中 if len(collectedParts) > 0 { finalResponse = mergeCollectedPartsToResponse(finalResponse, collectedParts) } // 序列化为 JSON(Gemini 格式) geminiBody, err := json.Marshal(finalResponse) if err != nil { return nil, nil, fmt.Errorf("failed to marshal gemini response: %w", err) } // 转换 Gemini 响应为 Claude 格式 claudeResp, agUsage, err := antigravity.TransformGeminiToClaude(geminiBody, originalModel) if err != nil { logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] transform_error error=%v body=%s", err, string(geminiBody)) return nil, nil, fmt.Errorf("failed to parse upstream response: %w", err) } // 转换为 service.ClaudeUsage usage := &ClaudeUsage{ InputTokens: agUsage.InputTokens, OutputTokens: agUsage.OutputTokens, CacheCreationInputTokens: agUsage.CacheCreationInputTokens, CacheReadInputTokens: agUsage.CacheReadInputTokens, ImageOutputTokens: agUsage.ImageOutputTokens, } return claudeResp, &antigravityStreamResult{usage: usage, firstTokenMs: firstTokenMs}, nil } // handleClaudeStreamToNonStreaming 收集上游流式响应,转换为 Claude 非流式格式返回 // 用于处理客户端非流式请求但上游只支持流式的情况 func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) (*antigravityStreamResult, error) { claudeResp, streamRes, err := s.collectClaudeStreamResponse(c, resp, startTime, originalModel) if err != nil { var failoverErr *UpstreamFailoverError if errors.As(err, &failoverErr) { return nil, err } if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { return nil, err } errMsg := "Failed to parse upstream response" errType := "upstream_error" if strings.Contains(err.Error(), "stream data interval timeout") { errMsg = "Upstream stream data interval timeout" errType = "upstream_timeout" } else if errors.Is(err, bufio.ErrTooLong) { errMsg = "Upstream response line too long" errType = "response_too_large" } return nil, s.writeClaudeError(c, http.StatusBadGateway, errType, errMsg) } c.Data(http.StatusOK, "application/json", claudeResp) return streamRes, nil } // handleClaudeStreamingResponse 处理 Claude 流式响应(Gemini SSE → Claude SSE 转换) func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) (*antigravityStreamResult, error) { c.Header("Content-Type", "text/event-stream") c.Header("Cache-Control", "no-cache") c.Header("Connection", "keep-alive") c.Header("X-Accel-Buffering", "no") c.Status(http.StatusOK) flusher, ok := c.Writer.(http.Flusher) if !ok { return nil, errors.New("streaming not supported") } processor := antigravity.NewStreamingProcessor(originalModel) var firstTokenMs *int // 使用 Scanner 并限制单行大小,避免 ReadString 无上限导致 OOM scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { maxLineSize = s.settingService.cfg.Gateway.MaxLineSize } scanBuf := getSSEScannerBuf64K() scanner.Buffer(scanBuf[:0], maxLineSize) // 辅助函数:转换 antigravity.ClaudeUsage 到 service.ClaudeUsage convertUsage := func(agUsage *antigravity.ClaudeUsage) *ClaudeUsage { if agUsage == nil { return &ClaudeUsage{} } return &ClaudeUsage{ InputTokens: agUsage.InputTokens, OutputTokens: agUsage.OutputTokens, CacheCreationInputTokens: agUsage.CacheCreationInputTokens, CacheReadInputTokens: agUsage.CacheReadInputTokens, } } type scanEvent struct { line string err error } // 独立 goroutine 读取上游,避免读取阻塞影响超时处理 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.settingService.cfg != nil && s.settingService.cfg.Gateway.StreamDataIntervalTimeout > 0 { streamInterval = time.Duration(s.settingService.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 } // 下游 keepalive:防止代理/Cloudflare Tunnel 因连接空闲而断开 keepaliveInterval := time.Duration(0) if s.settingService.cfg != nil && s.settingService.cfg.Gateway.StreamKeepaliveInterval > 0 { keepaliveInterval = time.Duration(s.settingService.cfg.Gateway.StreamKeepaliveInterval) * time.Second } var keepaliveTicker *time.Ticker if keepaliveInterval > 0 { keepaliveTicker = time.NewTicker(keepaliveInterval) defer keepaliveTicker.Stop() } var keepaliveCh <-chan time.Time if keepaliveTicker != nil { keepaliveCh = keepaliveTicker.C } lastDataAt := time.Now() cw := newAntigravityClientWriter(c.Writer, flusher, "antigravity claude") // 仅发送一次错误事件,避免多次写入导致协议混乱 errorEventSent := false sendErrorEvent := func(reason string) { if errorEventSent || cw.Disconnected() { return } errorEventSent = true _, _ = fmt.Fprintf(c.Writer, "event: error\ndata: {\"error\":\"%s\"}\n\n", reason) flusher.Flush() } // finishUsage 是获取 processor 最终 usage 的辅助函数 finishUsage := func() *ClaudeUsage { _, agUsage := processor.Finish() return convertUsage(agUsage) } for { select { case ev, ok := <-events: if !ok { // 上游完成,发送结束事件 finalEvents, agUsage := processor.Finish() if len(finalEvents) > 0 { cw.Write(finalEvents) } else if !processor.MessageStartSent() && !cw.Disconnected() { // 整个流未收到任何可解析的上游数据(全部 SSE 行均无法被 JSON 解析), // 触发 failover 在同账号重试,避免向客户端发出缺少 message_start 的残缺流 logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Claude-Stream] empty stream response (no valid events parsed), triggering failover") return nil, &UpstreamFailoverError{ StatusCode: http.StatusBadGateway, ResponseBody: []byte(`{"error":"empty stream response from upstream"}`), RetryableOnSameAccount: true, } } return &antigravityStreamResult{usage: convertUsage(agUsage), firstTokenMs: firstTokenMs, clientDisconnect: cw.Disconnected()}, nil } if ev.err != nil { if disconnect, handled := handleStreamReadError(ev.err, cw.Disconnected(), "antigravity claude"); handled { return &antigravityStreamResult{usage: finishUsage(), firstTokenMs: firstTokenMs, clientDisconnect: disconnect}, nil } if errors.Is(ev.err, bufio.ErrTooLong) { logger.LegacyPrintf("service.antigravity_gateway", "SSE line too long (antigravity): max_size=%d error=%v", maxLineSize, ev.err) sendErrorEvent("response_too_large") return &antigravityStreamResult{usage: convertUsage(nil), firstTokenMs: firstTokenMs}, ev.err } sendErrorEvent("stream_read_error") return nil, fmt.Errorf("stream read error: %w", ev.err) } lastDataAt = time.Now() s.observeAntigravityGeminiSSELine(c, ev.line) // 处理 SSE 行,转换为 Claude 格式 claudeEvents := processor.ProcessLine(strings.TrimRight(ev.line, "\r\n")) if len(claudeEvents) > 0 { if firstTokenMs == nil { ms := int(time.Since(startTime).Milliseconds()) firstTokenMs = &ms } cw.Write(claudeEvents) } case <-intervalCh: lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt)) if time.Since(lastRead) < streamInterval { continue } if cw.Disconnected() { logger.LegacyPrintf("service.antigravity_gateway", "Upstream timeout after client disconnect (antigravity claude), returning collected usage") return &antigravityStreamResult{usage: finishUsage(), firstTokenMs: firstTokenMs, clientDisconnect: true}, nil } logger.LegacyPrintf("service.antigravity_gateway", "Stream data interval timeout (antigravity)") sendErrorEvent("stream_timeout") return &antigravityStreamResult{usage: convertUsage(nil), firstTokenMs: firstTokenMs}, fmt.Errorf("stream data interval timeout") case <-keepaliveCh: if cw.Disconnected() { continue } if time.Since(lastDataAt) < keepaliveInterval { continue } // SSE ping 事件:Anthropic 原生格式,客户端会正确处理, // 同时保持连接活跃防止 Cloudflare Tunnel 等代理断开 if !cw.Fprintf("event: ping\ndata: {\"type\": \"ping\"}\n\n") { logger.LegacyPrintf("service.antigravity_gateway", "Client disconnected during keepalive ping (antigravity claude), continuing to drain upstream for billing") continue } } } } func (s *AntigravityGatewayService) extractImageInputSize(body []byte) string { var req antigravity.GeminiRequest if err := json.Unmarshal(body, &req); err != nil { return "" } if req.GenerationConfig != nil && req.GenerationConfig.ImageConfig != nil { return strings.TrimSpace(req.GenerationConfig.ImageConfig.ImageSize) } return "" } // isImageGenerationModel 判断模型是否为图片生成模型 // 支持的模型:gemini-3.1-flash-image, gemini-3-pro-image, gemini-2.5-flash-image 等 func isImageGenerationModel(model string) bool { modelLower := strings.ToLower(model) // 移除 models/ 前缀 modelLower = strings.TrimPrefix(modelLower, "models/") // 精确匹配或前缀匹配 return modelLower == "gemini-3.1-flash-image" || modelLower == "gemini-3.1-flash-image-preview" || strings.HasPrefix(modelLower, "gemini-3.1-flash-image-") || modelLower == "gemini-3-pro-image" || modelLower == "gemini-3-pro-image-preview" || strings.HasPrefix(modelLower, "gemini-3-pro-image-") || modelLower == "gemini-2.5-flash-image" || modelLower == "gemini-2.5-flash-image-preview" || strings.HasPrefix(modelLower, "gemini-2.5-flash-image-") }