Skip to content
Closed

1.5 #3167

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
21 changes: 9 additions & 12 deletions .github/workflows/docker-image-arm64.yml
Original file line number Diff line number Diff line change
@@ -1,9 +1,6 @@
name: Publish Docker image (Multi Registries, native amd64+arm64)

on:
push:
tags:
- '*'
workflow_dispatch:
inputs:
tag:
Expand Down Expand Up @@ -78,7 +75,7 @@ jobs:
uses: docker/metadata-action@v5
with:
images: |
calciumion/new-api
ahmczsy/new-api
# ghcr.io/${{ env.GHCR_REPOSITORY }}

- name: Build & push single-arch (to both registries)
Expand All @@ -88,8 +85,8 @@ jobs:
platforms: ${{ matrix.platform }}
push: true
tags: |
calciumion/new-api:${{ env.TAG }}-${{ matrix.arch }}
calciumion/new-api:latest-${{ matrix.arch }}
ahmczsy/new-api:${{ env.TAG }}-${{ matrix.arch }}
ahmczsy/new-api:latest-${{ matrix.arch }}
# ghcr.io/${{ env.GHCR_REPOSITORY }}:${{ env.TAG }}-${{ matrix.arch }}
# ghcr.io/${{ env.GHCR_REPOSITORY }}:latest-${{ matrix.arch }}
labels: ${{ steps.meta.outputs.labels }}
Expand Down Expand Up @@ -124,16 +121,16 @@ jobs:
- name: Create & push manifest (Docker Hub - version)
run: |
docker buildx imagetools create \
-t calciumion/new-api:${TAG} \
calciumion/new-api:${TAG}-amd64 \
calciumion/new-api:${TAG}-arm64
-t ahmczsy/new-api:${TAG} \
ahmczsy/new-api:${TAG}-amd64 \
ahmczsy/new-api:${TAG}-arm64

- name: Create & push manifest (Docker Hub - latest)
run: |
docker buildx imagetools create \
-t calciumion/new-api:latest \
calciumion/new-api:latest-amd64 \
calciumion/new-api:latest-arm64
-t ahmczsy/new-api:latest \
ahmczsy/new-api:latest-amd64 \
ahmczsy/new-api:latest-arm64

# ---- GHCR ----
# - name: Log in to GHCR
Expand Down
7 changes: 7 additions & 0 deletions dto/claude.go
Original file line number Diff line number Diff line change
Expand Up @@ -477,6 +477,13 @@ type ClaudeResponse struct {
Message *ClaudeMediaMessage `json:"message,omitempty"`
}

func (c *ClaudeResponse) ResetModel(newModel string) {
c.Model = newModel
if c.Message != nil {
c.Message.Model = newModel
}
}

// set index
func (c *ClaudeResponse) SetIndex(i int) {
c.Index = &i
Expand Down
2 changes: 2 additions & 0 deletions dto/error.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package dto

import (
"encoding/json"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/types"
Expand Down Expand Up @@ -43,6 +44,7 @@ func (e GeneralErrorResponse) TryToOpenAIError() *types.OpenAIError {
if len(e.Error) > 0 {
err := common.Unmarshal(e.Error, &openAIError)
if err == nil && openAIError.Message != "" {
openAIError.Message = strings.ReplaceAll(openAIError.Message, "Anthropic", "KernelCat")
return &openAIError
}
}
Expand Down
4 changes: 4 additions & 0 deletions dto/openai_response.go
Original file line number Diff line number Diff line change
Expand Up @@ -386,6 +386,10 @@ type ResponsesStreamResponse struct {
Part *ResponsesReasoningSummaryPart `json:"part,omitempty"`
}

func (resp *ResponsesStreamResponse) NeedResetModel() bool {
return resp.Response != nil && resp.Response.Model != ""
}

// GetOpenAIError 从动态错误类型中提取OpenAIError结构
func GetOpenAIError(errorField any) *types.OpenAIError {
if errorField == nil {
Expand Down
46 changes: 45 additions & 1 deletion relay/channel/claude/relay-claude.go
Original file line number Diff line number Diff line change
Expand Up @@ -708,6 +708,12 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
FormatClaudeResponseInfo(&claudeResponse, nil, claudeInfo)

if claudeResponse.Type == "message_start" {
newData, err := resetMessageStartData(c, info, data)
if err != nil {
common.SysLog("error resetMessageStartData stream response: " + err.Error())
return types.NewError(err, types.ErrorCodeBadResponseBody)
}
data = newData
// message_start, 获取usage
if claudeResponse.Message != nil {
info.UpstreamModelName = claudeResponse.Message.Model
Expand All @@ -721,6 +727,7 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
}
helper.ClaudeChunkData(c, claudeResponse, data)
} else if info.RelayFormat == types.RelayFormatOpenAI {
claudeResponse.ResetModel(info.OriginModelName)
response := StreamResponseClaude2OpenAI(&claudeResponse)

if !FormatClaudeResponseInfo(&claudeResponse, response, claudeInfo) {
Expand All @@ -735,6 +742,27 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
return nil
}

func resetMessageStartData(c *gin.Context, info *relaycommon.RelayInfo, data string) (string, error) {
tmpMap := map[string]any{}
if err := common.UnmarshalJsonStr(data, &tmpMap); err != nil {
return "", err
}
if tmpMap["message"] == nil {
return data, nil
}
messageMap, ok := tmpMap["message"].(map[string]any)
if !ok {
return data, nil
}
messageMap["model"] = info.OriginModelName
tmpMap["message"] = messageMap
bytes, err := common.Marshal(tmpMap)
if err != nil {
return "", err
}
return string(bytes), nil
}

func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, claudeInfo *ClaudeResponseInfo) {
if claudeInfo.Usage.PromptTokens == 0 {
//上游出错
Expand Down Expand Up @@ -809,14 +837,21 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
var responseData []byte
switch info.RelayFormat {
case types.RelayFormatOpenAI:
//返回的真实模型改为重定向模型
claudeResponse.ResetModel(info.OriginModelName)

openaiResponse := ResponseClaude2OpenAI(&claudeResponse)
openaiResponse.Usage = *claudeInfo.Usage
responseData, err = json.Marshal(openaiResponse)
if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody)
}
case types.RelayFormatClaude:
responseData = data
newData, err := resetModel(data, info.OriginModelName)
if err != nil {
return types.NewError(err, types.ErrorCodeBadResponseBody)
}
responseData = newData
}

if claudeResponse.Usage != nil && claudeResponse.Usage.ServerToolUse != nil && claudeResponse.Usage.ServerToolUse.WebSearchRequests > 0 {
Expand All @@ -827,6 +862,15 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud
return nil
}

func resetModel(data []byte, newModel string) ([]byte, error) {
tmpMap := map[string]any{}
if err := common.Unmarshal(data, &tmpMap); err != nil {
return nil, err
}
tmpMap["model"] = newModel
return common.Marshal(tmpMap)
}

func ClaudeHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.Usage, *types.NewAPIError) {
defer service.CloseResponseBodyGracefully(resp)

Expand Down
63 changes: 62 additions & 1 deletion relay/channel/openai/relay-openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,10 @@ func sendStreamData(c *gin.Context, info *relaycommon.RelayInfo, data string, fo
if data == "" {
return nil
}

newData, err := resetStreamModel(data, info.OriginModelName)
if err == nil {
data = newData
}
if !forceFormat && !thinkToContent {
return helper.StringData(c, data)
}
Expand Down Expand Up @@ -269,14 +272,21 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
bodyMap["usage"] = simpleResponse.Usage
bodyMap["model"] = info.OriginModelName
responseBody, _ = common.Marshal(bodyMap)
}
if forceFormat {
simpleResponse.Model = info.OriginModelName
responseBody, err = common.Marshal(simpleResponse)
if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
}
} else {
newBody, err := resetModel(responseBody, info.OriginModelName)
if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
}
responseBody = newBody
break
}
case types.RelayFormatClaude:
Expand All @@ -300,6 +310,57 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
return &simpleResponse.Usage, nil
}

func resetModel(data []byte, newModel string) ([]byte, error) {
tmpMap := map[string]any{}
if err := common.Unmarshal(data, &tmpMap); err != nil {
return nil, err
}
tmpMap["model"] = newModel
return common.Marshal(tmpMap)
}

func resetStreamModel(data string, newModel string) (string, error) {
tmpMap := map[string]any{}
if err := common.UnmarshalJsonStr(data, &tmpMap); err != nil {
return "", err
}
if tmpMap["model"] == "" {
return data, nil
}
tmpMap["model"] = newModel
marshal, err := common.Marshal(tmpMap)
if err != nil {
return "", err
}
return string(marshal), nil
}

func resetResponseStreamModel(data string, newModel string) (string, error) {
tmpMap := map[string]any{}
if err := common.UnmarshalJsonStr(data, &tmpMap); err != nil {
return "", err
}
if tmpMap["response"] == nil {
return data, nil
}
responseMap, ok := tmpMap["response"].(map[string]any)
if !ok {
return data, nil
}
if responseMap["model"] == "" {
return data, nil
}

responseMap["model"] = newModel
tmpMap["response"] = responseMap

marshal, err := common.Marshal(tmpMap)
if err != nil {
return "", err
}
return string(marshal), nil
}

func streamTTSResponse(c *gin.Context, resp *http.Response) {
c.Writer.WriteHeaderNow()

Expand Down
13 changes: 12 additions & 1 deletion relay/channel/openai/relay_responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,10 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
c.Set("image_generation_call_quality", responsesResponse.GetQuality())
c.Set("image_generation_call_size", responsesResponse.GetSize())
}

bytes, err := resetModel(responseBody, info.OriginModelName)
if err == nil {
responseBody = bytes
}
// 写入新的 response body
service.IOCopyBytesGracefully(c, resp, responseBody)

Expand Down Expand Up @@ -84,6 +87,14 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
// 检查当前数据是否包含 completed 状态和 usage 信息
var streamResponse dto.ResponsesStreamResponse
if err := common.UnmarshalJsonStr(data, &streamResponse); err == nil {

if streamResponse.NeedResetModel() {
newData, err := resetResponseStreamModel(data, info.OriginModelName)
if err == nil {
data = newData
}
}

sendResponsesStreamData(c, streamResponse, data)
switch streamResponse.Type {
case "response.completed":
Expand Down