Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions dto/channel_settings.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,13 @@ type ChannelSettings struct {
PassThroughBodyEnabled bool `json:"pass_through_body_enabled,omitempty"`
SystemPrompt string `json:"system_prompt,omitempty"`
SystemPromptOverride bool `json:"system_prompt_override,omitempty"`
// TrustUpstreamUsage, when enabled, makes the relay prefer the usage
// reported by the upstream in streaming responses over the locally
// streamed token count. Streaming paths no longer buffer the full
// response text regardless of this flag; the local count is always
// available as a bounded-memory fallback. Defaults to false, so the
// locally streamed count is used unless the upstream usage is trusted.
TrustUpstreamUsage bool `json:"trust_upstream_usage,omitempty"`
}

type VertexKeyType string
Expand Down
2 changes: 0 additions & 2 deletions relay/channel/aws/relay-aws.go
Original file line number Diff line number Diff line change
Expand Up @@ -236,7 +236,6 @@ func awsHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (*types
ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(),
Model: info.UpstreamModelName,
ResponseText: strings.Builder{},
Usage: &dto.Usage{},
}

Expand Down Expand Up @@ -268,7 +267,6 @@ func awsStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (
ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(),
Model: info.UpstreamModelName,
ResponseText: strings.Builder{},
Usage: &dto.Usage{},
}

Expand Down
36 changes: 22 additions & 14 deletions relay/channel/claude/relay-claude.go
Original file line number Diff line number Diff line change
Expand Up @@ -583,12 +583,22 @@ func ResponseClaude2OpenAI(claudeResponse *dto.ClaudeResponse) *dto.OpenAITextRe
}

type ClaudeResponseInfo struct {
ResponseId string
Created int64
Model string
ResponseText strings.Builder
Usage *dto.Usage
Done bool
ResponseId string
Created int64
Model string
// usageAcc 流式累计 completion token(text + thinking 分离计数),
// 替代原先用 strings.Builder 累积整段响应文本再估算的做法,避免大响应内存堆积。
usageAcc *service.UsageAccumulator
Usage *dto.Usage
Done bool
}

// ensureUsageAcc 懒初始化 usageAcc(首次累积时按 Model 创建)。
func (cri *ClaudeResponseInfo) ensureUsageAcc() *service.UsageAccumulator {
if cri.usageAcc == nil {
cri.usageAcc = service.NewUsageAccumulator(cri.Model)
}
return cri.usageAcc
}

func cacheCreationTokensForOpenAIUsage(usage *dto.Usage) int {
Expand Down Expand Up @@ -738,10 +748,10 @@ func FormatClaudeResponseInfo(claudeResponse *dto.ClaudeResponse, oaiResponse *d
} else if claudeResponse.Type == "content_block_delta" {
if claudeResponse.Delta != nil {
if claudeResponse.Delta.Text != nil {
claudeInfo.ResponseText.WriteString(*claudeResponse.Delta.Text)
claudeInfo.ensureUsageAcc().Feed(*claudeResponse.Delta.Text)
}
if claudeResponse.Delta.Thinking != nil {
claudeInfo.ResponseText.WriteString(*claudeResponse.Delta.Thinking)
claudeInfo.ensureUsageAcc().FeedReasoning(*claudeResponse.Delta.Thinking)
}
}
} else if claudeResponse.Type == "message_delta" {
Expand Down Expand Up @@ -840,13 +850,13 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau
common.SysLog("claude response usage is not complete, maybe upstream error")
}
// 只补缺失字段,不整份覆盖——保留 message_start 已拿到的 cache 字段
fallback := service.ResponseText2Usage(c, claudeInfo.ResponseText.String(), info.UpstreamModelName, info.GetEstimatePromptTokens())
fallbackCompletion := claudeInfo.ensureUsageAcc().LocalCompletionTokens()
if claudeInfo.Usage.CompletionTokens == 0 ||
(!claudeInfo.Done && fallback.CompletionTokens > claudeInfo.Usage.CompletionTokens) {
claudeInfo.Usage.CompletionTokens = fallback.CompletionTokens
(!claudeInfo.Done && fallbackCompletion > claudeInfo.Usage.CompletionTokens) {
claudeInfo.Usage.CompletionTokens = fallbackCompletion
}
if claudeInfo.Usage.PromptTokens == 0 {
claudeInfo.Usage.PromptTokens = fallback.PromptTokens
claudeInfo.Usage.PromptTokens = info.GetEstimatePromptTokens()
}
claudeInfo.Usage.TotalTokens = claudeInfo.Usage.PromptTokens + claudeInfo.Usage.CompletionTokens
}
Expand Down Expand Up @@ -874,7 +884,6 @@ func ClaudeStreamHandler(c *gin.Context, resp *http.Response, info *relaycommon.
ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(),
Model: info.UpstreamModelName,
ResponseText: strings.Builder{},
Usage: &dto.Usage{},
}
var err *types.NewAPIError
Expand Down Expand Up @@ -943,7 +952,6 @@ func ClaudeHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayI
ResponseId: helper.GetResponseID(c),
Created: common.GetTimestamp(),
Model: info.UpstreamModelName,
ResponseText: strings.Builder{},
Usage: &dto.Usage{},
}
responseBody, err := io.ReadAll(resp.Body)
Expand Down
10 changes: 5 additions & 5 deletions relay/channel/claude/relay_claude_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@ package claude

