diff --git a/server/internal/pool/free_models.go b/server/internal/pool/free_models.go index a414cfc..8f81d54 100644 --- a/server/internal/pool/free_models.go +++ b/server/internal/pool/free_models.go @@ -99,7 +99,7 @@ func (r *FreeModelsRefresher) refresh(ctx context.Context) { catalogID := "free/" + m.ID _ = r.q.UpsertModelCatalog(ctx, store.UpsertModelCatalogParams{ ID: catalogID, - Label: catalogID, + Label: prettifyFreeLabel(m.ID), IsChat: 1, Tier: nullString("free"), RawJson: "{}", @@ -110,3 +110,20 @@ func (r *FreeModelsRefresher) refresh(ctx context.Context) { } r.log.Info("free provider: models refreshed", "count", upserted) } + +// prettifyFreeLabel converts slug-style model IDs into readable labels. +// Same logic as nvidia's prettifyModelLabel but without org prefix stripping +// since free model IDs don't have org prefixes. +func prettifyFreeLabel(id string) string { + parts := strings.FieldsFunc(id, func(r rune) bool { + return r == '-' || r == '_' + }) + var words []string + for _, p := range parts { + if p == "" { + continue + } + words = append(words, strings.ToUpper(p[:1])+p[1:]) + } + return strings.Join(words, " ") +} diff --git a/server/internal/pool/providers/nvidia/nvidia.go b/server/internal/pool/providers/nvidia/nvidia.go index de6ba78..af68819 100644 --- a/server/internal/pool/providers/nvidia/nvidia.go +++ b/server/internal/pool/providers/nvidia/nvidia.go @@ -15,6 +15,7 @@ import ( "regexp" "strconv" "strings" + "sync" "time" "github.com/taciturnaxolotl/potluck/internal/pool" @@ -89,10 +90,19 @@ func (NvidiaModelFetcher) FetchModels(ctx context.Context, httpClient *http.Clie return nil, fmt.Errorf("fetch models list: %w", err) } + // Probe each model to check if it's accessible with this API key. + // NVIDIA has no entitlement API — the only way to know is to try. + accessible := probeModels(ctx, httpClient, baseURL, apiKey, models) + now := time.Now().Unix() var out []store.UpsertModelCatalogParams for _, m := range models { + // Skip models that aren't accessible on this account. + if !accessible[m.ID] { + continue + } + // Namespace all NVIDIA model IDs with "nvidia/" prefix so they don't // collide with other providers and the UI can filter by provider. modelID := m.ID @@ -101,7 +111,7 @@ func (NvidiaModelFetcher) FetchModels(ctx context.Context, httpClient *http.Clie } params := store.UpsertModelCatalogParams{ ID: modelID, - Label: m.ID, + Label: prettifyModelLabel(m.ID), Description: "", IsChat: 1, RawJson: "{}", @@ -133,6 +143,57 @@ type nvidiaModel struct { ID string `json:"id"` } +// probeModels checks which models are accessible with the given API key by +// sending a minimal chat completion request to each. Returns a set of +// accessible model IDs. Uses bounded concurrency (10 workers) to avoid +// hammering the API. Models returning 200 or 429 are considered accessible; +// 403/404/402 are not. +func probeModels(ctx context.Context, httpClient *http.Client, baseURL, apiKey string, models []nvidiaModel) map[string]bool { + const maxWorkers = 10 + accessible := make(map[string]bool, len(models)) + var mu sync.Mutex + + sem := make(chan struct{}, maxWorkers) + var wg sync.WaitGroup + + for _, m := range models { + wg.Add(1) + sem <- struct{}{} // acquire worker slot + go func(modelID string) { + defer wg.Done() + defer func() { <-sem }() // release worker slot + + probeCtx, cancel := context.WithTimeout(ctx, 8*time.Second) + defer cancel() + + body := []byte(fmt.Sprintf(`{"model":%q,"messages":[{"role":"user","content":"hi"}],"max_tokens":1}`, modelID)) + req, err := http.NewRequestWithContext(probeCtx, http.MethodPost, baseURL+"/chat/completions", strings.NewReader(string(body))) + if err != nil { + return + } + req.Header.Set("Authorization", "Bearer "+apiKey) + req.Header.Set("Content-Type", "application/json") + + resp, err := httpClient.Do(req) + if err != nil { + return + } + io.Copy(io.Discard, resp.Body) + resp.Body.Close() + + // 200 = accessible, 429 = rate limited but accessible + if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusTooManyRequests { + mu.Lock() + accessible[modelID] = true + mu.Unlock() + } + }(m.ID) + } + + wg.Wait() + return accessible +} + func fetchNvidiaModels(ctx context.Context, httpClient *http.Client, baseURL, apiKey string) ([]nvidiaModel, error) { ctx, cancel := context.WithTimeout(ctx, 20*time.Second) defer cancel() @@ -316,3 +377,30 @@ func cleanDescription(desc string) string { return truncate(text, 200) } + +// prettifyModelLabel converts slug-style model IDs into human-readable labels. +// e.g. "meta/llama-3.1-8b-instruct" → "Llama 3.1 8B Instruct" +// "deepseek-ai/deepseek-v4-flash" → "Deepseek V4 Flash" +// "google/gemma-3n-e4b-it" → "Gemma 3n E4b IT" +func prettifyModelLabel(id string) string { + // Strip org prefix (e.g. "meta/", "deepseek-ai/"). + if idx := strings.IndexByte(id, '/'); idx > 0 { + id = id[idx+1:] + } + + // Split on hyphens and underscores. + parts := strings.FieldsFunc(id, func(r rune) bool { + return r == '-' || r == '_' + }) + + var words []string + for _, p := range parts { + if p == "" { + continue + } + // Capitalize first letter, keep rest as-is (preserves "3.1", "8B", "V4"). + words = append(words, strings.ToUpper(p[:1])+p[1:]) + } + + return strings.Join(words, " ") +} diff --git a/web/src/routes/chat/ModelPicker.svelte b/web/src/routes/chat/ModelPicker.svelte index 368b47e..8c72d29 100644 --- a/web/src/routes/chat/ModelPicker.svelte +++ b/web/src/routes/chat/ModelPicker.svelte @@ -17,25 +17,47 @@ let search = $state(''); let pickerEl = $state(null); - let freeModels = $derived(models.filter((m) => m.id.startsWith('free/'))); - let paidModels = $derived(models.filter((m) => !m.id.startsWith('free/'))); - - let filteredFree = $derived( - search.trim() - ? freeModels.filter((m) => (m.label || m.id).toLowerCase().includes(search.toLowerCase())) - : freeModels - ); - let filteredPaid = $derived( - search.trim() - ? paidModels.filter((m) => (m.label || m.id).toLowerCase().includes(search.toLowerCase())) - : paidModels - ); + function modelProvider(id: string): string { + const idx = id.indexOf('/'); + return idx > 0 ? id.slice(0, idx) : 'pioneer'; + } + + // Group models by provider, preserving a stable order. + let providerGroups = $derived.by(() => { + const map = new Map(); + for (const m of models) { + const p = modelProvider(m.id); + if (!map.has(p)) map.set(p, []); + map.get(p)!.push(m); + } + // Sort groups: pioneer first, then alphabetical. + const entries = [...map.entries()].sort(([a], [b]) => { + if (a === 'pioneer') return -1; + if (b === 'pioneer') return 1; + return a.localeCompare(b); + }); + return entries; + }); + + let filteredGroups = $derived.by(() => { + const q = search.trim().toLowerCase(); + if (!q) return providerGroups; + return providerGroups + .map(([provider, ms]) => [provider, ms.filter((m) => + (m.label || m.id).toLowerCase().includes(q) + )] as [string, Model[]]) + .filter(([, ms]) => ms.length > 0); + }); function display(id: string) { if (!id) return 'pick model'; const m = models.find((x) => x.id === id); - if (m) return m.label || id.replace(/^free\//, ''); - return id.replace(/^free\//, ''); + return m?.label || stripProvider(id); + } + + function stripProvider(id: string): string { + const idx = id.indexOf('/'); + return idx > 0 ? id.slice(idx + 1) : id; } function pick(id: string) { @@ -70,10 +92,8 @@ aria-haspopup="menu" aria-expanded={open} > - {#if selectedModel?.startsWith('free/')} - - {:else if selectedModel} - + {#if selectedModel} + {modelProvider(selectedModel)} {/if} {display(selectedModel)}
free
- {#each filteredFree as m (m.id)} - - {/each} - {/if} - {#if filteredPaid.length > 0} -
pool
- {#each filteredPaid as m (m.id)} + {#each filteredGroups as [provider, groupModels]} +
{provider}
+ {#each groupModels as m (m.id)} + >{m.label || stripProvider(m.id)} {/each} - {/if} + {/each} {#if models.length === 0} loading… - {:else if filteredFree.length === 0 && filteredPaid.length === 0} + {:else if filteredGroups.length === 0} no matches {/if} @@ -154,14 +163,16 @@ cursor: default; } - .tier-dot { - width: 5px; - height: 5px; - border-radius: 50%; - flex-shrink: 0; + .provider-badge { + font-size: 0.6rem; + font-weight: 600; + text-transform: uppercase; + letter-spacing: 0.04em; + padding: 0.1em 0.35em; + border-radius: 3px; + background: var(--bg-sidebar); + color: var(--text-muted); } - .tier-dot.free { background: #4ade80; } - .tier-dot.pool { background: var(--accent); } .chip-label { max-width: 180px; @@ -205,42 +216,40 @@ text-align: left; background: none; border: none; - border-radius: var(--radius-sm); - padding: 0.32rem 0.5rem; + border-radius: 4px; + padding: 0.3rem 0.5rem; font-family: var(--font-mono); font-size: 0.75rem; color: var(--text); cursor: pointer; - transition: background 60ms; white-space: nowrap; overflow: hidden; text-overflow: ellipsis; } - .model-opt:hover { background: var(--bg-page); } + .model-opt:hover { + background: var(--bg-sidebar); + } .model-opt.sel { color: var(--accent); - font-weight: 500; + font-weight: 600; } .model-search-wrap { - padding: 0.3rem 0.3rem 0.2rem; - border-bottom: 1px solid var(--border); - margin-bottom: 0.2rem; + padding: 0.2rem 0.2rem 0.3rem; } - .model-search { - display: block; width: 100%; - box-sizing: border-box; - background: var(--bg-page); - border: 1px solid var(--border); - border-radius: var(--radius-sm); - padding: 0.28rem 0.5rem; - font-family: var(--font-mono); + padding: 0.3rem 0.5rem; font-size: 0.75rem; + font-family: var(--font-mono); + border: 1px solid var(--border); + border-radius: 4px; + background: var(--bg-page); color: var(--text); + box-sizing: border-box; + } + .model-search:focus { outline: none; + border-color: var(--accent); } - .model-search::placeholder { color: var(--text-faint); } - .model-search:focus { border-color: var(--accent); }