Skip to content
Open
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
25 changes: 21 additions & 4 deletions relay/channel/xai/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package xai

import (
"errors"
"fmt"
"io"
"net/http"
"strings"
Expand All @@ -10,6 +11,7 @@ import (
"github.com/QuantumNous/new-api/relay/channel"
"github.com/QuantumNous/new-api/relay/channel/openai"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/types"

"github.com/QuantumNous/new-api/relay/constant"
Expand All @@ -26,10 +28,20 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt
return nil, errors.New("not implemented")
}

func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
//TODO implement me
//panic("implement me")
return nil, errors.New("not available")
func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) {
if request == nil {
return nil, errors.New("request is nil")
}
openAIRequest, err := service.ClaudeToOpenAIRequest(*request, info)
if err != nil {
return nil, err
}
if info != nil && info.ChannelMeta != nil && info.SupportStreamOptions && info.IsStream {
openAIRequest.StreamOptions = &dto.StreamOptions{
IncludeUsage: true,
}
}
return a.ConvertOpenAIRequest(c, info, openAIRequest)
}

func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
Expand All @@ -51,6 +63,11 @@ func (a *Adaptor) Init(info *relaycommon.RelayInfo) {
}

func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
if info.RelayFormat == types.RelayFormatClaude &&
info.RelayMode != constant.RelayModeResponses &&
info.RelayMode != constant.RelayModeResponsesCompact {
return fmt.Sprintf("%s/v1/chat/completions", strings.TrimRight(info.ChannelBaseUrl, "/")), nil
}
return relaycommon.GetFullRequestURL(info.ChannelBaseUrl, info.RequestURLPath, info.ChannelType), nil
}

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

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

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

func TestConvertClaudeRequestUsesOpenAICompatibleChatRequest(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
stream := true
maxTokens := uint(128)
temperature := 0.2
info := &relaycommon.RelayInfo{
IsStream: true,
RelayFormat: types.RelayFormatClaude,
ChannelMeta: &relaycommon.ChannelMeta{
SupportStreamOptions: true,
UpstreamModelName: "grok-4.3-fast",
},
}
info.ClaudeConvertInfo = &relaycommon.ClaudeConvertInfo{
LastMessagesType: relaycommon.LastMessageTypeNone,
}

converted, err := (&Adaptor{}).ConvertClaudeRequest(c, info, &dto.ClaudeRequest{
Model: "grok-4.3-fast",
MaxTokens: &maxTokens,
Stream: &stream,
Temperature: &temperature,
Messages: []dto.ClaudeMessage{
{
Role: "user",
Content: "hello",
},
},
})

require.NoError(t, err)
openAIRequest, ok := converted.(*dto.GeneralOpenAIRequest)
require.True(t, ok)
assert.Equal(t, "grok-4.3-fast", openAIRequest.Model)
require.Len(t, openAIRequest.Messages, 1)
assert.Equal(t, "user", openAIRequest.Messages[0].Role)
assert.Equal(t, "hello", openAIRequest.Messages[0].StringContent())
require.NotNil(t, openAIRequest.Stream)
assert.True(t, *openAIRequest.Stream)
require.NotNil(t, openAIRequest.StreamOptions)
assert.True(t, openAIRequest.StreamOptions.IncludeUsage)
require.NotNil(t, openAIRequest.MaxTokens)
assert.Equal(t, maxTokens, *openAIRequest.MaxTokens)
require.NotNil(t, openAIRequest.Temperature)
assert.Equal(t, temperature, *openAIRequest.Temperature)
}

func TestGetRequestURLUsesChatCompletionsForClaudeFormat(t *testing.T) {
info := &relaycommon.RelayInfo{
RelayFormat: types.RelayFormatClaude,
RelayMode: relayconstant.RelayModeChatCompletions,
RequestURLPath: "/v1/messages",
ChannelMeta: &relaycommon.ChannelMeta{
ChannelBaseUrl: "https://api.x.ai",
},
}

requestURL, err := (&Adaptor{}).GetRequestURL(info)

require.NoError(t, err)
assert.Equal(t, "https://api.x.ai/v1/chat/completions", requestURL)
}

func TestXAIHandlerConvertsClaudeFormatResponse(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
info := &relaycommon.RelayInfo{
RelayFormat: types.RelayFormatClaude,
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "grok-4.3-fast",
},
ClaudeConvertInfo: &relaycommon.ClaudeConvertInfo{},
}
responseBody := `{"id":"chatcmpl-1","object":"chat.completion","created":1,"model":"grok-4.3-fast","choices":[{"index":0,"message":{"role":"assistant","content":"pong"},"finish_reason":"stop"}],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}`
resp := &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(responseBody)),
}

usage, newAPIError := xAIHandler(c, info, resp)

require.Nil(t, newAPIError)
require.NotNil(t, usage)
assert.Equal(t, 3, usage.PromptTokens)
assert.Equal(t, 2, usage.CompletionTokens)
assert.Equal(t, http.StatusOK, recorder.Code)
assert.Contains(t, recorder.Body.String(), `"type":"message"`)
assert.Contains(t, recorder.Body.String(), `"content":[{"type":"text","text":"pong"}]`)
assert.Contains(t, recorder.Body.String(), `"input_tokens":3`)
}

func TestXAIStreamHandlerConvertsClaudeFormatResponse(t *testing.T) {
gin.SetMode(gin.TestMode)
originalStreamingTimeout := constant.StreamingTimeout
constant.StreamingTimeout = 30
t.Cleanup(func() {
constant.StreamingTimeout = originalStreamingTimeout
})
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
info := &relaycommon.RelayInfo{
RelayFormat: types.RelayFormatClaude,
IsStream: true,
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "grok-4.3-fast",
},
ClaudeConvertInfo: &relaycommon.ClaudeConvertInfo{
LastMessagesType: relaycommon.LastMessageTypeNone,
},
}
info.SetEstimatePromptTokens(3)
responseBody := strings.Join([]string{
`data: {"id":"chatcmpl-1","object":"chat.completion.chunk","created":1,"model":"grok-4.3-fast","choices":[{"index":0,"delta":{"content":"pong"},"finish_reason":null}],"usage":null}`,
"",
`data: {"id":"chatcmpl-1","object":"chat.completion.chunk","created":1,"model":"grok-4.3-fast","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":3,"completion_tokens":0,"total_tokens":5}}`,
"",
"data: [DONE]",
"",
}, "\n")
resp := &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(responseBody)),
}

usage, newAPIError := xAIStreamHandler(c, info, resp)

require.Nil(t, newAPIError)
require.NotNil(t, usage)
assert.Equal(t, 3, usage.PromptTokens)
assert.Equal(t, 2, usage.CompletionTokens)
body := recorder.Body.String()
assert.Contains(t, body, "event: message_start")
assert.Contains(t, body, "event: content_block_delta")
assert.Contains(t, body, "event: message_stop")
assert.Contains(t, body, `"text":"pong"`)
assert.Contains(t, body, `"input_tokens":3`)
assert.NotContains(t, body, "[DONE]")
}
47 changes: 42 additions & 5 deletions relay/channel/xai/text.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
var responseTextBuilder strings.Builder
var toolCount int
var containStreamUsage bool
var lastStreamData string

helper.SetEventStreamHeaders(c)

Expand All @@ -60,8 +61,23 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
}

openaiResponse := streamResponseXAI2OpenAI(xAIResp, usage)
if openaiResponse == nil {
return
}
_ = openai.ProcessStreamResponse(*openaiResponse, &responseTextBuilder, &toolCount)
if err := helper.ObjectData(c, openaiResponse); err != nil {
openaiResponseData, err := common.Marshal(openaiResponse)
if err != nil {
common.SysLog(err.Error())
sr.Error(err)
return
}
lastStreamData = string(openaiResponseData)
if info.RelayFormat == types.RelayFormatClaude {
err = openai.HandleStreamFormat(c, info, lastStreamData, false, false)
} else {
err = helper.ObjectData(c, openaiResponse)
}
if err != nil {
common.SysLog(err.Error())
sr.Error(err)
}
Expand All @@ -72,7 +88,13 @@ func xAIStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
usage.CompletionTokens += toolCount * 7
}

helper.Done(c)
if info.RelayFormat == types.RelayFormatClaude {
if lastStreamData != "" {
openai.HandleFinalResponse(c, info, lastStreamData, "", 0, info.UpstreamModelName, "", usage, containStreamUsage)
}
} else {
helper.Done(c)
}
service.CloseResponseBodyGracefully(resp)
return usage, nil
}
Expand All @@ -94,13 +116,28 @@ func xAIHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response
xaiResponse.Usage.CompletionTokenDetails.TextTokens = xaiResponse.Usage.CompletionTokens - xaiResponse.Usage.CompletionTokenDetails.ReasoningTokens
}

// new body
encodeJson, err := common.Marshal(xaiResponse)
openAIResponse := dto.OpenAITextResponse{
Id: xaiResponse.Id,
Object: xaiResponse.Object,
Created: xaiResponse.Created,
Model: xaiResponse.Model,
Choices: xaiResponse.Choices,
}
if xaiResponse.Usage != nil {
openAIResponse.Usage = *xaiResponse.Usage
}
Comment on lines +126 to +128

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟠 Major | ⚡ Quick win

Preserve nullable usage semantics when upstream omits usage.

At Line 142, returning &openAIResponse.Usage makes usage non-nil even when Line 126 is not entered (xaiResponse.Usage == nil). That turns “usage absent” into synthetic zero-token usage, which can break downstream usage/billing fallback logic.

Suggested fix
 	openAIResponse := dto.OpenAITextResponse{
 		Id:      xaiResponse.Id,
 		Object:  xaiResponse.Object,
 		Created: xaiResponse.Created,
 		Model:   xaiResponse.Model,
 		Choices: xaiResponse.Choices,
 	}
+	var usagePtr *dto.Usage
 	if xaiResponse.Usage != nil {
 		openAIResponse.Usage = *xaiResponse.Usage
+		usagePtr = &openAIResponse.Usage
 	}
@@
-	return &openAIResponse.Usage, nil
+	return usagePtr, nil
 }

Also applies to: 142-142

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@relay/channel/xai/text.go` around lines 126 - 128, The issue is that when
xaiResponse.Usage is nil, the condition at the if statement is not entered, but
then at line 142 the code returns a pointer to openAIResponse.Usage which
creates a non-nil pointer to a zero-valued struct instead of preserving the nil
semantics. To fix this, modify the code to only return the Usage pointer when it
was actually populated from the upstream response. Make the return at line 142
conditional so that it returns nil for Usage when xaiResponse.Usage was nil
(i.e., the if block was not entered), thereby preserving the nullable semantics
and preventing synthetic zero-token usage from being created when usage data is
absent from upstream.


var responseObject any = xaiResponse
if info.RelayFormat == types.RelayFormatClaude {
responseObject = service.ResponseOpenAI2Claude(&openAIResponse, info)
}

encodeJson, err := common.Marshal(responseObject)
if err != nil {
return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
}

service.IOCopyBytesGracefully(c, resp, encodeJson)

return xaiResponse.Usage, nil
return &openAIResponse.Usage, nil
}