mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 17:08:33 +08:00
fix(antigravity):
1. harden OpenAI compatibility forwarding 2. reject usage-only non-stream responses
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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",
|
||||
)
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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 转换)
|
||||
|
||||
Reference in New Issue
Block a user