diff --git a/backend/internal/service/antigravity_gateway_gemini.go b/backend/internal/service/antigravity_gateway_gemini.go index 2b622afd37..77a7be3df7 100644 --- a/backend/internal/service/antigravity_gateway_gemini.go +++ b/backend/internal/service/antigravity_gateway_gemini.go @@ -91,7 +91,7 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate) return nil, s.writeGoogleError(c, http.StatusForbidden, fmt.Sprintf("model %s not in whitelist", originalModel)) } - billingModel := mappedModel + forwardedModel := mappedModel // 获取 access_token if s.tokenProvider == nil { @@ -204,7 +204,7 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co if err == nil && fallbackResp.StatusCode < 400 { _ = resp.Body.Close() resp = fallbackResp - billingModel = fallbackModel + forwardedModel = fallbackModel } else if fallbackResp != nil { _ = fallbackResp.Body.Close() } @@ -436,7 +436,7 @@ handleSuccess: RequestID: requestID, Usage: *usage, Model: originalModel, - UpstreamModel: billingModel, + UpstreamModel: forwardedModel, UpstreamResponseModel: observedUpstreamResponseModel(c), UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), Stream: stream, diff --git a/backend/internal/service/antigravity_gateway_service_test.go b/backend/internal/service/antigravity_gateway_service_test.go index 02c2d85e36..00fee73a11 100644 --- a/backend/internal/service/antigravity_gateway_service_test.go +++ b/backend/internal/service/antigravity_gateway_service_test.go @@ -200,13 +200,18 @@ func (c *recordingInternal500CounterCache) ResetInternal500Count(_ context.Conte return nil } -type antigravitySettingRepoStub struct{} +type antigravitySettingRepoStub struct { + values map[string]string +} func (s *antigravitySettingRepoStub) Get(ctx context.Context, key string) (*Setting, error) { panic("unexpected Get call") } func (s *antigravitySettingRepoStub) GetValue(ctx context.Context, key string) (string, error) { + if value, ok := s.values[key]; ok { + return value, nil + } return "", ErrSettingNotFound } @@ -862,6 +867,73 @@ func TestAntigravityGatewayService_ForwardGemini_BillsWithMappedModel(t *testing require.Equal(t, mappedModel, result.UpstreamModel) } +func TestAntigravityGatewayService_ForwardGemini_FallbackReportsActualUpstreamModel(t *testing.T) { + gin.SetMode(gin.TestMode) + writer := httptest.NewRecorder() + c, _ := gin.CreateTestContext(writer) + + body := []byte(`{"contents":[{"role":"user","parts":[{"text":"hello"}]}]}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1beta/models/gemini-primary:generateContent", bytes.NewReader(body)) + + const ( + originalModel = "gemini-primary" + mappedModel = "gemini-primary-upstream" + fallbackModel = "gemini-fallback-upstream" + ) + upstream := &queuedHTTPUpstreamStub{responses: []*http.Response{ + { + StatusCode: http.StatusNotFound, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"code":404,"message":"model not found"}}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + `data: {"response":{"modelVersion":"gemini-fallback-upstream","candidates":[{"content":{"parts":[{"text":"ok"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":3}}}` + "\n\n", + )), + }, + }} + settings := &antigravitySettingRepoStub{values: map[string]string{ + SettingKeyEnableModelFallback: "true", + SettingKeyFallbackModelAntigravity: fallbackModel, + }} + svc := &AntigravityGatewayService{ + settingService: NewSettingService(settings, &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}), + tokenProvider: &AntigravityTokenProvider{}, + httpUpstream: upstream, + } + account := &Account{ + ID: 9, + Name: "acc-gemini-fallback", + Platform: PlatformAntigravity, + Type: AccountTypeOAuth, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "token", + "project_id": "proj", + "model_mapping": map[string]any{ + originalModel: mappedModel, + }, + }, + } + + result, err := svc.ForwardGemini(context.Background(), c, account, originalModel, "generateContent", true, body, false) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, originalModel, result.Model) + require.Equal(t, fallbackModel, result.UpstreamModel) + require.Equal(t, fallbackModel, result.UpstreamResponseModel) + require.False(t, result.UpstreamResponseModelConflict) + mismatch := upstreamModelMismatch(result.UpstreamModel, result.UpstreamResponseModel) + require.NotNil(t, mismatch) + require.False(t, *mismatch) + require.Len(t, upstream.requestBodies, 2) + require.Contains(t, string(upstream.requestBodies[0]), `"model":"`+mappedModel+`"`) + require.Contains(t, string(upstream.requestBodies[1]), `"model":"`+fallbackModel+`"`) +} + func TestAntigravityGatewayService_ForwardGemini_RetriesCorruptedThoughtSignature(t *testing.T) { gin.SetMode(gin.TestMode) writer := httptest.NewRecorder() diff --git a/backend/internal/service/antigravity_gateway_streaming.go b/backend/internal/service/antigravity_gateway_streaming.go index 0b3a672853..2b68fc115a 100644 --- a/backend/internal/service/antigravity_gateway_streaming.go +++ b/backend/internal/service/antigravity_gateway_streaming.go @@ -35,11 +35,11 @@ func (s *AntigravityGatewayService) observeAntigravityGeminiSSELine(c *gin.Conte if payload == "" || payload == "[DONE]" { return } - raw := []byte(payload) - if inner, err := s.unwrapV1InternalResponse(raw); err == nil && len(inner) > 0 { - raw = inner - } - observer.ObserveGemini(raw) + // Observe the original payload: ObserveGemini supports both the v1internal + // wrapper and direct Gemini response shapes. The main stream handler will + // unwrap the same line for business processing, so unwrapping here would be + // duplicate work on every SSE event. + observer.ObserveGemini([]byte(payload)) } // antigravityClientWriter 封装流式响应的客户端写入,自动检测断开并标记。 diff --git a/backend/internal/service/upstream_response_model.go b/backend/internal/service/upstream_response_model.go index 2452c382d0..cc282c228f 100644 --- a/backend/internal/service/upstream_response_model.go +++ b/backend/internal/service/upstream_response_model.go @@ -53,34 +53,21 @@ func normalizeObservedUpstreamResponseModel(model string) string { } func (o *upstreamResponseModelObserver) ObserveOpenAI(payload []byte, eventType string) { - if len(payload) == 0 || !gjson.ValidBytes(payload) { - return - } - model := firstTrimmedGJSONModel( - gjson.GetBytes(payload, "response.model"), - gjson.GetBytes(payload, "model"), - ) + model := firstValidTrimmedGJSONModel(payload, "response.model", "model") o.Observe(model, isUpstreamResponseModelTerminalEvent(eventType)) } func (o *upstreamResponseModelObserver) ObserveAnthropic(payload []byte) { - if len(payload) == 0 || !gjson.ValidBytes(payload) { - return - } - model := firstTrimmedGJSONModel( - gjson.GetBytes(payload, "message.model"), - gjson.GetBytes(payload, "model"), - ) + model := firstValidTrimmedGJSONModel(payload, "message.model", "model") o.Observe(model, false) } func (o *upstreamResponseModelObserver) ObserveGemini(payload []byte) { - if len(payload) == 0 || !gjson.ValidBytes(payload) { - return - } - model := firstTrimmedGJSONModel( - gjson.GetBytes(payload, "modelVersion"), - gjson.GetBytes(payload, "response.modelVersion"), + model := firstValidTrimmedGJSONModel( + payload, + "modelVersion", + "response.modelVersion", + "response.response.modelVersion", ) // Gemini streaming has no universal terminal event carrying modelVersion; // treating each declaration as terminal retains the latest chunk. @@ -139,12 +126,22 @@ func observeOpenAISSEBody(observer *upstreamResponseModelObserver, body string) }) } -func firstTrimmedGJSONModel(values ...gjson.Result) string { - for _, value := range values { +func firstValidTrimmedGJSONModel(payload []byte, paths ...string) string { + if len(payload) == 0 { + return "" + } + for _, path := range paths { + value := gjson.GetBytes(payload, path) if !value.Exists() || value.Type != gjson.String { continue } if model := strings.TrimSpace(value.String()); model != "" { + // Validate only after finding a candidate. This avoids a full validation + // pass on the common model-free delta path while still rejecting malformed + // payloads that appear to declare a model. + if !gjson.ValidBytes(payload) { + return "" + } return model } } diff --git a/backend/internal/service/upstream_response_model_bench_test.go b/backend/internal/service/upstream_response_model_bench_test.go new file mode 100644 index 0000000000..93d6854f72 --- /dev/null +++ b/backend/internal/service/upstream_response_model_bench_test.go @@ -0,0 +1,98 @@ +package service + +import ( + "strings" + "testing" + + "github.com/tidwall/gjson" +) + +var ( + upstreamResponseModelBenchmarkSink string + + upstreamResponseModelBenchmarkTerminal = []byte(`{"type":"response.completed","response":{"id":"resp_123","model":"gpt-5.5-2026-04-23"},"usage":{"input_tokens":128,"output_tokens":64}}`) + upstreamResponseModelBenchmarkDelta = []byte(`{"type":"response.output_text.delta","delta":"hello"}`) + upstreamResponseModelBenchmarkWrapper = []byte(`{"response":{"response":{"modelVersion":"gemini-3-pro","candidates":[]}}}`) +) + +func BenchmarkUpstreamResponseModelOpenAI(b *testing.B) { + tests := []struct { + name string + payload []byte + }{ + {name: "terminal_with_model", payload: upstreamResponseModelBenchmarkTerminal}, + {name: "delta_without_model", payload: upstreamResponseModelBenchmarkDelta}, + } + + for _, tt := range tests { + b.Run(tt.name, func(b *testing.B) { + b.Run("legacy", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + upstreamResponseModelBenchmarkSink = benchmarkLegacyOpenAIModel(tt.payload) + } + }) + b.Run("optimized", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + upstreamResponseModelBenchmarkSink = firstValidTrimmedGJSONModel(tt.payload, "response.model", "model") + } + }) + }) + } +} + +func BenchmarkUpstreamResponseModelAntigravityWrapper(b *testing.B) { + b.Run("legacy", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + upstreamResponseModelBenchmarkSink = benchmarkLegacyWrappedGeminiModel(upstreamResponseModelBenchmarkWrapper) + } + }) + b.Run("optimized", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + upstreamResponseModelBenchmarkSink = firstValidTrimmedGJSONModel( + upstreamResponseModelBenchmarkWrapper, + "modelVersion", + "response.modelVersion", + "response.response.modelVersion", + ) + } + }) +} + +func benchmarkLegacyOpenAIModel(payload []byte) string { + if len(payload) == 0 || !gjson.ValidBytes(payload) { + return "" + } + return benchmarkFirstTrimmedGJSONModel( + gjson.GetBytes(payload, "response.model"), + gjson.GetBytes(payload, "model"), + ) +} + +func benchmarkLegacyWrappedGeminiModel(payload []byte) string { + if inner := gjson.GetBytes(payload, "response"); inner.Exists() { + payload = []byte(inner.Raw) + } + if len(payload) == 0 || !gjson.ValidBytes(payload) { + return "" + } + return benchmarkFirstTrimmedGJSONModel( + gjson.GetBytes(payload, "modelVersion"), + gjson.GetBytes(payload, "response.modelVersion"), + ) +} + +func benchmarkFirstTrimmedGJSONModel(values ...gjson.Result) string { + for _, value := range values { + if !value.Exists() || value.Type != gjson.String { + continue + } + if model := strings.TrimSpace(value.String()); model != "" { + return model + } + } + return "" +} diff --git a/backend/internal/service/upstream_response_model_test.go b/backend/internal/service/upstream_response_model_test.go index 37d382e851..a831e74e65 100644 --- a/backend/internal/service/upstream_response_model_test.go +++ b/backend/internal/service/upstream_response_model_test.go @@ -67,6 +67,56 @@ func TestObserveOpenAISSEBodyIgnoresMalformedPayload(t *testing.T) { require.False(t, observer.Conflict()) } +func TestObserveAntigravityGeminiSSELineReadsWrapperModelWithoutUnwrap(t *testing.T) { + tests := []struct { + name string + payload string + want string + }{ + { + name: "top-level sibling", + payload: `{"modelVersion":"gemini-3-pro","response":{"candidates":[]}}`, + want: "gemini-3-pro", + }, + { + name: "single wrapper", + payload: `{"response":{"modelVersion":"gemini-3-pro","candidates":[]}}`, + want: "gemini-3-pro", + }, + { + name: "nested response after one wrapper", + payload: `{"response":{"response":{"modelVersion":"gemini-3-pro","candidates":[]}}}`, + want: "gemini-3-pro", + }, + { + name: "outer declaration takes precedence", + payload: `{"modelVersion":"gemini-outer","response":{"modelVersion":"gemini-inner","candidates":[]}}`, + want: "gemini-outer", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(nil) + beginUpstreamResponseModelObservation(c) + + svc := &AntigravityGatewayService{} + svc.observeAntigravityGeminiSSELine(c, "data: "+tt.payload) + + require.Equal(t, tt.want, observedUpstreamResponseModel(c)) + require.False(t, observedUpstreamResponseModelConflict(c)) + }) + } +} + +func TestUpstreamResponseModelObserverRejectsMalformedJSONWithModelField(t *testing.T) { + observer := &upstreamResponseModelObserver{} + observer.ObserveOpenAI([]byte(`{"response":{"model":"gpt-5.4"}`), "response.completed") + + require.Empty(t, observer.Model()) +} + func TestUpstreamResponseModelObserverBoundsUntrustedModelName(t *testing.T) { observer := &upstreamResponseModelObserver{} observer.Observe(" "+strings.Repeat("模", upstreamResponseModelMaxLength+1)+" ", false)