mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-01 15:02:58 +08:00
fix(gateway): normalize Claude Code 1m model suffix
This commit is contained in:
@@ -166,6 +166,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
||||
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
|
||||
return
|
||||
}
|
||||
body = parsedReq.Body.Bytes()
|
||||
reqModel := parsedReq.Model
|
||||
reqStream := parsedReq.Stream
|
||||
reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", reqStream))
|
||||
@@ -1882,6 +1883,7 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) {
|
||||
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
|
||||
return
|
||||
}
|
||||
body = parsedReq.Body.Bytes()
|
||||
// count_tokens 走 messages 严格校验时,复用已解析请求,避免二次反序列化。
|
||||
SetClaudeCodeClientContext(c, body, parsedReq)
|
||||
reqLog = reqLog.With(zap.String("model", parsedReq.Model), zap.Bool("stream", parsedReq.Stream))
|
||||
|
||||
@@ -67,6 +67,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
|
||||
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
|
||||
return
|
||||
}
|
||||
body = parsedReq.Body.Bytes()
|
||||
if parsedReq.Model == "" {
|
||||
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
|
||||
return
|
||||
|
||||
@@ -161,6 +161,18 @@ func setGatewayRequestRanges(parsed *ParsedRequest, protocol string, jsonStr str
|
||||
}
|
||||
}
|
||||
|
||||
const claudeCodeLongContextModelSuffix = "[1m]"
|
||||
|
||||
// Claude Code treats [1m] as a client-side context selector and normally removes it
|
||||
// before provider requests. Normalize leaked suffixes, including its duplicated form.
|
||||
func normalizeClaudeCodeLongContextModel(model string) string {
|
||||
for len(model) > len(claudeCodeLongContextModelSuffix) &&
|
||||
strings.EqualFold(model[len(model)-len(claudeCodeLongContextModelSuffix):], claudeCodeLongContextModelSuffix) {
|
||||
model = model[:len(model)-len(claudeCodeLongContextModelSuffix)]
|
||||
}
|
||||
return model
|
||||
}
|
||||
|
||||
// parseGatewayRequestCurrentBody 只做标量和 raw range 轻量解析,不恢复 system/messages 对象图。
|
||||
func parseGatewayRequestCurrentBody(parsed *ParsedRequest, protocol string) error {
|
||||
if parsed == nil || parsed.Body == nil {
|
||||
@@ -183,6 +195,19 @@ func parseGatewayRequestCurrentBody(parsed *ParsedRequest, protocol string) erro
|
||||
return fmt.Errorf("invalid model field type")
|
||||
}
|
||||
parsed.Model = modelResult.String()
|
||||
if protocol == domain.PlatformAnthropic {
|
||||
normalizedModel := normalizeClaudeCodeLongContextModel(parsed.Model)
|
||||
if normalizedModel != parsed.Model {
|
||||
normalizedBody, err := sjson.SetBytes(bodyBytes, "model", normalizedModel)
|
||||
if err != nil {
|
||||
return fmt.Errorf("normalize model field: %w", err)
|
||||
}
|
||||
parsed.Body.Replace(normalizedBody)
|
||||
bodyBytes = normalizedBody
|
||||
jsonStr = *(*string)(unsafe.Pointer(&bodyBytes))
|
||||
parsed.Model = normalizedModel
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
streamResult := gjson.Get(jsonStr, "stream")
|
||||
|
||||
@@ -77,6 +77,40 @@ func TestParseGatewayRequest_InvalidStreamType(t *testing.T) {
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestParseGatewayRequest_AnthropicNormalizesClaudeCodeLongContextModelSuffix(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
model string
|
||||
want string
|
||||
}{
|
||||
{name: "lowercase suffix", model: "claude-opus-4-8[1m]", want: "claude-opus-4-8"},
|
||||
{name: "uppercase suffix", model: "claude-opus-4-8[1M]", want: "claude-opus-4-8"},
|
||||
{name: "duplicated suffix", model: "claude-opus-4-8[1M][1m]", want: "claude-opus-4-8"},
|
||||
{name: "suffix in middle", model: "claude-opus-4-8[1m]-preview", want: "claude-opus-4-8[1m]-preview"},
|
||||
{name: "suffix only", model: "[1m]", want: "[1m]"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
body := []byte(fmt.Sprintf(`{"model":%q,"system":"test","messages":[{"role":"user","content":"hi"}]}`, tt.model))
|
||||
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.want, parsed.Model)
|
||||
require.Equal(t, tt.want, gjson.GetBytes(parsed.Body.Bytes(), "model").String())
|
||||
require.Equal(t, `"test"`, string(parsed.SystemRaw()))
|
||||
require.NotEmpty(t, parsed.MessagesRaw())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseGatewayRequest_NonAnthropicPreservesClaudeCodeLongContextModelSuffix(t *testing.T) {
|
||||
body := []byte(`{"model":"claude-opus-4-8[1m]","input":"hi"}`)
|
||||
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "responses")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "claude-opus-4-8[1m]", parsed.Model)
|
||||
require.Equal(t, "claude-opus-4-8[1m]", gjson.GetBytes(parsed.Body.Bytes(), "model").String())
|
||||
}
|
||||
|
||||
func TestParseGatewayRequest_ResponsesInput(t *testing.T) {
|
||||
body := []byte(`{"model":"gpt-5.1","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"hello"}]}]}`)
|
||||
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "responses")
|
||||
|
||||
Reference in New Issue
Block a user