diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index a0e8ec1d4e..22a9c16f5e 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.145 +0.1.146 diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index b20d9ef652..0caa7f718b 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -158,6 +158,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { bodyRef := service.NewRequestBodyRef(body) parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic) if err != nil { + logRequestBodyParseFailure(reqLog, body, err) h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } @@ -1796,6 +1797,7 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) { bodyRef := service.NewRequestBodyRef(body) parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic) if err != nil { + logRequestBodyParseFailure(reqLog, body, err) h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index d0ecc01e6a..03ceb0d952 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -64,6 +64,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { // Validate JSON if !gjson.ValidBytes(body) { + logRequestBodyParseFailure(reqLog, body, nil) h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 4a8d752193..f5ee18b722 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -64,6 +64,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) { // Validate JSON if !gjson.ValidBytes(body) { + logRequestBodyParseFailure(reqLog, body, nil) h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index baff1dcbd6..847d386cde 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -64,6 +64,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { } if !gjson.ValidBytes(body) { + logRequestBodyParseFailure(reqLog, body, nil) h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index 8be533c723..56d775eb7c 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -60,6 +60,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { return } if !gjson.ValidBytes(body) { + logRequestBodyParseFailure(reqLog, body, nil) h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index fc9c4d5df7..9a6709cc4f 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -64,6 +64,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { bodyRef := service.NewRequestBodyRef(body) parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic) if err != nil { + logRequestBodyParseFailure(reqLog, body, err) h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 7f097afa4b..551de2dd61 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -218,6 +218,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { // 校验请求体 JSON 合法性 if !gjson.ValidBytes(body) { + logRequestBodyParseFailure(reqLog, body, nil) h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } @@ -697,6 +698,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { } if !gjson.ValidBytes(body) { + logRequestBodyParseFailure(reqLog, body, nil) h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return } diff --git a/backend/internal/handler/request_body_parse_log.go b/backend/internal/handler/request_body_parse_log.go new file mode 100644 index 0000000000..c38a7f9bd5 --- /dev/null +++ b/backend/internal/handler/request_body_parse_log.go @@ -0,0 +1,54 @@ +package handler + +import ( + "strconv" + + "github.com/Wei-Shaw/sub2api/internal/service" + "go.uber.org/zap" +) + +// parseFailureSnippetLen bounds the head/tail snippets logged on body parse +// failure. 256 bytes is enough to see the structural context (model field, +// first content block / trailing brace) without dumping user payloads. +const parseFailureSnippetLen = 256 + +// logRequestBodyParseFailure records the real reason a request body failed +// JSON parsing/validation. The client keeps receiving the generic +// "Failed to parse request body"; the sanitized diagnostics (underlying +// error with byte offset, body length, escaped head/tail snippets) land in +// the server log only, so operators can distinguish genuinely invalid JSON +// from a truncated or partially consumed body. +// +// err may be nil for call sites that validate with gjson.ValidBytes directly; +// the diagnostic error is derived from the body in that case. +func logRequestBodyParseFailure(reqLog *zap.Logger, body []byte, err error) { + if reqLog == nil { + return + } + if err == nil { + err = service.DescribeInvalidJSON(body) + } + + head := body + var tail []byte + if len(body) > parseFailureSnippetLen { + head = body[:parseFailureSnippetLen] + tail = body[len(body)-parseFailureSnippetLen:] + } + + fields := []zap.Field{ + zap.Error(err), + zap.Int("body_len", len(body)), + zap.String("body_head", sanitizeBodySnippet(head)), + } + if len(tail) > 0 { + fields = append(fields, zap.String("body_tail", sanitizeBodySnippet(tail))) + } + reqLog.Warn("parse request body failed", fields...) +} + +// sanitizeBodySnippet escapes control characters and invalid UTF-8 so the +// snippet is always a single printable log line. +func sanitizeBodySnippet(b []byte) string { + return strconv.Quote(string(b)) +} diff --git a/backend/internal/handler/request_body_parse_log_test.go b/backend/internal/handler/request_body_parse_log_test.go new file mode 100644 index 0000000000..c1477eb4d7 --- /dev/null +++ b/backend/internal/handler/request_body_parse_log_test.go @@ -0,0 +1,100 @@ +//go:build unit + +package handler + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" + "go.uber.org/zap" + "go.uber.org/zap/zaptest/observer" +) + +func newObservedLogger(t *testing.T) (*zap.Logger, *observer.ObservedLogs) { + t.Helper() + core, logs := observer.New(zap.WarnLevel) + return zap.New(core), logs +} + +func loggedFields(t *testing.T, logs *observer.ObservedLogs) map[string]any { + t.Helper() + entries := logs.All() + require.Len(t, entries, 1) + fields := map[string]any{} + for _, f := range entries[0].Context { + switch f.Key { + case "body_len": + fields[f.Key] = int(f.Integer) + case "error": + fields[f.Key] = f.Interface.(error).Error() + default: + fields[f.Key] = f.String + } + } + return fields +} + +func TestLogRequestBodyParseFailure_DerivesErrorWhenNil(t *testing.T) { + log, logs := newObservedLogger(t) + body := []byte(`{"model": bad}`) + + logRequestBodyParseFailure(log, body, nil) + + fields := loggedFields(t, logs) + require.Equal(t, len(body), fields["body_len"]) + require.Contains(t, fields["error"], "invalid json") + require.Contains(t, fields["error"], "offset=11") +} + +func TestLogRequestBodyParseFailure_ShortBodyHasNoTail(t *testing.T) { + log, logs := newObservedLogger(t) + body := []byte(`{"broken":`) + + logRequestBodyParseFailure(log, body, nil) + + fields := loggedFields(t, logs) + require.Contains(t, fields, "body_head") + require.NotContains(t, fields, "body_tail") + require.Contains(t, fields["body_head"].(string), `{\"broken\":`) +} + +func TestLogRequestBodyParseFailure_LargeBodyBoundedSnippets(t *testing.T) { + log, logs := newObservedLogger(t) + // ~1MB body: head must show the structural prefix, tail the trailing bytes, + // and neither snippet may exceed the configured bound (plus quoting overhead). + body := []byte(`{"model":"claude-sonnet-4-6","big":"` + strings.Repeat("A", 1<<20) + `"`) + + logRequestBodyParseFailure(log, body, nil) + + fields := loggedFields(t, logs) + require.Equal(t, len(body), fields["body_len"]) + head := fields["body_head"].(string) + tail := fields["body_tail"].(string) + require.Contains(t, head, "claude-sonnet-4-6") + require.Contains(t, tail, "AAA") + require.NotContains(t, tail, "claude-sonnet-4-6") + // strconv.Quote adds surrounding quotes and escapes; 4x is a generous cap. + require.LessOrEqual(t, len(head), parseFailureSnippetLen*4) + require.LessOrEqual(t, len(tail), parseFailureSnippetLen*4) +} + +func TestLogRequestBodyParseFailure_EscapesControlCharacters(t *testing.T) { + log, logs := newObservedLogger(t) + body := []byte("{\"model\":\x01\n\"x\"}") + + logRequestBodyParseFailure(log, body, nil) + + fields := loggedFields(t, logs) + head := fields["body_head"].(string) + require.NotContains(t, head, "\n") + require.NotContains(t, head, "\x01") + require.Contains(t, head, `\n`) + require.Contains(t, head, `\x01`) +} + +func TestLogRequestBodyParseFailure_NilLoggerNoPanic(t *testing.T) { + require.NotPanics(t, func() { + logRequestBodyParseFailure(nil, []byte(`{`), nil) + }) +} diff --git a/backend/internal/pkg/xai/models.go b/backend/internal/pkg/xai/models.go index 4902fcb94f..a5b800cf2c 100644 --- a/backend/internal/pkg/xai/models.go +++ b/backend/internal/pkg/xai/models.go @@ -12,6 +12,7 @@ type Model struct { var defaultModels = []Model{ {ID: "grok-4.3", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"}, {ID: "grok-build-0.1", Object: "model", OwnedBy: "xai", DisplayName: "Grok Build 0.1"}, + {ID: "grok-composer-2.5-fast", Object: "model", OwnedBy: "xai", DisplayName: "Grok Composer 2.5 Fast"}, {ID: "grok-4.20-0309-reasoning", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Reasoning"}, {ID: "grok-4.20-0309-non-reasoning", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Non Reasoning"}, {ID: "grok-4.20-multi-agent-0309", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Multi Agent"}, @@ -46,6 +47,7 @@ func DefaultModelMapping() map[string]string { mapping["grok"] = "grok-4.3" mapping["grok-latest"] = "grok-4.3" mapping["grok-build"] = "grok-build-0.1" + mapping["grok-composer"] = "grok-composer-2.5-fast" mapping["grok-4.20-reasoning"] = "grok-4.20-0309-reasoning" mapping["grok-4.20-non-reasoning"] = "grok-4.20-0309-non-reasoning" return mapping diff --git a/backend/internal/pkg/xai/oauth_test.go b/backend/internal/pkg/xai/oauth_test.go index 68a4fea240..28609a08fa 100644 --- a/backend/internal/pkg/xai/oauth_test.go +++ b/backend/internal/pkg/xai/oauth_test.go @@ -210,6 +210,7 @@ func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) { require.Equal(t, "grok-4.3", mapping["grok"]) require.Equal(t, "grok-4.3", mapping["grok-latest"]) require.Equal(t, "grok-build-0.1", mapping["grok-build"]) + require.Equal(t, "grok-composer-2.5-fast", mapping["grok-composer"]) require.Equal(t, "grok-4.20-0309-reasoning", mapping["grok-4.20-reasoning"]) require.Equal(t, "grok-4.20-0309-non-reasoning", mapping["grok-4.20-non-reasoning"]) require.Equal(t, "grok-4.20-multi-agent-0309", mapping["grok-4.20-multi-agent-0309"]) diff --git a/backend/internal/service/antigravity_gateway_service.go b/backend/internal/service/antigravity_gateway_service.go index aa4cab22d7..585170438e 100644 --- a/backend/internal/service/antigravity_gateway_service.go +++ b/backend/internal/service/antigravity_gateway_service.go @@ -153,14 +153,24 @@ type antigravityRetryLoopResult struct { } // resolveAntigravityForwardBaseURL 解析转发用 base URL。 -// 默认使用 daily(ForwardBaseURLs 的首个地址);当环境变量为 prod 时使用第二个地址。 +// +// 默认使用生产端点 cloudcode-pa.googleapis.com(antigravity.BaseURLs 的首个地址, +// 与账号 OAuth 登录/测试连接所用的 antigravity.BaseURL 一致)。 +// +// 历史上这里改用 ForwardBaseURLs()(把 daily/sandbox 排到首位)并默认取首个地址, +// 导致网关把带生产 OAuth token 的请求发到 daily-cloudcode-pa.sandbox.googleapis.com, +// 上游拒绝 → 账号被 401「Invalid bearer token」/502 打入临时不可调度且无法恢复 +// (见 #3611 / #2962)。后台「测试连接」用的是生产端点,所以「测试成功但网关 401」。 +// +// daily/sandbox 端点仅供内部联调,需显式设置 +// GATEWAY_ANTIGRAVITY_FORWARD_BASE_URL=daily(或 sandbox)才启用。 func resolveAntigravityForwardBaseURL() string { - baseURLs := antigravity.ForwardBaseURLs() + baseURLs := antigravity.BaseURLs if len(baseURLs) == 0 { return "" } mode := strings.ToLower(strings.TrimSpace(os.Getenv(antigravityForwardBaseURLEnv))) - if mode == "prod" && len(baseURLs) > 1 { + if (mode == "daily" || mode == "sandbox") && len(baseURLs) > 1 { return baseURLs[1] } return baseURLs[0] diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go index a90714ca1d..1665b1fe47 100644 --- a/backend/internal/service/gateway_request.go +++ b/backend/internal/service/gateway_request.go @@ -3,6 +3,7 @@ package service import ( "bytes" "encoding/json" + "errors" "fmt" "math" "regexp" @@ -168,7 +169,7 @@ func parseGatewayRequestCurrentBody(parsed *ParsedRequest, protocol string) erro bodyBytes := parsed.Body.Bytes() if !gjson.ValidBytes(bodyBytes) { - return fmt.Errorf("invalid json") + return DescribeInvalidJSON(bodyBytes) } // 只在当前函数内零拷贝读取 JSON 字段;ReplaceBody 后必须重新进入本函数刷新派生状态。 @@ -216,6 +217,26 @@ func refreshGatewayRequestRanges(parsed *ParsedRequest, protocol string) error { return parseGatewayRequestCurrentBody(parsed, protocol) } +// DescribeInvalidJSON returns a diagnostic error for a request body that +// failed JSON validation. It re-parses with encoding/json (failure path only) +// to pinpoint the first offending byte, so operators can distinguish genuinely +// invalid JSON from a truncated / partially consumed body. The error carries +// only length/offset/character information — never body content — so callers +// may safely wrap or log it. +func DescribeInvalidJSON(body []byte) error { + var raw json.RawMessage + if err := json.Unmarshal(body, &raw); err != nil { + var syntaxErr *json.SyntaxError + if errors.As(err, &syntaxErr) { + return fmt.Errorf("invalid json (len=%d, offset=%d): %s", len(body), syntaxErr.Offset, syntaxErr.Error()) + } + return fmt.Errorf("invalid json (len=%d): %w", len(body), err) + } + // gjson rejected the body but encoding/json accepted it (divergent edge + // cases, e.g. certain malformed UTF-8 sequences); report the basics. + return fmt.Errorf("invalid json (len=%d)", len(body)) +} + // ParsedRequest 保存网关请求的预解析结果 // // 性能优化说明: diff --git a/backend/internal/service/gateway_request_invalid_json_test.go b/backend/internal/service/gateway_request_invalid_json_test.go new file mode 100644 index 0000000000..cc69a41e72 --- /dev/null +++ b/backend/internal/service/gateway_request_invalid_json_test.go @@ -0,0 +1,52 @@ +//go:build unit + +package service + +import ( + "fmt" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/domain" + "github.com/stretchr/testify/require" +) + +func TestDescribeInvalidJSON_TruncatedBody(t *testing.T) { + // Simulates a body cut off mid-stream (e.g. partially consumed by middleware). + body := []byte(`{"model":"claude-sonnet-4-6","messages":[{"role":"user","content":"hi`) + + err := DescribeInvalidJSON(body) + + require.Error(t, err) + require.Contains(t, err.Error(), fmt.Sprintf("len=%d", len(body))) + require.Contains(t, err.Error(), "unexpected end of JSON input") +} + +func TestDescribeInvalidJSON_InvalidCharacterWithOffset(t *testing.T) { + body := []byte(`{"model": bad}`) + + err := DescribeInvalidJSON(body) + + require.Error(t, err) + require.Contains(t, err.Error(), "offset=11") + require.Contains(t, err.Error(), "invalid character") +} + +func TestDescribeInvalidJSON_DoesNotLeakBodyContent(t *testing.T) { + secret := "sk-super-secret-value" + body := []byte(`{"api_key":"` + secret + `","broken":`) + + err := DescribeInvalidJSON(body) + + require.Error(t, err) + require.NotContains(t, err.Error(), secret) +} + +func TestParseGatewayRequest_InvalidJSONErrorIsDiagnostic(t *testing.T) { + body := []byte(`{"model":"claude-sonnet-4-6","messages":[`) + + _, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic) + + require.Error(t, err) + require.True(t, strings.HasPrefix(err.Error(), "invalid json (len="), "error should carry diagnostics, got: %s", err.Error()) +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index dcaf3a645c..2ca554edb3 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -5105,6 +5105,13 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A if err := replaceBody(StripEmptyTextBlocks(body)); err != nil { return nil, err } + // Pre-filter: strip web-search history blocks the upstream cannot accept + // (emulation-synthesized server_tool_use / web_search_tool_result always; + // genuine ones additionally for passback-required upstreams). See + // FilterWebSearchHistoryBlocks. reqModel 此时已是映射后的模型 ID。 + if err := replaceBody(FilterWebSearchHistoryBlocks(body, reqModel)); err != nil { + return nil, err + } // Pre-filter: remove thinking blocks with missing/invalid signatures before forwarding. // Clients (e.g. Claude Code) sometimes send multi-turn conversations where a historical // assistant message contains a thinking block that is missing the required "signature" field, @@ -5688,6 +5695,11 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput( } // Pre-filter: strip empty text blocks (including nested in tool_result) to prevent upstream 400. input.Body = StripEmptyTextBlocks(input.Body) + // Pre-filter: strip web-search history blocks the upstream cannot accept + // (emulation-synthesized ones always; genuine ones additionally for + // passback-required third-party upstreams such as GLM/Kimi/DeepSeek, + // which reject server_tool_use with 400). input.RequestModel 已是映射后的模型 ID。 + input.Body = FilterWebSearchHistoryBlocks(input.Body, input.RequestModel) if input.Parsed != nil { // 透传分支也会改写实际 wire body,成功 usage hash 依赖这里同步当前 body。 if err := input.Parsed.ReplaceBody(input.Body); err != nil { diff --git a/backend/internal/service/gateway_websearch_block_filter.go b/backend/internal/service/gateway_websearch_block_filter.go new file mode 100644 index 0000000000..a0706c5c2d --- /dev/null +++ b/backend/internal/service/gateway_websearch_block_filter.go @@ -0,0 +1,138 @@ +package service + +import ( + "bytes" + "encoding/json" + "strings" + "unsafe" + + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" +) + +const ( + blockTypeServerToolUse = "server_tool_use" + blockTypeWebSearchToolResult = "web_search_tool_result" +) + +// Fast-path byte patterns: both block types only ever appear as quoted JSON +// string values, so a raw substring check is a safe pre-filter regardless of +// key/value spacing. +var ( + patternServerToolUse = []byte(`"server_tool_use"`) + patternWebSearchToolResult = []byte(`"web_search_tool_result"`) +) + +// FilterWebSearchHistoryBlocks removes web-search content blocks from +// historical messages when the upstream cannot accept them: +// +// 1. Emulation-synthesized blocks — server_tool_use / web_search_tool_result +// whose tool-use ID carries webSearchToolUseIDPrefix — are fabricated +// locally by the web-search emulation (gateway_websearch_emulation.go). +// No upstream ever issued them, so clients replaying the conversation +// (e.g. Claude Code) poison every follow-up request. They are stripped +// for all upstreams. +// 2. For passback-required upstreams (DeepSeek/Kimi/GLM …, see +// ResolveThinkingProtocol) all server_tool_use / web_search_tool_result +// blocks are stripped: these upstreams only accept +// text/thinking/image/tool_use/tool_result and reject anything else with +// 400 "invalid value: `server_tool_use`". anthropic-strict and unknown +// upstreams keep genuine blocks untouched. +// +// The emulated assistant turn always carries a trailing text summary, so the +// search context survives the strip. A message whose content would become +// empty gets a placeholder text block (mirroring FilterThinkingBlocksForRetry). +// Returns the original body unchanged when nothing needs stripping. +func FilterWebSearchHistoryBlocks(body []byte, mappedModel string) []byte { + if !bytes.Contains(body, patternServerToolUse) && !bytes.Contains(body, patternWebSearchToolResult) { + return body + } + + stripAll := ResolveThinkingProtocol(mappedModel) == ThinkingProtocolPassbackRequired + + jsonStr := *(*string)(unsafe.Pointer(&body)) + msgsRes := gjson.Get(jsonStr, "messages") + if !msgsRes.Exists() || !msgsRes.IsArray() { + return body + } + + var messages []any + if err := json.Unmarshal(sliceRawFromBody(body, msgsRes), &messages); err != nil { + return body + } + + modified := false + for _, msg := range messages { + msgMap, ok := msg.(map[string]any) + if !ok { + continue + } + content, ok := msgMap["content"].([]any) + if !ok { + continue + } + + // 延迟分配:只有命中需剥离的块才构建新 slice。 + var newContent []any + for i, block := range content { + blockMap, isMap := block.(map[string]any) + if isMap && shouldStripWebSearchBlock(blockMap, stripAll) { + if newContent == nil { + newContent = make([]any, 0, len(content)) + newContent = append(newContent, content[:i]...) + } + continue + } + if newContent != nil { + newContent = append(newContent, block) + } + } + if newContent == nil { + continue + } + modified = true + if len(newContent) == 0 { + role, _ := msgMap["role"].(string) + placeholder := "(content removed)" + if role == "assistant" { + placeholder = "(assistant content removed)" + } + newContent = []any{map[string]any{"type": "text", "text": placeholder}} + } + msgMap["content"] = newContent + } + + if !modified { + return body + } + + msgsBytes, err := json.Marshal(messages) + if err != nil { + return body + } + out, err := sjson.SetRawBytes(body, "messages", msgsBytes) + if err != nil { + return body + } + return out +} + +func shouldStripWebSearchBlock(block map[string]any, stripAll bool) bool { + blockType, _ := block["type"].(string) + switch blockType { + case blockTypeServerToolUse: + if stripAll { + return true + } + id, _ := block["id"].(string) + return strings.HasPrefix(id, webSearchToolUseIDPrefix) + case blockTypeWebSearchToolResult: + if stripAll { + return true + } + id, _ := block["tool_use_id"].(string) + return strings.HasPrefix(id, webSearchToolUseIDPrefix) + default: + return false + } +} diff --git a/backend/internal/service/gateway_websearch_block_filter_test.go b/backend/internal/service/gateway_websearch_block_filter_test.go new file mode 100644 index 0000000000..cebda4d37b --- /dev/null +++ b/backend/internal/service/gateway_websearch_block_filter_test.go @@ -0,0 +1,140 @@ +//go:build unit + +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +// emulatedWebSearchBody is a follow-up /v1/messages request whose history +// contains an assistant turn synthesized by the web-search emulation +// (server_tool_use + web_search_tool_result with the local srvtoolu_ws_ ID +// prefix, followed by the text summary). +const emulatedWebSearchBody = `{"model":"claude-sonnet-4-6","max_tokens":1024,"messages":[` + + `{"role":"user","content":[{"type":"text","text":"search the weather"}]},` + + `{"role":"assistant","content":[` + + `{"type":"server_tool_use","id":"srvtoolu_ws_0123456789abcdef","name":"web_search","input":{"query":"weather"}},` + + `{"type":"web_search_tool_result","tool_use_id":"srvtoolu_ws_0123456789abcdef","content":[{"type":"web_search_result","url":"https://example.com","title":"Weather"}]},` + + `{"type":"text","text":"Here are the search results for \"weather\":"}]},` + + `{"role":"user","content":[{"type":"text","text":"thanks, continue"}]}]}` + +// genuineWebSearchBody carries real Anthropic web-search blocks (upstream IDs +// do NOT have the local srvtoolu_ws_ prefix). +const genuineWebSearchBody = `{"model":"claude-sonnet-4-6","max_tokens":1024,"messages":[` + + `{"role":"user","content":[{"type":"text","text":"search"}]},` + + `{"role":"assistant","content":[` + + `{"type":"server_tool_use","id":"srvtoolu_01ABCDEF","name":"web_search","input":{"query":"weather"}},` + + `{"type":"web_search_tool_result","tool_use_id":"srvtoolu_01ABCDEF","content":[{"type":"web_search_result","url":"https://example.com","title":"Weather"}]},` + + `{"type":"text","text":"summary with citations"}]}]}` + +func collectContentTypes(t *testing.T, body []byte) []string { + t.Helper() + var types []string + for _, msg := range gjson.GetBytes(body, "messages").Array() { + for _, block := range msg.Get("content").Array() { + types = append(types, block.Get("type").String()) + } + } + return types +} + +func TestFilterWebSearchHistoryBlocks_StripsEmulatedBlocksForAnthropicStrict(t *testing.T) { + out := FilterWebSearchHistoryBlocks([]byte(emulatedWebSearchBody), "claude-sonnet-4-6") + + require.Equal(t, []string{"text", "text", "text"}, collectContentTypes(t, out)) + // The emulated text summary must survive so the search context is preserved. + require.Contains(t, string(out), "Here are the search results") + require.NotContains(t, string(out), "srvtoolu_ws_") + require.True(t, gjson.ValidBytes(out)) +} + +func TestFilterWebSearchHistoryBlocks_KeepsGenuineBlocksForAnthropicStrict(t *testing.T) { + body := []byte(genuineWebSearchBody) + out := FilterWebSearchHistoryBlocks(body, "claude-sonnet-4-6") + + require.Equal(t, string(body), string(out)) +} + +func TestFilterWebSearchHistoryBlocks_StripsAllBlocksForPassbackRequired(t *testing.T) { + // GLM only accepts text/thinking/image/tool_use/tool_result and rejects + // server_tool_use with 400, so genuine blocks must be stripped as well. + out := FilterWebSearchHistoryBlocks([]byte(genuineWebSearchBody), "glm-4.7") + + require.Equal(t, []string{"text", "text"}, collectContentTypes(t, out)) + require.NotContains(t, string(out), "server_tool_use") + require.NotContains(t, string(out), "web_search_tool_result") + require.Contains(t, string(out), "summary with citations") +} + +func TestFilterWebSearchHistoryBlocks_StripsEmulatedBlocksForUnknownModel(t *testing.T) { + out := FilterWebSearchHistoryBlocks([]byte(emulatedWebSearchBody), "totally-unknown-model") + + require.Equal(t, []string{"text", "text", "text"}, collectContentTypes(t, out)) + require.NotContains(t, string(out), "srvtoolu_ws_") +} + +func TestFilterWebSearchHistoryBlocks_KeepsGenuineBlocksForUnknownModel(t *testing.T) { + body := []byte(genuineWebSearchBody) + out := FilterWebSearchHistoryBlocks(body, "totally-unknown-model") + + require.Equal(t, string(body), string(out)) +} + +func TestFilterWebSearchHistoryBlocks_NoWebSearchBlocksFastPath(t *testing.T) { + body := []byte(`{"model":"claude-sonnet-4-6","messages":[{"role":"user","content":[{"type":"text","text":"hi"}]}]}`) + out := FilterWebSearchHistoryBlocks(body, "claude-sonnet-4-6") + + require.Equal(t, string(body), string(out)) +} + +func TestFilterWebSearchHistoryBlocks_EmptiedMessageGetsPlaceholder(t *testing.T) { + body := []byte(`{"model":"glm-4.7","messages":[` + + `{"role":"user","content":[{"type":"text","text":"search"}]},` + + `{"role":"assistant","content":[` + + `{"type":"server_tool_use","id":"srvtoolu_01X","name":"web_search","input":{"query":"q"}},` + + `{"type":"web_search_tool_result","tool_use_id":"srvtoolu_01X","content":[]}]}]}`) + + out := FilterWebSearchHistoryBlocks(body, "glm-4.7") + + msgs := gjson.GetBytes(out, "messages").Array() + require.Len(t, msgs, 2) + assistant := msgs[1] + require.Equal(t, "assistant", assistant.Get("role").String()) + content := assistant.Get("content").Array() + require.Len(t, content, 1) + require.Equal(t, "text", content[0].Get("type").String()) + require.Equal(t, "(assistant content removed)", content[0].Get("text").String()) +} + +func TestFilterWebSearchHistoryBlocks_StringContentUntouched(t *testing.T) { + // A string mentioning the pattern inside a text value must not trigger a rewrite. + body := []byte(`{"model":"claude-sonnet-4-6","messages":[` + + `{"role":"user","content":"please explain \"server_tool_use\" blocks"}]}`) + + out := FilterWebSearchHistoryBlocks(body, "claude-sonnet-4-6") + + require.Equal(t, string(body), string(out)) +} + +func TestFilterWebSearchHistoryBlocks_InvalidMessagesUnchanged(t *testing.T) { + body := []byte(`{"model":"claude-sonnet-4-6","messages":"server_tool_use"}`) + out := FilterWebSearchHistoryBlocks(body, "claude-sonnet-4-6") + + require.Equal(t, string(body), string(out)) +} + +func TestFilterWebSearchHistoryBlocks_PreservesOtherToolBlocks(t *testing.T) { + body := []byte(`{"model":"glm-4.7","messages":[` + + `{"role":"assistant","content":[` + + `{"type":"tool_use","id":"toolu_01A","name":"get_weather","input":{}},` + + `{"type":"server_tool_use","id":"srvtoolu_ws_abc","name":"web_search","input":{"query":"q"}},` + + `{"type":"text","text":"result"}]},` + + `{"role":"user","content":[{"type":"tool_result","tool_use_id":"toolu_01A","content":"sunny"}]}]}`) + + out := FilterWebSearchHistoryBlocks(body, "glm-4.7") + + require.Equal(t, []string{"tool_use", "text", "tool_result"}, collectContentTypes(t, out)) +} diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 348213a992..023440f94f 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -105,6 +105,33 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( return nil, policyErr } upstreamBody = updatedBody + + // Grok Composer does not accept image_url parts directly, but Grok Build + // can describe the images first. Bridge only this exact failure mode. + token, tokenKind, err := s.GetAccessToken(ctx, account) + if err != nil { + return nil, err + } + if strings.TrimSpace(token) == "" { + return nil, fmt.Errorf("account %d missing %s credential", account.ID, tokenKind) + } + + var bridgeUsage OpenAIUsage + if account.Platform == PlatformGrok { + bridgedBody, usage, bridged, bridgeErr := s.bridgeGrokComposerImageInputs(ctx, c, account, upstreamBody, token) + if bridgeErr != nil { + var failoverErr *UpstreamFailoverError + if !errors.As(bridgeErr, &failoverErr) && c != nil && c.Writer != nil && !c.Writer.Written() { + writeChatCompletionsError(c, http.StatusBadGateway, "upstream_error", bridgeErr.Error()) + } + return nil, bridgeErr + } + if bridged { + upstreamBody = bridgedBody + addOpenAIUsage(&bridgeUsage, usage) + } + } + if clientStream { var usageErr error upstreamBody, usageErr = ensureOpenAIChatStreamUsage(upstreamBody) @@ -122,14 +149,6 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( ) // 5. Build upstream request - token, tokenKind, err := s.GetAccessToken(ctx, account) - if err != nil { - return nil, err - } - if strings.TrimSpace(token) == "" { - return nil, fmt.Errorf("account %d missing %s credential", account.ID, tokenKind) - } - targetURL, err := s.rawChatCompletionsURL(account) if err != nil { return nil, err @@ -245,10 +264,17 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } // 8. Forward response + var result *OpenAIForwardResult + var forwardErr error if clientStream { - return s.streamRawChatCompletions(c, resp, account, originalModel, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime, len(body)) + result, forwardErr = s.streamRawChatCompletions(c, resp, account, originalModel, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime, len(body)) + } else { + result, forwardErr = s.bufferRawChatCompletions(c, resp, originalModel, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime) } - return s.bufferRawChatCompletions(c, resp, originalModel, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime) + if result != nil { + addOpenAIUsage(&result.Usage, bridgeUsage) + } + return result, forwardErr } func (s *OpenAIGatewayService) rawChatCompletionsURL(account *Account) (string, error) { diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 4961b9c589..4a0ad06d46 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -10,12 +10,18 @@ import ( "strings" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" "github.com/tidwall/sjson" ) +const ( + grokComposerImageBridgeVisionModel = "grok-build-0.1" + grokComposerImageBridgeMaxOutputTokens = 512 +) + func (s *OpenAIGatewayService) forwardGrokResponses( ctx context.Context, c *gin.Context, @@ -309,6 +315,303 @@ func shouldDropGrokToolChoice(toolChoice gjson.Result, tools []json.RawMessage) return false } +func (s *OpenAIGatewayService) bridgeGrokComposerImageInputs( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + token string, +) ([]byte, OpenAIUsage, bool, error) { + if !shouldBridgeGrokComposerImageInputs(body) { + return body, OpenAIUsage{}, false, nil + } + + var reqBody map[string]any + if err := json.Unmarshal(body, &reqBody); err != nil { + return body, OpenAIUsage{}, false, fmt.Errorf("parse grok composer image bridge request: %w", err) + } + + imageURLs := collectGrokComposerImageURLs(reqBody) + if len(imageURLs) == 0 { + return body, OpenAIUsage{}, false, nil + } + + descriptions := make([]string, 0, len(imageURLs)) + var bridgeUsage OpenAIUsage + for index, imageURL := range imageURLs { + description, usage, err := s.describeGrokComposerImage(ctx, c, account, token, imageURL, index+1) + if err != nil { + return body, bridgeUsage, false, err + } + descriptions = append(descriptions, description) + addOpenAIUsage(&bridgeUsage, usage) + } + + if !rewriteGrokComposerImagesAsText(reqBody, descriptions) { + return body, bridgeUsage, false, nil + } + bridgedBody, err := marshalOpenAIUpstreamJSON(reqBody) + if err != nil { + return body, bridgeUsage, false, fmt.Errorf("serialize grok composer image bridge request: %w", err) + } + return bridgedBody, bridgeUsage, true, nil +} + +func shouldBridgeGrokComposerImageInputs(body []byte) bool { + if len(body) == 0 || !isGrokComposerModel(gjson.GetBytes(body, "model").String()) { + return false + } + messages := gjson.GetBytes(body, "messages") + if !messages.Exists() { + return false + } + return openAIJSONValueMayContainImageInput(messages) +} + +func isGrokComposerModel(model string) bool { + model = strings.TrimSpace(strings.ToLower(model)) + if model == "" { + return false + } + if strings.Contains(model, "/") { + parts := strings.Split(model, "/") + model = strings.TrimSpace(parts[len(parts)-1]) + } + return strings.Contains(model, "composer") +} + +func collectGrokComposerImageURLs(reqBody map[string]any) []string { + messages, ok := reqBody["messages"].([]any) + if !ok { + return nil + } + + var imageURLs []string + for _, msg := range messages { + msgMap, ok := msg.(map[string]any) + if !ok { + continue + } + parts, ok := msgMap["content"].([]any) + if !ok { + continue + } + for _, part := range parts { + if imageURL := grokComposerImageURLFromPart(part); imageURL != "" { + imageURLs = append(imageURLs, imageURL) + } + } + } + return imageURLs +} + +func grokComposerImageURLFromPart(part any) string { + partMap, ok := part.(map[string]any) + if !ok { + return "" + } + if strings.TrimSpace(strings.ToLower(fmt.Sprint(partMap["type"]))) != "image_url" { + return "" + } + switch imageURL := partMap["image_url"].(type) { + case string: + return normalizeGrokComposerImageURL(imageURL) + case map[string]any: + raw, _ := imageURL["url"].(string) + return normalizeGrokComposerImageURL(raw) + default: + return "" + } +} + +func normalizeGrokComposerImageURL(raw string) string { + trimmed := strings.TrimSpace(raw) + if trimmed == "" || isEmptyBase64DataURI(trimmed) { + return "" + } + return trimmed +} + +func (s *OpenAIGatewayService) describeGrokComposerImage( + ctx context.Context, + c *gin.Context, + account *Account, + token string, + imageURL string, + index int, +) (string, OpenAIUsage, error) { + body, err := buildGrokComposerImageDescriptionBody(imageURL, index) + if err != nil { + return "", OpenAIUsage{}, err + } + + upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token) + releaseUpstreamCtx() + if err != nil { + return "", OpenAIUsage{}, fmt.Errorf("build grok composer image bridge request: %w", err) + } + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + return "", OpenAIUsage{}, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 400 { + respBody := s.readUpstreamErrorBody(resp) + s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + upstreamMsg := sanitizeUpstreamErrorMessage(extractUpstreamErrorMessage(respBody)) + if upstreamMsg == "" { + upstreamMsg = fmt.Sprintf("xAI image bridge upstream returned status %d", resp.StatusCode) + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")), + Kind: "failover", + Message: upstreamMsg, + }) + s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + if s.shouldFailoverUpstreamError(resp.StatusCode) { + return "", OpenAIUsage{}, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + return "", OpenAIUsage{}, fmt.Errorf("grok composer image bridge upstream error: %s", upstreamMsg) + } + + s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, nil) + if err != nil { + return "", OpenAIUsage{}, fmt.Errorf("read grok composer image bridge response: %w", err) + } + + var parsed apicompat.ResponsesResponse + if err := json.Unmarshal(respBody, &parsed); err != nil { + return "", OpenAIUsage{}, fmt.Errorf("parse grok composer image bridge response: %w", err) + } + description := strings.TrimSpace(grokResponsesOutputText(&parsed)) + if description == "" { + return "", copyOpenAIUsageFromResponsesUsage(parsed.Usage), fmt.Errorf("grok composer image bridge returned empty description") + } + return description, copyOpenAIUsageFromResponsesUsage(parsed.Usage), nil +} + +func buildGrokComposerImageDescriptionBody(imageURL string, index int) ([]byte, error) { + prompt := fmt.Sprintf("Describe image %d in concise, factual text for a downstream coding/composer model. Include visible text, UI elements, diagrams, errors, and spatial relationships. Do not mention that you are an image analysis bridge.", index) + req := map[string]any{ + "model": grokComposerImageBridgeVisionModel, + "stream": false, + "store": false, + "max_output_tokens": grokComposerImageBridgeMaxOutputTokens, + "input": []any{ + map[string]any{ + "type": "message", + "role": "user", + "content": []any{ + map[string]any{"type": "input_text", "text": prompt}, + map[string]any{"type": "input_image", "image_url": imageURL}, + }, + }, + }, + } + return marshalOpenAIUpstreamJSON(req) +} + +func grokResponsesOutputText(resp *apicompat.ResponsesResponse) string { + if resp == nil { + return "" + } + var parts []string + for _, output := range resp.Output { + for _, content := range output.Content { + if content.Type == "output_text" || content.Type == "text" || content.Type == "input_text" { + if text := strings.TrimSpace(content.Text); text != "" { + parts = append(parts, text) + } + } + } + } + return strings.Join(parts, "\n\n") +} + +func rewriteGrokComposerImagesAsText(reqBody map[string]any, descriptions []string) bool { + messages, ok := reqBody["messages"].([]any) + if !ok { + return false + } + + imageIndex := 0 + changed := false + for _, msg := range messages { + msgMap, ok := msg.(map[string]any) + if !ok { + continue + } + parts, ok := msgMap["content"].([]any) + if !ok { + continue + } + var textParts []string + messageChanged := false + for _, part := range parts { + if imageURL := grokComposerImageURLFromPart(part); imageURL != "" { + if imageIndex < len(descriptions) { + textParts = append(textParts, fmt.Sprintf("Image %d description: %s", imageIndex+1, strings.TrimSpace(descriptions[imageIndex]))) + } + imageIndex++ + messageChanged = true + continue + } + if text := grokComposerTextFromPart(part); text != "" { + textParts = append(textParts, text) + } + } + if messageChanged { + msgMap["content"] = strings.Join(textParts, "\n\n") + changed = true + } + } + return changed +} + +func grokComposerTextFromPart(part any) string { + partMap, ok := part.(map[string]any) + if !ok { + return "" + } + partType := strings.TrimSpace(strings.ToLower(fmt.Sprint(partMap["type"]))) + switch partType { + case "text", "input_text": + text, _ := partMap["text"].(string) + return strings.TrimSpace(text) + default: + return "" + } +} + +func addOpenAIUsage(dst *OpenAIUsage, usage OpenAIUsage) { + if dst == nil { + return + } + dst.InputTokens += usage.InputTokens + dst.ImageInputTokens += usage.ImageInputTokens + dst.OutputTokens += usage.OutputTokens + dst.CacheCreationInputTokens += usage.CacheCreationInputTokens + dst.CacheReadInputTokens += usage.CacheReadInputTokens + dst.ImageOutputTokens += usage.ImageOutputTokens +} + func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) { targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) if err != nil { diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index eae424ce9a..d012ab533a 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -651,6 +651,76 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te require.NotNil(t, repo.updates[53][grokQuotaSnapshotExtraKey]) } +func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok-composer-2.5-fast","messages":[{"role":"system","content":"You are concise."},{"role":"user","content":[{"type":"text","text":"What is shown?"},{"type":"image_url","image_url":{"url":"data:image/png;base64,QUJD"}}]}],"stream":false}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 55, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{55: account}, + }, + } + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}, "xai-request-id": []string{"vision-req"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_vision","object":"response","model":"grok-build-0.1","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"A small diagram with ABC letters."}]}],"usage":{"input_tokens":11,"output_tokens":7,"total_tokens":18}}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "X-Request-Id": []string{"composer-req"}, + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"9"}, + "X-Ratelimit-Limit-Tokens": []string{"1000"}, + "X-Ratelimit-Remaining-Tokens": []string{"980"}, + }, + Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl_composer","object":"chat.completion","model":"grok-composer-2.5-fast","choices":[{"index":0,"message":{"role":"assistant","content":"It shows ABC."},"finish_reason":"stop"}],"usage":{"prompt_tokens":3,"completion_tokens":5,"total_tokens":8}}`)), + }, + }} + svc := &OpenAIGatewayService{ + cfg: rawChatCompletionsTestConfig(), + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.requests, 2) + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.requests[0].URL.String()) + require.Equal(t, "grok-build-0.1", gjson.GetBytes(upstream.bodies[0], "model").String()) + require.Equal(t, "input_image", gjson.GetBytes(upstream.bodies[0], "input.0.content.1.type").String()) + require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.requests[1].URL.String()) + require.Equal(t, "grok-composer-2.5-fast", gjson.GetBytes(upstream.bodies[1], "model").String()) + require.False(t, strings.Contains(string(upstream.bodies[1]), "image_url")) + require.Contains(t, gjson.GetBytes(upstream.bodies[1], "messages.1.content").String(), "Image 1 description") + require.Contains(t, gjson.GetBytes(upstream.bodies[1], "messages.1.content").String(), "A small diagram with ABC letters.") + require.Equal(t, 14, result.Usage.InputTokens) + require.Equal(t, 12, result.Usage.OutputTokens) + require.Equal(t, "It shows ABC.", gjson.Get(recorder.Body.String(), "choices.0.message.content").String()) + require.NotNil(t, repo.updates[55][grokQuotaSnapshotExtraKey]) +} + func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 645b31992a..fd104db8b6 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -5668,7 +5668,13 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r if isEventStreamResponse(resp.Header) { return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel) } - bodyLooksLikeSSE := bytes.Contains(body, []byte("data:")) || bytes.Contains(body, []byte("event:")) + // bodyLooksLikeSSE is a line-level heuristic: real SSE framing requires + // "data:"/"event:" field names at the very start of a physical line. A + // plain bytes.Contains scan would also match ordinary JSON responses + // whose string content merely echoes the literal text "data:" or + // "event:" (e.g. compact tool output), causing those JSON bodies to be + // misrouted into handleSSEToJSON and lose their usage accounting. + bodyLooksLikeSSE := bodyHasSSEFraming(body) // For OAuth accounts, also fall back to a body-content heuristic because // the upstream may omit the Content-Type header while still sending SSE. @@ -5718,6 +5724,22 @@ func isEventStreamResponse(header http.Header) bool { return strings.Contains(contentType, "text/event-stream") } +// bodyHasSSEFraming reports whether body contains genuine SSE framing by +// scanning for physical lines that begin with the "data:" or "event:" +// field names, per the SSE spec. Unlike a raw substring scan, this does not +// match when those strings only appear embedded inside JSON string values +// (e.g. "data: foo" quoted as part of an assistant text field), since such +// occurrences never start a physical line in a valid JSON encoding. +func bodyHasSSEFraming(body []byte) bool { + for _, line := range bytes.Split(body, []byte("\n")) { + line = bytes.TrimRight(line, "\r") + if bytes.HasPrefix(line, []byte("data:")) || bytes.HasPrefix(line, []byte("event:")) { + return true + } + } + return false +} + func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Context, body []byte, originalModel, mappedModel string) (*openaiNonStreamingResult, error) { bodyText := string(body) finalResponse, ok := extractCodexFinalResponse(bodyText) diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index c11d78e55c..b3e9889a7d 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -2739,6 +2739,41 @@ func TestHandleNonStreamingResponse_APIKeyFallsBackToSSEBodyWhenContentTypeIsWro require.Equal(t, "hello", gjson.Get(rec.Body.String(), "output.0.content.0.text").String()) } +func TestHandleNonStreamingResponse_OAuthJSONBodyWithDataEventTextKeepsJSONUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/compact", nil) + + svc := &OpenAIGatewayService{cfg: &config.Config{}} + // Plain JSON compact response whose output text happens to contain the + // literal substrings "data:" and "event:" (e.g. echoing shell/log output). + // This must NOT be misdetected as SSE framing: it has a top-level usage + // object and no upstream text/event-stream Content-Type. + jsonBody := `{"id":"resp_oauth_compact","object":"response","model":"gpt-5.4","status":"completed",` + + `"output":[{"type":"message","content":[{"type":"output_text",` + + `"text":"processing data: 1,2,3 then event: click finished"}]}],` + + `"usage":{"input_tokens":11,"output_tokens":22,"total_tokens":33}}` + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(jsonBody)), + } + account := &Account{ID: 146, Type: AccountTypeOAuth} + + result, err := svc.handleNonStreamingResponse(context.Background(), resp, c, account, "gpt-5.4", "gpt-5.4") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 11, result.InputTokens) + require.Equal(t, 22, result.OutputTokens) + // Response must remain the original JSON body (not routed through the SSE + // path, which would rewrite/lose the body or usage). + require.Equal(t, "application/json", rec.Header().Get("Content-Type")) + require.Equal(t, "resp_oauth_compact", gjson.Get(rec.Body.String(), "id").String()) + require.Equal(t, int64(33), gjson.Get(rec.Body.String(), "usage.total_tokens").Int()) + require.Contains(t, rec.Body.String(), "processing data: 1,2,3 then event: click finished") +} + func TestHandleSSEToJSON_ReconstructsImageGenerationOutputItemDone(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() diff --git a/deploy/Dockerfile b/deploy/Dockerfile index aacd121c12..d83b0e25b8 100644 --- a/deploy/Dockerfile +++ b/deploy/Dockerfile @@ -18,9 +18,11 @@ ARG GOSUMDB=sum.golang.google.cn FROM ${NODE_IMAGE} AS frontend-builder WORKDIR /app/frontend +ENV NODE_OPTIONS=--max-old-space-size=1536 -# Install pnpm. Keep this aligned with CI to avoid lockfile metadata drift. -RUN corepack enable && corepack prepare pnpm@9 --activate +# Install pnpm. Keep this pinned to the lockfile-compatible major version so +# Docker builds remain reproducible when pnpm changes config validation rules. +RUN corepack enable && corepack prepare pnpm@9.15.9 --activate # Install dependencies first (better caching) COPY frontend/package.json frontend/pnpm-lock.yaml ./ diff --git a/frontend/src/composables/useModelWhitelist.ts b/frontend/src/composables/useModelWhitelist.ts index deb6f9d83d..28bc1d28ad 100644 --- a/frontend/src/composables/useModelWhitelist.ts +++ b/frontend/src/composables/useModelWhitelist.ts @@ -137,12 +137,14 @@ const metaModels = [ const xaiModels = [ 'grok-4.3', 'grok-build-0.1', + 'grok-composer-2.5-fast', 'grok-4.20-0309-reasoning', 'grok-4.20-0309-non-reasoning', 'grok-4.20-multi-agent-0309', 'grok', 'grok-latest', 'grok-build', + 'grok-composer', 'grok-4.20-reasoning', 'grok-4.20-non-reasoning', 'grok-imagine', @@ -297,6 +299,7 @@ const grokPresetMappings = [ { label: 'Grok 4.3', from: 'grok-4.3', to: 'grok-4.3', color: 'bg-slate-100 text-slate-700 hover:bg-slate-200 dark:bg-slate-800/50 dark:text-slate-300' }, { label: 'Grok Latest', from: 'grok-latest', to: 'grok-4.3', color: 'bg-emerald-100 text-emerald-700 hover:bg-emerald-200 dark:bg-emerald-900/30 dark:text-emerald-400' }, { label: 'Build 0.1', from: 'grok-build', to: 'grok-build-0.1', color: 'bg-cyan-100 text-cyan-700 hover:bg-cyan-200 dark:bg-cyan-900/30 dark:text-cyan-400' }, + { label: 'Composer 2.5', from: 'grok-composer', to: 'grok-composer-2.5-fast', color: 'bg-teal-100 text-teal-700 hover:bg-teal-200 dark:bg-teal-900/30 dark:text-teal-400' }, { label: '4.20 Reasoning', from: 'grok-4.20-reasoning', to: 'grok-4.20-0309-reasoning', color: 'bg-indigo-100 text-indigo-700 hover:bg-indigo-200 dark:bg-indigo-900/30 dark:text-indigo-400' }, { label: '4.20 Non Reasoning', from: 'grok-4.20-non-reasoning', to: 'grok-4.20-0309-non-reasoning', color: 'bg-violet-100 text-violet-700 hover:bg-violet-200 dark:bg-violet-900/30 dark:text-violet-400' }, { label: 'Imagine Image', from: 'grok-imagine', to: 'grok-imagine-image-quality', color: 'bg-sky-100 text-sky-700 hover:bg-sky-200 dark:bg-sky-900/30 dark:text-sky-400' },