From 71d7f86883db55411dc3df9c8245132d4b172b7e Mon Sep 17 00:00:00 2001 From: chinnsenn Date: Sun, 26 Jul 2026 19:37:24 +0900 Subject: [PATCH 1/3] fix(antigravity): 1. harden OpenAI compatibility forwarding 2. reject usage-only non-stream responses --- backend/internal/handler/endpoint.go | 24 +- backend/internal/handler/endpoint_test.go | 34 + .../gateway_handler_chat_completions.go | 11 + .../gateway_handler_error_fallback_test.go | 17 + .../handler/gateway_handler_responses.go | 16 +- ...openai_gateway_credential_failover_test.go | 22 + .../handler/openai_gateway_handler.go | 4 +- .../service/antigravity_gateway_compat.go | 526 +++++++++++++++ .../antigravity_gateway_compat_stream.go | 473 ++++++++++++++ .../antigravity_gateway_compat_test.go | 617 ++++++++++++++++++ .../service/antigravity_gateway_streaming.go | 94 ++- 11 files changed, 1810 insertions(+), 28 deletions(-) create mode 100644 backend/internal/service/antigravity_gateway_compat.go create mode 100644 backend/internal/service/antigravity_gateway_compat_stream.go create mode 100644 backend/internal/service/antigravity_gateway_compat_test.go diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index 0e0d28bb81..df8776981c 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -31,9 +31,12 @@ const ( EndpointGeminiModels = "/v1beta/models" ) +const EndpointAntigravityGenerateContent = "/v1internal:streamGenerateContent" + // gin.Context keys used by the middleware and helpers below. const ( - ctxKeyInboundEndpoint = "_gateway_inbound_endpoint" + ctxKeyInboundEndpoint = "_gateway_inbound_endpoint" + ctxKeyActualUpstreamEndpoint = "_gateway_actual_upstream_endpoint" ) // ────────────────────────────────────────────────────────── @@ -297,6 +300,13 @@ func GetInboundEndpoint(c *gin.Context) string { // and the account platform. Handlers call this after scheduling an // account, passing account.Platform. func GetUpstreamEndpoint(c *gin.Context, platform string) string { + if c != nil { + if value, ok := c.Get(ctxKeyActualUpstreamEndpoint); ok { + if endpoint, ok := value.(string); ok && endpoint != "" { + return endpoint + } + } + } inbound := GetInboundEndpoint(c) rawPath := "" if c != nil && c.Request != nil && c.Request.URL != nil { @@ -304,3 +314,15 @@ func GetUpstreamEndpoint(c *gin.Context, platform string) string { } return DeriveUpstreamEndpoint(inbound, rawPath, platform) } + +func setActualUpstreamEndpoint(c *gin.Context, endpoint string) { + if c != nil { + c.Set(ctxKeyActualUpstreamEndpoint, strings.TrimSpace(endpoint)) + } +} + +func shouldUseAntigravityCompat(account *service.Account) bool { + return account != nil && + account.Platform == service.PlatformAntigravity && + account.Type == service.AccountTypeOAuth +} diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go index c8ae566ad5..3f7b994702 100644 --- a/backend/internal/handler/endpoint_test.go +++ b/backend/internal/handler/endpoint_test.go @@ -146,6 +146,40 @@ func TestDeriveUpstreamEndpoint(t *testing.T) { } } +func TestShouldUseAntigravityCompat(t *testing.T) { + tests := []struct { + name string + account *service.Account + want bool + }{ + {"oauth", &service.Account{Platform: service.PlatformAntigravity, Type: service.AccountTypeOAuth}, true}, + {"setup token", &service.Account{Platform: service.PlatformAntigravity, Type: service.AccountTypeSetupToken}, false}, + {"upstream", &service.Account{Platform: service.PlatformAntigravity, Type: service.AccountTypeUpstream}, false}, + {"api key", &service.Account{Platform: service.PlatformAntigravity, Type: service.AccountTypeAPIKey}, false}, + {"anthropic oauth", &service.Account{Platform: service.PlatformAnthropic, Type: service.AccountTypeOAuth}, false}, + {"nil", nil, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, shouldUseAntigravityCompat(tt.account)) + }) + } +} + +func TestGetUpstreamEndpointPrefersRuntimeOverride(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, EndpointChatCompletions, nil) + c.Set(ctxKeyInboundEndpoint, EndpointChatCompletions) + + setActualUpstreamEndpoint(c, EndpointAntigravityGenerateContent) + require.Equal(t, EndpointAntigravityGenerateContent, GetUpstreamEndpoint(c, service.PlatformAntigravity)) + + setActualUpstreamEndpoint(c, "") + require.Equal(t, EndpointMessages, GetUpstreamEndpoint(c, service.PlatformAntigravity)) +} + func TestResolveOpenAIUpstreamEndpointPrefersForwardResult(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index 7061325f8a..3bbac76adb 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -240,6 +240,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { forwardBody = h.gatewayService.ReplaceModelInBody(body, channelMapping.MappedModel) } var result *service.ForwardResult + setActualUpstreamEndpoint(c, "") if account.Platform == service.PlatformGemini { if h.geminiCompatService == nil { h.chatCompletionsErrorResponse(c, http.StatusBadGateway, "upstream_error", "Gemini compatibility service is not configured") @@ -249,6 +250,16 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { return } result, err = h.geminiCompatService.ForwardAsChatCompletions(c.Request.Context(), c, account, forwardBody) + } else if shouldUseAntigravityCompat(account) { + if h.antigravityGatewayService == nil { + h.chatCompletionsErrorResponse(c, http.StatusBadGateway, "upstream_error", "Antigravity compatibility service is not configured") + if accountReleaseFunc != nil { + accountReleaseFunc() + } + return + } + setActualUpstreamEndpoint(c, EndpointAntigravityGenerateContent) + result, err = h.antigravityGatewayService.ForwardAsChatCompletions(c.Request.Context(), c, account, forwardBody, parsedReq) } else { result, err = h.gatewayService.ForwardAsChatCompletions(c.Request.Context(), c, account, forwardBody, parsedReq) } diff --git a/backend/internal/handler/gateway_handler_error_fallback_test.go b/backend/internal/handler/gateway_handler_error_fallback_test.go index 40ba79f7f1..a38580dfc6 100644 --- a/backend/internal/handler/gateway_handler_error_fallback_test.go +++ b/backend/internal/handler/gateway_handler_error_fallback_test.go @@ -8,6 +8,7 @@ import ( "strings" "testing" + "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -53,6 +54,22 @@ func TestGatewayEnsureForwardErrorResponse_AppendsSSEAfterWritten(t *testing.T) assert.Contains(t, w.Body.String(), `data: {"type":"error"`) } +func TestGatewayEnsureForwardErrorResponse_SkipsCommittedSSEError(t *testing.T) { + gin.SetMode(gin.TestMode) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, EndpointResponses, nil) + c.Header("Content-Type", "text/event-stream") + _, _ = c.Writer.WriteString("event: error\ndata: {\"type\":\"error\"}\n\n") + service.MarkResponseCommitted(c) + + h := &GatewayHandler{} + wrote := h.ensureForwardErrorResponse(c, true) + + require.False(t, wrote) + require.Equal(t, 1, strings.Count(w.Body.String(), "event: error")) +} + // case B 回归:Anthropic-backed /responses,Writer 已被写过时 // ensureForwardErrorResponse 仍要发 response.failed。 func TestGatewayEnsureForwardErrorResponse_ResponsesRouteAfterWrittenEmitsResponseFailed(t *testing.T) { diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 8f5d8531bb..5653438613 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -228,7 +228,21 @@ func (h *GatewayHandler) Responses(c *gin.Context) { if channelMapping.Mapped { forwardBody = h.gatewayService.ReplaceModelInBody(body, channelMapping.MappedModel) } - result, err := h.gatewayService.ForwardAsResponses(requestCtx, c, account, forwardBody, parsedReq) + var result *service.ForwardResult + setActualUpstreamEndpoint(c, "") + if shouldUseAntigravityCompat(account) { + if h.antigravityGatewayService == nil { + h.responsesErrorResponse(c, http.StatusBadGateway, "upstream_error", "Antigravity compatibility service is not configured") + if accountReleaseFunc != nil { + accountReleaseFunc() + } + return + } + setActualUpstreamEndpoint(c, EndpointAntigravityGenerateContent) + result, err = h.antigravityGatewayService.ForwardAsResponses(requestCtx, c, account, forwardBody, parsedReq) + } else { + result, err = h.gatewayService.ForwardAsResponses(requestCtx, c, account, forwardBody, parsedReq) + } if accountReleaseFunc != nil { accountReleaseFunc() diff --git a/backend/internal/handler/openai_gateway_credential_failover_test.go b/backend/internal/handler/openai_gateway_credential_failover_test.go index d650ef6a21..046b7ec52a 100644 --- a/backend/internal/handler/openai_gateway_credential_failover_test.go +++ b/backend/internal/handler/openai_gateway_credential_failover_test.go @@ -42,6 +42,28 @@ func TestGatewayChatCredentialStopDoesNotSelectAnotherAccountAndReturnsSafe503(t require.NotContains(t, recorder.Body.String(), "client_secret") } +func TestGatewayChatAntigravityCredentialFailureReturnsActionableMessage(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + + (&GatewayHandler{}).handleCCFailoverExhausted(c, &service.UpstreamFailoverError{ + StatusCode: http.StatusUnauthorized, + Stage: service.GatewayFailureStageAccountAuth, + Scope: service.GatewayFailureScopeAccount, + Reason: service.AntigravityCredentialRejectedReason, + NextAccountAction: service.NextAccountRetry, + ClientStatusCode: http.StatusBadGateway, + ClientMessage: service.AntigravityCredentialRejectedClientMessage, + ResponseBody: []byte(`{"error":{"message":"Invalid bearer token","refresh_token":"must-not-leak"}}`), + }, false) + + require.Equal(t, http.StatusBadGateway, recorder.Code) + require.Contains(t, recorder.Body.String(), service.AntigravityCredentialRejectedClientMessage) + require.NotContains(t, strings.ToLower(recorder.Body.String()), "bearer") + require.NotContains(t, strings.ToLower(recorder.Body.String()), "refresh_token") +} + func TestGatewayChatInferenceExhaustionRestoresRetryAfter(t *testing.T) { gin.SetMode(gin.TestMode) recorder := httptest.NewRecorder() diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 761f291e90..9b550ffa7d 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -2227,7 +2227,9 @@ func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverE } func credentialFailoverClientResponse(failoverErr *service.UpstreamFailoverError) (int, string) { - _ = failoverErr + if failoverErr != nil && failoverErr.Reason == service.AntigravityCredentialRejectedReason { + return http.StatusBadGateway, service.AntigravityCredentialRejectedClientMessage + } return http.StatusServiceUnavailable, service.GrokCredentialUnavailableClientMessage } diff --git a/backend/internal/service/antigravity_gateway_compat.go b/backend/internal/service/antigravity_gateway_compat.go new file mode 100644 index 0000000000..af1b8ebbe4 --- /dev/null +++ b/backend/internal/service/antigravity_gateway_compat.go @@ -0,0 +1,526 @@ +package service + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/gin-gonic/gin" +) + +type antigravityCompatProtocol uint8 + +const ( + antigravityCompatChatCompletions antigravityCompatProtocol = iota + antigravityCompatResponses +) + +const ( + // AntigravityCredentialRejectedClientMessage 是可安全返回给客户端的认证修复提示。 + AntigravityCredentialRejectedClientMessage = "Antigravity rejected the OAuth credential after refresh; reauthorize the account and verify project_id" + // AntigravityCredentialRejectedReason 标识上游拒绝已刷新 OAuth 凭据。 + AntigravityCredentialRejectedReason GatewayFailureReason = "antigravity_oauth_credential_rejected" +) + +type antigravityCompatRequest struct { + protocol antigravityCompatProtocol + originalBody []byte + claudeBody []byte + originalModel string + clientStream bool + includeUsage bool + startTime time.Time + reasoningEffort *string +} + +type antigravityCompatUpstreamCall struct { + request antigravityCompatRequest + billingModel string + prefix string + proxyURL string + accessToken string + geminiBody []byte +} + +// ForwardAsChatCompletions 使用 Antigravity 原生 OAuth 账号转发 Chat Completions 请求。 +func (s *AntigravityGatewayService) ForwardAsChatCompletions( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + _ *ParsedRequest, +) (*ForwardResult, error) { + if err := s.validateAntigravityCompatAccount(c, account); err != nil { + return nil, err + } + + var request apicompat.ChatCompletionsRequest + if json.Unmarshal(body, &request) != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") + } + if strings.TrimSpace(request.Model) == "" { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", "model is required") + } + + responsesRequest, err := apicompat.ChatCompletionsToResponses(&request) + if err != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + } + claudeRequest, err := apicompat.ResponsesToAnthropicRequest(responsesRequest) + if err != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + } + preserveChatCompletionTokenLimit(&request, claudeRequest) + claudeRequest.Stream = request.Stream + claudeBody, err := json.Marshal(claudeRequest) + if err != nil { + return nil, fmt.Errorf("marshal anthropic request: %w", err) + } + + return s.forwardAntigravityCompat(ctx, c, account, antigravityCompatRequest{ + protocol: antigravityCompatChatCompletions, + originalBody: body, + claudeBody: claudeBody, + originalModel: request.Model, + clientStream: request.Stream, + includeUsage: request.StreamOptions != nil && request.StreamOptions.IncludeUsage, + startTime: time.Now(), + reasoningEffort: extractCCReasoningEffortFromBody(body), + }) +} + +// ForwardAsResponses 使用 Antigravity 原生 OAuth 账号转发 Responses 请求。 +func (s *AntigravityGatewayService) ForwardAsResponses( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + _ *ParsedRequest, +) (*ForwardResult, error) { + if err := s.validateAntigravityCompatAccount(c, account); err != nil { + return nil, err + } + + var request apicompat.ResponsesRequest + if json.Unmarshal(body, &request) != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") + } + if strings.TrimSpace(request.Model) == "" { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", "model is required") + } + + claudeRequest, err := apicompat.ResponsesToAnthropicRequest(&request) + if err != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + } + claudeRequest.Stream = request.Stream + claudeBody, err := json.Marshal(claudeRequest) + if err != nil { + return nil, fmt.Errorf("marshal anthropic request: %w", err) + } + + return s.forwardAntigravityCompat(ctx, c, account, antigravityCompatRequest{ + protocol: antigravityCompatResponses, + originalBody: body, + claudeBody: claudeBody, + originalModel: request.Model, + clientStream: request.Stream, + startTime: time.Now(), + reasoningEffort: ExtractResponsesReasoningEffortFromBody(body), + }) +} + +func (s *AntigravityGatewayService) validateAntigravityCompatAccount(c *gin.Context, account *Account) error { + if account != nil && account.Platform == PlatformAntigravity && account.Type == AccountTypeOAuth { + return nil + } + return s.writeAntigravityCompatError( + c, + http.StatusBadRequest, + "invalid_request_error", + "native OAuth account required for antigravity compatibility mode", + ) +} + +func preserveChatCompletionTokenLimit(request *apicompat.ChatCompletionsRequest, claudeRequest *apicompat.AnthropicRequest) { + if request == nil || claudeRequest == nil { + return + } + limit := request.MaxTokens + if request.MaxCompletionTokens != nil { + limit = request.MaxCompletionTokens + } + if limit != nil && *limit > 0 { + claudeRequest.MaxTokens = *limit + } +} + +func (s *AntigravityGatewayService) forwardAntigravityCompat( + ctx context.Context, + c *gin.Context, + account *Account, + request antigravityCompatRequest, +) (*ForwardResult, error) { + call, err := s.prepareAntigravityCompatCall(ctx, c, account, request) + if err != nil { + return nil, err + } + + result, err := s.antigravityRetryLoop(antigravityRetryLoopParams{ + ctx: ctx, + prefix: call.prefix, + account: account, + proxyURL: call.proxyURL, + accessToken: call.accessToken, + action: "streamGenerateContent", + body: call.geminiBody, + c: c, + httpUpstream: s.httpUpstream, + settingService: s.settingService, + accountRepo: s.accountRepo, + handleError: s.handleUpstreamError, + requestedModel: request.originalModel, + isStickySession: false, + groupID: 0, + sessionHash: "", + }) + if err != nil { + return nil, s.handleAntigravityCompatTransportError(c, err) + } + + return s.consumeAntigravityCompatResponse(ctx, c, account, call, result.resp) +} + +func (s *AntigravityGatewayService) prepareAntigravityCompatCall( + ctx context.Context, + c *gin.Context, + account *Account, + request antigravityCompatRequest, +) (*antigravityCompatUpstreamCall, error) { + var claudeRequest antigravity.ClaudeRequest + if json.Unmarshal(request.claudeBody, &claudeRequest) != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body") + } + + mappedModel := s.getMappedModel(account, request.originalModel) + if mappedModel == "" { + MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate) + message := fmt.Sprintf("model %s not in whitelist", request.originalModel) + return nil, s.writeAntigravityCompatError(c, http.StatusForbidden, "permission_error", message) + } + thinkingEnabled := claudeRequest.Thinking != nil && + (claudeRequest.Thinking.Type == "enabled" || claudeRequest.Thinking.Type == "adaptive") + mappedModel = applyThinkingModelSuffix(mappedModel, thinkingEnabled) + + if s.tokenProvider == nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadGateway, "api_error", "Antigravity token provider not configured") + } + accessToken, err := s.tokenProvider.GetAccessToken(ctx, account) + if err != nil { + return nil, &UpstreamFailoverError{ + StatusCode: http.StatusBadGateway, + ResponseBody: []byte(`{"error":{"type":"authentication_error","message":"Failed to get upstream access token"},"type":"error"}`), + } + } + + projectID, err := resolveAntigravityProjectID(account) + if err != nil { + _ = s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + return nil, err + } + geminiBody, err := s.buildAntigravityCompatGeminiBody(ctx, request.claudeBody, &claudeRequest, projectID, mappedModel) + if err != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadRequest, "invalid_request_error", "Invalid request") + } + + request.reasoningEffort = ApplyThinkingEnabledFallback(request.reasoningEffort, request.originalBody, mappedModel) + return &antigravityCompatUpstreamCall{ + request: request, + billingModel: mappedModel, + prefix: logPrefix(getSessionID(c), account.Name), + proxyURL: antigravityCompatProxyURL(account), + accessToken: accessToken, + geminiBody: geminiBody, + }, nil +} + +func (s *AntigravityGatewayService) buildAntigravityCompatGeminiBody( + ctx context.Context, + claudeBody []byte, + claudeRequest *antigravity.ClaudeRequest, + projectID string, + mappedModel string, +) ([]byte, error) { + if strings.HasPrefix(strings.ToLower(mappedModel), "gemini-") { + body, err := convertClaudeMessagesToGeminiGenerateContent(claudeBody) + if err != nil { + return nil, err + } + body = ensureGeminiFunctionCallThoughtSignatures(body) + body, err = injectIdentityPatchToGeminiRequest(body) + if err != nil { + return nil, err + } + if cleaned, cleanErr := cleanGeminiRequest(body); cleanErr == nil { + body = cleaned + } + return s.wrapV1InternalRequest(projectID, mappedModel, body) + } + + options := s.getClaudeTransformOptions(ctx) + options.EnableIdentityPatch = true + return antigravity.TransformClaudeToGeminiWithOptions(claudeRequest, projectID, mappedModel, options) +} + +func antigravityCompatProxyURL(account *Account) string { + if account.ProxyID == nil || account.Proxy == nil { + return "" + } + return account.Proxy.URL() +} + +func (s *AntigravityGatewayService) handleAntigravityCompatTransportError(c *gin.Context, err error) error { + if switchErr, ok := IsAntigravityAccountSwitchError(err); ok { + return &UpstreamFailoverError{ + StatusCode: http.StatusServiceUnavailable, + ForceCacheBilling: switchErr.IsStickySession, + } + } + if c.Request.Context().Err() != nil { + return s.writeAntigravityCompatError(c, http.StatusBadGateway, "client_disconnected", "Client disconnected before upstream response") + } + return s.writeAntigravityCompatError(c, http.StatusBadGateway, "upstream_error", "Upstream request failed after retries") +} + +func (s *AntigravityGatewayService) consumeAntigravityCompatResponse( + ctx context.Context, + c *gin.Context, + account *Account, + call *antigravityCompatUpstreamCall, + resp *http.Response, +) (*ForwardResult, error) { + defer func() { _ = resp.Body.Close() }() + if resp.StatusCode >= http.StatusBadRequest { + return nil, s.handleAntigravityCompatHTTPError(ctx, c, account, call, resp) + } + + requestID := resp.Header.Get("x-request-id") + if requestID != "" { + c.Header("x-request-id", requestID) + } + streamResult, err := s.consumeAntigravityCompatSuccess(c, call, resp) + if err != nil { + return nil, err + } + if streamResult.usage == nil { + streamResult.usage = &ClaudeUsage{} + } + + return &ForwardResult{ + RequestID: requestID, + Usage: *streamResult.usage, + Model: call.request.originalModel, + UpstreamModel: call.billingModel, + Stream: call.request.clientStream, + Duration: time.Since(call.request.startTime), + FirstTokenMs: streamResult.firstTokenMs, + ReasoningEffort: call.request.reasoningEffort, + ClientDisconnect: streamResult.clientDisconnect, + }, nil +} + +func (s *AntigravityGatewayService) consumeAntigravityCompatSuccess( + c *gin.Context, + call *antigravityCompatUpstreamCall, + resp *http.Response, +) (*antigravityStreamResult, error) { + if call.request.clientStream { + if call.request.protocol == antigravityCompatChatCompletions { + return s.handleChatCompletionsStreamingFromAntigravity( + c, + resp, + call.request.startTime, + call.request.originalModel, + call.request.includeUsage, + ) + } + return s.handleResponsesStreamingFromAntigravity(c, resp, call.request.startTime, call.request.originalModel) + } + + if call.request.protocol == antigravityCompatChatCompletions { + return s.handleChatCompletionsNonStreamingFromAntigravity(c, resp, call.request.startTime, call.request.originalModel) + } + return s.handleResponsesNonStreamingFromAntigravity(c, resp, call.request.startTime, call.request.originalModel) +} + +func (s *AntigravityGatewayService) handleAntigravityCompatHTTPError( + ctx context.Context, + c *gin.Context, + account *Account, + call *antigravityCompatUpstreamCall, + resp *http.Response, +) error { + body := s.readUpstreamErrorBody(resp) + s.handleUpstreamError( + ctx, + call.prefix, + account, + resp.StatusCode, + resp.Header, + body, + call.request.originalModel, + 0, + "", + false, + ) + if s.shouldFailoverUpstreamError(resp.StatusCode) { + message := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractAntigravityErrorMessage(body))) + event := OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + Kind: "failover", + Message: message, + Detail: s.getUpstreamErrorDetail(body), + } + if resp.StatusCode == http.StatusUnauthorized { + event.Stage = string(GatewayFailureStageAccountAuth) + event.Scope = string(GatewayFailureScopeAccount) + event.Reason = string(AntigravityCredentialRejectedReason) + appendOpsUpstreamError(c, event) + return antigravityCredentialRejectedError(resp, body) + } + appendOpsUpstreamError(c, event) + return &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: body, + ResponseHeaders: resp.Header.Clone(), + } + } + return s.writeMappedAntigravityCompatError(c, account, resp.StatusCode, resp.Header.Get("x-request-id"), body) +} + +func antigravityCredentialRejectedError(resp *http.Response, body []byte) *UpstreamFailoverError { + return &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: body, + ResponseHeaders: resp.Header.Clone(), + Stage: GatewayFailureStageAccountAuth, + Scope: GatewayFailureScopeAccount, + Reason: AntigravityCredentialRejectedReason, + NextAccountAction: NextAccountRetry, + ClientStatusCode: http.StatusBadGateway, + ClientMessage: AntigravityCredentialRejectedClientMessage, + } +} + +func (s *AntigravityGatewayService) writeAntigravityCompatError( + c *gin.Context, + status int, + errType string, + message string, +) error { + MarkResponseCommitted(c) + c.JSON(status, gin.H{ + "error": gin.H{ + "message": message, + "type": errType, + "param": nil, + "code": nil, + }, + }) + return errors.New(message) +} + +func (s *AntigravityGatewayService) writeMappedAntigravityCompatError( + c *gin.Context, + account *Account, + upstreamStatus int, + upstreamRequestID string, + body []byte, +) error { + MarkResponseCommitted(c) + message := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractAntigravityErrorMessage(body))) + setOpsUpstreamError(c, upstreamStatus, message, s.getUpstreamErrorDetail(body)) + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: upstreamStatus, + UpstreamRequestID: upstreamRequestID, + Kind: "http_error", + Message: message, + }) + c.JSON(mapUpstreamStatusCode(upstreamStatus), gin.H{ + "error": gin.H{ + "message": getPassthroughOrDefault(message, "Upstream request failed"), + "type": "upstream_error", + "param": nil, + "code": nil, + }, + }) + return fmt.Errorf("upstream error: %d %s", upstreamStatus, message) +} + +func (s *AntigravityGatewayService) handleChatCompletionsNonStreamingFromAntigravity( + c *gin.Context, + resp *http.Response, + startTime time.Time, + originalModel string, +) (*antigravityStreamResult, error) { + claudeResponse, result, err := s.collectClaudeStreamResponse(resp, startTime, originalModel) + if err != nil { + return nil, s.mapAntigravityCompatCollectionError(c, err) + } + var anthropicResponse apicompat.AnthropicResponse + if json.Unmarshal(claudeResponse, &anthropicResponse) != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadGateway, "upstream_error", "Failed to parse upstream response") + } + responsesResponse := apicompat.AnthropicToResponsesResponse(&anthropicResponse) + c.JSON(http.StatusOK, apicompat.ResponsesToChatCompletions(responsesResponse, originalModel)) + return result, nil +} + +func (s *AntigravityGatewayService) handleResponsesNonStreamingFromAntigravity( + c *gin.Context, + resp *http.Response, + startTime time.Time, + originalModel string, +) (*antigravityStreamResult, error) { + claudeResponse, result, err := s.collectClaudeStreamResponse(resp, startTime, originalModel) + if err != nil { + return nil, s.mapAntigravityCompatCollectionError(c, err) + } + var anthropicResponse apicompat.AnthropicResponse + if json.Unmarshal(claudeResponse, &anthropicResponse) != nil { + return nil, s.writeAntigravityCompatError(c, http.StatusBadGateway, "upstream_error", "Failed to parse upstream response") + } + c.JSON(http.StatusOK, apicompat.AnthropicToResponsesResponse(&anthropicResponse)) + return result, nil +} + +func (s *AntigravityGatewayService) mapAntigravityCompatCollectionError(c *gin.Context, err error) error { + var failoverError *UpstreamFailoverError + if errors.As(err, &failoverError) { + return err + } + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return err + } + if strings.Contains(err.Error(), "stream data interval timeout") { + return s.writeAntigravityCompatError(c, http.StatusBadGateway, "upstream_timeout", "Upstream stream data interval timeout") + } + if errors.Is(err, bufio.ErrTooLong) { + return s.writeAntigravityCompatError(c, http.StatusBadGateway, "response_too_large", "Upstream response line too long") + } + return s.writeAntigravityCompatError(c, http.StatusBadGateway, "upstream_error", "Failed to parse upstream response") +} diff --git a/backend/internal/service/antigravity_gateway_compat_stream.go b/backend/internal/service/antigravity_gateway_compat_stream.go new file mode 100644 index 0000000000..3243193ceb --- /dev/null +++ b/backend/internal/service/antigravity_gateway_compat_stream.go @@ -0,0 +1,473 @@ +package service + +import ( + "bufio" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/gin-gonic/gin" +) + +type antigravityCompatStreamAdapter interface { + Emit(*apicompat.AnthropicStreamEvent, *antigravityClientWriter) + Finalize(*antigravityClientWriter) + WriteError(*antigravityClientWriter, string) +} + +type antigravityChatStreamAdapter struct { + anthropicState *apicompat.AnthropicEventToResponsesState + chatState *apicompat.ResponsesEventToChatState +} + +func newAntigravityChatStreamAdapter(model string, includeUsage bool) *antigravityChatStreamAdapter { + anthropicState := apicompat.NewAnthropicEventToResponsesState() + anthropicState.Model = model + chatState := apicompat.NewResponsesEventToChatState() + chatState.Model = model + chatState.IncludeUsage = includeUsage + return &antigravityChatStreamAdapter{ + anthropicState: anthropicState, + chatState: chatState, + } +} + +func (a *antigravityChatStreamAdapter) Emit(event *apicompat.AnthropicStreamEvent, writer *antigravityClientWriter) { + for _, responseEvent := range apicompat.AnthropicEventToResponsesEvents(event, a.anthropicState) { + a.emitResponseEvent(&responseEvent, writer) + } +} + +func (a *antigravityChatStreamAdapter) Finalize(writer *antigravityClientWriter) { + for _, responseEvent := range apicompat.FinalizeAnthropicResponsesStream(a.anthropicState) { + a.emitResponseEvent(&responseEvent, writer) + } + for _, chunk := range apicompat.FinalizeResponsesChatStream(a.chatState) { + if data, err := apicompat.ChatChunkToSSE(chunk); err == nil { + writer.Write([]byte(data)) + } + } + writer.Write([]byte("data: [DONE]\n\n")) +} + +func (a *antigravityChatStreamAdapter) WriteError(writer *antigravityClientWriter, reason string) { + writer.Fprintf("data: {\"error\":{\"message\":%q,\"type\":\"upstream_error\"}}\n\n", reason) +} + +func (a *antigravityChatStreamAdapter) emitResponseEvent(event *apicompat.ResponsesStreamEvent, writer *antigravityClientWriter) { + for _, chunk := range apicompat.ResponsesEventToChatChunks(event, a.chatState) { + if data, err := apicompat.ChatChunkToSSE(chunk); err == nil { + writer.Write([]byte(data)) + } + } +} + +type antigravityResponsesStreamAdapter struct { + anthropicState *apicompat.AnthropicEventToResponsesState +} + +func newAntigravityResponsesStreamAdapter(model string) *antigravityResponsesStreamAdapter { + state := apicompat.NewAnthropicEventToResponsesState() + state.Model = model + return &antigravityResponsesStreamAdapter{anthropicState: state} +} + +func (a *antigravityResponsesStreamAdapter) Emit(event *apicompat.AnthropicStreamEvent, writer *antigravityClientWriter) { + for _, responseEvent := range apicompat.AnthropicEventToResponsesEvents(event, a.anthropicState) { + a.emitResponseEvent(responseEvent, writer) + } +} + +func (a *antigravityResponsesStreamAdapter) Finalize(writer *antigravityClientWriter) { + for _, responseEvent := range apicompat.FinalizeAnthropicResponsesStream(a.anthropicState) { + a.emitResponseEvent(responseEvent, writer) + } +} + +func (a *antigravityResponsesStreamAdapter) WriteError(writer *antigravityClientWriter, reason string) { + writer.Fprintf("event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"upstream_error\",\"message\":%q}}\n\n", reason) +} + +func (a *antigravityResponsesStreamAdapter) emitResponseEvent(event apicompat.ResponsesStreamEvent, writer *antigravityClientWriter) { + if data, err := apicompat.ResponsesEventToSSE(event); err == nil { + writer.Write([]byte(data)) + } +} + +type antigravityCompatScanEvent struct { + line string + err error +} + +type antigravityCompatStreamSession struct { + processor *antigravity.StreamingProcessor + adapter antigravityCompatStreamAdapter + writer *antigravityClientWriter + usage *ClaudeUsage + pendingEvents []apicompat.AnthropicStreamEvent + firstTokenMs *int + startTime time.Time + meaningfulData bool +} + +func newAntigravityCompatStreamSession( + model string, + startTime time.Time, + adapter antigravityCompatStreamAdapter, + writer *antigravityClientWriter, +) *antigravityCompatStreamSession { + return &antigravityCompatStreamSession{ + processor: antigravity.NewStreamingProcessor(model), + adapter: adapter, + writer: writer, + usage: &ClaudeUsage{}, + startTime: startTime, + } +} + +func (s *antigravityCompatStreamSession) consume(line string) { + claudeEvents := s.processor.ProcessLine(strings.TrimRight(line, "\r\n")) + if len(claudeEvents) == 0 { + return + } + s.consumeClaudeEvents(claudeEvents) +} + +func (s *antigravityCompatStreamSession) hasMeaningfulData() bool { + return s.meaningfulData +} + +func (s *antigravityCompatStreamSession) finish() *antigravityStreamResult { + finalEvents, usage := s.processor.Finish() + mergeAntigravityCompatUsage(s.usage, usage) + s.consumeClaudeEvents(finalEvents) + s.adapter.Finalize(s.writer) + return s.result(s.writer.Disconnected()) +} + +func (s *antigravityCompatStreamSession) collectResult(clientDisconnect bool) *antigravityStreamResult { + _, usage := s.processor.Finish() + mergeAntigravityCompatUsage(s.usage, usage) + return s.result(clientDisconnect) +} + +func (s *antigravityCompatStreamSession) result(clientDisconnect bool) *antigravityStreamResult { + return &antigravityStreamResult{ + usage: s.usage, + firstTokenMs: s.firstTokenMs, + clientDisconnect: clientDisconnect, + } +} + +func (s *antigravityCompatStreamSession) consumeClaudeEvents(data []byte) { + var eventType string + for _, line := range strings.Split(string(data), "\n") { + line = strings.TrimSpace(line) + switch { + case strings.HasPrefix(line, "event:"): + eventType = strings.TrimSpace(strings.TrimPrefix(line, "event:")) + case strings.HasPrefix(line, "data:"): + s.consumeClaudeData(eventType, strings.TrimSpace(strings.TrimPrefix(line, "data:"))) + } + } +} + +func (s *antigravityCompatStreamSession) consumeClaudeData(eventType, payload string) { + var event apicompat.AnthropicStreamEvent + if json.Unmarshal([]byte(payload), &event) != nil { + return + } + if event.Type == "" { + event.Type = eventType + } + if event.Usage != nil { + mergeAnthropicUsage(s.usage, *event.Usage) + } + if event.Message != nil { + mergeAnthropicUsage(s.usage, event.Message.Usage) + } + s.emitOrBuffer(event) +} + +func (s *antigravityCompatStreamSession) emitOrBuffer(event apicompat.AnthropicStreamEvent) { + if s.meaningfulData { + s.adapter.Emit(&event, s.writer) + return + } + + s.pendingEvents = append(s.pendingEvents, event) + if !isMeaningfulAntigravityCompatEvent(&event) { + return + } + + s.meaningfulData = true + ms := int(time.Since(s.startTime).Milliseconds()) + s.firstTokenMs = &ms + for i := range s.pendingEvents { + s.adapter.Emit(&s.pendingEvents[i], s.writer) + } + s.pendingEvents = nil +} + +func isMeaningfulAntigravityCompatEvent(event *apicompat.AnthropicStreamEvent) bool { + if event == nil { + return false + } + if event.Type == "message_stop" { + return true + } + if event.ContentBlock != nil { + block := event.ContentBlock + return block.Type == "tool_use" || + block.Text != "" || + block.Thinking != "" || + block.Signature != "" || + block.Source != nil + } + if event.Delta != nil { + delta := event.Delta + return delta.Text != "" || + delta.PartialJSON != "" || + delta.Thinking != "" || + delta.Signature != "" || + delta.StopReason != "" + } + return false +} + +func mergeAntigravityCompatUsage(dst *ClaudeUsage, src *antigravity.ClaudeUsage) { + if dst == nil || src == nil { + return + } + dst.InputTokens = src.InputTokens + dst.OutputTokens = src.OutputTokens + dst.CacheCreationInputTokens = src.CacheCreationInputTokens + dst.CacheReadInputTokens = src.CacheReadInputTokens + dst.ImageOutputTokens = src.ImageOutputTokens +} + +func (s *AntigravityGatewayService) handleAntigravityCompatStream( + c *gin.Context, + resp *http.Response, + startTime time.Time, + originalModel string, + adapter antigravityCompatStreamAdapter, + prefix string, +) (*antigravityStreamResult, error) { + flusher, ok := c.Writer.(http.Flusher) + if !ok { + return nil, errors.New("streaming not supported") + } + + writer := newAntigravityClientWriter(c.Writer, flusher, prefix) + writer.beforeFirstWrite = func() { + 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) + } + session := newAntigravityCompatStreamSession(originalModel, startTime, adapter, writer) + events, stopScanner, maxLineSize := s.startAntigravityCompatScanner(resp.Body) + defer stopScanner() + + timeout := s.antigravityCompatStreamTimeout() + timeoutTimer, timeoutCh := newAntigravityCompatTimer(timeout) + if timeoutTimer != nil { + defer timeoutTimer.Stop() + } + keepaliveTicker, keepaliveCh := s.newAntigravityCompatKeepaliveTicker() + if keepaliveTicker != nil { + defer keepaliveTicker.Stop() + } + + for { + select { + case event, open := <-events: + if !open { + if !session.hasMeaningfulData() && !writer.Disconnected() { + return nil, antigravityCompatEmptyStreamError() + } + return session.finish(), nil + } + if event.err != nil { + return s.handleAntigravityCompatReadError(c, session, event.err, maxLineSize, prefix) + } + resetAntigravityCompatTimer(timeoutTimer, timeout) + session.consume(event.line) + + case <-timeoutCh: + if writer.Disconnected() { + return session.collectResult(true), nil + } + if !session.hasMeaningfulData() { + return nil, antigravityCompatEmptyStreamError() + } + logger.LegacyPrintf("service.antigravity_gateway", "Stream data interval timeout (%s)", prefix) + writeAntigravityCompatStreamError(c, adapter, writer, "stream_timeout") + return session.collectResult(false), fmt.Errorf("stream data interval timeout") + + case <-keepaliveCh: + if session.hasMeaningfulData() && !writer.Disconnected() { + writer.Write([]byte(": ping\n\n")) + } + } + } +} + +func (s *AntigravityGatewayService) startAntigravityCompatScanner( + body io.Reader, +) (<-chan antigravityCompatScanEvent, func(), int) { + maxLineSize := defaultMaxLineSize + if s.settingService != nil && s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.settingService.cfg.Gateway.MaxLineSize + } + scanner := bufio.NewScanner(body) + scanBuf := getSSEScannerBuf64K() + scanner.Buffer(scanBuf[:0], maxLineSize) + + events := make(chan antigravityCompatScanEvent, 16) + done := make(chan struct{}) + go func() { + defer putSSEScannerBuf64K(scanBuf) + defer close(events) + send := func(event antigravityCompatScanEvent) bool { + select { + case events <- event: + return true + case <-done: + return false + } + } + for scanner.Scan() { + if !send(antigravityCompatScanEvent{line: scanner.Text()}) { + return + } + } + if err := scanner.Err(); err != nil { + send(antigravityCompatScanEvent{err: err}) + } + }() + return events, func() { close(done) }, maxLineSize +} + +func (s *AntigravityGatewayService) antigravityCompatStreamTimeout() time.Duration { + if s.settingService == nil || s.settingService.cfg == nil { + return 0 + } + return time.Duration(s.settingService.cfg.Gateway.StreamDataIntervalTimeout) * time.Second +} + +func (s *AntigravityGatewayService) newAntigravityCompatKeepaliveTicker() (*time.Ticker, <-chan time.Time) { + if s.settingService == nil || s.settingService.cfg == nil { + return nil, nil + } + interval := time.Duration(s.settingService.cfg.Gateway.StreamKeepaliveInterval) * time.Second + if interval <= 0 { + return nil, nil + } + ticker := time.NewTicker(interval) + return ticker, ticker.C +} + +func newAntigravityCompatTimer(timeout time.Duration) (*time.Timer, <-chan time.Time) { + if timeout <= 0 { + return nil, nil + } + timer := time.NewTimer(timeout) + return timer, timer.C +} + +func resetAntigravityCompatTimer(timer *time.Timer, timeout time.Duration) { + if timer == nil { + return + } + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timer.Reset(timeout) +} + +func (s *AntigravityGatewayService) handleAntigravityCompatReadError( + c *gin.Context, + session *antigravityCompatStreamSession, + err error, + maxLineSize int, + prefix string, +) (*antigravityStreamResult, error) { + if !session.hasMeaningfulData() && !session.writer.Disconnected() { + return nil, antigravityCompatEmptyStreamError() + } + if disconnect, handled := handleStreamReadError(err, session.writer.Disconnected(), prefix); handled { + return session.collectResult(disconnect), nil + } + if errors.Is(err, bufio.ErrTooLong) { + logger.LegacyPrintf("service.antigravity_gateway", "SSE line too long (%s): max_size=%d error=%v", prefix, maxLineSize, err) + writeAntigravityCompatStreamError(c, session.adapter, session.writer, "response_too_large") + return session.result(false), err + } + writeAntigravityCompatStreamError(c, session.adapter, session.writer, "stream_read_error") + return nil, fmt.Errorf("stream read error: %w", err) +} + +func writeAntigravityCompatStreamError( + c *gin.Context, + adapter antigravityCompatStreamAdapter, + writer *antigravityClientWriter, + reason string, +) { + adapter.WriteError(writer, reason) + MarkResponseCommitted(c) +} + +func antigravityCompatEmptyStreamError() error { + logger.LegacyPrintf("service.antigravity_gateway", "Empty Antigravity compatibility stream, triggering failover") + return &UpstreamFailoverError{ + StatusCode: http.StatusBadGateway, + ResponseBody: []byte(`{"error":"empty stream response from upstream"}`), + RetryableOnSameAccount: true, + } +} + +func (s *AntigravityGatewayService) handleChatCompletionsStreamingFromAntigravity( + c *gin.Context, + resp *http.Response, + startTime time.Time, + originalModel string, + includeUsage bool, +) (*antigravityStreamResult, error) { + return s.handleAntigravityCompatStream( + c, + resp, + startTime, + originalModel, + newAntigravityChatStreamAdapter(originalModel, includeUsage), + "antigravity chat completions stream", + ) +} + +func (s *AntigravityGatewayService) handleResponsesStreamingFromAntigravity( + c *gin.Context, + resp *http.Response, + startTime time.Time, + originalModel string, +) (*antigravityStreamResult, error) { + return s.handleAntigravityCompatStream( + c, + resp, + startTime, + originalModel, + newAntigravityResponsesStreamAdapter(originalModel), + "antigravity responses stream", + ) +} diff --git a/backend/internal/service/antigravity_gateway_compat_test.go b/backend/internal/service/antigravity_gateway_compat_test.go new file mode 100644 index 0000000000..ff3d195cbf --- /dev/null +++ b/backend/internal/service/antigravity_gateway_compat_test.go @@ -0,0 +1,617 @@ +package service + +import ( + "bytes" + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +type antigravityCompatTokenCache struct { + token string +} + +type antigravityCompatErrorReader struct { + data []byte + off int + err error +} + +func (r *antigravityCompatErrorReader) Read(p []byte) (int, error) { + if r.off < len(r.data) { + n := copy(p, r.data[r.off:]) + r.off += n + return n, nil + } + return 0, r.err +} + +func (r *antigravityCompatErrorReader) Close() error { return nil } + +func (c *antigravityCompatTokenCache) GetAccessToken(context.Context, string) (string, error) { + return c.token, nil +} + +func (c *antigravityCompatTokenCache) SetAccessToken(context.Context, string, string, time.Duration) error { + return nil +} + +func (c *antigravityCompatTokenCache) DeleteAccessToken(context.Context, string) error { + return nil +} + +func (c *antigravityCompatTokenCache) AcquireRefreshLock(context.Context, string, time.Duration) (bool, error) { + return true, nil +} + +func (c *antigravityCompatTokenCache) ReleaseRefreshLock(context.Context, string) error { + return nil +} + +func newAntigravityCompatService(cfg config.GatewayConfig, upstream HTTPUpstream) *AntigravityGatewayService { + tokenProvider := NewAntigravityTokenProvider( + nil, + &antigravityCompatTokenCache{token: "fresh-oauth-token"}, + nil, + ) + return NewAntigravityGatewayService( + nil, + nil, + nil, + tokenProvider, + nil, + upstream, + NewSettingService(&antigravitySettingRepoStub{}, &config.Config{Gateway: cfg}), + nil, + ) +} + +func newAntigravityCompatAccount(accountType string) *Account { + return &Account{ + ID: 3757, + Name: "antigravity-compat", + Platform: PlatformAntigravity, + Type: accountType, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "stale-account-token", + "project_id": "project-3757", + "model_mapping": map[string]any{ + "gemini-3.1-pro-high": "gemini-3.1-pro-high", + "claude-sonnet-4-5": "claude-sonnet-4-5", + }, + }, + } +} + +func newAntigravityCompatContext(method, path string, body []byte) (*gin.Context, *httptest.ResponseRecorder) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(method, path, bytes.NewReader(body)) + return c, recorder +} + +func antigravityCompatSuccessResponse() *http.Response { + body := `data: {"response":{"responseId":"resp_3757","candidates":[{"content":{"parts":[{"text":"ok"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":3}}}` + "\n\n" + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "X-Request-Id": []string{"request-3757"}, + }, + Body: io.NopCloser(strings.NewReader(body)), + } +} + +func TestAntigravityCompatOAuthUsesNativeTokenAndRoute(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + path string + body []byte + call func(*AntigravityGatewayService, context.Context, *gin.Context, *Account, []byte) (*ForwardResult, error) + }{ + { + name: "chat completions", + path: "/v1/chat/completions", + body: []byte(`{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"Reply exactly: ok"}]}`), + call: func(svc *AntigravityGatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + return svc.ForwardAsChatCompletions(ctx, c, account, body, nil) + }, + }, + { + name: "responses", + path: "/v1/responses", + body: []byte(`{"model":"gemini-3.1-pro-high","input":"Reply exactly: ok"}`), + call: func(svc *AntigravityGatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + return svc.ForwardAsResponses(ctx, c, account, body, nil) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var authorization string + var upstreamPath string + var upstreamAlt string + upstream := &queuedHTTPUpstreamStub{ + responses: []*http.Response{antigravityCompatSuccessResponse()}, + onCall: func(req *http.Request, _ *queuedHTTPUpstreamStub) { + authorization = req.Header.Get("Authorization") + upstreamPath = req.URL.Path + upstreamAlt = req.URL.Query().Get("alt") + }, + } + svc := newAntigravityCompatService( + config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, + upstream, + ) + c, recorder := newAntigravityCompatContext(http.MethodPost, tt.path, tt.body) + + result, err := tt.call(svc, context.Background(), c, newAntigravityCompatAccount(AccountTypeOAuth), tt.body) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "Bearer fresh-oauth-token", authorization) + require.Equal(t, "/v1internal:streamGenerateContent", upstreamPath) + require.Equal(t, "sse", upstreamAlt) + require.Equal(t, "request-3757", result.RequestID) + require.Equal(t, 8, result.Usage.InputTokens) + require.Equal(t, 3, result.Usage.OutputTokens) + require.Equal(t, http.StatusOK, recorder.Code) + require.Contains(t, recorder.Body.String(), "ok") + if tt.name == "chat completions" { + require.Equal(t, "stop", gjson.Get(recorder.Body.String(), "choices.0.finish_reason").String()) + require.Equal(t, int64(8), gjson.Get(recorder.Body.String(), "usage.prompt_tokens").Int()) + require.Equal(t, int64(3), gjson.Get(recorder.Body.String(), "usage.completion_tokens").Int()) + } + }) + } +} + +func TestAntigravityCompatRejectsUnsupportedAccountType(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + path string + accountType string + call func(*AntigravityGatewayService, context.Context, *gin.Context, *Account, []byte) (*ForwardResult, error) + }{ + { + name: "chat completions upstream", + path: "/v1/chat/completions", + accountType: AccountTypeUpstream, + call: func(svc *AntigravityGatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + return svc.ForwardAsChatCompletions(ctx, c, account, body, nil) + }, + }, + { + name: "responses setup token", + path: "/v1/responses", + accountType: AccountTypeSetupToken, + call: func(svc *AntigravityGatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + return svc.ForwardAsResponses(ctx, c, account, body, nil) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body := []byte(`{"model":"gemini-3.1-pro-high"}`) + c, recorder := newAntigravityCompatContext(http.MethodPost, tt.path, body) + + result, err := tt.call(&AntigravityGatewayService{}, context.Background(), c, newAntigravityCompatAccount(tt.accountType), body) + + require.Error(t, err) + require.Nil(t, result) + require.Equal(t, http.StatusBadRequest, recorder.Code) + require.Contains(t, recorder.Body.String(), "native OAuth account required for antigravity compatibility mode") + }) + } +} + +func TestAntigravityCompatPreservesChatTokenLimit(t *testing.T) { + gin.SetMode(gin.TestMode) + tests := []struct { + name string + body string + want int64 + }{ + { + name: "legacy max_tokens below bridge floor", + body: `{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"ok"}],"max_tokens":8}`, + want: 8, + }, + { + name: "max_completion_tokens takes precedence", + body: `{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"ok"}],"max_tokens":8,"max_completion_tokens":13}`, + want: 13, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &queuedHTTPUpstreamStub{responses: []*http.Response{antigravityCompatSuccessResponse()}} + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, upstream) + body := []byte(tt.body) + c, _ := newAntigravityCompatContext(http.MethodPost, "/v1/chat/completions", body) + + result, err := svc.ForwardAsChatCompletions( + context.Background(), + c, + newAntigravityCompatAccount(AccountTypeOAuth), + body, + nil, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.requestBodies, 1) + require.Equal(t, tt.want, gjson.GetBytes(upstream.requestBodies[0], "request.generationConfig.maxOutputTokens").Int()) + }) + } +} + +func TestAntigravityCompatRoutesByMappedModelFamily(t *testing.T) { + gin.SetMode(gin.TestMode) + tests := []struct { + model string + wantSessionID bool + }{ + {model: "gemini-3.1-pro-high", wantSessionID: false}, + {model: "claude-sonnet-4-5", wantSessionID: true}, + } + + for _, tt := range tests { + t.Run(tt.model, func(t *testing.T) { + upstream := &queuedHTTPUpstreamStub{responses: []*http.Response{antigravityCompatSuccessResponse()}} + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, upstream) + body := []byte(`{"model":"` + tt.model + `","messages":[{"role":"user","content":"ok"}],"max_tokens":8}`) + c, _ := newAntigravityCompatContext(http.MethodPost, "/v1/chat/completions", body) + + result, err := svc.ForwardAsChatCompletions( + context.Background(), + c, + newAntigravityCompatAccount(AccountTypeOAuth), + body, + nil, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.requestBodies, 1) + require.Equal(t, tt.model, gjson.GetBytes(upstream.requestBodies[0], "model").String()) + require.Equal(t, tt.wantSessionID, gjson.GetBytes(upstream.requestBodies[0], "request.sessionId").Exists()) + }) + } +} + +func TestAntigravityCompatUnauthorizedIsCredentialFailure(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &queuedHTTPUpstreamStub{responses: []*http.Response{{ + StatusCode: http.StatusUnauthorized, + Header: http.Header{"X-Request-Id": []string{"auth-3757"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"Invalid bearer token"}}`)), + }}} + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, upstream) + body := []byte(`{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"ok"}]}`) + c, recorder := newAntigravityCompatContext(http.MethodPost, "/v1/chat/completions", body) + + result, err := svc.ForwardAsChatCompletions( + context.Background(), + c, + newAntigravityCompatAccount(AccountTypeOAuth), + body, + nil, + ) + + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureStageAccountAuth, failoverErr.Stage) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, AntigravityCredentialRejectedReason, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Equal(t, http.StatusBadGateway, failoverErr.ClientStatusCode) + require.Equal(t, AntigravityCredentialRejectedClientMessage, failoverErr.ClientMessage) + require.Equal(t, "auth-3757", failoverErr.ResponseHeaders.Get("X-Request-Id")) + require.Empty(t, recorder.Body.String()) +} + +func TestAntigravityCompatEmptyStreamTriggersFailover(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + run func(*AntigravityGatewayService, *gin.Context, *http.Response) (*antigravityStreamResult, error) + }{ + { + name: "chat completions", + run: func(svc *AntigravityGatewayService, c *gin.Context, resp *http.Response) (*antigravityStreamResult, error) { + return svc.handleChatCompletionsStreamingFromAntigravity(c, resp, time.Now(), "gemini-3.1-pro-high", true) + }, + }, + { + name: "responses", + run: func(svc *AntigravityGatewayService, c *gin.Context, resp *http.Response) (*antigravityStreamResult, error) { + return svc.handleResponsesStreamingFromAntigravity(c, resp, time.Now(), "gemini-3.1-pro-high") + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, nil) + c, recorder := newAntigravityCompatContext(http.MethodPost, "/", nil) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader("data: malformed\n\ndata: [DONE]\n\n")), + } + + result, err := tt.run(svc, c, resp) + + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.True(t, failoverErr.RetryableOnSameAccount) + require.Empty(t, recorder.Body.String()) + require.Empty(t, recorder.Header().Get("Content-Type")) + }) + } +} + +func TestAntigravityCompatUsageOnlyStreamTriggersFailover(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + run func(*AntigravityGatewayService, *gin.Context, *http.Response) (*antigravityStreamResult, error) + }{ + { + name: "chat completions", + run: func(svc *AntigravityGatewayService, c *gin.Context, resp *http.Response) (*antigravityStreamResult, error) { + return svc.handleChatCompletionsStreamingFromAntigravity(c, resp, time.Now(), "gemini-3.1-pro-high", true) + }, + }, + { + name: "responses", + run: func(svc *AntigravityGatewayService, c *gin.Context, resp *http.Response) (*antigravityStreamResult, error) { + return svc.handleResponsesStreamingFromAntigravity(c, resp, time.Now(), "gemini-3.1-pro-high") + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, nil) + c, recorder := newAntigravityCompatContext(http.MethodPost, "/", nil) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + `data: {"response":{"responseId":"resp_3757","usageMetadata":{"promptTokenCount":8}}}` + "\n\n", + )), + } + + result, err := tt.run(svc, c, resp) + + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.True(t, failoverErr.RetryableOnSameAccount) + require.Empty(t, recorder.Body.String()) + require.Empty(t, recorder.Header().Get("Content-Type")) + }) + } +} + +func TestAntigravityCompatUsageOnlyNonStreamingTriggersFailover(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + path string + body []byte + call func(*AntigravityGatewayService, context.Context, *gin.Context, *Account, []byte) (*ForwardResult, error) + }{ + { + name: "chat completions", + path: "/v1/chat/completions", + body: []byte(`{"model":"gemini-3.1-pro-high","messages":[{"role":"user","content":"ok"}]}`), + call: func(svc *AntigravityGatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + return svc.ForwardAsChatCompletions(ctx, c, account, body, nil) + }, + }, + { + name: "responses", + path: "/v1/responses", + body: []byte(`{"model":"gemini-3.1-pro-high","input":"ok"}`), + call: func(svc *AntigravityGatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + return svc.ForwardAsResponses(ctx, c, account, body, nil) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &queuedHTTPUpstreamStub{responses: []*http.Response{{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + `data: {"response":{"responseId":"resp_3757","usageMetadata":{"promptTokenCount":8}}}` + "\n\n", + )), + }}} + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, upstream) + c, recorder := newAntigravityCompatContext(http.MethodPost, tt.path, tt.body) + + result, err := tt.call( + svc, + context.Background(), + c, + newAntigravityCompatAccount(AccountTypeOAuth), + tt.body, + ) + + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.True(t, failoverErr.RetryableOnSameAccount) + require.Empty(t, recorder.Body.String()) + require.Empty(t, recorder.Header().Get("Content-Type")) + }) + } +} + +func TestAntigravityCompatChatStreamMapsToolCallAndUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, nil) + c, recorder := newAntigravityCompatContext(http.MethodPost, "/v1/chat/completions", nil) + body := `data: {"response":{"responseId":"resp_3757","candidates":[{"content":{"parts":[{"functionCall":{"id":"call_3757","name":"get_weather","args":{"city":"Tokyo"}}}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":3}}}` + "\n\n" + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + } + + result, err := svc.handleChatCompletionsStreamingFromAntigravity( + c, + resp, + time.Now(), + "gemini-3.1-pro-high", + true, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 8, result.usage.InputTokens) + require.Equal(t, 3, result.usage.OutputTokens) + require.Contains(t, recorder.Body.String(), `"tool_calls"`) + require.Contains(t, recorder.Body.String(), `"get_weather"`) + require.Contains(t, recorder.Body.String(), `"finish_reason":"tool_calls"`) + require.Contains(t, recorder.Body.String(), `"prompt_tokens":8`) + require.Equal(t, 1, strings.Count(recorder.Body.String(), "data: [DONE]")) +} + +func TestAntigravityCompatFirstEventTimeoutTriggersFailover(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newAntigravityCompatService( + config.GatewayConfig{MaxLineSize: defaultMaxLineSize, StreamDataIntervalTimeout: 1}, + nil, + ) + c, recorder := newAntigravityCompatContext(http.MethodPost, "/v1/chat/completions", nil) + reader, writer := io.Pipe() + resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: reader} + type outcome struct { + result *antigravityStreamResult + err error + } + done := make(chan outcome, 1) + + go func() { + result, err := svc.handleChatCompletionsStreamingFromAntigravity(c, resp, time.Now(), "gemini-3.1-pro-high", false) + done <- outcome{result: result, err: err} + }() + + select { + case got := <-done: + require.Nil(t, got.result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, got.err, &failoverErr) + require.True(t, failoverErr.RetryableOnSameAccount) + require.Empty(t, recorder.Header().Get("Content-Type")) + case <-time.After(2 * time.Second): + _ = writer.Close() + _ = reader.Close() + t.Fatal("compat stream ignored StreamDataIntervalTimeout") + } + _ = writer.Close() + _ = reader.Close() +} + +func TestAntigravityCompatClientDisconnectDrainsUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, nil) + c, _ := newAntigravityCompatContext(http.MethodPost, "/v1/chat/completions", nil) + c.Writer = &antigravityFailingWriter{ResponseWriter: c.Writer, failAfter: 0} + body := strings.Join([]string{ + `data: {"response":{"responseId":"resp_3757","candidates":[{"content":{"parts":[{"text":"partial"}]}}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":1}}}`, + "", + `data: {"response":{"responseId":"resp_3757","candidates":[{"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":15}}}`, + "", + }, "\n") + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(body)), + } + + result, err := svc.handleChatCompletionsStreamingFromAntigravity(c, resp, time.Now(), "gemini-3.1-pro-high", false) + + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, result.clientDisconnect) + require.Equal(t, 15, result.usage.OutputTokens) +} + +func TestAntigravityCompatStreamErrorCommitsSingleTerminalFrame(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newAntigravityCompatService(config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, nil) + c, recorder := newAntigravityCompatContext(http.MethodPost, "/v1/responses", nil) + body := []byte(`data: {"response":{"responseId":"resp_3757","candidates":[{"content":{"parts":[{"text":"partial"}]}}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":1}}}` + "\n\n") + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: &antigravityCompatErrorReader{ + data: body, + err: io.ErrUnexpectedEOF, + }, + } + + result, err := svc.handleResponsesStreamingFromAntigravity(c, resp, time.Now(), "gemini-3.1-pro-high") + + require.Error(t, err) + require.Nil(t, result) + require.True(t, IsResponseCommitted(c)) + require.Equal(t, 1, strings.Count(recorder.Body.String(), "event: error")) +} + +func TestAntigravityCompatKeepaliveAfterFirstEvent(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newAntigravityCompatService( + config.GatewayConfig{MaxLineSize: defaultMaxLineSize, StreamKeepaliveInterval: 1}, + nil, + ) + c, recorder := newAntigravityCompatContext(http.MethodPost, "/v1/responses", nil) + reader, writer := io.Pipe() + resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: reader} + done := make(chan error, 1) + + go func() { + _, err := svc.handleResponsesStreamingFromAntigravity(c, resp, time.Now(), "gemini-3.1-pro-high") + done <- err + }() + _, err := io.WriteString( + writer, + `data: {"response":{"responseId":"resp_3757","candidates":[{"content":{"parts":[{"text":"partial"}]}}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":1}}}`+"\n\n", + ) + require.NoError(t, err) + time.Sleep(1200 * time.Millisecond) + require.NoError(t, writer.Close()) + require.NoError(t, <-done) + require.Contains(t, recorder.Body.String(), ": ping\n\n") + require.Contains(t, recorder.Header().Get("Content-Type"), "text/event-stream") + require.NoError(t, reader.Close()) +} diff --git a/backend/internal/service/antigravity_gateway_streaming.go b/backend/internal/service/antigravity_gateway_streaming.go index 1a6c59f617..2810162c45 100644 --- a/backend/internal/service/antigravity_gateway_streaming.go +++ b/backend/internal/service/antigravity_gateway_streaming.go @@ -25,10 +25,11 @@ type antigravityStreamResult struct { // antigravityClientWriter 封装流式响应的客户端写入,自动检测断开并标记。 // 断开后所有写入操作变为 no-op,调用方通过 Disconnected() 判断是否继续 drain 上游。 type antigravityClientWriter struct { - w gin.ResponseWriter - flusher http.Flusher - disconnected bool - prefix string // 日志前缀,标识来源方法 + w gin.ResponseWriter + flusher http.Flusher + disconnected bool + prefix string // 日志前缀,标识来源方法 + beforeFirstWrite func() } func newAntigravityClientWriter(w gin.ResponseWriter, flusher http.Flusher, prefix string) *antigravityClientWriter { @@ -40,6 +41,7 @@ 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 @@ -53,6 +55,7 @@ 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 @@ -63,6 +66,15 @@ func (cw *antigravityClientWriter) Fprintf(format string, args ...any) bool { 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) @@ -751,9 +763,9 @@ func (s *AntigravityGatewayService) writeGoogleError(c *gin.Context, status int, return fmt.Errorf("%s", message) } -// handleClaudeStreamToNonStreaming 收集上游流式响应,转换为 Claude 非流式格式返回 +// collectClaudeStreamResponse 收集上游流式响应,转换为 Claude 非流式格式返回 // 用于处理客户端非流式请求但上游只支持流式的情况 -func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) (*antigravityStreamResult, error) { +func (s *AntigravityGatewayService) collectClaudeStreamResponse(resp *http.Response, startTime time.Time, originalModel string) ([]byte, *antigravityStreamResult, error) { scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { @@ -766,6 +778,7 @@ func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Cont 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 @@ -827,7 +840,7 @@ func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Cont 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, ev.err + return nil, nil, ev.err } line := ev.line @@ -853,21 +866,23 @@ func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Cont continue } - // 记录首 token 时间 - if firstTokenMs == nil { - ms := int(time.Since(startTime).Milliseconds()) - firstTokenMs = &ms - } - last = parsed // 保留最后一个有 parts 的响应,并收集所有 parts - if parts := extractGeminiParts(parsed); len(parts) > 0 { + 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)) @@ -875,24 +890,24 @@ func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Cont continue } logger.LegacyPrintf("service.antigravity_gateway", "Stream data interval timeout (antigravity claude non-stream)") - return nil, fmt.Errorf("stream data interval timeout") + return nil, nil, fmt.Errorf("stream data interval timeout") } } returnResponse: - // 选择最后一个有效响应 - finalResponse := pickGeminiCollectResult(last, lastWithParts) - // 处理空响应情况 — 触发同账号重试 + failover 切换账号 - if last == nil && lastWithParts == nil { + if !meaningfulResponse { logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] warning: empty stream response (claude non-stream), triggering failover") - return nil, &UpstreamFailoverError{ + 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) @@ -901,27 +916,56 @@ returnResponse: // 序列化为 JSON(Gemini 格式) geminiBody, err := json.Marshal(finalResponse) if err != nil { - return nil, fmt.Errorf("failed to marshal gemini response: %w", err) + 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, s.writeClaudeError(c, http.StatusBadGateway, "upstream_error", "Failed to parse upstream response") + return nil, nil, fmt.Errorf("failed to parse upstream response: %w", err) } - c.Data(http.StatusOK, "application/json", claudeResp) - // 转换为 service.ClaudeUsage usage := &ClaudeUsage{ InputTokens: agUsage.InputTokens, OutputTokens: agUsage.OutputTokens, CacheCreationInputTokens: agUsage.CacheCreationInputTokens, CacheReadInputTokens: agUsage.CacheReadInputTokens, + ImageOutputTokens: agUsage.ImageOutputTokens, } - return &antigravityStreamResult{usage: usage, firstTokenMs: firstTokenMs}, nil + 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(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 转换) From 3e081061165a9a51378f41aee43295a50c6060a6 Mon Sep 17 00:00:00 2001 From: chinnsenn Date: Sun, 26 Jul 2026 22:09:26 +0900 Subject: [PATCH 2/3] fix(gemini): preserve Hermes web search functions --- .../service/gemini_messages_compat_service.go | 17 ++--- .../gemini_messages_compat_service_test.go | 64 +++++++++++++++++++ 2 files changed, 70 insertions(+), 11 deletions(-) diff --git a/backend/internal/service/gemini_messages_compat_service.go b/backend/internal/service/gemini_messages_compat_service.go index 1a3821847d..30f635643f 100644 --- a/backend/internal/service/gemini_messages_compat_service.go +++ b/backend/internal/service/gemini_messages_compat_service.go @@ -3428,17 +3428,12 @@ func normalizeGeminiRequestForAIStudio(body []byte) []byte { func isClaudeWebSearchToolMap(tool map[string]any) bool { toolType, _ := tool["type"].(string) - if strings.HasPrefix(toolType, "web_search") || toolType == "google_search" { - return true - } - - name, _ := tool["name"].(string) - switch strings.TrimSpace(name) { - case "web_search", "google_search", "web_search_20250305": - return true - default: - return false - } + // A function named web_search is still a client-side function. This is + // especially important for Chat Completions clients such as Hermes, whose + // built-in runtime tools are represented as ordinary function tools. + // Promote only explicitly typed server-side search tools to Gemini's + // built-in googleSearch tool. + return strings.HasPrefix(toolType, "web_search") || toolType == "google_search" } // cleanToolSchema 清理工具的 JSON Schema,移除 Gemini 不支持的字段 diff --git a/backend/internal/service/gemini_messages_compat_service_test.go b/backend/internal/service/gemini_messages_compat_service_test.go index b2fc1c110b..24061593fb 100644 --- a/backend/internal/service/gemini_messages_compat_service_test.go +++ b/backend/internal/service/gemini_messages_compat_service_test.go @@ -170,6 +170,70 @@ func TestGeminiForwardAsChatCompletions_StreamsOpenAIChunksFromGeminiSSE(t *test require.Contains(t, out, "data: [DONE]") } +func TestGeminiForwardAsChatCompletions_FunctionNamedWebSearchStaysClientSide(t *testing.T) { + gin.SetMode(gin.TestMode) + + httpStub := &geminiCompatHTTPUpstreamStub{ + response: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"candidates":[{"content":{"parts":[{"text":"hello"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":2,"candidatesTokenCount":1}}`, + )), + }, + } + svc := &GeminiMessagesCompatService{ + httpUpstream: httpStub, + cfg: &config.Config{}, + } + account := &Account{ + ID: 103, + Platform: PlatformGemini, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "gemini-api-key", + }, + Concurrency: 1, + } + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + body := []byte(`{ + "model":"gemini-3.6-flash-high", + "messages":[{"role":"user","content":"search and read"}], + "tools":[ + {"type":"function","function":{"name":"web_search","description":"Search through the Hermes client","parameters":{"type":"object","properties":{"query":{"type":"string"}},"required":["query"]}}}, + {"type":"function","function":{"name":"read_file","description":"Read a local file","parameters":{"type":"object","properties":{"path":{"type":"string"}},"required":["path"]}}} + ] + }`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, http.StatusOK, rec.Code) + require.NotNil(t, httpStub.lastReq) + + postedBody, err := io.ReadAll(httpStub.lastReq.Body) + require.NoError(t, err) + + var posted map[string]any + require.NoError(t, json.Unmarshal(postedBody, &posted)) + tools, ok := posted["tools"].([]any) + require.True(t, ok) + require.Len(t, tools, 1, "Chat Completions function tools must not be promoted to Gemini built-ins by name") + + functionTool, ok := tools[0].(map[string]any) + require.True(t, ok) + functionDecls, ok := functionTool["functionDeclarations"].([]any) + require.True(t, ok) + require.Len(t, functionDecls, 2) + require.Equal(t, "web_search", functionDecls[0].(map[string]any)["name"]) + require.Equal(t, "read_file", functionDecls[1].(map[string]any)["name"]) + require.NotContains(t, functionTool, "googleSearch") + require.NotContains(t, functionTool, "google_search") +} + // TestConvertClaudeToolsToGeminiTools_CustomType 测试custom类型工具转换 func TestConvertClaudeToolsToGeminiTools_CustomType(t *testing.T) { tests := []struct { From cc84cd8b4c0aaae83544e46bb76c1f0ae8e49c5e Mon Sep 17 00:00:00 2001 From: chinnsenn Date: Mon, 27 Jul 2026 13:19:00 +0900 Subject: [PATCH 3/3] test(gemini): check function declaration assertions --- .../service/gemini_messages_compat_service_test.go | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/backend/internal/service/gemini_messages_compat_service_test.go b/backend/internal/service/gemini_messages_compat_service_test.go index 24061593fb..59f83a6835 100644 --- a/backend/internal/service/gemini_messages_compat_service_test.go +++ b/backend/internal/service/gemini_messages_compat_service_test.go @@ -228,8 +228,12 @@ func TestGeminiForwardAsChatCompletions_FunctionNamedWebSearchStaysClientSide(t functionDecls, ok := functionTool["functionDeclarations"].([]any) require.True(t, ok) require.Len(t, functionDecls, 2) - require.Equal(t, "web_search", functionDecls[0].(map[string]any)["name"]) - require.Equal(t, "read_file", functionDecls[1].(map[string]any)["name"]) + webSearchDecl, ok := functionDecls[0].(map[string]any) + require.True(t, ok) + readFileDecl, ok := functionDecls[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "web_search", webSearchDecl["name"]) + require.Equal(t, "read_file", readFileDecl["name"]) require.NotContains(t, functionTool, "googleSearch") require.NotContains(t, functionTool, "google_search") }