import (
"encoding/base64"
"strings"
"testing"

"github.com/QuantumNous/new-api/dto"
Expand Down Expand Up @@ -161,8 +160,8 @@ func TestFormatClaudeResponseInfo_NilClaudeInfo(t *testing.T) {
func TestFormatClaudeResponseInfo_ContentBlockDelta(t *testing.T) {
text := "hello"
claudeInfo := &ClaudeResponseInfo{
Usage: &dto.Usage{},
ResponseText: strings.Builder{},
Model: "claude-3-5-sonnet",
Usage: &dto.Usage{},
}
claudeResponse := &dto.ClaudeResponse{
Type: "content_block_delta",
Expand All @@ -175,8 +174,9 @@ func TestFormatClaudeResponseInfo_ContentBlockDelta(t *testing.T) {
if !ok {
t.Fatal("expected true")
}
if claudeInfo.ResponseText.String() != "hello" {
t.Errorf("ResponseText = %q, want %q", claudeInfo.ResponseText.String(), "hello")
// 文本通过流式累计器计数,应反映已喂入的 "hello"
if got := claudeInfo.ensureUsageAcc().LocalCompletionTokens(); got <= 0 {
t.Errorf("LocalCompletionTokens = %d, want > 0 for %q", got, text)
}
}

Expand Down
147 changes: 147 additions & 0 deletions relay/channel/claude/stream_handler_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,147 @@
package claude

import (
"io"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"

"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestMain(m *testing.M) {
// StreamScannerHandler uses time.NewTicker(StreamingTimeout); avoid zero-interval panic.
if constant.StreamingTimeout <= 0 {
constant.StreamingTimeout = 300
}
os.Exit(m.Run())
}

func newStreamTestInfo(model string, trust bool) *relaycommon.RelayInfo {
info := &relaycommon.RelayInfo{
RelayFormat: types.RelayFormatClaude,
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: model,
ChannelSetting: dto.ChannelSettings{TrustUpstreamUsage: trust},
},
}
info.SetEstimatePromptTokens(100)
return info
}

func newSSEResp(sse string) *http.Response {
return &http.Response{
Body: io.NopCloser(strings.NewReader(sse)),
StatusCode: http.StatusOK,
}
}

func newStreamTestCtx() *gin.Context {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
return c
}

// 上游提供 usage 且 trust=true:应直接用上游 usage(含 output_tokens)。
func TestClaudeStreamHandler_TrustUpstreamUsage(t *testing.T) {
c := newStreamTestCtx()
info := newStreamTestInfo("claude-3-5-sonnet", true)
sse := `data: {"type":"message_start","message":{"id":"msg_1","model":"claude-3-5-sonnet","usage":{"input_tokens":100,"output_tokens":1}}}
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello world this is the answer"}}
data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"input_tokens":100,"output_tokens":42}}
data: {"type":"message_stop"}
`
usage, apiErr := ClaudeStreamHandler(c, newSSEResp(sse), info)
require.Nil(t, apiErr)
require.NotNil(t, usage)
require.Equal(t, 100, usage.PromptTokens)
require.Equal(t, 42, usage.CompletionTokens, "trust=true 应采用上游 output_tokens=42")
}

// 上游未给 output_tokens(异常/中断):应回退到本地流式估算,且不为 0。
func TestClaudeStreamHandler_LocalFallback(t *testing.T) {
c := newStreamTestCtx()
info := newStreamTestInfo("claude-3-5-sonnet", false)
// message_delta 不带 usage,message_stop 前断;本地需要根据文本估算
sse := `data: {"type":"message_start","message":{"id":"msg_1","model":"claude-3-5-sonnet","usage":{"input_tokens":100,"output_tokens":0}}}
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello world this is a fairly long answer text"}}
`
usage, apiErr := ClaudeStreamHandler(c, newSSEResp(sse), info)
require.Nil(t, apiErr)
require.NotNil(t, usage)
require.Greater(t, usage.CompletionTokens, 0, "上游未给 usage 时本地估算应 > 0")
}

// thinking + text 分离计数:thinking 也应计入 completion。
func TestClaudeStreamHandler_ThinkingCounted(t *testing.T) {
c := newStreamTestCtx()
info := newStreamTestInfo("claude-3-5-sonnet", false)
textOnly := `data: {"type":"message_start","message":{"id":"m","model":"claude-3-5-sonnet","usage":{"input_tokens":10,"output_tokens":0}}}
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"visible answer"}}
`
withThinking := `data: {"type":"message_start","message":{"id":"m","model":"claude-3-5-sonnet","usage":{"input_tokens":10,"output_tokens":0}}}
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"visible answer"}}
data: {"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"some internal reasoning content here that is fairly long"}}
`
u1, e1 := ClaudeStreamHandler(newStreamTestCtx(), newSSEResp(textOnly), info)
require.Nil(t, e1)
u2, e2 := ClaudeStreamHandler(c, newSSEResp(withThinking), newStreamTestInfo("claude-3-5-sonnet", false))
require.Nil(t, e2)
require.Greater(t, u2.CompletionTokens, u1.CompletionTokens, "带 thinking 的 completion 应更大(thinking 被计入)")
}

// cache 字段(read/creation)必须从 message_start 正确传递到最终 usage,
// 不被本次累积重构破坏。
func TestClaudeStreamHandler_CacheTokensPreserved(t *testing.T) {
c := newStreamTestCtx()
info := newStreamTestInfo("claude-3-5-sonnet", true)
sse := `data: {"type":"message_start","message":{"id":"m","model":"claude-3-5-sonnet","usage":{"input_tokens":50,"output_tokens":1,"cache_read_input_tokens":4096,"cache_creation_input_tokens":256}}}
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"answer text here"}}
data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"input_tokens":50,"output_tokens":20}}
data: {"type":"message_stop"}
`
usage, apiErr := ClaudeStreamHandler(c, newSSEResp(sse), info)
require.Nil(t, apiErr)
require.NotNil(t, usage)
require.Equal(t, 4096, usage.PromptTokensDetails.CachedTokens, "cache_read 应保留")
require.Equal(t, 256, usage.PromptTokensDetails.CachedCreationTokens, "cache_creation 应保留")
require.Equal(t, 20, usage.CompletionTokens, "trust=true 用上游 output_tokens")
}

// 使用从生产 sub2api 抓取的【真实】Claude SSE 响应(含 event: 行、ping、
// cache_creation 嵌套结构、末尾空白),验证 handler 在真实上游格式下正确工作。
// 这不是臆想的格式——是 2026-06 实际抓包内容(已脱敏 id)。
func TestClaudeStreamHandler_RealUpstreamFormat(t *testing.T) {
c := newStreamTestCtx()
info := newStreamTestInfo("claude-haiku-4-5", true)
// 注意:真实流每个 data 行后有尾随空格、event: 行穿插、有 ping 事件。
sse := "event: message_start\n" +
`data: {"type":"message_start","message":{"model":"claude-haiku-4-5","id":"msg_x","type":"message","role":"assistant","content":[],"usage":{"input_tokens":8,"cache_creation_input_tokens":0,"cache_read_input_tokens":0,"cache_creation":{"ephemeral_5m_input_tokens":0,"ephemeral_1h_input_tokens":0},"output_tokens":1}} }` + "\n\n" +
"event: content_block_start\n" +
`data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""} }` + "\n\n" +
"event: ping\n" +
`data: {"type": "ping"}` + "\n\n" +
"event: content_block_delta\n" +
`data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hey"} }` + "\n\n" +
"event: content_block_delta\n" +
`data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"! How's it going?"} }` + "\n\n" +
"event: content_block_stop\n" +
`data: {"type":"content_block_stop","index":0 }` + "\n\n" +
"event: message_delta\n" +
`data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"input_tokens":8,"output_tokens":11}}` + "\n\n" +
"event: message_stop\n" +
`data: {"type":"message_stop"}` + "\n\n"
usage, apiErr := ClaudeStreamHandler(c, newSSEResp(sse), info)
require.Nil(t, apiErr)
require.NotNil(t, usage)
require.Equal(t, 8, usage.PromptTokens, "真实 message_start input_tokens")
require.Equal(t, 11, usage.CompletionTokens, "真实 message_delta output_tokens (trust=true)")
}
10 changes: 7 additions & 3 deletions relay/channel/cloudflare/relay_cloudflare.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res

helper.SetEventStreamHeaders(c)
id := helper.GetResponseID(c)
var responseText string
usageAcc := service.NewUsageAccumulator(info.UpstreamModelName)
isFirst := true

for scanner.Scan() {
Expand All @@ -58,7 +58,7 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res
}
for _, choice := range response.Choices {
choice.Delta.Role = "assistant"
responseText += choice.Delta.GetContentString()
usageAcc.Feed(choice.Delta.GetContentString())
}
response.Id = id
response.Model = info.UpstreamModelName
Expand All @@ -75,7 +75,11 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res
if err := scanner.Err(); err != nil {
logger.LogError(c, "error_scanning_stream_response: "+err.Error())
}
usage := service.ResponseText2Usage(c, responseText, info.UpstreamModelName, info.GetEstimatePromptTokens())
usage := &dto.Usage{
PromptTokens: info.GetEstimatePromptTokens(),
CompletionTokens: usageAcc.LocalCompletionTokens(),
}
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
if info.ShouldIncludeUsage {
response := helper.GenerateFinalUsageResponse(id, info.StartTime.Unix(), info.UpstreamModelName, *usage)
err := helper.ObjectData(c, response)
Expand Down
62 changes: 62 additions & 0 deletions relay/channel/cloudflare/stream_handler_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
package cloudflare

import (
"io"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"
"time"

"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestMain(m *testing.M) {
if constant.StreamingTimeout <= 0 {
constant.StreamingTimeout = 300
}
os.Exit(m.Run())
}

func streamInfo(model string) *relaycommon.RelayInfo {
info := &relaycommon.RelayInfo{
RelayFormat: types.RelayFormatOpenAI,
StartTime: time.Now(),
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: model,
ChannelSetting: dto.ChannelSettings{},
},
}
info.SetEstimatePromptTokens(100)
return info
}

func sseResp(s string) *http.Response {
return &http.Response{Body: io.NopCloser(strings.NewReader(s)), StatusCode: http.StatusOK}
}

func streamCtx() *gin.Context {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
return c
}

// cloudflare 流式(无上游 usage):本地估算 completion > 0。注意返回值顺序 (error, usage)。
func TestCfStreamHandler_LocalEstimate(t *testing.T) {
sse := `data: {"id":"x","choices":[{"delta":{"role":"assistant","content":"Hello world this is"}}]}
data: {"id":"x","choices":[{"delta":{"content":" a cloudflare answer"}}]}
data: [DONE]
`
apiErr, usage := cfStreamHandler(streamCtx(), streamInfo("@cf/meta/llama"), sseResp(sse))
require.Nil(t, apiErr)
require.NotNil(t, usage)
require.Greater(t, usage.CompletionTokens, 0)
require.Equal(t, 100, usage.PromptTokens)
}
Loading