diff --git a/server/cmd/server/main.go b/server/cmd/server/main.go index 2cc15e4..0b6fed6 100644 --- a/server/cmd/server/main.go +++ b/server/cmd/server/main.go @@ -32,7 +32,6 @@ import ( "github.com/taciturnaxolotl/potluck/internal/migrations" "github.com/taciturnaxolotl/potluck/internal/money" "github.com/taciturnaxolotl/potluck/internal/pool" - "github.com/taciturnaxolotl/potluck/internal/provider" "github.com/taciturnaxolotl/potluck/internal/provider/registry" "github.com/taciturnaxolotl/potluck/internal/store" "github.com/taciturnaxolotl/potluck/internal/stream" @@ -97,16 +96,6 @@ func main() { cfg.Spend.MaxConcurrentStreams, ) hub := stream.NewHub(q) - pioneer := provider.New(cfg.Pioneer.BaseURL) - - var freeProvider *provider.Client - if cfg.FreeProvider.Enabled() { - // provider.Client appends /v1/chat/completions itself, so strip any trailing - // /v1 from the configured URL (e.g. http://host:8000/v1 → http://host:8000). - freeBase := strings.TrimSuffix(strings.TrimRight(cfg.FreeProvider.BaseURL, "/"), "/v1") - freeProvider = provider.New(freeBase) - log.Info("free provider configured", "base_url", cfg.FreeProvider.BaseURL) - } // Build multi-provider registry from DB. Falls back to a pioneer-only // registry if the providers table hasn't been migrated yet. @@ -138,12 +127,7 @@ func main() { go reconciler.Run(reconcilerCtx) // Start the models refresher: updates models_catalog hourly. - go pool.NewModelsRefresher(q, keyPool, log.Default()).Run(reconcilerCtx) - - // Start the free models refresher if a free provider is configured. - if cfg.FreeProvider.Enabled() { - go pool.NewFreeModelsRefresher(cfg.FreeProvider.BaseURL, q, log.Default()).Run(reconcilerCtx) - } + go pool.NewModelsRefresher(q, keyPool, reg, log.Default()).Run(reconcilerCtx) // Start the allocation recomputer: every N minutes, refresh per-user // daily allowances using smartAllocate so behavior changes (light user @@ -181,23 +165,19 @@ func main() { r.Post("/auth/logout", hcaLogoutHandler(authSvc, cfg.IsProduction())) apiSrv := &web.Server{ - Q: q, - Auth: authSvc, - Ledger: ledg, - Hub: hub, - Pool: keyPool, - Provider: pioneer, - FreeProvider: freeProvider, - Registry: reg, + Q: q, + Auth: authSvc, + Ledger: ledg, + Hub: hub, + Pool: keyPool, + Registry: reg, } v1Srv := &v1.Server{ - Q: q, - Auth: authSvc, - Ledger: ledg, - Provider: pioneer, - Pool: keyPool, - FreeProvider: freeProvider, - Registry: reg, + Q: q, + Auth: authSvc, + Ledger: ledg, + Pool: keyPool, + Registry: reg, } // /api/* — cookie-authenticated, internal surface. diff --git a/server/internal/api/v1/server.go b/server/internal/api/v1/server.go index a68ef99..2d4dfe6 100644 --- a/server/internal/api/v1/server.go +++ b/server/internal/api/v1/server.go @@ -15,20 +15,17 @@ import ( "github.com/taciturnaxolotl/potluck/internal/auth" "github.com/taciturnaxolotl/potluck/internal/ledger" "github.com/taciturnaxolotl/potluck/internal/pool" - "github.com/taciturnaxolotl/potluck/internal/provider" "github.com/taciturnaxolotl/potluck/internal/provider/registry" "github.com/taciturnaxolotl/potluck/internal/store" ) // Server bundles the deps the v1 handlers need. type Server struct { - Q *store.Queries - Auth *auth.Service - Ledger *ledger.Service - Provider *provider.Client // legacy; being replaced by Registry - Pool *pool.Manager - FreeProvider *provider.Client // nil when free provider is not configured - Registry *registry.Registry // multi-provider registry (new) + Q *store.Queries + Auth *auth.Service + Ledger *ledger.Service + Pool *pool.Manager + Registry *registry.Registry // multi-provider registry } // Mount installs the v1 routes onto r. The caller chains the bearer-auth diff --git a/server/internal/api/web/web.go b/server/internal/api/web/web.go index ed37046..7c9b943 100644 --- a/server/internal/api/web/web.go +++ b/server/internal/api/web/web.go @@ -22,7 +22,6 @@ import ( "github.com/taciturnaxolotl/potluck/internal/auth" "github.com/taciturnaxolotl/potluck/internal/ledger" "github.com/taciturnaxolotl/potluck/internal/pool" - "github.com/taciturnaxolotl/potluck/internal/provider" "github.com/taciturnaxolotl/potluck/internal/provider/registry" "github.com/taciturnaxolotl/potluck/internal/store" "github.com/taciturnaxolotl/potluck/internal/stream" @@ -30,14 +29,12 @@ import ( // Server bundles the deps the web handlers need. type Server struct { - Q *store.Queries - Auth *auth.Service - Ledger *ledger.Service - Hub *stream.Hub - Pool *pool.Manager - Provider *provider.Client // Pioneer upstream (legacy; being replaced by Registry) - FreeProvider *provider.Client // self-hosted free endpoint; nil if not configured - Registry *registry.Registry // multi-provider registry (new) + Q *store.Queries + Auth *auth.Service + Ledger *ledger.Service + Hub *stream.Hub + Pool *pool.Manager + Registry *registry.Registry // multi-provider registry } // Mount registers /api/* routes on r. The caller wraps with cookie-auth diff --git a/server/internal/pool/model_fetcher.go b/server/internal/pool/model_fetcher.go new file mode 100644 index 0000000..dc564f3 --- /dev/null +++ b/server/internal/pool/model_fetcher.go @@ -0,0 +1,296 @@ +package pool + +// model_fetcher.go — provider-specific model catalog fetching. +// Each provider type implements ModelFetcher to pull its available models. + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "io" + "math" + "net/http" + "time" + + "github.com/taciturnaxolotl/potluck/internal/store" +) + +// Shared helper functions for model catalog upserts. +func nullInt64(v int64) sql.NullInt64 { + if v == 0 { + return sql.NullInt64{} + } + return sql.NullInt64{Int64: v, Valid: true} +} + +func nullString(s string) sql.NullString { + if s == "" { + return sql.NullString{} + } + return sql.NullString{String: s, Valid: true} +} + +// Pioneer API response types. +type v1Model struct { + ID string `json:"id"` + DisplayName string `json:"display_name"` + MaxInputTokens int `json:"max_input_tokens"` + MaxTokens int `json:"max_tokens"` + IsChatModel bool `json:"is_chat_model"` + Capabilities struct { + Thinking struct { + Supported bool `json:"supported"` + } `json:"thinking"` + ImageInput struct { + Supported bool `json:"supported"` + } `json:"image_input"` + StructuredOutputs struct { + Supported bool `json:"supported"` + } `json:"structured_outputs"` + } `json:"capabilities"` + Tier string `json:"tier"` + Object string `json:"object"` +} + +type baseModelResp struct { + ID string `json:"id"` + Label string `json:"label"` + Description string `json:"description"` + ContextWindow int64 `json:"context_window"` + InputPricePerMil float64 `json:"input_price_per_million"` + OutputPricePerMil *float64 `json:"output_price_per_million"` + License string `json:"license"` + Tier string `json:"tier"` + IsChatModel bool `json:"is_chat_model"` +} + +// ModelFetcher fetches the model catalog for a specific provider. +type ModelFetcher interface { + // FetchModels returns the list of models available from this provider. + // apiKey may be empty for providers that don't require auth for model listing. + FetchModels(ctx context.Context, httpClient *http.Client, baseURL, apiKey string) ([]store.UpsertModelCatalogParams, error) +} + +// GetModelFetcher returns the appropriate model fetcher for a provider type. +func GetModelFetcher(providerType string) ModelFetcher { + switch providerType { + case "openai_compat": + return PioneerModelFetcher{} + case "free": + return OpenAICompatModelFetcher{} + default: + return nil + } +} + +// PioneerModelFetcher fetches models from pioneer.ai's /base-models + /v1/models. +type PioneerModelFetcher struct{} + +func (PioneerModelFetcher) FetchModels(ctx context.Context, httpClient *http.Client, baseURL, apiKey string) ([]store.UpsertModelCatalogParams, error) { + // Fetch /base-models (no auth required). + baseModels, err := fetchPioneerBaseModels(ctx, httpClient, baseURL) + if err != nil { + // Non-fatal — try /v1/models alone. + baseModels = nil + } + + // Fetch /v1/models (requires auth). + var v1Models []v1Model + if apiKey != "" { + v1Models, err = fetchPioneerV1Models(ctx, httpClient, baseURL, apiKey) + if err != nil && len(baseModels) == 0 { + return nil, fmt.Errorf("both base-models and v1/models failed: %w", err) + } + } + + now := time.Now().Unix() + baseByID := map[string]baseModelResp{} + for _, m := range baseModels { + baseByID[m.ID] = m + } + v1ByID := map[string]v1Model{} + for _, m := range v1Models { + v1ByID[m.ID] = m + } + + var out []store.UpsertModelCatalogParams + + // Primary: /base-models enriched with /v1 data. + for _, bm := range baseModels { + vm := v1ByID[bm.ID] + params := baseModelToParams(bm, vm, now) + out = append(out, params) + } + + // Also include models only in /v1. + for _, vm := range v1Models { + if _, ok := baseByID[vm.ID]; ok { + continue + } + params := v1OnlyModelToParams(vm, now) + out = append(out, params) + } + + return out, nil +} + +// OpenAICompatModelFetcher fetches models from any OpenAI-compatible /v1/models endpoint. +type OpenAICompatModelFetcher struct{} + +func (OpenAICompatModelFetcher) FetchModels(ctx context.Context, httpClient *http.Client, baseURL, apiKey string) ([]store.UpsertModelCatalogParams, error) { + url := baseURL + "/v1/models" + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + if apiKey != "" { + req.Header.Set("Authorization", "Bearer "+apiKey) + } + resp, err := httpClient.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + if resp.StatusCode/100 != 2 { + return nil, fmt.Errorf("/v1/models: HTTP %d", resp.StatusCode) + } + b, _ := io.ReadAll(resp.Body) + var out struct { + Data []struct { + ID string `json:"id"` + } `json:"data"` + } + if err := json.Unmarshal(b, &out); err != nil { + return nil, err + } + + now := time.Now().Unix() + var params []store.UpsertModelCatalogParams + for _, m := range out.Data { + rawJSON, _ := json.Marshal(m) + params = append(params, store.UpsertModelCatalogParams{ + ID: m.ID, + Label: m.ID, + Description: "", + IsChat: 1, + RawJson: string(rawJSON), + RefreshedAt: now, + }) + } + return params, nil +} + +// Helper functions for pioneer model conversion. + +func baseModelToParams(bm baseModelResp, vm v1Model, now int64) store.UpsertModelCatalogParams { + label := bm.Label + desc := bm.Description + tier := bm.Tier + + var inputMicros, outputMicros int64 + if bm.InputPricePerMil > 0 { + inputMicros = int64(math.Round(bm.InputPricePerMil * 1_000_000)) + } + if bm.OutputPricePerMil != nil && *bm.OutputPricePerMil > 0 { + outputMicros = int64(math.Round(*bm.OutputPricePerMil * 1_000_000)) + } + + isChat := int64(0) + if bm.IsChatModel { + isChat = 1 + } + + var ctxWindow, maxOutput int64 + var rawJSON []byte + if vm.ID != "" { + if vm.DisplayName != "" { + label = vm.DisplayName + } + if vm.Tier != "" { + tier = vm.Tier + } + ctxWindow = int64(vm.MaxInputTokens) + maxOutput = int64(vm.MaxTokens) + rawJSON, _ = json.Marshal(vm) + } else { + ctxWindow = bm.ContextWindow + rawJSON, _ = json.Marshal(bm) + } + + return store.UpsertModelCatalogParams{ + ID: bm.ID, + Label: label, + Description: desc, + ContextWindow: nullInt64(ctxWindow), + MaxOutputTokens: nullInt64(maxOutput), + IsChat: isChat, + Tier: nullString(tier), + InputPricePerMillionMicros: nullInt64(inputMicros), + OutputPricePerMillionMicros: nullInt64(outputMicros), + RawJson: string(rawJSON), + RefreshedAt: now, + } +} + +func v1OnlyModelToParams(vm v1Model, now int64) store.UpsertModelCatalogParams { + rawJSON, _ := json.Marshal(vm) + return store.UpsertModelCatalogParams{ + ID: vm.ID, + Label: vm.DisplayName, + ContextWindow: nullInt64(int64(vm.MaxInputTokens)), + MaxOutputTokens: nullInt64(int64(vm.MaxTokens)), + IsChat: 1, + Tier: nullString(vm.Tier), + RawJson: string(rawJSON), + RefreshedAt: now, + } +} + +func fetchPioneerBaseModels(ctx context.Context, httpClient *http.Client, baseURL string) ([]baseModelResp, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL+"/base-models", nil) + if err != nil { + return nil, err + } + resp, err := httpClient.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + if resp.StatusCode/100 != 2 { + return nil, fmt.Errorf("base-models: HTTP %d", resp.StatusCode) + } + b, _ := io.ReadAll(resp.Body) + var out struct { + Models []baseModelResp `json:"models"` + } + if err := json.Unmarshal(b, &out); err != nil { + return nil, err + } + return out.Models, nil +} + +func fetchPioneerV1Models(ctx context.Context, httpClient *http.Client, baseURL, apiKey string) ([]v1Model, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL+"/v1/models", nil) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+apiKey) + resp, err := httpClient.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + if resp.StatusCode/100 != 2 { + b, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("v1/models: HTTP %d: %s", resp.StatusCode, b) + } + b, _ := io.ReadAll(resp.Body) + var out struct { + Data []v1Model `json:"data"` + } + if err := json.Unmarshal(b, &out); err != nil { + return nil, err + } + return out.Data, nil +} diff --git a/server/internal/pool/models_refresher.go b/server/internal/pool/models_refresher.go index bfb46b0..efb9fe8 100644 --- a/server/internal/pool/models_refresher.go +++ b/server/internal/pool/models_refresher.go @@ -1,39 +1,36 @@ package pool -// models_refresher.go — hourly refresh of the pioneer model catalog. -// Fetches /v1/models and /base-models, upserts into models_catalog, -// rotates through available pool keys on failure. +// models_refresher.go — hourly refresh of model catalogs across all providers. +// Dispatches to provider-specific ModelFetcher implementations. import ( "context" - "database/sql" - "encoding/json" - "fmt" - "io" - "math" "net/http" "time" charmlog "charm.land/log/v2" + "github.com/taciturnaxolotl/potluck/internal/provider/registry" "github.com/taciturnaxolotl/potluck/internal/store" ) const ModelsRefreshInterval = time.Hour -// ModelsRefresher fetches the pioneer model catalog on a background ticker. +// ModelsRefresher fetches model catalogs on a background ticker. type ModelsRefresher struct { q *store.Queries manager *Manager + reg *registry.Registry httpClient *http.Client log *charmlog.Logger } // NewModelsRefresher creates a ModelsRefresher. -func NewModelsRefresher(q *store.Queries, m *Manager, log *charmlog.Logger) *ModelsRefresher { +func NewModelsRefresher(q *store.Queries, m *Manager, reg *registry.Registry, log *charmlog.Logger) *ModelsRefresher { return &ModelsRefresher{ q: q, manager: m, + reg: reg, httpClient: &http.Client{Timeout: 20 * time.Second}, log: log, } @@ -42,7 +39,7 @@ func NewModelsRefresher(q *store.Queries, m *Manager, log *charmlog.Logger) *Mod // Run starts the refresh loop. Call in a goroutine. func (r *ModelsRefresher) Run(ctx context.Context) { r.log.Info("models refresher starting", "interval", ModelsRefreshInterval) - r.refresh(ctx) + r.refreshAll(ctx) t := time.NewTicker(ModelsRefreshInterval) defer t.Stop() @@ -51,232 +48,44 @@ func (r *ModelsRefresher) Run(ctx context.Context) { case <-ctx.Done(): return case <-t.C: - r.refresh(ctx) + r.refreshAll(ctx) } } } -type v1Model struct { - ID string `json:"id"` - DisplayName string `json:"display_name"` - MaxInputTokens int `json:"max_input_tokens"` - MaxTokens int `json:"max_tokens"` - IsChatModel bool `json:"is_chat_model"` // not in the response but we infer - Capabilities struct { - Thinking struct { - Supported bool `json:"supported"` - } `json:"thinking"` - ImageInput struct { - Supported bool `json:"supported"` - } `json:"image_input"` - StructuredOutputs struct { - Supported bool `json:"supported"` - } `json:"structured_outputs"` - } `json:"capabilities"` - Tier string `json:"tier"` - Object string `json:"object"` -} - -type baseModelResp struct { - ID string `json:"id"` - Label string `json:"label"` - Description string `json:"description"` - ContextWindow int64 `json:"context_window"` - InputPricePerMil float64 `json:"input_price_per_million"` - OutputPricePerMil *float64 `json:"output_price_per_million"` - License string `json:"license"` - Tier string `json:"tier"` - IsChatModel bool `json:"is_chat_model"` -} - -func (r *ModelsRefresher) refresh(ctx context.Context) { - ctx, cancel := context.WithTimeout(ctx, 30*time.Second) - defer cancel() - - // Pick a key to authenticate. Rotate through keys on error. - apiKey := r.pickKey(ctx) - - // Fetch /base-models (no auth required but nice to have for rate limits). - baseModels, err := r.fetchBaseModels(ctx) - if err != nil { - r.log.Warn("models refresher: base-models fetch failed", "err", err) - // Non-fatal — we can still upsert from /v1/models. - } - baseByID := map[string]baseModelResp{} - for _, m := range baseModels { - baseByID[m.ID] = m - } - - // Fetch /v1/models (requires auth). - if apiKey == "" { - r.log.Warn("models refresher: no healthy key available, skipping /v1/models") - return - } - v1Models, err := r.fetchV1Models(ctx, apiKey) - if err != nil { - r.log.Warn("models refresher: /v1/models fetch failed", "err", err) - return - } - - now := time.Now().Unix() - upserted := 0 - - // Build v1 index for enrichment. - v1ByID := map[string]v1Model{} - for _, m := range v1Models { - v1ByID[m.ID] = m - } - - // Primary source: /base-models (includes models not in /v1, e.g. off-plan). - for _, bm := range baseModels { - vm := v1ByID[bm.ID] - - label := bm.Label - desc := bm.Description - tier := bm.Tier - - var inputMicros, outputMicros int64 - if bm.InputPricePerMil > 0 { - inputMicros = int64(math.Round(bm.InputPricePerMil * 1_000_000)) - } - if bm.OutputPricePerMil != nil && *bm.OutputPricePerMil > 0 { - outputMicros = int64(math.Round(*bm.OutputPricePerMil * 1_000_000)) - } - - isChat := int64(0) - if bm.IsChatModel { - isChat = 1 +// refreshAll iterates over all active providers and refreshes their model catalogs. +func (r *ModelsRefresher) refreshAll(ctx context.Context) { + providers := r.reg.List() + for _, p := range providers { + fetcher := GetModelFetcher(string(p.Type)) + if fetcher == nil { + continue } - // Enrich with /v1 data where available. - var ctxWindow, maxOutput int64 - var rawJSON []byte - if vm.ID != "" { - if vm.DisplayName != "" { - label = vm.DisplayName - } - if vm.Tier != "" { - tier = vm.Tier - } - ctxWindow = int64(vm.MaxInputTokens) - maxOutput = int64(vm.MaxTokens) - rawJSON, _ = json.Marshal(vm) - } else { - ctxWindow = bm.ContextWindow - rawJSON, _ = json.Marshal(bm) + // Pick a key for this provider (needed for auth on some endpoints). + var apiKey string + sel, err := r.manager.PickForProvider(ctx, p.ID) + if err == nil && sel != nil { + apiKey = sel.APIKey() } - _ = r.q.UpsertModelCatalog(ctx, store.UpsertModelCatalogParams{ - ID: bm.ID, - Label: label, - Description: desc, - ContextWindow: nullInt64(ctxWindow), - MaxOutputTokens: nullInt64(maxOutput), - IsChat: isChat, - Tier: nullString(tier), - InputPricePerMillionMicros: nullInt64(inputMicros), - OutputPricePerMillionMicros: nullInt64(outputMicros), - RawJson: string(rawJSON), - RefreshedAt: now, - }) - upserted++ - } + ctx, cancel := context.WithTimeout(ctx, 30*time.Second) + models, err := fetcher.FetchModels(ctx, r.httpClient, p.BaseURL, apiKey) + cancel() - // Also upsert any models that only exist in /v1 (unlikely but don't lose them). - for _, vm := range v1Models { - if _, ok := baseByID[vm.ID]; ok { - continue // already handled above + if err != nil { + r.log.Warn("models refresher: fetch failed", + "provider", p.ID, "err", err) + continue } - label := vm.DisplayName - var inputMicros, outputMicros int64 - isChat := int64(1) - rawJSON, _ := json.Marshal(vm) - - _ = r.q.UpsertModelCatalog(ctx, store.UpsertModelCatalogParams{ - ID: vm.ID, - Label: label, - ContextWindow: nullInt64(int64(vm.MaxInputTokens)), - MaxOutputTokens: nullInt64(int64(vm.MaxTokens)), - IsChat: isChat, - Tier: nullString(vm.Tier), - InputPricePerMillionMicros: nullInt64(inputMicros), - OutputPricePerMillionMicros: nullInt64(outputMicros), - RawJson: string(rawJSON), - RefreshedAt: now, - }) - upserted++ - } - r.log.Info("models refresher: catalog updated", "base_models", len(baseModels), "v1_models", len(v1Models), "upserted", upserted) -} - -// pickKey returns a decrypted pioneer API key from the pool, or "" if none available. -func (r *ModelsRefresher) pickKey(ctx context.Context) string { - sel, err := r.manager.Pick(ctx) - if err != nil { - return "" - } - return sel.APIKey() -} - -func (r *ModelsRefresher) fetchBaseModels(ctx context.Context) ([]baseModelResp, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://api.pioneer.ai/base-models", nil) - if err != nil { - return nil, err - } - resp, err := r.httpClient.Do(req) - if err != nil { - return nil, err - } - defer resp.Body.Close() - if resp.StatusCode/100 != 2 { - return nil, fmt.Errorf("base-models: HTTP %d", resp.StatusCode) - } - b, _ := io.ReadAll(resp.Body) - var out struct { - Models []baseModelResp `json:"models"` - } - if err := json.Unmarshal(b, &out); err != nil { - return nil, err - } - return out.Models, nil -} - -func (r *ModelsRefresher) fetchV1Models(ctx context.Context, apiKey string) ([]v1Model, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://api.pioneer.ai/v1/models", nil) - if err != nil { - return nil, err - } - req.Header.Set("Authorization", "Bearer "+apiKey) - resp, err := r.httpClient.Do(req) - if err != nil { - return nil, err - } - defer resp.Body.Close() - if resp.StatusCode/100 != 2 { - b, _ := io.ReadAll(resp.Body) - return nil, fmt.Errorf("v1/models: HTTP %d: %s", resp.StatusCode, b) - } - b, _ := io.ReadAll(resp.Body) - var out struct { - Data []v1Model `json:"data"` - } - if err := json.Unmarshal(b, &out); err != nil { - return nil, err - } - return out.Data, nil -} - -func nullInt64(v int64) sql.NullInt64 { - if v == 0 { - return sql.NullInt64{} - } - return sql.NullInt64{Int64: v, Valid: true} -} - -func nullString(s string) sql.NullString { - if s == "" { - return sql.NullString{} + upserted := 0 + for _, params := range models { + if err := r.q.UpsertModelCatalog(ctx, params); err == nil { + upserted++ + } + } + r.log.Info("models refresher: catalog updated", + "provider", p.ID, "models", len(models), "upserted", upserted) } - return sql.NullString{String: s, Valid: true} } diff --git a/server/internal/pool/pool.go b/server/internal/pool/pool.go index 4880d62..920a694 100644 --- a/server/internal/pool/pool.go +++ b/server/internal/pool/pool.go @@ -120,6 +120,23 @@ func (m *Manager) PickForUser(ctx context.Context, userID string) (*Selection, e return &Selection{keyID: key.ID, plaintext: plaintext}, nil } +// PickForProvider picks any healthy key for a specific provider. +// Used by the models refresher to authenticate model listing requests. +func (m *Manager) PickForProvider(ctx context.Context, providerID string) (*Selection, error) { + key, err := m.q.PickPoolKeyForProvider(ctx, providerID) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, ErrNoKeys + } + return nil, fmt.Errorf("pool: pick key for provider %s: %w", providerID, err) + } + plaintext, err := m.Decrypt(key.KeyCiphertext) + if err != nil { + return nil, fmt.Errorf("pool: decrypt key %s: %w", key.ID, err) + } + return &Selection{keyID: key.ID, plaintext: plaintext}, nil +} + // RecordSpend records spend for the key used in a request. // Call this after the upstream call settles (use context.Background() — the // request context may already be canceled). diff --git a/server/internal/pool/provider_strategy.go b/server/internal/pool/provider_strategy.go new file mode 100644 index 0000000..7c3dc10 --- /dev/null +++ b/server/internal/pool/provider_strategy.go @@ -0,0 +1,142 @@ +package pool + +import ( + "context" + "database/sql" + "net/http" + + "github.com/taciturnaxolotl/potluck/internal/store" +) + +// HealthChecker probes a pool key's upstream provider to determine if it's +// still valid and what its current usage/limits are. Each provider type +// implements this differently (pioneer uses /billing/plan-info; others may +// use /v1/models or have no health check at all). +type HealthChecker interface { + // CheckHealth probes a single key. Returns nil result for transient + // errors (caller should retry next tick). Non-nil result with an error + // means permanent failure. + CheckHealth(ctx context.Context, httpClient *http.Client, apiKey string) (*HealthResult, error) + + // AcceptedPlan reports whether a payment plan is allowed in the pool. + // Return true for all plans if the provider doesn't have plan tiers. + AcceptedPlan(plan string) bool +} + +// HealthResult holds the outcome of a health check probe. +type HealthResult struct { + // HTTP status flags — checked before other fields. + HTTP401 bool // unauthorized / exhausted + HTTP503 bool // transient service unavailable + + // Provider-specific metadata. Fields that don't apply to a provider + // should be left zero/empty. + PaymentPlan string + CreditLimitMicros int64 + RemainingMicros int64 + TodayMicros int64 + TeamID string +} + +// BillingIngestor pulls billing data from an upstream provider and writes +// it to pool_key_billing_rows with user attribution. Providers without a +// billing API can return nil (no-op) and rely on token-count estimation. +type BillingIngestor interface { + // IngestBilling fetches new billing rows for a key and writes them to + // the database. Called after a successful health check. + IngestBilling(ctx context.Context, q *store.Queries, key store.PoolKey, apiKey string, httpClient *http.Client) error +} + +// NoopBillingIngestor is a BillingIngestor that does nothing. Used for +// providers without a billing API. +type NoopBillingIngestor struct{} + +func (NoopBillingIngestor) IngestBilling(context.Context, *store.Queries, store.PoolKey, string, *http.Client) error { + return nil +} + +// NoopHealthChecker always reports healthy. Used for providers without +// a health check endpoint (e.g., self-hosted free models). +type NoopHealthChecker struct{} + +func (NoopHealthChecker) CheckHealth(context.Context, *http.Client, string) (*HealthResult, error) { + return &HealthResult{}, nil +} + +func (NoopHealthChecker) AcceptedPlan(string) bool { return true } + +// PioneerHealthChecker implements HealthChecker for pioneer.ai keys. +type PioneerHealthChecker struct{} + +func (PioneerHealthChecker) AcceptedPlan(plan string) bool { + switch plan { + case "pro", "pro_legacy", "partner": + return true + default: + return false + } +} + +func (PioneerHealthChecker) CheckHealth(ctx context.Context, httpClient *http.Client, apiKey string) (*HealthResult, error) { + result := probePlanInfo(ctx, httpClient, apiKey) + if result.err != nil { + return nil, result.err + } + return &HealthResult{ + HTTP401: result.http401, + HTTP503: result.http503, + PaymentPlan: result.plan.PaymentPlan, + CreditLimitMicros: result.creditLimitMicros, + RemainingMicros: result.remainingMicros, + TodayMicros: result.todayMicros, + TeamID: result.teamID, + }, nil +} + +// PioneerBillingIngestor implements BillingIngestor for pioneer.ai keys. +type PioneerBillingIngestor struct{} + +func (PioneerBillingIngestor) IngestBilling(ctx context.Context, q *store.Queries, key store.PoolKey, apiKey string, httpClient *http.Client) error { + // Pioneer billing ingestion is currently tightly coupled to the Reconciler + // (it uses r.q and r.httpClient). This will be fully decoupled when we + // refactor billing_ingest.go. For now, return nil and let the reconciler + // call ingestBillingRows directly for pioneer keys. + return nil +} + +// GetHealthChecker returns the appropriate health checker for a provider type. +func GetHealthChecker(providerType string) HealthChecker { + switch providerType { + case "openai_compat": + // Currently only pioneer uses openai_compat; when we add more, + // this will need per-provider dispatch (e.g., by provider ID). + return PioneerHealthChecker{} + default: + return NoopHealthChecker{} + } +} + +// GetBillingIngestor returns the appropriate billing ingestor for a provider type. +func GetBillingIngestor(providerType string) BillingIngestor { + switch providerType { + case "openai_compat": + return PioneerBillingIngestor{} + default: + return NoopBillingIngestor{} + } +} + +// nullStr and nullInt helpers for DB params. +func nullStr(s string) sql.NullString { + if s == "" { + return sql.NullString{} + } + return sql.NullString{String: s, Valid: true} +} + +func nullInt(v int64) sql.NullInt64 { + if v == 0 { + return sql.NullInt64{} + } + return sql.NullInt64{Int64: v, Valid: true} +} diff --git a/server/internal/pool/reconciler.go b/server/internal/pool/reconciler.go index 6e530b3..af5bb47 100644 --- a/server/internal/pool/reconciler.go +++ b/server/internal/pool/reconciler.go @@ -229,17 +229,25 @@ func (r *Reconciler) probeKey(ctx context.Context, key store.PoolKey) { return } - result := probePlanInfo(ctx, r.httpClient, plaintext) + checker := GetHealthChecker(key.ProviderID) + healthResult, err := checker.CheckHealth(ctx, r.httpClient, plaintext) now := time.Now().Unix() + if err != nil { + // Network/decode error — transient, leave health unchanged, retry next tick. + r.log.Warn("reconciler: probe transient error", + "key_id", key.ID, "label", key.Label, "err", err) + return + } + switch { - case result.http503: - // Pioneer auth service down — transient, don't touch health. - r.log.Warn("reconciler: pioneer auth down (503), skipping key", + case healthResult.HTTP503: + // Provider auth service down — transient, don't touch health. + r.log.Warn("reconciler: provider auth down (503), skipping key", "key_id", key.ID, "label", key.Label) return - case result.http401: + case healthResult.HTTP401: // Key exhausted or revoked — mark unauthorized. var unhealthySince sql.NullInt64 if key.PioneerHealth == HealthUnauthorized && key.PioneerUnhealthySince.Valid { @@ -263,27 +271,21 @@ func (r *Reconciler) probeKey(ctx context.Context, key store.PoolKey) { ID: key.ID, }) return - - case result.err != nil: - // Network/decode error — transient, leave health unchanged, retry next tick. - r.log.Warn("reconciler: probe transient error", - "key_id", key.ID, "label", key.Label, "err", result.err) - return } // Successful probe — validate plan and update. - if !AcceptedPlan(result.plan.PaymentPlan) { + if !checker.AcceptedPlan(healthResult.PaymentPlan) { r.log.Warn("reconciler: unsupported plan, marking unauthorized", "key_id", key.ID, "label", key.Label, - "plan", result.plan.PaymentPlan) + "plan", healthResult.PaymentPlan) _ = r.q.UpdatePoolKeyHealth(ctx, store.UpdatePoolKeyHealthParams{ PioneerHealth: HealthUnauthorized, PioneerUnhealthySince: sql.NullInt64{Int64: now, Valid: true}, - PioneerTeamID: nullStr(result.teamID), - PioneerPaymentPlan: nullStr(result.plan.PaymentPlan), - PioneerCreditLimitMicros: nullInt(result.creditLimitMicros), - PioneerRemainingMicros: nullInt(result.remainingMicros), - TodayMicros: result.todayMicros, + PioneerTeamID: nullStr(healthResult.TeamID), + PioneerPaymentPlan: nullStr(healthResult.PaymentPlan), + PioneerCreditLimitMicros: nullInt(healthResult.CreditLimitMicros), + PioneerRemainingMicros: nullInt(healthResult.RemainingMicros), + TodayMicros: healthResult.TodayMicros, LastBillingSyncAt: sql.NullInt64{Int64: now, Valid: true}, ID: key.ID, }) @@ -300,37 +302,37 @@ func (r *Reconciler) probeKey(ctx context.Context, key store.PoolKey) { _ = r.q.UpdatePoolKeyHealth(ctx, store.UpdatePoolKeyHealthParams{ PioneerHealth: HealthHealthy, PioneerUnhealthySince: sql.NullInt64{}, // NULL — clear it - PioneerTeamID: nullStr(result.teamID), - PioneerPaymentPlan: nullStr(result.plan.PaymentPlan), - PioneerCreditLimitMicros: nullInt(result.creditLimitMicros), - PioneerRemainingMicros: nullInt(result.remainingMicros), - TodayMicros: result.todayMicros, + PioneerTeamID: nullStr(healthResult.TeamID), + PioneerPaymentPlan: nullStr(healthResult.PaymentPlan), + PioneerCreditLimitMicros: nullInt(healthResult.CreditLimitMicros), + PioneerRemainingMicros: nullInt(healthResult.RemainingMicros), + TodayMicros: healthResult.TodayMicros, LastBillingSyncAt: sql.NullInt64{Int64: now, Valid: true}, ID: key.ID, }) - // Keep max_micros in sync with pioneer's credit limit so the pool always + // Keep max_micros in sync with provider's credit limit so the pool always // knows the real ceiling. shared_micros is clamped to the new max if it // would otherwise exceed it. - if result.creditLimitMicros > 0 && result.creditLimitMicros != key.PioneerCreditLimitMicros.Int64 { + if healthResult.CreditLimitMicros > 0 && healthResult.CreditLimitMicros != key.PioneerCreditLimitMicros.Int64 { _ = r.q.SyncPoolKeyMaxFromCreditLimit(ctx, store.SyncPoolKeyMaxFromCreditLimitParams{ - MaxMicros: result.creditLimitMicros, - SharedMicros: result.creditLimitMicros, - SharedMicros_2: result.creditLimitMicros, + MaxMicros: healthResult.CreditLimitMicros, + SharedMicros: healthResult.CreditLimitMicros, + SharedMicros_2: healthResult.CreditLimitMicros, ID: key.ID, }) r.log.Info("reconciler: credit limit changed, updated max_micros", "key_id", key.ID, "label", key.Label, "old_micros", key.PioneerCreditLimitMicros.Int64, - "new_micros", result.creditLimitMicros, + "new_micros", healthResult.CreditLimitMicros, ) } r.log.Info("reconciler: key synced", "key_id", key.ID, "label", key.Label, - "today_usd", fmt.Sprintf("%.2f", float64(result.todayMicros)/1_000_000), - "remaining_usd", fmt.Sprintf("%.2f", float64(result.remainingMicros)/1_000_000), + "today_usd", fmt.Sprintf("%.2f", float64(healthResult.TodayMicros)/1_000_000), + "remaining_usd", fmt.Sprintf("%.2f", float64(healthResult.RemainingMicros)/1_000_000), ) // Ingest new billing rows and attribute to users. @@ -345,17 +347,3 @@ func (r *Reconciler) probeKey(ctx context.Context, key store.PoolKey) { }) } } - -func nullStr(s string) sql.NullString { - if s == "" { - return sql.NullString{} - } - return sql.NullString{String: s, Valid: true} -} - -func nullInt(v int64) sql.NullInt64 { - if v == 0 { - return sql.NullInt64{} - } - return sql.NullInt64{Int64: v, Valid: true} -}