diff --git a/controller/task.go b/controller/task.go index eac7db153b48..93d7c1528187 100644 --- a/controller/task.go +++ b/controller/task.go @@ -89,6 +89,14 @@ func tasksToDto(tasks []*model.Task, fillUser bool) []*dto.TaskDto { } } result[i] = relay.TaskModel2Dto(task) + if !fillUser { + // Strip upstream model name from user-facing responses. + if props, ok := result[i].Properties.(model.Properties); ok { + stripped := props + stripped.UpstreamModelName = "" + result[i].Properties = stripped + } + } } return result } diff --git a/model/log.go b/model/log.go index 8ec7807e0339..a1f0e617e2f0 100644 --- a/model/log.go +++ b/model/log.go @@ -62,6 +62,9 @@ func formatUserLogs(logs []*Log, startIdx int) { delete(otherMap, "admin_info") // delete(otherMap, "reject_reason") delete(otherMap, "stream_status") + // Remove model mapping fields visible only to admins. + delete(otherMap, "is_model_mapped") + delete(otherMap, "upstream_model_name") } logs[i].Other = common.MapToJsonStr(otherMap) logs[i].Id = startIdx + i + 1 diff --git a/relay/channel/claude/relay-claude.go b/relay/channel/claude/relay-claude.go index e177e56dab14..fa928ab356a3 100644 --- a/relay/channel/claude/relay-claude.go +++ b/relay/channel/claude/relay-claude.go @@ -434,10 +434,13 @@ func RequestOpenAI2ClaudeMessage(c *gin.Context, textRequest dto.GeneralOpenAIRe return &claudeRequest, nil } -func StreamResponseClaude2OpenAI(claudeResponse *dto.ClaudeResponse) *dto.ChatCompletionsStreamResponse { +func StreamResponseClaude2OpenAI(claudeResponse *dto.ClaudeResponse, callerModel string) *dto.ChatCompletionsStreamResponse { var response dto.ChatCompletionsStreamResponse response.Object = "chat.completion.chunk" - response.Model = claudeResponse.Model + response.Model = callerModel + if response.Model == "" { + response.Model = claudeResponse.Model + } response.Choices = make([]dto.ChatCompletionsStreamResponseChoice, 0) tools := make([]dto.ToolCallResponse, 0) fcIdx := 0 @@ -451,7 +454,6 @@ func StreamResponseClaude2OpenAI(claudeResponse *dto.ClaudeResponse) *dto.ChatCo if claudeResponse.Type == "message_start" { if claudeResponse.Message != nil { response.Id = claudeResponse.Message.Id - response.Model = claudeResponse.Message.Model } //claudeUsage = &claudeResponse.Message.Usage choice.Delta.SetContentString("") @@ -518,7 +520,7 @@ func StreamResponseClaude2OpenAI(claudeResponse *dto.ClaudeResponse) *dto.ChatCo return &response } -func ResponseClaude2OpenAI(claudeResponse *dto.ClaudeResponse) *dto.OpenAITextResponse { +func ResponseClaude2OpenAI(claudeResponse *dto.ClaudeResponse, callerModel string) *dto.OpenAITextResponse { choices := make([]dto.OpenAITextResponseChoice, 0) fullTextResponse := dto.OpenAITextResponse{ Id: fmt.Sprintf("chatcmpl-%s", common.GetUUID()), @@ -575,7 +577,11 @@ func ResponseClaude2OpenAI(claudeResponse *dto.ClaudeResponse) *dto.OpenAITextRe if thinkingContent != "" { choice.Message.ReasoningContent = &thinkingContent } - fullTextResponse.Model = claudeResponse.Model + model := callerModel + if model == "" { + model = claudeResponse.Model + } + fullTextResponse.Model = model choices = append(choices, choice) fullTextResponse.Choices = choices return &fullTextResponse @@ -778,7 +784,10 @@ func FormatClaudeResponseInfo(claudeResponse *dto.ClaudeResponse, oaiResponse *d if oaiResponse != nil { oaiResponse.Id = claudeInfo.ResponseId oaiResponse.Created = claudeInfo.Created - oaiResponse.Model = claudeInfo.Model + // Preserve caller model: only set model from claudeInfo if oaiResponse doesn't have one + if oaiResponse.Model == "" { + oaiResponse.Model = claudeInfo.Model + } } return true } @@ -816,7 +825,8 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud } helper.ClaudeChunkData(c, claudeResponse, data) } else if info.RelayFormat == types.RelayFormatOpenAI { - response := StreamResponseClaude2OpenAI(&claudeResponse) + callerModel := relaycommon.GetCallerModelName(c, info) + response := StreamResponseClaude2OpenAI(&claudeResponse, callerModel) if !FormatClaudeResponseInfo(&claudeResponse, response, claudeInfo) { return nil @@ -858,7 +868,8 @@ func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, clau } else if info.RelayFormat == types.RelayFormatOpenAI { if info.ShouldIncludeUsage { openAIUsage := buildOpenAIStyleUsageFromClaudeUsage(claudeInfo.Usage) - response := helper.GenerateFinalUsageResponse(claudeInfo.ResponseId, claudeInfo.Created, info.UpstreamModelName, openAIUsage) + callerModel := relaycommon.GetCallerModelName(c, info) + response := helper.GenerateFinalUsageResponse(claudeInfo.ResponseId, claudeInfo.Created, callerModel, openAIUsage) err := helper.ObjectData(c, response) if err != nil { common.SysLog("send final response failed: " + err.Error()) @@ -917,7 +928,8 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud var responseData []byte switch info.RelayFormat { case types.RelayFormatOpenAI: - openaiResponse := ResponseClaude2OpenAI(&claudeResponse) + callerModel := relaycommon.GetCallerModelName(c, info) + openaiResponse := ResponseClaude2OpenAI(&claudeResponse, callerModel) openaiResponse.Usage = buildOpenAIStyleUsageFromClaudeUsage(claudeInfo.Usage) responseData, err = json.Marshal(openaiResponse) if err != nil { diff --git a/relay/channel/cloudflare/relay_cloudflare.go b/relay/channel/cloudflare/relay_cloudflare.go index a543c8fda4b2..0686b62afc51 100644 --- a/relay/channel/cloudflare/relay_cloudflare.go +++ b/relay/channel/cloudflare/relay_cloudflare.go @@ -61,7 +61,7 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res responseText += choice.Delta.GetContentString() } response.Id = id - response.Model = info.UpstreamModelName + response.Model = relaycommon.GetCallerModelName(c, info) err = helper.ObjectData(c, response) if isFirst { isFirst = false @@ -77,7 +77,7 @@ func cfStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res } usage := service.ResponseText2Usage(c, responseText, info.UpstreamModelName, info.GetEstimatePromptTokens()) if info.ShouldIncludeUsage { - response := helper.GenerateFinalUsageResponse(id, info.StartTime.Unix(), info.UpstreamModelName, *usage) + response := helper.GenerateFinalUsageResponse(id, info.StartTime.Unix(), relaycommon.GetCallerModelName(c, info), *usage) err := helper.ObjectData(c, response) if err != nil { logger.LogError(c, "error_rendering_final_usage_response: "+err.Error()) @@ -101,7 +101,7 @@ func cfHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) if err != nil { return types.NewError(err, types.ErrorCodeBadResponseBody), nil } - response.Model = info.UpstreamModelName + response.Model = relaycommon.GetCallerModelName(c, info) var responseText string for _, choice := range response.Choices { responseText += choice.Message.StringContent() diff --git a/relay/channel/cohere/relay-cohere.go b/relay/channel/cohere/relay-cohere.go index c205e1063363..919f7b272781 100644 --- a/relay/channel/cohere/relay-cohere.go +++ b/relay/channel/cohere/relay-cohere.go @@ -128,7 +128,7 @@ func cohereStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http openaiResp.Id = responseId openaiResp.Created = createdTime openaiResp.Object = "chat.completion.chunk" - openaiResp.Model = info.UpstreamModelName + openaiResp.Model = relaycommon.GetCallerModelName(c, info) if cohereResp.IsFinished { finishReason := stopReasonCohere2OpenAI(cohereResp.FinishReason) openaiResp.Choices = []dto.ChatCompletionsStreamResponseChoice{ @@ -193,7 +193,7 @@ func cohereHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo openaiResp.Id = cohereResp.ResponseId openaiResp.Created = createdTime openaiResp.Object = "chat.completion" - openaiResp.Model = info.UpstreamModelName + openaiResp.Model = relaycommon.GetCallerModelName(c, info) openaiResp.Usage = usage openaiResp.Choices = []dto.OpenAITextResponseChoice{ diff --git a/relay/channel/coze/relay-coze.go b/relay/channel/coze/relay-coze.go index 69ebd8a684c2..6f48d07ee2de 100644 --- a/relay/channel/coze/relay-coze.go +++ b/relay/channel/coze/relay-coze.go @@ -55,7 +55,7 @@ func cozeChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Res // convert coze response to openai response var response dto.TextResponse var cozeResponse CozeChatDetailResponse - response.Model = info.UpstreamModelName + response.Model = relaycommon.GetCallerModelName(c, info) err = json.Unmarshal(responseBody, &cozeResponse) if err != nil { return nil, types.NewError(err, types.ErrorCodeBadResponseBody) @@ -165,7 +165,7 @@ func handleCozeEvent(c *gin.Context, event string, data string, responseText *st usage.TotalTokens = chatData.Usage.TokenCount finishReason := "stop" - stopResponse := helper.GenerateStopResponse(id, common.GetTimestamp(), info.UpstreamModelName, finishReason) + stopResponse := helper.GenerateStopResponse(id, common.GetTimestamp(), relaycommon.GetCallerModelName(c, info), finishReason) helper.ObjectData(c, stopResponse) case "conversation.message.delta": @@ -190,7 +190,7 @@ func handleCozeEvent(c *gin.Context, event string, data string, responseText *st Id: id, Object: "chat.completion.chunk", Created: common.GetTimestamp(), - Model: info.UpstreamModelName, + Model: relaycommon.GetCallerModelName(c, info), } choice := dto.ChatCompletionsStreamResponseChoice{ diff --git a/relay/channel/gemini/relay-gemini.go b/relay/channel/gemini/relay-gemini.go index 355c75d71b7c..d284e1c3845a 100644 --- a/relay/channel/gemini/relay-gemini.go +++ b/relay/channel/gemini/relay-gemini.go @@ -1338,7 +1338,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * response.Id = id response.Created = createAt - response.Model = info.UpstreamModelName + response.Model = relaycommon.GetCallerModelName(c, info) for choiceIdx := range response.Choices { choiceKey := response.Choices[choiceIdx].Index for toolIdx := range response.Choices[choiceIdx].Delta.ToolCalls { @@ -1365,7 +1365,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * logger.LogDebug(c, fmt.Sprintf("info.SendResponseCount = %d", info.SendResponseCount)) if info.SendResponseCount == 0 { // send first response - emptyResponse := helper.GenerateStartEmptyResponse(id, createAt, info.UpstreamModelName, nil) + emptyResponse := helper.GenerateStartEmptyResponse(id, createAt, relaycommon.GetCallerModelName(c, info), nil) if response.IsToolCall() { if len(emptyResponse.Choices) > 0 && len(response.Choices) > 0 { toolCalls := response.Choices[0].Delta.ToolCalls @@ -1399,7 +1399,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * logger.LogError(c, err.Error()) } if isStop { - _ = handleStream(c, info, helper.GenerateStopResponse(id, createAt, info.UpstreamModelName, finishReason)) + _ = handleStream(c, info, helper.GenerateStopResponse(id, createAt, relaycommon.GetCallerModelName(c, info), finishReason)) } return true }) @@ -1408,7 +1408,7 @@ func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp * return usage, err } - response := helper.GenerateFinalUsageResponse(id, createAt, info.UpstreamModelName, *usage) + response := helper.GenerateFinalUsageResponse(id, createAt, relaycommon.GetCallerModelName(c, info), *usage) handleErr := handleFinalStream(c, info, response) if handleErr != nil { common.SysLog("send final response failed: " + handleErr.Error()) @@ -1466,7 +1466,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R return &usage, nil } fullTextResponse := responseGeminiChat2OpenAI(c, &geminiResponse) - fullTextResponse.Model = info.UpstreamModelName + fullTextResponse.Model = relaycommon.GetCallerModelName(c, info) usage := buildUsageFromGeminiMetadata(geminiResponse.UsageMetadata, info.GetEstimatePromptTokens()) fullTextResponse.Usage = usage @@ -1510,7 +1510,7 @@ func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h openAIResponse := dto.OpenAIEmbeddingResponse{ Object: "list", Data: make([]dto.OpenAIEmbeddingResponseItem, 0, len(geminiResponse.Embeddings)), - Model: info.UpstreamModelName, + Model: relaycommon.GetCallerModelName(c, info), } for i, embedding := range geminiResponse.Embeddings { diff --git a/relay/channel/ollama/relay-ollama.go b/relay/channel/ollama/relay-ollama.go index 975c244c3582..21816f1f53ae 100644 --- a/relay/channel/ollama/relay-ollama.go +++ b/relay/channel/ollama/relay-ollama.go @@ -272,7 +272,7 @@ func ollamaEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h data = append(data, dto.OpenAIEmbeddingResponseItem{Index: i, Object: "embedding", Embedding: emb}) } usage := &dto.Usage{PromptTokens: oResp.PromptEvalCount, CompletionTokens: 0, TotalTokens: oResp.PromptEvalCount} - embResp := &dto.OpenAIEmbeddingResponse{Object: "list", Data: data, Model: info.UpstreamModelName, Usage: *usage} + embResp := &dto.OpenAIEmbeddingResponse{Object: "list", Data: data, Model: relaycommon.GetCallerModelName(c, info), Usage: *usage} out, _ := common.Marshal(embResp) service.IOCopyBytesGracefully(c, resp, out) return usage, nil diff --git a/relay/channel/ollama/stream.go b/relay/channel/ollama/stream.go index 43e024deafd7..c99966a23244 100644 --- a/relay/channel/ollama/stream.go +++ b/relay/channel/ollama/stream.go @@ -72,7 +72,7 @@ func ollamaStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http helper.SetEventStreamHeaders(c) scanner := bufio.NewScanner(resp.Body) usage := &dto.Usage{} - var model = info.UpstreamModelName + var model = relaycommon.GetCallerModelName(c, info) var responseId = common.GetUUID() var created = time.Now().Unix() var toolCallIndex int @@ -93,7 +93,7 @@ func ollamaStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http return usage, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } if chunk.Model != "" { - model = chunk.Model + // Don't override caller model with upstream model } created = toUnix(chunk.CreatedAt) @@ -259,9 +259,9 @@ func ollamaChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R } } - model := lastChunk.Model + model := relaycommon.GetCallerModelName(c, info) if model == "" { - model = info.UpstreamModelName + model = lastChunk.Model } created := toUnix(lastChunk.CreatedAt) usage := &dto.Usage{PromptTokens: lastChunk.PromptEvalCount, CompletionTokens: lastChunk.EvalCount, TotalTokens: lastChunk.PromptEvalCount + lastChunk.EvalCount} diff --git a/relay/channel/openai/chat_via_responses.go b/relay/channel/openai/chat_via_responses.go index 2c0752275daa..2c45c17c0a23 100644 --- a/relay/channel/openai/chat_via_responses.go +++ b/relay/channel/openai/chat_via_responses.go @@ -65,6 +65,15 @@ func OaiResponsesToChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } + // Override model to caller model + callerModel := relaycommon.GetCallerModelName(c, info) + if callerModel != "" { + chatResp.Model = callerModel + } + if err != nil { + return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + if usage == nil || usage.TotalTokens == 0 { text := service.ExtractOutputTextFromResponses(&responsesResp) usage = service.ResponseText2Usage(c, text, info.UpstreamModelName, info.GetEstimatePromptTokens()) @@ -99,7 +108,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo responseId := helper.GetResponseID(c) createAt := time.Now().Unix() - model := info.UpstreamModelName + model := relaycommon.GetCallerModelName(c, info) var ( usage = &dto.Usage{} @@ -312,9 +321,6 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo switch streamResp.Type { case "response.created": if streamResp.Response != nil { - if streamResp.Response.Model != "" { - model = streamResp.Response.Model - } if streamResp.Response.CreatedAt != 0 { createAt = int64(streamResp.Response.CreatedAt) } @@ -444,9 +450,6 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo case "response.completed": if streamResp.Response != nil { - if streamResp.Response.Model != "" { - model = streamResp.Response.Model - } if streamResp.Response.CreatedAt != 0 { createAt = int64(streamResp.Response.CreatedAt) } diff --git a/relay/channel/openai/helper.go b/relay/channel/openai/helper.go index 08811a77205a..2e29277ad7b5 100644 --- a/relay/channel/openai/helper.go +++ b/relay/channel/openai/helper.go @@ -179,7 +179,7 @@ func handleLastResponse(lastStreamData string, responseId *string, createAt *int *responseId = lastStreamResponse.Id *createAt = lastStreamResponse.Created *systemFingerprint = lastStreamResponse.GetSystemFingerprint() - *model = lastStreamResponse.Model + // Don't overwrite model with upstream model; caller model is set upstream if service.ValidUsage(lastStreamResponse.Usage) { *containStreamUsage = true diff --git a/relay/channel/openai/relay-openai.go b/relay/channel/openai/relay-openai.go index a85751844c0b..bef12ebde59c 100644 --- a/relay/channel/openai/relay-openai.go +++ b/relay/channel/openai/relay-openai.go @@ -27,7 +27,14 @@ func sendStreamData(c *gin.Context, info *relaycommon.RelayInfo, data string, fo return nil } + // Patch top-level model to caller model for raw passthrough path if !forceFormat && !thinkToContent { + callerModel := relaycommon.GetCallerModelName(c, info) + patched, changed, err := relaycommon.PatchTopLevelModelRaw(common.StringToByteSlice(data), callerModel) + if err == nil && changed { + relaycommon.MarkResponseBodyRewritten(c) + return helper.StringData(c, string(patched)) + } return helper.StringData(c, data) } @@ -111,7 +118,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re defer service.CloseResponseBodyGracefully(resp) - model := info.UpstreamModelName + model := relaycommon.GetCallerModelName(c, info) var responseId string var createAt int64 = 0 var systemFingerprint string @@ -123,8 +130,10 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re var lastStreamData string var secondLastStreamData string // 存储倒数第二个stream data,用于音频模型 - // 检查是否为音频模型 - isAudioModel := strings.Contains(strings.ToLower(model), "audio") + // Audio model detection uses upstream model name, not caller model. + // When model mapping is active the caller name may not contain "audio" + // even though the upstream model is an audio model. + isAudioModel := strings.Contains(strings.ToLower(info.UpstreamModelName), "audio") helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { if lastStreamData != "" { @@ -294,6 +303,15 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo responseBody = geminiRespStr } + callerModel := relaycommon.GetCallerModelName(c, info) + if callerModel != "" { + patched, changed, err := relaycommon.PatchTopLevelModelRaw(responseBody, callerModel) + if err == nil && changed { + relaycommon.MarkResponseBodyRewritten(c) + responseBody = patched + } + } + service.IOCopyBytesGracefully(c, resp, responseBody) return &simpleResponse.Usage, nil diff --git a/relay/channel/openai/relay_openai_test.go b/relay/channel/openai/relay_openai_test.go new file mode 100644 index 000000000000..ef2b7a5a398d --- /dev/null +++ b/relay/channel/openai/relay_openai_test.go @@ -0,0 +1,200 @@ +package openai + +import ( + "bytes" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/dto" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/types" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +// buildAudioSSE constructs an SSE stream body that mimics OpenAI audio model +// streaming behavior: usage info appears in the second-to-last data chunk. +// +// Layout: chunk1 (content) -> chunk2 (content + usage) -> chunk3 (finish) -> [DONE] +// After processing, secondLastStreamData = chunk2 (carries usage). +func buildAudioSSE(upstreamModel string, usage dto.Usage) []byte { + var b bytes.Buffer + + // Chunk 1: content delta + b.WriteString(fmt.Sprintf(`data: {"id":"chatcmpl-test","object":"chat.completion.chunk","model":"%s","choices":[{"index":0,"delta":{"content":"Hello"}}]}`, upstreamModel)) + b.WriteString("\n\n") + + // Chunk 2: content + usage (second-to-last -> becomes secondLastStreamData) + b.WriteString(fmt.Sprintf(`data: {"id":"chatcmpl-test","object":"chat.completion.chunk","model":"%s","choices":[{"index":0,"delta":{"content":" world"}}],"usage":{"prompt_tokens":%d,"completion_tokens":%d,"total_tokens":%d}}`, + upstreamModel, usage.PromptTokens, usage.CompletionTokens, usage.TotalTokens)) + b.WriteString("\n\n") + + // Chunk 3: finish (last -> becomes lastStreamData, no usage) + b.WriteString(fmt.Sprintf(`data: {"id":"chatcmpl-test","object":"chat.completion.chunk","model":"%s","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}`, upstreamModel)) + b.WriteString("\n\n") + + b.WriteString("data: [DONE]\n\n") + return b.Bytes() +} + +type nopCloser struct{ io.Reader } + +func (nopCloser) Close() error { return nil } + +// TestOaiStreamHandler_AudioModelDetection_MappedModel verifies that when the +// caller model name does not contain "audio" but the upstream model is an audio +// model (model_mapping scenario), the handler still extracts usage from the +// second-to-last SSE chunk instead of falling back to text-based estimation. +func TestOaiStreamHandler_AudioModelDetection_MappedModel(t *testing.T) { + gin.SetMode(gin.TestMode) + + oldStreamingTimeout := constant.StreamingTimeout + constant.StreamingTimeout = 300 + t.Cleanup(func() { constant.StreamingTimeout = oldStreamingTimeout }) + + callerModel := "my-voice-bot" + upstreamModel := "gpt-4o-audio-preview" + + upstreamUsage := dto.Usage{ + PromptTokens: 100, + CompletionTokens: 50, + TotalTokens: 150, + } + + sseBody := buildAudioSSE(upstreamModel, upstreamUsage) + resp := &http.Response{ + Body: nopCloser{strings.NewReader(string(sseBody))}, + Header: make(http.Header), + } + resp.Header.Set("Content-Type", "text/event-stream") + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + + info := &relaycommon.RelayInfo{ + RelayFormat: types.RelayFormatOpenAI, + OriginModelName: callerModel, + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: upstreamModel, + }, + RelayMode: relayconstant.RelayModeChatCompletions, + } + info.SetEstimatePromptTokens(10) + + usage, apiErr := OaiStreamHandler(c, info, resp) + require.Nil(t, apiErr, "OaiStreamHandler should not return an error") + require.NotNil(t, usage, "usage should not be nil") + + require.Equal(t, 100, usage.PromptTokens, + "PromptTokens should match upstream audio usage (not estimated)") + require.Equal(t, 50, usage.CompletionTokens, + "CompletionTokens should match upstream audio usage (not estimated)") + require.Equal(t, 150, usage.TotalTokens, + "TotalTokens should match upstream audio usage (not estimated)") +} + +// TestOaiStreamHandler_AudioModelDetection_CallerModelContainsAudio verifies +// that when the caller model itself contains "audio", usage is also extracted +// correctly. This serves as a control case ensuring both paths converge. +func TestOaiStreamHandler_AudioModelDetection_CallerModelContainsAudio(t *testing.T) { + gin.SetMode(gin.TestMode) + + oldStreamingTimeout := constant.StreamingTimeout + constant.StreamingTimeout = 300 + t.Cleanup(func() { constant.StreamingTimeout = oldStreamingTimeout }) + + callerModel := "gpt-4o-audio-preview" + upstreamModel := "gpt-4o-audio-preview" + + upstreamUsage := dto.Usage{ + PromptTokens: 100, + CompletionTokens: 50, + TotalTokens: 150, + } + + sseBody := buildAudioSSE(upstreamModel, upstreamUsage) + resp := &http.Response{ + Body: nopCloser{strings.NewReader(string(sseBody))}, + Header: make(http.Header), + } + resp.Header.Set("Content-Type", "text/event-stream") + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + + info := &relaycommon.RelayInfo{ + RelayFormat: types.RelayFormatOpenAI, + OriginModelName: callerModel, + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: upstreamModel, + }, + RelayMode: relayconstant.RelayModeChatCompletions, + } + info.SetEstimatePromptTokens(10) + + usage, apiErr := OaiStreamHandler(c, info, resp) + require.Nil(t, apiErr) + require.NotNil(t, usage) + + require.Equal(t, 100, usage.PromptTokens) + require.Equal(t, 50, usage.CompletionTokens) + require.Equal(t, 150, usage.TotalTokens) +} + +// TestOaiStreamHandler_NonAudioModel_SkipsSecondLastUsage verifies that for a +// non-audio model, the second-to-last usage is NOT extracted (the fallback +// text-based estimation path is used instead). +func TestOaiStreamHandler_NonAudioModel_SkipsSecondLastUsage(t *testing.T) { + gin.SetMode(gin.TestMode) + + oldStreamingTimeout := constant.StreamingTimeout + constant.StreamingTimeout = 300 + t.Cleanup(func() { constant.StreamingTimeout = oldStreamingTimeout }) + + callerModel := "gpt-4o" + upstreamModel := "gpt-4o" + + upstreamUsage := dto.Usage{ + PromptTokens: 100, + CompletionTokens: 50, + TotalTokens: 150, + } + + sseBody := buildAudioSSE(upstreamModel, upstreamUsage) + resp := &http.Response{ + Body: nopCloser{strings.NewReader(string(sseBody))}, + Header: make(http.Header), + } + resp.Header.Set("Content-Type", "text/event-stream") + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + + info := &relaycommon.RelayInfo{ + RelayFormat: types.RelayFormatOpenAI, + OriginModelName: callerModel, + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: upstreamModel, + }, + RelayMode: relayconstant.RelayModeChatCompletions, + } + info.SetEstimatePromptTokens(10) + + usage, apiErr := OaiStreamHandler(c, info, resp) + require.Nil(t, apiErr) + require.NotNil(t, usage) + + // Non-audio model should NOT use the second-to-last chunk's usage; + // it falls through to text-based estimation or last-chunk usage. + require.Equal(t, 10, usage.PromptTokens, + "Non-audio model should use estimated prompt tokens, not second-to-last usage") +} diff --git a/relay/channel/openai/relay_responses.go b/relay/channel/openai/relay_responses.go index 2665b8d027e9..cfa605098b7a 100644 --- a/relay/channel/openai/relay_responses.go +++ b/relay/channel/openai/relay_responses.go @@ -41,6 +41,14 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http } // 写入新的 response body + callerModel := relaycommon.GetCallerModelName(c, info) + if callerModel != "" { + patched, changed, err := relaycommon.PatchTopLevelModelRaw(responseBody, callerModel) + if err == nil && changed { + relaycommon.MarkResponseBodyRewritten(c) + responseBody = patched + } + } service.IOCopyBytesGracefully(c, resp, responseBody) // compute usage @@ -79,8 +87,19 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp var usage = &dto.Usage{} var responseTextBuilder strings.Builder + callerModel := relaycommon.GetCallerModelName(c, info) + helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { + // Patch response.model in raw SSE data before sending + if callerModel != "" { + patched, changed, err := relaycommon.PatchResponsesEventModelRaw(common.StringToByteSlice(data), callerModel) + if err == nil && changed { + relaycommon.MarkResponseBodyRewritten(c) + data = string(patched) + } + } + // 检查当前数据是否包含 completed 状态和 usage 信息 var streamResponse dto.ResponsesStreamResponse if err := common.UnmarshalJsonStr(data, &streamResponse); err != nil { diff --git a/relay/channel/xai/text.go b/relay/channel/xai/text.go index f9a8ee2e6f96..ced35f5be116 100644 --- a/relay/channel/xai/text.go +++ b/relay/channel/xai/text.go @@ -16,18 +16,22 @@ import ( "github.com/gin-gonic/gin" ) -func streamResponseXAI2OpenAI(xAIResp *dto.ChatCompletionsStreamResponse, usage *dto.Usage) *dto.ChatCompletionsStreamResponse { +func streamResponseXAI2OpenAI(xAIResp *dto.ChatCompletionsStreamResponse, usage *dto.Usage, callerModel string) *dto.ChatCompletionsStreamResponse { if xAIResp == nil { return nil } if xAIResp.Usage != nil { xAIResp.Usage.CompletionTokens = usage.CompletionTokens } + model := callerModel + if model == "" { + model = xAIResp.Model + } openAIResp := &dto.ChatCompletionsStreamResponse{ Id: xAIResp.Id, Object: xAIResp.Object, Created: xAIResp.Created, - Model: xAIResp.Model, + Model: model, Choices: xAIResp.Choices, Usage: xAIResp.Usage, } @@ -41,6 +45,8 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re var toolCount int var containStreamUsage bool + callerModel := relaycommon.GetCallerModelName(c, info) + helper.SetEventStreamHeaders(c) helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { @@ -59,7 +65,7 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re usage.CompletionTokens = usage.TotalTokens - usage.PromptTokens } - openaiResponse := streamResponseXAI2OpenAI(xAIResp, usage) + openaiResponse := streamResponseXAI2OpenAI(xAIResp, usage, callerModel) _ = openai.ProcessStreamResponse(*openaiResponse, &responseTextBuilder, &toolCount) if err := helper.ObjectData(c, openaiResponse); err != nil { common.SysLog(err.Error()) @@ -95,6 +101,10 @@ func xAIHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response } // new body + callerModel := relaycommon.GetCallerModelName(c, info) + if callerModel != "" { + xaiResponse.Model = callerModel + } encodeJson, err := common.Marshal(xaiResponse) if err != nil { return nil, types.NewError(err, types.ErrorCodeBadResponseBody) diff --git a/relay/common/response_model_alias.go b/relay/common/response_model_alias.go new file mode 100644 index 000000000000..1a026ee01c03 --- /dev/null +++ b/relay/common/response_model_alias.go @@ -0,0 +1,83 @@ +package common + +import ( + "github.com/QuantumNous/new-api/constant" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" +) + +// GetCallerModelName returns the model name that should be shown to the API caller. +// Priority: Gin context "original_model" -> info.OriginModelName -> info.UpstreamModelName. +func GetCallerModelName(c *gin.Context, info *RelayInfo) string { + if c != nil { + if v := c.GetString(string(constant.ContextKeyOriginalModel)); v != "" { + return v + } + } + if info != nil && info.OriginModelName != "" { + return info.OriginModelName + } + if info != nil { + return info.UpstreamModelName + } + return "" +} + +// PatchTopLevelModelRaw rewrites the top-level "model" field in a raw JSON payload. +// Returns (patched data, whether a change was made, error). +func PatchTopLevelModelRaw(data []byte, callerModel string) ([]byte, bool, error) { + if len(data) == 0 || callerModel == "" { + return data, false, nil + } + + existing := gjson.GetBytes(data, "model") + if !existing.Exists() || existing.String() == callerModel { + return data, false, nil + } + + patched, err := sjson.SetBytes(data, "model", callerModel) + if err != nil { + return data, false, err + } + return patched, true, nil +} + +// PatchResponsesEventModelRaw rewrites response.model in a Responses API SSE event payload. +// It applies to any event that contains a response.model field, regardless of event type. +// Returns (patched data, whether a change was made, error). +func PatchResponsesEventModelRaw(data []byte, callerModel string) ([]byte, bool, error) { + if len(data) == 0 || callerModel == "" { + return data, false, nil + } + + existing := gjson.GetBytes(data, "response.model") + if !existing.Exists() || existing.String() == callerModel { + return data, false, nil + } + + patched, err := sjson.SetBytes(data, "response.model", callerModel) + if err != nil { + return data, false, err + } + return patched, true, nil +} + +const responseBodyRewrittenKey = "response_body_rewritten" + +// MarkResponseBodyRewritten sets a context flag indicating the response body was modified. +func MarkResponseBodyRewritten(c *gin.Context) { + if c != nil { + c.Set(responseBodyRewrittenKey, true) + } +} + +// IsResponseBodyRewritten returns true if the response body was rewritten. +func IsResponseBodyRewritten(c *gin.Context) bool { + if c == nil { + return false + } + v, _ := c.Get(responseBodyRewrittenKey) + b, ok := v.(bool) + return ok && b +} diff --git a/service/http.go b/service/http.go index d9818cfe158d..a4c0ec1a424f 100644 --- a/service/http.go +++ b/service/http.go @@ -9,6 +9,7 @@ import ( "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/logger" + relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/gin-gonic/gin" ) @@ -38,6 +39,13 @@ func ShouldCopyUpstreamHeader(c *gin.Context, k string, v []string) bool { } return false } + // When the response body has been rewritten, strip integrity headers + // that would be invalid against the modified content. + if relaycommon.IsResponseBodyRewritten(c) { + if strings.EqualFold(k, "ETag") || strings.EqualFold(k, "Content-MD5") { + return false + } + } return true } diff --git a/web/classic/src/components/settings/OtherSetting.jsx b/web/classic/src/components/settings/OtherSetting.jsx index 56049093b2fa..87e3ee183b7e 100644 --- a/web/classic/src/components/settings/OtherSetting.jsx +++ b/web/classic/src/components/settings/OtherSetting.jsx @@ -239,39 +239,47 @@ const OtherSetting = () => { // Option 1: Use a public CORS proxy service // const proxyUrl = 'https://cors-anywhere.herokuapp.com/'; // const res = await API.get( - // `${proxyUrl}https://api.github.com/repos/Calcium-Ion/new-api/releases/latest`, + // `${proxyUrl}https://api.github.com/repos/Wischoicer-Xian/new-api/releases/latest`, // ); // Option 2: Use the JSON proxy approach which often works better with GitHub API - const res = await fetch( - 'https://api.github.com/repos/Calcium-Ion/new-api/releases/latest', + const response = await fetch( + 'https://api.github.com/repos/QuantumNous/new-api/releases/latest', { headers: { Accept: 'application/json', 'Content-Type': 'application/json', - // Adding User-Agent which is often required by GitHub API 'User-Agent': 'new-api-update-checker', }, }, - ).then((response) => response.json()); + ); - // Option 3: Use a local proxy endpoint - // Create a cached version of the response to avoid frequent GitHub API calls - // const res = await API.get('/api/status/github-latest-release'); + if (!response.ok) { + throw new Error( + response.status === 404 + ? '当前仓库暂无发布版本,无法检查更新' + : `GitHub API 请求失败 (${response.status})`, + ); + } + const res = await response.json(); const { tag_name, body } = res; + if (!tag_name) { + throw new Error('返回的发布信息格式异常'); + } + if (tag_name === statusState?.status?.version) { showSuccess(`已是最新版本:${tag_name}`); } else { setUpdateData({ tag_name: tag_name, - content: marked.parse(body), + content: marked.parse(body || ''), }); setShowUpdateModal(true); } } catch (error) { console.error('Failed to check for updates:', error); - showError('检查更新失败,请稍后再试'); + showError(error.message || '检查更新失败,请稍后再试'); } finally { setLoadingInput((loadingInput) => ({ ...loadingInput, @@ -343,7 +351,7 @@ const OtherSetting = () => { // Function to open GitHub release page const openGitHubRelease = () => { window.open( - `https://github.com/Calcium-Ion/new-api/releases/tag/${updateData.tag_name}`, + `https://github.com/QuantumNous/new-api/releases/tag/${updateData.tag_name}`, '_blank', ); }; diff --git a/web/classic/src/components/table/usage-logs/UsageLogsColumnDefs.jsx b/web/classic/src/components/table/usage-logs/UsageLogsColumnDefs.jsx index 07e6fbb9896d..2551f6ed6af4 100644 --- a/web/classic/src/components/table/usage-logs/UsageLogsColumnDefs.jsx +++ b/web/classic/src/components/table/usage-logs/UsageLogsColumnDefs.jsx @@ -269,9 +269,10 @@ function renderBillingTag(record, t) { return null; } -function renderModelName(record, copyText, t) { +function renderModelName(record, copyText, t, isAdminUser = true) { let other = getLogOther(record.other); let modelMapped = + isAdminUser && other?.is_model_mapped && other?.upstream_model_name && other?.upstream_model_name !== ''; @@ -688,7 +689,7 @@ export const getLogsColumns = ({ record.type === 2 || record.type === 5 || record.type === 6 ? ( - <>{renderModelName(record, copyText, t)} + <>{renderModelName(record, copyText, t, isAdminUser)} ) : ( <> ); diff --git a/web/classic/src/hooks/usage-logs/useUsageLogsData.jsx b/web/classic/src/hooks/usage-logs/useUsageLogsData.jsx index 78975dd634f7..7301b90bae1b 100644 --- a/web/classic/src/hooks/usage-logs/useUsageLogsData.jsx +++ b/web/classic/src/hooks/usage-logs/useUsageLogsData.jsx @@ -449,6 +449,7 @@ export const useLogsData = () => { } if (logs[i].type === 2) { let modelMapped = + isAdminUser && other?.is_model_mapped && other?.upstream_model_name && other?.upstream_model_name !== ''; diff --git a/web/default/src/features/system-settings/maintenance/update-checker-section.tsx b/web/default/src/features/system-settings/maintenance/update-checker-section.tsx index ffba90d38d62..058c4783ae56 100644 --- a/web/default/src/features/system-settings/maintenance/update-checker-section.tsx +++ b/web/default/src/features/system-settings/maintenance/update-checker-section.tsx @@ -62,7 +62,7 @@ export function UpdateCheckerSection({ setChecking(true) try { const response = await fetch( - 'https://api.github.com/repos/Calcium-Ion/new-api/releases/latest', + 'https://api.github.com/repos/QuantumNous/new-api/releases/latest', { headers: { Accept: 'application/vnd.github+json', @@ -72,7 +72,13 @@ export function UpdateCheckerSection({ ) if (!response.ok) { - throw new Error(t('Failed to contact GitHub releases API')) + throw new Error( + response.status === 404 + ? t('No release found for this repository yet.') + : t('Failed to contact GitHub releases API ({{status}})', { + status: response.status, + }), + ) } const data = (await response.json()) as ReleaseInfo diff --git a/web/default/src/features/usage-logs/components/columns/common-logs-columns.tsx b/web/default/src/features/usage-logs/components/columns/common-logs-columns.tsx index d512f79fc25d..5c8e48171aaa 100644 --- a/web/default/src/features/usage-logs/components/columns/common-logs-columns.tsx +++ b/web/default/src/features/usage-logs/components/columns/common-logs-columns.tsx @@ -527,7 +527,7 @@ export function useCommonLogsColumns(isAdmin: boolean): ColumnDef[] { const log = row.original if (!isDisplayableLogType(log.type)) return null - const modelInfo = formatModelName(log) + const modelInfo = formatModelName(log, isAdmin) return (
diff --git a/web/default/src/features/usage-logs/components/dialogs/details-dialog.tsx b/web/default/src/features/usage-logs/components/dialogs/details-dialog.tsx index 2785a5528795..a23500a0ec35 100644 --- a/web/default/src/features/usage-logs/components/dialogs/details-dialog.tsx +++ b/web/default/src/features/usage-logs/components/dialogs/details-dialog.tsx @@ -829,8 +829,8 @@ export function DetailsDialog(props: DetailsDialogProps) { /> )} - {/* Model mapping */} - {other?.is_model_mapped && other?.upstream_model_name && ( + {/* Model mapping (admin only) */} + {props.isAdmin && other?.is_model_mapped && other?.upstream_model_name && (