From bf1707366b46be175c279fb614806eadc98e7259 Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?=E6=9D=8E=E9=BE=99=200668001470?=
Date: Mon, 16 Mar 2026 09:16:33 +0800
Subject: [PATCH] fix(providers): sync cloudflare model IDs with new provider
matrix
---
pkg/providers/anthropic_messages/provider.go | 415 ++++++++++++
.../anthropic_messages/provider_test.go | 622 ++++++++++++++++++
pkg/providers/azure/provider.go | 150 +++++
pkg/providers/azure/provider_test.go | 232 +++++++
pkg/providers/common/common.go | 380 +++++++++++
pkg/providers/common/common_test.go | 558 ++++++++++++++++
pkg/providers/factory_provider.go | 44 +-
pkg/providers/factory_provider_test.go | 120 ++++
8 files changed, 2519 insertions(+), 2 deletions(-)
create mode 100644 pkg/providers/anthropic_messages/provider.go
create mode 100644 pkg/providers/anthropic_messages/provider_test.go
create mode 100644 pkg/providers/azure/provider.go
create mode 100644 pkg/providers/azure/provider_test.go
create mode 100644 pkg/providers/common/common.go
create mode 100644 pkg/providers/common/common_test.go
diff --git a/pkg/providers/anthropic_messages/provider.go b/pkg/providers/anthropic_messages/provider.go
new file mode 100644
index 000000000..8a83a7058
--- /dev/null
+++ b/pkg/providers/anthropic_messages/provider.go
@@ -0,0 +1,415 @@
+// PicoClaw - Ultra-lightweight personal AI agent
+// License: MIT
+//
+// Copyright (c) 2026 PicoClaw contributors
+
+package anthropicmessages
+
+import (
+ "bytes"
+ "context"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "net/url"
+ "strings"
+ "time"
+
+ "github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
+)
+
+type (
+ ToolCall = protocoltypes.ToolCall
+ FunctionCall = protocoltypes.FunctionCall
+ LLMResponse = protocoltypes.LLMResponse
+ UsageInfo = protocoltypes.UsageInfo
+ Message = protocoltypes.Message
+ ToolDefinition = protocoltypes.ToolDefinition
+ ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
+)
+
+const (
+ defaultAPIVersion = "2023-06-01"
+ defaultBaseURL = "https://api.anthropic.com/v1"
+ defaultRequestTimeout = 120 * time.Second
+)
+
+// Provider implements Anthropic Messages API via HTTP (without SDK).
+// It supports custom endpoints that use Anthropic's native message format.
+type Provider struct {
+ apiKey string
+ apiBase string
+ httpClient *http.Client
+}
+
+// NewProvider creates a new Anthropic Messages API provider.
+func NewProvider(apiKey, apiBase string) *Provider {
+ return NewProviderWithTimeout(apiKey, apiBase, 0)
+}
+
+// NewProviderWithTimeout creates a provider with custom request timeout.
+func NewProviderWithTimeout(apiKey, apiBase string, timeoutSeconds int) *Provider {
+ baseURL := normalizeBaseURL(apiBase)
+ timeout := defaultRequestTimeout
+ if timeoutSeconds > 0 {
+ timeout = time.Duration(timeoutSeconds) * time.Second
+ }
+
+ return &Provider{
+ apiKey: apiKey,
+ apiBase: baseURL,
+ httpClient: &http.Client{
+ Timeout: timeout,
+ },
+ }
+}
+
+// Chat sends messages to the Anthropic Messages API and returns the response.
+func (p *Provider) Chat(
+ ctx context.Context,
+ messages []Message,
+ tools []ToolDefinition,
+ model string,
+ options map[string]any,
+) (*LLMResponse, error) {
+ if p.apiKey == "" {
+ return nil, fmt.Errorf("API key not configured")
+ }
+
+ // Build request body
+ requestBody, err := buildRequestBody(messages, tools, model, options)
+ if err != nil {
+ return nil, fmt.Errorf("building request body: %w", err)
+ }
+
+ // Serialize to JSON
+ jsonBody, err := json.Marshal(requestBody)
+ if err != nil {
+ return nil, fmt.Errorf("serializing request body: %w", err)
+ }
+
+ // Build request URL
+ endpointURL, err := url.JoinPath(p.apiBase, "messages")
+ if err != nil {
+ return nil, fmt.Errorf("building endpoint URL: %w", err)
+ }
+
+ // Create HTTP request
+ req, err := http.NewRequestWithContext(ctx, "POST", endpointURL, bytes.NewReader(jsonBody))
+ if err != nil {
+ return nil, fmt.Errorf("creating HTTP request: %w", err)
+ }
+
+ // Set headers
+ req.Header.Set("Content-Type", "application/json")
+ req.Header.Set("X-API-Key", p.apiKey) //nolint:canonicalheader // Anthropic API requires exact header name
+ req.Header.Set("Anthropic-Version", defaultAPIVersion)
+
+ // Execute request
+ resp, err := p.httpClient.Do(req)
+ if err != nil {
+ return nil, fmt.Errorf("executing HTTP request: %w", err)
+ }
+ defer resp.Body.Close()
+
+ // Read response body
+ body, err := io.ReadAll(resp.Body)
+ if err != nil {
+ return nil, fmt.Errorf("reading response body: %w", err)
+ }
+
+ // Check for HTTP errors with detailed messages
+ switch resp.StatusCode {
+ case http.StatusUnauthorized:
+ return nil, fmt.Errorf("authentication failed (401): check your API key")
+ case http.StatusTooManyRequests:
+ return nil, fmt.Errorf("rate limited (429): %s", string(body))
+ case http.StatusBadRequest:
+ return nil, fmt.Errorf("bad request (400): %s", string(body))
+ case http.StatusNotFound:
+ return nil, fmt.Errorf("endpoint not found (404): %s", string(body))
+ case http.StatusInternalServerError:
+ return nil, fmt.Errorf("internal server error (500): %s", string(body))
+ case http.StatusServiceUnavailable:
+ return nil, fmt.Errorf("service unavailable (503): %s", string(body))
+ default:
+ if resp.StatusCode != http.StatusOK {
+ return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body))
+ }
+ }
+
+ // Parse response
+ return parseResponseBody(body)
+}
+
+// GetDefaultModel returns the default model for this provider.
+func (p *Provider) GetDefaultModel() string {
+ return "claude-sonnet-4.6"
+}
+
+// buildRequestBody converts internal message format to Anthropic Messages API format.
+func buildRequestBody(
+ messages []Message,
+ tools []ToolDefinition,
+ model string,
+ options map[string]any,
+) (map[string]any, error) {
+ // max_tokens is required and guaranteed by agent loop
+ maxTokens, ok := asInt(options["max_tokens"])
+ if !ok {
+ return nil, fmt.Errorf("max_tokens is required in options")
+ }
+
+ result := map[string]any{
+ "model": model,
+ "max_tokens": int64(maxTokens),
+ "messages": []any{},
+ }
+
+ // Set temperature from options
+ if temp, ok := asFloat(options["temperature"]); ok {
+ result["temperature"] = temp
+ }
+
+ // Process messages
+ var systemPrompt string
+ var apiMessages []any
+
+ for _, msg := range messages {
+ switch msg.Role {
+ case "system":
+ // Accumulate system messages
+ if systemPrompt != "" {
+ systemPrompt += "\n\n" + msg.Content
+ } else {
+ systemPrompt = msg.Content
+ }
+
+ case "user":
+ if msg.ToolCallID != "" {
+ // Tool result message
+ content := []map[string]any{
+ {
+ "type": "tool_result",
+ "tool_use_id": msg.ToolCallID,
+ "content": msg.Content,
+ },
+ }
+ apiMessages = append(apiMessages, map[string]any{
+ "role": "user",
+ "content": content,
+ })
+ } else {
+ // Regular user message
+ apiMessages = append(apiMessages, map[string]any{
+ "role": "user",
+ "content": msg.Content,
+ })
+ }
+
+ case "assistant":
+ content := []any{}
+
+ // Add text content if present
+ if msg.Content != "" {
+ content = append(content, map[string]any{
+ "type": "text",
+ "text": msg.Content,
+ })
+ }
+
+ // Add tool_use blocks
+ for _, tc := range msg.ToolCalls {
+ toolUse := map[string]any{
+ "type": "tool_use",
+ "id": tc.ID,
+ "name": tc.Name,
+ "input": tc.Arguments,
+ }
+ content = append(content, toolUse)
+ }
+
+ apiMessages = append(apiMessages, map[string]any{
+ "role": "assistant",
+ "content": content,
+ })
+
+ case "tool":
+ // Tool result (alternative format)
+ content := []map[string]any{
+ {
+ "type": "tool_result",
+ "tool_use_id": msg.ToolCallID,
+ "content": msg.Content,
+ },
+ }
+ apiMessages = append(apiMessages, map[string]any{
+ "role": "user",
+ "content": content,
+ })
+ }
+ }
+
+ result["messages"] = apiMessages
+
+ // Set system prompt if present
+ if systemPrompt != "" {
+ result["system"] = systemPrompt
+ }
+
+ // Add tools if present
+ if len(tools) > 0 {
+ result["tools"] = buildTools(tools)
+ }
+
+ return result, nil
+}
+
+// buildTools converts tool definitions to Anthropic format.
+func buildTools(tools []ToolDefinition) []any {
+ result := make([]any, len(tools))
+ for i, tool := range tools {
+ toolDef := map[string]any{
+ "name": tool.Function.Name,
+ "description": tool.Function.Description,
+ "input_schema": tool.Function.Parameters,
+ }
+ result[i] = toolDef
+ }
+ return result
+}
+
+// parseResponseBody parses Anthropic Messages API response.
+func parseResponseBody(body []byte) (*LLMResponse, error) {
+ var resp anthropicMessageResponse
+ if err := json.Unmarshal(body, &resp); err != nil {
+ return nil, fmt.Errorf("parsing JSON response: %w", err)
+ }
+
+ // Extract content and tool calls
+ var content strings.Builder
+ toolCalls := make([]ToolCall, 0) // Initialize as empty slice (not nil) for consistent JSON serialization
+
+ for _, block := range resp.Content {
+ switch block.Type {
+ case "text":
+ content.WriteString(block.Text)
+ case "tool_use":
+ argsJSON, _ := json.Marshal(block.Input)
+ toolCalls = append(toolCalls, ToolCall{
+ ID: block.ID,
+ Name: block.Name,
+ Arguments: block.Input,
+ Function: &FunctionCall{
+ Name: block.Name,
+ Arguments: string(argsJSON),
+ },
+ })
+ }
+ }
+
+ // Map stop_reason
+ finishReason := "stop"
+ switch resp.StopReason {
+ case "tool_use":
+ finishReason = "tool_calls"
+ case "max_tokens":
+ finishReason = "length"
+ case "end_turn":
+ finishReason = "stop"
+ case "stop_sequence":
+ finishReason = "stop"
+ }
+
+ return &LLMResponse{
+ Content: content.String(),
+ ToolCalls: toolCalls,
+ FinishReason: finishReason,
+ Usage: &UsageInfo{
+ PromptTokens: int(resp.Usage.InputTokens),
+ CompletionTokens: int(resp.Usage.OutputTokens),
+ TotalTokens: int(resp.Usage.InputTokens + resp.Usage.OutputTokens),
+ },
+ }, nil
+}
+
+// normalizeBaseURL ensures the base URL is properly formatted.
+// It removes /v1 suffix if present (to avoid duplication) and always appends /v1.
+// This handles edge cases like "https://api.example.com/v1/proxy" correctly.
+func normalizeBaseURL(apiBase string) string {
+ base := strings.TrimSpace(apiBase)
+ if base == "" {
+ return defaultBaseURL
+ }
+
+ // Remove trailing slashes
+ base = strings.TrimRight(base, "/")
+
+ // Remove /v1 suffix if present (will be re-added)
+ // This prevents duplication for URLs like "https://api.example.com/v1/proxy"
+ if before, ok := strings.CutSuffix(base, "/v1"); ok {
+ base = before
+ }
+
+ // Ensure we don't have an empty string after cutting
+ if base == "" {
+ return defaultBaseURL
+ }
+
+ // Add /v1 suffix (required by Anthropic Messages API)
+ return base + "/v1"
+}
+
+// Helper functions for type conversion
+
+func asInt(v any) (int, bool) {
+ switch val := v.(type) {
+ case int:
+ return val, true
+ case float64:
+ return int(val), true
+ case int64:
+ return int(val), true
+ default:
+ return 0, false
+ }
+}
+
+func asFloat(v any) (float64, bool) {
+ switch val := v.(type) {
+ case float64:
+ return val, true
+ case int:
+ return float64(val), true
+ case int64:
+ return float64(val), true
+ default:
+ return 0, false
+ }
+}
+
+// Anthropic API response structures
+
+type anthropicMessageResponse struct {
+ ID string `json:"id"`
+ Type string `json:"type"`
+ Role string `json:"role"`
+ Content []contentBlock `json:"content"`
+ StopReason string `json:"stop_reason"`
+ Model string `json:"model"`
+ Usage usageInfo `json:"usage"`
+}
+
+type contentBlock struct {
+ Type string `json:"type"`
+ Text string `json:"text,omitempty"`
+ ID string `json:"id,omitempty"`
+ Name string `json:"name,omitempty"`
+ Input map[string]any `json:"input,omitempty"`
+}
+
+type usageInfo struct {
+ InputTokens int64 `json:"input_tokens"`
+ OutputTokens int64 `json:"output_tokens"`
+}
diff --git a/pkg/providers/anthropic_messages/provider_test.go b/pkg/providers/anthropic_messages/provider_test.go
new file mode 100644
index 000000000..da4213e92
--- /dev/null
+++ b/pkg/providers/anthropic_messages/provider_test.go
@@ -0,0 +1,622 @@
+// PicoClaw - Ultra-lightweight personal AI agent
+// License: MIT
+//
+// Copyright (c) 2026 PicoClaw contributors
+
+package anthropicmessages
+
+import (
+ "context"
+ "encoding/json"
+ "reflect"
+ "strings"
+ "testing"
+)
+
+func TestBuildRequestBody(t *testing.T) {
+ tests := []struct {
+ name string
+ messages []Message
+ tools []ToolDefinition
+ model string
+ options map[string]any
+ want map[string]any
+ wantErr bool
+ }{
+ {
+ name: "basic user message",
+ messages: []Message{
+ {Role: "user", Content: "Hello, world!"},
+ },
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 8192,
+ },
+ want: map[string]any{
+ "model": "test-model",
+ "max_tokens": int64(8192),
+ "messages": []any{
+ map[string]any{
+ "role": "user",
+ "content": "Hello, world!",
+ },
+ },
+ },
+ },
+ {
+ name: "user and assistant messages",
+ messages: []Message{
+ {Role: "user", Content: "What is 2+2?"},
+ {Role: "assistant", Content: "4"},
+ },
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 8192,
+ },
+ want: map[string]any{
+ "model": "test-model",
+ "max_tokens": int64(8192),
+ "messages": []any{
+ map[string]any{
+ "role": "user",
+ "content": "What is 2+2?",
+ },
+ map[string]any{
+ "role": "assistant",
+ "content": []any{
+ map[string]any{
+ "type": "text",
+ "text": "4",
+ },
+ },
+ },
+ },
+ },
+ },
+ {
+ name: "with system message",
+ messages: []Message{
+ {Role: "system", Content: "You are a helpful assistant."},
+ {Role: "user", Content: "Hello"},
+ },
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 8192,
+ },
+ want: map[string]any{
+ "model": "test-model",
+ "max_tokens": int64(8192),
+ "system": "You are a helpful assistant.",
+ "messages": []any{
+ map[string]any{
+ "role": "user",
+ "content": "Hello",
+ },
+ },
+ },
+ },
+ {
+ name: "with custom max_tokens and temperature",
+ messages: []Message{
+ {Role: "user", Content: "Test"},
+ },
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 2048,
+ "temperature": 0.5,
+ },
+ want: map[string]any{
+ "model": "test-model",
+ "max_tokens": int64(2048),
+ "temperature": 0.5,
+ "messages": []any{
+ map[string]any{
+ "role": "user",
+ "content": "Test",
+ },
+ },
+ },
+ },
+ {
+ name: "missing max_tokens returns error",
+ messages: []Message{
+ {Role: "user", Content: "Test"},
+ },
+ model: "test-model",
+ options: map[string]any{},
+ want: nil,
+ wantErr: true,
+ },
+ {
+ name: "with tools",
+ messages: []Message{
+ {Role: "user", Content: "What's the weather?"},
+ },
+ tools: []ToolDefinition{
+ {
+ Function: ToolFunctionDefinition{
+ Name: "get_weather",
+ Description: "Get current weather",
+ Parameters: map[string]any{
+ "type": "object",
+ "properties": map[string]any{
+ "location": map[string]any{
+ "type": "string",
+ "description": "City name",
+ },
+ },
+ },
+ },
+ },
+ },
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 8192,
+ },
+ want: map[string]any{
+ "model": "test-model",
+ "max_tokens": int64(8192),
+ "messages": []any{
+ map[string]any{
+ "role": "user",
+ "content": "What's the weather?",
+ },
+ },
+ "tools": []any{
+ map[string]any{
+ "name": "get_weather",
+ "description": "Get current weather",
+ "input_schema": map[string]any{
+ "type": "object",
+ "properties": map[string]any{
+ "location": map[string]any{
+ "type": "string",
+ "description": "City name",
+ },
+ },
+ },
+ },
+ },
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got, err := buildRequestBody(tt.messages, tt.tools, tt.model, tt.options)
+ if (err != nil) != tt.wantErr {
+ t.Errorf("buildRequestBody() error = %v, wantErr %v", err, tt.wantErr)
+ return
+ }
+ if !reflect.DeepEqual(got, tt.want) {
+ gotJSON, _ := json.MarshalIndent(got, "", " ")
+ wantJSON, _ := json.MarshalIndent(tt.want, "", " ")
+ t.Errorf("buildRequestBody() mismatch:\ngot:\n%s\nwant:\n%s", gotJSON, wantJSON)
+ }
+ })
+ }
+}
+
+func TestParseResponseBody(t *testing.T) {
+ tests := []struct {
+ name string
+ body []byte
+ want *LLMResponse
+ wantErr bool
+ }{
+ {
+ name: "basic text response",
+ body: []byte(`{
+ "id": "msg-123",
+ "type": "message",
+ "role": "assistant",
+ "content": [
+ {"type": "text", "text": "Hello, how can I help?"}
+ ],
+ "stop_reason": "end_turn",
+ "model": "test-model",
+ "usage": {
+ "input_tokens": 10,
+ "output_tokens": 5
+ }
+ }`),
+ want: &LLMResponse{
+ Content: "Hello, how can I help?",
+ ToolCalls: []ToolCall{},
+ FinishReason: "stop",
+ Usage: &UsageInfo{
+ PromptTokens: 10,
+ CompletionTokens: 5,
+ TotalTokens: 15,
+ },
+ Reasoning: "",
+ ReasoningDetails: nil,
+ },
+ wantErr: false,
+ },
+ {
+ name: "response with tool use",
+ body: []byte(`{
+ "id": "msg-456",
+ "type": "message",
+ "role": "assistant",
+ "content": [
+ {"type": "text", "text": "I'll check the weather for you."},
+ {
+ "type": "tool_use",
+ "id": "toolu-123",
+ "name": "get_weather",
+ "input": {"location": "Tokyo"}
+ }
+ ],
+ "stop_reason": "tool_use",
+ "model": "test-model",
+ "usage": {
+ "input_tokens": 20,
+ "output_tokens": 15
+ }
+ }`),
+ want: &LLMResponse{
+ Content: "I'll check the weather for you.",
+ ToolCalls: []ToolCall{
+ {
+ ID: "toolu-123",
+ Name: "get_weather",
+ Arguments: map[string]any{
+ "location": "Tokyo",
+ },
+ Function: &FunctionCall{
+ Name: "get_weather",
+ Arguments: `{"location":"Tokyo"}`,
+ },
+ },
+ },
+ FinishReason: "tool_calls",
+ Usage: &UsageInfo{
+ PromptTokens: 20,
+ CompletionTokens: 15,
+ TotalTokens: 35,
+ },
+ Reasoning: "",
+ ReasoningDetails: nil,
+ },
+ wantErr: false,
+ },
+ {
+ name: "invalid JSON",
+ body: []byte(`invalid json`),
+ want: nil,
+ wantErr: true,
+ },
+ {
+ name: "max_tokens stop reason",
+ body: []byte(`{
+ "id": "msg-789",
+ "type": "message",
+ "role": "assistant",
+ "content": [
+ {"type": "text", "text": "Partial response"}
+ ],
+ "stop_reason": "max_tokens",
+ "model": "test-model",
+ "usage": {
+ "input_tokens": 100,
+ "output_tokens": 4096
+ }
+ }`),
+ want: &LLMResponse{
+ Content: "Partial response",
+ ToolCalls: []ToolCall{},
+ FinishReason: "length",
+ Usage: &UsageInfo{
+ PromptTokens: 100,
+ CompletionTokens: 4096,
+ TotalTokens: 4196,
+ },
+ Reasoning: "",
+ ReasoningDetails: nil,
+ },
+ wantErr: false,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got, err := parseResponseBody(tt.body)
+ if (err != nil) != tt.wantErr {
+ t.Errorf("parseResponseBody() error = %v, wantErr %v", err, tt.wantErr)
+ return
+ }
+ if err != nil {
+ return
+ }
+
+ // Compare individual fields
+ if got.Content != tt.want.Content {
+ t.Errorf("Content = %q, want %q", got.Content, tt.want.Content)
+ }
+ if got.FinishReason != tt.want.FinishReason {
+ t.Errorf("FinishReason = %q, want %q", got.FinishReason, tt.want.FinishReason)
+ }
+ if got.Usage == nil && tt.want.Usage != nil {
+ t.Errorf("Usage = nil, want non-nil")
+ } else if got.Usage != nil && tt.want.Usage == nil {
+ t.Errorf("Usage = non-nil, want nil")
+ } else if got.Usage != nil && tt.want.Usage != nil {
+ if got.Usage.PromptTokens != tt.want.Usage.PromptTokens {
+ t.Errorf("Usage.PromptTokens = %d, want %d", got.Usage.PromptTokens, tt.want.Usage.PromptTokens)
+ }
+ if got.Usage.CompletionTokens != tt.want.Usage.CompletionTokens {
+ t.Errorf("Usage.CompletionTokens = %d, want %d",
+ got.Usage.CompletionTokens, tt.want.Usage.CompletionTokens)
+ }
+ if got.Usage.TotalTokens != tt.want.Usage.TotalTokens {
+ t.Errorf("Usage.TotalTokens = %d, want %d", got.Usage.TotalTokens, tt.want.Usage.TotalTokens)
+ }
+ }
+ if len(got.ToolCalls) != len(tt.want.ToolCalls) {
+ t.Errorf("ToolCalls length = %d, want %d", len(got.ToolCalls), len(tt.want.ToolCalls))
+ } else {
+ for i := range got.ToolCalls {
+ if got.ToolCalls[i].ID != tt.want.ToolCalls[i].ID {
+ t.Errorf("ToolCalls[%d].ID = %q, want %q",
+ i, got.ToolCalls[i].ID, tt.want.ToolCalls[i].ID)
+ }
+ if got.ToolCalls[i].Name != tt.want.ToolCalls[i].Name {
+ t.Errorf("ToolCalls[%d].Name = %q, want %q",
+ i, got.ToolCalls[i].Name, tt.want.ToolCalls[i].Name)
+ }
+ }
+ }
+ })
+ }
+}
+
+func TestNormalizeBaseURL(t *testing.T) {
+ tests := []struct {
+ name string
+ apiBase string
+ expected string
+ }{
+ {
+ name: "empty string defaults to official API",
+ apiBase: "",
+ expected: "https://api.anthropic.com/v1",
+ },
+ {
+ name: "URL without /v1 gets it appended",
+ apiBase: "https://api.example.com/anthropic",
+ expected: "https://api.example.com/anthropic/v1",
+ },
+ {
+ name: "URL with /v1 remains unchanged",
+ apiBase: "https://api.example.com/v1",
+ expected: "https://api.example.com/v1",
+ },
+ {
+ name: "URL with trailing slash gets cleaned",
+ apiBase: "https://api.example.com/anthropic/",
+ expected: "https://api.example.com/anthropic/v1",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got := normalizeBaseURL(tt.apiBase)
+ if got != tt.expected {
+ t.Errorf("normalizeBaseURL(%q) = %q, want %q", tt.apiBase, got, tt.expected)
+ }
+ })
+ }
+}
+
+func TestNewProvider(t *testing.T) {
+ provider := NewProvider("test-key", "https://api.example.com")
+ if provider == nil {
+ t.Fatal("NewProvider() returned nil")
+ }
+ if provider.apiKey != "test-key" {
+ t.Errorf("provider.apiKey = %q, want %q", provider.apiKey, "test-key")
+ }
+ if provider.apiBase != "https://api.example.com/v1" {
+ t.Errorf("provider.apiBase = %q, want %q", provider.apiBase, "https://api.example.com/v1")
+ }
+}
+
+func TestGetDefaultModel(t *testing.T) {
+ provider := NewProvider("test-key", "")
+ got := provider.GetDefaultModel()
+ expected := "claude-sonnet-4.6"
+ if got != expected {
+ t.Errorf("GetDefaultModel() = %q, want %q", got, expected)
+ }
+}
+
+// TestBuildRequestBodyEdgeCases tests edge cases for buildRequestBody.
+func TestBuildRequestBodyEdgeCases(t *testing.T) {
+ tests := []struct {
+ name string
+ messages []Message
+ tools []ToolDefinition
+ model string
+ options map[string]any
+ wantErr bool
+ }{
+ {
+ name: "empty message list",
+ messages: []Message{},
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 8192,
+ },
+ wantErr: false,
+ },
+ {
+ name: "very long system message",
+ messages: []Message{
+ {Role: "system", Content: strings.Repeat("This is a very long system prompt. ", 1000)},
+ {Role: "user", Content: "Hello"},
+ },
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 8192,
+ },
+ wantErr: false,
+ },
+ {
+ name: "multiple consecutive system messages",
+ messages: []Message{
+ {Role: "system", Content: "First system message"},
+ {Role: "system", Content: "Second system message"},
+ {Role: "system", Content: "Third system message"},
+ {Role: "user", Content: "Hello"},
+ },
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 8192,
+ },
+ wantErr: false,
+ },
+ {
+ name: "tool result without tool call",
+ messages: []Message{
+ {Role: "user", Content: "Use a tool"},
+ {Role: "assistant", Content: "", ToolCalls: []ToolCall{
+ {ID: "tool-1", Name: "test_tool", Arguments: map[string]any{"arg": "value"}},
+ }},
+ {Role: "user", ToolCallID: "tool-1", Content: "Tool result"},
+ },
+ model: "test-model",
+ options: map[string]any{
+ "max_tokens": 8192,
+ },
+ wantErr: false,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got, err := buildRequestBody(tt.messages, tt.tools, tt.model, tt.options)
+ if (err != nil) != tt.wantErr {
+ t.Errorf("buildRequestBody() error = %v, wantErr %v", err, tt.wantErr)
+ return
+ }
+ if err != nil {
+ return
+ }
+
+ // Verify basic structure
+ if got == nil {
+ t.Error("buildRequestBody() returned nil")
+ return
+ }
+ if got["model"] != tt.model {
+ t.Errorf("model = %v, want %v", got["model"], tt.model)
+ }
+ })
+ }
+}
+
+// TestParseResponseBodyEdgeCases tests edge cases for parseResponseBody.
+func TestParseResponseBodyEdgeCases(t *testing.T) {
+ tests := []struct {
+ name string
+ body []byte
+ wantErr bool
+ check func(*testing.T, *LLMResponse)
+ }{
+ {
+ name: "empty content blocks",
+ body: []byte(`{
+ "id": "msg-empty",
+ "type": "message",
+ "role": "assistant",
+ "content": [],
+ "stop_reason": "end_turn",
+ "model": "test-model",
+ "usage": {"input_tokens": 5, "output_tokens": 0}
+ }`),
+ wantErr: false,
+ check: func(t *testing.T, resp *LLMResponse) {
+ if resp.Content != "" {
+ t.Errorf("Content = %q, want empty string", resp.Content)
+ }
+ if len(resp.ToolCalls) != 0 {
+ t.Errorf("ToolCalls length = %d, want 0", len(resp.ToolCalls))
+ }
+ },
+ },
+ {
+ name: "multiple tool use blocks",
+ body: []byte(`{
+ "id": "msg-multi",
+ "type": "message",
+ "role": "assistant",
+ "content": [
+ {"type": "tool_use", "id": "tool-1", "name": "func1", "input": {"arg": "val1"}},
+ {"type": "tool_use", "id": "tool-2", "name": "func2", "input": {"arg": "val2"}}
+ ],
+ "stop_reason": "tool_use",
+ "model": "test-model",
+ "usage": {"input_tokens": 10, "output_tokens": 20}
+ }`),
+ wantErr: false,
+ check: func(t *testing.T, resp *LLMResponse) {
+ if len(resp.ToolCalls) != 2 {
+ t.Errorf("ToolCalls length = %d, want 2", len(resp.ToolCalls))
+ }
+ },
+ },
+ {
+ name: "malformed JSON response",
+ body: []byte(`{invalid json`),
+ wantErr: true,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got, err := parseResponseBody(tt.body)
+ if (err != nil) != tt.wantErr {
+ t.Errorf("parseResponseBody() error = %v, wantErr %v", err, tt.wantErr)
+ return
+ }
+ if tt.check != nil && err == nil {
+ tt.check(t, got)
+ }
+ })
+ }
+}
+
+// TestProviderChatErrors tests error handling in Chat.
+// Note: apiBase check removed as it's dead code - normalizeBaseURL() always provides a default.
+func TestProviderChatErrors(t *testing.T) {
+ tests := []struct {
+ name string
+ apiKey string
+ messages []Message
+ wantErrMsg string
+ }{
+ {
+ name: "missing API key",
+ apiKey: "",
+ messages: []Message{{Role: "user", Content: "Test"}},
+ wantErrMsg: "API key not configured",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ // Create provider using constructor to ensure proper initialization
+ provider := NewProvider(tt.apiKey, "https://api.example.com")
+
+ _, err := provider.Chat(context.Background(), tt.messages, nil, "test-model", nil)
+ if err == nil {
+ t.Fatal("Chat() expected error, got nil")
+ }
+ if err.Error() != tt.wantErrMsg {
+ t.Errorf("Chat() error = %q, want %q", err.Error(), tt.wantErrMsg)
+ }
+ })
+ }
+}
diff --git a/pkg/providers/azure/provider.go b/pkg/providers/azure/provider.go
new file mode 100644
index 000000000..e0ddbbde4
--- /dev/null
+++ b/pkg/providers/azure/provider.go
@@ -0,0 +1,150 @@
+package azure
+
+import (
+ "bytes"
+ "context"
+ "encoding/json"
+ "fmt"
+ "net/http"
+ "net/url"
+ "strings"
+ "time"
+
+ "github.com/sipeed/picoclaw/pkg/providers/common"
+ "github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
+)
+
+type (
+ LLMResponse = protocoltypes.LLMResponse
+ Message = protocoltypes.Message
+ ToolDefinition = protocoltypes.ToolDefinition
+)
+
+const (
+ // azureAPIVersion is the Azure OpenAI API version used for all requests.
+ azureAPIVersion = "2024-10-21"
+ defaultRequestTimeout = common.DefaultRequestTimeout
+)
+
+// Provider implements the LLM provider interface for Azure OpenAI endpoints.
+// It handles Azure-specific authentication (api-key header), URL construction
+// (deployment-based), and request body formatting (max_completion_tokens, no model field).
+type Provider struct {
+ apiKey string
+ apiBase string
+ httpClient *http.Client
+}
+
+// Option configures the Azure Provider.
+type Option func(*Provider)
+
+// WithRequestTimeout sets the HTTP request timeout.
+func WithRequestTimeout(timeout time.Duration) Option {
+ return func(p *Provider) {
+ if timeout > 0 {
+ p.httpClient.Timeout = timeout
+ }
+ }
+}
+
+// NewProvider creates a new Azure OpenAI provider.
+func NewProvider(apiKey, apiBase, proxy string, opts ...Option) *Provider {
+ p := &Provider{
+ apiKey: apiKey,
+ apiBase: strings.TrimRight(apiBase, "/"),
+ httpClient: common.NewHTTPClient(proxy),
+ }
+
+ for _, opt := range opts {
+ if opt != nil {
+ opt(p)
+ }
+ }
+
+ return p
+}
+
+// NewProviderWithTimeout creates a new Azure OpenAI provider with a custom request timeout in seconds.
+func NewProviderWithTimeout(apiKey, apiBase, proxy string, requestTimeoutSeconds int) *Provider {
+ return NewProvider(
+ apiKey, apiBase, proxy,
+ WithRequestTimeout(time.Duration(requestTimeoutSeconds)*time.Second),
+ )
+}
+
+// Chat sends a chat completion request to the Azure OpenAI endpoint.
+// The model parameter is used as the Azure deployment name in the URL.
+func (p *Provider) Chat(
+ ctx context.Context,
+ messages []Message,
+ tools []ToolDefinition,
+ model string,
+ options map[string]any,
+) (*LLMResponse, error) {
+ if p.apiBase == "" {
+ return nil, fmt.Errorf("Azure API base not configured")
+ }
+
+ // model is the deployment name for Azure OpenAI
+ deployment := model
+
+ // Build Azure-specific URL safely using url.JoinPath and query encoding
+ // to prevent path traversal or query injection via deployment names.
+ base, err := url.JoinPath(p.apiBase, "openai/deployments", deployment, "chat/completions")
+ if err != nil {
+ return nil, fmt.Errorf("failed to build Azure request URL: %w", err)
+ }
+ requestURL := base + "?api-version=" + azureAPIVersion
+
+ // Build request body — no "model" field (Azure infers from deployment URL)
+ requestBody := map[string]any{
+ "messages": common.SerializeMessages(messages),
+ }
+
+ if len(tools) > 0 {
+ requestBody["tools"] = tools
+ requestBody["tool_choice"] = "auto"
+ }
+
+ // Azure OpenAI always uses max_completion_tokens
+ if maxTokens, ok := common.AsInt(options["max_tokens"]); ok {
+ requestBody["max_completion_tokens"] = maxTokens
+ }
+
+ if temperature, ok := common.AsFloat(options["temperature"]); ok {
+ requestBody["temperature"] = temperature
+ }
+
+ jsonData, err := json.Marshal(requestBody)
+ if err != nil {
+ return nil, fmt.Errorf("failed to marshal request: %w", err)
+ }
+
+ req, err := http.NewRequestWithContext(ctx, "POST", requestURL, bytes.NewReader(jsonData))
+ if err != nil {
+ return nil, fmt.Errorf("failed to create request: %w", err)
+ }
+
+ // Azure uses api-key header instead of Authorization: Bearer
+ req.Header.Set("Content-Type", "application/json")
+ if p.apiKey != "" {
+ req.Header.Set("Api-Key", p.apiKey)
+ }
+
+ resp, err := p.httpClient.Do(req)
+ if err != nil {
+ return nil, fmt.Errorf("failed to send request: %w", err)
+ }
+ defer resp.Body.Close()
+
+ if resp.StatusCode != http.StatusOK {
+ return nil, common.HandleErrorResponse(resp, p.apiBase)
+ }
+
+ return common.ReadAndParseResponse(resp, p.apiBase)
+}
+
+// GetDefaultModel returns an empty string as Azure deployments are user-configured.
+func (p *Provider) GetDefaultModel() string {
+ return ""
+}
diff --git a/pkg/providers/azure/provider_test.go b/pkg/providers/azure/provider_test.go
new file mode 100644
index 000000000..531b81296
--- /dev/null
+++ b/pkg/providers/azure/provider_test.go
@@ -0,0 +1,232 @@
+package azure
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+)
+
+// writeValidResponse writes a minimal valid Azure OpenAI chat completion response.
+func writeValidResponse(w http.ResponseWriter) {
+ 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)
+}
+
+func TestProviderChat_AzureURLConstruction(t *testing.T) {
+ var capturedPath string
+ var capturedAPIVersion string
+
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ capturedPath = r.URL.Path
+ capturedAPIVersion = r.URL.Query().Get("api-version")
+ writeValidResponse(w)
+ }))
+ defer server.Close()
+
+ p := NewProvider("test-key", server.URL, "")
+ _, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "my-gpt5-deployment", nil)
+ if err != nil {
+ t.Fatalf("Chat() error = %v", err)
+ }
+
+ wantPath := "/openai/deployments/my-gpt5-deployment/chat/completions"
+ if capturedPath != wantPath {
+ t.Errorf("URL path = %q, want %q", capturedPath, wantPath)
+ }
+ if capturedAPIVersion != azureAPIVersion {
+ t.Errorf("api-version = %q, want %q", capturedAPIVersion, azureAPIVersion)
+ }
+}
+
+func TestProviderChat_AzureAuthHeader(t *testing.T) {
+ var capturedAPIKey string
+ var capturedAuth string
+
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ capturedAPIKey = r.Header.Get("Api-Key")
+ capturedAuth = r.Header.Get("Authorization")
+ writeValidResponse(w)
+ }))
+ defer server.Close()
+
+ p := NewProvider("test-azure-key", server.URL, "")
+ _, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "deployment", nil)
+ if err != nil {
+ t.Fatalf("Chat() error = %v", err)
+ }
+
+ if capturedAPIKey != "test-azure-key" {
+ t.Errorf("api-key header = %q, want %q", capturedAPIKey, "test-azure-key")
+ }
+ if capturedAuth != "" {
+ t.Errorf("Authorization header should be empty, got %q", capturedAuth)
+ }
+}
+
+func TestProviderChat_AzureOmitsModelFromBody(t *testing.T) {
+ var requestBody map[string]any
+
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ json.NewDecoder(r.Body).Decode(&requestBody)
+ writeValidResponse(w)
+ }))
+ defer server.Close()
+
+ p := NewProvider("test-key", server.URL, "")
+ _, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "deployment", nil)
+ if err != nil {
+ t.Fatalf("Chat() error = %v", err)
+ }
+
+ if _, exists := requestBody["model"]; exists {
+ t.Error("request body should not contain 'model' field for Azure OpenAI")
+ }
+}
+
+func TestProviderChat_AzureUsesMaxCompletionTokens(t *testing.T) {
+ var requestBody map[string]any
+
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ json.NewDecoder(r.Body).Decode(&requestBody)
+ writeValidResponse(w)
+ }))
+ defer server.Close()
+
+ p := NewProvider("test-key", server.URL, "")
+ _, err := p.Chat(
+ t.Context(),
+ []Message{{Role: "user", Content: "hi"}},
+ nil,
+ "deployment",
+ map[string]any{"max_tokens": 2048},
+ )
+ if err != nil {
+ t.Fatalf("Chat() error = %v", err)
+ }
+
+ if _, exists := requestBody["max_completion_tokens"]; !exists {
+ t.Error("request body should contain 'max_completion_tokens'")
+ }
+ if _, exists := requestBody["max_tokens"]; exists {
+ t.Error("request body should not contain 'max_tokens'")
+ }
+}
+
+func TestProviderChat_AzureHTTPError(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized)
+ }))
+ defer server.Close()
+
+ p := NewProvider("bad-key", server.URL, "")
+ _, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "deployment", nil)
+ if err == nil {
+ t.Fatal("expected error, got nil")
+ }
+}
+
+func TestProviderChat_AzureParseToolCalls(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_calls": []map[string]any{
+ {
+ "id": "call_1",
+ "type": "function",
+ "function": map[string]any{
+ "name": "get_weather",
+ "arguments": `{"city":"Seattle"}`,
+ },
+ },
+ },
+ },
+ "finish_reason": "tool_calls",
+ },
+ },
+ }
+ w.Header().Set("Content-Type", "application/json")
+ json.NewEncoder(w).Encode(resp)
+ }))
+ defer server.Close()
+
+ p := NewProvider("test-key", server.URL, "")
+ out, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "weather?"}}, nil, "deployment", 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))
+ }
+ if out.ToolCalls[0].Name != "get_weather" {
+ t.Errorf("ToolCalls[0].Name = %q, want %q", out.ToolCalls[0].Name, "get_weather")
+ }
+}
+
+func TestProvider_AzureEmptyAPIBase(t *testing.T) {
+ p := NewProvider("test-key", "", "")
+ _, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "deployment", nil)
+ if err == nil {
+ t.Fatal("expected error for empty API base")
+ }
+}
+
+func TestProvider_AzureRequestTimeoutDefault(t *testing.T) {
+ p := NewProvider("test-key", "https://example.com", "")
+ if p.httpClient.Timeout != defaultRequestTimeout {
+ t.Errorf("timeout = %v, want %v", p.httpClient.Timeout, defaultRequestTimeout)
+ }
+}
+
+func TestProvider_AzureRequestTimeoutOverride(t *testing.T) {
+ p := NewProvider("test-key", "https://example.com", "", WithRequestTimeout(300*time.Second))
+ if p.httpClient.Timeout != 300*time.Second {
+ t.Errorf("timeout = %v, want %v", p.httpClient.Timeout, 300*time.Second)
+ }
+}
+
+func TestProvider_AzureNewProviderWithTimeout(t *testing.T) {
+ p := NewProviderWithTimeout("test-key", "https://example.com", "", 180)
+ if p.httpClient.Timeout != 180*time.Second {
+ t.Errorf("timeout = %v, want %v", p.httpClient.Timeout, 180*time.Second)
+ }
+}
+
+func TestProviderChat_AzureDeploymentNameEscaped(t *testing.T) {
+ var capturedPath string
+
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ capturedPath = r.URL.RawPath // use RawPath to see percent-encoding
+ if capturedPath == "" {
+ capturedPath = r.URL.Path
+ }
+ writeValidResponse(w)
+ }))
+ defer server.Close()
+
+ p := NewProvider("test-key", server.URL, "")
+
+ // Deployment name with characters that could cause path injection
+ _, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "my deploy/../../admin", nil)
+ if err != nil {
+ t.Fatalf("Chat() error = %v", err)
+ }
+
+ // The slash and special chars in the deployment name must be escaped, not treated as path separators
+ if capturedPath == "/openai/deployments/my deploy/../../admin/chat/completions" {
+ t.Fatal("deployment name was interpolated without escaping — path injection possible")
+ }
+}
diff --git a/pkg/providers/common/common.go b/pkg/providers/common/common.go
new file mode 100644
index 000000000..23680a1bf
--- /dev/null
+++ b/pkg/providers/common/common.go
@@ -0,0 +1,380 @@
+// PicoClaw - Ultra-lightweight personal AI agent
+// License: MIT
+//
+// Copyright (c) 2026 PicoClaw contributors
+
+// Package common provides shared utilities used by multiple LLM provider
+// implementations (openai_compat, azure, etc.).
+package common
+
+import (
+ "bufio"
+ "bytes"
+ "encoding/json"
+ "fmt"
+ "io"
+ "log"
+ "net/http"
+ "net/url"
+ "strings"
+ "time"
+
+ "github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
+)
+
+// Re-export protocol types used across providers.
+type (
+ ToolCall = protocoltypes.ToolCall
+ FunctionCall = protocoltypes.FunctionCall
+ LLMResponse = protocoltypes.LLMResponse
+ UsageInfo = protocoltypes.UsageInfo
+ Message = protocoltypes.Message
+ ToolDefinition = protocoltypes.ToolDefinition
+ ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
+ ExtraContent = protocoltypes.ExtraContent
+ GoogleExtra = protocoltypes.GoogleExtra
+ ReasoningDetail = protocoltypes.ReasoningDetail
+)
+
+const DefaultRequestTimeout = 120 * time.Second
+
+// NewHTTPClient creates an *http.Client with an optional proxy and the default timeout.
+func NewHTTPClient(proxy string) *http.Client {
+ client := &http.Client{
+ Timeout: DefaultRequestTimeout,
+ }
+ if proxy != "" {
+ parsed, err := url.Parse(proxy)
+ if err == nil {
+ // Preserve http.DefaultTransport settings (TLS, HTTP/2, timeouts, etc.)
+ if base, ok := http.DefaultTransport.(*http.Transport); ok {
+ tr := base.Clone()
+ tr.Proxy = http.ProxyURL(parsed)
+ client.Transport = tr
+ } else {
+ // Fallback: minimal transport if DefaultTransport is not *http.Transport.
+ client.Transport = &http.Transport{
+ Proxy: http.ProxyURL(parsed),
+ }
+ }
+ } else {
+ log.Printf("common: invalid proxy URL %q: %v", proxy, err)
+ }
+ }
+ return client
+}
+
+// --- Message serialization ---
+
+// openaiMessage is the wire-format message for OpenAI-compatible APIs.
+// It mirrors protocoltypes.Message but omits SystemParts, which is an
+// internal field that would be unknown to third-party endpoints.
+type openaiMessage struct {
+ Role string `json:"role"`
+ Content string `json:"content"`
+ ReasoningContent string `json:"reasoning_content,omitempty"`
+ ToolCalls []ToolCall `json:"tool_calls,omitempty"`
+ ToolCallID string `json:"tool_call_id,omitempty"`
+}
+
+// SerializeMessages converts internal Message structs to the OpenAI wire format.
+// - Strips SystemParts (unknown to third-party endpoints)
+// - Converts messages with Media to multipart content format (text + image_url parts)
+// - Preserves ToolCallID, ToolCalls, and ReasoningContent for all messages
+func SerializeMessages(messages []Message) []any {
+ out := make([]any, 0, len(messages))
+ for _, m := range messages {
+ if len(m.Media) == 0 {
+ out = append(out, openaiMessage{
+ Role: m.Role,
+ Content: m.Content,
+ ReasoningContent: m.ReasoningContent,
+ ToolCalls: m.ToolCalls,
+ ToolCallID: m.ToolCallID,
+ })
+ continue
+ }
+
+ // Multipart content format for messages with media
+ parts := make([]map[string]any, 0, 1+len(m.Media))
+ if m.Content != "" {
+ parts = append(parts, map[string]any{
+ "type": "text",
+ "text": m.Content,
+ })
+ }
+ for _, mediaURL := range m.Media {
+ if strings.HasPrefix(mediaURL, "data:image/") {
+ parts = append(parts, map[string]any{
+ "type": "image_url",
+ "image_url": map[string]any{
+ "url": mediaURL,
+ },
+ })
+ }
+ }
+
+ msg := map[string]any{
+ "role": m.Role,
+ "content": parts,
+ }
+ if m.ToolCallID != "" {
+ msg["tool_call_id"] = m.ToolCallID
+ }
+ if len(m.ToolCalls) > 0 {
+ msg["tool_calls"] = m.ToolCalls
+ }
+ if m.ReasoningContent != "" {
+ msg["reasoning_content"] = m.ReasoningContent
+ }
+ out = append(out, msg)
+ }
+ return out
+}
+
+// --- Response parsing ---
+
+// ParseResponse parses a JSON chat completion response body into an LLMResponse.
+func ParseResponse(body io.Reader) (*LLMResponse, error) {
+ var apiResponse struct {
+ Choices []struct {
+ Message struct {
+ Content string `json:"content"`
+ ReasoningContent string `json:"reasoning_content"`
+ Reasoning string `json:"reasoning"`
+ ReasoningDetails []ReasoningDetail `json:"reasoning_details"`
+ ToolCalls []struct {
+ ID string `json:"id"`
+ Type string `json:"type"`
+ Function *struct {
+ Name string `json:"name"`
+ Arguments json.RawMessage `json:"arguments"`
+ } `json:"function"`
+ ExtraContent *struct {
+ Google *struct {
+ ThoughtSignature string `json:"thought_signature"`
+ } `json:"google"`
+ } `json:"extra_content"`
+ } `json:"tool_calls"`
+ } `json:"message"`
+ FinishReason string `json:"finish_reason"`
+ } `json:"choices"`
+ Usage *UsageInfo `json:"usage"`
+ }
+
+ if err := json.NewDecoder(body).Decode(&apiResponse); err != nil {
+ return nil, fmt.Errorf("failed to decode response: %w", err)
+ }
+
+ if len(apiResponse.Choices) == 0 {
+ return &LLMResponse{
+ Content: "",
+ FinishReason: "stop",
+ }, nil
+ }
+
+ choice := apiResponse.Choices[0]
+ toolCalls := make([]ToolCall, 0, len(choice.Message.ToolCalls))
+ for _, tc := range choice.Message.ToolCalls {
+ arguments := make(map[string]any)
+ name := ""
+
+ // Extract thought_signature from Gemini/Google-specific extra content
+ thoughtSignature := ""
+ if tc.ExtraContent != nil && tc.ExtraContent.Google != nil {
+ thoughtSignature = tc.ExtraContent.Google.ThoughtSignature
+ }
+
+ if tc.Function != nil {
+ name = tc.Function.Name
+ arguments = DecodeToolCallArguments(tc.Function.Arguments, name)
+ }
+
+ toolCall := ToolCall{
+ ID: tc.ID,
+ Name: name,
+ Arguments: arguments,
+ ThoughtSignature: thoughtSignature,
+ }
+
+ if thoughtSignature != "" {
+ toolCall.ExtraContent = &ExtraContent{
+ Google: &GoogleExtra{
+ ThoughtSignature: thoughtSignature,
+ },
+ }
+ }
+
+ toolCalls = append(toolCalls, toolCall)
+ }
+
+ return &LLMResponse{
+ Content: choice.Message.Content,
+ ReasoningContent: choice.Message.ReasoningContent,
+ Reasoning: choice.Message.Reasoning,
+ ReasoningDetails: choice.Message.ReasoningDetails,
+ ToolCalls: toolCalls,
+ FinishReason: choice.FinishReason,
+ Usage: apiResponse.Usage,
+ }, nil
+}
+
+// DecodeToolCallArguments decodes a tool call's arguments from raw JSON.
+func DecodeToolCallArguments(raw json.RawMessage, name string) map[string]any {
+ arguments := make(map[string]any)
+ raw = bytes.TrimSpace(raw)
+ if len(raw) == 0 || bytes.Equal(raw, []byte("null")) {
+ return arguments
+ }
+
+ var decoded any
+ if err := json.Unmarshal(raw, &decoded); err != nil {
+ log.Printf("common: failed to decode tool call arguments payload for %q: %v", name, err)
+ arguments["raw"] = string(raw)
+ return arguments
+ }
+
+ switch v := decoded.(type) {
+ case string:
+ if strings.TrimSpace(v) == "" {
+ return arguments
+ }
+ if err := json.Unmarshal([]byte(v), &arguments); err != nil {
+ log.Printf("common: failed to decode tool call arguments for %q: %v", name, err)
+ arguments["raw"] = v
+ }
+ return arguments
+ case map[string]any:
+ return v
+ default:
+ log.Printf("common: unsupported tool call arguments type for %q: %T", name, decoded)
+ arguments["raw"] = string(raw)
+ return arguments
+ }
+}
+
+// --- HTTP response helpers ---
+
+// HandleErrorResponse reads a non-200 response body and returns an appropriate error.
+func HandleErrorResponse(resp *http.Response, apiBase string) error {
+ contentType := resp.Header.Get("Content-Type")
+ body, readErr := io.ReadAll(io.LimitReader(resp.Body, 256))
+ if readErr != nil {
+ return fmt.Errorf("failed to read response: %w", readErr)
+ }
+ if LooksLikeHTML(body, contentType) {
+ return WrapHTMLResponseError(resp.StatusCode, body, contentType, apiBase)
+ }
+ return fmt.Errorf(
+ "API request failed:\n Status: %d\n Body: %s",
+ resp.StatusCode,
+ ResponsePreview(body, 128),
+ )
+}
+
+// ReadAndParseResponse peeks at the response body to detect HTML errors,
+// then parses the JSON response into an LLMResponse.
+func ReadAndParseResponse(resp *http.Response, apiBase string) (*LLMResponse, error) {
+ contentType := resp.Header.Get("Content-Type")
+ reader := bufio.NewReader(resp.Body)
+ prefix, err := reader.Peek(256)
+ if err != nil && err != io.EOF && err != bufio.ErrBufferFull {
+ return nil, fmt.Errorf("failed to inspect response: %w", err)
+ }
+ if LooksLikeHTML(prefix, contentType) {
+ return nil, WrapHTMLResponseError(resp.StatusCode, prefix, contentType, apiBase)
+ }
+ out, err := ParseResponse(reader)
+ if err != nil {
+ return nil, fmt.Errorf("failed to parse JSON response: %w", err)
+ }
+ return out, nil
+}
+
+// LooksLikeHTML checks if the response body appears to be HTML.
+func LooksLikeHTML(body []byte, contentType string) bool {
+ contentType = strings.ToLower(strings.TrimSpace(contentType))
+ if strings.Contains(contentType, "text/html") || strings.Contains(contentType, "application/xhtml+xml") {
+ return true
+ }
+ prefix := bytes.ToLower(leadingTrimmedPrefix(body, 128))
+ return bytes.HasPrefix(prefix, []byte(""
+ }
+ if len(trimmed) <= maxLen {
+ return string(trimmed)
+ }
+ return string(trimmed[:maxLen]) + "..."
+}
+
+func leadingTrimmedPrefix(body []byte, maxLen int) []byte {
+ i := 0
+ for i < len(body) {
+ switch body[i] {
+ case ' ', '\t', '\n', '\r', '\f', '\v':
+ i++
+ default:
+ end := i + maxLen
+ if end > len(body) {
+ end = len(body)
+ }
+ return body[i:end]
+ }
+ }
+ return nil
+}
+
+// --- Numeric helpers ---
+
+// AsInt converts various numeric types to int.
+func AsInt(v any) (int, bool) {
+ switch val := v.(type) {
+ case int:
+ return val, true
+ case int64:
+ return int(val), true
+ case float64:
+ return int(val), true
+ case float32:
+ return int(val), true
+ default:
+ return 0, false
+ }
+}
+
+// AsFloat converts various numeric types to float64.
+func AsFloat(v any) (float64, bool) {
+ switch val := v.(type) {
+ case float64:
+ return val, true
+ case float32:
+ return float64(val), true
+ case int:
+ return float64(val), true
+ case int64:
+ return float64(val), true
+ default:
+ return 0, false
+ }
+}
diff --git a/pkg/providers/common/common_test.go b/pkg/providers/common/common_test.go
new file mode 100644
index 000000000..bb7e7434d
--- /dev/null
+++ b/pkg/providers/common/common_test.go
@@ -0,0 +1,558 @@
+package common
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "net/url"
+ "strings"
+ "testing"
+
+ "github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
+)
+
+// --- NewHTTPClient tests ---
+
+func TestNewHTTPClient_DefaultTimeout(t *testing.T) {
+ client := NewHTTPClient("")
+ if client.Timeout != DefaultRequestTimeout {
+ t.Errorf("timeout = %v, want %v", client.Timeout, DefaultRequestTimeout)
+ }
+}
+
+func TestNewHTTPClient_WithProxy(t *testing.T) {
+ client := NewHTTPClient("http://127.0.0.1:8080")
+ transport, ok := client.Transport.(*http.Transport)
+ if !ok || transport == nil {
+ t.Fatalf("expected http.Transport with proxy, got %T", client.Transport)
+ }
+ req := &http.Request{URL: &url.URL{Scheme: "https", Host: "api.example.com"}}
+ gotProxy, err := transport.Proxy(req)
+ if err != nil {
+ t.Fatalf("proxy function error: %v", err)
+ }
+ if gotProxy == nil || gotProxy.String() != "http://127.0.0.1:8080" {
+ t.Errorf("proxy = %v, want http://127.0.0.1:8080", gotProxy)
+ }
+}
+
+func TestNewHTTPClient_NoProxy(t *testing.T) {
+ client := NewHTTPClient("")
+ if client.Transport != nil {
+ t.Errorf("expected nil transport without proxy, got %T", client.Transport)
+ }
+}
+
+func TestNewHTTPClient_InvalidProxy(t *testing.T) {
+ // Should not panic, just log and return client without proxy
+ client := NewHTTPClient("://bad-url")
+ if client == nil {
+ t.Fatal("expected non-nil client even with invalid proxy")
+ }
+}
+
+// --- SerializeMessages tests ---
+
+func TestSerializeMessages_PlainText(t *testing.T) {
+ messages := []Message{
+ {Role: "user", Content: "hello"},
+ {Role: "assistant", Content: "hi", ReasoningContent: "thinking..."},
+ }
+ result := SerializeMessages(messages)
+
+ data, _ := json.Marshal(result)
+ var msgs []map[string]any
+ json.Unmarshal(data, &msgs)
+
+ if msgs[0]["content"] != "hello" {
+ t.Errorf("expected plain string content, got %v", msgs[0]["content"])
+ }
+ if msgs[1]["reasoning_content"] != "thinking..." {
+ t.Errorf("reasoning_content not preserved, got %v", msgs[1]["reasoning_content"])
+ }
+}
+
+func TestSerializeMessages_WithMedia(t *testing.T) {
+ messages := []Message{
+ {Role: "user", Content: "describe this", Media: []string{"data:image/png;base64,abc123"}},
+ }
+ result := SerializeMessages(messages)
+
+ data, _ := json.Marshal(result)
+ var msgs []map[string]any
+ json.Unmarshal(data, &msgs)
+
+ content, ok := msgs[0]["content"].([]any)
+ if !ok {
+ t.Fatalf("expected array content for media message, got %T", msgs[0]["content"])
+ }
+ if len(content) != 2 {
+ t.Fatalf("expected 2 content parts, got %d", len(content))
+ }
+}
+
+func TestSerializeMessages_MediaWithToolCallID(t *testing.T) {
+ messages := []Message{
+ {Role: "tool", Content: "result", Media: []string{"data:image/png;base64,xyz"}, ToolCallID: "call_1"},
+ }
+ result := SerializeMessages(messages)
+
+ data, _ := json.Marshal(result)
+ var msgs []map[string]any
+ json.Unmarshal(data, &msgs)
+
+ if msgs[0]["tool_call_id"] != "call_1" {
+ t.Errorf("tool_call_id not preserved, got %v", msgs[0]["tool_call_id"])
+ }
+}
+
+func TestSerializeMessages_StripsSystemParts(t *testing.T) {
+ messages := []Message{
+ {
+ Role: "system",
+ Content: "you are helpful",
+ SystemParts: []protocoltypes.ContentBlock{
+ {Type: "text", Text: "you are helpful"},
+ },
+ },
+ }
+ result := SerializeMessages(messages)
+
+ data, _ := json.Marshal(result)
+ if strings.Contains(string(data), "system_parts") {
+ t.Error("system_parts should not appear in serialized output")
+ }
+}
+
+// --- ParseResponse tests ---
+
+func TestParseResponse_BasicContent(t *testing.T) {
+ body := `{"choices":[{"message":{"content":"hello world"},"finish_reason":"stop"}]}`
+ out, err := ParseResponse(strings.NewReader(body))
+ if err != nil {
+ t.Fatalf("ParseResponse() error = %v", err)
+ }
+ if out.Content != "hello world" {
+ t.Errorf("Content = %q, want %q", out.Content, "hello world")
+ }
+ if out.FinishReason != "stop" {
+ t.Errorf("FinishReason = %q, want %q", out.FinishReason, "stop")
+ }
+}
+
+func TestParseResponse_EmptyChoices(t *testing.T) {
+ body := `{"choices":[]}`
+ out, err := ParseResponse(strings.NewReader(body))
+ if err != nil {
+ t.Fatalf("ParseResponse() error = %v", err)
+ }
+ if out.Content != "" {
+ t.Errorf("Content = %q, want empty", out.Content)
+ }
+ if out.FinishReason != "stop" {
+ t.Errorf("FinishReason = %q, want %q", out.FinishReason, "stop")
+ }
+}
+
+func TestParseResponse_WithToolCalls(t *testing.T) {
+ body := `{"choices":[{"message":{"content":"","tool_calls":[{"id":"call_1","type":"function","function":{"name":"get_weather","arguments":"{\"city\":\"SF\"}"}}]},"finish_reason":"tool_calls"}]}`
+ out, err := ParseResponse(strings.NewReader(body))
+ if err != nil {
+ t.Fatalf("ParseResponse() error = %v", err)
+ }
+ if len(out.ToolCalls) != 1 {
+ t.Fatalf("len(ToolCalls) = %d, want 1", len(out.ToolCalls))
+ }
+ if out.ToolCalls[0].Name != "get_weather" {
+ t.Errorf("ToolCalls[0].Name = %q, want %q", out.ToolCalls[0].Name, "get_weather")
+ }
+ if out.ToolCalls[0].Arguments["city"] != "SF" {
+ t.Errorf("ToolCalls[0].Arguments[city] = %v, want SF", out.ToolCalls[0].Arguments["city"])
+ }
+}
+
+func TestParseResponse_WithUsage(t *testing.T) {
+ body := `{"choices":[{"message":{"content":"ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15}}`
+ out, err := ParseResponse(strings.NewReader(body))
+ if err != nil {
+ t.Fatalf("ParseResponse() error = %v", err)
+ }
+ if out.Usage == nil {
+ t.Fatal("Usage is nil")
+ }
+ if out.Usage.PromptTokens != 10 {
+ t.Errorf("PromptTokens = %d, want 10", out.Usage.PromptTokens)
+ }
+}
+
+func TestParseResponse_WithReasoningContent(t *testing.T) {
+ body := `{"choices":[{"message":{"content":"2","reasoning_content":"Let me think... 1+1=2"},"finish_reason":"stop"}]}`
+ out, err := ParseResponse(strings.NewReader(body))
+ if err != nil {
+ t.Fatalf("ParseResponse() error = %v", err)
+ }
+ if out.ReasoningContent != "Let me think... 1+1=2" {
+ t.Errorf("ReasoningContent = %q, want %q", out.ReasoningContent, "Let me think... 1+1=2")
+ }
+}
+
+func TestParseResponse_InvalidJSON(t *testing.T) {
+ _, err := ParseResponse(strings.NewReader("not json"))
+ if err == nil {
+ t.Fatal("expected error for invalid JSON")
+ }
+}
+
+// --- DecodeToolCallArguments tests ---
+
+func TestDecodeToolCallArguments_ObjectJSON(t *testing.T) {
+ raw := json.RawMessage(`{"city":"Seattle","units":"metric"}`)
+ args := DecodeToolCallArguments(raw, "test")
+ if args["city"] != "Seattle" {
+ t.Errorf("city = %v, want Seattle", args["city"])
+ }
+ if args["units"] != "metric" {
+ t.Errorf("units = %v, want metric", args["units"])
+ }
+}
+
+func TestDecodeToolCallArguments_StringJSON(t *testing.T) {
+ raw := json.RawMessage(`"{\"city\":\"SF\"}"`)
+ args := DecodeToolCallArguments(raw, "test")
+ if args["city"] != "SF" {
+ t.Errorf("city = %v, want SF", args["city"])
+ }
+}
+
+func TestDecodeToolCallArguments_EmptyInput(t *testing.T) {
+ args := DecodeToolCallArguments(nil, "test")
+ if len(args) != 0 {
+ t.Errorf("expected empty map, got %v", args)
+ }
+}
+
+func TestDecodeToolCallArguments_NullInput(t *testing.T) {
+ args := DecodeToolCallArguments(json.RawMessage(`null`), "test")
+ if len(args) != 0 {
+ t.Errorf("expected empty map, got %v", args)
+ }
+}
+
+func TestDecodeToolCallArguments_InvalidJSON(t *testing.T) {
+ args := DecodeToolCallArguments(json.RawMessage(`not-json`), "test")
+ if _, ok := args["raw"]; !ok {
+ t.Error("expected 'raw' fallback key for invalid JSON")
+ }
+}
+
+func TestDecodeToolCallArguments_EmptyStringJSON(t *testing.T) {
+ args := DecodeToolCallArguments(json.RawMessage(`" "`), "test")
+ if len(args) != 0 {
+ t.Errorf("expected empty map for whitespace string, got %v", args)
+ }
+}
+
+// --- HandleErrorResponse tests ---
+
+func TestHandleErrorResponse_JSONError(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ w.WriteHeader(http.StatusBadRequest)
+ w.Write([]byte(`{"error":"bad request"}`))
+ }))
+ defer server.Close()
+
+ resp, err := http.Get(server.URL)
+ if err != nil {
+ t.Fatalf("http.Get() error = %v", err)
+ }
+ defer resp.Body.Close()
+ err = HandleErrorResponse(resp, server.URL)
+ if err == nil {
+ t.Fatal("expected error")
+ }
+ if !strings.Contains(err.Error(), "400") {
+ t.Errorf("error should contain status code, got %v", err)
+ }
+ if strings.Contains(err.Error(), "HTML") {
+ t.Errorf("should not mention HTML for JSON error, got %v", err)
+ }
+}
+
+func TestHandleErrorResponse_HTMLError(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "text/html")
+ w.WriteHeader(http.StatusBadGateway)
+ w.Write([]byte("bad gateway"))
+ }))
+ defer server.Close()
+
+ resp, err := http.Get(server.URL)
+ if err != nil {
+ t.Fatalf("http.Get() error = %v", err)
+ }
+ defer resp.Body.Close()
+ err = HandleErrorResponse(resp, server.URL)
+ if err == nil {
+ t.Fatal("expected error")
+ }
+ if !strings.Contains(err.Error(), "HTML instead of JSON") {
+ t.Errorf("expected HTML error message, got %v", err)
+ }
+}
+
+// --- ReadAndParseResponse tests ---
+
+func TestReadAndParseResponse_ValidJSON(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ w.Write([]byte(`{"choices":[{"message":{"content":"ok"},"finish_reason":"stop"}]}`))
+ }))
+ defer server.Close()
+
+ resp, err := http.Get(server.URL)
+ if err != nil {
+ t.Fatalf("http.Get() error = %v", err)
+ }
+ defer resp.Body.Close()
+ out, err := ReadAndParseResponse(resp, server.URL)
+ if err != nil {
+ t.Fatalf("ReadAndParseResponse() error = %v", err)
+ }
+ if out.Content != "ok" {
+ t.Errorf("Content = %q, want %q", out.Content, "ok")
+ }
+}
+
+func TestReadAndParseResponse_HTMLResponse(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "text/html")
+ w.Write([]byte("login page"))
+ }))
+ defer server.Close()
+
+ resp, err := http.Get(server.URL)
+ if err != nil {
+ t.Fatalf("http.Get() error = %v", err)
+ }
+ defer resp.Body.Close()
+ _, err = ReadAndParseResponse(resp, server.URL)
+ if err == nil {
+ t.Fatal("expected error for HTML response")
+ }
+ if !strings.Contains(err.Error(), "HTML instead of JSON") {
+ t.Errorf("expected HTML error, got %v", err)
+ }
+}
+
+// --- LooksLikeHTML tests ---
+
+func TestLooksLikeHTML_ContentTypeHTML(t *testing.T) {
+ if !LooksLikeHTML(nil, "text/html; charset=utf-8") {
+ t.Error("expected true for text/html content type")
+ }
+}
+
+func TestLooksLikeHTML_ContentTypeXHTML(t *testing.T) {
+ if !LooksLikeHTML(nil, "application/xhtml+xml") {
+ t.Error("expected true for xhtml content type")
+ }
+}
+
+func TestLooksLikeHTML_BodyPrefix(t *testing.T) {
+ tests := []struct {
+ name string
+ body string
+ }{
+ {"doctype", ""},
+ {"html tag", ""},
+ {"head tag", ""},
+ {"body tag", "content"},
+ {"whitespace before", " \n\t"},
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ if !LooksLikeHTML([]byte(tt.body), "application/json") {
+ t.Errorf("expected true for body %q", tt.body)
+ }
+ })
+ }
+}
+
+func TestLooksLikeHTML_NotHTML(t *testing.T) {
+ if LooksLikeHTML([]byte(`{"error":"bad"}`), "application/json") {
+ t.Error("expected false for JSON body")
+ }
+}
+
+// --- ResponsePreview tests ---
+
+func TestResponsePreview_Short(t *testing.T) {
+ got := ResponsePreview([]byte("hello"), 128)
+ if got != "hello" {
+ t.Errorf("got %q, want %q", got, "hello")
+ }
+}
+
+func TestResponsePreview_Truncated(t *testing.T) {
+ body := strings.Repeat("a", 200)
+ got := ResponsePreview([]byte(body), 128)
+ if len(got) != 131 { // 128 + "..."
+ t.Errorf("len = %d, want 131", len(got))
+ }
+ if !strings.HasSuffix(got, "...") {
+ t.Error("expected ... suffix")
+ }
+}
+
+func TestResponsePreview_Empty(t *testing.T) {
+ got := ResponsePreview([]byte(""), 128)
+ if got != "" {
+ t.Errorf("got %q, want %q", got, "")
+ }
+}
+
+func TestResponsePreview_Whitespace(t *testing.T) {
+ got := ResponsePreview([]byte(" \n\t "), 128)
+ if got != "" {
+ t.Errorf("got %q, want %q for whitespace-only body", got, "")
+ }
+}
+
+// --- AsInt tests ---
+
+func TestAsInt(t *testing.T) {
+ tests := []struct {
+ name string
+ val any
+ want int
+ ok bool
+ }{
+ {"int", 42, 42, true},
+ {"int64", int64(99), 99, true},
+ {"float64", float64(512), 512, true},
+ {"float32", float32(256), 256, true},
+ {"string", "nope", 0, false},
+ {"nil", nil, 0, false},
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got, ok := AsInt(tt.val)
+ if ok != tt.ok || got != tt.want {
+ t.Errorf("AsInt(%v) = (%d, %v), want (%d, %v)", tt.val, got, ok, tt.want, tt.ok)
+ }
+ })
+ }
+}
+
+// --- AsFloat tests ---
+
+func TestAsFloat(t *testing.T) {
+ tests := []struct {
+ name string
+ val any
+ want float64
+ ok bool
+ }{
+ {"float64", float64(0.7), 0.7, true},
+ {"float32", float32(0.5), float64(float32(0.5)), true},
+ {"int", 1, 1.0, true},
+ {"int64", int64(100), 100.0, true},
+ {"string", "nope", 0, false},
+ {"nil", nil, 0, false},
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got, ok := AsFloat(tt.val)
+ if ok != tt.ok || got != tt.want {
+ t.Errorf("AsFloat(%v) = (%f, %v), want (%f, %v)", tt.val, got, ok, tt.want, tt.ok)
+ }
+ })
+ }
+}
+
+// --- WrapHTMLResponseError tests ---
+
+func TestWrapHTMLResponseError(t *testing.T) {
+ err := WrapHTMLResponseError(502, []byte("bad"), "text/html", "https://api.example.com")
+ if err == nil {
+ t.Fatal("expected error")
+ }
+ msg := err.Error()
+ if !strings.Contains(msg, "502") {
+ t.Errorf("expected status code in error, got %v", msg)
+ }
+ if !strings.Contains(msg, "https://api.example.com") {
+ t.Errorf("expected api base in error, got %v", msg)
+ }
+ if !strings.Contains(msg, "HTML instead of JSON") {
+ t.Errorf("expected HTML mention in error, got %v", msg)
+ }
+}
+
+// --- HandleErrorResponse with read failure ---
+
+func TestHandleErrorResponse_EmptyBody(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ w.WriteHeader(http.StatusInternalServerError)
+ // empty body
+ }))
+ defer server.Close()
+
+ resp, err := http.Get(server.URL)
+ if err != nil {
+ t.Fatalf("http.Get() error = %v", err)
+ }
+ defer resp.Body.Close()
+ err = HandleErrorResponse(resp, server.URL)
+ if err == nil {
+ t.Fatal("expected error")
+ }
+ if !strings.Contains(err.Error(), "500") {
+ t.Errorf("expected status code, got %v", err)
+ }
+}
+
+// --- ReadAndParseResponse with invalid JSON ---
+
+func TestReadAndParseResponse_InvalidJSON(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ w.Write([]byte("not valid json"))
+ }))
+ defer server.Close()
+
+ resp, err := http.Get(server.URL)
+ if err != nil {
+ t.Fatalf("http.Get() error = %v", err)
+ }
+ defer resp.Body.Close()
+ _, err = ReadAndParseResponse(resp, server.URL)
+ if err == nil {
+ t.Fatal("expected error for invalid JSON")
+ }
+}
+
+// --- ParseResponse with thought_signature (Google/Gemini) ---
+
+func TestParseResponse_WithThoughtSignature(t *testing.T) {
+ body := `{"choices":[{"message":{"content":"","tool_calls":[{"id":"call_1","type":"function","function":{"name":"test_tool","arguments":"{}"},"extra_content":{"google":{"thought_signature":"sig123"}}}]},"finish_reason":"tool_calls"}]}`
+ out, err := ParseResponse(strings.NewReader(body))
+ if err != nil {
+ t.Fatalf("ParseResponse() error = %v", err)
+ }
+ if len(out.ToolCalls) != 1 {
+ t.Fatalf("len(ToolCalls) = %d, want 1", len(out.ToolCalls))
+ }
+ if out.ToolCalls[0].ThoughtSignature != "sig123" {
+ t.Errorf("ThoughtSignature = %q, want %q", out.ToolCalls[0].ThoughtSignature, "sig123")
+ }
+ if out.ToolCalls[0].ExtraContent == nil || out.ToolCalls[0].ExtraContent.Google == nil {
+ t.Fatal("ExtraContent.Google is nil")
+ }
+ if out.ToolCalls[0].ExtraContent.Google.ThoughtSignature != "sig123" {
+ t.Errorf("ExtraContent.Google.ThoughtSignature = %q, want %q",
+ out.ToolCalls[0].ExtraContent.Google.ThoughtSignature, "sig123")
+ }
+}
diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go
index c269ae664..c7833f6e9 100644
--- a/pkg/providers/factory_provider.go
+++ b/pkg/providers/factory_provider.go
@@ -10,6 +10,8 @@ import (
"strings"
"github.com/sipeed/picoclaw/pkg/config"
+ anthropicmessages "github.com/sipeed/picoclaw/pkg/providers/anthropic_messages"
+ "github.com/sipeed/picoclaw/pkg/providers/azure"
)
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
@@ -56,7 +58,8 @@ func ExtractProtocol(model string) (protocol, modelID string) {
// CreateProviderFromConfig creates a provider based on the ModelConfig.
// It uses the protocol prefix in the Model field to determine which provider to create.
-// Supported protocols: openai, litellm, anthropic, antigravity, claude-cli, codex-cli, github-copilot
+// Supported protocols: openai, litellm, anthropic, anthropic-messages, antigravity,
+// claude-cli, codex-cli, github-copilot
// Returns the provider, the model ID (without protocol prefix), and any error.
func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) {
if cfg == nil {
@@ -95,10 +98,28 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout,
), modelID, nil
+ case "azure", "azure-openai":
+ // Azure OpenAI uses deployment-based URLs, api-key header auth,
+ // and always sends max_completion_tokens.
+ if cfg.APIKey == "" {
+ return nil, "", fmt.Errorf("api_key is required for azure protocol")
+ }
+ if cfg.APIBase == "" {
+ return nil, "", fmt.Errorf(
+ "api_base is required for azure protocol (e.g., https://your-resource.openai.azure.com)",
+ )
+ }
+ return azure.NewProviderWithTimeout(
+ cfg.APIKey,
+ cfg.APIBase,
+ cfg.Proxy,
+ cfg.RequestTimeout,
+ ), modelID, nil
+
case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
"vivgrid", "volcengine", "vllm", "qwen", "mistral", "avian",
- "minimax":
+ "minimax", "longcat", "modelscope":
// All other OpenAI-compatible HTTP providers
if cfg.APIKey == "" && cfg.APIBase == "" {
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
@@ -140,6 +161,21 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout,
), modelID, nil
+ case "anthropic-messages":
+ // Anthropic Messages API with native format (HTTP-based, no SDK)
+ apiBase := cfg.APIBase
+ if apiBase == "" {
+ apiBase = "https://api.anthropic.com/v1"
+ }
+ if cfg.APIKey == "" {
+ return nil, "", fmt.Errorf("api_key is required for anthropic-messages protocol (model: %s)", cfg.Model)
+ }
+ return anthropicmessages.NewProviderWithTimeout(
+ cfg.APIKey,
+ apiBase,
+ cfg.RequestTimeout,
+ ), modelID, nil
+
case "antigravity":
return NewAntigravityProvider(), modelID, nil
@@ -218,6 +254,10 @@ func getDefaultAPIBase(protocol string) string {
return "https://api.avian.io/v1"
case "minimax":
return "https://api.minimaxi.com/v1"
+ case "longcat":
+ return "https://api.longcat.chat/openai"
+ case "modelscope":
+ return "https://api-inference.modelscope.cn/v1"
default:
return ""
}
diff --git a/pkg/providers/factory_provider_test.go b/pkg/providers/factory_provider_test.go
index 7195d7cd1..57b3b48aa 100644
--- a/pkg/providers/factory_provider_test.go
+++ b/pkg/providers/factory_provider_test.go
@@ -145,6 +145,8 @@ func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) {
{"vllm", "vllm"},
{"deepseek", "deepseek"},
{"ollama", "ollama"},
+ {"longcat", "longcat"},
+ {"modelscope", "modelscope"},
}
for _, tt := range tests {
@@ -194,6 +196,58 @@ func TestCreateProviderFromConfig_LiteLLM(t *testing.T) {
}
}
+func TestCreateProviderFromConfig_LongCat(t *testing.T) {
+ cfg := &config.ModelConfig{
+ ModelName: "test-longcat",
+ Model: "longcat/LongCat-Flash-Thinking",
+ APIKey: "test-key",
+ APIBase: "https://api.longcat.chat/openai",
+ }
+
+ provider, modelID, err := CreateProviderFromConfig(cfg)
+ if err != nil {
+ t.Fatalf("CreateProviderFromConfig() error = %v", err)
+ }
+ if provider == nil {
+ t.Fatal("CreateProviderFromConfig() returned nil provider")
+ }
+ if modelID != "LongCat-Flash-Thinking" {
+ t.Errorf("modelID = %q, want %q", modelID, "LongCat-Flash-Thinking")
+ }
+ if _, ok := provider.(*HTTPProvider); !ok {
+ t.Fatalf("expected *HTTPProvider, got %T", provider)
+ }
+}
+
+func TestCreateProviderFromConfig_ModelScope(t *testing.T) {
+ cfg := &config.ModelConfig{
+ ModelName: "test-modelscope",
+ Model: "modelscope/Qwen/Qwen3-235B-A22B-Instruct-2507",
+ APIKey: "test-key",
+ APIBase: "https://api-inference.modelscope.cn/v1",
+ }
+
+ provider, modelID, err := CreateProviderFromConfig(cfg)
+ if err != nil {
+ t.Fatalf("CreateProviderFromConfig() error = %v", err)
+ }
+ if provider == nil {
+ t.Fatal("CreateProviderFromConfig() returned nil provider")
+ }
+ if modelID != "Qwen/Qwen3-235B-A22B-Instruct-2507" {
+ t.Errorf("modelID = %q, want %q", modelID, "Qwen/Qwen3-235B-A22B-Instruct-2507")
+ }
+ if _, ok := provider.(*HTTPProvider); !ok {
+ t.Fatalf("expected *HTTPProvider, got %T", provider)
+ }
+}
+
+func TestGetDefaultAPIBase_ModelScope(t *testing.T) {
+ if got := getDefaultAPIBase("modelscope"); got != "https://api-inference.modelscope.cn/v1" {
+ t.Fatalf("getDefaultAPIBase(%q) = %q, want %q", "modelscope", got, "https://api-inference.modelscope.cn/v1")
+ }
+}
+
func TestCreateProviderFromConfig_Anthropic(t *testing.T) {
cfg := &config.ModelConfig{
ModelName: "test-anthropic",
@@ -349,3 +403,69 @@ func TestCreateProviderFromConfig_RequestTimeoutPropagation(t *testing.T) {
t.Fatalf("Chat() error = %q, want timeout-related error", errMsg)
}
}
+
+func TestCreateProviderFromConfig_Azure(t *testing.T) {
+ cfg := &config.ModelConfig{
+ ModelName: "azure-gpt5",
+ Model: "azure/my-gpt5-deployment",
+ APIKey: "test-azure-key",
+ APIBase: "https://my-resource.openai.azure.com",
+ }
+
+ provider, modelID, err := CreateProviderFromConfig(cfg)
+ if err != nil {
+ t.Fatalf("CreateProviderFromConfig() error = %v", err)
+ }
+ if provider == nil {
+ t.Fatal("CreateProviderFromConfig() returned nil provider")
+ }
+ if modelID != "my-gpt5-deployment" {
+ t.Errorf("modelID = %q, want %q", modelID, "my-gpt5-deployment")
+ }
+}
+
+func TestCreateProviderFromConfig_AzureOpenAIAlias(t *testing.T) {
+ cfg := &config.ModelConfig{
+ ModelName: "azure-gpt4",
+ Model: "azure-openai/my-deployment",
+ APIKey: "test-azure-key",
+ APIBase: "https://my-resource.openai.azure.com",
+ }
+
+ provider, modelID, err := CreateProviderFromConfig(cfg)
+ if err != nil {
+ t.Fatalf("CreateProviderFromConfig() error = %v", err)
+ }
+ if provider == nil {
+ t.Fatal("CreateProviderFromConfig() returned nil provider")
+ }
+ if modelID != "my-deployment" {
+ t.Errorf("modelID = %q, want %q", modelID, "my-deployment")
+ }
+}
+
+func TestCreateProviderFromConfig_AzureMissingAPIKey(t *testing.T) {
+ cfg := &config.ModelConfig{
+ ModelName: "azure-gpt5",
+ Model: "azure/my-gpt5-deployment",
+ APIBase: "https://my-resource.openai.azure.com",
+ }
+
+ _, _, err := CreateProviderFromConfig(cfg)
+ if err == nil {
+ t.Fatal("CreateProviderFromConfig() expected error for missing API key")
+ }
+}
+
+func TestCreateProviderFromConfig_AzureMissingAPIBase(t *testing.T) {
+ cfg := &config.ModelConfig{
+ ModelName: "azure-gpt5",
+ Model: "azure/my-gpt5-deployment",
+ APIKey: "test-azure-key",
+ }
+
+ _, _, err := CreateProviderFromConfig(cfg)
+ if err == nil {
+ t.Fatal("CreateProviderFromConfig() expected error for missing API base")
+ }
+}