feat(providers): robust minimax tool call and thinking extraction
This commit is contained in:
parent
8a1fb03974
commit
96809986de
2 changed files with 443 additions and 2 deletions
|
|
@ -5,6 +5,7 @@ import (
|
|||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"html"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
|
|
@ -15,6 +16,22 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
|
||||
)
|
||||
|
||||
const (
|
||||
thinkOpenTag = "<think>"
|
||||
thinkCloseTag = "</think>"
|
||||
miniMaxOpen = "[TOOL_CALL]"
|
||||
miniMaxCloseA = "</minimax:tool_call>"
|
||||
miniMaxCloseB = "[/minimax:tool_call]"
|
||||
invokeOpen = "<invoke name=\""
|
||||
invokeClose = "</invoke>"
|
||||
paramOpen = "<parameter name=\""
|
||||
paramClose = "</parameter>"
|
||||
|
||||
maxReasoningBlocks = 10
|
||||
maxMiniMaxToolCalls = 20
|
||||
maxParameters = 50
|
||||
)
|
||||
|
||||
type (
|
||||
ToolCall = protocoltypes.ToolCall
|
||||
FunctionCall = protocoltypes.FunctionCall
|
||||
|
|
@ -231,7 +248,154 @@ func parseResponse(body []byte) (*LLMResponse, error) {
|
|||
}
|
||||
|
||||
choice := apiResponse.Choices[0]
|
||||
content := choice.Message.Content
|
||||
reasoningContent := choice.Message.ReasoningContent
|
||||
toolCalls := make([]ToolCall, 0, len(choice.Message.ToolCalls))
|
||||
|
||||
// --- 1. Extract reasoning blocks from content if present ---
|
||||
// This handles models that put reasoning in the main content block.
|
||||
// We handle multiple blocks and join them, with a safety limit.
|
||||
for i := 0; i < maxReasoningBlocks; i++ {
|
||||
start := strings.Index(content, thinkOpenTag)
|
||||
if start == -1 {
|
||||
break
|
||||
}
|
||||
// Search for the closing tag starting after the opening tag
|
||||
endRel := strings.Index(content[start:], thinkCloseTag)
|
||||
if endRel == -1 {
|
||||
// Malformed or cut off; leave it to avoid data loss
|
||||
break
|
||||
}
|
||||
end := start + endRel
|
||||
extracted := strings.TrimSpace(content[start+len(thinkOpenTag) : end])
|
||||
if reasoningContent == "" {
|
||||
reasoningContent = extracted
|
||||
} else if extracted != "" {
|
||||
reasoningContent += "\n\n" + extracted
|
||||
}
|
||||
content = strings.TrimSpace(content[:start] + content[end+len(thinkCloseTag):])
|
||||
}
|
||||
|
||||
// --- 2. Extract MiniMax-style [TOOL_CALL] XML if present ---
|
||||
for miniMaxIdx := 0; miniMaxIdx < maxMiniMaxToolCalls; miniMaxIdx++ {
|
||||
tagStart := strings.Index(content, miniMaxOpen)
|
||||
if tagStart == -1 {
|
||||
break
|
||||
}
|
||||
|
||||
// Search for the earliest valid closing tag after the opening tag
|
||||
angleIdx := strings.Index(content[tagStart:], miniMaxCloseA)
|
||||
bracketIdx := strings.Index(content[tagStart:], miniMaxCloseB)
|
||||
|
||||
tagEndIdx := -1
|
||||
tagLen := 0
|
||||
|
||||
if angleIdx != -1 && (bracketIdx == -1 || angleIdx < bracketIdx) {
|
||||
tagEndIdx = tagStart + angleIdx
|
||||
tagLen = len(miniMaxCloseA)
|
||||
} else if bracketIdx != -1 {
|
||||
tagEndIdx = tagStart + bracketIdx
|
||||
tagLen = len(miniMaxCloseB)
|
||||
}
|
||||
|
||||
// If no closing tag is found, the string is malformed or cut off. Break to avoid infinite loop.
|
||||
if tagEndIdx == -1 {
|
||||
break
|
||||
}
|
||||
|
||||
// Calculate indices for XML content extraction
|
||||
xmlBodyStart := tagStart + len(miniMaxOpen)
|
||||
if xmlBodyStart > tagEndIdx {
|
||||
break
|
||||
}
|
||||
xmlPart := content[xmlBodyStart:tagEndIdx]
|
||||
|
||||
// Very basic XML-ish parsing for MiniMax format
|
||||
// Extract name: <invoke name="([^"]+)">
|
||||
nameStart := strings.Index(xmlPart, invokeOpen)
|
||||
if nameStart != -1 {
|
||||
nameStart += len(invokeOpen)
|
||||
nameEnd := strings.Index(xmlPart[nameStart:], "\"")
|
||||
invokeEnd := strings.Index(xmlPart, invokeClose)
|
||||
|
||||
if nameEnd != -1 && invokeEnd != -1 {
|
||||
toolName := xmlPart[nameStart : nameStart+nameEnd]
|
||||
args := make(map[string]any)
|
||||
|
||||
// Extract parameters: <parameter name="([^"]+)">([^<]*)</parameter>
|
||||
paramsPart := xmlPart[nameStart+nameEnd:]
|
||||
malformedParams := false
|
||||
for pCount := 0; pCount < maxParameters; pCount++ {
|
||||
pStart := strings.Index(paramsPart, paramOpen)
|
||||
if pStart == -1 {
|
||||
break
|
||||
}
|
||||
pStart += len(paramOpen)
|
||||
|
||||
// Bounds check for parameter name
|
||||
if pStart >= len(paramsPart) {
|
||||
malformedParams = true
|
||||
break
|
||||
}
|
||||
pNameEnd := strings.Index(paramsPart[pStart:], "\"")
|
||||
if pNameEnd == -1 {
|
||||
malformedParams = true
|
||||
break
|
||||
}
|
||||
pName := paramsPart[pStart : pStart+pNameEnd]
|
||||
|
||||
// Search for value start ">"
|
||||
valMarkerIdx := strings.Index(paramsPart[pStart+pNameEnd:], ">")
|
||||
if valMarkerIdx == -1 {
|
||||
malformedParams = true
|
||||
break
|
||||
}
|
||||
valueStart := pStart + pNameEnd + valMarkerIdx + 1
|
||||
if valueStart > len(paramsPart) {
|
||||
malformedParams = true
|
||||
break
|
||||
}
|
||||
|
||||
// Search for closing </parameter>
|
||||
valueEndMarkerIdx := strings.Index(paramsPart[valueStart:], paramClose)
|
||||
if valueEndMarkerIdx == -1 {
|
||||
malformedParams = true
|
||||
break
|
||||
}
|
||||
valueEnd := valueStart + valueEndMarkerIdx
|
||||
|
||||
// Extract and decode entities
|
||||
value := html.UnescapeString(paramsPart[valueStart:valueEnd])
|
||||
args[pName] = value
|
||||
|
||||
// Advance to the next parameter
|
||||
nextParamStart := valueEnd + len(paramClose)
|
||||
if nextParamStart > len(paramsPart) {
|
||||
paramsPart = ""
|
||||
break
|
||||
}
|
||||
paramsPart = paramsPart[nextParamStart:]
|
||||
}
|
||||
|
||||
if !malformedParams {
|
||||
toolCalls = append(toolCalls, ToolCall{
|
||||
ID: fmt.Sprintf("minimax-%d-%d", time.Now().UnixNano(), miniMaxIdx),
|
||||
Name: toolName,
|
||||
Arguments: args,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
content = strings.TrimSpace(content[:tagStart] + content[tagEndIdx+tagLen:])
|
||||
|
||||
// Override FinishReason so the loop agent knows to execute it
|
||||
if choice.FinishReason != "tool_calls" {
|
||||
choice.FinishReason = "tool_calls"
|
||||
}
|
||||
}
|
||||
|
||||
// --- 3. Process standard OpenAI JSON tool calls ---
|
||||
for _, tc := range choice.Message.ToolCalls {
|
||||
arguments := make(map[string]any)
|
||||
name := ""
|
||||
|
|
@ -272,8 +436,8 @@ func parseResponse(body []byte) (*LLMResponse, error) {
|
|||
}
|
||||
|
||||
return &LLMResponse{
|
||||
Content: choice.Message.Content,
|
||||
ReasoningContent: choice.Message.ReasoningContent,
|
||||
Content: content,
|
||||
ReasoningContent: reasoningContent,
|
||||
ToolCalls: toolCalls,
|
||||
FinishReason: choice.FinishReason,
|
||||
Usage: apiResponse.Usage,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import (
|
|||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
|
@ -327,6 +328,235 @@ func TestNormalizeModel_UsesAPIBase(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestProviderChat_ParsesMiniMaxToolCallsAndThinking(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{
|
||||
"content": "<think>\nThinking about the news...\n</think>\n\nI will search for the news.\n\n[TOOL_CALL]<invoke name=\"web_search\"><parameter name=\"count\">5</parameter><parameter name=\"query\">AI news</parameter></invoke></minimax:tool_call>",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
p := NewProvider("key", server.URL, "")
|
||||
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "minimax-m2.5", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
if !strings.Contains(out.ReasoningContent, "Thinking about the news") {
|
||||
t.Errorf("ReasoningContent = %q, want it to contain reasoning", out.ReasoningContent)
|
||||
}
|
||||
if out.Content != "I will search for the news." {
|
||||
t.Errorf("Content = %q, want %q", out.Content, "I will search for the news.")
|
||||
}
|
||||
if len(out.ToolCalls) != 1 {
|
||||
t.Fatalf("len(ToolCalls) = %d, want 1", len(out.ToolCalls))
|
||||
}
|
||||
if out.ToolCalls[0].Name != "web_search" {
|
||||
t.Errorf("ToolCalls[0].Name = %q, want %q", out.ToolCalls[0].Name, "web_search")
|
||||
}
|
||||
if out.ToolCalls[0].Arguments["query"] != "AI news" {
|
||||
t.Errorf("ToolCalls[0].Arguments[query] = %v, want AI news", out.ToolCalls[0].Arguments["query"])
|
||||
}
|
||||
if out.ToolCalls[0].Arguments["count"] != "5" {
|
||||
t.Errorf("ToolCalls[0].Arguments[count] = %v, want 5", out.ToolCalls[0].Arguments["count"])
|
||||
}
|
||||
if out.FinishReason != "tool_calls" {
|
||||
t.Errorf("FinishReason = %q, want %q", out.FinishReason, "tool_calls")
|
||||
}
|
||||
|
||||
// Verify tags are removed from content/reasoning
|
||||
if strings.Contains(out.Content, "[TOOL_CALL]") {
|
||||
t.Errorf("Content contains [TOOL_CALL]")
|
||||
}
|
||||
if strings.Contains(out.ReasoningContent, "<think>") {
|
||||
t.Errorf("ReasoningContent contains <think>")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderChat_ParsesMiniMaxMultipleToolCalls(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{
|
||||
"content": "[TOOL_CALL]<invoke name=\"web_search\"><parameter name=\"query\">AI news</parameter></invoke></minimax:tool_call>\n\n[TOOL_CALL]<invoke name=\"translate\"><parameter name=\"text\">hello</parameter><parameter name=\"target_lang\">es</parameter></invoke></minimax:tool_call>",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
p := NewProvider("key", server.URL, "")
|
||||
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "minimax-m2.5", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
if len(out.ToolCalls) != 2 {
|
||||
t.Fatalf("len(ToolCalls) = %d, want 2", len(out.ToolCalls))
|
||||
}
|
||||
if out.ToolCalls[0].Name != "web_search" || out.ToolCalls[1].Name != "translate" {
|
||||
t.Errorf("ToolCall names mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderChat_ParsesMiniMaxToolCallsWithSpecialCharsAndEmptyValues(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{
|
||||
"content": "[TOOL_CALL]<invoke name=\"web_search\"><parameter name=\"query\">AI \"news\" & <trends></parameter><parameter name=\"optional\"></parameter></invoke></minimax:tool_call>",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
p := NewProvider("key", server.URL, "")
|
||||
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "minimax-m2.5", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
if len(out.ToolCalls) != 1 {
|
||||
t.Fatalf("len(ToolCalls) = %d, want 1", len(out.ToolCalls))
|
||||
}
|
||||
query := out.ToolCalls[0].Arguments["query"].(string)
|
||||
if query != `AI "news" & <trends>` {
|
||||
t.Errorf("query = %q, want decoded XML special chars", query)
|
||||
}
|
||||
|
||||
optionalVal, ok := out.ToolCalls[0].Arguments["optional"]
|
||||
if !ok {
|
||||
t.Fatalf(`optional argument missing; want present with value ""`)
|
||||
} else if optionalStr, ok := optionalVal.(string); !ok || optionalStr != "" {
|
||||
t.Errorf("optional = %#v, want empty string", optionalVal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderChat_MergesJSONAndMiniMaxToolCalls(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{
|
||||
"content": "[TOOL_CALL]<invoke name=\"web_search\"><parameter name=\"query\">AI news</parameter></invoke></minimax:tool_call>",
|
||||
"tool_calls": []map[string]any{
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": map[string]any{
|
||||
"name": "get_weather",
|
||||
"arguments": `{"location":"San Francisco"}`,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
p := NewProvider("key", server.URL, "")
|
||||
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "minimax-m2.5", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
if len(out.ToolCalls) != 2 {
|
||||
t.Fatalf("len(ToolCalls) = %d, want 2 (1 JSON tool call + 1 MiniMax tool call)", len(out.ToolCalls))
|
||||
}
|
||||
|
||||
// Verify both names exist
|
||||
names := map[string]bool{out.ToolCalls[0].Name: true, out.ToolCalls[1].Name: true}
|
||||
if !names["web_search"] || !names["get_weather"] {
|
||||
t.Errorf("merged tool calls mismatch: %v", names)
|
||||
}
|
||||
|
||||
if out.FinishReason != "tool_calls" {
|
||||
t.Fatalf("FinishReason = %q, want %q", out.FinishReason, "tool_calls")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderChat_ParsesMiniMaxToolCallsWithAlternativeClosingTag(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{
|
||||
"content": "[TOOL_CALL]<invoke name=\"web_search\"><parameter name=\"query\">AI news</parameter></invoke>[/minimax:tool_call]",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
p := NewProvider("key", server.URL, "")
|
||||
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "minimax-m2.5", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
if len(out.ToolCalls) != 1 {
|
||||
t.Fatalf("len(ToolCalls) = %d, want 1", len(out.ToolCalls))
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderChat_ParsesMiniMaxToolCallsWithMalformedXML(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{
|
||||
"content": "[TOOL_CALL]<invoke name=\"web_search\"><parameter name=\"query\">AI news</minimax:tool_call>",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
p := NewProvider("key", server.URL, "")
|
||||
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "minimax-m2.5", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
if len(out.ToolCalls) != 0 {
|
||||
t.Errorf("len(ToolCalls) = %d, want 0 for malformed XML", len(out.ToolCalls))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
func TestProvider_RequestTimeoutDefault(t *testing.T) {
|
||||
p := NewProviderWithMaxTokensFieldAndTimeout("key", "https://example.com/v1", "", "", 0)
|
||||
if p.httpClient.Timeout != defaultRequestTimeout {
|
||||
|
|
@ -361,3 +591,50 @@ func TestProvider_FunctionalOptionRequestTimeoutNonPositive(t *testing.T) {
|
|||
t.Fatalf("http timeout = %v, want %v", p.httpClient.Timeout, defaultRequestTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProviderChat_ParsesMiniMaxToolCallsGeneralization(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
resp := map[string]any{
|
||||
"choices": []map[string]any{
|
||||
{
|
||||
"message": map[string]any{
|
||||
"content": "[TOOL_CALL]<invoke name=\"shell_execute\"><parameter name=\"command\">ls -la</parameter><parameter name=\"timeout\">30</parameter></invoke></minimax:tool_call>\n" +
|
||||
"[TOOL_CALL]<invoke name=\"read_file\"><parameter name=\"path\">/etc/passwd</parameter></invoke>[/minimax:tool_call]",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
p := NewProvider("key", server.URL, "")
|
||||
out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "minimax-m2.5", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Chat() error = %v", err)
|
||||
}
|
||||
|
||||
if len(out.ToolCalls) != 2 {
|
||||
t.Fatalf("len(ToolCalls) = %d, want 2", len(out.ToolCalls))
|
||||
}
|
||||
|
||||
// Verify tool 1: shell_execute
|
||||
tc1 := out.ToolCalls[0]
|
||||
if tc1.Name != "shell_execute" {
|
||||
t.Errorf("tc1.Name = %q, want shell_execute", tc1.Name)
|
||||
}
|
||||
if tc1.Arguments["command"] != "ls -la" || tc1.Arguments["timeout"] != "30" {
|
||||
t.Errorf("tc1 arguments mismatch: %v", tc1.Arguments)
|
||||
}
|
||||
|
||||
// Verify tool 2: read_file
|
||||
tc2 := out.ToolCalls[1]
|
||||
if tc2.Name != "read_file" {
|
||||
t.Errorf("tc2.Name = %q, want read_file", tc2.Name)
|
||||
}
|
||||
if tc2.Arguments["path"] != "/etc/passwd" {
|
||||
t.Errorf("tc2 arguments mismatch: %v", tc2.Arguments)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue