feat: add MiniMax provider with SSE streaming support
Add minimax protocol with non-standard endpoint path (/text/chatcompletion_v2) and SSE streaming accumulation. The openai_compat provider gains configurable EndpointPath, Stream options, and a full SSE parser that handles incremental tool-call assembly. Add "stream" bool field to ModelConfig so any OpenAI-compatible provider can opt into streaming via config.json. MiniMax defaults to stream=true; others default to false. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
7f0c1cdc57
commit
cda7df202c
6 changed files with 341 additions and 10 deletions
|
|
@ -35,6 +35,12 @@
|
||||||
"model": "deepseek/deepseek-chat",
|
"model": "deepseek/deepseek-chat",
|
||||||
"api_key": "sk-your-deepseek-key"
|
"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_name": "loadbalanced-gpt4",
|
||||||
"model": "openai/gpt-5.2",
|
"model": "openai/gpt-5.2",
|
||||||
|
|
|
||||||
|
|
@ -397,6 +397,7 @@ type ModelConfig struct {
|
||||||
// Optional optimizations
|
// Optional optimizations
|
||||||
RPM int `json:"rpm,omitempty"` // Requests per minute limit
|
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")
|
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.
|
// Validate checks if the ModelConfig has all required fields.
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ import (
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"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.
|
// 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 == "" {
|
if apiBase == "" {
|
||||||
apiBase = getDefaultAPIBase(protocol)
|
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",
|
case "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
||||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||||
|
|
@ -97,7 +116,10 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
if apiBase == "" {
|
if apiBase == "" {
|
||||||
apiBase = getDefaultAPIBase(protocol)
|
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":
|
case "anthropic":
|
||||||
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
|
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
|
||||||
|
|
@ -186,7 +208,17 @@ func getDefaultAPIBase(protocol string) string {
|
||||||
return "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
return "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||||
case "vllm":
|
case "vllm":
|
||||||
return "http://localhost:8000/v1"
|
return "http://localhost:8000/v1"
|
||||||
|
case "minimax":
|
||||||
|
return "https://api.minimax.io/v1"
|
||||||
default:
|
default:
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// boolDefault dereferences a *bool, returning def when nil.
|
||||||
|
func boolDefault(p *bool, def bool) bool {
|
||||||
|
if p != nil {
|
||||||
|
return *p
|
||||||
|
}
|
||||||
|
return def
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
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)
|
resp, err := p.delegate.Chat(ctx, messages, tools, model, options)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
package openai_compat
|
package openai_compat
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
|
@ -30,17 +31,36 @@ type (
|
||||||
type Provider struct {
|
type Provider struct {
|
||||||
apiKey string
|
apiKey string
|
||||||
apiBase 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)
|
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
|
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 {
|
func NewProvider(apiKey, apiBase, proxy string) *Provider {
|
||||||
return NewProviderWithMaxTokensField(apiKey, apiBase, proxy, "")
|
return NewProviderWithMaxTokensField(apiKey, apiBase, proxy, "")
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField string) *Provider {
|
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{
|
client := &http.Client{
|
||||||
Timeout: 120 * time.Second,
|
Timeout: timeout,
|
||||||
}
|
}
|
||||||
|
|
||||||
if proxy != "" {
|
if proxy != "" {
|
||||||
|
|
@ -54,10 +74,17 @@ func NewProviderWithMaxTokensField(apiKey, apiBase, proxy, maxTokensField string
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
endpointPath := opts.EndpointPath
|
||||||
|
if endpointPath == "" {
|
||||||
|
endpointPath = "/chat/completions"
|
||||||
|
}
|
||||||
|
|
||||||
return &Provider{
|
return &Provider{
|
||||||
apiKey: apiKey,
|
apiKey: apiKey,
|
||||||
apiBase: strings.TrimRight(apiBase, "/"),
|
apiBase: strings.TrimRight(apiBase, "/"),
|
||||||
maxTokensField: maxTokensField,
|
endpointPath: endpointPath,
|
||||||
|
maxTokensField: opts.MaxTokensField,
|
||||||
|
stream: opts.Stream,
|
||||||
httpClient: client,
|
httpClient: client,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -111,12 +138,16 @@ func (p *Provider) Chat(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if p.stream {
|
||||||
|
requestBody["stream"] = true
|
||||||
|
}
|
||||||
|
|
||||||
jsonData, err := json.Marshal(requestBody)
|
jsonData, err := json.Marshal(requestBody)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
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 {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -132,15 +163,20 @@ func (p *Provider) Chat(
|
||||||
}
|
}
|
||||||
defer resp.Body.Close()
|
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)
|
body, err := io.ReadAll(resp.Body)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
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)
|
return parseResponse(body)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -240,7 +276,7 @@ func normalizeModel(model, apiBase string) string {
|
||||||
|
|
||||||
prefix := strings.ToLower(model[:idx])
|
prefix := strings.ToLower(model[:idx])
|
||||||
switch prefix {
|
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:]
|
return model[idx+1:]
|
||||||
default:
|
default:
|
||||||
return model
|
return model
|
||||||
|
|
@ -276,3 +312,129 @@ func asFloat(v any) (float64, bool) {
|
||||||
return 0, false
|
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
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ package openai_compat
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
|
@ -281,3 +282,126 @@ func TestNormalizeModel_UsesAPIBase(t *testing.T) {
|
||||||
t.Fatalf("normalizeModel(openrouter) = %q, want %q", got, "openrouter/auto")
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue