Previously, NewAntigravityProvider() would create a provider even when OAuth credentials didn't exist, only failing when Chat() was called. This prevented the fallback mechanism from working correctly. Changes: - NewAntigravityProvider() now returns (*AntigravityProvider, error) - Validates credentials exist at creation time (not during Chat()) - Updated CreateProviderFromConfig() to handle the error - Updated test to expect credential validation error This ensures that when antigravity OAuth is logged out or credentials are missing, the provider creation fails immediately, allowing the fallback chain to try alternative models/providers. Fixes issue where gemini-flash model with expired antigravity OAuth would not fall back to other configured providers.
216 lines
6.2 KiB
Go
216 lines
6.2 KiB
Go
// PicoClaw - Ultra-lightweight personal AI agent
|
|
// License: MIT
|
|
//
|
|
// Copyright (c) 2026 PicoClaw contributors
|
|
|
|
package providers
|
|
|
|
import (
|
|
"fmt"
|
|
"strings"
|
|
|
|
"github.com/sipeed/picoclaw/pkg/config"
|
|
)
|
|
|
|
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
|
|
func createClaudeAuthProvider() (LLMProvider, error) {
|
|
cred, err := getCredential("anthropic")
|
|
if err != nil {
|
|
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
|
}
|
|
if cred == nil {
|
|
return nil, fmt.Errorf("no credentials for anthropic. Run: picoclaw auth login --provider anthropic")
|
|
}
|
|
return NewClaudeProviderWithTokenSource(cred.AccessToken, createClaudeTokenSource()), nil
|
|
}
|
|
|
|
// createCodexAuthProvider creates a Codex provider using OAuth credentials from auth store.
|
|
func createCodexAuthProvider() (LLMProvider, error) {
|
|
cred, err := getCredential("openai")
|
|
if err != nil {
|
|
return nil, fmt.Errorf("loading auth credentials: %w", err)
|
|
}
|
|
if cred == nil {
|
|
return nil, fmt.Errorf("no credentials for openai. Run: picoclaw auth login --provider openai")
|
|
}
|
|
return NewCodexProviderWithTokenSource(cred.AccessToken, cred.AccountID, createCodexTokenSource()), nil
|
|
}
|
|
|
|
// ExtractProtocol extracts the protocol prefix and model identifier from a model string.
|
|
// If no prefix is specified, it defaults to "openai".
|
|
// Examples:
|
|
// - "openai/gpt-4o" -> ("openai", "gpt-4o")
|
|
// - "anthropic/claude-sonnet-4.6" -> ("anthropic", "claude-sonnet-4.6")
|
|
// - "gpt-4o" -> ("openai", "gpt-4o") // default protocol
|
|
func ExtractProtocol(model string) (protocol, modelID string) {
|
|
model = strings.TrimSpace(model)
|
|
protocol, modelID, found := strings.Cut(model, "/")
|
|
if !found {
|
|
return "openai", model
|
|
}
|
|
return protocol, modelID
|
|
}
|
|
|
|
// 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, anthropic, 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 {
|
|
return nil, "", fmt.Errorf("config is nil")
|
|
}
|
|
|
|
if cfg.Model == "" {
|
|
return nil, "", fmt.Errorf("model is required")
|
|
}
|
|
|
|
protocol, modelID := ExtractProtocol(cfg.Model)
|
|
|
|
switch protocol {
|
|
case "openai":
|
|
// OpenAI with OAuth/token auth (Codex-style)
|
|
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
|
|
provider, err := createCodexAuthProvider()
|
|
if err != nil {
|
|
return nil, "", err
|
|
}
|
|
return provider, modelID, nil
|
|
}
|
|
// OpenAI with API key
|
|
if cfg.APIKey == "" && cfg.APIBase == "" {
|
|
return nil, "", fmt.Errorf("api_key or api_base is required for HTTP-based protocol %q", protocol)
|
|
}
|
|
apiBase := cfg.APIBase
|
|
if apiBase == "" {
|
|
apiBase = getDefaultAPIBase(protocol)
|
|
}
|
|
return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
|
|
cfg.APIKey,
|
|
apiBase,
|
|
cfg.Proxy,
|
|
cfg.MaxTokensField,
|
|
cfg.RequestTimeout,
|
|
), modelID, nil
|
|
|
|
case "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
|
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
|
"volcengine", "vllm", "qwen", "mistral":
|
|
// 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)
|
|
}
|
|
apiBase := cfg.APIBase
|
|
if apiBase == "" {
|
|
apiBase = getDefaultAPIBase(protocol)
|
|
}
|
|
return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
|
|
cfg.APIKey,
|
|
apiBase,
|
|
cfg.Proxy,
|
|
cfg.MaxTokensField,
|
|
cfg.RequestTimeout,
|
|
), modelID, nil
|
|
|
|
case "anthropic":
|
|
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
|
|
// Use OAuth credentials from auth store
|
|
provider, err := createClaudeAuthProvider()
|
|
if err != nil {
|
|
return nil, "", err
|
|
}
|
|
return provider, modelID, nil
|
|
}
|
|
// Use API key with HTTP API
|
|
apiBase := cfg.APIBase
|
|
if apiBase == "" {
|
|
apiBase = "https://api.anthropic.com/v1"
|
|
}
|
|
if cfg.APIKey == "" {
|
|
return nil, "", fmt.Errorf("api_key is required for anthropic protocol (model: %s)", cfg.Model)
|
|
}
|
|
return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
|
|
cfg.APIKey,
|
|
apiBase,
|
|
cfg.Proxy,
|
|
cfg.MaxTokensField,
|
|
cfg.RequestTimeout,
|
|
), modelID, nil
|
|
|
|
case "antigravity":
|
|
provider, err := NewAntigravityProvider()
|
|
if err != nil {
|
|
return nil, "", err
|
|
}
|
|
return provider, modelID, nil
|
|
|
|
case "claude-cli", "claudecli":
|
|
workspace := cfg.Workspace
|
|
if workspace == "" {
|
|
workspace = "."
|
|
}
|
|
return NewClaudeCliProvider(workspace), modelID, nil
|
|
|
|
case "codex-cli", "codexcli":
|
|
workspace := cfg.Workspace
|
|
if workspace == "" {
|
|
workspace = "."
|
|
}
|
|
return NewCodexCliProvider(workspace), modelID, nil
|
|
|
|
case "github-copilot", "copilot":
|
|
apiBase := cfg.APIBase
|
|
if apiBase == "" {
|
|
apiBase = "localhost:4321"
|
|
}
|
|
connectMode := cfg.ConnectMode
|
|
if connectMode == "" {
|
|
connectMode = "grpc"
|
|
}
|
|
provider, err := NewGitHubCopilotProvider(apiBase, connectMode, modelID)
|
|
if err != nil {
|
|
return nil, "", err
|
|
}
|
|
return provider, modelID, nil
|
|
|
|
default:
|
|
return nil, "", fmt.Errorf("unknown protocol %q in model %q", protocol, cfg.Model)
|
|
}
|
|
}
|
|
|
|
// getDefaultAPIBase returns the default API base URL for a given protocol.
|
|
func getDefaultAPIBase(protocol string) string {
|
|
switch protocol {
|
|
case "openai":
|
|
return "https://api.openai.com/v1"
|
|
case "openrouter":
|
|
return "https://openrouter.ai/api/v1"
|
|
case "groq":
|
|
return "https://api.groq.com/openai/v1"
|
|
case "zhipu":
|
|
return "https://open.bigmodel.cn/api/paas/v4"
|
|
case "gemini":
|
|
return "https://generativelanguage.googleapis.com/v1beta"
|
|
case "nvidia":
|
|
return "https://integrate.api.nvidia.com/v1"
|
|
case "ollama":
|
|
return "http://localhost:11434/v1"
|
|
case "moonshot":
|
|
return "https://api.moonshot.cn/v1"
|
|
case "shengsuanyun":
|
|
return "https://router.shengsuanyun.com/api/v1"
|
|
case "deepseek":
|
|
return "https://api.deepseek.com/v1"
|
|
case "cerebras":
|
|
return "https://api.cerebras.ai/v1"
|
|
case "volcengine":
|
|
return "https://ark.cn-beijing.volces.com/api/v3"
|
|
case "qwen":
|
|
return "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
|
case "vllm":
|
|
return "http://localhost:8000/v1"
|
|
case "mistral":
|
|
return "https://api.mistral.ai/v1"
|
|
default:
|
|
return ""
|
|
}
|
|
}
|