diff --git a/pkg/providers/openai_compat/provider.go b/pkg/providers/openai_compat/provider.go index 7dace71f2..dfc73153f 100644 --- a/pkg/providers/openai_compat/provider.go +++ b/pkg/providers/openai_compat/provider.go @@ -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 = "" + thinkCloseTag = "" + miniMaxOpen = "[TOOL_CALL]" + miniMaxCloseA = "" + miniMaxCloseB = "[/minimax:tool_call]" + invokeOpen = " tagEndIdx { + break + } + xmlPart := content[xmlBodyStart:tagEndIdx] + + // Very basic XML-ish parsing for MiniMax format + // Extract 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: ([^<]*) + 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 + 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, diff --git a/pkg/providers/openai_compat/provider_test.go b/pkg/providers/openai_compat/provider_test.go index 7247fea3e..55cd3d8c0 100644 --- a/pkg/providers/openai_compat/provider_test.go +++ b/pkg/providers/openai_compat/provider_test.go @@ -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": "\nThinking about the news...\n\n\nI will search for the news.\n\n[TOOL_CALL]5AI news", + }, + "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, "") { + t.Errorf("ReasoningContent contains ") + } +} + +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]AI news\n\n[TOOL_CALL]helloes", + }, + "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]AI \"news\" & <trends>", + }, + "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" & ` { + 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]AI news", + "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]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) != 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]AI news", + }, + "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]ls -la30\n" + + "[TOOL_CALL]/etc/passwd[/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) + } +}