Commit 6e10f9bc by CaIon

fix(relay): preserve Kimi K3 dynamic tool loading messages

Kimi K3 injects tools mid-conversation via a system message that carries
a `tools` array. `dto.Message` had no such field, so the tools were
silently dropped during the parse/re-marshal round trip and the upstream
rejected the request with `'tool_choice'='required' requires a 'tools'
field`.

- add `Message.Tools` (json.RawMessage passthrough)
- omit the `content` key only for tool-loading messages with nil content,
  as Kimi rejects `tools` next to `content`; all other messages keep
  emitting `"content": null`
- count message-level tools in token estimation
- skip tool-loading messages in channel system prompt injection and make
  the compatible handler reuse applySystemPromptIfNeeded
- add kimi-k3 to the moonshot model list

Fixes #7235
parent 3b465226
package moonshot package moonshot
var ModelList = []string{ var ModelList = []string{
"kimi-k3",
"kimi-k2.5", "kimi-k2.5",
"kimi-k2-0905-preview", "kimi-k2-0905-preview",
"kimi-k2-turbo-preview", "kimi-k2-turbo-preview",
......
...@@ -29,9 +29,11 @@ func applySystemPromptIfNeeded(c *gin.Context, info *relaycommon.RelayInfo, requ ...@@ -29,9 +29,11 @@ func applySystemPromptIfNeeded(c *gin.Context, info *relaycommon.RelayInfo, requ
systemRole := request.GetSystemRoleName() systemRole := request.GetSystemRoleName()
// A Kimi K3 dynamic tool loading message ({"role":"system","tools":[...]})
// declares tools rather than a system prompt and must never receive content.
containSystemPrompt := false containSystemPrompt := false
for _, message := range request.Messages { for _, message := range request.Messages {
if message.Role == systemRole { if message.Role == systemRole && len(message.Tools) == 0 {
containSystemPrompt = true containSystemPrompt = true
break break
} }
...@@ -51,7 +53,7 @@ func applySystemPromptIfNeeded(c *gin.Context, info *relaycommon.RelayInfo, requ ...@@ -51,7 +53,7 @@ func applySystemPromptIfNeeded(c *gin.Context, info *relaycommon.RelayInfo, requ
common.SetContextKey(c, constant.ContextKeySystemPromptOverride, true) common.SetContextKey(c, constant.ContextKeySystemPromptOverride, true)
for i, message := range request.Messages { for i, message := range request.Messages {
if message.Role != systemRole { if message.Role != systemRole || len(message.Tools) > 0 {
continue continue
} }
if message.IsStringContent() { if message.IsStringContent() {
......
package relay package relay
import ( import (
"encoding/json"
"io" "io"
"math" "math"
"net/http" "net/http"
...@@ -153,3 +154,72 @@ func TestTextRequestViaResponsesConvertsClaudeDirectly(t *testing.T) { ...@@ -153,3 +154,72 @@ func TestTextRequestViaResponsesConvertsClaudeDirectly(t *testing.T) {
require.Len(t, response.Content, 1) require.Len(t, response.Content, 1)
assert.Equal(t, "ok", response.Content[0].GetText()) assert.Equal(t, "ok", response.Content[0].GetText())
} }
func TestApplySystemPromptIfNeededSkipsToolLoadingMessages(t *testing.T) {
tools := json.RawMessage(`[{"type":"function","function":{"name":"get_current_time","parameters":{"type":"object","properties":{"city":{"type":"string"}}}}}]`)
toolLoading := dto.Message{Role: "system", Tools: tools}
user := dto.Message{Role: "user", Content: "What time is it in Beijing?"}
tests := []struct {
name string
messages []dto.Message
wantMessages []dto.Message
wantOverride bool
}{
{
name: "tool loading message alone is not a system prompt",
messages: []dto.Message{toolLoading, user},
wantMessages: []dto.Message{
{Role: "system", Content: "Answer in English."},
toolLoading,
user,
},
},
{
name: "override targets the real system prompt only",
messages: []dto.Message{toolLoading, {Role: "system", Content: "You are Kimi."}, user},
wantMessages: []dto.Message{
toolLoading,
{Role: "system", Content: "Answer in English.\nYou are Kimi."},
user,
},
wantOverride: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(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{
ChannelMeta: &relaycommon.ChannelMeta{
ChannelSetting: dto.ChannelSettings{
SystemPrompt: "Answer in English.",
SystemPromptOverride: true,
},
},
}
request := &dto.GeneralOpenAIRequest{
Model: "kimi-k3",
Messages: append([]dto.Message(nil), tt.messages...),
}
applySystemPromptIfNeeded(c, info, request)
require.Len(t, request.Messages, len(tt.wantMessages))
for i, want := range tt.wantMessages {
got := request.Messages[i]
assert.Equal(t, want.Role, got.Role, "message %d role", i)
assert.Equal(t, want.Content, got.Content, "message %d content", i)
if len(want.Tools) > 0 {
assert.JSONEq(t, string(want.Tools), string(got.Tools), "message %d tools", i)
} else {
assert.Empty(t, got.Tools, "message %d tools", i)
}
}
_, overrideSet := common.GetContextKey(c, constant.ContextKeySystemPromptOverride)
assert.Equal(t, tt.wantOverride, overrideSet)
})
}
}
...@@ -115,46 +115,8 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types ...@@ -115,46 +115,8 @@ func TextHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types
} }
relaycommon.AppendRequestConversionFromRequest(info, convertedRequest) relaycommon.AppendRequestConversionFromRequest(info, convertedRequest)
if info.ChannelSetting.SystemPrompt != "" { if req, ok := convertedRequest.(*dto.GeneralOpenAIRequest); ok {
// 如果有系统提示,则将其添加到请求中 applySystemPromptIfNeeded(c, info, req)
request, ok := convertedRequest.(*dto.GeneralOpenAIRequest)
if ok {
containSystemPrompt := false
for _, message := range request.Messages {
if message.Role == request.GetSystemRoleName() {
containSystemPrompt = true
break
}
}
if !containSystemPrompt {
// 如果没有系统提示,则添加系统提示
systemMessage := dto.Message{
Role: request.GetSystemRoleName(),
Content: info.ChannelSetting.SystemPrompt,
}
request.Messages = append([]dto.Message{systemMessage}, request.Messages...)
} else if info.ChannelSetting.SystemPromptOverride {
common.SetContextKey(c, constant.ContextKeySystemPromptOverride, true)
// 如果有系统提示,且允许覆盖,则拼接到前面
for i, message := range request.Messages {
if message.Role == request.GetSystemRoleName() {
if message.IsStringContent() {
request.Messages[i].SetStringContent(info.ChannelSetting.SystemPrompt + "\n" + message.StringContent())
} else {
contents := message.ParseContent()
contents = append([]dto.MediaContent{
{
Type: dto.ContentTypeText,
Text: info.ChannelSetting.SystemPrompt,
},
}, contents...)
request.Messages[i].Content = contents
}
break
}
}
}
}
} }
jsonData, err := common.Marshal(convertedRequest) jsonData, err := common.Marshal(convertedRequest)
......
...@@ -119,7 +119,43 @@ func (r GeneralOpenAIRequest) MarshalJSON() ([]byte, error) { ...@@ -119,7 +119,43 @@ func (r GeneralOpenAIRequest) MarshalJSON() ([]byte, error) {
if !IsQwenThinkingBudgetModel(r.Model) { if !IsQwenThinkingBudgetModel(r.Model) {
r.ThinkingBudget = nil r.ThinkingBudget = nil
} }
return kitutil.Marshal((*Alias)(&r))
hasToolLoadingMessage := false
for _, message := range r.Messages {
if len(message.Tools) > 0 && message.Content == nil {
hasToolLoadingMessage = true
break
}
}
if !hasToolLoadingMessage {
return kitutil.Marshal((*Alias)(&r))
}
// Kimi K3 dynamic tool loading: a system message that carries tools must not
// carry a content key at all, otherwise the upstream rejects it. Only those
// messages drop the key; every other message keeps emitting "content": null.
type toolLoadingMessage struct {
Message
Content any `json:"content,omitempty"`
}
messages := make([]json.RawMessage, 0, len(r.Messages))
for _, message := range r.Messages {
var encoded []byte
var err error
if len(message.Tools) > 0 && message.Content == nil {
encoded, err = kitutil.Marshal(toolLoadingMessage{Message: message})
} else {
encoded, err = kitutil.Marshal(message)
}
if err != nil {
return nil, err
}
messages = append(messages, encoded)
}
return kitutil.Marshal(struct {
*Alias
Messages []json.RawMessage `json:"messages,omitempty"`
}{Alias: (*Alias)(&r), Messages: messages})
} }
func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta { func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta {
...@@ -155,9 +191,18 @@ func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta { ...@@ -155,9 +191,18 @@ func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta {
tokenCountMeta.MaxTokens = int(maxTokens) tokenCountMeta.MaxTokens = int(maxTokens)
} }
var dynamicTools []ToolCallRequest
for _, message := range r.Messages { for _, message := range r.Messages {
tokenCountMeta.MessagesCount++ tokenCountMeta.MessagesCount++
texts = append(texts, message.Role) texts = append(texts, message.Role)
if len(message.Tools) > 0 {
// Kimi K3 dynamic tool loading: tools declared on a message are
// visible to the model and are counted like top-level tools.
var messageTools []ToolCallRequest
if err := kitutil.Unmarshal(message.Tools, &messageTools); err == nil {
dynamicTools = append(dynamicTools, messageTools...)
}
}
if message.Content != nil { if message.Content != nil {
if message.Name != nil { if message.Name != nil {
tokenCountMeta.NameCount++ tokenCountMeta.NameCount++
...@@ -189,22 +234,23 @@ func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta { ...@@ -189,22 +234,23 @@ func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta {
} }
} }
if r.Tools != nil { tools := r.Tools
openaiTools := r.Tools if len(dynamicTools) > 0 {
for _, tool := range openaiTools { tools = append(dynamicTools, r.Tools...)
tokenCountMeta.ToolsCount++ }
texts = append(texts, tool.Function.Name) for _, tool := range tools {
if tool.Function.Description != "" { tokenCountMeta.ToolsCount++
texts = append(texts, tool.Function.Description) texts = append(texts, tool.Function.Name)
} if tool.Function.Description != "" {
if tool.Function.Parameters != nil { texts = append(texts, tool.Function.Description)
texts = append(texts, fmt.Sprintf("%v", tool.Function.Parameters)) }
} if tool.Function.Parameters != nil {
texts = append(texts, fmt.Sprintf("%v", tool.Function.Parameters))
} }
//toolTokens := CountTokenInput(countStr, request.Model)
//tkm += 8
//tkm += toolTokens
} }
//toolTokens := CountTokenInput(countStr, request.Model)
//tkm += 8
//tkm += toolTokens
tokenCountMeta.CombineText = strings.Join(texts, "\n") tokenCountMeta.CombineText = strings.Join(texts, "\n")
tokenCountMeta.Files = fileMeta tokenCountMeta.Files = fileMeta
return &tokenCountMeta return &tokenCountMeta
...@@ -378,6 +424,9 @@ type Message struct { ...@@ -378,6 +424,9 @@ 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"`
// Tools carries Kimi K3 dynamic tool loading declarations on a system message.
// Same shape as the top-level tools array; passthrough-only for OpenAI-compatible upstreams.
Tools json.RawMessage `json:"tools,omitempty"`
// Annotations is an official Chat response field. Keeping it on the shared // Annotations is an official Chat response field. Keeping it on the shared
// message type also preserves annotations when clients replay assistant output. // message type also preserves annotations when clients replay assistant output.
Annotations json.RawMessage `json:"annotations,omitempty"` Annotations json.RawMessage `json:"annotations,omitempty"`
......
...@@ -248,3 +248,62 @@ func TestIsOpenAIGPT5Model(t *testing.T) { ...@@ -248,3 +248,62 @@ func TestIsOpenAIGPT5Model(t *testing.T) {
}) })
} }
} }
func TestGeneralOpenAIRequestPreserveMessageLevelTools(t *testing.T) {
raw := []byte(`{
"model":"kimi-k3",
"tool_choice":"required",
"tools":[{"type":"function","function":{"name":"get_weather","description":"Get the weather","parameters":{"type":"object","properties":{"city":{"type":"string"}}}}}],
"messages":[
{"role":"system","content":"You are Kimi."},
{"role":"user","content":"What time is it in Beijing?"},
{"role":"system","tools":[{"type":"function","function":{"name":"get_current_time","description":"Get the current time of a city","parameters":{"type":"object","properties":{"city":{"type":"string"}},"required":["city"]}}}]},
{"role":"assistant","content":null,"tool_calls":[{"id":"call_1","type":"function","function":{"name":"get_current_time","arguments":"{\"city\":\"Beijing\"}"}}]},
{"role":"system","content":"","tools":[{"type":"function","function":{"name":"lookup_order","parameters":{"type":"object"}}}]}
]
}`)
var req GeneralOpenAIRequest
require.NoError(t, kitutil.Unmarshal(raw, &req))
require.Len(t, req.Messages, 5)
encoded, err := kitutil.Marshal(req)
require.NoError(t, err)
messages := gjson.GetBytes(encoded, "messages").Array()
require.Len(t, messages, 5)
assert.Equal(t, "required", gjson.GetBytes(encoded, "tool_choice").String())
assert.JSONEq(t, gjson.GetBytes(raw, "tools").Raw, gjson.GetBytes(encoded, "tools").Raw)
// Regular messages keep their content untouched.
assert.Equal(t, "You are Kimi.", messages[0].Get("content").String())
assert.False(t, messages[0].Get("tools").Exists())
assert.Equal(t, "What time is it in Beijing?", messages[1].Get("content").String())
// Kimi K3 dynamic tool loading message: tools preserved byte-for-byte, no content key at all.
assert.JSONEq(t, gjson.GetBytes(raw, "messages.2.tools").Raw, messages[2].Get("tools").Raw)
assert.False(t, messages[2].Get("content").Exists())
assert.Equal(t, "system", messages[2].Get("role").String())
// Assistant tool call replay still emits an explicit "content": null.
assistantContent := messages[3].Get("content")
assert.True(t, assistantContent.Exists())
assert.Equal(t, gjson.Null, assistantContent.Type)
assert.JSONEq(t, gjson.GetBytes(raw, "messages.3.tool_calls").Raw, messages[3].Get("tool_calls").Raw)
// Explicit empty content next to tools is forwarded as-is for the upstream to judge.
emptyContent := messages[4].Get("content")
assert.True(t, emptyContent.Exists())
assert.Equal(t, gjson.String, emptyContent.Type)
assert.Equal(t, "", emptyContent.String())
assert.JSONEq(t, gjson.GetBytes(raw, "messages.4.tools").Raw, messages[4].Get("tools").Raw)
// Token estimation sees message-level tools alongside the top-level ones.
meta := req.GetTokenCountMeta()
assert.Equal(t, 3, meta.ToolsCount)
assert.Equal(t, 5, meta.MessagesCount)
assert.Contains(t, meta.CombineText, "get_weather")
assert.Contains(t, meta.CombineText, "get_current_time")
assert.Contains(t, meta.CombineText, "Get the current time of a city")
assert.Contains(t, meta.CombineText, "lookup_order")
}
package kitutil
import (
"encoding/json"
"io"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// recordingCodec notes which Codec method each helper reaches and delegates to
// the standard-library default so the helpers still produce real results.
type recordingCodec struct {
calls []string
}
func (c *recordingCodec) Marshal(v any) ([]byte, error) {
c.calls = append(c.calls, "Marshal")
return stdCodec{}.Marshal(v)
}
func (c *recordingCodec) Unmarshal(data []byte, v any) error {
c.calls = append(c.calls, "Unmarshal")
return stdCodec{}.Unmarshal(data, v)
}
func (c *recordingCodec) Decode(r io.Reader, v any) error {
c.calls = append(c.calls, "Decode")
return stdCodec{}.Decode(r, v)
}
func (c *recordingCodec) Valid(data []byte) bool {
c.calls = append(c.calls, "Valid")
return stdCodec{}.Valid(data)
}
func TestJSONHelpersRouteThroughInjectedCodec(t *testing.T) {
t.Cleanup(func() { SetCodec(stdCodec{}) })
fake := &recordingCodec{}
SetCodec(fake)
// A nil codec must not displace the installed one.
SetCodec(nil)
encoded, err := Marshal(map[string]int{"a": 1})
require.NoError(t, err)
assert.Equal(t, `{"a":1}`, string(encoded))
var fromBytes map[string]int
require.NoError(t, Unmarshal([]byte(`{"b":2}`), &fromBytes))
assert.Equal(t, map[string]int{"b": 2}, fromBytes)
var fromString map[string]int
require.NoError(t, UnmarshalJsonStr(`{"c":3}`, &fromString))
assert.Equal(t, map[string]int{"c": 3}, fromString)
var decoded map[string]int
require.NoError(t, DecodeJson(strings.NewReader(`{"d":4}`), &decoded))
assert.Equal(t, map[string]int{"d": 4}, decoded)
assert.True(t, Valid([]byte(`[]`)))
assert.False(t, Valid([]byte(`[`)))
converted, err := Any2Type[map[string]int](map[string]any{"e": 5})
require.NoError(t, err)
assert.Equal(t, map[string]int{"e": 5}, converted)
assert.Equal(t, "hello", JsonRawMessageToString(json.RawMessage(`"hello"`)))
assert.Equal(t, []string{
"Marshal", // Marshal
"Unmarshal", // Unmarshal
"Unmarshal", // UnmarshalJsonStr
"Decode", // DecodeJson
"Valid", // Valid (well-formed)
"Valid", // Valid (malformed)
"Marshal", // Any2Type encode
"Unmarshal", // Any2Type decode
"Unmarshal", // JsonRawMessageToString string literal
}, fake.calls)
}
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