mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge remote-tracking branch 'origin/main' into feature/batch-image-foundation
# Conflicts: # deploy/Dockerfile
This commit is contained in:
@@ -1 +1 @@
|
||||
0.1.145
|
||||
0.1.146
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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"])
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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 保存网关请求的预解析结果
|
||||
//
|
||||
// 性能优化说明:
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
+4
-2
@@ -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 ./
|
||||
|
||||
@@ -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' },
|
||||
|
||||
Reference in New Issue
Block a user