fix(provider): sync lmstudio probing and model normalization
This commit is contained in:
parent
9ee422c12d
commit
3e587cf1d4
5 changed files with 102 additions and 14 deletions
|
|
@ -341,6 +341,20 @@ func isEmptyAPIKeyAllowed(protocol string) bool {
|
|||
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.
|
||||
func getDefaultAPIBase(protocol string) string {
|
||||
meta, ok := protocolMetaByName[protocol]
|
||||
|
|
|
|||
|
|
@ -42,6 +42,23 @@ type Option func(*Provider)
|
|||
|
||||
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 {
|
||||
return func(p *Provider) {
|
||||
p.maxTokensField = maxTokensField
|
||||
|
|
@ -397,13 +414,11 @@ func normalizeModel(model, apiBase string) string {
|
|||
}
|
||||
|
||||
prefix := strings.ToLower(before)
|
||||
switch prefix {
|
||||
case "litellm", "moonshot", "nvidia", "groq", "ollama", "deepseek", "google",
|
||||
"openrouter", "zhipu", "mistral", "vivgrid", "minimax", "novita":
|
||||
if _, ok := stripModelPrefixProviders[prefix]; ok {
|
||||
return after
|
||||
default:
|
||||
return model
|
||||
}
|
||||
|
||||
return model
|
||||
}
|
||||
|
||||
func buildToolsList(tools []ToolDefinition, nativeSearch bool) []any {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
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",
|
||||
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",
|
||||
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" {
|
||||
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" {
|
||||
t.Fatalf("normalizeModel(openrouter) = %q, want %q", got, "openrouter/auto")
|
||||
}
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import (
|
|||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/config"
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
const modelProbeTimeout = 800 * time.Millisecond
|
||||
|
|
@ -60,10 +61,14 @@ func requiresRuntimeProbe(m *config.ModelConfig) bool {
|
|||
return true
|
||||
}
|
||||
|
||||
switch modelProtocol(m.Model) {
|
||||
protocol := modelProtocol(m.Model)
|
||||
|
||||
switch protocol {
|
||||
case "claude-cli", "claudecli", "codex-cli", "codexcli", "github-copilot", "copilot":
|
||||
return true
|
||||
case "ollama", "vllm":
|
||||
}
|
||||
|
||||
if providers.IsEmptyAPIKeyAllowedForProtocol(protocol) {
|
||||
apiBase := strings.TrimSpace(m.APIBase)
|
||||
return apiBase == "" || hasLocalAPIBase(apiBase)
|
||||
}
|
||||
|
|
@ -81,7 +86,7 @@ func probeLocalModelAvailability(m *config.ModelConfig) bool {
|
|||
switch protocol {
|
||||
case "ollama":
|
||||
return probeOllamaModelFunc(apiBase, modelID)
|
||||
case "vllm":
|
||||
case "vllm", "lmstudio":
|
||||
return probeOpenAICompatibleModelFunc(apiBase, modelID, m.APIKey())
|
||||
case "github-copilot", "copilot":
|
||||
return probeTCPServiceFunc(apiBase)
|
||||
|
|
@ -100,11 +105,12 @@ func modelProbeAPIBase(m *config.ModelConfig) string {
|
|||
return normalizeModelProbeAPIBase(apiBase)
|
||||
}
|
||||
|
||||
switch modelProtocol(m.Model) {
|
||||
case "ollama":
|
||||
return "http://localhost:11434/v1"
|
||||
case "vllm":
|
||||
return "http://localhost:8000/v1"
|
||||
protocol := modelProtocol(m.Model)
|
||||
if providers.IsEmptyAPIKeyAllowedForProtocol(protocol) {
|
||||
return providers.DefaultAPIBaseForProtocol(protocol)
|
||||
}
|
||||
|
||||
switch protocol {
|
||||
case "github-copilot", "copilot":
|
||||
return "localhost:4321"
|
||||
default:
|
||||
|
|
|
|||
|
|
@ -35,3 +35,48 @@ func TestProbeLocalModelAvailability_OpenAICompatibleIncludesAPIKey(t *testing.T
|
|||
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")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue