feat(providers): robust minimax tool call and thinking extraction

This commit is contained in:
Andrew 2026-02-26 10:43:11 +00:00
parent 8a1fb03974
commit 96809986de
2 changed files with 443 additions and 2 deletions

View file

@ -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,

View file

@ -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\" &amp; &lt;trends&gt;</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)
}
}