Skip to content

Commit bac4080

Browse files
fix(provider): preserve requested model in antigravity and sora
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
1 parent 4edcfe1 commit bac4080

4 files changed

Lines changed: 46 additions & 29 deletions

File tree

backend/internal/service/antigravity_gateway_service.go

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1742,7 +1742,8 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context,
17421742
return &ForwardResult{
17431743
RequestID: requestID,
17441744
Usage: *usage,
1745-
Model: billingModel, // 使用映射模型用于计费和日志
1745+
Model: originalModel,
1746+
UpstreamModel: billingModel,
17461747
Stream: claudeReq.Stream,
17471748
Duration: time.Since(startTime),
17481749
FirstTokenMs: firstTokenMs,
@@ -2435,7 +2436,8 @@ handleSuccess:
24352436
return &ForwardResult{
24362437
RequestID: requestID,
24372438
Usage: *usage,
2438-
Model: billingModel,
2439+
Model: originalModel,
2440+
UpstreamModel: billingModel,
24392441
Stream: stream,
24402442
Duration: time.Since(startTime),
24412443
FirstTokenMs: firstTokenMs,

backend/internal/service/antigravity_gateway_service_test.go

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -542,7 +542,8 @@ func TestAntigravityGatewayService_Forward_BillsWithMappedModel(t *testing.T) {
542542
result, err := svc.Forward(context.Background(), c, account, body, false)
543543
require.NoError(t, err)
544544
require.NotNil(t, result)
545-
require.Equal(t, mappedModel, result.Model)
545+
require.Equal(t, "claude-sonnet-4-5", result.Model)
546+
require.Equal(t, mappedModel, result.UpstreamModel)
546547
}
547548

548549
// TestAntigravityGatewayService_ForwardGemini_BillsWithMappedModel
@@ -594,7 +595,8 @@ func TestAntigravityGatewayService_ForwardGemini_BillsWithMappedModel(t *testing
594595
result, err := svc.ForwardGemini(context.Background(), c, account, "gemini-2.5-flash", "generateContent", true, body, false)
595596
require.NoError(t, err)
596597
require.NotNil(t, result)
597-
require.Equal(t, mappedModel, result.Model)
598+
require.Equal(t, "gemini-2.5-flash", result.Model)
599+
require.Equal(t, mappedModel, result.UpstreamModel)
598600
}
599601

600602
func TestAntigravityGatewayService_ForwardGemini_RetriesCorruptedThoughtSignature(t *testing.T) {
@@ -664,7 +666,8 @@ func TestAntigravityGatewayService_ForwardGemini_RetriesCorruptedThoughtSignatur
664666
result, err := svc.ForwardGemini(context.Background(), c, account, originalModel, "streamGenerateContent", true, body, false)
665667
require.NoError(t, err)
666668
require.NotNil(t, result)
667-
require.Equal(t, mappedModel, result.Model)
669+
require.Equal(t, originalModel, result.Model)
670+
require.Equal(t, mappedModel, result.UpstreamModel)
668671
require.Len(t, upstream.requestBodies, 2, "signature error should trigger exactly one retry")
669672

670673
firstReq := string(upstream.requestBodies[0])

backend/internal/service/sora_gateway_service.go

Lines changed: 30 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -148,10 +148,13 @@ func (s *SoraGatewayService) Forward(ctx context.Context, c *gin.Context, accoun
148148
s.writeSoraError(c, http.StatusBadRequest, "invalid_request_error", "model is required", clientStream)
149149
return nil, errors.New("model is required")
150150
}
151+
originalModel := reqModel
151152

152153
mappedModel := account.GetMappedModel(reqModel)
154+
var upstreamModel string
153155
if mappedModel != "" && mappedModel != reqModel {
154156
reqModel = mappedModel
157+
upstreamModel = mappedModel
155158
}
156159

157160
modelCfg, ok := GetSoraModelConfig(reqModel)
@@ -213,13 +216,14 @@ func (s *SoraGatewayService) Forward(ctx context.Context, c *gin.Context, accoun
213216
c.JSON(http.StatusOK, buildSoraNonStreamResponse(content, reqModel))
214217
}
215218
return &ForwardResult{
216-
RequestID: "",
217-
Model: reqModel,
218-
Stream: clientStream,
219-
Duration: time.Since(startTime),
220-
FirstTokenMs: firstTokenMs,
221-
Usage: ClaudeUsage{},
222-
MediaType: "prompt",
219+
RequestID: "",
220+
Model: originalModel,
221+
UpstreamModel: upstreamModel,
222+
Stream: clientStream,
223+
Duration: time.Since(startTime),
224+
FirstTokenMs: firstTokenMs,
225+
Usage: ClaudeUsage{},
226+
MediaType: "prompt",
223227
}, nil
224228
}
225229

@@ -269,13 +273,14 @@ func (s *SoraGatewayService) Forward(ctx context.Context, c *gin.Context, accoun
269273
c.JSON(http.StatusOK, resp)
270274
}
271275
return &ForwardResult{
272-
RequestID: "",
273-
Model: reqModel,
274-
Stream: clientStream,
275-
Duration: time.Since(startTime),
276-
FirstTokenMs: firstTokenMs,
277-
Usage: ClaudeUsage{},
278-
MediaType: "prompt",
276+
RequestID: "",
277+
Model: originalModel,
278+
UpstreamModel: upstreamModel,
279+
Stream: clientStream,
280+
Duration: time.Since(startTime),
281+
FirstTokenMs: firstTokenMs,
282+
Usage: ClaudeUsage{},
283+
MediaType: "prompt",
279284
}, nil
280285
}
281286
if characterResult != nil && strings.TrimSpace(characterResult.Username) != "" {
@@ -419,16 +424,17 @@ func (s *SoraGatewayService) Forward(ctx context.Context, c *gin.Context, accoun
419424
}
420425

421426
return &ForwardResult{
422-
RequestID: taskID,
423-
Model: reqModel,
424-
Stream: clientStream,
425-
Duration: time.Since(startTime),
426-
FirstTokenMs: firstTokenMs,
427-
Usage: ClaudeUsage{},
428-
MediaType: mediaType,
429-
MediaURL: firstMediaURL(finalURLs),
430-
ImageCount: imageCount,
431-
ImageSize: imageSize,
427+
RequestID: taskID,
428+
Model: originalModel,
429+
UpstreamModel: upstreamModel,
430+
Stream: clientStream,
431+
Duration: time.Since(startTime),
432+
FirstTokenMs: firstTokenMs,
433+
Usage: ClaudeUsage{},
434+
MediaType: mediaType,
435+
MediaURL: firstMediaURL(finalURLs),
436+
ImageCount: imageCount,
437+
ImageSize: imageSize,
432438
}, nil
433439
}
434440

backend/internal/service/sora_gateway_service_test.go

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -144,6 +144,11 @@ func TestSoraGatewayService_ForwardPromptEnhance(t *testing.T) {
144144
ID: 1,
145145
Platform: PlatformSora,
146146
Status: StatusActive,
147+
Credentials: map[string]any{
148+
"model_mapping": map[string]any{
149+
"prompt-enhance-short-10s": "prompt-enhance-short-15s",
150+
},
151+
},
147152
}
148153
body := []byte(`{"model":"prompt-enhance-short-10s","messages":[{"role":"user","content":"cat running"}],"stream":false}`)
149154

@@ -152,6 +157,7 @@ func TestSoraGatewayService_ForwardPromptEnhance(t *testing.T) {
152157
require.NotNil(t, result)
153158
require.Equal(t, "prompt", result.MediaType)
154159
require.Equal(t, "prompt-enhance-short-10s", result.Model)
160+
require.Equal(t, "prompt-enhance-short-15s", result.UpstreamModel)
155161
}
156162

157163
func TestSoraGatewayService_ForwardStoryboardPrompt(t *testing.T) {

0 commit comments

Comments
 (0)