Commit 0ed497f0 by Calcium-Ion Committed by GitHub

feat(relay): hosted-tool conversion fidelity, reasoning normalization, and…

feat(relay): hosted-tool conversion fidelity, reasoning normalization, and billing usage integrity (#7137)

* feat(relaykit): preserve hosted tools across conversions

- add protocol-neutral hosted-tool DTOs, conversion metadata, and loss policies
- bridge citations, grounding metadata, and hosted-tool stream lifecycles
- document the public conversion behavior and channel policy controls

* refactor(relaykit): normalize reasoning and thinking intent

- centralize provider-neutral reasoning intent, effort, and budget mappings
- parse model suffixes at the host entry boundary while preserving provider-owned tails
- keep adaptive Claude thinking and explicit zero-token compatibility consistent

* fix(billing): preserve authoritative usage across relay hops

- carry native BillingUsage sidecars through direct and streamed protocol bridges
- merge partial and terminal usage monotonically with safe fallback settlement
- retain cache metadata, penultimate usage, and per-call Gemini tool surcharges

* feat(relay): bridge Responses with Claude and Gemini protocols

- add direct request, response, and stream converters across supported relay formats
- expose Claude count_tokens and Chat-to-Responses compatibility endpoints
- carry conversion diagnostics through the host while retaining the curated public goldens

* fix(relay): wire relaykit conversions into host channels

- connect handlers, adaptors, and channel settings to the standalone conversion layer
- keep model mapping, pricing identity, retries, and provider-specific suffix behavior aligned
- ignore local audit artifacts and retain focused public regression coverage
parent b7017c25
.idea .idea
.review
.vscode .vscode
.zed .zed
.history .history
...@@ -20,7 +21,7 @@ tiktoken_cache ...@@ -20,7 +21,7 @@ tiktoken_cache
.gocache .gocache
.gomodcache/ .gomodcache/
.cache .cache
plans .plans
.claude .claude
.cursor .cursor
...@@ -37,7 +38,7 @@ skills-lock.json ...@@ -37,7 +38,7 @@ skills-lock.json
# Local-only live probes and scratch test workspaces. # Local-only live probes and scratch test workspaces.
.local-tests/ .local-tests/
service/relayconvert/chat_responses_live_local_test.go relaykit/relayconvert/chat_responses_live_local_test.go
service/openaicompat/chat_responses_live_local_test.go service/openaicompat/chat_responses_live_local_test.go
go.work go.work
go.work.sum go.work.sum
...@@ -259,6 +259,13 @@ func testChannel(ctx context.Context, channel *model.Channel, testUserID int, te ...@@ -259,6 +259,13 @@ func testChannel(ctx context.Context, channel *model.Channel, testUserID int, te
newAPIError: types.NewError(err, types.ErrorCodeChannelModelMappedError), newAPIError: types.NewError(err, types.ErrorCodeChannelModelMappedError),
} }
} }
if err = helper.ApplyReasoningModelSuffix(info, request); err != nil {
return testResult{
context: c,
localErr: err,
newAPIError: types.NewErrorWithStatusCode(err, types.ErrorCodeConvertRequestFailed, http.StatusBadRequest, types.ErrOptionWithSkipRetry()),
}
}
testModel = info.UpstreamModelName testModel = info.UpstreamModelName
// 更新请求中的模型名称 // 更新请求中的模型名称
...@@ -943,7 +950,7 @@ func testChannelForHealthCheck(ctx context.Context, channel *model.Channel, test ...@@ -943,7 +950,7 @@ func testChannelForHealthCheck(ctx context.Context, channel *model.Channel, test
} }
if allowDisable && isChannelEnabled && shouldBanChannel && channel.GetAutoBan() { if allowDisable && isChannelEnabled && shouldBanChannel && channel.GetAutoBan() {
processChannelError(result.context, *types.NewChannelError(channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey, common.GetContextKeyString(result.context, constant.ContextKeyChannelKey), channel.GetAutoBan()), newAPIError) processChannelError(result.context, *types.NewChannelError(channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey, common.GetContextKeyString(result.context, constant.ContextKeyChannelKey), channel.GetAutoBan()), newAPIError, nil)
summary.Disabled++ summary.Disabled++
} }
......
...@@ -238,7 +238,7 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { ...@@ -238,7 +238,7 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
newAPIError = service.NormalizeViolationFeeError(newAPIError) newAPIError = service.NormalizeViolationFeeError(newAPIError)
relayInfo.LastError = newAPIError relayInfo.LastError = newAPIError
processChannelError(c, *types.NewChannelError(channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey, common.GetContextKeyString(c, constant.ContextKeyChannelKey), channel.GetAutoBan()), newAPIError) processChannelError(c, *types.NewChannelError(channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey, common.GetContextKeyString(c, constant.ContextKeyChannelKey), channel.GetAutoBan()), newAPIError, relayInfo)
if !shouldRetry(c, newAPIError, common.RetryTimes-retryParam.GetRetry()) { if !shouldRetry(c, newAPIError, common.RetryTimes-retryParam.GetRetry()) {
break break
...@@ -257,6 +257,38 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { ...@@ -257,6 +257,38 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
} }
} }
// CountClaudeTokens implements Anthropic's token-counting utility endpoint.
// It deliberately skips upstream generation and billing; callers use this
// endpoint to size prompts before creating a Message.
func CountClaudeTokens(c *gin.Context) {
request, err := helper.GetAndValidateClaudeRequest(c)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{
"type": "error",
"error": gin.H{
"type": "invalid_request_error",
"message": common.MessageWithRequestId(err.Error(), c.GetString(common.RequestIdKey)),
},
})
return
}
info := relaycommon.GenRelayInfoClaude(c, request)
inputTokens, err := service.CountRequestToken(c, request.GetTokenCountMeta(), info)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"type": "error",
"error": gin.H{
"type": "api_error",
"message": common.MessageWithRequestId(err.Error(), c.GetString(common.RequestIdKey)),
},
})
return
}
c.JSON(http.StatusOK, gin.H{"input_tokens": inputTokens})
}
var upgrader = websocket.Upgrader{ var upgrader = websocket.Upgrader{
Subprotocols: []string{"realtime"}, // WS 握手支持的协议,如果有使用 Sec-WebSocket-Protocol,则必须在此声明对应的 Protocol TODO add other protocol Subprotocols: []string{"realtime"}, // WS 握手支持的协议,如果有使用 Sec-WebSocket-Protocol,则必须在此声明对应的 Protocol TODO add other protocol
CheckOrigin: func(r *http.Request) bool { CheckOrigin: func(r *http.Request) bool {
...@@ -362,7 +394,7 @@ func shouldRetry(c *gin.Context, openaiErr *types.NewAPIError, retryTimes int) b ...@@ -362,7 +394,7 @@ func shouldRetry(c *gin.Context, openaiErr *types.NewAPIError, retryTimes int) b
return operation_setting.ShouldRetryByStatusCode(code) return operation_setting.ShouldRetryByStatusCode(code)
} }
func processChannelError(c *gin.Context, channelError types.ChannelError, err *types.NewAPIError) { func processChannelError(c *gin.Context, channelError types.ChannelError, err *types.NewAPIError, relayInfo *relaycommon.RelayInfo) {
logger.LogError(c, fmt.Sprintf("channel error (channel #%d, status code: %d): %s", channelError.ChannelId, err.StatusCode, common.LocalLogPreview(err.Error()))) logger.LogError(c, fmt.Sprintf("channel error (channel #%d, status code: %d): %s", channelError.ChannelId, err.StatusCode, common.LocalLogPreview(err.Error())))
// 不要使用context获取渠道信息,异步处理时可能会出现渠道信息不一致的情况 // 不要使用context获取渠道信息,异步处理时可能会出现渠道信息不一致的情况
// do not use context to get channel info, there may be inconsistent channel info when processing asynchronously // do not use context to get channel info, there may be inconsistent channel info when processing asynchronously
...@@ -392,6 +424,14 @@ func processChannelError(c *gin.Context, channelError types.ChannelError, err *t ...@@ -392,6 +424,14 @@ func processChannelError(c *gin.Context, channelError types.ChannelError, err *t
other["channel_type"] = c.GetInt("channel_type") other["channel_type"] = c.GetInt("channel_type")
adminInfo := make(map[string]interface{}) adminInfo := make(map[string]interface{})
adminInfo["use_channel"] = c.GetStringSlice("use_channel") adminInfo["use_channel"] = c.GetStringSlice("use_channel")
if relayInfo != nil {
if diagnostics := relayInfo.ConversionDiagnostics(); len(diagnostics) > 0 {
adminInfo["conversion_diagnostics"] = diagnostics
}
if relayInfo.ConversionDiagnosticsTruncated() {
adminInfo["conversion_diagnostics_truncated"] = true
}
}
isMultiKey := common.GetContextKeyBool(c, constant.ContextKeyChannelIsMultiKey) isMultiKey := common.GetContextKeyBool(c, constant.ContextKeyChannelIsMultiKey)
if isMultiKey { if isMultiKey {
adminInfo["is_multi_key"] = true adminInfo["is_multi_key"] = true
...@@ -655,7 +695,8 @@ func executeTaskSubmissionWith( ...@@ -655,7 +695,8 @@ func executeTaskSubmissionWith(
processChannelError(c, processChannelError(c,
*types.NewChannelError(channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey, *types.NewChannelError(channel.Id, channel.Type, channel.Name, channel.ChannelInfo.IsMultiKey,
common.GetContextKeyString(c, constant.ContextKeyChannelKey), channel.GetAutoBan()), common.GetContextKeyString(c, constant.ContextKeyChannelKey), channel.GetAutoBan()),
types.NewOpenAIError(taskErr.Error, types.ErrorCodeBadResponseStatusCode, taskErr.StatusCode)) types.NewOpenAIError(taskErr.Error, types.ErrorCodeBadResponseStatusCode, taskErr.StatusCode),
relayInfo)
} }
willRetry := shouldRetryTaskRelay(c, channel.Id, taskErr, common.RetryTimes-retryParam.GetRetry()) willRetry := shouldRetryTaskRelay(c, channel.Id, taskErr, common.RetryTimes-retryParam.GetRetry())
......
package controller
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestCountClaudeTokensReturnsInputTokensWhenRelayCountingDisabled(t *testing.T) {
gin.SetMode(gin.TestMode)
originalCountToken := constant.CountToken
constant.CountToken = false
t.Cleanup(func() {
constant.CountToken = originalCountToken
})
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Request = httptest.NewRequest(
http.MethodPost,
"/v1/messages/count_tokens?beta=true",
strings.NewReader(`{
"model":"gemini-3.6-flash",
"messages":[{"role":"user","content":"count this prompt"}],
"tools":[{"name":"lookup","description":"Look up a value","input_schema":{"type":"object","properties":{"query":{"type":"string"}}}}]
}`),
)
ctx.Request.Header.Set("Content-Type", "application/json")
common.SetContextKey(ctx, constant.ContextKeyOriginalModel, "gemini-3.6-flash")
CountClaudeTokens(ctx)
require.Equal(t, http.StatusOK, recorder.Code)
var response struct {
InputTokens int `json:"input_tokens"`
}
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
assert.Positive(t, response.InputTokens)
}
func TestCountClaudeTokensRejectsMissingMessages(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Request = httptest.NewRequest(
http.MethodPost,
"/v1/messages/count_tokens",
strings.NewReader(`{"model":"gemini-3.6-flash"}`),
)
ctx.Request.Header.Set("Content-Type", "application/json")
CountClaudeTokens(ctx)
require.Equal(t, http.StatusBadRequest, recorder.Code)
var response struct {
Type string `json:"type"`
Error struct {
Type string `json:"type"`
Message string `json:"message"`
} `json:"error"`
}
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
assert.Equal(t, "error", response.Type)
assert.Equal(t, "invalid_request_error", response.Error.Type)
assert.Contains(t, response.Error.Message, "messages")
}
...@@ -989,6 +989,9 @@ func (channel *Channel) ValidateSettings() error { ...@@ -989,6 +989,9 @@ func (channel *Channel) ValidateSettings() error {
return err return err
} }
} }
if err := channelOtherSettings.ValidateToolLossPolicy(); err != nil {
return err
}
if channel.Type == constant.ChannelTypeAdvancedCustom { if channel.Type == constant.ChannelTypeAdvancedCustom {
if channelOtherSettings.AdvancedCustom == nil { if channelOtherSettings.AdvancedCustom == nil {
return fmt.Errorf("advanced_custom is required") return fmt.Errorf("advanced_custom is required")
......
...@@ -39,6 +39,10 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt ...@@ -39,6 +39,10 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt
} }
func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) { func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) {
claudeAdaptor := claude.Adaptor{}
if _, err := claudeAdaptor.ConvertClaudeRequest(c, info, request); err != nil {
return nil, err
}
for i, message := range request.Messages { for i, message := range request.Messages {
updated := false updated := false
if !message.IsStringContent() { if !message.IsStringContent() {
......
...@@ -357,7 +357,7 @@ func TestAwsStreamHandlerUsesFinalUpstreamUsage(t *testing.T) { ...@@ -357,7 +357,7 @@ func TestAwsStreamHandlerUsesFinalUpstreamUsage(t *testing.T) {
assert.Contains(t, recorder.Body.String(), "[DONE]") assert.Contains(t, recorder.Body.String(), "[DONE]")
} }
func TestAwsStreamHandlerStopsAtClientCancellationAndKeepsPartialBillingUsage(t *testing.T) { func TestAwsStreamHandlerStopsAtClientCancellation(t *testing.T) {
originalRelayTimeout := common.RelayTimeout originalRelayTimeout := common.RelayTimeout
common.RelayTimeout = 0 common.RelayTimeout = 0
t.Cleanup(func() { t.Cleanup(func() {
...@@ -439,12 +439,6 @@ func TestAwsStreamHandlerStopsAtClientCancellationAndKeepsPartialBillingUsage(t ...@@ -439,12 +439,6 @@ func TestAwsStreamHandlerStopsAtClientCancellationAndKeepsPartialBillingUsage(t
require.ErrorIs(t, upstreamContext.Err(), context.Canceled) require.ErrorIs(t, upstreamContext.Err(), context.Canceled)
require.Nil(t, result.err) require.Nil(t, result.err)
require.NotNil(t, result.usage) require.NotNil(t, result.usage)
require.NotNil(t, result.usage.BillingUsage)
require.NotNil(t, result.usage.BillingUsage.ClaudeUsage)
assert.Equal(t, dto.BillingUsageSourceClaudeMessages, result.usage.BillingUsage.Source)
assert.Equal(t, dto.BillingUsageSemanticAnthropic, result.usage.BillingUsage.Semantic)
assert.Equal(t, 100, result.usage.BillingUsage.ClaudeUsage.InputTokens)
assert.Equal(t, 1, result.usage.BillingUsage.ClaudeUsage.OutputTokens)
assert.Equal(t, bodyLengthBeforeCancel, responseWriter.Body.Len()) assert.Equal(t, bodyLengthBeforeCancel, responseWriter.Body.Len())
assert.NotContains(t, responseWriter.Body.String(), "[DONE]") assert.NotContains(t, responseWriter.Body.String(), "[DONE]")
......
...@@ -12,6 +12,7 @@ import ( ...@@ -12,6 +12,7 @@ import (
"github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert" "github.com/QuantumNous/new-api/relaykit/relayconvert"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/model_setting" "github.com/QuantumNous/new-api/setting/model_setting"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
...@@ -26,6 +27,22 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt ...@@ -26,6 +27,22 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt
} }
func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) { func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) {
if request.MaxTokens != nil && *request.MaxTokens == 0 {
request.MaxTokens = nil
}
if err := relayconvert.ApplyClaudeThinkingModel(request, info); err != nil {
return nil, err
}
if request.MaxTokens == nil {
defaultMaxTokens := uint(model_setting.GetClaudeSettings().GetDefaultMaxTokens(request.Model))
request.MaxTokens = &defaultMaxTokens
}
// ApplyClaudeThinkingModel no longer rewrites request.Model. Do not write
// a still-suffixed name back over the entry-normalized UpstreamModelName
// (AWS/Vertex look up getAwsModelID / claudeModelMap from that field).
if info.UpstreamModelName == "" {
info.UpstreamModelName = request.Model
}
return request, nil return request, nil
} }
...@@ -96,7 +113,7 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn ...@@ -96,7 +113,7 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn
if request == nil { if request == nil {
return nil, errors.New("request is nil") return nil, errors.New("request is nil")
} }
result, err := relayconvert.ConvertRequest(c, info, types.RelayFormatClaude, request) result, err := service.ConvertRequest(c, info, types.RelayFormatClaude, request)
if err != nil { if err != nil {
return nil, err return nil, err
} }
...@@ -113,8 +130,15 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela ...@@ -113,8 +130,15 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela
} }
func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) { func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) {
// TODO implement me result, err := service.ConvertRequest(c, info, types.RelayFormatClaude, &request)
return nil, errors.New("not implemented") if err != nil {
return nil, err
}
claudeRequest, ok := result.Value.(*dto.ClaudeRequest)
if !ok {
return nil, fmt.Errorf("expected Anthropic Messages request, got %T", result.Value)
}
return claudeRequest, nil
} }
func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (any, error) { func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (any, error) {
...@@ -123,6 +147,9 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request ...@@ -123,6 +147,9 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) { func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
info.FinalRequestRelayFormat = types.RelayFormatClaude info.FinalRequestRelayFormat = types.RelayFormatClaude
if info.RelayFormat == types.RelayFormatOpenAIResponses && info.IsStream {
return ClaudeResponsesStreamHandler(c, resp, info)
}
if info.IsStream { if info.IsStream {
return ClaudeStreamHandler(c, resp, info) return ClaudeStreamHandler(c, resp, info)
} else { } else {
......
package claude
import (
"net/http/httptest"
"testing"
"github.com/QuantumNous/new-api/common"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/setting/model_setting"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestConvertClaudeRequestTreatsZeroMaxTokensAsUnset(t *testing.T) {
zero := uint(0)
req := &dto.ClaudeRequest{
Model: "claude-sonnet-4-5",
MaxTokens: &zero,
Messages: []dto.ClaudeMessage{
{Role: "user", Content: "hello"},
},
}
info := &relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-sonnet-4-5",
},
}
out, err := (&Adaptor{}).ConvertClaudeRequest(nil, info, req)
require.NoError(t, err)
converted, ok := out.(*dto.ClaudeRequest)
require.True(t, ok)
require.NotNil(t, converted.MaxTokens)
assert.Equal(t, uint(model_setting.GetClaudeSettings().GetDefaultMaxTokens(req.Model)), *converted.MaxTokens)
}
func TestConvertClaudeRequestZeroMaxTokensStillRaisesThinkingBudget(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
zero := uint(0)
original := &dto.ClaudeRequest{
Model: "claude-3-7-sonnet-thinking",
MaxTokens: &zero,
Messages: []dto.ClaudeMessage{
{Role: "user", Content: "hello"},
},
}
info := &relaycommon.RelayInfo{
OriginModelName: "claude-3-7-sonnet-thinking",
Request: original,
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-3-7-sonnet-thinking",
},
}
outbound, err := common.DeepCopy(original)
require.NoError(t, err)
require.NoError(t, helper.ModelMappedHelper(c, info, outbound))
require.NoError(t, helper.ApplyReasoningModelSuffix(info, outbound))
out, err := (&Adaptor{}).ConvertClaudeRequest(nil, info, outbound)
require.NoError(t, err)
converted, ok := out.(*dto.ClaudeRequest)
require.True(t, ok)
assert.Equal(t, "claude-3-7-sonnet", converted.Model)
require.NotNil(t, converted.Thinking)
require.NotNil(t, converted.MaxTokens)
assert.Greater(t, *converted.MaxTokens, uint(1024))
}
func TestConvertClaudeRequestDoesNotOverwriteTrimmedUpstreamModelName(t *testing.T) {
req := &dto.ClaudeRequest{
Model: "claude-3-7-sonnet-thinking",
Messages: []dto.ClaudeMessage{
{Role: "user", Content: "hello"},
},
}
info := &relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-3-7-sonnet",
},
}
_, err := (&Adaptor{}).ConvertClaudeRequest(nil, info, req)
require.NoError(t, err)
assert.Equal(t, "claude-3-7-sonnet", info.UpstreamModelName)
}
package claude package claude
import ( import (
"fmt"
"io" "io"
"net/http" "net/http"
"strings" "strings"
...@@ -19,6 +20,8 @@ import ( ...@@ -19,6 +20,8 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
const claudeToChatStreamStateKey = "relaykit.claude_to_chat_stream_state"
func stopReasonClaude2OpenAI(reason string) string { func stopReasonClaude2OpenAI(reason string) string {
return relayconvert.StopReasonClaudeToOpenAI(reason) return relayconvert.StopReasonClaudeToOpenAI(reason)
} }
...@@ -117,7 +120,14 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -117,7 +120,14 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
countClaudeStreamBillableTools(c, info, &claudeResponse) countClaudeStreamBillableTools(c, info, &claudeResponse)
helper.ClaudeChunkData(c, claudeResponse, data) helper.ClaudeChunkData(c, claudeResponse, data)
} else if info.RelayFormat == types.RelayFormatOpenAI { } else if info.RelayFormat == types.RelayFormatOpenAI {
response := StreamResponseClaude2OpenAI(&claudeResponse) state, err := claudeToChatStreamState(c)
if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody)
}
response, err := state.ConvertChunk(&claudeResponse)
if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody)
}
if !FormatClaudeResponseInfo(&claudeResponse, response, claudeInfo) { if !FormatClaudeResponseInfo(&claudeResponse, response, claudeInfo) {
return nil return nil
...@@ -125,6 +135,9 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -125,6 +135,9 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
countClaudeStreamBillableTools(c, info, &claudeResponse) countClaudeStreamBillableTools(c, info, &claudeResponse)
if response == nil {
return nil
}
err = helper.ObjectData(c, response) err = helper.ObjectData(c, response)
if err != nil { if err != nil {
logger.LogError(c, "send_stream_response_failed: "+err.Error()) logger.LogError(c, "send_stream_response_failed: "+err.Error())
...@@ -133,6 +146,20 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -133,6 +146,20 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
return nil return nil
} }
func claudeToChatStreamState(c *gin.Context) (*relayconvert.ClaudeToChatStreamState, error) {
if value, ok := c.Get(claudeToChatStreamStateKey); ok {
state, ok := value.(*relayconvert.ClaudeToChatStreamState)
if !ok || state == nil {
return nil, fmt.Errorf("invalid Claude-to-Chat stream state %T", value)
}
return state, nil
}
state := relayconvert.NewClaudeToChatStreamState()
c.Set(claudeToChatStreamStateKey, state)
return state, nil
}
func countClaudeStreamBillableTools(c *gin.Context, info *relaycommon.RelayInfo, claudeResponse *dto.ClaudeResponse) { func countClaudeStreamBillableTools(c *gin.Context, info *relaycommon.RelayInfo, claudeResponse *dto.ClaudeResponse) {
if claudeResponse == nil { if claudeResponse == nil {
return return
...@@ -172,9 +199,7 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau ...@@ -172,9 +199,7 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau
if claudeInfo.Usage != nil { if claudeInfo.Usage != nil {
claudeInfo.Usage.UsageSemantic = "anthropic" claudeInfo.Usage.UsageSemantic = "anthropic"
} }
if claudeInfo.Usage != nil && claudeInfo.Usage.BillingUsage == nil { relayconvert.FinalizeClaudeStreamBillingUsage(claudeInfo)
claudeInfo.Usage.BillingUsage = dto.NewClaudeMessagesBillingUsage(buildMessageDeltaPatchUsage(nil, claudeInfo))
}
if info.RelayFormat == types.RelayFormatClaude { if info.RelayFormat == types.RelayFormatClaude {
// //
...@@ -232,7 +257,10 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -232,7 +257,10 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
claudeInfo.Usage.CompletionTokens = claudeResponse.Usage.OutputTokens claudeInfo.Usage.CompletionTokens = claudeResponse.Usage.OutputTokens
claudeInfo.Usage.TotalTokens = claudeResponse.Usage.InputTokens + claudeResponse.Usage.OutputTokens claudeInfo.Usage.TotalTokens = claudeResponse.Usage.InputTokens + claudeResponse.Usage.OutputTokens
claudeInfo.Usage.UsageSemantic = "anthropic" claudeInfo.Usage.UsageSemantic = "anthropic"
claudeInfo.Usage.BillingUsage = dto.NewClaudeMessagesBillingUsage(claudeResponse.Usage) claudeInfo.Usage.BillingUsage = dto.CloneBillingUsage(claudeResponse.Usage.BillingUsage)
if claudeInfo.Usage.BillingUsage == nil {
claudeInfo.Usage.BillingUsage = dto.NewClaudeMessagesBillingUsage(claudeResponse.Usage)
}
claudeInfo.Usage.PromptTokensDetails.CachedTokens = claudeResponse.Usage.CacheReadInputTokens claudeInfo.Usage.PromptTokensDetails.CachedTokens = claudeResponse.Usage.CacheReadInputTokens
claudeInfo.Usage.PromptTokensDetails.CachedCreationTokens = claudeResponse.Usage.CacheCreationInputTokens claudeInfo.Usage.PromptTokensDetails.CachedCreationTokens = claudeResponse.Usage.CacheCreationInputTokens
claudeInfo.Usage.ClaudeCacheCreation5mTokens = claudeResponse.Usage.GetCacheCreation5mTokens() claudeInfo.Usage.ClaudeCacheCreation5mTokens = claudeResponse.Usage.GetCacheCreation5mTokens()
...@@ -247,6 +275,22 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud ...@@ -247,6 +275,22 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody) return types.NewError(err, types.ErrorCodeBadResponseBody)
} }
case types.RelayFormatOpenAIResponses:
convertResult, err := service.ConvertResponse(c, info, types.RelayFormatOpenAIResponses, &claudeResponse)
if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody)
}
responsesResponse, ok := convertResult.Value.(*dto.OpenAIResponsesResponse)
if !ok {
return types.NewError(fmt.Errorf("expected OpenAI Responses response, got %T", convertResult.Value), types.ErrorCodeBadResponseBody)
}
if responseID := helper.GetResponseID(c); responseID != "" {
responsesResponse.ID = responseID
}
responseData, err = common.Marshal(responsesResponse)
if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody)
}
case types.RelayFormatClaude: case types.RelayFormatClaude:
responseData = data responseData = data
} }
......
package claude package claude
import ( import (
"net/http/httptest"
"strings" "strings"
"testing" "testing"
"github.com/QuantumNous/new-api/common"
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert" "github.com/QuantumNous/new-api/relaykit/relayconvert"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
...@@ -323,8 +327,27 @@ func TestBuildOpenAIStyleUsageFromClaudeUsageDefaultsAggregateCacheCreationTo5m( ...@@ -323,8 +327,27 @@ func TestBuildOpenAIStyleUsageFromClaudeUsageDefaultsAggregateCacheCreationTo5m(
require.Equal(t, 0, openAIUsage.ClaudeCacheCreation1hTokens) require.Equal(t, 0, openAIUsage.ClaudeCacheCreation1hTokens)
} }
func applyOpenAIChatReasoningThroughHandlerOrder(t *testing.T, original dto.GeneralOpenAIRequest) (*dto.GeneralOpenAIRequest, *relaycommon.RelayInfo) {
t.Helper()
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
info := &relaycommon.RelayInfo{
OriginModelName: original.Model,
Request: &original,
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: original.Model,
},
}
outbound, err := common.DeepCopy(&original)
require.NoError(t, err)
require.NoError(t, helper.ModelMappedHelper(c, info, outbound))
require.NoError(t, helper.ApplyReasoningModelSuffix(info, outbound))
return outbound, info
}
func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48HighUsesAdaptiveThinking(t *testing.T) { func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48HighUsesAdaptiveThinking(t *testing.T) {
request := dto.GeneralOpenAIRequest{ original := dto.GeneralOpenAIRequest{
Model: "claude-opus-4-8-high", Model: "claude-opus-4-8-high",
Temperature: commonPointer(0.7), Temperature: commonPointer(0.7),
TopP: commonPointer(0.9), TopP: commonPointer(0.9),
...@@ -337,7 +360,8 @@ func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48HighUsesAdaptiveThinking( ...@@ -337,7 +360,8 @@ func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48HighUsesAdaptiveThinking(
}, },
} }
claudeRequest, err := relayconvert.OpenAIChatRequestToClaudeMessages(nil, &relaycommon.RelayInfo{}, request) outbound, info := applyOpenAIChatReasoningThroughHandlerOrder(t, original)
claudeRequest, err := relayconvert.OpenAIChatRequestToClaudeMessages(nil, info, *outbound)
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, "claude-opus-4-8", claudeRequest.Model) require.Equal(t, "claude-opus-4-8", claudeRequest.Model)
require.NotNil(t, claudeRequest.Thinking) require.NotNil(t, claudeRequest.Thinking)
...@@ -350,7 +374,7 @@ func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48HighUsesAdaptiveThinking( ...@@ -350,7 +374,7 @@ func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48HighUsesAdaptiveThinking(
} }
func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48ThinkingUsesAdaptiveHighEffort(t *testing.T) { func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48ThinkingUsesAdaptiveHighEffort(t *testing.T) {
request := dto.GeneralOpenAIRequest{ original := dto.GeneralOpenAIRequest{
Model: "claude-opus-4-8-thinking", Model: "claude-opus-4-8-thinking",
Temperature: commonPointer(0.7), Temperature: commonPointer(0.7),
TopP: commonPointer(0.9), TopP: commonPointer(0.9),
...@@ -363,7 +387,8 @@ func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48ThinkingUsesAdaptiveHighE ...@@ -363,7 +387,8 @@ func TestOpenAIChatRequestToClaudeMessages_ClaudeOpus48ThinkingUsesAdaptiveHighE
}, },
} }
claudeRequest, err := relayconvert.OpenAIChatRequestToClaudeMessages(nil, &relaycommon.RelayInfo{}, request) outbound, info := applyOpenAIChatReasoningThroughHandlerOrder(t, original)
claudeRequest, err := relayconvert.OpenAIChatRequestToClaudeMessages(nil, info, *outbound)
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, "claude-opus-4-8", claudeRequest.Model) require.Equal(t, "claude-opus-4-8", claudeRequest.Model)
require.NotNil(t, claudeRequest.Thinking) require.NotNil(t, claudeRequest.Thinking)
......
package claude
import (
"fmt"
"net/http"
"strings"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/logger"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert"
"github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service"
"github.com/gin-gonic/gin"
)
func ClaudeResponsesStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.Usage, *types.NewAPIError) {
responseID := helper.GetResponseID(c)
created := common.GetTimestamp()
state, err := relayconvert.NewResponseStreamState(types.RelayFormatClaude, types.RelayFormatOpenAIResponses, relayconvert.ResponseStreamOptions{
ID: responseID,
Model: info.UpstreamModelName,
Created: created,
EmitSequenceNumber: true,
})
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
}
hostedBridge := relayconvert.NewClaudeHostedStreamBridge()
claudeInfo := &ClaudeResponseInfo{
ResponseId: responseID,
Created: created,
Model: info.UpstreamModelName,
ResponseText: strings.Builder{},
Usage: &dto.Usage{},
}
var streamErr *types.NewAPIError
// streamFailed means a Responses-native terminal error was sent successfully.
// In that case the scanner stops without a transport error and the partial
// upstream usage remains billable.
streamFailed := false
sendResponsesEvent := func(eventType string, payload dto.ResponsesStreamResponse) bool {
payload.Type = eventType
data, err := common.Marshal(payload)
if err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
return false
}
if err := helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: eventType}, string(data)); err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
return false
}
return true
}
sendResult := func(result relayconvert.ResponseResult) bool {
event, ok := result.Value.(relayconvert.ChatToResponsesStreamEvent)
if !ok {
streamErr = types.NewOpenAIError(
fmt.Errorf("expected OpenAI Responses stream event, got %T", result.Value),
types.ErrorCodeBadResponse,
http.StatusInternalServerError,
)
return false
}
return sendResponsesEvent(event.Type, event.Payload)
}
failResponsesStream := func(err error) bool {
failureResults, handled := state.FailResponsesStream("server_error", err.Error(), "")
if !handled {
return false
}
for _, result := range failureResults {
if !sendResult(result) {
return true
}
}
streamFailed = true
return true
}
helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) {
var claudeResponse dto.ClaudeResponse
if err := common.UnmarshalJsonStr(data, &claudeResponse); err != nil {
logger.LogError(c, "failed to unmarshal Claude stream event: "+err.Error())
if failResponsesStream(err) {
// A nil streamErr here is intentional: the protocol-level failure
// event was delivered, so only the scanner needs to stop.
sr.Stop(streamErr)
return
}
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
sr.Stop(streamErr)
return
}
if claudeError := claudeResponse.GetClaudeError(); claudeError != nil && claudeError.Type != "" {
if failResponsesStream(fmt.Errorf("%s", claudeError.Message)) {
sr.Stop(streamErr)
return
}
streamErr = types.WithClaudeError(*claudeError, http.StatusInternalServerError)
sr.Stop(streamErr)
return
}
if claudeResponse.StopReason != "" {
maybeMarkClaudeRefusal(c, claudeResponse.StopReason)
}
if claudeResponse.Delta != nil && claudeResponse.Delta.StopReason != nil {
maybeMarkClaudeRefusal(c, *claudeResponse.Delta.StopReason)
}
if claudeResponse.Type == "message_start" && claudeResponse.Message != nil {
info.UpstreamModelName = claudeResponse.Message.Model
}
FormatClaudeResponseInfo(&claudeResponse, nil, claudeInfo)
countClaudeStreamBillableTools(c, info, &claudeResponse)
hostedEvents, consumed, err := hostedBridge.Convert(&claudeResponse, state)
if err != nil {
if failResponsesStream(err) {
sr.Stop(streamErr)
return
}
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
sr.Stop(streamErr)
return
}
for _, event := range hostedEvents {
if !sendResponsesEvent(event.Type, event.Payload) {
sr.Stop(streamErr)
return
}
}
if consumed {
return
}
results, err := service.ConvertStreamResponseChunk(c, info, state, &claudeResponse)
if err != nil {
if failResponsesStream(err) {
sr.Stop(streamErr)
return
}
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
sr.Stop(streamErr)
return
}
for _, result := range results {
if !sendResult(result) {
sr.Stop(streamErr)
return
}
}
})
if streamErr != nil {
return nil, streamErr
}
if streamFailed {
return claudeInfo.Usage, nil
}
HandleStreamFinalResponse(c, info, claudeInfo)
openAIUsage := buildOpenAIStyleUsageFromClaudeUsage(claudeInfo.Usage)
state.SetUsage(&openAIUsage)
finalResults, err := service.FinalizeStreamResponse(c, info, state)
if err != nil {
if failResponsesStream(err) {
return claudeInfo.Usage, streamErr
}
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
}
for _, result := range finalResults {
if !sendResult(result) {
return nil, streamErr
}
}
return claudeInfo.Usage, nil
}
...@@ -13,8 +13,8 @@ import ( ...@@ -13,8 +13,8 @@ import (
"github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert" "github.com/QuantumNous/new-api/relaykit/relayconvert"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/model_setting" "github.com/QuantumNous/new-api/setting/model_setting"
"github.com/QuantumNous/new-api/setting/reasoning"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/samber/lo" "github.com/samber/lo"
...@@ -24,6 +24,9 @@ type Adaptor struct { ...@@ -24,6 +24,9 @@ type Adaptor struct {
} }
func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.GeminiChatRequest) (any, error) { func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.GeminiChatRequest) (any, error) {
if err := relayconvert.ApplyGeminiThinkingConfigChecked(request, info); err != nil {
return nil, err
}
if len(request.Contents) > 0 { if len(request.Contents) > 0 {
for i, content := range request.Contents { for i, content := range request.Contents {
if i == 0 { if i == 0 {
...@@ -44,7 +47,7 @@ func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayIn ...@@ -44,7 +47,7 @@ func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayIn
} }
func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, req *dto.ClaudeRequest) (any, error) { func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, req *dto.ClaudeRequest) (any, error) {
result, err := relayconvert.ConvertRequest(c, info, types.RelayFormatGemini, req) result, err := service.ConvertRequest(c, info, types.RelayFormatGemini, req)
if err != nil { if err != nil {
return nil, err return nil, err
} }
...@@ -132,21 +135,6 @@ func (a *Adaptor) Init(info *relaycommon.RelayInfo) { ...@@ -132,21 +135,6 @@ func (a *Adaptor) Init(info *relaycommon.RelayInfo) {
func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) { func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
if model_setting.GetGeminiSettings().ThinkingAdapterEnabled &&
!model_setting.ShouldPreserveThinkingSuffix(info.OriginModelName) {
// 新增逻辑:处理 -thinking-<budget> 格式
if strings.Contains(info.UpstreamModelName, "-thinking-") {
parts := strings.Split(info.UpstreamModelName, "-thinking-")
info.UpstreamModelName = parts[0]
} else if strings.HasSuffix(info.UpstreamModelName, "-thinking") { // 旧的适配
info.UpstreamModelName = strings.TrimSuffix(info.UpstreamModelName, "-thinking")
} else if strings.HasSuffix(info.UpstreamModelName, "-nothinking") {
info.UpstreamModelName = strings.TrimSuffix(info.UpstreamModelName, "-nothinking")
} else if baseModel, level, ok := reasoning.TrimEffortSuffix(info.UpstreamModelName); ok && level != "" {
info.UpstreamModelName = baseModel
}
}
version := model_setting.GetGeminiVersionSetting(info.UpstreamModelName) version := model_setting.GetGeminiVersionSetting(info.UpstreamModelName)
if strings.HasPrefix(info.UpstreamModelName, "imagen") { if strings.HasPrefix(info.UpstreamModelName, "imagen") {
...@@ -183,7 +171,7 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn ...@@ -183,7 +171,7 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn
if request == nil { if request == nil {
return nil, errors.New("request is nil") return nil, errors.New("request is nil")
} }
result, err := relayconvert.ConvertRequest(c, info, types.RelayFormatGemini, request) result, err := service.ConvertRequest(c, info, types.RelayFormatGemini, request)
if err != nil { if err != nil {
return nil, err return nil, err
} }
...@@ -239,7 +227,7 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela ...@@ -239,7 +227,7 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela
} }
func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) { func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) {
result, err := relayconvert.ConvertRequest(c, info, types.RelayFormatGemini, &request) result, err := service.ConvertRequest(c, info, types.RelayFormatGemini, &request)
if err != nil { if err != nil {
return nil, err return nil, err
} }
......
...@@ -34,6 +34,7 @@ func GeminiTextGenerationHandler(c *gin.Context, info *relaycommon.RelayInfo, re ...@@ -34,6 +34,7 @@ func GeminiTextGenerationHandler(c *gin.Context, info *relaycommon.RelayInfo, re
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
countGeminiBillableFunctionCalls(info, &geminiResponse)
if len(geminiResponse.Candidates) == 0 && geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { if len(geminiResponse.Candidates) == 0 && geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil {
common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason)) common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason))
......
...@@ -55,10 +55,13 @@ func patchGeminiZeroCompletionUsage(c *gin.Context, info *relaycommon.RelayInfo, ...@@ -55,10 +55,13 @@ func patchGeminiZeroCompletionUsage(c *gin.Context, info *relaycommon.RelayInfo,
usage.CompletionTokens = imageCount * 1400 usage.CompletionTokens = imageCount * 1400
} }
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
// Overwrite the metadata-derived billing usage: effectiveBillingUsage prefers // Settlement prefers BillingUsage, so fill the missing completion in the
// BillingUsage during settlement, so keeping the prompt-only metadata there // original upstream dialect without discarding cache or modality details.
// would still bill zero completion tokens. if usage.BillingUsage != nil {
usage.BillingUsage = dto.NewEstimatedGeminiChatBillingUsage(usage) usage.BillingUsage = dto.CloneBillingUsageWithEstimatedCompletion(usage.BillingUsage, usage.CompletionTokens)
} else {
usage.BillingUsage = dto.NewEstimatedGeminiChatBillingUsage(usage)
}
} }
func geminiResponseUsageText(response *dto.GeminiChatResponse) string { func geminiResponseUsageText(response *dto.GeminiChatResponse) string {
...@@ -88,6 +91,23 @@ func markGeminiGoogleSearchCall(c *gin.Context, response *dto.GeminiChatResponse ...@@ -88,6 +91,23 @@ func markGeminiGoogleSearchCall(c *gin.Context, response *dto.GeminiChatResponse
} }
} }
func countGeminiBillableFunctionCalls(info *relaycommon.RelayInfo, response *dto.GeminiChatResponse) {
if info == nil || response == nil {
return
}
for _, candidate := range response.Candidates {
for _, part := range candidate.Content.Parts {
if part.FunctionCall == nil {
continue
}
if part.FunctionCall.WillContinue != nil && *part.FunctionCall.WillContinue {
continue
}
info.CountBillableToolCall(dto.BuildInCallFunctionCall, part.FunctionCall.FunctionName)
}
}
}
func buildUsageFromGeminiResponse(c *gin.Context, info *relaycommon.RelayInfo, response *dto.GeminiChatResponse) dto.Usage { func buildUsageFromGeminiResponse(c *gin.Context, info *relaycommon.RelayInfo, response *dto.GeminiChatResponse) dto.Usage {
metadata := response.GetUsageMetadata() metadata := response.GetUsageMetadata()
if dto.HasGeminiUsageMetadataTokens(metadata) { if dto.HasGeminiUsageMetadataTokens(metadata) {
...@@ -148,12 +168,15 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -148,12 +168,15 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
var usage = &dto.Usage{} var usage = &dto.Usage{}
var imageCount int var imageCount int
var hasBillableUsageMetadata bool var hasBillableUsageMetadata bool
var streamErr error
var accumulatedUsageMetadata *dto.GeminiUsageMetadata
responseText := strings.Builder{} responseText := strings.Builder{}
helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) {
var geminiResponse dto.GeminiChatResponse var geminiResponse dto.GeminiChatResponse
if err := common.UnmarshalJsonStr(data, &geminiResponse); err != nil { if err := common.UnmarshalJsonStr(data, &geminiResponse); err != nil {
sr.Stop(fmt.Errorf("unmarshal: %w", err)) streamErr = fmt.Errorf("unmarshal Gemini stream response: %w", err)
sr.Stop(streamErr)
return return
} }
...@@ -162,6 +185,7 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -162,6 +185,7 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
} }
markGeminiGoogleSearchCall(c, &geminiResponse) markGeminiGoogleSearchCall(c, &geminiResponse)
countGeminiBillableFunctionCalls(info, &geminiResponse)
// 统计图片数量 // 统计图片数量
for _, candidate := range geminiResponse.Candidates { for _, candidate := range geminiResponse.Candidates {
...@@ -177,13 +201,19 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -177,13 +201,19 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
// 更新使用量统计 // 更新使用量统计
if metadata := geminiResponse.GetUsageMetadata(); dto.HasGeminiUsageMetadataTokens(metadata) { if metadata := geminiResponse.GetUsageMetadata(); dto.HasGeminiUsageMetadataTokens(metadata) {
mappedUsage := buildUsageFromGeminiMetadata(metadata, info.GetEstimatePromptTokens()) accumulatedUsageMetadata = dto.MergeGeminiUsageMetadataNonZero(accumulatedUsageMetadata, metadata)
mappedUsage := buildUsageFromGeminiMetadata(accumulatedUsageMetadata, info.GetEstimatePromptTokens())
*usage = mappedUsage *usage = mappedUsage
hasBillableUsageMetadata = true hasBillableUsageMetadata = true
} }
if !callback(data, &geminiResponse) { if !callback(data, &geminiResponse) {
sr.Stop(fmt.Errorf("gemini callback stopped")) if isGeminiDownstreamStop(c, info) {
sr.Stop(nil)
return
}
streamErr = errors.New("Gemini stream callback stopped")
sr.Stop(streamErr)
} }
}) })
...@@ -203,9 +233,24 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -203,9 +233,24 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
patchGeminiZeroCompletionUsage(c, info, usage, responseText.String(), imageCount) patchGeminiZeroCompletionUsage(c, info, usage, responseText.String(), imageCount)
} }
if streamErr != nil {
return usage, types.NewOpenAIError(streamErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
if info.StreamStatus != nil && !info.StreamStatus.IsNormalEnd() {
logger.LogWarn(c, fmt.Sprintf("Gemini stream ended unexpectedly: %s", info.StreamStatus.Summary()))
}
return usage, nil return usage, nil
} }
func isGeminiDownstreamStop(c *gin.Context, info *relaycommon.RelayInfo) bool {
if c != nil && c.Request != nil && c.Request.Context().Err() != nil {
return true
}
return info != nil && info.StreamStatus != nil &&
info.StreamStatus.EndReason == relaycommon.StreamEndReasonClientGone
}
func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
id := helper.GetResponseID(c) id := helper.GetResponseID(c)
createAt := common.GetTimestamp() createAt := common.GetTimestamp()
...@@ -323,6 +368,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R ...@@ -323,6 +368,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
markGeminiGoogleSearchCall(c, &geminiResponse) markGeminiGoogleSearchCall(c, &geminiResponse)
countGeminiBillableFunctionCalls(info, &geminiResponse)
if len(geminiResponse.Candidates) == 0 { if len(geminiResponse.Candidates) == 0 {
usage := buildUsageFromGeminiResponse(c, info, &geminiResponse) usage := buildUsageFromGeminiResponse(c, info, &geminiResponse)
...@@ -371,7 +417,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R ...@@ -371,7 +417,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
case types.RelayFormatClaude: case types.RelayFormatClaude:
convertResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatClaude, fullTextResponse) convertResult, err := service.ConvertResponse(c, info, types.RelayFormatClaude, fullTextResponse)
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
......
...@@ -32,6 +32,7 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h ...@@ -32,6 +32,7 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
markGeminiGoogleSearchCall(c, &geminiResponse) markGeminiGoogleSearchCall(c, &geminiResponse)
countGeminiBillableFunctionCalls(info, &geminiResponse)
if len(geminiResponse.Candidates) == 0 { if len(geminiResponse.Candidates) == 0 {
usage := buildUsageFromGeminiResponse(c, info, &geminiResponse) usage := buildUsageFromGeminiResponse(c, info, &geminiResponse)
if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil {
...@@ -50,15 +51,9 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h ...@@ -50,15 +51,9 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h
) )
} }
chatResp := responseGeminiChat2OpenAI(c, &geminiResponse)
chatResp.Model = info.UpstreamModelName
if responseID := helper.GetResponseID(c); responseID != "" {
chatResp.Id = responseID
}
usage := buildUsageFromGeminiResponse(c, info, &geminiResponse) usage := buildUsageFromGeminiResponse(c, info, &geminiResponse)
chatResp.Usage = usage
convertResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatOpenAIResponses, chatResp) convertResult, err := service.ConvertResponse(c, info, types.RelayFormatOpenAIResponses, &geminiResponse)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
...@@ -66,10 +61,11 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h ...@@ -66,10 +61,11 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h
if !ok { if !ok {
return nil, types.NewOpenAIError(fmt.Errorf("expected OpenAI responses response, got %T", convertResult.Value), types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(fmt.Errorf("expected OpenAI responses response, got %T", convertResult.Value), types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
responsesUsage := convertResult.Usage if responseID := helper.GetResponseID(c); responseID != "" {
if responsesUsage == nil || responsesUsage.TotalTokens == 0 { responsesResp.ID = responseID
responsesResp.Usage = relayconvert.UsageFromChatUsage(&usage)
} }
responsesResp.Model = info.UpstreamModelName
responsesResp.Usage = relayconvert.UsageFromChatUsage(&usage)
responseBody, err = common.Marshal(responsesResp) responseBody, err = common.Marshal(responsesResp)
if err != nil { if err != nil {
...@@ -82,17 +78,16 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h ...@@ -82,17 +78,16 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h
func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
responseID := helper.GetResponseID(c) responseID := helper.GetResponseID(c)
created := common.GetTimestamp() created := common.GetTimestamp()
state, err := relayconvert.NewResponseStreamState(types.RelayFormatOpenAI, types.RelayFormatOpenAIResponses, relayconvert.ResponseStreamOptions{ state, err := relayconvert.NewResponseStreamState(types.RelayFormatGemini, types.RelayFormatOpenAIResponses, relayconvert.ResponseStreamOptions{
ID: responseID, ID: responseID,
Model: info.UpstreamModelName, Model: info.UpstreamModelName,
Created: created, Created: created,
EmitSequenceNumber: true,
}) })
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
} }
finishReason := constant.FinishReasonStop hostedBridge := relayconvert.NewGeminiHostedStreamBridge()
toolCallIndexByChoice := make(map[int]map[string]int)
nextToolCallIndexByChoice := make(map[int]int)
var streamErr *types.NewAPIError var streamErr *types.NewAPIError
sendEvent := func(event relayconvert.ChatToResponsesStreamEvent) bool { sendEvent := func(event relayconvert.ChatToResponsesStreamEvent) bool {
...@@ -101,12 +96,37 @@ func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, r ...@@ -101,12 +96,37 @@ func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, r
streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError) streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
return false return false
} }
helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data)) if err := helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data)); err != nil {
if info.StreamStatus != nil {
info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonClientGone, err)
}
return false
}
return true
}
failResponsesStream := func(err error) bool {
failureResults, handled := state.FailResponsesStream("server_error", err.Error(), "")
if !handled {
return false
}
for _, result := range failureResults {
event, ok := result.Value.(relayconvert.ChatToResponsesStreamEvent)
if !ok {
streamErr = types.NewOpenAIError(fmt.Errorf("expected OAI responses stream event, got %T", result.Value), types.ErrorCodeBadResponse, http.StatusInternalServerError)
return true
}
if !sendEvent(event) {
return true
}
}
return true return true
} }
sendChunk := func(chunk *dto.ChatCompletionsStreamResponse) bool { sendChunk := func(chunk *dto.GeminiChatResponse) bool {
results, err := relayconvert.ConvertStreamResponseChunk(c, info, state, chunk) results, err := service.ConvertStreamResponseChunk(c, info, state, chunk)
if err != nil { if err != nil {
if failResponsesStream(err) {
return false
}
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
return false return false
} }
...@@ -123,58 +143,46 @@ func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, r ...@@ -123,58 +143,46 @@ func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, r
return true return true
} }
usage, streamAPIError := geminiStreamHandler(c, info, resp, func(data string, geminiResponse *dto.GeminiChatResponse) bool { usage, streamAPIError := geminiStreamHandler(c, info, resp, func(_ string, geminiResponse *dto.GeminiChatResponse) bool {
response, isStop := streamResponseGeminiChat2OpenAI(geminiResponse) hostedBridge.Observe(geminiResponse)
response.Id = responseID return sendChunk(geminiResponse)
response.Created = created
response.Model = info.UpstreamModelName
if response.IsToolCall() {
finishReason = constant.FinishReasonToolCalls
}
for choiceIdx := range response.Choices {
choiceKey := response.Choices[choiceIdx].Index
for toolIdx := range response.Choices[choiceIdx].Delta.ToolCalls {
tool := &response.Choices[choiceIdx].Delta.ToolCalls[toolIdx]
if tool.ID == "" {
continue
}
indexByID := toolCallIndexByChoice[choiceKey]
if indexByID == nil {
indexByID = make(map[string]int)
toolCallIndexByChoice[choiceKey] = indexByID
}
if idx, ok := indexByID[tool.ID]; ok {
tool.SetIndex(idx)
continue
}
idx := nextToolCallIndexByChoice[choiceKey]
nextToolCallIndexByChoice[choiceKey] = idx + 1
indexByID[tool.ID] = idx
tool.SetIndex(idx)
}
}
if !sendChunk(response) {
return false
}
if isStop {
return sendChunk(helper.GenerateStopResponse(responseID, created, info.UpstreamModelName, finishReason))
}
return true
}) })
if streamAPIError != nil { if streamAPIError != nil {
if failResponsesStream(streamAPIError) && streamErr == nil {
return usage, nil
}
return usage, streamAPIError return usage, streamAPIError
} }
if info.StreamStatus != nil && !info.StreamStatus.IsNormalEnd() {
if info.StreamStatus.EndReason != relaycommon.StreamEndReasonClientGone {
failResponsesStream(fmt.Errorf("gemini stream ended unexpectedly: %s", info.StreamStatus.Summary()))
}
return usage, nil
}
if streamErr != nil { if streamErr != nil {
return nil, streamErr return nil, streamErr
} }
hostedEvents, err := hostedBridge.Finalize(state)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
}
for _, event := range hostedEvents {
if !sendEvent(event) {
if streamErr != nil {
return usage, streamErr
}
return usage, nil
}
}
if usage != nil { if usage != nil {
state.SetUsage(usage) state.SetUsage(usage)
} }
finalResults, err := relayconvert.FinalizeStreamResponse(c, info, state) finalResults, err := service.FinalizeStreamResponse(c, info, state)
if err != nil { if err != nil {
if failResponsesStream(err) {
return usage, streamErr
}
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
} }
for _, result := range finalResults { for _, result := range finalResults {
...@@ -183,7 +191,10 @@ func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, r ...@@ -183,7 +191,10 @@ func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, r
return nil, types.NewOpenAIError(fmt.Errorf("expected OAI responses stream event, got %T", result.Value), types.ErrorCodeBadResponse, http.StatusInternalServerError) return nil, types.NewOpenAIError(fmt.Errorf("expected OAI responses stream event, got %T", result.Value), types.ErrorCodeBadResponse, http.StatusInternalServerError)
} }
if !sendEvent(event) { if !sendEvent(event) {
return nil, streamErr if streamErr != nil {
return usage, streamErr
}
return usage, nil
} }
} }
return usage, nil return usage, nil
......
...@@ -75,14 +75,14 @@ func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayIn ...@@ -75,14 +75,14 @@ func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayIn
if request == nil { if request == nil {
return nil, errors.New("request is nil") return nil, errors.New("request is nil")
} }
return request, nil return a.claudeAdaptor.ConvertClaudeRequest(c, info, request)
} }
func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.GeminiChatRequest) (any, error) { func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.GeminiChatRequest) (any, error) {
if request == nil { if request == nil {
return nil, errors.New("request is nil") return nil, errors.New("request is nil")
} }
return request, nil return a.geminiAdaptor.ConvertGeminiRequest(c, info, request)
} }
func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest) (any, error) { func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest) (any, error) {
......
...@@ -41,33 +41,10 @@ func OaiResponsesToChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp ...@@ -41,33 +41,10 @@ func OaiResponsesToChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) return nil, types.WithOpenAIError(*oaiError, resp.StatusCode)
} }
chatResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatOpenAI, &responsesResp) responseValue, usage, err := convertResponsesResponseForClient(c, info, &responsesResp)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
chatResp, ok := chatResult.Value.(*dto.OpenAITextResponse)
if !ok {
return nil, types.NewOpenAIError(fmt.Errorf("expected OpenAI chat response, got %T", chatResult.Value), types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
if chatID := helper.GetResponseID(c); chatID != "" {
chatResp.Id = chatID
}
usage := chatResult.Usage
if usage == nil || usage.TotalTokens == 0 {
text := service.ExtractOutputTextFromResponses(&responsesResp)
usage = service.ResponseText2Usage(c, text, info.UpstreamModelName, info.GetEstimatePromptTokens())
chatResp.Usage = *usage
}
responseValue := any(chatResp)
if info.RelayFormat != types.RelayFormatOpenAI {
targetResult, err := relayconvert.ConvertResponse(c, info, info.RelayFormat, chatResp)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
responseValue = targetResult.Value
}
responseBody, err := common.Marshal(responseValue) responseBody, err := common.Marshal(responseValue)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
...@@ -150,39 +127,39 @@ func OaiResponsesToChatBufferedStreamHandler(c *gin.Context, info *relaycommon.R ...@@ -150,39 +127,39 @@ func OaiResponsesToChatBufferedStreamHandler(c *gin.Context, info *relaycommon.R
} }
accumulator.SupplementResponseOutput(finalResponse) accumulator.SupplementResponseOutput(finalResponse)
chatResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatOpenAI, finalResponse) responseValue, usage, err := convertResponsesResponseForClient(c, info, finalResponse)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
chatResp, ok := chatResult.Value.(*dto.OpenAITextResponse) responseBody, err := common.Marshal(responseValue)
if !ok { if err != nil {
return nil, types.NewOpenAIError(fmt.Errorf("expected OpenAI chat response, got %T", chatResult.Value), types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
} }
if chatID := helper.GetResponseID(c); chatID != "" {
chatResp.Id = chatID service.IOCopyBytesGracefully(c, resp, responseBody)
return usage, nil
}
func convertResponsesResponseForClient(c *gin.Context, info *relaycommon.RelayInfo, response *dto.OpenAIResponsesResponse) (any, *dto.Usage, error) {
if responseID := helper.GetResponseID(c); responseID != "" {
response.ID = responseID
} }
usage := chatResult.Usage
usage := relayconvert.UsageFromResponsesUsage(response.Usage)
if usage == nil || usage.TotalTokens == 0 { if usage == nil || usage.TotalTokens == 0 {
text := service.ExtractOutputTextFromResponses(finalResponse) text := service.ExtractOutputTextFromResponses(response)
usage = service.ResponseText2Usage(c, text, info.UpstreamModelName, info.GetEstimatePromptTokens()) usage = service.ResponseText2Usage(c, text, info.UpstreamModelName, info.GetEstimatePromptTokens())
chatResp.Usage = *usage response.Usage = relayconvert.UsageFromChatUsage(usage)
} }
responseValue := any(chatResp) result, err := service.ConvertResponse(c, info, info.RelayFormat, response)
if info.RelayFormat != types.RelayFormatOpenAI {
targetResult, err := relayconvert.ConvertResponse(c, info, info.RelayFormat, chatResp)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
responseValue = targetResult.Value
}
responseBody, err := common.Marshal(responseValue)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError) return nil, nil, err
} }
if result.Usage != nil && result.Usage.TotalTokens != 0 {
service.IOCopyBytesGracefully(c, resp, responseBody) usage = result.Usage
return usage, nil }
return result.Value, usage, nil
} }
func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
...@@ -293,7 +270,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo ...@@ -293,7 +270,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
return return
} }
results, err := relayconvert.ConvertStreamResponseChunk(c, info, state, &streamResp) results, err := service.ConvertStreamResponseChunk(c, info, state, &streamResp)
if err != nil { if err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
sr.Stop(streamErr) sr.Stop(streamErr)
...@@ -320,7 +297,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo ...@@ -320,7 +297,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
if info.RelayFormat == types.RelayFormatClaude && info.ClaudeConvertInfo != nil { if info.RelayFormat == types.RelayFormatClaude && info.ClaudeConvertInfo != nil {
info.ClaudeConvertInfo.Usage = usage info.ClaudeConvertInfo.Usage = usage
} }
finalResults, err := relayconvert.FinalizeStreamResponse(c, info, state) finalResults, err := service.FinalizeStreamResponse(c, info, state)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
} }
......
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/constant"
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
...@@ -171,6 +172,50 @@ func TestOaiResponsesToChatBufferedStreamHandlerReturnsJSONFromSSE(t *testing.T) ...@@ -171,6 +172,50 @@ func TestOaiResponsesToChatBufferedStreamHandlerReturnsJSONFromSSE(t *testing.T)
require.Contains(t, got, `"finish_reason":"tool_calls"`) require.Contains(t, got, `"finish_reason":"tool_calls"`)
} }
func TestOaiResponsesToChatBufferedStreamHandlerPreservesInterleavedClaudeContent(t *testing.T) {
oldMode := gin.Mode()
gin.SetMode(gin.TestMode)
t.Cleanup(func() { gin.SetMode(oldMode) })
body := strings.Join([]string{
`data: {"type":"response.output_item.added","output_index":0,"item":{"type":"reasoning","id":"rs_1","summary":[]}}`,
`data: {"type":"response.reasoning_summary_text.delta","output_index":0,"item_id":"rs_1","delta":"**Planning file inspection**"}`,
`data: {"type":"response.output_item.added","output_index":1,"item":{"type":"message","id":"msg_1","role":"assistant","content":[]}}`,
`data: {"type":"response.output_text.delta","output_index":1,"item_id":"msg_1","delta":"I’ll inspect the starter repository."}`,
`data: {"type":"response.output_item.added","output_index":2,"item":{"type":"reasoning","id":"rs_2","summary":[]}}`,
`data: {"type":"response.reasoning_summary_text.delta","output_index":2,"item_id":"rs_2","delta":"**Clarifying environment task requirements**"}`,
`data: {"type":"response.output_item.added","output_index":3,"item":{"type":"message","id":"msg_2","role":"assistant","content":[]}}`,
`data: {"type":"response.output_text.delta","output_index":3,"item_id":"msg_2","delta":"What would you like me to build?"}`,
`data: {"type":"response.done","response":{"id":"resp_1","model":"gpt-test","status":"completed","usage":{"input_tokens":1,"output_tokens":2,"total_tokens":3}}}`,
`data: [DONE]`,
``,
}, "\n")
c, recorder, resp, info := newResponsesChatTestContext(t, body, false)
info.RelayFormat = types.RelayFormatClaude
usage, apiErr := OaiResponsesToChatBufferedStreamHandler(c, info, resp)
require.Nil(t, apiErr)
require.NotNil(t, usage)
assert.Equal(t, 3, usage.TotalTokens)
var claudeResponse dto.ClaudeResponse
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &claudeResponse))
require.Len(t, claudeResponse.Content, 4)
assert.Equal(t, []string{"thinking", "text", "thinking", "text"}, []string{
claudeResponse.Content[0].Type,
claudeResponse.Content[1].Type,
claudeResponse.Content[2].Type,
claudeResponse.Content[3].Type,
})
require.NotNil(t, claudeResponse.Content[0].Thinking)
require.NotNil(t, claudeResponse.Content[2].Thinking)
assert.Equal(t, "**Planning file inspection**", *claudeResponse.Content[0].Thinking)
assert.Equal(t, "I’ll inspect the starter repository.", claudeResponse.Content[1].GetText())
assert.Equal(t, "**Clarifying environment task requirements**", *claudeResponse.Content[2].Thinking)
assert.Equal(t, "What would you like me to build?", claudeResponse.Content[3].GetText())
}
func TestOaiChatToResponsesStreamHandlerConvertsSSEOrderAndUsage(t *testing.T) { func TestOaiChatToResponsesStreamHandlerConvertsSSEOrderAndUsage(t *testing.T) {
oldMode := gin.Mode() oldMode := gin.Mode()
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
......
...@@ -19,16 +19,20 @@ import ( ...@@ -19,16 +19,20 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
const chatToGeminiStreamStateKey = "relaykit.chat_to_gemini_stream_state"
// 辅助函数 // 辅助函数
func HandleStreamFormat(c *gin.Context, info *relaycommon.RelayInfo, data string, forceFormat bool, thinkToContent bool) error { func HandleStreamFormat(c *gin.Context, info *relaycommon.RelayInfo, data string, forceFormat bool, thinkToContent bool) error {
info.SendResponseCount++
switch info.RelayFormat { switch info.RelayFormat {
case types.RelayFormatOpenAI: case types.RelayFormatOpenAI:
info.SendResponseCount++
return sendStreamData(c, info, data, forceFormat, thinkToContent) return sendStreamData(c, info, data, forceFormat, thinkToContent)
case types.RelayFormatClaude: case types.RelayFormatClaude:
info.SendResponseCount++
return handleClaudeFormat(c, data, info) return handleClaudeFormat(c, data, info)
case types.RelayFormatGemini: case types.RelayFormatGemini:
// The stateful relaykit path owns its chunk counter so multi-hop and
// direct conversions observe the same stream state semantics.
return handleGeminiFormat(c, data, info) return handleGeminiFormat(c, data, info)
} }
return nil return nil
...@@ -41,9 +45,9 @@ func handleClaudeFormat(c *gin.Context, data string, info *relaycommon.RelayInfo ...@@ -41,9 +45,9 @@ func handleClaudeFormat(c *gin.Context, data string, info *relaycommon.RelayInfo
} }
if streamResponse.Usage != nil { if streamResponse.Usage != nil {
info.ClaudeConvertInfo.Usage = streamResponse.Usage info.EnsureClaudeConvertInfo().Usage = streamResponse.Usage
} }
result, err := relayconvert.ConvertStreamResponse(c, info, types.RelayFormatClaude, &streamResponse) result, err := service.ConvertStreamResponse(c, info, types.RelayFormatClaude, &streamResponse)
if err != nil { if err != nil {
return err return err
} }
...@@ -64,29 +68,55 @@ func handleGeminiFormat(c *gin.Context, data string, info *relaycommon.RelayInfo ...@@ -64,29 +68,55 @@ func handleGeminiFormat(c *gin.Context, data string, info *relaycommon.RelayInfo
return err return err
} }
result, err := relayconvert.ConvertStreamResponse(c, info, types.RelayFormatGemini, &streamResponse) state, err := chatToGeminiStreamState(c, &streamResponse)
if err != nil { if err != nil {
return err return err
} }
geminiResponse, ok := result.Value.(*dto.GeminiChatResponse) results, err := service.ConvertStreamResponseChunk(c, info, state, &streamResponse)
if !ok { if err != nil {
return fmt.Errorf("expected Gemini stream response, got %T", result.Value) return err
} }
return sendGeminiStreamResults(c, results)
}
// 如果返回 nil,表示没有实际内容,跳过发送 func chatToGeminiStreamState(c *gin.Context, streamResponse *dto.ChatCompletionsStreamResponse) (*relayconvert.ResponseStreamState, error) {
if geminiResponse == nil { if value, ok := c.Get(chatToGeminiStreamStateKey); ok {
return nil state, ok := value.(*relayconvert.ResponseStreamState)
if !ok || state == nil {
return nil, fmt.Errorf("invalid Chat-to-Gemini stream state %T", value)
}
return state, nil
} }
geminiResponseStr, err := common.Marshal(geminiResponse) state, err := relayconvert.NewResponseStreamState(types.RelayFormatOpenAI, types.RelayFormatGemini, relayconvert.ResponseStreamOptions{
ID: streamResponse.Id,
Model: streamResponse.Model,
Created: streamResponse.Created,
})
if err != nil { if err != nil {
logger.LogError(c, "failed to marshal gemini response: "+err.Error()) return nil, err
return err
} }
c.Set(chatToGeminiStreamStateKey, state)
return state, nil
}
// send gemini format response func sendGeminiStreamResults(c *gin.Context, results []relayconvert.ResponseResult) error {
c.Render(-1, common.CustomEvent{Data: "data: " + string(geminiResponseStr)}) for _, result := range results {
_ = helper.FlushWriter(c) geminiResponse, ok := result.Value.(*dto.GeminiChatResponse)
if !ok {
return fmt.Errorf("expected Gemini stream response, got %T", result.Value)
}
if geminiResponse == nil {
continue
}
data, err := common.Marshal(geminiResponse)
if err != nil {
logger.LogError(c, "failed to marshal gemini response: "+err.Error())
return err
}
c.Render(-1, common.CustomEvent{Data: "data: " + string(data)})
_ = helper.FlushWriter(c)
}
return nil return nil
} }
...@@ -148,7 +178,7 @@ func handleLastResponse(lastStreamData string, responseId *string, createAt *int ...@@ -148,7 +178,7 @@ func handleLastResponse(lastStreamData string, responseId *string, createAt *int
if service.ValidUsage(lastStreamResponse.Usage) { if service.ValidUsage(lastStreamResponse.Usage) {
*containStreamUsage = true *containStreamUsage = true
*usage = lastStreamResponse.Usage *usage = dto.MergeUsageNonZero(*usage, lastStreamResponse.Usage)
if !info.ShouldIncludeUsage { if !info.ShouldIncludeUsage {
*shouldSendLastResp = lo.SomeBy(lastStreamResponse.Choices, func(choice dto.ChatCompletionsStreamResponseChoice) bool { *shouldSendLastResp = lo.SomeBy(lastStreamResponse.Choices, func(choice dto.ChatCompletionsStreamResponseChoice) bool {
return choice.Delta.GetContentString() != "" || choice.Delta.GetReasoningContent() != "" return choice.Delta.GetContentString() != "" || choice.Delta.GetReasoningContent() != ""
...@@ -181,7 +211,7 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream ...@@ -181,7 +211,7 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream
info.ClaudeConvertInfo.Usage = usage info.ClaudeConvertInfo.Usage = usage
result, err := relayconvert.ConvertStreamResponse(c, info, types.RelayFormatClaude, &streamResponse) result, err := service.ConvertStreamResponse(c, info, types.RelayFormatClaude, &streamResponse)
if err != nil { if err != nil {
common.SysLog("error converting Claude stream response: " + err.Error()) common.SysLog("error converting Claude stream response: " + err.Error())
return return
...@@ -203,36 +233,31 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream ...@@ -203,36 +233,31 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream
return return
} }
// 这里处理的是 openai 最后一个流响应,其 delta 为空,有 finish_reason 字段 state, err := chatToGeminiStreamState(c, &streamResponse)
// 因此相比较于 google 官方的流响应,由 openai 转换而来会多一个 parts 为空,finishReason 为 STOP 的响应
// 而包含最后一段文本输出的响应(倒数第二个)的 finishReason 为 null
// 暂不知是否有程序会不兼容。
result, err := relayconvert.ConvertStreamResponse(c, info, types.RelayFormatGemini, &streamResponse)
if err != nil { if err != nil {
common.SysLog("error converting Gemini stream response: " + err.Error()) common.SysLog("error creating Gemini stream state: " + err.Error())
return return
} }
geminiResponse, ok := result.Value.(*dto.GeminiChatResponse) state.SetUsage(usage)
if !ok {
common.SysLog(fmt.Sprintf("expected Gemini stream response, got %T", result.Value)) results, err := service.ConvertStreamResponseChunk(c, info, state, &streamResponse)
if err != nil {
common.SysLog("error converting final Gemini stream response: " + err.Error())
return return
} }
if err := sendGeminiStreamResults(c, results); err != nil {
// openai 流响应开头的空数据 common.SysLog("error sending final Gemini stream response: " + err.Error())
if geminiResponse == nil {
return return
} }
geminiResponseStr, err := common.Marshal(geminiResponse) results, err = service.FinalizeStreamResponse(c, info, state)
if err != nil { if err != nil {
common.SysLog("error marshalling gemini response: " + err.Error()) common.SysLog("error finalizing Gemini stream response: " + err.Error())
return return
} }
if err := sendGeminiStreamResults(c, results); err != nil {
// 发送最终的 Gemini 响应 common.SysLog("error sending finalized Gemini stream response: " + err.Error())
c.Render(-1, common.CustomEvent{Data: "data: " + string(geminiResponseStr)}) }
_ = helper.FlushWriter(c)
} }
} }
......
...@@ -13,7 +13,6 @@ import ( ...@@ -13,7 +13,6 @@ import (
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper" "github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/service"
...@@ -118,13 +117,10 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -118,13 +117,10 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
var toolCount int var toolCount int
var usage = &dto.Usage{} var usage = &dto.Usage{}
var lastStreamData string var lastStreamData string
var secondLastStreamData string // 存储倒数第二个stream data,用于音频模型 var secondLastStreamData string // 保留倒数第二个stream data;部分兼容网关把完整usage放在倒数第二个事件
seenStreamToolCalls := make(map[string]struct{}) seenStreamToolCalls := make(map[string]struct{})
var streamFunctionCallNames []string var streamFunctionCallNames []string
// 检查是否为音频模型
isAudioModel := strings.Contains(strings.ToLower(model), "audio")
helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) {
if lastStreamData != "" { if lastStreamData != "" {
if err := HandleStreamFormat(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent); err != nil { if err := HandleStreamFormat(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent); err != nil {
...@@ -133,8 +129,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -133,8 +129,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
} }
} }
if len(data) > 0 { if len(data) > 0 {
// 对音频模型,保存倒数第二个stream data if lastStreamData != "" {
if isAudioModel && lastStreamData != "" {
secondLastStreamData = lastStreamData secondLastStreamData = lastStreamData
} }
...@@ -147,31 +142,36 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -147,31 +142,36 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
} }
}) })
// 对音频模型,从倒数第二个stream data中提取usage信息 // 处理最后的响应
if isAudioModel && secondLastStreamData != "" { shouldSendLastResp := true
if err := handleLastResponse(lastStreamData, &responseId, &createAt, &systemFingerprint, &model, &usage,
&containStreamUsage, info, &shouldSendLastResp); err != nil {
logger.LogError(c, fmt.Sprintf("error handling last response: %s, lastStreamData: [%s]", err.Error(), lastStreamData))
}
// 部分兼容网关把完整的累计usage附在倒数第二个事件上,随后发送一个空的最后事件。
// 仅当最后一个事件没有有效usage时,回退到倒数第二个事件的完整快照。
usageFrame := lastStreamData
if !containStreamUsage && secondLastStreamData != "" {
var streamResp struct { var streamResp struct {
Usage *dto.Usage `json:"usage"` Usage *dto.Usage `json:"usage"`
} }
err := common.Unmarshal([]byte(secondLastStreamData), &streamResp) err := common.Unmarshal([]byte(secondLastStreamData), &streamResp)
if err == nil && streamResp.Usage != nil && service.ValidUsage(streamResp.Usage) { if err == nil && streamResp.Usage != nil &&
usage = streamResp.Usage streamResp.Usage.PromptTokens > 0 &&
(streamResp.Usage.CompletionTokens > 0 || streamResp.Usage.TotalTokens > 0) {
usage = dto.MergeUsageNonZero(usage, streamResp.Usage)
containStreamUsage = true containStreamUsage = true
usageFrame = secondLastStreamData
if common.DebugEnabled { if common.DebugEnabled {
logger.LogDebug(c, "Audio model usage extracted from second last SSE: PromptTokens=%d, CompletionTokens=%d, TotalTokens=%d, InputTokens=%d, OutputTokens=%d", logger.LogDebug(c, "usage extracted from second last SSE: PromptTokens=%d, CompletionTokens=%d, TotalTokens=%d, InputTokens=%d, OutputTokens=%d",
usage.PromptTokens, usage.CompletionTokens, usage.TotalTokens, usage.PromptTokens, usage.CompletionTokens, usage.TotalTokens,
usage.InputTokens, usage.OutputTokens) usage.InputTokens, usage.OutputTokens)
} }
} }
} }
// 处理最后的响应
shouldSendLastResp := true
if err := handleLastResponse(lastStreamData, &responseId, &createAt, &systemFingerprint, &model, &usage,
&containStreamUsage, info, &shouldSendLastResp); err != nil {
logger.LogError(c, fmt.Sprintf("error handling last response: %s, lastStreamData: [%s]", err.Error(), lastStreamData))
}
if info.RelayFormat == types.RelayFormatOpenAI { if info.RelayFormat == types.RelayFormatOpenAI {
if shouldSendLastResp { if shouldSendLastResp {
_ = sendStreamData(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent) _ = sendStreamData(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent)
...@@ -183,7 +183,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re ...@@ -183,7 +183,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
usage.CompletionTokens += toolCount * 7 usage.CompletionTokens += toolCount * 7
} }
applyUsagePostProcessing(info, usage, common.StringToByteSlice(lastStreamData)) applyUsagePostProcessing(info, usage, common.StringToByteSlice(usageFrame))
for _, name := range streamFunctionCallNames { for _, name := range streamFunctionCallNames {
info.CountBillableToolCall(dto.BuildInCallFunctionCall, name) info.CountBillableToolCall(dto.BuildInCallFunctionCall, name)
...@@ -201,7 +201,7 @@ func collectStreamFunctionCallNames(data string, seen map[string]struct{}, names ...@@ -201,7 +201,7 @@ func collectStreamFunctionCallNames(data string, seen map[string]struct{}, names
} }
for _, choice := range streamResponse.Choices { for _, choice := range streamResponse.Choices {
for i, tc := range choice.Delta.ToolCalls { for i, tc := range choice.Delta.ToolCalls {
name := tc.Function.Name name := strings.TrimSpace(tc.Function.Name)
if name == "" { if name == "" {
continue continue
} }
...@@ -209,11 +209,30 @@ func collectStreamFunctionCallNames(data string, seen map[string]struct{}, names ...@@ -209,11 +209,30 @@ func collectStreamFunctionCallNames(data string, seen map[string]struct{}, names
if tc.Index != nil { if tc.Index != nil {
toolIdx = *tc.Index toolIdx = *tc.Index
} }
key := fmt.Sprintf("%d-%d", choice.Index, toolIdx) fallbackKey := fmt.Sprintf("index\x00%d\x00%d\x00%s", choice.Index, toolIdx, name)
if _, ok := seen[key]; ok { activeKey := fmt.Sprintf("active\x00%d\x00%d\x00%s", choice.Index, toolIdx, name)
continue callID := strings.TrimSpace(tc.ID)
if callID != "" {
idKey := fmt.Sprintf("id\x00%d\x00%s", choice.Index, callID)
if _, ok := seen[idKey]; ok {
continue
}
seen[idKey] = struct{}{}
seen[activeKey] = struct{}{}
if _, delayedID := seen[fallbackKey]; delayedID {
delete(seen, fallbackKey)
continue
}
} else {
if _, ok := seen[fallbackKey]; ok {
continue
}
if _, ok := seen[activeKey]; ok {
continue
}
seen[fallbackKey] = struct{}{}
seen[activeKey] = struct{}{}
} }
seen[key] = struct{}{}
*names = append(*names, name) *names = append(*names, name)
} }
} }
...@@ -280,11 +299,12 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo ...@@ -280,11 +299,12 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
completionTokens += ctkm completionTokens += ctkm
} }
} }
simpleResponse.Usage = dto.Usage{ fallbackUsage := &dto.Usage{
PromptTokens: info.GetEstimatePromptTokens(), PromptTokens: info.GetEstimatePromptTokens(),
CompletionTokens: completionTokens, CompletionTokens: completionTokens,
TotalTokens: info.GetEstimatePromptTokens() + completionTokens, TotalTokens: info.GetEstimatePromptTokens() + completionTokens,
} }
simpleResponse.Usage = *fallbackUsage
usageModified = true usageModified = true
} }
...@@ -310,7 +330,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo ...@@ -310,7 +330,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
break break
} }
case types.RelayFormatClaude: case types.RelayFormatClaude:
convertResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatClaude, &simpleResponse) convertResult, err := service.ConvertResponse(c, info, types.RelayFormatClaude, &simpleResponse)
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
...@@ -320,7 +340,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo ...@@ -320,7 +340,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
} }
responseBody = claudeRespStr responseBody = claudeRespStr
case types.RelayFormatGemini: case types.RelayFormatGemini:
convertResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatGemini, &simpleResponse) convertResult, err := service.ConvertResponse(c, info, types.RelayFormatGemini, &simpleResponse)
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody) return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
} }
......
...@@ -11,6 +11,7 @@ import ( ...@@ -11,6 +11,7 @@ import (
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper" "github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/service"
...@@ -38,16 +39,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -38,16 +39,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
service.IOCopyBytesGracefully(c, resp, responseBody) service.IOCopyBytesGracefully(c, resp, responseBody)
// compute usage // compute usage
usage := dto.Usage{} usage := relayconvert.NormalizeResponsesUsage(responsesResponse.Usage)
if responsesResponse.Usage != nil {
usage.PromptTokens = responsesResponse.Usage.InputTokens
usage.CompletionTokens = responsesResponse.Usage.OutputTokens
usage.TotalTokens = responsesResponse.Usage.TotalTokens
if responsesResponse.Usage.InputTokensDetails != nil {
usage.PromptTokensDetails.CachedTokens = responsesResponse.Usage.InputTokensDetails.CachedTokens
usage.PromptTokensDetails.CacheWriteTokens = responsesResponse.Usage.InputTokensDetails.CacheWriteTokens
}
}
// Count actual tool invocations from Output (not tool declarations). // Count actual tool invocations from Output (not tool declarations).
for _, output := range responsesResponse.Output { for _, output := range responsesResponse.Output {
switch output.Type { switch output.Type {
...@@ -69,7 +61,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http ...@@ -69,7 +61,7 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
} }
imageCounter.Commit(info) imageCounter.Commit(info)
return &usage, nil return usage, nil
} }
func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
...@@ -99,19 +91,8 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp ...@@ -99,19 +91,8 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
case "response.completed", "response.done": case "response.completed", "response.done":
if streamResponse.Response != nil { if streamResponse.Response != nil {
if streamResponse.Response.Usage != nil { if streamResponse.Response.Usage != nil {
if streamResponse.Response.Usage.InputTokens != 0 { incomingUsage := relayconvert.NormalizeResponsesUsage(streamResponse.Response.Usage)
usage.PromptTokens = streamResponse.Response.Usage.InputTokens usage = dto.MergeUsageNonZero(usage, incomingUsage)
}
if streamResponse.Response.Usage.OutputTokens != 0 {
usage.CompletionTokens = streamResponse.Response.Usage.OutputTokens
}
if streamResponse.Response.Usage.TotalTokens != 0 {
usage.TotalTokens = streamResponse.Response.Usage.TotalTokens
}
if streamResponse.Response.Usage.InputTokensDetails != nil {
usage.PromptTokensDetails.CachedTokens = streamResponse.Response.Usage.InputTokensDetails.CachedTokens
usage.PromptTokensDetails.CacheWriteTokens = streamResponse.Response.Usage.InputTokensDetails.CacheWriteTokens
}
} }
if !imageCommitted { if !imageCommitted {
if relaycommon.IsNonBillableResponsesStatus(streamResponse.Response.Status) { if relaycommon.IsNonBillableResponsesStatus(streamResponse.Response.Status) {
...@@ -173,6 +154,9 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp ...@@ -173,6 +154,9 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
} }
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
if usage.BillingUsage != nil {
usage.BillingUsage = dto.CloneBillingUsageWithEstimatedCompletion(usage.BillingUsage, usage.CompletionTokens)
}
return usage, nil return usage, nil
} }
...@@ -38,7 +38,7 @@ func OaiChatToResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp ...@@ -38,7 +38,7 @@ func OaiChatToResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
if responseID := helper.GetResponseID(c); responseID != "" { if responseID := helper.GetResponseID(c); responseID != "" {
chatResp.Id = responseID chatResp.Id = responseID
} }
convertResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatOpenAIResponses, &chatResp) convertResult, err := service.ConvertResponse(c, info, types.RelayFormatOpenAIResponses, &chatResp)
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
} }
...@@ -70,8 +70,9 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo ...@@ -70,8 +70,9 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
responseID := helper.GetResponseID(c) responseID := helper.GetResponseID(c)
state, err := relayconvert.NewResponseStreamState(types.RelayFormatOpenAI, types.RelayFormatOpenAIResponses, relayconvert.ResponseStreamOptions{ state, err := relayconvert.NewResponseStreamState(types.RelayFormatOpenAI, types.RelayFormatOpenAIResponses, relayconvert.ResponseStreamOptions{
ID: responseID, ID: responseID,
Model: info.UpstreamModelName, Model: info.UpstreamModelName,
EmitSequenceNumber: true,
}) })
if err != nil { if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
...@@ -84,7 +85,27 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo ...@@ -84,7 +85,27 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError) streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
return false return false
} }
helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data)) if err := helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data)); err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
return false
}
return true
}
failResponsesStream := func(err error) bool {
failureResults, handled := state.FailResponsesStream("server_error", err.Error(), "")
if !handled {
return false
}
for _, result := range failureResults {
event, ok := result.Value.(relayconvert.ChatToResponsesStreamEvent)
if !ok {
streamErr = types.NewOpenAIError(fmt.Errorf("expected OAI responses stream event, got %T", result.Value), types.ErrorCodeBadResponse, http.StatusInternalServerError)
return true
}
if !sendEvent(event) {
return true
}
}
return true return true
} }
...@@ -97,6 +118,10 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo ...@@ -97,6 +118,10 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
var errorResp dto.OpenAITextResponse var errorResp dto.OpenAITextResponse
if err := common.UnmarshalJsonStr(data, &errorResp); err == nil { if err := common.UnmarshalJsonStr(data, &errorResp); err == nil {
if oaiError := errorResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" { if oaiError := errorResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
if failResponsesStream(fmt.Errorf("%s", oaiError.Message)) {
sr.Stop(streamErr)
return
}
streamErr = types.WithOpenAIError(*oaiError, resp.StatusCode) streamErr = types.WithOpenAIError(*oaiError, resp.StatusCode)
sr.Stop(streamErr) sr.Stop(streamErr)
return return
...@@ -106,12 +131,21 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo ...@@ -106,12 +131,21 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
var chunk dto.ChatCompletionsStreamResponse var chunk dto.ChatCompletionsStreamResponse
if err := common.UnmarshalJsonStr(data, &chunk); err != nil { if err := common.UnmarshalJsonStr(data, &chunk); err != nil {
logger.LogError(c, "failed to unmarshal chat stream response: "+err.Error()) logger.LogError(c, "failed to unmarshal chat stream response: "+err.Error())
sr.Error(err) if failResponsesStream(err) {
sr.Stop(streamErr)
return
}
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
sr.Stop(streamErr)
return return
} }
results, err := relayconvert.ConvertStreamResponseChunk(c, info, state, &chunk) results, err := service.ConvertStreamResponseChunk(c, info, state, &chunk)
if err != nil { if err != nil {
if failResponsesStream(err) {
sr.Stop(streamErr)
return
}
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
sr.Stop(streamErr) sr.Stop(streamErr)
return return
...@@ -140,8 +174,11 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo ...@@ -140,8 +174,11 @@ func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
state.SetUsage(usage) state.SetUsage(usage)
} }
finalResults, err := relayconvert.FinalizeStreamResponse(c, info, state) finalResults, err := service.FinalizeStreamResponse(c, info, state)
if err != nil { if err != nil {
if failResponsesStream(err) {
return usage, streamErr
}
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
} }
for _, result := range finalResults { for _, result := range finalResults {
......
package sub2api package sub2api
import ( import (
"encoding/json"
"testing" "testing"
"github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/constant"
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
relayconstant "github.com/QuantumNous/new-api/relay/constant" relayconstant "github.com/QuantumNous/new-api/relay/constant"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
...@@ -44,3 +46,40 @@ func TestAdaptorInheritsNewAPIResponsesCompactSupport(t *testing.T) { ...@@ -44,3 +46,40 @@ func TestAdaptorInheritsNewAPIResponsesCompactSupport(t *testing.T) {
assert.Equal(t, "sub2api", adaptor.GetChannelName()) assert.Equal(t, "sub2api", adaptor.GetChannelName())
assert.Empty(t, adaptor.GetModelList()) assert.Empty(t, adaptor.GetModelList())
} }
func TestConvertClaudeRequestPreservesAdaptiveThinkingForCompatibleModel(t *testing.T) {
adaptor := &Adaptor{}
maxTokens := uint(8192)
temperature := 0.2
topP := 0.99
request := &dto.ClaudeRequest{
Model: "gpt-5.6-sol",
MaxTokens: &maxTokens,
Temperature: &temperature,
TopP: &topP,
Thinking: &dto.Thinking{Type: "adaptive", Display: "summarized"},
OutputConfig: json.RawMessage(`{"effort":"xhigh","provider_option":true}`),
Messages: []dto.ClaudeMessage{
{Role: "user", Content: "hello"},
},
}
info := &relaycommon.RelayInfo{
OriginModelName: "gpt-5.6-sol",
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeSub2API,
},
}
converted, err := adaptor.ConvertClaudeRequest(nil, info, request)
require.NoError(t, err)
assert.Same(t, request, converted)
require.NotNil(t, request.Thinking)
assert.Equal(t, "adaptive", request.Thinking.Type)
assert.Equal(t, "summarized", request.Thinking.Display)
assert.JSONEq(t, `{"effort":"xhigh","provider_option":true}`, string(request.OutputConfig))
assert.Same(t, &temperature, request.Temperature)
assert.Same(t, &topP, request.TopP)
assert.Equal(t, "xhigh", info.ReasoningEffort)
assert.Equal(t, "gpt-5.6-sol", info.UpstreamModelName)
}
...@@ -18,7 +18,6 @@ import ( ...@@ -18,7 +18,6 @@ import (
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/model_setting" "github.com/QuantumNous/new-api/setting/model_setting"
"github.com/QuantumNous/new-api/setting/reasoning"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/samber/lo" "github.com/samber/lo"
...@@ -56,15 +55,16 @@ type Adaptor struct { ...@@ -56,15 +55,16 @@ type Adaptor struct {
} }
func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.GeminiChatRequest) (any, error) { func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.GeminiChatRequest) (any, error) {
// Vertex AI does not support functionResponse.id; keep it stripped here for consistency. // Vertex AI's generateContent schema does not expose the Gemini API's
// function-call identity fields. Strip both sides at this provider boundary.
if model_setting.GetGeminiSettings().RemoveFunctionResponseIdEnabled { if model_setting.GetGeminiSettings().RemoveFunctionResponseIdEnabled {
removeFunctionResponseID(request) removeFunctionCallIDs(request)
} }
geminiAdaptor := gemini.Adaptor{} geminiAdaptor := gemini.Adaptor{}
return geminiAdaptor.ConvertGeminiRequest(c, info, request) return geminiAdaptor.ConvertGeminiRequest(c, info, request)
} }
func removeFunctionResponseID(request *dto.GeminiChatRequest) { func removeFunctionCallIDs(request *dto.GeminiChatRequest) {
if request == nil { if request == nil {
return return
} }
...@@ -76,10 +76,10 @@ func removeFunctionResponseID(request *dto.GeminiChatRequest) { ...@@ -76,10 +76,10 @@ func removeFunctionResponseID(request *dto.GeminiChatRequest) {
} }
for j := range request.Contents[i].Parts { for j := range request.Contents[i].Parts {
part := &request.Contents[i].Parts[j] part := &request.Contents[i].Parts[j]
if part.FunctionResponse == nil { if part.FunctionCall != nil {
continue part.FunctionCall.ID = ""
} }
if len(part.FunctionResponse.ID) > 0 { if part.FunctionResponse != nil && len(part.FunctionResponse.ID) > 0 {
part.FunctionResponse.ID = nil part.FunctionResponse.ID = nil
} }
} }
...@@ -88,12 +88,16 @@ func removeFunctionResponseID(request *dto.GeminiChatRequest) { ...@@ -88,12 +88,16 @@ func removeFunctionResponseID(request *dto.GeminiChatRequest) {
if len(request.Requests) > 0 { if len(request.Requests) > 0 {
for i := range request.Requests { for i := range request.Requests {
removeFunctionResponseID(&request.Requests[i]) removeFunctionCallIDs(&request.Requests[i])
} }
} }
} }
func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) { func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) {
claudeAdaptor := claude.Adaptor{}
if _, err := claudeAdaptor.ConvertClaudeRequest(c, info, request); err != nil {
return nil, err
}
if v, ok := claudeModelMap[info.UpstreamModelName]; ok { if v, ok := claudeModelMap[info.UpstreamModelName]; ok {
c.Set("request_model", v) c.Set("request_model", v)
} else { } else {
...@@ -170,21 +174,6 @@ func (a *Adaptor) getRequestUrl(info *relaycommon.RelayInfo, modelName, suffix s ...@@ -170,21 +174,6 @@ func (a *Adaptor) getRequestUrl(info *relaycommon.RelayInfo, modelName, suffix s
func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) { func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
suffix := "" suffix := ""
if a.RequestMode == RequestModeGemini { if a.RequestMode == RequestModeGemini {
if model_setting.GetGeminiSettings().ThinkingAdapterEnabled &&
!model_setting.ShouldPreserveThinkingSuffix(info.OriginModelName) {
// 新增逻辑:处理 -thinking-<budget> 格式
if strings.Contains(info.UpstreamModelName, "-thinking-") {
parts := strings.Split(info.UpstreamModelName, "-thinking-")
info.UpstreamModelName = parts[0]
} else if strings.HasSuffix(info.UpstreamModelName, "-thinking") { // 旧的适配
info.UpstreamModelName = strings.TrimSuffix(info.UpstreamModelName, "-thinking")
} else if strings.HasSuffix(info.UpstreamModelName, "-nothinking") {
info.UpstreamModelName = strings.TrimSuffix(info.UpstreamModelName, "-nothinking")
} else if baseModel, level, ok := reasoning.TrimEffortSuffix(info.UpstreamModelName); ok && level != "" {
info.UpstreamModelName = baseModel
}
}
if info.IsStream { if info.IsStream {
suffix = "streamGenerateContent?alt=sse" suffix = "streamGenerateContent?alt=sse"
} else { } else {
...@@ -310,6 +299,9 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn ...@@ -310,6 +299,9 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn
if !ok { if !ok {
return nil, fmt.Errorf("expected Gemini generateContent request, got %T", result.Value) return nil, fmt.Errorf("expected Gemini generateContent request, got %T", result.Value)
} }
if model_setting.GetGeminiSettings().RemoveFunctionResponseIdEnabled {
removeFunctionCallIDs(geminiRequest)
}
c.Set("request_model", request.Model) c.Set("request_model", request.Model)
return geminiRequest, nil return geminiRequest, nil
} else if a.RequestMode == RequestModeOpenSource { } else if a.RequestMode == RequestModeOpenSource {
......
...@@ -28,7 +28,8 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt ...@@ -28,7 +28,8 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt
} }
func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, req *dto.ClaudeRequest) (any, error) { func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, req *dto.ClaudeRequest) (any, error) {
return req, nil claudeAdaptor := claude.Adaptor{}
return claudeAdaptor.ConvertClaudeRequest(c, info, req)
} }
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
......
...@@ -70,30 +70,35 @@ func applySystemPromptIfNeeded(c *gin.Context, info *relaycommon.RelayInfo, requ ...@@ -70,30 +70,35 @@ func applySystemPromptIfNeeded(c *gin.Context, info *relaycommon.RelayInfo, requ
} }
} }
func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, adaptor channel.Adaptor, request *dto.GeneralOpenAIRequest) (*dto.Usage, *types.NewAPIError) { func textRequestViaResponses(c *gin.Context, info *relaycommon.RelayInfo, adaptor channel.Adaptor, request any) (*dto.Usage, *types.NewAPIError) {
chatJSON, err := common.Marshal(request) paramOverrideApplied := false
if err != nil { if chatRequest, ok := request.(*dto.GeneralOpenAIRequest); ok {
return nil, types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) chatJSON, err := common.Marshal(chatRequest)
} if err != nil {
return nil, types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
chatJSON, err = relaycommon.RemoveDisabledFields(chatJSON, info.ChannelOtherSettings, info.ChannelSetting.PassThroughBodyEnabled) }
if err != nil {
return nil, types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
}
if len(info.ParamOverride) > 0 { chatJSON, err = relaycommon.RemoveDisabledFields(chatJSON, info.ChannelOtherSettings, info.ChannelSetting.PassThroughBodyEnabled)
chatJSON, err = relaycommon.ApplyParamOverrideWithRelayInfo(chatJSON, info)
if err != nil { if err != nil {
return nil, newAPIErrorFromParamOverride(err) return nil, types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
}
if len(info.ParamOverride) > 0 {
chatJSON, err = relaycommon.ApplyParamOverrideWithRelayInfo(chatJSON, info)
if err != nil {
return nil, newAPIErrorFromParamOverride(err)
}
paramOverrideApplied = true
} }
}
var overriddenChatReq dto.GeneralOpenAIRequest var overriddenChatReq dto.GeneralOpenAIRequest
if err := common.Unmarshal(chatJSON, &overriddenChatReq); err != nil { if err := common.Unmarshal(chatJSON, &overriddenChatReq); err != nil {
return nil, types.NewError(err, types.ErrorCodeChannelParamOverrideInvalid, types.ErrOptionWithSkipRetry()) return nil, types.NewError(err, types.ErrorCodeChannelParamOverrideInvalid, types.ErrOptionWithSkipRetry())
}
request = &overriddenChatReq
} }
result, err := service.ConvertRequestVia(c, info, &overriddenChatReq, types.RelayFormatOpenAI, types.RelayFormatOpenAIResponses) result, err := service.ConvertRequest(c, info, types.RelayFormatOpenAIResponses, request)
if err != nil { if err != nil {
return nil, types.NewErrorWithStatusCode(err, types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry()) return nil, types.NewErrorWithStatusCode(err, types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
} }
...@@ -101,7 +106,10 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad ...@@ -101,7 +106,10 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad
if !ok { if !ok {
return nil, types.NewError(fmt.Errorf("expected OpenAI responses request, got %T", result.Value), types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return nil, types.NewError(fmt.Errorf("expected OpenAI responses request, got %T", result.Value), types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
return relayResponsesRequest(c, info, adaptor, responsesReq, paramOverrideApplied)
}
func relayResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, adaptor channel.Adaptor, responsesReq *dto.OpenAIResponsesRequest, paramOverrideApplied bool) (*dto.Usage, *types.NewAPIError) {
savedRelayMode := info.RelayMode savedRelayMode := info.RelayMode
savedRequestURLPath := info.RequestURLPath savedRequestURLPath := info.RequestURLPath
defer func() { defer func() {
...@@ -114,7 +122,7 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad ...@@ -114,7 +122,7 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad
convertedRequest, err := adaptor.ConvertOpenAIResponsesRequest(c, info, *responsesReq) convertedRequest, err := adaptor.ConvertOpenAIResponsesRequest(c, info, *responsesReq)
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return nil, newConvertRequestFailedError(c, info, err)
} }
relaycommon.AppendRequestConversionFromRequest(info, convertedRequest) relaycommon.AppendRequestConversionFromRequest(info, convertedRequest)
...@@ -127,6 +135,12 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad ...@@ -127,6 +135,12 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad
if err != nil { if err != nil {
return nil, types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return nil, types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
} }
if !paramOverrideApplied && len(info.ParamOverride) > 0 {
jsonData, err = relaycommon.ApplyParamOverrideWithRelayInfo(jsonData, info)
if err != nil {
return nil, newAPIErrorFromParamOverride(err)
}
}
body, closer, err := relaycommon.NewOutboundJSONBody(jsonData) body, closer, err := relaycommon.NewOutboundJSONBody(jsonData)
if err != nil { if err != nil {
......
package relay package relay
import ( import (
"io"
"math" "math"
"net/http"
"net/http/httptest"
"testing" "testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
openaichannel "github.com/QuantumNous/new-api/relay/channel/openai"
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/types" relayconstant "github.com/QuantumNous/new-api/relay/constant"
"github.com/QuantumNous/new-api/relaykit/dto"
relaytypes "github.com/QuantumNous/new-api/relaykit/types"
hosttypes "github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
...@@ -31,7 +41,7 @@ func TestIsResponsesEventStreamContentType(t *testing.T) { ...@@ -31,7 +41,7 @@ func TestIsResponsesEventStreamContentType(t *testing.T) {
func TestRecalcQuotaFromRatiosIgnoresInvalidMultipliers(t *testing.T) { func TestRecalcQuotaFromRatiosIgnoresInvalidMultipliers(t *testing.T) {
info := &relaycommon.RelayInfo{ info := &relaycommon.RelayInfo{
PriceData: types.PriceData{ PriceData: hosttypes.PriceData{
Quota: 100, Quota: 100,
}, },
} }
...@@ -52,7 +62,7 @@ func TestRecalcQuotaFromRatiosIgnoresInvalidMultipliers(t *testing.T) { ...@@ -52,7 +62,7 @@ func TestRecalcQuotaFromRatiosIgnoresInvalidMultipliers(t *testing.T) {
func TestRecalcQuotaFromRatiosRejectsAllInvalidAdjustedRatios(t *testing.T) { func TestRecalcQuotaFromRatiosRejectsAllInvalidAdjustedRatios(t *testing.T) {
info := &relaycommon.RelayInfo{ info := &relaycommon.RelayInfo{
PriceData: types.PriceData{ PriceData: hosttypes.PriceData{
Quota: 100, Quota: 100,
}, },
} }
...@@ -69,3 +79,77 @@ func TestRecalcQuotaFromRatiosRejectsAllInvalidAdjustedRatios(t *testing.T) { ...@@ -69,3 +79,77 @@ func TestRecalcQuotaFromRatiosRejectsAllInvalidAdjustedRatios(t *testing.T) {
assert.Equal(t, 0, quota) assert.Equal(t, 0, quota)
assert.True(t, info.PriceData.HasOtherRatio("duration")) assert.True(t, info.PriceData.HasOtherRatio("duration"))
} }
func TestTextRequestViaResponsesConvertsClaudeDirectly(t *testing.T) {
type capturedRequest struct {
path string
body []byte
}
captured := make(chan capturedRequest, 1)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, err := io.ReadAll(r.Body)
if err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
captured <- capturedRequest{path: r.URL.Path, body: body}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{
"id":"resp_1",
"object":"response",
"status":"completed",
"model":"gpt-5.6-sol",
"output":[{"type":"message","id":"msg_1","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],
"usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}
}`))
}))
defer server.Close()
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
c.Request.Header.Set("Content-Type", "application/json")
info := &relaycommon.RelayInfo{
RelayMode: relayconstant.RelayModeChatCompletions,
RelayFormat: relaytypes.RelayFormatClaude,
OriginModelName: "gpt-5.6-sol",
RequestConversionChain: []relaytypes.RelayFormat{relaytypes.RelayFormatClaude},
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeOpenAI,
ChannelBaseUrl: server.URL,
ApiKey: "test-key",
UpstreamModelName: "gpt-5.6-sol",
},
}
adaptor := &openaichannel.Adaptor{}
adaptor.Init(info)
request := &dto.ClaudeRequest{
Model: "gpt-5.6-sol",
Thinking: &dto.Thinking{Type: "adaptive", Display: "summarized"},
Messages: []dto.ClaudeMessage{{Role: "user", Content: "hello"}},
}
usage, apiErr := textRequestViaResponses(c, info, adaptor, request)
require.Nil(t, apiErr)
require.NotNil(t, usage)
assert.Equal(t, 5, usage.TotalTokens)
assert.Equal(t, []relaytypes.RelayFormat{relaytypes.RelayFormatClaude, relaytypes.RelayFormatOpenAIResponses}, info.RequestConversionChain)
upstream := <-captured
assert.Equal(t, "/v1/responses", upstream.path)
var upstreamBody map[string]any
require.NoError(t, common.Unmarshal(upstream.body, &upstreamBody))
assert.NotContains(t, upstreamBody, "messages")
reasoning, ok := upstreamBody["reasoning"].(map[string]any)
require.True(t, ok)
assert.Equal(t, "high", reasoning["effort"])
assert.Equal(t, "detailed", reasoning["summary"])
var response dto.ClaudeResponse
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
require.Len(t, response.Content, 1)
assert.Equal(t, "ok", response.Content[0].GetText())
}
package relay package relay
import ( import (
"encoding/json"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
...@@ -16,7 +15,6 @@ import ( ...@@ -16,7 +15,6 @@ import (
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/model_setting" "github.com/QuantumNous/new-api/setting/model_setting"
"github.com/QuantumNous/new-api/setting/reasoning"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -40,6 +38,9 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -40,6 +38,9 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
if err = helper.ApplyReasoningModelSuffix(info, request); err != nil {
return newConvertRequestFailedError(c, info, err)
}
adaptor := GetAdaptor(info.ApiType) adaptor := GetAdaptor(info.ApiType)
if adaptor == nil { if adaptor == nil {
...@@ -47,71 +48,6 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -47,71 +48,6 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
} }
adaptor.Init(info) adaptor.Init(info)
if request.MaxTokens == nil || *request.MaxTokens == 0 {
defaultMaxTokens := uint(model_setting.GetClaudeSettings().GetDefaultMaxTokens(request.Model))
request.MaxTokens = &defaultMaxTokens
}
if baseModel, effortLevel, ok := reasoning.TrimEffortSuffix(request.Model); ok && effortLevel != "" &&
(strings.HasPrefix(request.Model, "claude-opus-4-6") ||
strings.HasPrefix(request.Model, "claude-opus-4-7") ||
strings.HasPrefix(request.Model, "claude-opus-4-8")) {
request.Model = baseModel
request.Thinking = &dto.Thinking{
Type: "adaptive",
}
request.OutputConfig = json.RawMessage(fmt.Sprintf(`{"effort":"%s"}`, effortLevel))
if strings.HasPrefix(request.Model, "claude-opus-4-7") ||
strings.HasPrefix(request.Model, "claude-opus-4-8") {
// Opus 4.7/4.8 reject non-default temperature/top_p/top_k with 400
// and defaults display to "omitted"; restore the 4.6 visible summary.
request.Thinking.Display = "summarized"
request.Temperature = nil
request.TopP = nil
request.TopK = nil
} else {
request.Temperature = common.GetPointer[float64](1.0)
}
info.UpstreamModelName = request.Model
} else if model_setting.GetClaudeSettings().ThinkingAdapterEnabled &&
strings.HasSuffix(request.Model, "-thinking") {
if request.Thinking == nil {
baseModel := strings.TrimSuffix(request.Model, "-thinking")
if strings.HasPrefix(baseModel, "claude-opus-4-7") ||
strings.HasPrefix(baseModel, "claude-opus-4-8") {
// Opus 4.7/4.8 reject thinking.type="enabled"; use adaptive at high effort.
request.Thinking = &dto.Thinking{Type: "adaptive", Display: "summarized"}
request.OutputConfig = json.RawMessage(`{"effort":"high"}`)
request.Temperature = nil
request.TopP = nil
request.TopK = nil
} else {
// 因为BudgetTokens 必须大于1024
if request.MaxTokens == nil || *request.MaxTokens < 1280 {
request.MaxTokens = common.GetPointer[uint](1280)
}
// BudgetTokens 为 max_tokens 的 80%
request.Thinking = &dto.Thinking{
Type: "enabled",
BudgetTokens: common.GetPointer[int](int(float64(*request.MaxTokens) * model_setting.GetClaudeSettings().ThinkingAdapterBudgetTokensPercentage)),
}
// TODO: 临时处理
// https://docs.anthropic.com/en/docs/build-with-claude/extended-thinking#important-considerations-when-using-extended-thinking
request.Temperature = common.GetPointer[float64](1.0)
}
}
if !model_setting.ShouldPreserveThinkingSuffix(info.OriginModelName) {
request.Model = strings.TrimSuffix(request.Model, "-thinking")
}
info.UpstreamModelName = request.Model
}
if !model_setting.GetGlobalSettings().PassThroughRequestEnabled && !info.ChannelSetting.PassThroughBodyEnabled {
if effort := request.GetEfforts(); effort != "" {
info.SetReasoningEffort(effort)
}
}
if info.ChannelSetting.SystemPrompt != "" { if info.ChannelSetting.SystemPrompt != "" {
if request.System == nil { if request.System == nil {
request.SetStringSystem(info.ChannelSetting.SystemPrompt) request.SetStringSystem(info.ChannelSetting.SystemPrompt)
...@@ -140,16 +76,7 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -140,16 +76,7 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
if !model_setting.GetGlobalSettings().PassThroughRequestEnabled && if !model_setting.GetGlobalSettings().PassThroughRequestEnabled &&
!info.ChannelSetting.PassThroughBodyEnabled && !info.ChannelSetting.PassThroughBodyEnabled &&
service.ShouldChatCompletionsUseResponsesGlobal(info.ChannelId, info.ChannelType, info.OriginModelName) { service.ShouldChatCompletionsUseResponsesGlobal(info.ChannelId, info.ChannelType, info.OriginModelName) {
result, convErr := service.ConvertRequest(c, info, types.RelayFormatOpenAI, request) usage, newApiErr := textRequestViaResponses(c, info, adaptor, request)
if convErr != nil {
return types.NewError(convErr, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
}
openAIRequest, ok := result.Value.(*dto.GeneralOpenAIRequest)
if !ok {
return types.NewError(fmt.Errorf("expected OpenAI chat completions request, got %T", result.Value), types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
}
usage, newApiErr := chatCompletionsViaResponses(c, info, adaptor, openAIRequest)
if newApiErr != nil { if newApiErr != nil {
return newApiErr return newApiErr
} }
...@@ -168,7 +95,7 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -168,7 +95,7 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
} else { } else {
convertedRequest, err := adaptor.ConvertClaudeRequest(c, info, request) convertedRequest, err := adaptor.ConvertClaudeRequest(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return newConvertRequestFailedError(c, info, err)
} }
relaycommon.AppendRequestConversionFromRequest(info, convertedRequest) relaycommon.AppendRequestConversionFromRequest(info, convertedRequest)
jsonData, err := common.Marshal(convertedRequest) jsonData, err := common.Marshal(convertedRequest)
......
package common
import (
"context"
"fmt"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/relaykit/types"
"github.com/gin-gonic/gin"
)
const maxConversionDiagnostics = 32
type conversionDiagnosticKey struct {
code string
path string
severity types.ConversionDiagnosticSeverity
from types.RelayFormat
to types.RelayFormat
}
// RecordConversionDiagnostics retains conversion losses for the consume log
// and emits one request-correlated warning per distinct diagnostic. The cap
// prevents malformed streams from growing request state without bound.
func (info *RelayInfo) RecordConversionDiagnostics(ctx context.Context, diagnostics []types.ConversionDiagnostic) {
if info == nil || len(diagnostics) == 0 {
return
}
if ginCtx, ok := ctx.(*gin.Context); ok && ginCtx == nil {
ctx = nil
}
if info.conversionDiagnosticKeys == nil {
info.conversionDiagnosticKeys = make(map[conversionDiagnosticKey]struct{})
}
for _, diagnostic := range diagnostics {
key := conversionDiagnosticKey{
code: diagnostic.Code,
path: diagnostic.Path,
severity: diagnostic.Severity,
from: diagnostic.From,
to: diagnostic.To,
}
if _, exists := info.conversionDiagnosticKeys[key]; exists {
continue
}
if len(info.conversionDiagnostics) >= maxConversionDiagnostics {
if !info.conversionDiagnosticsTruncated {
info.conversionDiagnosticsTruncated = true
logger.LogWarn(ctx, fmt.Sprintf("conversion diagnostics truncated after %d distinct entries", maxConversionDiagnostics))
}
continue
}
info.conversionDiagnosticKeys[key] = struct{}{}
info.conversionDiagnostics = append(info.conversionDiagnostics, diagnostic)
logger.LogWarn(ctx, fmt.Sprintf(
"conversion diagnostic: code=%q severity=%q from=%q to=%q path=%q message=%q",
diagnostic.Code, diagnostic.Severity, diagnostic.From, diagnostic.To, diagnostic.Path, diagnostic.Message,
))
}
}
func (info *RelayInfo) ConversionDiagnostics() []types.ConversionDiagnostic {
if info == nil || len(info.conversionDiagnostics) == 0 {
return nil
}
return append([]types.ConversionDiagnostic(nil), info.conversionDiagnostics...)
}
func (info *RelayInfo) ConversionDiagnosticsTruncated() bool {
return info != nil && info.conversionDiagnosticsTruncated
}
...@@ -10,6 +10,7 @@ import ( ...@@ -10,6 +10,7 @@ import (
"strings" "strings"
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
kitreasoning "github.com/QuantumNous/new-api/relaykit/relayconvert/reasoning"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/samber/lo" "github.com/samber/lo"
"github.com/tidwall/gjson" "github.com/tidwall/gjson"
...@@ -224,22 +225,73 @@ func syncReasoningEffortAfterParamOverride(info *RelayInfo, before, after []byte ...@@ -224,22 +225,73 @@ func syncReasoningEffortAfterParamOverride(info *RelayInfo, before, after []byte
} }
func extractReasoningEffortFromJSON(format types.RelayFormat, data []byte) (string, bool) { func extractReasoningEffortFromJSON(format types.RelayFormat, data []byte) (string, bool) {
var paths []string
switch format { switch format {
case types.RelayFormatOpenAI: case types.RelayFormatOpenAI:
paths = []string{"reasoning_effort", "reasoning.effort"} if effort, exists := firstStringValue(data, "reasoning_effort"); exists && effort != "" {
return effort, true
}
if enabled := gjson.GetBytes(data, "reasoning.enabled"); enabled.Exists() {
if enabled.Type != gjson.True && enabled.Type != gjson.False {
return "", true
}
if !enabled.Bool() {
return string(kitreasoning.EffortNone), true
}
if effort, exists := firstStringValue(data, "reasoning.effort"); exists && effort != "" {
return effort, true
}
if budget := gjson.GetBytes(data, "reasoning.max_tokens"); budget.Exists() {
return reasoningEffortFromBudgetValue(budget)
}
return string(kitreasoning.EffortHigh), true
}
if effort, exists := firstStringValue(data, "reasoning.effort"); exists && effort != "" {
return effort, true
}
if budget := gjson.GetBytes(data, "reasoning.max_tokens"); budget.Exists() {
return reasoningEffortFromBudgetValue(budget)
}
return "", false
case types.RelayFormatOpenAIResponses: case types.RelayFormatOpenAIResponses:
paths = []string{"reasoning.effort"} return firstStringValue(data, "reasoning.effort")
case types.RelayFormatClaude: case types.RelayFormatClaude:
paths = []string{"output_config.effort"} if effort, exists := firstStringValue(data, "output_config.effort"); exists && effort != "" {
return effort, true
}
thinkingType, hasThinkingType := firstStringValue(data, "thinking.type")
if thinkingType == "disabled" {
return string(kitreasoning.EffortNone), true
}
if budget := gjson.GetBytes(data, "thinking.budget_tokens"); budget.Exists() {
return reasoningEffortFromBudgetValue(budget)
}
if thinkingType == "enabled" || thinkingType == "adaptive" {
return string(kitreasoning.EffortHigh), true
}
return "", hasThinkingType
case types.RelayFormatGemini: case types.RelayFormatGemini:
paths = []string{ level, hasLevel := firstStringValue(data,
"generationConfig.thinkingConfig.thinkingLevel", "generationConfig.thinkingConfig.thinkingLevel",
"generation_config.thinking_config.thinking_level", "generation_config.thinking_config.thinking_level",
)
if level != "" {
return level, true
}
for _, path := range []string{
"generationConfig.thinkingConfig.thinkingBudget",
"generation_config.thinking_config.thinking_budget",
} {
if budget := gjson.GetBytes(data, path); budget.Exists() {
return reasoningEffortFromBudgetValue(budget)
}
} }
return "", hasLevel
default: default:
return "", false return "", false
} }
}
func firstStringValue(data []byte, paths ...string) (string, bool) {
for _, path := range paths { for _, path := range paths {
value := gjson.GetBytes(data, path) value := gjson.GetBytes(data, path)
if !value.Exists() { if !value.Exists() {
...@@ -253,6 +305,25 @@ func extractReasoningEffortFromJSON(format types.RelayFormat, data []byte) (stri ...@@ -253,6 +305,25 @@ func extractReasoningEffortFromJSON(format types.RelayFormat, data []byte) (stri
return "", false return "", false
} }
func reasoningEffortFromBudgetValue(value gjson.Result) (string, bool) {
if value.Type != gjson.Number {
return "", true
}
budget := value.Float()
switch {
case budget == 0:
return string(kitreasoning.EffortNone), true
case budget < 0:
return string(kitreasoning.EffortHigh), true
case budget <= 1024:
return string(kitreasoning.EffortLow), true
case budget <= 8192:
return string(kitreasoning.EffortMedium), true
default:
return string(kitreasoning.EffortHigh), true
}
}
func shouldEnableParamOverrideAudit(paramOverride map[string]interface{}) bool { func shouldEnableParamOverrideAudit(paramOverride map[string]interface{}) bool {
if common.DebugEnabled { if common.DebugEnabled {
return true return true
......
...@@ -14,6 +14,7 @@ import ( ...@@ -14,6 +14,7 @@ import (
relayconstant "github.com/QuantumNous/new-api/relay/constant" relayconstant "github.com/QuantumNous/new-api/relay/constant"
"github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert/convmeta" "github.com/QuantumNous/new-api/relaykit/relayconvert/convmeta"
kitreasoning "github.com/QuantumNous/new-api/relaykit/relayconvert/reasoning"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/setting/model_setting" "github.com/QuantumNous/new-api/setting/model_setting"
hosttypes "github.com/QuantumNous/new-api/types" hosttypes "github.com/QuantumNous/new-api/types"
...@@ -98,25 +99,34 @@ type RelayInfo struct { ...@@ -98,25 +99,34 @@ type RelayInfo struct {
UsePrice bool UsePrice bool
RelayMode int RelayMode int
OriginModelName string OriginModelName string
RequestURLPath string
RequestHeaders map[string]string // BillingModelName is the pricing identity for this request. It is kept
ShouldIncludeUsage bool // separate from OriginModelName and UpstreamModelName so virtual pricing
DisablePing bool // 是否禁止向下游发送自定义 Ping // aliases never participate in channel selection or upstream routing.
ClientWs *websocket.Conn BillingModelName string
TargetWs *websocket.Conn
InputAudioFormat string RequestURLPath string
OutputAudioFormat string RequestHeaders map[string]string
RealtimeTools []dto.RealTimeTool ShouldIncludeUsage bool
IsFirstRequest bool DisablePing bool // 是否禁止向下游发送自定义 Ping
AudioUsage bool ClientWs *websocket.Conn
ReasoningEffort string TargetWs *websocket.Conn
UserSetting dto.UserSetting InputAudioFormat string
UserEmail string OutputAudioFormat string
UserQuota int RealtimeTools []dto.RealTimeTool
RelayFormat types.RelayFormat IsFirstRequest bool
SendResponseCount int AudioUsage bool
ReceivedResponseCount int ReasoningEffort string
FinalPreConsumedQuota int // 最终预消耗的配额 // ReasoningConversion is the suffix-derived reasoning intent attached
// after model mapping. Converters read it via ReasoningState().
ReasoningConversion *dto.ReasoningConversionState
UserSetting dto.UserSetting
UserEmail string
UserQuota int
RelayFormat types.RelayFormat
SendResponseCount int
ReceivedResponseCount int
FinalPreConsumedQuota int // 最终预消耗的配额
// ForcePreConsume 为 true 时禁用 BillingSession 的信任额度旁路, // ForcePreConsume 为 true 时禁用 BillingSession 的信任额度旁路,
// 强制预扣全额。用于异步任务(视频/音乐生成等),因为请求返回后任务仍在运行, // 强制预扣全额。用于异步任务(视频/音乐生成等),因为请求返回后任务仍在运行,
// 必须在提交前锁定全额。 // 必须在提交前锁定全额。
...@@ -176,6 +186,10 @@ type RelayInfo struct { ...@@ -176,6 +186,10 @@ type RelayInfo struct {
// convOptions caches the converter settings snapshot (see ConvOptions). // convOptions caches the converter settings snapshot (see ConvOptions).
convOptions *convmeta.Options convOptions *convmeta.Options
conversionDiagnostics []types.ConversionDiagnostic
conversionDiagnosticKeys map[conversionDiagnosticKey]struct{}
conversionDiagnosticsTruncated bool
ThinkingContentInfo ThinkingContentInfo
TokenCountMeta TokenCountMeta
*ClaudeConvertInfo *ClaudeConvertInfo
...@@ -186,6 +200,9 @@ type RelayInfo struct { ...@@ -186,6 +200,9 @@ type RelayInfo struct {
} }
func (info *RelayInfo) InitChannelMeta(c *gin.Context) { func (info *RelayInfo) InitChannelMeta(c *gin.Context) {
info.FinalRequestRelayFormat = ""
info.RequestConversionChain = nil
info.InitRequestConversionChain()
channelType := common.GetContextKeyInt(c, constant.ContextKeyChannelType) channelType := common.GetContextKeyInt(c, constant.ContextKeyChannelType)
paramOverride := common.GetContextKeyStringMap(c, constant.ContextKeyChannelParamOverride) paramOverride := common.GetContextKeyStringMap(c, constant.ContextKeyChannelParamOverride)
headerOverride := common.GetContextKeyStringMap(c, constant.ContextKeyChannelHeaderOverride) headerOverride := common.GetContextKeyStringMap(c, constant.ContextKeyChannelHeaderOverride)
...@@ -236,8 +253,10 @@ func (info *RelayInfo) InitChannelMeta(c *gin.Context) { ...@@ -236,8 +253,10 @@ func (info *RelayInfo) InitChannelMeta(c *gin.Context) {
info.convOptions = nil info.convOptions = nil
if model_setting.GetGlobalSettings().PassThroughRequestEnabled || channelMeta.ChannelSetting.PassThroughBodyEnabled { if model_setting.GetGlobalSettings().PassThroughRequestEnabled || channelMeta.ChannelSetting.PassThroughBodyEnabled {
info.ReasoningEffort = "" info.ReasoningEffort = ""
info.ReasoningConversion = nil
} else { } else {
info.ReasoningEffort = reasoningEffortFromRequest(info.Request) info.ReasoningEffort = reasoningEffortFromRequest(info.Request)
info.ReasoningConversion = nil
} }
// reset some fields based on channel meta // reset some fields based on channel meta
...@@ -261,6 +280,9 @@ func (info *RelayInfo) ToString() string { ...@@ -261,6 +280,9 @@ func (info *RelayInfo) ToString() string {
fmt.Fprintf(b, "IsPlayground: %t, ", info.IsPlayground) fmt.Fprintf(b, "IsPlayground: %t, ", info.IsPlayground)
fmt.Fprintf(b, "RequestURLPath: %q, ", info.RequestURLPath) fmt.Fprintf(b, "RequestURLPath: %q, ", info.RequestURLPath)
fmt.Fprintf(b, "OriginModelName: %q, ", info.OriginModelName) fmt.Fprintf(b, "OriginModelName: %q, ", info.OriginModelName)
if info.BillingModelName != "" && info.BillingModelName != info.OriginModelName {
fmt.Fprintf(b, "BillingModelName: %q, ", info.BillingModelName)
}
fmt.Fprintf(b, "EstimatePromptTokens: %d, ", info.estimatePromptTokens) fmt.Fprintf(b, "EstimatePromptTokens: %d, ", info.estimatePromptTokens)
fmt.Fprintf(b, "ShouldIncludeUsage: %t, ", info.ShouldIncludeUsage) fmt.Fprintf(b, "ShouldIncludeUsage: %t, ", info.ShouldIncludeUsage)
fmt.Fprintf(b, "DisablePing: %t, ", info.DisablePing) fmt.Fprintf(b, "DisablePing: %t, ", info.DisablePing)
...@@ -464,7 +486,10 @@ func reasoningEffortFromRequest(request dto.Request) string { ...@@ -464,7 +486,10 @@ func reasoningEffortFromRequest(request dto.Request) string {
} }
case *dto.GeminiChatRequest: case *dto.GeminiChatRequest:
if req != nil && req.GenerationConfig.ThinkingConfig != nil { if req != nil && req.GenerationConfig.ThinkingConfig != nil {
effort = req.GenerationConfig.ThinkingConfig.ThinkingLevel intent, err := kitreasoning.FromGemini(req)
if err == nil {
effort = string(kitreasoning.EffectiveEffort(intent))
}
} }
} }
return strings.TrimSpace(effort) return strings.TrimSpace(effort)
...@@ -739,6 +764,18 @@ func (info *RelayInfo) GetOriginModelName() string { ...@@ -739,6 +764,18 @@ func (info *RelayInfo) GetOriginModelName() string {
return info.OriginModelName return info.OriginModelName
} }
// GetBillingModelName returns the effective pricing identity without changing
// either the client-visible model or the model sent to the selected channel.
func (info *RelayInfo) GetBillingModelName() string {
if info == nil {
return ""
}
if info.BillingModelName != "" {
return info.BillingModelName
}
return info.OriginModelName
}
func (info *RelayInfo) GetUpstreamModelName() string { func (info *RelayInfo) GetUpstreamModelName() string {
if info == nil || info.ChannelMeta == nil { if info == nil || info.ChannelMeta == nil {
return "" return ""
...@@ -780,6 +817,13 @@ func (info *RelayInfo) SetReasoningEffort(effort string) { ...@@ -780,6 +817,13 @@ func (info *RelayInfo) SetReasoningEffort(effort string) {
info.ReasoningEffort = strings.TrimSpace(effort) info.ReasoningEffort = strings.TrimSpace(effort)
} }
func (info *RelayInfo) ReasoningState() *dto.ReasoningConversionState {
if info == nil {
return nil
}
return info.ReasoningConversion
}
func (info *RelayInfo) EnsureClaudeConvertInfo() *convmeta.ClaudeConvertInfo { func (info *RelayInfo) EnsureClaudeConvertInfo() *convmeta.ClaudeConvertInfo {
if info == nil { if info == nil {
return &convmeta.ClaudeConvertInfo{ return &convmeta.ClaudeConvertInfo{
...@@ -832,8 +876,12 @@ func (info *RelayInfo) ConvOptions() *convmeta.Options { ...@@ -832,8 +876,12 @@ func (info *RelayInfo) ConvOptions() *convmeta.Options {
}, },
OpenRouterDialect: info != nil && info.GetChannelType() == constant.ChannelTypeOpenRouter, OpenRouterDialect: info != nil && info.GetChannelType() == constant.ChannelTypeOpenRouter,
PreserveThinkingSuffix: model_setting.ShouldPreserveThinkingSuffix, PreserveThinkingSuffix: model_setting.ShouldPreserveThinkingSuffix,
PreserveEffortTail: model_setting.ShouldPreserveEffortTail,
} }
if info != nil { if info != nil {
if info.ChannelMeta != nil {
options.ToolLossPolicy = types.ConversionLossPolicy(info.ChannelOtherSettings.ToolLossPolicy)
}
info.convOptions = options info.convOptions = options
} }
return options return options
......
...@@ -56,6 +56,7 @@ func TestRelayInfoMetaTypedNilReceiver(t *testing.T) { ...@@ -56,6 +56,7 @@ func TestRelayInfoMetaTypedNilReceiver(t *testing.T) {
assert.Zero(t, meta.GetChannelType()) assert.Zero(t, meta.GetChannelType())
assert.False(t, meta.GetIsStream()) assert.False(t, meta.GetIsStream())
assert.Empty(t, meta.GetReasoningEffort()) assert.Empty(t, meta.GetReasoningEffort())
assert.Nil(t, meta.ReasoningState())
assert.Zero(t, meta.GetEstimatePromptTokens()) assert.Zero(t, meta.GetEstimatePromptTokens())
assert.Zero(t, meta.GetSendResponseCount()) assert.Zero(t, meta.GetSendResponseCount())
...@@ -81,6 +82,7 @@ func TestRelayInfoMetaTypedNilReceiver(t *testing.T) { ...@@ -81,6 +82,7 @@ func TestRelayInfoMetaTypedNilReceiver(t *testing.T) {
assert.NotNil(t, firstOptions.Gemini.SupportsImagine) assert.NotNil(t, firstOptions.Gemini.SupportsImagine)
assert.NotNil(t, firstOptions.Gemini.SafetySetting) assert.NotNil(t, firstOptions.Gemini.SafetySetting)
assert.NotNil(t, firstOptions.PreserveThinkingSuffix) assert.NotNil(t, firstOptions.PreserveThinkingSuffix)
assert.NotNil(t, firstOptions.PreserveEffortTail)
} }
func TestGenRelayInfoCapturesRequestReasoningEffort(t *testing.T) { func TestGenRelayInfoCapturesRequestReasoningEffort(t *testing.T) {
......
...@@ -46,7 +46,7 @@ func (info *RelayInfo) CountBillableToolCall(itemType string, functionName strin ...@@ -46,7 +46,7 @@ func (info *RelayInfo) CountBillableToolCall(itemType string, functionName strin
if _, reserved := reservedBillableToolNames[functionName]; reserved { if _, reserved := reservedBillableToolNames[functionName]; reserved {
return return
} }
if operation_setting.GetToolPriceForModel(functionName, info.OriginModelName) <= 0 { if operation_setting.GetToolPriceForModel(functionName, info.GetBillingModelName()) <= 0 {
return return
} }
info.incrementBillableToolCall(functionName) info.incrementBillableToolCall(functionName)
......
...@@ -43,6 +43,9 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types ...@@ -43,6 +43,9 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
if err = helper.ApplyReasoningModelSuffix(info, request); err != nil {
return newConvertRequestFailedError(c, info, err)
}
includeUsage := true includeUsage := true
// 判断用户是否需要返回使用情况 // 判断用户是否需要返回使用情况
...@@ -76,7 +79,7 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types ...@@ -76,7 +79,7 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types
!info.ChannelSetting.PassThroughBodyEnabled && !info.ChannelSetting.PassThroughBodyEnabled &&
service.ShouldChatCompletionsUseResponsesGlobal(info.ChannelId, info.ChannelType, info.OriginModelName) { service.ShouldChatCompletionsUseResponsesGlobal(info.ChannelId, info.ChannelType, info.OriginModelName) {
applySystemPromptIfNeeded(c, info, request) applySystemPromptIfNeeded(c, info, request)
usage, newApiErr := chatCompletionsViaResponses(c, info, adaptor, request) usage, newApiErr := textRequestViaResponses(c, info, adaptor, request)
if newApiErr != nil { if newApiErr != nil {
return newApiErr return newApiErr
} }
...@@ -108,7 +111,7 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types ...@@ -108,7 +111,7 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types
} else { } else {
convertedRequest, err := adaptor.ConvertOpenAIRequest(c, info, request) convertedRequest, err := adaptor.ConvertOpenAIRequest(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return newConvertRequestFailedError(c, info, err)
} }
relaycommon.AppendRequestConversionFromRequest(info, convertedRequest) relaycommon.AppendRequestConversionFromRequest(info, convertedRequest)
......
package relay
import (
"errors"
"net/http"
relaycommon "github.com/QuantumNous/new-api/relay/common"
kitreasoning "github.com/QuantumNous/new-api/relaykit/relayconvert/reasoning"
"github.com/QuantumNous/new-api/relaykit/types"
"github.com/gin-gonic/gin"
)
func newConvertRequestFailedError(c *gin.Context, info *relaycommon.RelayInfo, err error) *types.NewAPIError {
var loss *types.ConversionLossError
if errors.As(err, &loss) {
info.RecordConversionDiagnostics(c, loss.Diagnostics)
return types.NewErrorWithStatusCode(err, types.ErrorCodeConvertRequestFailed, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
}
if kitreasoning.IsClientError(err) {
return types.NewErrorWithStatusCode(err, types.ErrorCodeConvertRequestFailed, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
}
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry())
}
package relay
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/QuantumNous/new-api/common"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestOptInSafeToolLossRejectedAsBadRequestWithAdminDiagnostics(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
info := &relaycommon.RelayInfo{
OriginModelName: "gpt-4o",
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "gpt-4o",
ChannelOtherSettings: dto.ChannelOtherSettings{
ToolLossPolicy: string(types.ConversionLossPolicySafe),
},
},
}
tools, err := common.Marshal([]map[string]any{{"codeExecution": map[string]any{}}})
require.NoError(t, err)
req := &dto.GeminiChatRequest{
Contents: []dto.GeminiChatContent{
{Role: "user", Parts: []dto.GeminiPart{{Text: "run this"}}},
},
Tools: tools,
}
result, convErr := service.ConvertRequest(c, info, types.RelayFormatOpenAI, req)
require.Error(t, convErr)
var loss *types.ConversionLossError
require.ErrorAs(t, convErr, &loss)
require.NotEmpty(t, loss.Diagnostics)
require.NotNil(t, result)
apiErr := newConvertRequestFailedError(c, info, convErr)
require.NotNil(t, apiErr)
assert.Equal(t, http.StatusBadRequest, apiErr.StatusCode)
assert.Equal(t, types.ErrorCodeConvertRequestFailed, apiErr.GetErrorCode())
assert.True(t, types.IsSkipRetryError(apiErr))
diagnostics := info.ConversionDiagnostics()
require.NotEmpty(t, diagnostics)
assert.True(t, hasHostDiagnosticCode(diagnostics, "unsupported_hosted_tool"))
other := service.GenerateTextOtherInfo(c, info, 1, 1, 1, 0, 0, 0, 1)
adminInfo, ok := other["admin_info"].(map[string]interface{})
require.True(t, ok)
require.Contains(t, adminInfo, "conversion_diagnostics")
}
func hasHostDiagnosticCode(diagnostics []types.ConversionDiagnostic, code string) bool {
for _, diagnostic := range diagnostics {
if diagnostic.Code == code {
return true
}
}
return false
}
...@@ -12,7 +12,6 @@ import ( ...@@ -12,7 +12,6 @@ import (
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper" "github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/model_setting" "github.com/QuantumNous/new-api/setting/model_setting"
...@@ -20,37 +19,6 @@ import ( ...@@ -20,37 +19,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func isNoThinkingRequest(req *dto.GeminiChatRequest) bool {
if req.GenerationConfig.ThinkingConfig != nil && req.GenerationConfig.ThinkingConfig.ThinkingBudget != nil {
configBudget := req.GenerationConfig.ThinkingConfig.ThinkingBudget
if configBudget != nil && *configBudget == 0 {
// 如果思考预算为 0,则认为是非思考请求
return true
}
}
return false
}
func trimModelThinking(modelName string) string {
// 去除模型名称中的 -nothinking 后缀
if strings.HasSuffix(modelName, "-nothinking") {
return strings.TrimSuffix(modelName, "-nothinking")
}
// 去除模型名称中的 -thinking 后缀
if strings.HasSuffix(modelName, "-thinking") {
return strings.TrimSuffix(modelName, "-thinking")
}
// 去除模型名称中的 -thinking-number
if strings.Contains(modelName, "-thinking-") {
parts := strings.Split(modelName, "-thinking-")
if len(parts) > 1 {
return parts[0] + "-thinking"
}
}
return modelName
}
func GeminiHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { func GeminiHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) {
info.InitChannelMeta(c) info.InitChannelMeta(c)
...@@ -69,23 +37,8 @@ func GeminiHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -69,23 +37,8 @@ func GeminiHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
if err = helper.ApplyReasoningModelSuffix(info, request); err != nil {
if model_setting.GetGeminiSettings().ThinkingAdapterEnabled { return newConvertRequestFailedError(c, info, err)
if isNoThinkingRequest(request) {
// check is thinking
if !strings.Contains(info.OriginModelName, "-nothinking") {
// try to get no thinking model price
noThinkingModelName := info.OriginModelName + "-nothinking"
containPrice := helper.HasModelBillingConfig(noThinkingModelName)
if containPrice {
info.OriginModelName = noThinkingModelName
info.UpstreamModelName = noThinkingModelName
}
}
}
if request.GenerationConfig.ThinkingConfig == nil {
relayconvert.ApplyGeminiThinkingConfig(request, info)
}
} }
adaptor := GetAdaptor(info.ApiType) adaptor := GetAdaptor(info.ApiType)
...@@ -146,7 +99,7 @@ func GeminiHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ ...@@ -146,7 +99,7 @@ func GeminiHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ
// 使用 ConvertGeminiRequest 转换请求格式 // 使用 ConvertGeminiRequest 转换请求格式
convertedRequest, err := adaptor.ConvertGeminiRequest(c, info, request) convertedRequest, err := adaptor.ConvertGeminiRequest(c, info, request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return newConvertRequestFailedError(c, info, err)
} }
relaycommon.AppendRequestConversionFromRequest(info, convertedRequest) relaycommon.AppendRequestConversionFromRequest(info, convertedRequest)
jsonData, err := common.Marshal(convertedRequest) jsonData, err := common.Marshal(convertedRequest)
...@@ -245,6 +198,9 @@ func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo) (newAPI ...@@ -245,6 +198,9 @@ func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo) (newAPI
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
if err = helper.ApplyReasoningModelSuffix(info, req); err != nil {
return newConvertRequestFailedError(c, info, err)
}
req.SetModelName("models/" + info.UpstreamModelName) req.SetModelName("models/" + info.UpstreamModelName)
......
...@@ -71,13 +71,14 @@ func HandleGroupRatio(ctx *gin.Context, relayInfo *relaycommon.RelayInfo) hostty ...@@ -71,13 +71,14 @@ func HandleGroupRatio(ctx *gin.Context, relayInfo *relaycommon.RelayInfo) hostty
} }
func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens int, meta *types.TokenCountMeta) (hosttypes.PriceData, error) { func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens int, meta *types.TokenCountMeta) (hosttypes.PriceData, error) {
modelPrice, usePrice := ratio_setting.GetModelPrice(info.OriginModelName, false) billingModelName := info.GetBillingModelName()
modelPrice, usePrice := ratio_setting.GetModelPrice(billingModelName, false)
groupRatioInfo := HandleGroupRatio(c, info) groupRatioInfo := HandleGroupRatio(c, info)
// Check if this model uses tiered_expr billing // Check if this model uses tiered_expr billing
if billing_setting.GetBillingMode(info.OriginModelName) == billing_setting.BillingModeTieredExpr { if billing_setting.GetBillingMode(billingModelName) == billing_setting.BillingModeTieredExpr {
return modelPriceHelperTiered(c, info, promptTokens, meta, groupRatioInfo) return modelPriceHelperTiered(c, info, billingModelName, promptTokens, meta, groupRatioInfo)
} }
var preConsumedQuota int var preConsumedQuota int
...@@ -98,7 +99,7 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens ...@@ -98,7 +99,7 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
} }
var success bool var success bool
var matchName string var matchName string
modelRatio, success, matchName = ratio_setting.GetModelRatio(info.OriginModelName) modelRatio, success, matchName = ratio_setting.GetModelRatio(billingModelName)
if !success { if !success {
acceptUnsetRatio := false acceptUnsetRatio := false
if info.UserSetting.AcceptUnsetRatioModel { if info.UserSetting.AcceptUnsetRatioModel {
...@@ -108,15 +109,15 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens ...@@ -108,15 +109,15 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
return hosttypes.PriceData{}, modelPriceNotConfiguredError(matchName, info.UserId) return hosttypes.PriceData{}, modelPriceNotConfiguredError(matchName, info.UserId)
} }
} }
completionRatio = ratio_setting.GetCompletionRatio(info.OriginModelName) completionRatio = ratio_setting.GetCompletionRatio(billingModelName)
cacheRatio, _ = ratio_setting.GetCacheRatio(info.OriginModelName) cacheRatio, _ = ratio_setting.GetCacheRatio(billingModelName)
cacheCreationRatio, _ = ratio_setting.GetCreateCacheRatio(info.OriginModelName) cacheCreationRatio, _ = ratio_setting.GetCreateCacheRatio(billingModelName)
cacheCreationRatio5m = cacheCreationRatio cacheCreationRatio5m = cacheCreationRatio
// 固定1h和5min缓存写入价格的比例 // 固定1h和5min缓存写入价格的比例
cacheCreationRatio1h = cacheCreationRatio * claudeCacheCreation1hMultiplier cacheCreationRatio1h = cacheCreationRatio * claudeCacheCreation1hMultiplier
imageRatio, _ = ratio_setting.GetImageRatio(info.OriginModelName) imageRatio, _ = ratio_setting.GetImageRatio(billingModelName)
audioRatio = ratio_setting.GetAudioRatio(info.OriginModelName) audioRatio = ratio_setting.GetAudioRatio(billingModelName)
audioCompletionRatio = ratio_setting.GetAudioCompletionRatio(info.OriginModelName) audioCompletionRatio = ratio_setting.GetAudioCompletionRatio(billingModelName)
ratio := modelRatio * groupRatioInfo.GroupRatio ratio := modelRatio * groupRatioInfo.GroupRatio
quota, err := common.QuotaFromFloatStrict(float64(preConsumedTokens) * ratio) quota, err := common.QuotaFromFloatStrict(float64(preConsumedTokens) * ratio)
if err != nil { if err != nil {
...@@ -266,10 +267,10 @@ func HasModelBillingConfig(modelName string) bool { ...@@ -266,10 +267,10 @@ func HasModelBillingConfig(modelName string) bool {
return ok && strings.TrimSpace(expr) != "" return ok && strings.TrimSpace(expr) != ""
} }
func modelPriceHelperTiered(c *gin.Context, info *relaycommon.RelayInfo, promptTokens int, meta *types.TokenCountMeta, groupRatioInfo hosttypes.GroupRatioInfo) (hosttypes.PriceData, error) { func modelPriceHelperTiered(c *gin.Context, info *relaycommon.RelayInfo, billingModelName string, promptTokens int, meta *types.TokenCountMeta, groupRatioInfo hosttypes.GroupRatioInfo) (hosttypes.PriceData, error) {
exprStr, ok := billing_setting.GetBillingExpr(info.OriginModelName) exprStr, ok := billing_setting.GetBillingExpr(billingModelName)
if !ok { if !ok {
return hosttypes.PriceData{}, fmt.Errorf("model %s is configured as tiered_expr but has no billing expression", info.OriginModelName) return hosttypes.PriceData{}, fmt.Errorf("model %s is configured as tiered_expr but has no billing expression", billingModelName)
} }
estimatedCompletionTokens := meta.MaxTokens estimatedCompletionTokens := meta.MaxTokens
...@@ -288,7 +289,7 @@ func modelPriceHelperTiered(c *gin.Context, info *relaycommon.RelayInfo, promptT ...@@ -288,7 +289,7 @@ func modelPriceHelperTiered(c *gin.Context, info *relaycommon.RelayInfo, promptT
Len: float64(promptTokens), Len: float64(promptTokens),
}, requestInput) }, requestInput)
if err != nil { if err != nil {
return hosttypes.PriceData{}, fmt.Errorf("model %s tiered expr run failed: %w", info.OriginModelName, err) return hosttypes.PriceData{}, fmt.Errorf("model %s tiered expr run failed: %w", billingModelName, err)
} }
// Expression coefficients are $/1M tokens prices; convert to quota the same way per-call billing does. // Expression coefficients are $/1M tokens prices; convert to quota the same way per-call billing does.
...@@ -309,7 +310,7 @@ func modelPriceHelperTiered(c *gin.Context, info *relaycommon.RelayInfo, promptT ...@@ -309,7 +310,7 @@ func modelPriceHelperTiered(c *gin.Context, info *relaycommon.RelayInfo, promptT
exprHash := billingexpr.ExprHashString(exprStr) exprHash := billingexpr.ExprHashString(exprStr)
snapshot := &billingexpr.BillingSnapshot{ snapshot := &billingexpr.BillingSnapshot{
BillingMode: billing_setting.BillingModeTieredExpr, BillingMode: billing_setting.BillingModeTieredExpr,
ModelName: info.OriginModelName, ModelName: billingModelName,
ExprString: exprStr, ExprString: exprStr,
ExprHash: exprHash, ExprHash: exprHash,
GroupRatio: groupRatioInfo.GroupRatio, GroupRatio: groupRatioInfo.GroupRatio,
...@@ -330,7 +331,7 @@ func modelPriceHelperTiered(c *gin.Context, info *relaycommon.RelayInfo, promptT ...@@ -330,7 +331,7 @@ func modelPriceHelperTiered(c *gin.Context, info *relaycommon.RelayInfo, promptT
QuotaToPreConsume: preConsumedQuota, QuotaToPreConsume: preConsumedQuota,
} }
logger.LogDebug(c, "model_price_helper_tiered result: model=%s preConsume=%d quotaBeforeGroup=%.2f groupRatio=%.2f tier=%s", info.OriginModelName, preConsumedQuota, quotaBeforeGroup, groupRatioInfo.GroupRatio, trace.MatchedTier) logger.LogDebug(c, "model_price_helper_tiered result: model=%s preConsume=%d quotaBeforeGroup=%.2f groupRatio=%.2f tier=%s", billingModelName, preConsumedQuota, quotaBeforeGroup, groupRatioInfo.GroupRatio, trace.MatchedTier)
info.PriceData = priceData info.PriceData = priceData
return priceData, nil return priceData, nil
......
...@@ -8,11 +8,15 @@ import ( ...@@ -8,11 +8,15 @@ import (
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/pkg/billingexpr" "github.com/QuantumNous/new-api/pkg/billingexpr"
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
"github.com/QuantumNous/new-api/setting/billing_setting" "github.com/QuantumNous/new-api/setting/billing_setting"
"github.com/QuantumNous/new-api/setting/config" "github.com/QuantumNous/new-api/setting/config"
"github.com/QuantumNous/new-api/setting/model_setting"
"github.com/QuantumNous/new-api/setting/operation_setting"
"github.com/QuantumNous/new-api/setting/ratio_setting" "github.com/QuantumNous/new-api/setting/ratio_setting"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
...@@ -272,3 +276,97 @@ func TestModelPriceHelperRequestBillingRatiosOnlyApplyToFixedPrice(t *testing.T) ...@@ -272,3 +276,97 @@ func TestModelPriceHelperRequestBillingRatiosOnlyApplyToFixedPrice(t *testing.T)
require.Equal(t, common.QuotaClampOverflow, clamp.Kind) require.Equal(t, common.QuotaClampOverflow, clamp.Kind)
require.Nil(t, info.Billing) require.Nil(t, info.Billing)
} }
// Pricing at controller/relay.go runs before ApplyReasoningModelSuffix.
// Identity is GetBillingModelName() → OriginModelName (the suffixed client
// name), matching main's info.OriginModelName lookup. Wildcard entries such
// as gemini-2.5-flash-thinking-* depend on that unstripped origin form.
func TestModelPriceHelperUsesSuffixedOriginLikeMain(t *testing.T) {
gin.SetMode(gin.TestMode)
savedRatios := ratio_setting.ModelRatio2JSONString()
t.Cleanup(func() {
require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(savedRatios))
})
ratios := ratio_setting.GetModelRatioCopy()
ratios["gemini-2.5-flash"] = 0.15
ratios["gemini-2.5-flash-thinking-*"] = 0.075
ratioJSON, err := common.Marshal(ratios)
require.NoError(t, err)
require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(string(ratioJSON)))
oldSelfUse := operation_setting.SelfUseModeEnabled
operation_setting.SelfUseModeEnabled = true
t.Cleanup(func() { operation_setting.SelfUseModeEnabled = oldSelfUse })
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set("group", "default")
suffixed := &relaycommon.RelayInfo{
OriginModelName: "gemini-2.5-flash-thinking-8192",
UserGroup: "default",
UsingGroup: "default",
}
suffixedPrice, err := ModelPriceHelper(ctx, suffixed, 1000, &types.TokenCountMeta{})
require.NoError(t, err)
assert.Empty(t, suffixed.BillingModelName)
assert.Equal(t, "gemini-2.5-flash-thinking-8192", suffixed.GetBillingModelName())
assert.Equal(t, 0.075, suffixedPrice.ModelRatio)
base := &relaycommon.RelayInfo{
OriginModelName: "gemini-2.5-flash",
UserGroup: "default",
UsingGroup: "default",
}
basePrice, err := ModelPriceHelper(ctx, base, 1000, &types.TokenCountMeta{})
require.NoError(t, err)
assert.Empty(t, base.BillingModelName)
assert.Equal(t, "gemini-2.5-flash", base.GetBillingModelName())
assert.Equal(t, 0.15, basePrice.ModelRatio)
}
func TestModelPriceHelperNativeGeminiNoThinkingDoesNotAliasBillingModel(t *testing.T) {
gin.SetMode(gin.TestMode)
savedRatios := ratio_setting.ModelRatio2JSONString()
t.Cleanup(func() {
require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(savedRatios))
})
ratios := ratio_setting.GetModelRatioCopy()
ratios["gemini-3-pro"] = 1.25
ratioJSON, err := common.Marshal(ratios)
require.NoError(t, err)
require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(string(ratioJSON)))
oldSelfUse := operation_setting.SelfUseModeEnabled
operation_setting.SelfUseModeEnabled = true
t.Cleanup(func() { operation_setting.SelfUseModeEnabled = oldSelfUse })
geminiSettings := model_setting.GetGeminiSettings()
oldThinking := geminiSettings.ThinkingAdapterEnabled
geminiSettings.ThinkingAdapterEnabled = true
t.Cleanup(func() { geminiSettings.ThinkingAdapterEnabled = oldThinking })
budget := 0
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set("group", "default")
info := &relaycommon.RelayInfo{
OriginModelName: "gemini-3-pro",
UserGroup: "default",
UsingGroup: "default",
Request: &dto.GeminiChatRequest{
GenerationConfig: dto.GeminiChatGenerationConfig{
ThinkingConfig: &dto.GeminiThinkingConfig{
ThinkingBudget: &budget,
},
},
},
}
priceData, err := ModelPriceHelper(ctx, info, 1000, &types.TokenCountMeta{})
require.NoError(t, err)
assert.Empty(t, info.BillingModelName)
assert.Equal(t, "gemini-3-pro", info.GetBillingModelName())
assert.Equal(t, 1.25, priceData.ModelRatio)
assert.NotEqual(t, 37.5, priceData.ModelRatio)
}
package helper
import (
"strings"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/relaykit/relayconvert/convmeta"
"github.com/QuantumNous/new-api/relaykit/relayconvert/reasoning"
"github.com/QuantumNous/new-api/setting/model_setting"
)
// ApplyReasoningModelSuffix parses host-private reasoning suffixes from the
// origin and mapped model names, attaches the resulting intent to RelayInfo,
// and normalizes UpstreamModelName to the unsuffixed base. Optional outbound
// requests are the DeepCopy the handler will send upstream; they must be
// synced here because info.Request is the original, not that copy. Conflict
// between an explicit request field and a suffix is a client error.
func ApplyReasoningModelSuffix(info *relaycommon.RelayInfo, outbound ...dto.Request) error {
if info == nil {
return nil
}
passThrough := model_setting.GetGlobalSettings().PassThroughRequestEnabled
if info.ChannelMeta != nil && info.ChannelSetting.PassThroughBodyEnabled {
passThrough = true
}
if passThrough {
return nil
}
opts := info.ConvOptions()
origin := info.GetOriginModelName()
upstream := ""
if info.ChannelMeta != nil {
upstream = info.UpstreamModelName
}
if opts.ShouldPreserveThinkingSuffix(origin) || opts.ShouldPreserveThinkingSuffix(upstream) {
return nil
}
originBase, originIntent, originFound, err := parseHostModelSuffix(origin, opts)
if err != nil {
return reasoning.AsClientError(err)
}
upstreamBase, upstreamIntent, upstreamFound, err := parseHostModelSuffix(upstream, opts)
if err != nil {
return reasoning.AsClientError(err)
}
suffix := originIntent
if originFound && upstreamFound {
suffix, err = reasoning.MergeExplicitAndSuffix(originIntent, upstreamIntent, origin)
if err != nil {
return reasoning.AsClientError(err)
}
} else if upstreamFound {
suffix = upstreamIntent
}
explicit, err := explicitIntentFromRequest(info.Request)
if err != nil {
return reasoning.AsClientError(err)
}
conflictModel := upstream
if conflictModel == "" {
conflictModel = origin
}
if _, err = reasoning.MergeExplicitAndSuffix(explicit, suffix, conflictModel); err != nil {
return reasoning.AsClientError(err)
}
if !suffix.IsEmpty() {
info.ReasoningConversion = reasoning.StateFromIntent(suffix)
}
if upstreamFound && info.ChannelMeta != nil {
info.UpstreamModelName = upstreamBase
} else if !info.IsModelMapped && originFound && info.ChannelMeta != nil {
info.UpstreamModelName = originBase
}
// Handlers DeepCopy before this helper; info.Request is the original.
// Sync every outbound copy the caller is about to send upstream.
for _, outbound := range outbound {
if outbound != nil {
outbound.SetModelName(info.UpstreamModelName)
}
}
if info.Request != nil {
info.Request.SetModelName(info.UpstreamModelName)
}
return nil
}
func parseHostModelSuffix(name string, opts *convmeta.Options) (string, reasoning.Intent, bool, error) {
if name == "" {
return name, reasoning.Intent{}, false, nil
}
if strings.HasPrefix(name, "claude-") {
return reasoning.ParseClaudeModelSuffix(name, opts.Claude.ThinkingAdapterEnabled)
}
if strings.HasPrefix(name, "gemini-") {
if !opts.Gemini.ThinkingAdapterEnabled {
return name, reasoning.Intent{}, false, nil
}
return reasoning.ParseGeminiModelSuffix(name, true)
}
// deepseek-v4 effort tails are consumed by ParseDeepSeekV4ThinkingSuffix
// in the DeepSeek adaptor; stripping them here drops THINKING+effort.
if strings.HasPrefix(name, "deepseek-v4-") {
return name, reasoning.Intent{}, false, nil
}
effort, base := reasoning.ParseOpenAIReasoningEffortFromModelSuffix(name, opts.PreserveEffortTail)
if effort != "" {
parsed, err := reasoning.ParseEffort(effort)
if err != nil {
return name, reasoning.Intent{}, false, err
}
mode := reasoning.ModeEnabled
if parsed == reasoning.EffortNone {
mode = reasoning.ModeDisabled
}
return base, reasoning.Intent{Mode: mode, Effort: parsed, Source: reasoning.SourceSuffix}, true, nil
}
// Generic -thinking trim is OpenRouter-only. Volcengine/DeepSeek adaptors
// read the suffix off UpstreamModelName themselves.
if opts != nil && opts.OpenRouterDialect && strings.HasSuffix(name, "-thinking") {
return strings.TrimSuffix(name, "-thinking"), reasoning.Intent{Mode: reasoning.ModeEnabled, Source: reasoning.SourceSuffix}, true, nil
}
return name, reasoning.Intent{}, false, nil
}
func explicitIntentFromRequest(req dto.Request) (reasoning.Intent, error) {
switch r := req.(type) {
case *dto.ClaudeRequest:
return reasoning.FromClaude(r)
case *dto.GeminiChatRequest:
return reasoning.FromGemini(r)
case *dto.GeneralOpenAIRequest:
return reasoning.FromOpenAIChat(r)
case *dto.OpenAIResponsesRequest:
return reasoning.FromOpenAIResponses(r)
default:
return reasoning.Intent{}, nil
}
}
package helper
import (
"net/http/httptest"
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/setting/model_setting"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestApplyReasoningModelSuffixTrimsUpstreamAndAttachesState(t *testing.T) {
info := &relaycommon.RelayInfo{
OriginModelName: "claude-3-7-sonnet-thinking",
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-3-7-sonnet-thinking",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Equal(t, "claude-3-7-sonnet", info.UpstreamModelName)
require.NotNil(t, info.ReasoningConversion)
assert.Equal(t, "enabled", info.ReasoningConversion.Mode)
}
func TestApplyReasoningModelSuffixRetryKeepsEquivalentState(t *testing.T) {
info := &relaycommon.RelayInfo{
OriginModelName: "claude-opus-4-8-high",
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-opus-4-8-high",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
require.NotNil(t, info.ReasoningConversion)
firstMode := info.ReasoningConversion.Mode
firstEffort := info.ReasoningConversion.Effort
info.UpstreamModelName = info.OriginModelName
require.NoError(t, ApplyReasoningModelSuffix(info))
require.NotNil(t, info.ReasoningConversion)
assert.Equal(t, firstMode, info.ReasoningConversion.Mode)
assert.Equal(t, firstEffort, info.ReasoningConversion.Effort)
assert.Equal(t, "claude-opus-4-8", info.UpstreamModelName)
}
func TestApplyReasoningModelSuffixRetryClearsStateWhenNewChannelHasNoSuffix(t *testing.T) {
gin.SetMode(gin.TestMode)
req := &dto.ClaudeRequest{Model: "claude-3-7-sonnet"}
info := &relaycommon.RelayInfo{
OriginModelName: "claude-3-7-sonnet",
Request: req,
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-3-7-sonnet-thinking",
IsModelMapped: true,
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
require.NotNil(t, info.ReasoningState())
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Request = httptest.NewRequest("POST", "/v1/messages", nil)
common.SetContextKey(ctx, constant.ContextKeyOriginalModel, "claude-3-7-sonnet")
common.SetContextKey(ctx, constant.ContextKeyChannelType, constant.ChannelTypeAnthropic)
info.InitChannelMeta(ctx)
assert.Nil(t, info.ReasoningState())
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Nil(t, info.ReasoningState())
}
func TestApplyReasoningModelSuffixPassThroughDoesNotTrim(t *testing.T) {
settings := model_setting.GetGlobalSettings()
original := settings.PassThroughRequestEnabled
t.Cleanup(func() { settings.PassThroughRequestEnabled = original })
settings.PassThroughRequestEnabled = true
info := &relaycommon.RelayInfo{
OriginModelName: "claude-3-7-sonnet-thinking",
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-3-7-sonnet-thinking",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Equal(t, "claude-3-7-sonnet-thinking", info.UpstreamModelName)
assert.Nil(t, info.ReasoningConversion)
}
func TestApplyReasoningModelSuffixBlacklistDoesNotTrim(t *testing.T) {
settings := model_setting.GetGlobalSettings()
original := append([]string(nil), settings.ThinkingModelBlacklist...)
t.Cleanup(func() { settings.ThinkingModelBlacklist = original })
settings.ThinkingModelBlacklist = append(settings.ThinkingModelBlacklist, "claude-3-7-sonnet-thinking")
info := &relaycommon.RelayInfo{
OriginModelName: "claude-3-7-sonnet-thinking",
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-3-7-sonnet-thinking",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Equal(t, "claude-3-7-sonnet-thinking", info.UpstreamModelName)
assert.Nil(t, info.ReasoningConversion)
}
func TestApplyReasoningModelSuffixRejectsExplicitSuffixConflict(t *testing.T) {
info := &relaycommon.RelayInfo{
OriginModelName: "claude-3-7-sonnet-thinking",
Request: &dto.ClaudeRequest{
Model: "claude-3-7-sonnet-thinking",
Thinking: &dto.Thinking{Type: "disabled"},
},
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "claude-3-7-sonnet-thinking",
},
}
err := ApplyReasoningModelSuffix(info)
require.Error(t, err)
}
func TestApplyReasoningModelSuffixGeminiNoThinkingWhenAdapterEnabled(t *testing.T) {
settings := model_setting.GetGeminiSettings()
original := settings.ThinkingAdapterEnabled
t.Cleanup(func() { settings.ThinkingAdapterEnabled = original })
settings.ThinkingAdapterEnabled = true
info := &relaycommon.RelayInfo{
OriginModelName: "gemini-2.5-flash-nothinking",
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "gemini-2.5-flash-nothinking",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Equal(t, "gemini-2.5-flash", info.UpstreamModelName)
require.NotNil(t, info.ReasoningConversion)
assert.Equal(t, "disabled", info.ReasoningConversion.Mode)
assert.Equal(t, "none", info.ReasoningConversion.Effort)
}
func TestApplyReasoningModelSuffixPreservesEffortTailModelID(t *testing.T) {
info := &relaycommon.RelayInfo{
OriginModelName: "qwen-max",
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "qwen-max",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Equal(t, "qwen-max", info.UpstreamModelName)
assert.Nil(t, info.ReasoningConversion)
}
func TestApplyReasoningModelSuffixLeavesDeepSeekV4SuffixForAdaptor(t *testing.T) {
info := &relaycommon.RelayInfo{
OriginModelName: "deepseek-v4-chat-max",
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeDeepSeek,
UpstreamModelName: "deepseek-v4-chat-max",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Equal(t, "deepseek-v4-chat-max", info.UpstreamModelName)
assert.Nil(t, info.ReasoningConversion)
}
func TestApplyReasoningModelSuffixLeavesVolcengineDeepSeekThinkingForAdaptor(t *testing.T) {
info := &relaycommon.RelayInfo{
OriginModelName: "deepseek-r1-thinking",
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeVolcEngine,
UpstreamModelName: "deepseek-r1-thinking",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Equal(t, "deepseek-r1-thinking", info.UpstreamModelName)
assert.Nil(t, info.ReasoningConversion)
}
func TestApplyReasoningModelSuffixStillParsesOpenAIEffortTail(t *testing.T) {
info := &relaycommon.RelayInfo{
OriginModelName: "gpt-5.1-high",
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeOpenAI,
UpstreamModelName: "gpt-5.1-high",
},
}
require.NoError(t, ApplyReasoningModelSuffix(info))
assert.Equal(t, "gpt-5.1", info.UpstreamModelName)
require.NotNil(t, info.ReasoningConversion)
assert.Equal(t, "enabled", info.ReasoningConversion.Mode)
assert.Equal(t, "high", info.ReasoningConversion.Effort)
}
func TestApplyReasoningModelSuffixTrimsOpenRouterThinkingOnly(t *testing.T) {
openRouter := &relaycommon.RelayInfo{
OriginModelName: "some-model-thinking",
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeOpenRouter,
UpstreamModelName: "some-model-thinking",
},
}
require.NoError(t, ApplyReasoningModelSuffix(openRouter))
assert.Equal(t, "some-model", openRouter.UpstreamModelName)
require.NotNil(t, openRouter.ReasoningConversion)
assert.Equal(t, "enabled", openRouter.ReasoningConversion.Mode)
}
...@@ -70,6 +70,9 @@ func ResponsesHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError * ...@@ -70,6 +70,9 @@ func ResponsesHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry())
} }
if err = helper.ApplyReasoningModelSuffix(info, request); err != nil {
return newConvertRequestFailedError(c, info, err)
}
adaptor := GetAdaptor(info.ApiType) adaptor := GetAdaptor(info.ApiType)
if adaptor == nil { if adaptor == nil {
...@@ -86,7 +89,7 @@ func ResponsesHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError * ...@@ -86,7 +89,7 @@ func ResponsesHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *
} else { } else {
convertedRequest, err := adaptor.ConvertOpenAIResponsesRequest(c, info, *request) convertedRequest, err := adaptor.ConvertOpenAIResponsesRequest(c, info, *request)
if err != nil { if err != nil {
return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) return newConvertRequestFailedError(c, info, err)
} }
relaycommon.AppendRequestConversionFromRequest(info, convertedRequest) relaycommon.AppendRequestConversionFromRequest(info, convertedRequest)
jsonData, err := common.Marshal(convertedRequest) jsonData, err := common.Marshal(convertedRequest)
......
...@@ -199,6 +199,9 @@ meta := &convmeta.Values{ ...@@ -199,6 +199,9 @@ meta := &convmeta.Values{
- OpenAI Chat 或 OpenAI Responses 转 Claude 时,Claude 请求必须具有 `max_tokens`。源请求未提供时,需要配置 `Claude.DefaultMaxTokens`,否则转换会返回错误。 - OpenAI Chat 或 OpenAI Responses 转 Claude 时,Claude 请求必须具有 `max_tokens`。源请求未提供时,需要配置 `Claude.DefaultMaxTokens`,否则转换会返回错误。
- RelayKit 不负责选择渠道或映射模型名。调用转换前,应将请求中的 `Model` 设置为目标上游使用的模型名。 - RelayKit 不负责选择渠道或映射模型名。调用转换前,应将请求中的 `Model` 设置为目标上游使用的模型名。
- 自定义 `convmeta.Meta` 的指针实现必须保证所有方法对 nil receiver 安全,完整约束见 `convmeta.Meta` 的接口注释。 - 自定义 `convmeta.Meta` 的指针实现必须保证所有方法对 nil receiver 安全,完整约束见 `convmeta.Meta` 的接口注释。
- 工具损耗策略默认是 `allow`:跨协议转换会成功,损耗以诊断形式返回。`safe` / `strict` 只在请求阶段 opt-in 拒绝;响应和流式转换无论策略如何都不会因损耗失败。
- `ThinkingAdapterEnabled` 只控制是否把已解析的推理意图渲染到 Claude / Gemini 请求上。`-thinking` / `-nothinking` / effort 尾缀等命名约定不再由转换器自动解释。
- 若你的入口仍使用这些模型名后缀,请在调用转换前自行调用 `relayconvert/reasoning``Parse*` 帮助函数,把结果写成 `dto.ReasoningConversionState`,并通过 `convmeta.Meta.ReasoningState()``convmeta.Values.ReasoningConversion`)传入。同时把发给上游的模型名裁成无后缀基础名。
## 多模态内容 ## 多模态内容
......
...@@ -86,6 +86,10 @@ type ChannelOtherSettings struct { ...@@ -86,6 +86,10 @@ type ChannelOtherSettings struct {
UpstreamModelUpdateLastRemovedModels []string `json:"upstream_model_update_last_removed_models,omitempty"` // 上次检测到的可删除模型 UpstreamModelUpdateLastRemovedModels []string `json:"upstream_model_update_last_removed_models,omitempty"` // 上次检测到的可删除模型
UpstreamModelUpdateIgnoredModels []string `json:"upstream_model_update_ignored_models,omitempty"` // 手动忽略的模型 UpstreamModelUpdateIgnoredModels []string `json:"upstream_model_update_ignored_models,omitempty"` // 手动忽略的模型
AdvancedCustom *AdvancedCustomConfig `json:"advanced_custom,omitempty"` AdvancedCustom *AdvancedCustomConfig `json:"advanced_custom,omitempty"`
// ToolLossPolicy is a channel-level opt-in for request-phase conversion
// rejection. Empty follows the default allow policy. Accepted values:
// "", "allow", "safe", "strict".
ToolLossPolicy string `json:"tool_loss_policy,omitempty"`
} }
func (s *ChannelOtherSettings) IsOpenRouterEnterprise() bool { func (s *ChannelOtherSettings) IsOpenRouterEnterprise() bool {
...@@ -95,6 +99,20 @@ func (s *ChannelOtherSettings) IsOpenRouterEnterprise() bool { ...@@ -95,6 +99,20 @@ func (s *ChannelOtherSettings) IsOpenRouterEnterprise() bool {
return *s.OpenRouterEnterprise return *s.OpenRouterEnterprise
} }
// ValidateToolLossPolicy validates the channel-level request-phase tool-loss
// policy. Empty keeps the default allow policy.
func (s *ChannelOtherSettings) ValidateToolLossPolicy() error {
if s == nil {
return nil
}
switch strings.TrimSpace(s.ToolLossPolicy) {
case "", string(types.ConversionLossPolicyAllow), string(types.ConversionLossPolicySafe), string(types.ConversionLossPolicyStrict):
return nil
default:
return fmt.Errorf("invalid tool_loss_policy: %s", s.ToolLossPolicy)
}
}
const ( const (
advancedCustomConverterNone = "none" advancedCustomConverterNone = "none"
advancedCustomConverterClaudeMessagesToOpenAIChat = "anthropic_messages_to_openai_chat_completions" advancedCustomConverterClaudeMessagesToOpenAIChat = "anthropic_messages_to_openai_chat_completions"
......
...@@ -642,3 +642,15 @@ func TestChannelSettingsValidateHTTPTransport(t *testing.T) { ...@@ -642,3 +642,15 @@ func TestChannelSettingsValidateHTTPTransport(t *testing.T) {
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "http2_connection_shards") assert.Contains(t, err.Error(), "http2_connection_shards")
} }
func TestChannelOtherSettingsValidateToolLossPolicy(t *testing.T) {
require.NoError(t, (*ChannelOtherSettings)(nil).ValidateToolLossPolicy())
require.NoError(t, (&ChannelOtherSettings{}).ValidateToolLossPolicy())
require.NoError(t, (&ChannelOtherSettings{ToolLossPolicy: "allow"}).ValidateToolLossPolicy())
require.NoError(t, (&ChannelOtherSettings{ToolLossPolicy: "safe"}).ValidateToolLossPolicy())
require.NoError(t, (&ChannelOtherSettings{ToolLossPolicy: "strict"}).ValidateToolLossPolicy())
err := (&ChannelOtherSettings{ToolLossPolicy: "drop"}).ValidateToolLossPolicy()
require.Error(t, err)
assert.Contains(t, err.Error(), "tool_loss_policy")
}
...@@ -24,10 +24,24 @@ type ClaudeMediaMessage struct { ...@@ -24,10 +24,24 @@ type ClaudeMediaMessage struct {
PartialJson *string `json:"partial_json,omitempty"` PartialJson *string `json:"partial_json,omitempty"`
Role string `json:"role,omitempty"` Role string `json:"role,omitempty"`
Thinking *string `json:"thinking,omitempty"` Thinking *string `json:"thinking,omitempty"`
Data string `json:"data,omitempty"`
Signature string `json:"signature,omitempty"` Signature string `json:"signature,omitempty"`
Delta string `json:"delta,omitempty"` Delta string `json:"delta,omitempty"`
CacheControl json.RawMessage `json:"cache_control,omitempty"` CacheControl json.RawMessage `json:"cache_control,omitempty"`
// tool_calls
// Text blocks and citations_delta events.
Citations json.RawMessage `json:"citations,omitempty"`
Citation json.RawMessage `json:"citation,omitempty"`
// Server-tool and tool-result blocks.
Caller json.RawMessage `json:"caller,omitempty"`
ServerName string `json:"server_name,omitempty"`
IsError *bool `json:"is_error,omitempty"`
// ErrorCode is a relaykit compatibility extension. Claude places provider
// error codes inside nested tool-result error content.
ErrorCode string `json:"error_code,omitempty"`
// Tool-use and tool-result blocks.
Id string `json:"id,omitempty"` Id string `json:"id,omitempty"`
Name string `json:"name,omitempty"` Name string `json:"name,omitempty"`
Input any `json:"input,omitempty"` Input any `json:"input,omitempty"`
...@@ -173,6 +187,7 @@ type Tool struct { ...@@ -173,6 +187,7 @@ type Tool struct {
Name string `json:"name"` Name string `json:"name"`
Description string `json:"description,omitempty"` Description string `json:"description,omitempty"`
InputSchema map[string]interface{} `json:"input_schema"` InputSchema map[string]interface{} `json:"input_schema"`
Strict *bool `json:"strict,omitempty"`
} }
type InputSchema struct { type InputSchema struct {
...@@ -182,10 +197,14 @@ type InputSchema struct { ...@@ -182,10 +197,14 @@ type InputSchema struct {
} }
type ClaudeWebSearchTool struct { type ClaudeWebSearchTool struct {
Type string `json:"type"` Type string `json:"type"`
Name string `json:"name"` Name string `json:"name"`
MaxUses int `json:"max_uses,omitempty"` MaxUses int `json:"max_uses,omitempty"`
UserLocation *ClaudeWebSearchUserLocation `json:"user_location,omitempty"` AllowedDomains []string `json:"allowed_domains,omitempty"`
BlockedDomains []string `json:"blocked_domains,omitempty"`
AllowedCallers []string `json:"allowed_callers,omitempty"`
ResponseInclusion string `json:"response_inclusion,omitempty"`
UserLocation *ClaudeWebSearchUserLocation `json:"user_location,omitempty"`
} }
type ClaudeWebSearchUserLocation struct { type ClaudeWebSearchUserLocation struct {
...@@ -413,7 +432,7 @@ func (c *ClaudeRequest) GetTools() []any { ...@@ -413,7 +432,7 @@ func (c *ClaudeRequest) GetTools() []any {
func (c *ClaudeRequest) GetEfforts() string { func (c *ClaudeRequest) GetEfforts() string {
var OutputConfig OutputConfigForEffort var OutputConfig OutputConfigForEffort
if err := json.Unmarshal(c.OutputConfig, &OutputConfig); err == nil { if err := kitutil.Unmarshal(c.OutputConfig, &OutputConfig); err == nil {
effort := OutputConfig.Effort effort := OutputConfig.Effort
return effort return effort
} }
...@@ -596,5 +615,8 @@ func (u *ClaudeUsage) GetCacheCreationTotalTokens() int { ...@@ -596,5 +615,8 @@ func (u *ClaudeUsage) GetCacheCreationTotalTokens() int {
} }
type ClaudeServerToolUse struct { type ClaudeServerToolUse struct {
WebSearchRequests int `json:"web_search_requests"` WebSearchRequests int `json:"web_search_requests,omitempty"`
WebFetchRequests int `json:"web_fetch_requests,omitempty"`
CodeExecutionRequests int `json:"code_execution_requests,omitempty"`
ToolSearchRequests int `json:"tool_search_requests,omitempty"`
} }
...@@ -48,8 +48,9 @@ type ToolConfig struct { ...@@ -48,8 +48,9 @@ type ToolConfig struct {
} }
type FunctionCallingConfig struct { type FunctionCallingConfig struct {
Mode FunctionCallingConfigMode `json:"mode,omitempty"` Mode FunctionCallingConfigMode `json:"mode,omitempty"`
AllowedFunctionNames []string `json:"allowedFunctionNames,omitempty"` AllowedFunctionNames []string `json:"allowedFunctionNames,omitempty"`
StreamFunctionCallArguments *bool `json:"streamFunctionCallArguments,omitempty"`
} }
type FunctionCallingConfigMode string type FunctionCallingConfigMode string
...@@ -161,8 +162,8 @@ func (r *GeminiChatRequest) SetTools(tools []GeminiChatTool) { ...@@ -161,8 +162,8 @@ func (r *GeminiChatRequest) SetTools(tools []GeminiChatTool) {
} }
type GeminiThinkingConfig struct { type GeminiThinkingConfig struct {
IncludeThoughts bool `json:"includeThoughts,omitempty"` IncludeThoughts *bool `json:"includeThoughts,omitempty"`
ThinkingBudget *int `json:"thinkingBudget,omitempty"` ThinkingBudget *int `json:"thinkingBudget,omitempty"`
// TODO Conflict with thinkingbudget. // TODO Conflict with thinkingbudget.
ThinkingLevel string `json:"thinkingLevel,omitempty"` ThinkingLevel string `json:"thinkingLevel,omitempty"`
} }
...@@ -184,7 +185,7 @@ func (c *GeminiThinkingConfig) UnmarshalJSON(data []byte) error { ...@@ -184,7 +185,7 @@ func (c *GeminiThinkingConfig) UnmarshalJSON(data []byte) error {
*c = GeminiThinkingConfig(aux.Alias) *c = GeminiThinkingConfig(aux.Alias)
if aux.IncludeThoughtsSnake != nil { if aux.IncludeThoughtsSnake != nil {
c.IncludeThoughts = *aux.IncludeThoughtsSnake c.IncludeThoughts = aux.IncludeThoughtsSnake
} }
if aux.ThinkingBudgetSnake != nil { if aux.ThinkingBudgetSnake != nil {
...@@ -239,8 +240,21 @@ func (g *GeminiInlineData) UnmarshalJSON(data []byte) error { ...@@ -239,8 +240,21 @@ func (g *GeminiInlineData) UnmarshalJSON(data []byte) error {
} }
type FunctionCall struct { type FunctionCall struct {
FunctionName string `json:"name"` // ID is optional in the Gemini protocol and identifies the matching function response.
Arguments any `json:"args"` ID string `json:"id,omitempty"`
FunctionName string `json:"name"`
Arguments any `json:"args"`
PartialArgs []GeminiPartialArg `json:"partialArgs,omitempty"`
WillContinue *bool `json:"willContinue,omitempty"`
}
type GeminiPartialArg struct {
JSONPath string `json:"jsonPath"`
NumberValue *float64 `json:"numberValue,omitempty"`
StringValue *string `json:"stringValue,omitempty"`
BoolValue *bool `json:"boolValue,omitempty"`
NullValue json.RawMessage `json:"nullValue,omitempty"`
WillContinue *bool `json:"willContinue,omitempty"`
} }
type GeminiFunctionResponse struct { type GeminiFunctionResponse struct {
...@@ -320,11 +334,16 @@ type GeminiChatSafetySettings struct { ...@@ -320,11 +334,16 @@ type GeminiChatSafetySettings struct {
} }
type GeminiChatTool struct { type GeminiChatTool struct {
GoogleSearch any `json:"googleSearch,omitempty"` GoogleSearch any `json:"googleSearch,omitempty"`
GoogleSearchRetrieval any `json:"googleSearchRetrieval,omitempty"` GoogleSearchRetrieval any `json:"googleSearchRetrieval,omitempty"`
CodeExecution any `json:"codeExecution,omitempty"` GoogleMaps json.RawMessage `json:"googleMaps,omitempty"`
FunctionDeclarations any `json:"functionDeclarations,omitempty"` EnterpriseWebSearch json.RawMessage `json:"enterpriseWebSearch,omitempty"`
URLContext any `json:"urlContext,omitempty"` CodeExecution any `json:"codeExecution,omitempty"`
FunctionDeclarations any `json:"functionDeclarations,omitempty"`
URLContext any `json:"urlContext,omitempty"`
FileSearch json.RawMessage `json:"fileSearch,omitempty"`
ComputerUse json.RawMessage `json:"computerUse,omitempty"`
Retrieval json.RawMessage `json:"retrieval,omitempty"`
} }
type GeminiChatGenerationConfig struct { type GeminiChatGenerationConfig struct {
...@@ -447,7 +466,14 @@ type GeminiChatCandidate struct { ...@@ -447,7 +466,14 @@ type GeminiChatCandidate struct {
} }
type GeminiGroundingMetadata struct { type GeminiGroundingMetadata struct {
WebSearchQueries []string `json:"webSearchQueries,omitempty"` WebSearchQueries []string `json:"webSearchQueries,omitempty"`
RetrievalQueries []string `json:"retrievalQueries,omitempty"`
GroundingChunks json.RawMessage `json:"groundingChunks,omitempty"`
GroundingSupports json.RawMessage `json:"groundingSupports,omitempty"`
SearchEntryPoint json.RawMessage `json:"searchEntryPoint,omitempty"`
RetrievalMetadata json.RawMessage `json:"retrievalMetadata,omitempty"`
SourceFlaggingUris json.RawMessage `json:"sourceFlaggingUris,omitempty"`
GoogleMapsWidgetContextToken string `json:"googleMapsWidgetContextToken,omitempty"`
} }
type GeminiChatSafetyRating struct { type GeminiChatSafetyRating struct {
......
...@@ -81,7 +81,7 @@ type GeneralOpenAIRequest struct { ...@@ -81,7 +81,7 @@ type GeneralOpenAIRequest struct {
ExtraBody json.RawMessage `json:"extra_body,omitempty"` ExtraBody json.RawMessage `json:"extra_body,omitempty"`
//xai //xai
SearchParameters json.RawMessage `json:"search_parameters,omitempty"` SearchParameters json.RawMessage `json:"search_parameters,omitempty"`
// claude // OpenAI Chat web search.
WebSearchOptions *WebSearchOptions `json:"web_search_options,omitempty"` WebSearchOptions *WebSearchOptions `json:"web_search_options,omitempty"`
// OpenRouter Params // OpenRouter Params
Usage json.RawMessage `json:"usage,omitempty"` Usage json.RawMessage `json:"usage,omitempty"`
...@@ -108,6 +108,9 @@ type GeneralOpenAIRequest struct { ...@@ -108,6 +108,9 @@ type GeneralOpenAIRequest struct {
ReasoningSplit json.RawMessage `json:"reasoning_split,omitempty"` ReasoningSplit json.RawMessage `json:"reasoning_split,omitempty"`
// vLLM // vLLM
ThinkingTokenBudget json.RawMessage `json:"thinking_token_budget,omitempty"` ThinkingTokenBudget json.RawMessage `json:"thinking_token_budget,omitempty"`
// Internal conversion state; never serialized to an upstream protocol.
ReasoningConversion *ReasoningConversionState `json:"-"`
} }
func (r GeneralOpenAIRequest) MarshalJSON() ([]byte, error) { func (r GeneralOpenAIRequest) MarshalJSON() ([]byte, error) {
...@@ -266,6 +269,7 @@ type FunctionRequest struct { ...@@ -266,6 +269,7 @@ type FunctionRequest struct {
Name string `json:"name"` Name string `json:"name"`
Parameters any `json:"parameters,omitempty"` Parameters any `json:"parameters,omitempty"`
Arguments string `json:"arguments,omitempty"` Arguments string `json:"arguments,omitempty"`
Strict *bool `json:"strict,omitempty"`
} }
type StreamOptions struct { type StreamOptions struct {
...@@ -311,7 +315,10 @@ type Message struct { ...@@ -311,7 +315,10 @@ type Message struct {
Reasoning *string `json:"reasoning,omitempty"` Reasoning *string `json:"reasoning,omitempty"`
ToolCalls json.RawMessage `json:"tool_calls,omitempty"` ToolCalls json.RawMessage `json:"tool_calls,omitempty"`
ToolCallId string `json:"tool_call_id,omitempty"` ToolCallId string `json:"tool_call_id,omitempty"`
parsedContent []MediaContent // Annotations is an official Chat response field. Keeping it on the shared
// message type also preserves annotations when clients replay assistant output.
Annotations json.RawMessage `json:"annotations,omitempty"`
parsedContent []MediaContent
//parsedStringContent *string //parsedStringContent *string
} }
...@@ -485,14 +492,14 @@ func (m *Message) ParseToolCalls() []ToolCallRequest { ...@@ -485,14 +492,14 @@ func (m *Message) ParseToolCalls() []ToolCallRequest {
return nil return nil
} }
var toolCalls []ToolCallRequest var toolCalls []ToolCallRequest
if err := json.Unmarshal(m.ToolCalls, &toolCalls); err == nil { if err := kitutil.Unmarshal(m.ToolCalls, &toolCalls); err == nil {
return toolCalls return toolCalls
} }
return toolCalls return toolCalls
} }
func (m *Message) SetToolCalls(toolCalls any) { func (m *Message) SetToolCalls(toolCalls any) {
toolCallsJson, _ := json.Marshal(toolCalls) toolCallsJson, _ := kitutil.Marshal(toolCalls)
m.ToolCalls = toolCallsJson m.ToolCalls = toolCallsJson
} }
...@@ -562,6 +569,11 @@ func (m *Message) ParseContent() []MediaContent { ...@@ -562,6 +569,11 @@ func (m *Message) ParseContent() []MediaContent {
return contentList return contentList
} }
if content, ok := m.Content.([]MediaContent); ok {
m.parsedContent = content
return content
}
// 尝试解析为数组 // 尝试解析为数组
//var arrayContent []map[string]interface{} //var arrayContent []map[string]interface{}
...@@ -682,7 +694,7 @@ func (m *Message) ParseContent() []MediaContent { ...@@ -682,7 +694,7 @@ func (m *Message) ParseContent() []MediaContent {
} }
var stringContent string var stringContent string
if err := json.Unmarshal(m.Content, &stringContent); err == nil { if err := kitutil.Unmarshal(m.Content, &stringContent); err == nil {
m.parsedStringContent = &stringContent m.parsedStringContent = &stringContent
return stringContent return stringContent
} }
...@@ -707,14 +719,14 @@ func (m *Message) SetNullContent() { ...@@ -707,14 +719,14 @@ func (m *Message) SetNullContent() {
} }
func (m *Message) SetStringContent(content string) { func (m *Message) SetStringContent(content string) {
jsonContent, _ := json.Marshal(content) jsonContent, _ := kitutil.Marshal(content)
m.Content = jsonContent m.Content = jsonContent
m.parsedStringContent = &content m.parsedStringContent = &content
m.parsedContent = nil m.parsedContent = nil
} }
func (m *Message) SetMediaContent(content []MediaContent) { func (m *Message) SetMediaContent(content []MediaContent) {
jsonContent, _ := json.Marshal(content) jsonContent, _ := kitutil.Marshal(content)
m.Content = jsonContent m.Content = jsonContent
m.parsedContent = nil m.parsedContent = nil
m.parsedStringContent = nil m.parsedStringContent = nil
...@@ -725,7 +737,7 @@ func (m *Message) IsStringContent() bool { ...@@ -725,7 +737,7 @@ func (m *Message) IsStringContent() bool {
return true return true
} }
var stringContent string var stringContent string
if err := json.Unmarshal(m.Content, &stringContent); err == nil { if err := kitutil.Unmarshal(m.Content, &stringContent); err == nil {
m.parsedStringContent = &stringContent m.parsedStringContent = &stringContent
return true return true
} }
...@@ -741,7 +753,7 @@ func (m *Message) ParseContent() []MediaContent { ...@@ -741,7 +753,7 @@ func (m *Message) ParseContent() []MediaContent {
// 先尝试解析为字符串 // 先尝试解析为字符串
var stringContent string var stringContent string
if err := json.Unmarshal(m.Content, &stringContent); err == nil { if err := kitutil.Unmarshal(m.Content, &stringContent); err == nil {
contentList = []MediaContent{{ contentList = []MediaContent{{
Type: ContentTypeText, Type: ContentTypeText,
Text: stringContent, Text: stringContent,
...@@ -752,7 +764,7 @@ func (m *Message) ParseContent() []MediaContent { ...@@ -752,7 +764,7 @@ func (m *Message) ParseContent() []MediaContent {
// 尝试解析为数组 // 尝试解析为数组
var arrayContent []map[string]interface{} var arrayContent []map[string]interface{}
if err := json.Unmarshal(m.Content, &arrayContent); err == nil { if err := kitutil.Unmarshal(m.Content, &arrayContent); err == nil {
for _, contentItem := range arrayContent { for _, contentItem := range arrayContent {
contentType, ok := contentItem["type"].(string) contentType, ok := contentItem["type"].(string)
if !ok { if !ok {
...@@ -907,6 +919,9 @@ type OpenAIResponsesRequest struct { ...@@ -907,6 +919,9 @@ type OpenAIResponsesRequest struct {
ThinkingBudget json.RawMessage `json:"thinking_budget,omitempty"` ThinkingBudget json.RawMessage `json:"thinking_budget,omitempty"`
// perplexity // perplexity
Preset json.RawMessage `json:"preset,omitempty"` Preset json.RawMessage `json:"preset,omitempty"`
// Internal conversion state; never serialized to an upstream protocol.
ReasoningConversion *ReasoningConversionState `json:"-"`
} }
func (r OpenAIResponsesRequest) MarshalJSON() ([]byte, error) { func (r OpenAIResponsesRequest) MarshalJSON() ([]byte, error) {
......
...@@ -3,6 +3,7 @@ package dto ...@@ -3,6 +3,7 @@ package dto
import ( import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"strings"
kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil" kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil"
"github.com/QuantumNous/new-api/relaykit/types" "github.com/QuantumNous/new-api/relaykit/types"
...@@ -91,6 +92,10 @@ type ChatCompletionsStreamResponseChoiceDelta struct { ...@@ -91,6 +92,10 @@ type ChatCompletionsStreamResponseChoiceDelta struct {
Reasoning *string `json:"reasoning,omitempty"` Reasoning *string `json:"reasoning,omitempty"`
Role string `json:"role,omitempty"` Role string `json:"role,omitempty"`
ToolCalls []ToolCallResponse `json:"tool_calls,omitempty"` ToolCalls []ToolCallResponse `json:"tool_calls,omitempty"`
// Annotations is an OpenAI-compatible streaming extension supported by
// providers such as OpenRouter. Relaykit uses it to preserve streaming URL
// citations, including Claude round-trip metadata.
Annotations json.RawMessage `json:"annotations,omitempty"`
} }
func (c *ChatCompletionsStreamResponseChoiceDelta) SetContentString(s string) { func (c *ChatCompletionsStreamResponseChoiceDelta) SetContentString(s string) {
...@@ -325,17 +330,143 @@ type IncompleteDetails struct { ...@@ -325,17 +330,143 @@ type IncompleteDetails struct {
} }
type ResponsesOutput struct { type ResponsesOutput struct {
Type string `json:"type"` Type string `json:"type"`
ID string `json:"id"` ID string `json:"id"`
Status string `json:"status"` Status string `json:"status"`
Role string `json:"role"` Role string `json:"role"`
Content []ResponsesOutputContent `json:"content"` Content []ResponsesOutputContent `json:"content"`
Quality string `json:"quality"` Summary []ResponsesReasoningSummaryPart `json:"summary,omitempty"`
Size string `json:"size"` Quality string `json:"quality"`
Result string `json:"result,omitempty"` Size string `json:"size"`
CallId string `json:"call_id,omitempty"` Result string `json:"result,omitempty"`
Name string `json:"name,omitempty"` CallId string `json:"call_id,omitempty"`
Arguments json.RawMessage `json:"arguments,omitempty"` Name string `json:"name,omitempty"`
Arguments json.RawMessage `json:"arguments,omitempty"`
Action json.RawMessage `json:"action,omitempty"`
Queries json.RawMessage `json:"queries,omitempty"`
Results json.RawMessage `json:"results,omitempty"`
Sources json.RawMessage `json:"sources,omitempty"`
Code json.RawMessage `json:"code,omitempty"`
Outputs json.RawMessage `json:"outputs,omitempty"`
ContainerID string `json:"container_id,omitempty"`
PendingSafetyChecks json.RawMessage `json:"pending_safety_checks,omitempty"`
Caller json.RawMessage `json:"caller,omitempty"`
ServerLabel string `json:"server_label,omitempty"`
Output json.RawMessage `json:"output,omitempty"`
ItemError json.RawMessage `json:"error,omitempty"`
ApprovalRequestID string `json:"approval_request_id,omitempty"`
MCPTools json.RawMessage `json:"tools,omitempty"`
}
// MarshalJSON keeps hosted-tool variants within their protocol-specific
// schemas. ResponsesOutput also represents messages, images, and function
// calls, whose fields must not leak into web_search_call or mcp_call items.
func (r ResponsesOutput) MarshalJSON() ([]byte, error) {
switch r.Type {
case "web_search_call":
return kitutil.Marshal(struct {
Type string `json:"type"`
ID string `json:"id"`
Status string `json:"status,omitempty"`
Action json.RawMessage `json:"action,omitempty"`
}{Type: r.Type, ID: r.ID, Status: r.Status, Action: r.Action})
case "mcp_call":
return kitutil.Marshal(struct {
Type string `json:"type"`
ID string `json:"id"`
Name string `json:"name"`
ServerLabel string `json:"server_label"`
Arguments json.RawMessage `json:"arguments"`
Status string `json:"status,omitempty"`
Output json.RawMessage `json:"output,omitempty"`
Error json.RawMessage `json:"error,omitempty"`
ApprovalRequestID string `json:"approval_request_id,omitempty"`
}{
Type: r.Type,
ID: r.ID,
Name: r.Name,
ServerLabel: r.ServerLabel,
Arguments: r.Arguments,
Status: r.Status,
Output: r.Output,
Error: r.ItemError,
ApprovalRequestID: r.ApprovalRequestID,
})
default:
type responsesOutputAlias ResponsesOutput
return kitutil.Marshal(responsesOutputAlias(r))
}
}
// NormalizeResponsesWebSearchAction validates and canonicalizes the current
// Responses web_search_call action union. Claude emits {"query": ...}; the
// Responses representation additionally requires a discriminator.
func NormalizeResponsesWebSearchAction(raw json.RawMessage) (json.RawMessage, error) {
var action struct {
Type string `json:"type"`
Query string `json:"query"`
Queries []string `json:"queries"`
Sources json.RawMessage `json:"sources"`
URL string `json:"url"`
Pattern string `json:"pattern"`
}
if err := kitutil.Unmarshal(raw, &action); err != nil {
return nil, fmt.Errorf("decode Responses web-search action: %w", err)
}
action.Type = strings.TrimSpace(action.Type)
action.Query = strings.TrimSpace(action.Query)
action.URL = strings.TrimSpace(action.URL)
action.Pattern = strings.TrimSpace(action.Pattern)
for index := range action.Queries {
action.Queries[index] = strings.TrimSpace(action.Queries[index])
if action.Queries[index] == "" {
return nil, fmt.Errorf("Responses web-search action queries[%d] must not be empty", index)
}
}
if action.Type == "" && (action.Query != "" || len(action.Queries) > 0) {
action.Type = "search"
}
var canonical any
switch action.Type {
case "search":
if action.Query == "" && len(action.Queries) == 0 {
return nil, fmt.Errorf("Responses web-search action %q requires query or queries", action.Type)
}
if len(action.Sources) > 0 && kitutil.GetJsonType(action.Sources) != "array" && kitutil.GetJsonType(action.Sources) != "null" {
return nil, fmt.Errorf("Responses web-search action sources must be an array")
}
canonical = struct {
Type string `json:"type"`
Query string `json:"query,omitempty"`
Queries []string `json:"queries,omitempty"`
Sources json.RawMessage `json:"sources,omitempty"`
}{Type: action.Type, Query: action.Query, Queries: action.Queries, Sources: action.Sources}
case "open_page":
if action.URL == "" {
return nil, fmt.Errorf("Responses web-search action %q requires url", action.Type)
}
canonical = struct {
Type string `json:"type"`
URL string `json:"url"`
}{Type: action.Type, URL: action.URL}
case "find", "find_in_page":
if action.URL == "" || action.Pattern == "" {
return nil, fmt.Errorf("Responses web-search action %q requires url and pattern", action.Type)
}
canonical = struct {
Type string `json:"type"`
URL string `json:"url"`
Pattern string `json:"pattern"`
}{Type: "find_in_page", URL: action.URL, Pattern: action.Pattern}
default:
return nil, fmt.Errorf("unsupported Responses web-search action type %q", action.Type)
}
encoded, err := kitutil.Marshal(canonical)
if err != nil {
return nil, fmt.Errorf("encode Responses web-search action: %w", err)
}
return encoded, nil
} }
// ArgumentsString returns function call arguments in the string form expected by Chat Completions. // ArgumentsString returns function call arguments in the string form expected by Chat Completions.
...@@ -384,10 +515,20 @@ const ( ...@@ -384,10 +515,20 @@ const (
// ResponsesStreamResponse 用于处理 /v1/responses 流式响应 // ResponsesStreamResponse 用于处理 /v1/responses 流式响应
type ResponsesStreamResponse struct { type ResponsesStreamResponse struct {
Type string `json:"type"` Type string `json:"type"`
Response *OpenAIResponsesResponse `json:"response,omitempty"` Response *OpenAIResponsesResponse `json:"response,omitempty"`
Delta string `json:"delta,omitempty"` Code string `json:"code,omitempty"`
Item *ResponsesOutput `json:"item,omitempty"` Message string `json:"message,omitempty"`
Param string `json:"param,omitempty"`
Delta string `json:"delta,omitempty"`
Arguments *string `json:"arguments,omitempty"`
Name string `json:"name,omitempty"`
Text *string `json:"text,omitempty"`
Item *ResponsesOutput `json:"item,omitempty"`
SequenceNumber *int `json:"sequence_number,omitempty"`
Annotation json.RawMessage `json:"annotation,omitempty"`
AnnotationIndex *int `json:"annotation_index,omitempty"`
Obfuscation string `json:"obfuscation,omitempty"`
// - response.function_call_arguments.delta // - response.function_call_arguments.delta
// - response.function_call_arguments.done // - response.function_call_arguments.done
OutputIndex *int `json:"output_index,omitempty"` OutputIndex *int `json:"output_index,omitempty"`
......
package dto
// ReasoningConversionState carries provider-native reasoning controls between
// in-process conversion steps. It is not part of any provider wire protocol;
// request fields that reference it must use json:"-".
//
// Converters that rebuild an OpenAI request must copy this state so exact
// budgets and explicit include-thoughts choices survive multi-step routes.
type ReasoningConversionState struct {
Mode string
Effort string
BudgetTokens *int
IncludeThoughts *bool
}
package dto
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMergeClaudeUsageCacheCreationReplacesWholeObject(t *testing.T) {
t.Parallel()
merged := mergeClaudeUsageNonZero(
&ClaudeUsage{
CacheCreation: &ClaudeCacheCreationUsage{Ephemeral1hInputTokens: 1000},
},
&ClaudeUsage{
CacheCreation: &ClaudeCacheCreationUsage{
Ephemeral5mInputTokens: 1000,
Ephemeral1hInputTokens: 0,
},
},
)
require.NotNil(t, merged.CacheCreation)
assert.Equal(t, 1000, merged.CacheCreation.Ephemeral5mInputTokens)
assert.Equal(t, 0, merged.CacheCreation.Ephemeral1hInputTokens)
}
func TestMergeGeminiUsageMetadataCandidatesAndThoughtsReplacedAsPair(t *testing.T) {
t.Parallel()
merged := MergeGeminiUsageMetadataNonZero(
&GeminiUsageMetadata{
PromptTokenCount: 10,
ThoughtsTokenCount: 100,
},
&GeminiUsageMetadata{
PromptTokenCount: 10,
CandidatesTokenCount: 150,
ThoughtsTokenCount: 0,
TotalTokenCount: 160,
},
)
require.NotNil(t, merged)
assert.Equal(t, 150, merged.CandidatesTokenCount)
assert.Equal(t, 0, merged.ThoughtsTokenCount)
billing := NewGeminiChatBillingUsage(merged)
usage, ok := billing.CanonicalUsage()
require.True(t, ok)
assert.Equal(t, 150, usage.CompletionTokens)
}
func TestMergeUsageNonZeroKeepsPositiveValuesAndTakesMaxTotal(t *testing.T) {
t.Parallel()
merged := MergeUsageNonZero(
&Usage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15},
&Usage{PromptTokens: 0, CompletionTokens: 0, TotalTokens: 20},
)
require.NotNil(t, merged)
assert.Equal(t, 10, merged.PromptTokens)
assert.Equal(t, 5, merged.CompletionTokens)
assert.Equal(t, 20, merged.TotalTokens)
}
...@@ -16,6 +16,11 @@ func ClaudeStopReasonToOpenAIFinishReason(stopReason string) string { ...@@ -16,6 +16,11 @@ func ClaudeStopReasonToOpenAIFinishReason(stopReason string) string {
return "length" return "length"
case "tool_use": case "tool_use":
return "tool_calls" return "tool_calls"
case "pause_turn":
// Responses has no pause_turn finish reason. Treat the provider's
// resumable server-side loop as an incomplete response instead of a
// successful stop; the hosted output items preserve continuation state.
return "length"
case "refusal": case "refusal":
return types.FinishReasonContentFilter return types.FinishReasonContentFilter
default: default:
......
...@@ -8,6 +8,7 @@ import ( ...@@ -8,6 +8,7 @@ import (
"github.com/QuantumNous/new-api/relaykit/relayconvert/convmeta" "github.com/QuantumNous/new-api/relaykit/relayconvert/convmeta"
sharedclaude "github.com/QuantumNous/new-api/relaykit/relayconvert/internal/shared/claude" sharedclaude "github.com/QuantumNous/new-api/relaykit/relayconvert/internal/shared/claude"
kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil" kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil"
"github.com/QuantumNous/new-api/relaykit/relayconvert/reasoning"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
...@@ -80,6 +81,21 @@ func TestClaudeDefaultMaxTokensPresence(t *testing.T) { ...@@ -80,6 +81,21 @@ func TestClaudeDefaultMaxTokensPresence(t *testing.T) {
require.NotNil(t, got.MaxTokens) require.NotNil(t, got.MaxTokens)
assert.Equal(t, clientMaxTokens, *got.MaxTokens) assert.Equal(t, clientMaxTokens, *got.MaxTokens)
}) })
t.Run("client zero same as absent, hook fills", func(t *testing.T) {
clientMaxTokens := uint(0)
got, err := converter.convert(t, claudeDefaultsMeta(func(string) int { return 512 }), &clientMaxTokens)
require.NoError(t, err)
require.NotNil(t, got.MaxTokens)
assert.Equal(t, uint(512), *got.MaxTokens)
})
t.Run("client zero same as absent, no hook fails", func(t *testing.T) {
clientMaxTokens := uint(0)
got, err := converter.convert(t, &convmeta.Values{}, &clientMaxTokens)
require.ErrorIs(t, err, sharedclaude.ErrMissingMaxTokens)
assert.Nil(t, got)
})
}) })
} }
} }
...@@ -88,12 +104,18 @@ func TestClaudeDefaultMaxTokensPresence(t *testing.T) { ...@@ -88,12 +104,18 @@ func TestClaudeDefaultMaxTokensPresence(t *testing.T) {
// "-thinking" request without max_tokens must keep converting even when no // "-thinking" request without max_tokens must keep converting even when no
// DefaultMaxTokens hook is configured. // DefaultMaxTokens hook is configured.
func TestClaudeThinkingAdapterSatisfiesMaxTokensWithoutCallback(t *testing.T) { func TestClaudeThinkingAdapterSatisfiesMaxTokensWithoutCallback(t *testing.T) {
meta := &convmeta.Values{Options: &convmeta.Options{ _, intent, found, err := reasoning.ParseClaudeModelSuffix("claude-test-thinking", true)
Claude: convmeta.ClaudeOptions{ require.NoError(t, err)
ThinkingAdapterEnabled: true, require.True(t, found)
ThinkingAdapterBudgetTokensPercentage: 0.8, meta := &convmeta.Values{
ReasoningConversion: reasoning.StateFromIntent(intent),
Options: &convmeta.Options{
Claude: convmeta.ClaudeOptions{
ThinkingAdapterEnabled: true,
ThinkingAdapterBudgetTokensPercentage: 0.8,
},
}, },
}} }
got, err := OpenAIChatRequestToClaudeMessages(context.Background(), meta, dto.GeneralOpenAIRequest{ got, err := OpenAIChatRequestToClaudeMessages(context.Background(), meta, dto.GeneralOpenAIRequest{
Model: "claude-test-thinking", Model: "claude-test-thinking",
Messages: []dto.Message{ Messages: []dto.Message{
......
...@@ -28,6 +28,10 @@ type Meta interface { ...@@ -28,6 +28,10 @@ type Meta interface {
// SetReasoningEffort records the effort level a converter derived from a // SetReasoningEffort records the effort level a converter derived from a
// model-name suffix so downstream billing/logging can see it. // model-name suffix so downstream billing/logging can see it.
SetReasoningEffort(effort string) SetReasoningEffort(effort string)
// ReasoningState returns the suffix-derived reasoning intent attached at
// the host entry layer. Standalone callers that do not set it receive nil;
// converters then use only explicit request fields.
ReasoningState() *dto.ReasoningConversionState
GetEstimatePromptTokens() int GetEstimatePromptTokens() int
// EnsureClaudeConvertInfo lazily creates and returns the mutable // EnsureClaudeConvertInfo lazily creates and returns the mutable
...@@ -60,6 +64,20 @@ type ClaudeConvertInfo struct { ...@@ -60,6 +64,20 @@ type ClaudeConvertInfo struct {
ToolCallBaseIndex int ToolCallBaseIndex int
ToolCallMaxIndexOffset int ToolCallMaxIndexOffset int
ToolCalls []*ClaudeStreamToolCall
ToolCallByIndex map[int]*ClaudeStreamToolCall
ToolCallByID map[string]*ClaudeStreamToolCall
}
// ClaudeStreamToolCall tracks one OpenAI tool_calls entry while it is encoded
// as a Claude tool_use content block. Chat tool indexes and Claude content
// block indexes are separate domains, so the mapping must remain explicit.
type ClaudeStreamToolCall struct {
BlockIndex int
ID string
Name string
PendingArguments string
Started bool
} }
const ( const (
...@@ -79,6 +97,7 @@ type Values struct { ...@@ -79,6 +97,7 @@ type Values struct {
ChannelType int ChannelType int
IsStream bool IsStream bool
ReasoningEffort string ReasoningEffort string
ReasoningConversion *dto.ReasoningConversionState
EstimatePromptTokens int EstimatePromptTokens int
ClaudeConvertInfo *ClaudeConvertInfo ClaudeConvertInfo *ClaudeConvertInfo
...@@ -139,6 +158,13 @@ func (v *Values) SetReasoningEffort(effort string) { ...@@ -139,6 +158,13 @@ func (v *Values) SetReasoningEffort(effort string) {
} }
} }
func (v *Values) ReasoningState() *dto.ReasoningConversionState {
if v == nil {
return nil
}
return v.ReasoningConversion
}
func (v *Values) GetEstimatePromptTokens() int { func (v *Values) GetEstimatePromptTokens() int {
if v == nil { if v == nil {
return 0 return 0
...@@ -213,3 +239,11 @@ func OptionsOf(m Meta) *Options { ...@@ -213,3 +239,11 @@ func OptionsOf(m Meta) *Options {
} }
return m.ConvOptions() return m.ConvOptions()
} }
// ReasoningStateOf is a nil-safe reader for Meta.ReasoningState.
func ReasoningStateOf(m Meta) *dto.ReasoningConversionState {
if m == nil {
return nil
}
return m.ReasoningState()
}
...@@ -19,6 +19,7 @@ func TestValuesTypedNilMetaIsSafe(t *testing.T) { ...@@ -19,6 +19,7 @@ func TestValuesTypedNilMetaIsSafe(t *testing.T) {
assert.Zero(t, meta.GetChannelType()) assert.Zero(t, meta.GetChannelType())
assert.False(t, meta.GetIsStream()) assert.False(t, meta.GetIsStream())
assert.Empty(t, meta.GetReasoningEffort()) assert.Empty(t, meta.GetReasoningEffort())
assert.Nil(t, meta.ReasoningState())
assert.Zero(t, meta.GetEstimatePromptTokens()) assert.Zero(t, meta.GetEstimatePromptTokens())
assert.Zero(t, meta.GetSendResponseCount()) assert.Zero(t, meta.GetSendResponseCount())
......
package convmeta package convmeta
import "github.com/QuantumNous/new-api/relaykit/types"
// Options is the per-request snapshot of host configuration that converters // Options is the per-request snapshot of host configuration that converters
// consult. The host fills it from its settings system when constructing the // consult. The host fills it from its settings system when constructing the
// Meta (see relaycommon.RelayInfo.ConvOptions); relaykit users fill it // Meta (see relaycommon.RelayInfo.ConvOptions); relaykit users fill it
...@@ -8,6 +10,13 @@ type Options struct { ...@@ -8,6 +10,13 @@ type Options struct {
Claude ClaudeOptions Claude ClaudeOptions
Gemini GeminiOptions Gemini GeminiOptions
// ToolLossPolicy controls whether a cross-protocol conversion may omit or
// approximate built-in-tool semantics. The zero value uses the allow
// policy: conversion succeeds and every loss is returned as a diagnostic.
// safe/strict rejection is request-phase opt-in only; response and stream
// conversion never reject regardless of this field.
ToolLossPolicy types.ConversionLossPolicy
// OpenRouterDialect marks the upstream as OpenRouter's OpenAI-compatible // OpenRouterDialect marks the upstream as OpenRouter's OpenAI-compatible
// surface, which accepts extra fields (reasoning config, cache_control on // surface, which accepts extra fields (reasoning config, cache_control on
// system parts) that converters emit only for that dialect. The host sets // system parts) that converters emit only for that dialect. The host sets
...@@ -18,11 +27,16 @@ type Options struct { ...@@ -18,11 +27,16 @@ type Options struct {
// suffix must be kept on the outgoing model name (host blacklist lookup). // suffix must be kept on the outgoing model name (host blacklist lookup).
// Nil means "never preserve". // Nil means "never preserve".
PreserveThinkingSuffix func(modelName string) bool PreserveThinkingSuffix func(modelName string) bool
// PreserveEffortTail reports real model IDs whose names already end in an
// effort-like token (for example qwen-max). Nil means "never preserve".
PreserveEffortTail func(modelName string) bool
} }
type ClaudeOptions struct { type ClaudeOptions struct {
// ThinkingAdapterEnabled turns "-thinking"-suffixed OpenAI model names // ThinkingAdapterEnabled controls whether suffix-derived reasoning intent
// into Claude extended-thinking requests. // is rendered onto Claude thinking / output_config. Suffix parsing itself
// is the host entry layer's job (standalone users call Parse* themselves).
ThinkingAdapterEnabled bool ThinkingAdapterEnabled bool
// ThinkingAdapterBudgetTokensPercentage sizes thinking budget_tokens as a // ThinkingAdapterBudgetTokensPercentage sizes thinking budget_tokens as a
// fraction of max_tokens when the adapter fires. // fraction of max_tokens when the adapter fires.
...@@ -36,11 +50,16 @@ type ClaudeOptions struct { ...@@ -36,11 +50,16 @@ type ClaudeOptions struct {
// standalone relaykit users must supply one or guarantee max_tokens on // standalone relaykit users must supply one or guarantee max_tokens on
// every request. // every request.
DefaultMaxTokens func(modelName string) int DefaultMaxTokens func(modelName string) int
// WebSearchToolVersion selects the Claude hosted web-search tool version
// emitted by cross-protocol conversion. Empty keeps the compatibility
// baseline web_search_20250305.
WebSearchToolVersion string
} }
type GeminiOptions struct { type GeminiOptions struct {
// ThinkingAdapterEnabled maps -thinking/-nothinking/effort suffixes to // ThinkingAdapterEnabled controls whether suffix-derived reasoning intent
// Gemini thinkingConfig. // is rendered onto Gemini thinkingConfig. Suffix parsing itself is the
// host entry layer's job (standalone users call Parse* themselves).
ThinkingAdapterEnabled bool ThinkingAdapterEnabled bool
// ThinkingAdapterBudgetTokensPercentage sizes thinkingBudget as a fraction // ThinkingAdapterBudgetTokensPercentage sizes thinkingBudget as a fraction
// of maxOutputTokens when the adapter fires. // of maxOutputTokens when the adapter fires.
...@@ -77,3 +96,14 @@ func (o *GeminiOptions) SafetySettingFor(category string) string { ...@@ -77,3 +96,14 @@ func (o *GeminiOptions) SafetySettingFor(category string) string {
func (o *Options) ShouldPreserveThinkingSuffix(modelName string) bool { func (o *Options) ShouldPreserveThinkingSuffix(modelName string) bool {
return o != nil && o.PreserveThinkingSuffix != nil && o.PreserveThinkingSuffix(modelName) return o != nil && o.PreserveThinkingSuffix != nil && o.PreserveThinkingSuffix(modelName)
} }
func (o *Options) ShouldPreserveEffortTail(modelName string) bool {
return o != nil && o.PreserveEffortTail != nil && o.PreserveEffortTail(modelName)
}
func (o *Options) EffectiveToolLossPolicy() types.ConversionLossPolicy {
if o == nil || o.ToolLossPolicy == "" {
return types.ConversionLossPolicyAllow
}
return o.ToolLossPolicy
}
package relayconvert package relayconvert
// golden_test.go pins the byte-level output of every registered (from, to) // golden_test.go pins the byte-level output of selected public conversion
// conversion route so the relaykit extraction refactor can prove behavior is // routes. Run with -update to regenerate testdata/golden.
// unchanged at each phase. Run with -update to regenerate testdata/golden.
// //
// Volatile values (generated UUID-based ids, unix timestamps) are normalized // Volatile values (generated UUID-based ids, unix timestamps) are normalized
// before comparison so the snapshots are deterministic. // before comparison so the snapshots are deterministic.
...@@ -69,10 +68,30 @@ func checkGolden(t *testing.T, name string, got []byte) { ...@@ -69,10 +68,30 @@ func checkGolden(t *testing.T, name string, got []byte) {
return return
} }
want, err := os.ReadFile(path) want, err := os.ReadFile(path)
require.NoError(t, err, "golden file missing, run: go test ./service/relayconvert -run TestGolden -update") require.NoError(t, err, "golden file missing, run: cd relaykit && GOWORK=off go test ./relayconvert -run TestGolden -update")
require.Equal(t, string(want), string(got), "conversion output drifted from golden snapshot %s", path) require.Equal(t, string(want), string(got), "conversion output drifted from golden snapshot %s", path)
} }
func checkStreamEventsGolden(t *testing.T, name string, events []any) {
t.Helper()
got := marshalGolden(t, map[string]any{"events": events})
path := filepath.Join(goldenDir, name+".golden.json")
if *updateGolden {
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755))
require.NoError(t, os.WriteFile(path, got, 0o644))
return
}
wantData, err := os.ReadFile(path)
require.NoError(t, err, "golden file missing, run: cd relaykit && GOWORK=off go test ./relayconvert -run TestGolden -update")
var wantSnapshot map[string]json.RawMessage
require.NoError(t, json.Unmarshal(wantData, &wantSnapshot))
wantEvents, ok := wantSnapshot["events"]
require.True(t, ok, "stream golden snapshot %s has no events", path)
want := marshalGolden(t, map[string]json.RawMessage{"events": wantEvents})
require.Equal(t, string(want), string(got), "conversion events drifted from golden snapshot %s", path)
}
// goldenInfo mirrors the host's default converter options (new-api's // goldenInfo mirrors the host's default converter options (new-api's
// model_setting defaults at the time the snapshots were recorded) so the // model_setting defaults at the time the snapshots were recorded) so the
// golden files stay comparable across the extraction. // golden files stay comparable across the extraction.
...@@ -120,42 +139,6 @@ func fixtureRequests() map[types.RelayFormat]any { ...@@ -120,42 +139,6 @@ func fixtureRequests() map[types.RelayFormat]any {
"tool_choice": "auto" "tool_choice": "auto"
}`, openai) }`, openai)
claude := &dto.ClaudeRequest{}
mustUnmarshalFixture(`{
"model": "claude-test",
"max_tokens": 1024,
"stream": true,
"system": "You are a helpful assistant.",
"messages": [
{"role": "user", "content": [
{"type": "text", "text": "What is in this image?"},
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "aGVsbG8="}}
]},
{"role": "assistant", "content": [
{"type": "thinking", "thinking": "Let me look.", "signature": "sig"},
{"type": "tool_use", "id": "toolu_abc", "name": "get_weather", "input": {"city": "Paris"}}
]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_abc", "content": "15 degrees"}]}
],
"tools": [{"name": "get_weather", "description": "Get weather by city", "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}}],
"thinking": {"type": "enabled", "budget_tokens": 512}
}`, claude)
gemini := &dto.GeminiChatRequest{}
mustUnmarshalFixture(`{
"contents": [
{"role": "user", "parts": [
{"text": "What is in this image?"},
{"inlineData": {"mimeType": "image/png", "data": "aGVsbG8="}}
]},
{"role": "model", "parts": [{"functionCall": {"name": "get_weather", "args": {"city": "Paris"}}}]},
{"role": "user", "parts": [{"functionResponse": {"name": "get_weather", "response": {"result": "15 degrees"}}}]}
],
"systemInstruction": {"parts": [{"text": "You are a helpful assistant."}]},
"tools": [{"functionDeclarations": [{"name": "get_weather", "description": "Get weather by city", "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}}]}],
"generationConfig": {"maxOutputTokens": 1024, "temperature": 0.7}
}`, gemini)
responses := &dto.OpenAIResponsesRequest{} responses := &dto.OpenAIResponsesRequest{}
mustUnmarshalFixture(`{ mustUnmarshalFixture(`{
"model": "gpt-test", "model": "gpt-test",
...@@ -175,8 +158,6 @@ func fixtureRequests() map[types.RelayFormat]any { ...@@ -175,8 +158,6 @@ func fixtureRequests() map[types.RelayFormat]any {
return map[types.RelayFormat]any{ return map[types.RelayFormat]any{
types.RelayFormatOpenAI: openai, types.RelayFormatOpenAI: openai,
types.RelayFormatClaude: claude,
types.RelayFormatGemini: gemini,
types.RelayFormatOpenAIResponses: responses, types.RelayFormatOpenAIResponses: responses,
} }
} }
...@@ -308,8 +289,17 @@ func allFormats() []types.RelayFormat { ...@@ -308,8 +289,17 @@ func allFormats() []types.RelayFormat {
func TestGoldenRequestConversionMatrix(t *testing.T) { func TestGoldenRequestConversionMatrix(t *testing.T) {
requests := fixtureRequests() requests := fixtureRequests()
for _, from := range allFormats() { fromFormats := []types.RelayFormat{
for _, to := range allFormats() { types.RelayFormatOpenAI,
types.RelayFormatOpenAIResponses,
}
toFormats := []types.RelayFormat{
types.RelayFormatOpenAI,
types.RelayFormatClaude,
types.RelayFormatOpenAIResponses,
}
for _, from := range fromFormats {
for _, to := range toFormats {
if from == to { if from == to {
continue continue
} }
...@@ -330,7 +320,8 @@ func TestGoldenResponseConversionMatrix(t *testing.T) { ...@@ -330,7 +320,8 @@ func TestGoldenResponseConversionMatrix(t *testing.T) {
responses := fixtureResponses() responses := fixtureResponses()
for _, from := range allFormats() { for _, from := range allFormats() {
for _, to := range allFormats() { for _, to := range allFormats() {
if from == to { if from == to || to == types.RelayFormatGemini ||
(from == types.RelayFormatOpenAI && to == types.RelayFormatClaude) {
continue continue
} }
name := fmt.Sprintf("response/%s_to_%s", from, to) name := fmt.Sprintf("response/%s_to_%s", from, to)
...@@ -373,11 +364,9 @@ func TestGoldenStreamConversionMatrix(t *testing.T) { ...@@ -373,11 +364,9 @@ func TestGoldenStreamConversionMatrix(t *testing.T) {
outputs = append(outputs, r.Value) outputs = append(outputs, r.Value)
} }
snapshot := map[string]any{ // Billing usage has private, cross-module acceptance coverage. Keep
"events": outputs, // the public golden focused on client-visible stream events.
"usage": state.Usage(), checkStreamEventsGolden(t, name, outputs)
}
checkGolden(t, name, marshalGolden(t, snapshot))
}) })
} }
} }
......
package claudemessages
import (
"encoding/json"
"fmt"
"strings"
"unicode/utf8"
kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil"
)
func claudeCitationsToChat(raw json.RawMessage, text string, textOffset int) ([]any, error) {
if len(raw) == 0 {
return nil, nil
}
var citations []map[string]any
if err := kitutil.Unmarshal(raw, &citations); err != nil {
return nil, fmt.Errorf("invalid Claude citations: %w", err)
}
annotations := make([]any, 0, len(citations))
for _, citation := range citations {
url := strings.TrimSpace(kitutil.Interface2String(citation["url"]))
if url == "" {
continue
}
converted := map[string]any{
"url": url,
"title": strings.TrimSpace(kitutil.Interface2String(citation["title"])),
}
citedText := kitutil.Interface2String(citation["cited_text"])
if citedText != "" {
converted["cited_text"] = citedText
if index := strings.Index(text, citedText); index >= 0 {
startIndex := textOffset + utf8.RuneCountInString(text[:index])
converted["start_index"] = startIndex
converted["end_index"] = startIndex + utf8.RuneCountInString(citedText)
}
}
if encryptedIndex := kitutil.Interface2String(citation["encrypted_index"]); encryptedIndex != "" {
converted["encrypted_index"] = encryptedIndex
}
if converted["title"] == "" {
delete(converted, "title")
}
annotations = append(annotations, map[string]any{
"type": "url_citation",
"url_citation": converted,
})
}
return annotations, nil
}
func marshalChatAnnotations(annotations []any) (json.RawMessage, error) {
if len(annotations) == 0 {
return nil, nil
}
return kitutil.Marshal(annotations)
}
package claudemessages
import (
"testing"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestMessageStartZeroOutputSidecarRemainsRefreshable(t *testing.T) {
t.Parallel()
info := &ClaudeResponseInfo{Usage: &dto.Usage{}}
ok := FormatClaudeResponseInfo(&dto.ClaudeResponse{
Type: "message_start",
Message: &dto.ClaudeMediaMessage{
Id: "msg_1",
Model: "claude-test",
Usage: &dto.ClaudeUsage{
InputTokens: 10,
OutputTokens: 0,
BillingUsage: dto.NewClaudeMessagesBillingUsage(&dto.ClaudeUsage{
InputTokens: 10,
OutputTokens: 0,
}),
},
},
}, nil, info)
require.True(t, ok)
ok = FormatClaudeResponseInfo(&dto.ClaudeResponse{
Type: "message_delta",
Usage: &dto.ClaudeUsage{
OutputTokens: 42,
},
}, nil, info)
require.True(t, ok)
require.NotNil(t, info.Usage.BillingUsage)
require.NotNil(t, info.Usage.BillingUsage.ClaudeUsage)
assert.Equal(t, 42, info.Usage.BillingUsage.ClaudeUsage.OutputTokens)
}
func TestTerminalSidecarRemainsAuthoritativeAgainstFinalize(t *testing.T) {
t.Parallel()
info := &ClaudeResponseInfo{Usage: &dto.Usage{}}
ok := FormatClaudeResponseInfo(&dto.ClaudeResponse{
Type: "message_delta",
Usage: &dto.ClaudeUsage{
InputTokens: 10,
OutputTokens: 7,
BillingUsage: dto.NewClaudeMessagesBillingUsage(&dto.ClaudeUsage{
InputTokens: 10,
OutputTokens: 7,
}),
},
}, nil, info)
require.True(t, ok)
require.NotNil(t, info.Usage.BillingUsage)
require.NotNil(t, info.Usage.BillingUsage.ClaudeUsage)
info.Usage.CompletionTokens = 99
FinalizeClaudeStreamBillingUsage(info)
assert.Equal(t, 7, info.Usage.BillingUsage.ClaudeUsage.OutputTokens)
}
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment