fix(provider): sync lmstudio probing and model normalization

This commit is contained in:
lc6464 2026-03-30 22:43:28 +08:00
parent 9ee422c12d
commit 3e587cf1d4
No known key found for this signature in database
GPG key ID: 53C61B42FEC71D6D
5 changed files with 102 additions and 14 deletions

View file

@ -341,6 +341,20 @@ func isEmptyAPIKeyAllowed(protocol string) bool {
return ok && meta.emptyAPIKeyAllowed return ok && meta.emptyAPIKeyAllowed
} }
// IsEmptyAPIKeyAllowedForProtocol reports whether a protocol allows requests
// without api_key when using its default local endpoint.
func IsEmptyAPIKeyAllowedForProtocol(protocol string) bool {
protocol = strings.ToLower(strings.TrimSpace(protocol))
return isEmptyAPIKeyAllowed(protocol)
}
// DefaultAPIBaseForProtocol returns the configured default API base for a protocol.
// It returns empty string if the protocol has no default base.
func DefaultAPIBaseForProtocol(protocol string) string {
protocol = strings.ToLower(strings.TrimSpace(protocol))
return getDefaultAPIBase(protocol)
}
// 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 {
meta, ok := protocolMetaByName[protocol] meta, ok := protocolMetaByName[protocol]

View file

@ -42,6 +42,23 @@ type Option func(*Provider)
const defaultRequestTimeout = common.DefaultRequestTimeout const defaultRequestTimeout = common.DefaultRequestTimeout
var stripModelPrefixProviders = map[string]struct{}{
"litellm": {},
"moonshot": {},
"nvidia": {},
"groq": {},
"ollama": {},
"deepseek": {},
"google": {},
"openrouter": {},
"zhipu": {},
"mistral": {},
"vivgrid": {},
"minimax": {},
"novita": {},
"lmstudio": {},
}
func WithMaxTokensField(maxTokensField string) Option { func WithMaxTokensField(maxTokensField string) Option {
return func(p *Provider) { return func(p *Provider) {
p.maxTokensField = maxTokensField p.maxTokensField = maxTokensField
@ -397,13 +414,11 @@ func normalizeModel(model, apiBase string) string {
} }
prefix := strings.ToLower(before) prefix := strings.ToLower(before)
switch prefix { if _, ok := stripModelPrefixProviders[prefix]; ok {
case "litellm", "moonshot", "nvidia", "groq", "ollama", "deepseek", "google",
"openrouter", "zhipu", "mistral", "vivgrid", "minimax", "novita":
return after return after
default:
return model
} }
return model
} }
func buildToolsList(tools []ToolDefinition, nativeSearch bool) []any { func buildToolsList(tools []ToolDefinition, nativeSearch bool) []any {

View file

@ -432,7 +432,7 @@ func TestProviderChat_StripsMoonshotPrefixAndNormalizesKimiTemperature(t *testin
} }
} }
func TestProviderChat_StripsGroqOllamaDeepseekVivgridNovitaPrefixes(t *testing.T) { func TestProviderChat_StripsKnownProviderPrefixes(t *testing.T) {
var requestBody map[string]any var requestBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@ -474,6 +474,11 @@ func TestProviderChat_StripsGroqOllamaDeepseekVivgridNovitaPrefixes(t *testing.T
input: "ollama/qwen2.5:14b", input: "ollama/qwen2.5:14b",
wantModel: "qwen2.5:14b", wantModel: "qwen2.5:14b",
}, },
{
name: "strips lmstudio prefix and keeps nested model",
input: "lmstudio/openai/gpt-oss-20b",
wantModel: "openai/gpt-oss-20b",
},
{ {
name: "strips deepseek prefix", name: "strips deepseek prefix",
input: "deepseek/deepseek-chat", input: "deepseek/deepseek-chat",
@ -579,6 +584,9 @@ func TestNormalizeModel_UsesAPIBase(t *testing.T) {
if got := normalizeModel("deepseek/deepseek-chat", "https://api.deepseek.com/v1"); got != "deepseek-chat" { if got := normalizeModel("deepseek/deepseek-chat", "https://api.deepseek.com/v1"); got != "deepseek-chat" {
t.Fatalf("normalizeModel(deepseek) = %q, want %q", got, "deepseek-chat") t.Fatalf("normalizeModel(deepseek) = %q, want %q", got, "deepseek-chat")
} }
if got := normalizeModel("lmstudio/openai/gpt-oss-20b", "http://localhost:1234/v1"); got != "openai/gpt-oss-20b" {
t.Fatalf("normalizeModel(lmstudio) = %q, want %q", got, "openai/gpt-oss-20b")
}
if got := normalizeModel("openrouter/auto", "https://openrouter.ai/api/v1"); got != "openrouter/auto" { if got := normalizeModel("openrouter/auto", "https://openrouter.ai/api/v1"); got != "openrouter/auto" {
t.Fatalf("normalizeModel(openrouter) = %q, want %q", got, "openrouter/auto") t.Fatalf("normalizeModel(openrouter) = %q, want %q", got, "openrouter/auto")
} }

View file

@ -10,6 +10,7 @@ import (
"time" "time"
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/providers"
) )
const modelProbeTimeout = 800 * time.Millisecond const modelProbeTimeout = 800 * time.Millisecond
@ -60,10 +61,14 @@ func requiresRuntimeProbe(m *config.ModelConfig) bool {
return true return true
} }
switch modelProtocol(m.Model) { protocol := modelProtocol(m.Model)
switch protocol {
case "claude-cli", "claudecli", "codex-cli", "codexcli", "github-copilot", "copilot": case "claude-cli", "claudecli", "codex-cli", "codexcli", "github-copilot", "copilot":
return true return true
case "ollama", "vllm": }
if providers.IsEmptyAPIKeyAllowedForProtocol(protocol) {
apiBase := strings.TrimSpace(m.APIBase) apiBase := strings.TrimSpace(m.APIBase)
return apiBase == "" || hasLocalAPIBase(apiBase) return apiBase == "" || hasLocalAPIBase(apiBase)
} }
@ -81,7 +86,7 @@ func probeLocalModelAvailability(m *config.ModelConfig) bool {
switch protocol { switch protocol {
case "ollama": case "ollama":
return probeOllamaModelFunc(apiBase, modelID) return probeOllamaModelFunc(apiBase, modelID)
case "vllm": case "vllm", "lmstudio":
return probeOpenAICompatibleModelFunc(apiBase, modelID, m.APIKey()) return probeOpenAICompatibleModelFunc(apiBase, modelID, m.APIKey())
case "github-copilot", "copilot": case "github-copilot", "copilot":
return probeTCPServiceFunc(apiBase) return probeTCPServiceFunc(apiBase)
@ -100,11 +105,12 @@ func modelProbeAPIBase(m *config.ModelConfig) string {
return normalizeModelProbeAPIBase(apiBase) return normalizeModelProbeAPIBase(apiBase)
} }
switch modelProtocol(m.Model) { protocol := modelProtocol(m.Model)
case "ollama": if providers.IsEmptyAPIKeyAllowedForProtocol(protocol) {
return "http://localhost:11434/v1" return providers.DefaultAPIBaseForProtocol(protocol)
case "vllm": }
return "http://localhost:8000/v1"
switch protocol {
case "github-copilot", "copilot": case "github-copilot", "copilot":
return "localhost:4321" return "localhost:4321"
default: default:

View file

@ -35,3 +35,48 @@ func TestProbeLocalModelAvailability_OpenAICompatibleIncludesAPIKey(t *testing.T
t.Fatal("probeLocalModelAvailability() = false, want true when api_key is configured") t.Fatal("probeLocalModelAvailability() = false, want true when api_key is configured")
} }
} }
func TestRequiresRuntimeProbe_LMStudio(t *testing.T) {
if !requiresRuntimeProbe(&config.ModelConfig{Model: "lmstudio/openai/gpt-oss-20b"}) {
t.Fatal("requiresRuntimeProbe(lmstudio with default base) = false, want true")
}
if requiresRuntimeProbe(&config.ModelConfig{Model: "lmstudio/openai/gpt-oss-20b", APIBase: "https://api.example.com/v1"}) {
t.Fatal("requiresRuntimeProbe(lmstudio with remote base) = true, want false")
}
}
func TestModelProbeAPIBase_LMStudioDefault(t *testing.T) {
got := modelProbeAPIBase(&config.ModelConfig{Model: "lmstudio/openai/gpt-oss-20b"})
if got != "http://localhost:1234/v1" {
t.Fatalf("modelProbeAPIBase(lmstudio) = %q, want %q", got, "http://localhost:1234/v1")
}
}
func TestProbeLocalModelAvailability_LMStudioUsesOpenAICompatibleProbe(t *testing.T) {
originalProbe := probeOpenAICompatibleModelFunc
defer func() { probeOpenAICompatibleModelFunc = originalProbe }()
called := false
probeOpenAICompatibleModelFunc = func(apiBase, modelID, apiKey string) bool {
called = true
if apiBase != "http://localhost:1234/v1" {
t.Fatalf("apiBase = %q, want %q", apiBase, "http://localhost:1234/v1")
}
if modelID != "openai/gpt-oss-20b" {
t.Fatalf("modelID = %q, want %q", modelID, "openai/gpt-oss-20b")
}
if apiKey != "" {
t.Fatalf("apiKey = %q, want empty", apiKey)
}
return true
}
model := &config.ModelConfig{Model: "lmstudio/openai/gpt-oss-20b"}
if !probeLocalModelAvailability(model) {
t.Fatal("probeLocalModelAvailability(lmstudio) = false, want true")
}
if !called {
t.Fatal("probeOpenAICompatibleModelFunc was not called for lmstudio")
}
}