diff --git a/pkg/config/config.go b/pkg/config/config.go index 197b95973..a7965bbac 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -152,10 +152,12 @@ type ProvidersConfig struct { } type ProviderConfig struct { - APIKey string `json:"api_key" env:"PICOCLAW_PROVIDERS_{{.Name}}_API_KEY"` - APIBase string `json:"api_base" env:"PICOCLAW_PROVIDERS_{{.Name}}_API_BASE"` - Proxy string `json:"proxy,omitempty" env:"PICOCLAW_PROVIDERS_{{.Name}}_PROXY"` - AuthMethod string `json:"auth_method,omitempty" env:"PICOCLAW_PROVIDERS_{{.Name}}_AUTH_METHOD"` + APIKey string `json:"api_key" env:"PICOCLAW_PROVIDERS_{{.Name}}_API_KEY"` + APIBase string `json:"api_base" env:"PICOCLAW_PROVIDERS_{{.Name}}_API_BASE"` + API string `json:"api,omitempty" env:"PICOCLAW_PROVIDERS_{{.Name}}_API"` + Headers map[string]string `json:"headers,omitempty"` + Proxy string `json:"proxy,omitempty" env:"PICOCLAW_PROVIDERS_{{.Name}}_PROXY"` + AuthMethod string `json:"auth_method,omitempty" env:"PICOCLAW_PROVIDERS_{{.Name}}_AUTH_METHOD"` } type GatewayConfig struct { diff --git a/pkg/providers/http_provider.go b/pkg/providers/http_provider.go index 7179c4cc5..16fb99459 100644 --- a/pkg/providers/http_provider.go +++ b/pkg/providers/http_provider.go @@ -10,23 +10,39 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" "net/http" "net/url" "strings" + "github.com/openai/openai-go/v3/responses" "github.com/sipeed/picoclaw/pkg/auth" "github.com/sipeed/picoclaw/pkg/config" + "github.com/sipeed/picoclaw/pkg/logger" + "github.com/sipeed/picoclaw/pkg/utils" ) type HTTPProvider struct { apiKey string apiBase string + apiMode string + headers map[string]string httpClient *http.Client } -func NewHTTPProvider(apiKey, apiBase, proxy string) *HTTPProvider { +type httpProviderError struct { + statusCode int + body string + url string +} + +func (e *httpProviderError) Error() string { + return fmt.Sprintf("API error (%d): %s", e.statusCode, e.body) +} + +func NewHTTPProvider(apiKey, apiBase, proxy, apiMode string, headers map[string]string) *HTTPProvider { client := &http.Client{ Timeout: 0, } @@ -43,6 +59,8 @@ func NewHTTPProvider(apiKey, apiBase, proxy string) *HTTPProvider { return &HTTPProvider{ apiKey: apiKey, apiBase: apiBase, + apiMode: apiMode, + headers: headers, httpClient: client, } } @@ -52,74 +70,32 @@ func (p *HTTPProvider) Chat(ctx context.Context, messages []Message, tools []Too return nil, fmt.Errorf("API base not configured") } - // Strip provider prefix from model name (e.g., moonshot/kimi-k2.5 -> kimi-k2.5) - if idx := strings.Index(model, "/"); idx != -1 { - prefix := model[:idx] - if prefix == "moonshot" || prefix == "nvidia" { - model = model[idx+1:] + useResponses := shouldPreferResponses(model, p.apiMode) + if useResponses { + resp, err := p.chatWithResponses(ctx, messages, tools, model, options) + if err == nil { + return resp, nil } - } - - requestBody := map[string]interface{}{ - "model": model, - "messages": messages, - } - - if len(tools) > 0 { - requestBody["tools"] = tools - requestBody["tool_choice"] = "auto" - } - - if maxTokens, ok := options["max_tokens"].(int); ok { - lowerModel := strings.ToLower(model) - if strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "o1") { - requestBody["max_completion_tokens"] = maxTokens - } else { - requestBody["max_tokens"] = maxTokens + if shouldFallbackFromResponses(err) { + logger.DebugCF("provider", "Responses endpoint unsupported, falling back to chat/completions", map[string]interface{}{ + "model": model, + }) + return p.chatWithCompletions(ctx, messages, tools, model, options) } + return nil, err } - if temperature, ok := options["temperature"].(float64); ok { - lowerModel := strings.ToLower(model) - // Kimi k2 models only support temperature=1 - if strings.Contains(lowerModel, "kimi") && strings.Contains(lowerModel, "k2") { - requestBody["temperature"] = 1.0 - } else { - requestBody["temperature"] = temperature - } + resp, err := p.chatWithCompletions(ctx, messages, tools, model, options) + if err == nil { + return resp, nil } - - jsonData, err := json.Marshal(requestBody) - if err != nil { - return nil, fmt.Errorf("failed to marshal request: %w", err) + if shouldFallbackFromCompletions(err) { + logger.DebugCF("provider", "Chat/completions endpoint unsupported, falling back to responses", map[string]interface{}{ + "model": model, + }) + return p.chatWithResponses(ctx, messages, tools, model, options) } - - req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+"/chat/completions", bytes.NewReader(jsonData)) - if err != nil { - return nil, fmt.Errorf("failed to create request: %w", err) - } - - req.Header.Set("Content-Type", "application/json") - if p.apiKey != "" { - req.Header.Set("Authorization", "Bearer "+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() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return nil, fmt.Errorf("failed to read response: %w", err) - } - - if resp.StatusCode != http.StatusOK { - return nil, fmt.Errorf("API error: %s", string(body)) - } - - return p.parseResponse(body) + return nil, err } func (p *HTTPProvider) parseResponse(body []byte) (*LLMResponse, error) { @@ -196,6 +172,202 @@ func (p *HTTPProvider) GetDefaultModel() string { return "" } +func (p *HTTPProvider) applyHeaders(req *http.Request) { + if len(p.headers) > 0 { + for k, v := range p.headers { + req.Header.Set(k, v) + } + } + + if req.Header.Get("Content-Type") == "" { + req.Header.Set("Content-Type", "application/json") + } + if p.apiKey != "" && req.Header.Get("Authorization") == "" { + req.Header.Set("Authorization", "Bearer "+p.apiKey) + } +} + +func (p *HTTPProvider) chatWithCompletions(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) { + model = normalizeModelForHTTP(model) + requestBody := map[string]interface{}{ + "model": model, + "messages": messages, + } + + if len(tools) > 0 { + requestBody["tools"] = tools + requestBody["tool_choice"] = "auto" + } + + if maxTokens, ok := options["max_tokens"].(int); ok { + lowerModel := strings.ToLower(model) + if strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "o1") { + requestBody["max_completion_tokens"] = maxTokens + } else { + requestBody["max_tokens"] = maxTokens + } + } + + if temperature, ok := options["temperature"].(float64); ok { + lowerModel := strings.ToLower(model) + // Kimi k2 models only support temperature=1 + if strings.Contains(lowerModel, "kimi") && strings.Contains(lowerModel, "k2") { + requestBody["temperature"] = 1.0 + } else { + 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", p.apiBase+"/chat/completions", bytes.NewReader(jsonData)) + if err != nil { + return nil, fmt.Errorf("failed to create request: %w", err) + } + logger.DebugCF("provider", "HTTP request", map[string]interface{}{ + "url": req.URL.String(), + "method": req.Method, + }) + + p.applyHeaders(req) + + resp, err := p.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("failed to send request: %w", err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("failed to read response: %w", err) + } + + if resp.StatusCode != http.StatusOK { + logger.DebugCF("provider", "HTTP response error", map[string]interface{}{ + "status": resp.StatusCode, + "body": utils.Truncate(string(body), 500), + }) + return nil, &httpProviderError{statusCode: resp.StatusCode, body: string(body), url: req.URL.String()} + } + + return p.parseResponse(body) +} + +func (p *HTTPProvider) chatWithResponses(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) { + model = normalizeModelForHTTP(model) + params := buildCodexParams(messages, tools, model, stripTemperature(options)) + + jsonData, err := json.Marshal(params) + if err != nil { + return nil, fmt.Errorf("failed to marshal responses request: %w", err) + } + + req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+"/responses", bytes.NewReader(jsonData)) + if err != nil { + return nil, fmt.Errorf("failed to create request: %w", err) + } + logger.DebugCF("provider", "HTTP request", map[string]interface{}{ + "url": req.URL.String(), + "method": req.Method, + }) + + p.applyHeaders(req) + + resp, err := p.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("failed to send request: %w", err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("failed to read response: %w", err) + } + + if resp.StatusCode != http.StatusOK { + logger.DebugCF("provider", "HTTP response error", map[string]interface{}{ + "status": resp.StatusCode, + "body": utils.Truncate(string(body), 500), + }) + return nil, &httpProviderError{statusCode: resp.StatusCode, body: string(body), url: req.URL.String()} + } + + var apiResponse responses.Response + if err := json.Unmarshal(body, &apiResponse); err != nil { + return nil, fmt.Errorf("failed to unmarshal responses API response: %w", err) + } + return parseCodexResponse(&apiResponse), nil +} + +func normalizeModelForHTTP(model string) string { + // Strip provider prefix from model name (e.g., moonshot/kimi-k2.5 -> kimi-k2.5) + if idx := strings.Index(model, "/"); idx != -1 { + prefix := model[:idx] + if prefix == "moonshot" || prefix == "nvidia" { + return model[idx+1:] + } + } + return model +} + +func stripTemperature(options map[string]interface{}) map[string]interface{} { + if options == nil { + return nil + } + if _, ok := options["temperature"]; !ok { + return options + } + cleaned := make(map[string]interface{}, len(options)-1) + for k, v := range options { + if k == "temperature" { + continue + } + cleaned[k] = v + } + return cleaned +} + +func shouldPreferResponses(model, apiMode string) bool { + lowerMode := strings.ToLower(apiMode) + switch lowerMode { + case "openai-responses", "responses", "response": + return true + case "openai-completions", "chat-completions", "completions": + return false + } + + lower := strings.ToLower(model) + return strings.Contains(lower, "gpt-5") || strings.Contains(lower, "codex") || strings.Contains(lower, "o1") +} + +func shouldFallbackFromResponses(err error) bool { + var httpErr *httpProviderError + if errors.As(err, &httpErr) { + return isEndpointUnsupported(httpErr.statusCode) + } + return false +} + +func shouldFallbackFromCompletions(err error) bool { + var httpErr *httpProviderError + if errors.As(err, &httpErr) { + return isEndpointUnsupported(httpErr.statusCode) + } + return false +} + +func isEndpointUnsupported(statusCode int) bool { + switch statusCode { + case http.StatusNotFound, http.StatusMethodNotAllowed, http.StatusNotImplemented, http.StatusGone: + return true + default: + return false + } +} + func createClaudeAuthProvider() (LLMProvider, error) { cred, err := auth.GetCredential("anthropic") if err != nil { @@ -222,7 +394,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { model := cfg.Agents.Defaults.Model providerName := strings.ToLower(cfg.Agents.Defaults.Provider) - var apiKey, apiBase, proxy string + var apiKey, apiBase, proxy, apiMode string + var headers map[string]string lowerModel := strings.ToLower(model) @@ -233,6 +406,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { if cfg.Providers.Groq.APIKey != "" { apiKey = cfg.Providers.Groq.APIKey apiBase = cfg.Providers.Groq.APIBase + apiMode = cfg.Providers.Groq.API + headers = cfg.Providers.Groq.Headers if apiBase == "" { apiBase = "https://api.groq.com/openai/v1" } @@ -244,6 +419,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { } apiKey = cfg.Providers.OpenAI.APIKey apiBase = cfg.Providers.OpenAI.APIBase + apiMode = cfg.Providers.OpenAI.API + headers = cfg.Providers.OpenAI.Headers if apiBase == "" { apiBase = "https://api.openai.com/v1" } @@ -255,6 +432,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { } apiKey = cfg.Providers.Anthropic.APIKey apiBase = cfg.Providers.Anthropic.APIBase + apiMode = cfg.Providers.Anthropic.API + headers = cfg.Providers.Anthropic.Headers if apiBase == "" { apiBase = "https://api.anthropic.com/v1" } @@ -262,6 +441,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { case "openrouter": if cfg.Providers.OpenRouter.APIKey != "" { apiKey = cfg.Providers.OpenRouter.APIKey + apiMode = cfg.Providers.OpenRouter.API + headers = cfg.Providers.OpenRouter.Headers if cfg.Providers.OpenRouter.APIBase != "" { apiBase = cfg.Providers.OpenRouter.APIBase } else { @@ -272,6 +453,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { if cfg.Providers.Zhipu.APIKey != "" { apiKey = cfg.Providers.Zhipu.APIKey apiBase = cfg.Providers.Zhipu.APIBase + apiMode = cfg.Providers.Zhipu.API + headers = cfg.Providers.Zhipu.Headers if apiBase == "" { apiBase = "https://open.bigmodel.cn/api/paas/v4" } @@ -280,6 +463,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { if cfg.Providers.Gemini.APIKey != "" { apiKey = cfg.Providers.Gemini.APIKey apiBase = cfg.Providers.Gemini.APIBase + apiMode = cfg.Providers.Gemini.API + headers = cfg.Providers.Gemini.Headers if apiBase == "" { apiBase = "https://generativelanguage.googleapis.com/v1beta" } @@ -288,6 +473,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { if cfg.Providers.VLLM.APIBase != "" { apiKey = cfg.Providers.VLLM.APIKey apiBase = cfg.Providers.VLLM.APIBase + apiMode = cfg.Providers.VLLM.API + headers = cfg.Providers.VLLM.Headers } case "claude-cli", "claudecode", "claude-code": workspace := cfg.Agents.Defaults.Workspace @@ -305,6 +492,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { apiKey = cfg.Providers.Moonshot.APIKey apiBase = cfg.Providers.Moonshot.APIBase proxy = cfg.Providers.Moonshot.Proxy + apiMode = cfg.Providers.Moonshot.API + headers = cfg.Providers.Moonshot.Headers if apiBase == "" { apiBase = "https://api.moonshot.cn/v1" } @@ -312,6 +501,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { case strings.HasPrefix(model, "openrouter/") || strings.HasPrefix(model, "anthropic/") || strings.HasPrefix(model, "openai/") || strings.HasPrefix(model, "meta-llama/") || strings.HasPrefix(model, "deepseek/") || strings.HasPrefix(model, "google/"): apiKey = cfg.Providers.OpenRouter.APIKey proxy = cfg.Providers.OpenRouter.Proxy + apiMode = cfg.Providers.OpenRouter.API + headers = cfg.Providers.OpenRouter.Headers if cfg.Providers.OpenRouter.APIBase != "" { apiBase = cfg.Providers.OpenRouter.APIBase } else { @@ -325,6 +516,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { apiKey = cfg.Providers.Anthropic.APIKey apiBase = cfg.Providers.Anthropic.APIBase proxy = cfg.Providers.Anthropic.Proxy + apiMode = cfg.Providers.Anthropic.API + headers = cfg.Providers.Anthropic.Headers if apiBase == "" { apiBase = "https://api.anthropic.com/v1" } @@ -336,6 +529,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { apiKey = cfg.Providers.OpenAI.APIKey apiBase = cfg.Providers.OpenAI.APIBase proxy = cfg.Providers.OpenAI.Proxy + apiMode = cfg.Providers.OpenAI.API + headers = cfg.Providers.OpenAI.Headers if apiBase == "" { apiBase = "https://api.openai.com/v1" } @@ -344,6 +539,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { apiKey = cfg.Providers.Gemini.APIKey apiBase = cfg.Providers.Gemini.APIBase proxy = cfg.Providers.Gemini.Proxy + apiMode = cfg.Providers.Gemini.API + headers = cfg.Providers.Gemini.Headers if apiBase == "" { apiBase = "https://generativelanguage.googleapis.com/v1beta" } @@ -352,6 +549,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { apiKey = cfg.Providers.Zhipu.APIKey apiBase = cfg.Providers.Zhipu.APIBase proxy = cfg.Providers.Zhipu.Proxy + apiMode = cfg.Providers.Zhipu.API + headers = cfg.Providers.Zhipu.Headers if apiBase == "" { apiBase = "https://open.bigmodel.cn/api/paas/v4" } @@ -360,6 +559,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { apiKey = cfg.Providers.Groq.APIKey apiBase = cfg.Providers.Groq.APIBase proxy = cfg.Providers.Groq.Proxy + apiMode = cfg.Providers.Groq.API + headers = cfg.Providers.Groq.Headers if apiBase == "" { apiBase = "https://api.groq.com/openai/v1" } @@ -368,6 +569,8 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { apiKey = cfg.Providers.Nvidia.APIKey apiBase = cfg.Providers.Nvidia.APIBase proxy = cfg.Providers.Nvidia.Proxy + apiMode = cfg.Providers.Nvidia.API + headers = cfg.Providers.Nvidia.Headers if apiBase == "" { apiBase = "https://integrate.api.nvidia.com/v1" } @@ -376,11 +579,15 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { apiKey = cfg.Providers.VLLM.APIKey apiBase = cfg.Providers.VLLM.APIBase proxy = cfg.Providers.VLLM.Proxy + apiMode = cfg.Providers.VLLM.API + headers = cfg.Providers.VLLM.Headers default: if cfg.Providers.OpenRouter.APIKey != "" { apiKey = cfg.Providers.OpenRouter.APIKey proxy = cfg.Providers.OpenRouter.Proxy + apiMode = cfg.Providers.OpenRouter.API + headers = cfg.Providers.OpenRouter.Headers if cfg.Providers.OpenRouter.APIBase != "" { apiBase = cfg.Providers.OpenRouter.APIBase } else { @@ -400,5 +607,5 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { return nil, fmt.Errorf("no API base configured for provider (model: %s)", model) } - return NewHTTPProvider(apiKey, apiBase, proxy), nil + return NewHTTPProvider(apiKey, apiBase, proxy, apiMode, headers), nil }