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
|
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]
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue