refactor(provider): consolidate protocol metadata and local tests
This commit is contained in:
parent
47b95ef3fa
commit
9ee422c12d
2 changed files with 120 additions and 187 deletions
|
|
@ -17,6 +17,48 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/providers/bedrock"
|
"github.com/sipeed/picoclaw/pkg/providers/bedrock"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type protocolMeta struct {
|
||||||
|
defaultAPIBase string
|
||||||
|
emptyAPIKeyAllowed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
var protocolMetaByName = map[string]protocolMeta{
|
||||||
|
"openai": {defaultAPIBase: "https://api.openai.com/v1"},
|
||||||
|
"openrouter": {defaultAPIBase: "https://openrouter.ai/api/v1"},
|
||||||
|
"litellm": {defaultAPIBase: "http://localhost:4000/v1"},
|
||||||
|
"lmstudio": {defaultAPIBase: "http://localhost:1234/v1", emptyAPIKeyAllowed: true},
|
||||||
|
"novita": {defaultAPIBase: "https://api.novita.ai/openai"},
|
||||||
|
"groq": {defaultAPIBase: "https://api.groq.com/openai/v1"},
|
||||||
|
"zhipu": {defaultAPIBase: "https://open.bigmodel.cn/api/paas/v4"},
|
||||||
|
"gemini": {defaultAPIBase: "https://generativelanguage.googleapis.com/v1beta"},
|
||||||
|
"nvidia": {defaultAPIBase: "https://integrate.api.nvidia.com/v1"},
|
||||||
|
"ollama": {defaultAPIBase: "http://localhost:11434/v1", emptyAPIKeyAllowed: true},
|
||||||
|
"moonshot": {defaultAPIBase: "https://api.moonshot.cn/v1"},
|
||||||
|
"shengsuanyun": {defaultAPIBase: "https://router.shengsuanyun.com/api/v1"},
|
||||||
|
"deepseek": {defaultAPIBase: "https://api.deepseek.com/v1"},
|
||||||
|
"cerebras": {defaultAPIBase: "https://api.cerebras.ai/v1"},
|
||||||
|
"vivgrid": {defaultAPIBase: "https://api.vivgrid.com/v1"},
|
||||||
|
"volcengine": {defaultAPIBase: "https://ark.cn-beijing.volces.com/api/v3"},
|
||||||
|
"qwen": {defaultAPIBase: "https://dashscope.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"qwen-intl": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"qwen-international": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"dashscope-intl": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"qwen-us": {defaultAPIBase: "https://dashscope-us.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"dashscope-us": {defaultAPIBase: "https://dashscope-us.aliyuncs.com/compatible-mode/v1"},
|
||||||
|
"coding-plan": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"},
|
||||||
|
"alibaba-coding": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"},
|
||||||
|
"qwen-coding": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"},
|
||||||
|
"coding-plan-anthropic": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic"},
|
||||||
|
"alibaba-coding-anthropic": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic"},
|
||||||
|
"vllm": {defaultAPIBase: "http://localhost:8000/v1", emptyAPIKeyAllowed: true},
|
||||||
|
"mistral": {defaultAPIBase: "https://api.mistral.ai/v1"},
|
||||||
|
"avian": {defaultAPIBase: "https://api.avian.io/v1"},
|
||||||
|
"minimax": {defaultAPIBase: "https://api.minimaxi.com/v1"},
|
||||||
|
"longcat": {defaultAPIBase: "https://api.longcat.chat/openai"},
|
||||||
|
"modelscope": {defaultAPIBase: "https://api-inference.modelscope.cn/v1"},
|
||||||
|
"mimo": {defaultAPIBase: "https://api.xiaomimimo.com/v1"},
|
||||||
|
}
|
||||||
|
|
||||||
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
|
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
|
||||||
func createClaudeAuthProvider() (LLMProvider, error) {
|
func createClaudeAuthProvider() (LLMProvider, error) {
|
||||||
cred, err := getCredential("anthropic")
|
cred, err := getCredential("anthropic")
|
||||||
|
|
@ -295,74 +337,15 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func isEmptyAPIKeyAllowed(protocol string) bool {
|
func isEmptyAPIKeyAllowed(protocol string) bool {
|
||||||
switch protocol {
|
meta, ok := protocolMetaByName[protocol]
|
||||||
case "ollama", "vllm", "lmstudio":
|
return ok && meta.emptyAPIKeyAllowed
|
||||||
return true
|
|
||||||
default:
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// getDefaultAPIBase returns the default API base URL for a given protocol.
|
// getDefaultAPIBase returns the default API base URL for a given protocol.
|
||||||
func getDefaultAPIBase(protocol string) string {
|
func getDefaultAPIBase(protocol string) string {
|
||||||
switch protocol {
|
meta, ok := protocolMetaByName[protocol]
|
||||||
case "openai":
|
if !ok {
|
||||||
return "https://api.openai.com/v1"
|
|
||||||
case "openrouter":
|
|
||||||
return "https://openrouter.ai/api/v1"
|
|
||||||
case "litellm":
|
|
||||||
return "http://localhost:4000/v1"
|
|
||||||
case "lmstudio":
|
|
||||||
return "http://localhost:1234/v1"
|
|
||||||
case "novita":
|
|
||||||
return "https://api.novita.ai/openai"
|
|
||||||
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 "vivgrid":
|
|
||||||
return "https://api.vivgrid.com/v1"
|
|
||||||
case "volcengine":
|
|
||||||
return "https://ark.cn-beijing.volces.com/api/v3"
|
|
||||||
case "qwen":
|
|
||||||
return "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
|
||||||
case "qwen-intl", "qwen-international", "dashscope-intl":
|
|
||||||
return "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"
|
|
||||||
case "qwen-us", "dashscope-us":
|
|
||||||
return "https://dashscope-us.aliyuncs.com/compatible-mode/v1"
|
|
||||||
case "coding-plan", "alibaba-coding", "qwen-coding":
|
|
||||||
return "https://coding-intl.dashscope.aliyuncs.com/v1"
|
|
||||||
case "coding-plan-anthropic", "alibaba-coding-anthropic":
|
|
||||||
return "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic"
|
|
||||||
case "vllm":
|
|
||||||
return "http://localhost:8000/v1"
|
|
||||||
case "mistral":
|
|
||||||
return "https://api.mistral.ai/v1"
|
|
||||||
case "avian":
|
|
||||||
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"
|
|
||||||
case "mimo":
|
|
||||||
return "https://api.xiaomimimo.com/v1"
|
|
||||||
default:
|
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
return meta.defaultAPIBase
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -180,32 +180,66 @@ func TestCreateProviderFromConfig_LiteLLM(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_LMStudio_WithAPIKey(t *testing.T) {
|
func TestCreateProviderFromConfig_LocalProviders(t *testing.T) {
|
||||||
cfg := &config.ModelConfig{
|
tests := []struct {
|
||||||
ModelName: "test-lmstudio",
|
name string
|
||||||
Model: "lmstudio/openai/gpt-oss-20b",
|
modelName string
|
||||||
|
model string
|
||||||
|
apiKey string
|
||||||
|
wantModelID string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "LMStudio with API key",
|
||||||
|
modelName: "test-lmstudio",
|
||||||
|
model: "lmstudio/openai/gpt-oss-20b",
|
||||||
|
apiKey: "test-key",
|
||||||
|
wantModelID: "openai/gpt-oss-20b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "LMStudio without API key",
|
||||||
|
modelName: "test-lmstudio",
|
||||||
|
model: "lmstudio/openai/gpt-oss-20b",
|
||||||
|
apiKey: "",
|
||||||
|
wantModelID: "openai/gpt-oss-20b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Ollama with API key",
|
||||||
|
modelName: "test-ollama",
|
||||||
|
model: "ollama/llama3.1:8b",
|
||||||
|
apiKey: "test-key",
|
||||||
|
wantModelID: "llama3.1:8b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Ollama without API key",
|
||||||
|
modelName: "test-ollama",
|
||||||
|
model: "ollama/llama3.1:8b",
|
||||||
|
apiKey: "",
|
||||||
|
wantModelID: "llama3.1:8b",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "VLLM with API key",
|
||||||
|
modelName: "test-vllm",
|
||||||
|
model: "vllm/Qwen/Qwen3-8B",
|
||||||
|
apiKey: "test-key",
|
||||||
|
wantModelID: "Qwen/Qwen3-8B",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "VLLM without API key",
|
||||||
|
modelName: "test-vllm",
|
||||||
|
model: "vllm/Qwen/Qwen3-8B",
|
||||||
|
apiKey: "",
|
||||||
|
wantModelID: "Qwen/Qwen3-8B",
|
||||||
|
},
|
||||||
}
|
}
|
||||||
cfg.SetAPIKey("test-key")
|
|
||||||
|
|
||||||
provider, modelID, err := CreateProviderFromConfig(cfg)
|
for _, tt := range tests {
|
||||||
if err != nil {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
t.Fatalf("CreateProviderFromConfig() error = %v", err)
|
|
||||||
}
|
|
||||||
if provider == nil {
|
|
||||||
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
|
||||||
}
|
|
||||||
if modelID != "openai/gpt-oss-20b" {
|
|
||||||
t.Errorf("modelID = %q, want %q", modelID, "openai/gpt-oss-20b")
|
|
||||||
}
|
|
||||||
if _, ok := provider.(*HTTPProvider); !ok {
|
|
||||||
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_LMStudio_NoAPIKey(t *testing.T) {
|
|
||||||
cfg := &config.ModelConfig{
|
cfg := &config.ModelConfig{
|
||||||
ModelName: "test-lmstudio",
|
ModelName: tt.modelName,
|
||||||
Model: "lmstudio/openai/gpt-oss-20b",
|
Model: tt.model,
|
||||||
|
}
|
||||||
|
if tt.apiKey != "" {
|
||||||
|
cfg.SetAPIKey(tt.apiKey)
|
||||||
}
|
}
|
||||||
|
|
||||||
provider, modelID, err := CreateProviderFromConfig(cfg)
|
provider, modelID, err := CreateProviderFromConfig(cfg)
|
||||||
|
|
@ -215,97 +249,13 @@ func TestCreateProviderFromConfig_LMStudio_NoAPIKey(t *testing.T) {
|
||||||
if provider == nil {
|
if provider == nil {
|
||||||
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
t.Fatal("CreateProviderFromConfig() returned nil provider")
|
||||||
}
|
}
|
||||||
if modelID != "openai/gpt-oss-20b" {
|
if modelID != tt.wantModelID {
|
||||||
t.Errorf("modelID = %q, want %q", modelID, "openai/gpt-oss-20b")
|
t.Errorf("modelID = %q, want %q", modelID, tt.wantModelID)
|
||||||
}
|
}
|
||||||
if _, ok := provider.(*HTTPProvider); !ok {
|
if _, ok := provider.(*HTTPProvider); !ok {
|
||||||
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
||||||
}
|
}
|
||||||
}
|
})
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_Ollama_WithAPIKey(t *testing.T) {
|
|
||||||
cfg := &config.ModelConfig{
|
|
||||||
ModelName: "test-ollama",
|
|
||||||
Model: "ollama/llama3.1:8b",
|
|
||||||
}
|
|
||||||
cfg.SetAPIKey("test-key")
|
|
||||||
|
|
||||||
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 != "llama3.1:8b" {
|
|
||||||
t.Errorf("modelID = %q, want %q", modelID, "llama3.1:8b")
|
|
||||||
}
|
|
||||||
if _, ok := provider.(*HTTPProvider); !ok {
|
|
||||||
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_Ollama_NoAPIKey(t *testing.T) {
|
|
||||||
cfg := &config.ModelConfig{
|
|
||||||
ModelName: "test-ollama",
|
|
||||||
Model: "ollama/llama3.1:8b",
|
|
||||||
}
|
|
||||||
|
|
||||||
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 != "llama3.1:8b" {
|
|
||||||
t.Errorf("modelID = %q, want %q", modelID, "llama3.1:8b")
|
|
||||||
}
|
|
||||||
if _, ok := provider.(*HTTPProvider); !ok {
|
|
||||||
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_VLLM_WithAPIKey(t *testing.T) {
|
|
||||||
cfg := &config.ModelConfig{
|
|
||||||
ModelName: "test-vllm",
|
|
||||||
Model: "vllm/Qwen/Qwen3-8B",
|
|
||||||
}
|
|
||||||
cfg.SetAPIKey("test-key")
|
|
||||||
|
|
||||||
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-8B" {
|
|
||||||
t.Errorf("modelID = %q, want %q", modelID, "Qwen/Qwen3-8B")
|
|
||||||
}
|
|
||||||
if _, ok := provider.(*HTTPProvider); !ok {
|
|
||||||
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateProviderFromConfig_VLLM_NoAPIKey(t *testing.T) {
|
|
||||||
cfg := &config.ModelConfig{
|
|
||||||
ModelName: "test-vllm",
|
|
||||||
Model: "vllm/Qwen/Qwen3-8B",
|
|
||||||
}
|
|
||||||
|
|
||||||
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-8B" {
|
|
||||||
t.Errorf("modelID = %q, want %q", modelID, "Qwen/Qwen3-8B")
|
|
||||||
}
|
|
||||||
if _, ok := provider.(*HTTPProvider); !ok {
|
|
||||||
t.Fatalf("expected *HTTPProvider, got %T", provider)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue