diff --git a/config/config.example.json b/config/config.example.json index 8db9dd3e9..e4a16de68 100644 --- a/config/config.example.json +++ b/config/config.example.json @@ -35,6 +35,12 @@ "model": "deepseek/deepseek-chat", "api_key": "sk-your-deepseek-key" }, + { + "model_name": "minimax", + "model": "minimax/MiniMax-M1", + "api_key": "your-minimax-api-key", + "stream": true + }, { "model_name": "loadbalanced-gpt4", "model": "openai/gpt-5.2", diff --git a/pkg/config/config.go b/pkg/config/config.go index 6ce0b6ad6..aec93f663 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -397,6 +397,7 @@ type ModelConfig struct { // Optional optimizations RPM int `json:"rpm,omitempty"` // Requests per minute limit MaxTokensField string `json:"max_tokens_field,omitempty"` // Field name for max tokens (e.g., "max_completion_tokens") + Stream *bool `json:"stream,omitempty"` // Use SSE streaming (default: protocol-dependent) } // Validate checks if the ModelConfig has all required fields. diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go index 74fe8a36c..a0e44aca7 100644 --- a/pkg/providers/factory_provider.go +++ b/pkg/providers/factory_provider.go @@ -10,6 +10,7 @@ import ( "strings" "github.com/sipeed/picoclaw/pkg/config" + "github.com/sipeed/picoclaw/pkg/providers/openai_compat" ) // createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store. @@ -84,7 +85,25 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err if apiBase == "" { apiBase = getDefaultAPIBase(protocol) } - return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil + return NewHTTPProviderWithOptions(cfg.APIKey, apiBase, cfg.Proxy, openai_compat.Options{ + MaxTokensField: cfg.MaxTokensField, + Stream: boolDefault(cfg.Stream, false), + }), modelID, nil + + case "minimax": + // MiniMax uses a non-standard endpoint path and defaults to SSE streaming. + if cfg.APIKey == "" && cfg.APIBase == "" { + return nil, "", fmt.Errorf("api_key or api_base is required for minimax protocol") + } + apiBase := cfg.APIBase + if apiBase == "" { + apiBase = getDefaultAPIBase(protocol) + } + return NewHTTPProviderWithOptions(cfg.APIKey, apiBase, cfg.Proxy, openai_compat.Options{ + EndpointPath: "/text/chatcompletion_v2", + MaxTokensField: cfg.MaxTokensField, + Stream: boolDefault(cfg.Stream, true), + }), modelID, nil case "openrouter", "groq", "zhipu", "gemini", "nvidia", "ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras", @@ -97,7 +116,10 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err if apiBase == "" { apiBase = getDefaultAPIBase(protocol) } - return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil + return NewHTTPProviderWithOptions(cfg.APIKey, apiBase, cfg.Proxy, openai_compat.Options{ + MaxTokensField: cfg.MaxTokensField, + Stream: boolDefault(cfg.Stream, false), + }), modelID, nil case "anthropic": if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" { @@ -186,7 +208,17 @@ func getDefaultAPIBase(protocol string) string { return "https://dashscope.aliyuncs.com/compatible-mode/v1" case "vllm": return "http://localhost:8000/v1" + case "minimax": + return "https://api.minimax.io/v1" default: return "" } } + +// boolDefault dereferences a *bool, returning def when nil. +func boolDefault(p *bool, def bool) bool { + if p != nil { + return *p + } + return def +} diff --git a/pkg/providers/http_provider.go b/pkg/providers/http_provider.go index f669d74cb..46f921efd 100644 --- a/pkg/providers/http_provider.go +++ b/pkg/providers/http_provider.go @@ -28,6 +28,12 @@ func NewHTTPProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField st } } +func NewHTTPProviderWithOptions(apiKey, apiBase, proxy string, opts openai_compat.Options) *HTTPProvider { + return &HTTPProvider{ + delegate: openai_compat.NewProviderWithOptions(apiKey, apiBase, proxy, opts), + } +} + func (p *HTTPProvider) Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) { resp, err := p.delegate.Chat(ctx, messages, tools, model, options) if err != nil { diff --git a/pkg/providers/openai_compat/provider.go b/pkg/providers/openai_compat/provider.go index b8528953a..17ad2f3ed 100644 --- a/pkg/providers/openai_compat/provider.go +++ b/pkg/providers/openai_compat/provider.go @@ -1,6 +1,7 @@ package openai_compat import ( + "bufio" "bytes" "context" "encoding/json" @@ -30,17 +31,36 @@ type ( type Provider struct { apiKey string apiBase string + endpointPath string // API path appended to apiBase (default: "/chat/completions") maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models) + stream bool // Use SSE streaming internally (accumulates into a single LLMResponse) httpClient *http.Client } +// Options configures optional behaviour for the provider. +type Options struct { + EndpointPath string // API path appended to apiBase (default: "/chat/completions") + MaxTokensField string // Field name for max tokens parameter + Stream bool // Use SSE streaming internally +} + func NewProvider(apiKey, apiBase, proxy string) *Provider { return NewProviderWithMaxTokensField(apiKey, apiBase, proxy, "") } func NewProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField string) *Provider { + return NewProviderWithOptions(apiKey, apiBase, proxy, Options{ + MaxTokensField: maxTokensField, + }) +} + +func NewProviderWithOptions(apiKey, apiBase, proxy string, opts Options) *Provider { + timeout := 120 * time.Second + if opts.Stream { + timeout = 5 * time.Minute + } client := &http.Client{ - Timeout: 120 * time.Second, + Timeout: timeout, } if proxy != "" { @@ -54,10 +74,17 @@ func NewProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField string } } + endpointPath := opts.EndpointPath + if endpointPath == "" { + endpointPath = "/chat/completions" + } + return &Provider{ apiKey: apiKey, apiBase: strings.TrimRight(apiBase, "/"), - maxTokensField: maxTokensField, + endpointPath: endpointPath, + maxTokensField: opts.MaxTokensField, + stream: opts.Stream, httpClient: client, } } @@ -111,12 +138,16 @@ func (p *Provider) Chat( } } + if p.stream { + requestBody["stream"] = true + } + jsonData, err := json.Marshal(requestBody) if err != nil { return nil, fmt.Errorf("failed to marshal request: %w", err) } - req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+"/chat/completions", bytes.NewReader(jsonData)) + req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+p.endpointPath, bytes.NewReader(jsonData)) if err != nil { return nil, fmt.Errorf("failed to create request: %w", err) } @@ -132,15 +163,20 @@ func (p *Provider) Chat( } defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("API request failed:\n Status: %d\n Body: %s", resp.StatusCode, string(body)) + } + + if p.stream { + return parseStreamResponse(resp.Body) + } + body, err := io.ReadAll(resp.Body) if err != nil { return nil, fmt.Errorf("failed to read response: %w", err) } - if resp.StatusCode != http.StatusOK { - return nil, fmt.Errorf("API request failed:\n Status: %d\n Body: %s", resp.StatusCode, string(body)) - } - return parseResponse(body) } @@ -240,7 +276,7 @@ func normalizeModel(model, apiBase string) string { prefix := strings.ToLower(model[:idx]) switch prefix { - case "moonshot", "nvidia", "groq", "ollama", "deepseek", "google", "openrouter", "zhipu": + case "moonshot", "nvidia", "groq", "ollama", "deepseek", "google", "openrouter", "zhipu", "minimax": return model[idx+1:] default: return model @@ -276,3 +312,129 @@ func asFloat(v any) (float64, bool) { return 0, false } } + +// --- SSE streaming support --- + +type streamChunk struct { + Choices []streamChoice `json:"choices"` + Usage *UsageInfo `json:"usage"` +} + +type streamChoice struct { + Delta streamDelta `json:"delta"` + FinishReason string `json:"finish_reason"` +} + +type streamDelta struct { + Content string `json:"content"` + ToolCalls []streamDeltaTC `json:"tool_calls"` +} + +type streamDeltaTC struct { + Index int `json:"index"` + ID string `json:"id"` + Type string `json:"type"` + Function *streamDeltaFunction `json:"function"` +} + +type streamDeltaFunction struct { + Name string `json:"name"` + Arguments string `json:"arguments"` +} + +type streamToolCallAcc struct { + ID string + Name string + Arguments strings.Builder +} + +// parseStreamResponse reads an SSE (text/event-stream) response and +// accumulates it into a single LLMResponse. +func parseStreamResponse(r io.Reader) (*LLMResponse, error) { + scanner := bufio.NewScanner(r) + // Allow up to 1 MB per SSE line to handle large argument deltas. + scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) + + var content strings.Builder + var toolCalls []streamToolCallAcc + var finishReason string + var usage *UsageInfo + + for scanner.Scan() { + line := scanner.Text() + if !strings.HasPrefix(line, "data: ") { + continue + } + data := strings.TrimPrefix(line, "data: ") + if data == "[DONE]" { + break + } + + var chunk streamChunk + if err := json.Unmarshal([]byte(data), &chunk); err != nil { + continue // skip malformed chunks + } + + if len(chunk.Choices) == 0 { + if chunk.Usage != nil { + usage = chunk.Usage + } + continue + } + + choice := chunk.Choices[0] + if choice.Delta.Content != "" { + content.WriteString(choice.Delta.Content) + } + if choice.FinishReason != "" { + finishReason = choice.FinishReason + } + + // Accumulate streaming tool calls by index. + for _, tc := range choice.Delta.ToolCalls { + for len(toolCalls) <= tc.Index { + toolCalls = append(toolCalls, streamToolCallAcc{}) + } + if tc.ID != "" { + toolCalls[tc.Index].ID = tc.ID + } + if tc.Function != nil { + if tc.Function.Name != "" { + toolCalls[tc.Index].Name = tc.Function.Name + } + toolCalls[tc.Index].Arguments.WriteString(tc.Function.Arguments) + } + } + + if chunk.Usage != nil { + usage = chunk.Usage + } + } + if err := scanner.Err(); err != nil { + return nil, fmt.Errorf("reading stream: %w", err) + } + + result := &LLMResponse{ + Content: content.String(), + FinishReason: finishReason, + Usage: usage, + } + + for _, tc := range toolCalls { + arguments := make(map[string]any) + argStr := tc.Arguments.String() + if argStr != "" { + if err := json.Unmarshal([]byte(argStr), &arguments); err != nil { + log.Printf("openai_compat: failed to decode streamed tool call arguments for %q: %v", tc.Name, err) + arguments["raw"] = argStr + } + } + result.ToolCalls = append(result.ToolCalls, ToolCall{ + ID: tc.ID, + Name: tc.Name, + Arguments: arguments, + }) + } + + return result, nil +} diff --git a/pkg/providers/openai_compat/provider_test.go b/pkg/providers/openai_compat/provider_test.go index 42f9d42ab..64927d85c 100644 --- a/pkg/providers/openai_compat/provider_test.go +++ b/pkg/providers/openai_compat/provider_test.go @@ -2,6 +2,7 @@ package openai_compat import ( "encoding/json" + "fmt" "net/http" "net/http/httptest" "net/url" @@ -281,3 +282,126 @@ func TestNormalizeModel_UsesAPIBase(t *testing.T) { t.Fatalf("normalizeModel(openrouter) = %q, want %q", got, "openrouter/auto") } } + +func TestProviderChat_StreamingTextResponse(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/text/chatcompletion_v2" { + http.Error(w, "not found", http.StatusNotFound) + return + } + var body map[string]any + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + if body["stream"] != true { + t.Error("expected stream=true in request body") + } + + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(http.StatusOK) + flusher, _ := w.(http.Flusher) + + chunks := []string{ + `data: {"choices":[{"delta":{"content":"Hello"},"finish_reason":""}]}`, + `data: {"choices":[{"delta":{"content":" world"},"finish_reason":""}]}`, + `data: {"choices":[{"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":5,"completion_tokens":2,"total_tokens":7}}`, + `data: [DONE]`, + } + for _, c := range chunks { + fmt.Fprintln(w, c) + fmt.Fprintln(w) // blank line between events + if flusher != nil { + flusher.Flush() + } + } + })) + defer server.Close() + + p := NewProviderWithOptions("key", server.URL, "", Options{ + EndpointPath: "/text/chatcompletion_v2", + Stream: true, + }) + out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "MiniMax-M1", nil) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + if out.Content != "Hello world" { + t.Fatalf("Content = %q, want %q", out.Content, "Hello world") + } + if out.FinishReason != "stop" { + t.Fatalf("FinishReason = %q, want %q", out.FinishReason, "stop") + } + if out.Usage == nil || out.Usage.TotalTokens != 7 { + t.Fatalf("Usage.TotalTokens = %v, want 7", out.Usage) + } +} + +func TestProviderChat_StreamingToolCalls(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(http.StatusOK) + flusher, _ := w.(http.Flusher) + + chunks := []string{ + `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"get_weather","arguments":""}}]},"finish_reason":""}]}`, + `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"city\":"}}]},"finish_reason":""}]}`, + `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"SF\"}"}}]},"finish_reason":""}]}`, + `data: {"choices":[{"delta":{},"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":10,"completion_tokens":8,"total_tokens":18}}`, + `data: [DONE]`, + } + for _, c := range chunks { + fmt.Fprintln(w, c) + fmt.Fprintln(w) + if flusher != nil { + flusher.Flush() + } + } + })) + defer server.Close() + + p := NewProviderWithOptions("key", server.URL, "", Options{Stream: true}) + out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "weather?"}}, nil, "test", 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)) + } + tc := out.ToolCalls[0] + if tc.ID != "call_1" { + t.Fatalf("ToolCalls[0].ID = %q, want %q", tc.ID, "call_1") + } + if tc.Name != "get_weather" { + t.Fatalf("ToolCalls[0].Name = %q, want %q", tc.Name, "get_weather") + } + if tc.Arguments["city"] != "SF" { + t.Fatalf("ToolCalls[0].Arguments[city] = %v, want SF", tc.Arguments["city"]) + } +} + +func TestProviderChat_CustomEndpointPath(t *testing.T) { + var hitPath string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hitPath = r.URL.Path + resp := map[string]any{ + "choices": []map[string]any{ + {"message": map[string]any{"content": "ok"}, "finish_reason": "stop"}, + }, + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) + })) + defer server.Close() + + p := NewProviderWithOptions("key", server.URL, "", Options{ + EndpointPath: "/text/chatcompletion_v2", + }) + _, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "test", nil) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + if hitPath != "/text/chatcompletion_v2" { + t.Fatalf("endpoint path = %q, want %q", hitPath, "/text/chatcompletion_v2") + } +}