diff --git a/framework/modelcatalog/models.go b/framework/modelcatalog/models.go index 547679a80b6..b28454669c6 100644 --- a/framework/modelcatalog/models.go +++ b/framework/modelcatalog/models.go @@ -130,7 +130,7 @@ func (mc *ModelCatalog) GetProvidersForModel(model string) []schemas.ModelProvid } } - // Cross-provider special cases (preserved from pre-refactor models.go). + // Cross-provider special cases if _, ok := seen[schemas.OpenRouter]; !ok { openRouterModels := mc.GetModelsForProvider(schemas.OpenRouter) for _, p := range providers { diff --git a/framework/modelcatalog/pool.go b/framework/modelcatalog/pool.go index f4e316f27f4..79caf277fe5 100644 --- a/framework/modelcatalog/pool.go +++ b/framework/modelcatalog/pool.go @@ -13,6 +13,18 @@ func (mc *ModelCatalog) UpsertLive(provider schemas.ModelProvider, keyID string, mc.live.Upsert(provider, keyID, unfiltered, models) } +// UpsertLiveFromResponse extracts model IDs from a BifrostListModelsResponse +// (parsing "provider/model" prefixes, filtering by provider match, +// deduplicating) and pushes them into the live cache. A nil resp is a no-op +// so callers can't accidentally clear an existing cache entry by handing in +// a missing response. +func (mc *ModelCatalog) UpsertLiveFromResponse(provider schemas.ModelProvider, keyID string, unfiltered bool, resp *schemas.BifrostListModelsResponse) { + if resp == nil { + return + } + mc.live.Upsert(provider, keyID, unfiltered, extractModelIDs(resp, provider)) +} + // InvalidateLive drops both filtered + unfiltered live entries for one key. func (mc *ModelCatalog) InvalidateLive(provider schemas.ModelProvider, keyID string) { mc.live.Invalidate(provider, keyID) diff --git a/framework/modelcatalog/pool_test.go b/framework/modelcatalog/pool_test.go new file mode 100644 index 00000000000..b9dc4c5f0ff --- /dev/null +++ b/framework/modelcatalog/pool_test.go @@ -0,0 +1,169 @@ +package modelcatalog + +import ( + "slices" + "testing" + + "github.com/maximhq/bifrost/core/schemas" +) + +// TestUpsertLiveFromResponse_NilRespIsNoop guards the API surface: handing a +// nil response into UpsertLiveFromResponse must not clear an existing cache +// entry by storing an empty slice. Without the early-return, extractModelIDs +// would return nil and the live store would publish an empty model list, +// silently removing the provider's previously-fetched models from routing. +func TestUpsertLiveFromResponse_NilRespIsNoop(t *testing.T) { + mc := NewTestCatalog(nil) + mc.UpsertLive(schemas.OpenAI, "k1", false, []string{"gpt-4o", "o1"}) + + mc.UpsertLiveFromResponse(schemas.OpenAI, "k1", false, nil) + + got := mc.GetModelsForProvider(schemas.OpenAI) + slices.Sort(got) + want := []string{"gpt-4o", "o1"} + if !slices.Equal(got, want) { + t.Errorf("after UpsertLiveFromResponse(nil), GetModelsForProvider = %v, want %v (entry must survive nil resp)", got, want) + } +} + +// TestUpsertLiveFromResponse_PopulatesFromResponse covers the happy path: +// extractModelIDs strips the owning provider prefix and the resulting bare +// names land in the live cache for GetModelsForProvider. +func TestUpsertLiveFromResponse_PopulatesFromResponse(t *testing.T) { + mc := NewTestCatalog(nil) + resp := &schemas.BifrostListModelsResponse{ + Data: []schemas.Model{ + {ID: "openai/gpt-4o"}, + {ID: "openai/o1"}, + }, + } + mc.UpsertLiveFromResponse(schemas.OpenAI, "k1", false, resp) + + got := mc.GetModelsForProvider(schemas.OpenAI) + slices.Sort(got) + want := []string{"gpt-4o", "o1"} + if !slices.Equal(got, want) { + t.Errorf("GetModelsForProvider = %v, want %v", got, want) + } +} + +// TestExtractModelIDs_StripsOwningProviderPrefix verifies the canonical +// shape returned by every provider's ListModels — an ID prefixed with its +// own provider key — gets reduced to a bare model name. +func TestExtractModelIDs_StripsOwningProviderPrefix(t *testing.T) { + resp := &schemas.BifrostListModelsResponse{ + Data: []schemas.Model{ + {ID: "openai/gpt-4o"}, + {ID: "openai/o1"}, + }, + } + got := extractModelIDs(resp, schemas.OpenAI) + slices.Sort(got) + want := []string{"gpt-4o", "o1"} + if !slices.Equal(got, want) { + t.Errorf("extractModelIDs = %v, want %v", got, want) + } +} + +// TestExtractModelIDs_KeepsNestedProviderForGateway covers the +// gateway-provider shape (OpenRouter returns IDs like "openrouter/openai/gpt-4") +// — ParseModelString splits on the first slash, so the parsed prefix matches +// the owning provider and the remainder ("openai/gpt-4") is kept as-is. +func TestExtractModelIDs_KeepsNestedProviderForGateway(t *testing.T) { + resp := &schemas.BifrostListModelsResponse{ + Data: []schemas.Model{ + {ID: "openrouter/openai/gpt-4"}, + {ID: "openrouter/anthropic/claude-sonnet-4"}, + }, + } + got := extractModelIDs(resp, schemas.OpenRouter) + slices.Sort(got) + want := []string{"anthropic/claude-sonnet-4", "openai/gpt-4"} + if !slices.Equal(got, want) { + t.Errorf("extractModelIDs = %v, want %v", got, want) + } +} + +// TestExtractModelIDs_DropsForeignPrefix asserts the defensive filter: an +// ID prefixed with a different provider than the one being upserted is +// excluded. This shouldn't fire in practice (providers self-prefix their +// own list-models output before it reaches here), but the guard exists for +// malformed inputs and the test pins the behavior so refactors don't +// silently invert it. +func TestExtractModelIDs_DropsForeignPrefix(t *testing.T) { + resp := &schemas.BifrostListModelsResponse{ + Data: []schemas.Model{ + {ID: "openai/gpt-4o"}, + {ID: "anthropic/claude-sonnet"}, // foreign — should be dropped + }, + } + got := extractModelIDs(resp, schemas.OpenAI) + slices.Sort(got) + want := []string{"gpt-4o"} + if !slices.Equal(got, want) { + t.Errorf("extractModelIDs = %v, want %v (anthropic-prefixed entry must be dropped when caller asks for openai)", got, want) + } +} + +// TestExtractModelIDs_NilResp returns nil — the public wrapper relies on +// this to short-circuit cleanly when a list-models call returns no body. +func TestExtractModelIDs_NilResp(t *testing.T) { + if got := extractModelIDs(nil, schemas.OpenAI); got != nil { + t.Errorf("extractModelIDs(nil) = %v, want nil", got) + } +} + +// TestExtractModelIDs_Dedup keeps only one entry when the same bare model +// name appears twice in the response (one prefixed, one bare). +func TestExtractModelIDs_Dedup(t *testing.T) { + resp := &schemas.BifrostListModelsResponse{ + Data: []schemas.Model{ + {ID: "openai/gpt-4o"}, + {ID: "gpt-4o"}, + {ID: "openai/gpt-4o"}, + }, + } + got := extractModelIDs(resp, schemas.OpenAI) + if len(got) != 1 || got[0] != "gpt-4o" { + t.Errorf("extractModelIDs = %v, want [gpt-4o] (deduped)", got) + } +} + +// TestInvalidateLive_DropsBothFiltersForKey verifies the InvalidateLive +// forwarder reaches the live store and clears filtered + unfiltered entries +// for one (provider, keyID) pair in a single call. +func TestInvalidateLive_DropsBothFiltersForKey(t *testing.T) { + mc := NewTestCatalog(nil) + mc.UpsertLive(schemas.OpenAI, "k1", false, []string{"gpt-4o"}) + mc.UpsertLive(schemas.OpenAI, "k1", true, []string{"gpt-4o", "o1"}) + mc.UpsertLive(schemas.OpenAI, "k2", false, []string{"o1"}) + + mc.InvalidateLive(schemas.OpenAI, "k1") + + // k1 entries are gone; k2 survives. + if got := mc.GetModelsForProvider(schemas.OpenAI); !slices.Equal(got, []string{"o1"}) { + t.Errorf("filtered union after InvalidateLive(k1) = %v, want [o1] (k1 filtered dropped, k2 survives)", got) + } + if got := mc.GetUnfilteredModelsForProvider(schemas.OpenAI); len(got) != 0 { + t.Errorf("unfiltered union after InvalidateLive(k1) = %v, want [] (k1 unfiltered dropped; k2 has no unfiltered entry)", got) + } +} + +// TestInvalidateLiveProvider_DropsAcrossKeys verifies the provider-wide +// forwarder clears every (keyID, mode) combination for the provider. +func TestInvalidateLiveProvider_DropsAcrossKeys(t *testing.T) { + mc := NewTestCatalog(nil) + mc.UpsertLive(schemas.OpenAI, "k1", false, []string{"gpt-4o"}) + mc.UpsertLive(schemas.OpenAI, "k2", false, []string{"o1"}) + mc.UpsertLive(schemas.Anthropic, "k1", false, []string{"claude-sonnet"}) + + mc.InvalidateLiveProvider(schemas.OpenAI) + + if got := mc.GetModelsForProvider(schemas.OpenAI); len(got) != 0 { + t.Errorf("OpenAI after InvalidateLiveProvider = %v, want [] (every key dropped)", got) + } + // Other providers untouched. + if got := mc.GetModelsForProvider(schemas.Anthropic); !slices.Equal(got, []string{"claude-sonnet"}) { + t.Errorf("Anthropic after InvalidateLiveProvider(OpenAI) = %v, want [claude-sonnet]", got) + } +} diff --git a/framework/modelcatalog/shims.go b/framework/modelcatalog/shims.go deleted file mode 100644 index ef014f2fcc9..00000000000 --- a/framework/modelcatalog/shims.go +++ /dev/null @@ -1,43 +0,0 @@ -// Compatibility shims preserved to keep server.go and the enterprise -// transport compiling without code changes in the same commit as the -// internal refactor. Each shim mirrors the pre-refactor API by aggregating -// into a single live entry keyed by "" (empty keyID). -// -// The follow-up PR replaces the call sites with per-key fanout via -// BifrostListModelsRequest.KeyID and then DELETES THIS WHOLE FILE. -package modelcatalog - -import ( - "github.com/maximhq/bifrost/core/schemas" -) - -// UpsertModelDataForProvider stores the merged filtered response for the -// provider in a single aggregated live entry. modelsInKeys is retained as a -// fallback when modelData is empty (provider list-models failed or no keys -// configured) — matches the pre-refactor "trust the user-allowed list" path. -// -// Deprecated: shim. Use UpsertLive per key once the call site adopts -// BifrostListModelsRequest.KeyID. -func (mc *ModelCatalog) UpsertModelDataForProvider(provider schemas.ModelProvider, modelData *schemas.BifrostListModelsResponse, modelsInKeys []schemas.Model) { - models := extractModelIDs(modelData, provider) - if len(models) == 0 { - models = extractModelIDs(&schemas.BifrostListModelsResponse{Data: modelsInKeys}, provider) - } - mc.live.Upsert(provider, "", false, models) -} - -// UpsertUnfilteredModelDataForProvider stores the unfiltered provider -// response in a single aggregated entry. -// -// Deprecated: shim. See UpsertModelDataForProvider. -func (mc *ModelCatalog) UpsertUnfilteredModelDataForProvider(provider schemas.ModelProvider, modelData *schemas.BifrostListModelsResponse) { - models := extractModelIDs(modelData, provider) - mc.live.Upsert(provider, "", true, models) -} - -// DeleteModelDataForProvider drops every live entry for the provider. -// -// Deprecated: shim. Use InvalidateLiveProvider. -func (mc *ModelCatalog) DeleteModelDataForProvider(provider schemas.ModelProvider) { - mc.live.InvalidateProvider(provider) -} diff --git a/plugins/governance/resolver_test.go b/plugins/governance/resolver_test.go index f3135a1bdce..86fb7a40416 100644 --- a/plugins/governance/resolver_test.go +++ b/plugins/governance/resolver_test.go @@ -9,7 +9,6 @@ import ( "github.com/maximhq/bifrost/core/schemas" "github.com/maximhq/bifrost/framework/configstore" configstoreTables "github.com/maximhq/bifrost/framework/configstore/tables" - "github.com/maximhq/bifrost/framework/modelcatalog" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -36,87 +35,6 @@ func TestBudgetResolver_EvaluateRequest_AllowedRequest(t *testing.T) { assertVirtualKeyFound(t, result) } -// TestBudgetResolver_EvaluateRequest_WildcardAllowsCatalogOpaqueProvider verifies that a -// wildcard ("*") allow-list permits any model on a catalog-opaque provider (vLLM, whose -// self-hosted models are never in the bundled catalog), while leaving catalog-known providers -// (openai) fully intact — i.e. wildcard is still catalog-cross-checked for them. -func TestBudgetResolver_EvaluateRequest_WildcardAllowsCatalogOpaqueProvider(t *testing.T) { - logger := NewMockLogger() - - // Catalog knows openai/gpt-4o but has NO model list for vLLM. - mc := modelcatalog.NewTestCatalog(map[string]string{"openai/gpt-4o": "gpt-4o"}) - mc.UpsertModelDataForProvider(schemas.OpenAI, - &schemas.BifrostListModelsResponse{Data: []schemas.Model{{ID: "openai/gpt-4o"}}}, nil) - - // Non-nil inMemoryStore so isModelAllowed takes the catalog branch. - inMem := &mockInMemoryStore{ - configuredProviders: map[schemas.ModelProvider]configstore.ProviderConfig{ - schemas.VLLM: {}, - schemas.OpenAI: {}, - }, - } - - vk := buildVirtualKey("vk1", "sk-bf-test", "Test VK", true) - vk.ProviderConfigs = []configstoreTables.TableVirtualKeyProviderConfig{ - buildProviderConfig("vllm", []string{"*"}), - buildProviderConfig("openai", []string{"*"}), - } - store, err := NewLocalGovernanceStore(context.Background(), logger, nil, &configstore.GovernanceConfig{ - VirtualKeys: []configstoreTables.TableVirtualKey{*vk}, - }, mc) - require.NoError(t, err) - - resolver := NewBudgetResolver(store, mc, logger, inMem) - ctx := &schemas.BifrostContext{} - - // vLLM (catalog-opaque) + ["*"] + uncatalogued model -> allowed (the fix). - result := resolver.EvaluateVirtualKeyRequest(ctx, "sk-bf-test", schemas.VLLM, "my-self-hosted-llama", schemas.ChatCompletionRequest, false) - assertDecision(t, DecisionAllow, result) - - // openai intact: a real catalog model under ["*"] is still allowed. - result = resolver.EvaluateVirtualKeyRequest(ctx, "sk-bf-test", schemas.OpenAI, "gpt-4o", schemas.ChatCompletionRequest, false) - assertDecision(t, DecisionAllow, result) - - // openai intact: an unknown model under ["*"] is still catalog-cross-checked and blocked. - result = resolver.EvaluateVirtualKeyRequest(ctx, "sk-bf-test", schemas.OpenAI, "not-a-real-model", schemas.ChatCompletionRequest, false) - assertDecision(t, DecisionModelBlocked, result) -} - -// TestBudgetResolver_EvaluateRequest_WildcardOpaqueProviderRespectsBlacklist guards the ordering -// in isModelAllowed: the blacklist pass must run before the wildcard + catalog-opaque shortcut, -// so a blacklisted model is blocked on an opaque provider even under a ["*"] allow-list. -func TestBudgetResolver_EvaluateRequest_WildcardOpaqueProviderRespectsBlacklist(t *testing.T) { - logger := NewMockLogger() - - mc := modelcatalog.NewTestCatalog(nil) // catalog has no vLLM models -> opaque - inMem := &mockInMemoryStore{ - configuredProviders: map[schemas.ModelProvider]configstore.ProviderConfig{ - schemas.VLLM: {}, - }, - } - - vllmConfig := buildProviderConfig("vllm", []string{"*"}) - vllmConfig.BlacklistedModels = schemas.BlackList{"my-self-hosted-llama"} - - vk := buildVirtualKey("vk1", "sk-bf-test", "Test VK", true) - vk.ProviderConfigs = []configstoreTables.TableVirtualKeyProviderConfig{vllmConfig} - store, err := NewLocalGovernanceStore(context.Background(), logger, nil, &configstore.GovernanceConfig{ - VirtualKeys: []configstoreTables.TableVirtualKey{*vk}, - }, mc) - require.NoError(t, err) - - resolver := NewBudgetResolver(store, mc, logger, inMem) - ctx := &schemas.BifrostContext{} - - // Blacklisted model on the opaque provider is blocked despite the ["*"] allow-list. - result := resolver.EvaluateVirtualKeyRequest(ctx, "sk-bf-test", schemas.VLLM, "my-self-hosted-llama", schemas.ChatCompletionRequest, false) - assertDecision(t, DecisionModelBlocked, result) - - // A different (non-blacklisted) model on the same opaque provider is still allowed. - result = resolver.EvaluateVirtualKeyRequest(ctx, "sk-bf-test", schemas.VLLM, "another-local-model", schemas.ChatCompletionRequest, false) - assertDecision(t, DecisionAllow, result) -} - // TestBudgetResolver_EvaluateRequest_VirtualKeyNotFound tests missing VK func TestBudgetResolver_EvaluateRequest_VirtualKeyNotFound(t *testing.T) { logger := NewMockLogger() diff --git a/transports/bifrost-http/handlers/provider_keys.go b/transports/bifrost-http/handlers/provider_keys.go index cb8614c1ebc..449801dfb49 100644 --- a/transports/bifrost-http/handlers/provider_keys.go +++ b/transports/bifrost-http/handlers/provider_keys.go @@ -139,8 +139,10 @@ func (h *ProviderHandler) createProviderKey(ctx *fasthttp.RequestCtx) { return } - if err := h.attemptModelDiscovery(ctx, provider, providerConfig.CustomProviderConfig); err != nil { - logger.Warn("Model discovery failed for provider %s after key create: %v", provider, err) + if providerConfig.CustomProviderConfig == nil || !providerConfig.CustomProviderConfig.IsKeyLess { + if err := h.modelsManager.OnKeyAdded(ctx, provider, key); err != nil { + logger.Warn("Catalog refresh failed for provider %s after key create: %v", provider, err) + } } redactedKey, err := h.inMemoryStore.GetProviderKeyRedacted(provider, key.ID) @@ -244,8 +246,10 @@ func (h *ProviderHandler) updateProviderKey(ctx *fasthttp.RequestCtx) { return } - if err := h.attemptModelDiscovery(ctx, provider, providerConfig.CustomProviderConfig); err != nil { - logger.Warn("Model discovery failed for provider %s after key update: %v", provider, err) + if providerConfig.CustomProviderConfig == nil || !providerConfig.CustomProviderConfig.IsKeyLess { + if err := h.modelsManager.OnKeyUpdated(ctx, provider, mergedKey); err != nil { + logger.Warn("Catalog refresh failed for provider %s after key update: %v", provider, err) + } } redactedKey, err := h.inMemoryStore.GetProviderKeyRedacted(provider, keyID) @@ -305,8 +309,8 @@ func (h *ProviderHandler) deleteProviderKey(ctx *fasthttp.RequestCtx) { return } - if err := h.attemptModelDiscovery(ctx, provider, providerConfig.CustomProviderConfig); err != nil { - logger.Warn("Model discovery failed for provider %s after key delete: %v", provider, err) + if err := h.modelsManager.OnKeyDeleted(ctx, provider, keyID); err != nil { + logger.Warn("Catalog refresh failed for provider %s after key delete: %v", provider, err) } SendJSON(ctx, redactedKey) diff --git a/transports/bifrost-http/handlers/providers.go b/transports/bifrost-http/handlers/providers.go index 059f5e138f0..4d2ed7b0ae9 100644 --- a/transports/bifrost-http/handlers/providers.go +++ b/transports/bifrost-http/handlers/providers.go @@ -31,6 +31,9 @@ type ModelsManager interface { GetModelsForProvider(provider schemas.ModelProvider) []string GetUnfilteredModelsForProvider(provider schemas.ModelProvider) []string UpsertModelPricingAttributes(ctx context.Context, entries []ModelPricingAttributesEntry) error + OnKeyAdded(ctx context.Context, provider schemas.ModelProvider, key schemas.Key) error + OnKeyUpdated(ctx context.Context, provider schemas.ModelProvider, key schemas.Key) error + OnKeyDeleted(ctx context.Context, provider schemas.ModelProvider, keyID string) error } // ModelPricingAttributesEntry is the wire shape for PUT /api/models/catalog. diff --git a/transports/bifrost-http/handlers/providers_test.go b/transports/bifrost-http/handlers/providers_test.go index 26755d046ef..09398d9592f 100644 --- a/transports/bifrost-http/handlers/providers_test.go +++ b/transports/bifrost-http/handlers/providers_test.go @@ -54,6 +54,18 @@ func (m *mockModelsManager) UpsertModelPricingAttributes(_ context.Context, _ [] return nil } +func (m *mockModelsManager) OnKeyAdded(_ context.Context, _ schemas.ModelProvider, _ schemas.Key) error { + return nil +} + +func (m *mockModelsManager) OnKeyUpdated(_ context.Context, _ schemas.ModelProvider, _ schemas.Key) error { + return nil +} + +func (m *mockModelsManager) OnKeyDeleted(_ context.Context, _ schemas.ModelProvider, _ string) error { + return nil +} + // providerHandlerForTest builds a handler with fixed provider config and model sets. func providerHandlerForTest(provider schemas.ModelProvider, keys []schemas.Key, filtered, unfiltered []string) *ProviderHandler { return &ProviderHandler{ @@ -394,7 +406,7 @@ func TestListModelDetails_UnknownKeysDoNotFilter(t *testing.T) { []string{"gpt-4o", "gpt-4o-mini"}, []string{"gpt-4o", "gpt-4o-mini"}, ) - h.inMemoryStore.ModelCatalog = &modelcatalog.ModelCatalog{} + h.inMemoryStore.ModelCatalog = modelcatalog.NewTestCatalog(nil) ctx := &fasthttp.RequestCtx{} ctx.Request.Header.SetMethod("GET") @@ -425,7 +437,7 @@ func TestListModelDetails_SkipsUnknownKeysAndFiltersWithValid(t *testing.T) { []string{"gpt-4o", "gpt-4o-mini"}, []string{"gpt-4o", "gpt-4o-mini"}, ) - h.inMemoryStore.ModelCatalog = &modelcatalog.ModelCatalog{} + h.inMemoryStore.ModelCatalog = modelcatalog.NewTestCatalog(nil) ctx := &fasthttp.RequestCtx{} ctx.Request.Header.SetMethod("GET") @@ -462,7 +474,7 @@ func TestListModelDetails_SkipsDisabledKeysAndFiltersWithValid(t *testing.T) { []string{"gpt-4o", "gpt-4o-mini"}, []string{"gpt-4o", "gpt-4o-mini"}, ) - h.inMemoryStore.ModelCatalog = &modelcatalog.ModelCatalog{} + h.inMemoryStore.ModelCatalog = modelcatalog.NewTestCatalog(nil) ctx := &fasthttp.RequestCtx{} ctx.Request.Header.SetMethod("GET") @@ -498,7 +510,7 @@ func TestListModelDetails_UnfilteredIgnoresKeys(t *testing.T) { []string{"gpt-4o"}, []string{"gpt-4o", "gpt-4o-mini"}, ) - h.inMemoryStore.ModelCatalog = &modelcatalog.ModelCatalog{} + h.inMemoryStore.ModelCatalog = modelcatalog.NewTestCatalog(nil) ctx := &fasthttp.RequestCtx{} ctx.Request.Header.SetMethod("GET") diff --git a/transports/bifrost-http/server/server.go b/transports/bifrost-http/server/server.go index 56b380cf880..d78836649c3 100644 --- a/transports/bifrost-http/server/server.go +++ b/transports/bifrost-http/server/server.go @@ -100,6 +100,9 @@ type ServerCallbacks interface { RemoveModelConfig(ctx context.Context, id string) error ReloadProvider(ctx context.Context, provider schemas.ModelProvider) (*tables.TableProvider, error) RemoveProvider(ctx context.Context, provider schemas.ModelProvider) error + OnKeyAdded(ctx context.Context, provider schemas.ModelProvider, key schemas.Key) error + OnKeyUpdated(ctx context.Context, provider schemas.ModelProvider, key schemas.Key) error + OnKeyDeleted(ctx context.Context, provider schemas.ModelProvider, keyID string) error ReloadRoutingRule(ctx context.Context, id string) error RemoveRoutingRule(ctx context.Context, id string) error // MCP related callbacks @@ -624,84 +627,23 @@ func (s *BifrostHTTPServer) ReloadProvider(ctx context.Context, provider schemas } } - // Read current key count from in-memory store (providerInfo.Keys is not preloaded from DB) - inMemoryKeys, _ := s.Config.GetProviderKeysRaw(provider) - isKeylessProvider := providerInfo.CustomProviderConfig != nil && providerInfo.CustomProviderConfig.IsKeyLess - hasNoKeys := len(inMemoryKeys) == 0 && !isKeylessProvider - - // Getting allowed models from all provider keys (needed before model listing) - providerKeys, err := s.Config.ConfigStore.GetKeysByProvider(ctx, string(provider)) + // In-memory store holds the latest schemas.Key slice after the most recent + // CRUD write — read from there to avoid re-fetching + re-converting from DB. + inMemoryKeys, err := s.Config.GetProviderKeysRaw(provider) if err != nil { - return nil, fmt.Errorf("failed to update provider model catalog: failed to get keys by provider: %s", err) - } - - bfCtx := schemas.NewBifrostContext(ctx, time.Now().Add(15*time.Second)) - bfCtx.SetValue(schemas.BifrostContextKeySkipPluginPipeline, true) - bfCtx.SetValue(schemas.BifrostContextKeyValidateKeys, true) // Validate keys during provider add/update - defer bfCtx.Cancel() - - // Run filtered and unfiltered model listing concurrently - var ( - allModels *schemas.BifrostListModelsResponse - bifrostErr *schemas.BifrostError - unfilteredModels *schemas.BifrostListModelsResponse - listModelsErr *schemas.BifrostError - listWg sync.WaitGroup - ) - listWg.Add(2) - go func() { - defer listWg.Done() - allModels, bifrostErr = s.Client.ListModelsRequest(bfCtx, &schemas.BifrostListModelsRequest{ - Provider: provider, - }) - }() - go func() { - defer listWg.Done() - unfilteredModels, listModelsErr = s.Client.ListModelsRequest(bfCtx, &schemas.BifrostListModelsRequest{ - Provider: provider, - Unfiltered: true, - }) - }() - listWg.Wait() - - if allModels != nil && len(allModels.KeyStatuses) > 0 && s.Config.ConfigStore != nil { - s.updateKeyStatus(ctx, allModels.KeyStatuses) + return nil, fmt.Errorf("failed to read provider keys for %s: %w", provider, err) } - if bifrostErr != nil { - if len(bifrostErr.ExtraFields.KeyStatuses) > 0 && s.Config.ConfigStore != nil { - s.updateKeyStatus(ctx, bifrostErr.ExtraFields.KeyStatuses) - } + isKeylessProvider := providerInfo.CustomProviderConfig != nil && providerInfo.CustomProviderConfig.IsKeyLess + hasNoKeys := len(inMemoryKeys) == 0 && !isKeylessProvider - if hasNoKeys { - logger.Warn("model discovery skipped for provider %s: no keys configured", provider) - } else { - logger.Warn("failed to update provider model catalog: failed to list all models: %s. We are falling back onto the static datasheet", bifrost.GetErrorMessage(bifrostErr)) - } - // In case of error, we return an empty list of models, and fallback onto the static datasheet - allModels = &schemas.BifrostListModelsResponse{ - Data: make([]schemas.Model, 0), - } - } - modelsInKeys := make([]schemas.Model, 0) - for _, key := range providerKeys { - if key.Models.IsUnrestricted() { - continue - } - for _, model := range key.Models { - modelsInKeys = append(modelsInKeys, schemas.Model{ - ID: string(provider) + "/" + model, - }) - } - } - s.Config.ModelCatalog.UpsertModelDataForProvider(provider, allModels, modelsInKeys) - if listModelsErr != nil { - if hasNoKeys { - logger.Warn("unfiltered model discovery skipped for provider %s: no keys configured", provider) - } else { - logger.Error("failed to list unfiltered models for provider %s: %v: falling back onto the static datasheet", provider, bifrost.GetErrorMessage(listModelsErr)) - } + // Refresh keyconfig from the current key list, then drop any stale live + // entries (for keys removed in this update) before refetching per-key. + s.Config.ModelCatalog.SetKeyConfigForProvider(provider, inMemoryKeys) + s.Config.ModelCatalog.InvalidateLiveProvider(provider) + if hasNoKeys { + logger.Warn("model discovery skipped for provider %s: no keys configured", provider) } else { - s.Config.ModelCatalog.UpsertUnfilteredModelDataForProvider(provider, unfilteredModels) + s.RefreshLiveModelsForProvider(ctx, provider, inMemoryKeys) } return updatedProvider, nil } @@ -726,11 +668,84 @@ func (s *BifrostHTTPServer) RemoveProvider(ctx context.Context, provider schemas if s.Config == nil || s.Config.ModelCatalog == nil { return fmt.Errorf("pricing manager not found") } - s.Config.ModelCatalog.DeleteModelDataForProvider(provider) + s.Config.ModelCatalog.InvalidateLiveProvider(provider) + s.Config.ModelCatalog.RemoveKeyConfigForProvider(provider) return nil } +// OnKeyAdded refreshes the keyconfig snapshot and fetches list-models for the +// new key only — 2 calls instead of ReloadProvider's 2×N. Called by the key +// handler after a successful AddProviderKey write. +func (s *BifrostHTTPServer) OnKeyAdded(ctx context.Context, provider schemas.ModelProvider, key schemas.Key) error { + if s.Config == nil || s.Config.ModelCatalog == nil { + return fmt.Errorf("model catalog not found") + } + keys, err := s.Config.GetProviderKeysRaw(provider) + if err != nil { + return fmt.Errorf("failed to read provider keys for %s: %w", provider, err) + } + s.Config.ModelCatalog.SetKeyConfigForProvider(provider, keys) + // Keyless providers: empty keyID sentinel. + keyID := key.ID + if isKeylessProvider(provider, s.Config) { + keyID = "" + } + s.FetchAndStoreLiveForKey(ctx, provider, keyID) + return nil +} + +// OnKeyUpdated invalidates the affected key's live entries (the gate may have +// changed even when Value didn't), refreshes the keyconfig, then refetches +// for just that key. 2 calls regardless of N keys on the provider. +func (s *BifrostHTTPServer) OnKeyUpdated(ctx context.Context, provider schemas.ModelProvider, key schemas.Key) error { + if s.Config == nil || s.Config.ModelCatalog == nil { + return fmt.Errorf("model catalog not found") + } + keys, err := s.Config.GetProviderKeysRaw(provider) + if err != nil { + return fmt.Errorf("failed to read provider keys for %s: %w", provider, err) + } + s.Config.ModelCatalog.SetKeyConfigForProvider(provider, keys) + keyID := key.ID + if isKeylessProvider(provider, s.Config) { + keyID = "" + } + s.Config.ModelCatalog.InvalidateLive(provider, keyID) + s.FetchAndStoreLiveForKey(ctx, provider, keyID) + return nil +} + +// OnKeyDeleted invalidates the deleted key's live entries and refreshes the +// keyconfig. No list-models calls — the provider's remaining keys' cached +// entries stay valid. +func (s *BifrostHTTPServer) OnKeyDeleted(ctx context.Context, provider schemas.ModelProvider, keyID string) error { + if s.Config == nil || s.Config.ModelCatalog == nil { + return fmt.Errorf("model catalog not found") + } + keys, err := s.Config.GetProviderKeysRaw(provider) + if err != nil { + return fmt.Errorf("failed to read provider keys for %s: %w", provider, err) + } + s.Config.ModelCatalog.SetKeyConfigForProvider(provider, keys) + s.Config.ModelCatalog.InvalidateLive(provider, keyID) + return nil +} + +// isKeylessProvider returns true when the provider's config marks it +// keyless. Used to pick the live-cache key for OnKey* helpers: keyless +// providers cache under the empty-string sentinel. +func isKeylessProvider(provider schemas.ModelProvider, cfg *lib.Config) bool { + if cfg == nil { + return false + } + pc, err := cfg.GetProviderConfigRaw(provider) + if err != nil || pc == nil || pc.CustomProviderConfig == nil { + return false + } + return pc.CustomProviderConfig.IsKeyLess +} + // GetGovernanceData returns the governance data func (s *BifrostHTTPServer) GetGovernanceData(ctx context.Context) *governance.GovernanceData { // Use type-safe finder from Config @@ -895,51 +910,120 @@ func (s *BifrostHTTPServer) UpdateSyncConfig(ctx context.Context) error { return s.Config.ModelCatalog.UpdateSyncConfig(ctx, s.Config.FrameworkConfig.Pricing) } -func (s *BifrostHTTPServer) populateModelPoolWithListModels(ctx context.Context) error { - // Fetching keys for all providers and allowed models first - // Based on allowed models we will set the data in the model catalog +// RefreshLiveModelsForProvider runs filtered + unfiltered list-models for the +// provider, fanning out per key in parallel so the live cache ends up with +// per-(provider, keyID) entries. Keyless providers cache under the "" sentinel. +// +// Callers are responsible for invalidating stale entries first when keys +// have been removed from the provider's set. +func (s *BifrostHTTPServer) RefreshLiveModelsForProvider(ctx context.Context, provider schemas.ModelProvider, keys []schemas.Key) { + if len(keys) == 0 { + // Empty key slice + non-keyless provider would write under the "" sentinel + // reserved for keyless providers — colliding with the keyless namespace and + // triggering an unauthenticated fetch for a provider that requires a key. + if !isKeylessProvider(provider, s.Config) { + logger.Warn("model discovery skipped for provider %s: no keys configured", provider) + return + } + s.FetchAndStoreLiveForKey(ctx, provider, "") + return + } var wg sync.WaitGroup - for provider, providerConfig := range s.Config.Providers { + for _, key := range keys { wg.Add(1) - go func(provider schemas.ModelProvider, providerConfig configstore.ProviderConfig) { + go func(keyID string) { defer wg.Done() - bfCtx := schemas.NewBifrostContext(ctx, time.Now().Add(15*time.Second)) - bfCtx.SetValue(schemas.BifrostContextKeySkipPluginPipeline, true) - defer bfCtx.Cancel() - modelData, listModelsErr := s.Client.ListModelsRequest(bfCtx, &schemas.BifrostListModelsRequest{ - Provider: provider, - }) - if listModelsErr != nil { - logger.Error("failed to list models for provider %s: %v: falling back onto the static datasheet", provider, bifrost.GetErrorMessage(listModelsErr)) - } - allowedModels := make([]schemas.Model, 0) - for _, key := range providerConfig.Keys { - if key.Models.IsUnrestricted() { - continue - } - for _, model := range key.Models { - allowedModels = append(allowedModels, schemas.Model{ - ID: string(provider) + "/" + model, - }) - } - } - s.Config.ModelCatalog.UpsertModelDataForProvider(provider, modelData, allowedModels) - unfilteredModelData, listModelsErr := s.Client.ListModelsRequest(bfCtx, &schemas.BifrostListModelsRequest{ - Provider: provider, - Unfiltered: true, - }) - if listModelsErr != nil { - logger.Error("failed to list unfiltered models for provider %s: %v: falling back onto the static datasheet", provider, bifrost.GetErrorMessage(listModelsErr)) - } else { - s.Config.ModelCatalog.UpsertUnfilteredModelDataForProvider(provider, unfilteredModelData) - } - }(provider, providerConfig) + s.FetchAndStoreLiveForKey(ctx, provider, keyID) + }(key.ID) } wg.Wait() - return nil } -// ForceReloadPricing triggers an immediate pricing sync and resets the sync timer +// FetchAndStoreLiveForKey issues the filtered and unfiltered list-models +// calls for one (provider, keyID) in parallel and writes the results into +// the catalog. Errors are logged and surfaced via updateKeyStatus when the +// provider returns per-key statuses, but they do not abort the other call. +// keyID="" scopes to "no specific key" — used for keyless providers and as +// the legacy sentinel. Always validates keys for the providers that opt into +// the check (today: OpenRouter, whose /v1/models is unauthenticated) so the +// routing graph is the same at boot, after a key add, and after a reload — +// stale-but-routable behavior would diverge otherwise. +func (s *BifrostHTTPServer) FetchAndStoreLiveForKey(ctx context.Context, provider schemas.ModelProvider, keyID string) { + // Skip the fetch entirely when the provider has disabled list_models via + // allowed_requests — every per-(provider,keyID) call would just bounce with + // "operation not allowed", wasting two goroutines and one bfCtx per attempt. + if s.Config != nil { + if pc, err := s.Config.GetProviderConfigRaw(provider); err == nil && pc != nil && + pc.CustomProviderConfig != nil && + !pc.CustomProviderConfig.IsOperationAllowed(schemas.ListModelsRequest) { + return + } + } + // One BifrostContext per goroutine. BifrostContext.SetValue mutates state + // in place, so the request-scoped metadata core sets during a routing pass + // (RequestID, FallbackIndex, span IDs, ...) would otherwise bleed between + // the filtered and unfiltered calls and conflate them in logs/billing. + newListModelsCtx := func() *schemas.BifrostContext { + c := schemas.NewBifrostContext(ctx, time.Now().Add(15*time.Second)) + c.SetValue(schemas.BifrostContextKeySkipPluginPipeline, true) + c.SetValue(schemas.BifrostContextKeyValidateKeys, true) + return c + } + + var keyIDPtr *string + if keyID != "" { + keyIDPtr = &keyID + } + + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + bfCtx := newListModelsCtx() + defer bfCtx.Cancel() + resp, bfErr := s.Client.ListModelsRequest(bfCtx, &schemas.BifrostListModelsRequest{ + Provider: provider, + KeyID: keyIDPtr, + }) + if bfErr != nil { + logger.Warn("filtered list-models failed for provider %s key %s: %v: falling back onto the static datasheet", provider, keyID, bifrost.GetErrorMessage(bfErr)) + if len(bfErr.ExtraFields.KeyStatuses) > 0 && s.Config.ConfigStore != nil { + s.updateKeyStatus(ctx, bfErr.ExtraFields.KeyStatuses) + } + return + } + if resp == nil { + return + } + s.Config.ModelCatalog.UpsertLiveFromResponse(provider, keyID, false, resp) + if len(resp.KeyStatuses) > 0 && s.Config.ConfigStore != nil { + s.updateKeyStatus(ctx, resp.KeyStatuses) + } + }() + go func() { + defer wg.Done() + bfCtx := newListModelsCtx() + defer bfCtx.Cancel() + resp, bfErr := s.Client.ListModelsRequest(bfCtx, &schemas.BifrostListModelsRequest{ + Provider: provider, + KeyID: keyIDPtr, + Unfiltered: true, + }) + if bfErr != nil { + logger.Warn("unfiltered list-models failed for provider %s key %s: %v: falling back onto the static datasheet", provider, keyID, bifrost.GetErrorMessage(bfErr)) + return + } + if resp == nil { + return + } + s.Config.ModelCatalog.UpsertLiveFromResponse(provider, keyID, true, resp) + }() + wg.Wait() +} + +// ForceReloadPricing triggers an immediate pricing sync and resets the sync +// timer. No longer triggers a list-models refresh — pricing reload is now +// pricing-only. func (s *BifrostHTTPServer) ForceReloadPricing(ctx context.Context) error { if s.Config == nil { return fmt.Errorf("server config not initialized") @@ -948,12 +1032,13 @@ func (s *BifrostHTTPServer) ForceReloadPricing(ctx context.Context) error { if err := s.Config.ModelCatalog.ForceReloadPricing(ctx); err != nil { return fmt.Errorf("failed to force reload pricing: %w", err) } - return s.populateModelPoolWithListModels(ctx) } return nil } -// ReloadPricingFromDBAndPopulateModelPool reloads the pricing from DB and populates the model pool +// ReloadPricingFromDBAndPopulateModelPool reloads the pricing from DB. The +// list-models refresh that used to follow is gone — pricing reload is now +// pricing-only. func (s *BifrostHTTPServer) ReloadPricingFromDBAndPopulateModelPool(ctx context.Context) error { if s.Config == nil { return fmt.Errorf("server config not initialized") @@ -962,7 +1047,6 @@ func (s *BifrostHTTPServer) ReloadPricingFromDBAndPopulateModelPool(ctx context. if err := s.Config.ModelCatalog.ReloadFromDB(ctx); err != nil { return fmt.Errorf("failed to reload pricing from DB: %w", err) } - return s.populateModelPoolWithListModels(ctx) } return nil } @@ -1516,54 +1600,23 @@ func (s *BifrostHTTPServer) Bootstrap(ctx context.Context) error { // Sync plugin execution order from config to core (defensive — Init receives sorted list, // but this ensures order consistency if the loading path changes in the future) s.Client.ReorderPlugins(s.Config.GetPluginOrder()) - // List all models and add to model catalog with per-provider status tracking + // Seed the catalog: push the initial keyconfig snapshot and fetch per-key + // live models for every provider concurrently. logger.Info("listing all models and adding to model catalog") if s.Config.ModelCatalog != nil { - // Fetching keys for all providers and allowed models first - // Based on allowed models we will set the data in the model catalog + snapshot := make(map[schemas.ModelProvider][]schemas.Key, len(s.Config.Providers)) + for provider, providerConfig := range s.Config.Providers { + snapshot[provider] = providerConfig.Keys + } + s.Config.ModelCatalog.ReplaceKeyConfig(snapshot) + var wg sync.WaitGroup for provider, providerConfig := range s.Config.Providers { wg.Add(1) - go func(provider schemas.ModelProvider, providerConfig configstore.ProviderConfig) { + go func(p schemas.ModelProvider, keys []schemas.Key) { defer wg.Done() - bfCtx := schemas.NewBifrostContext(ctx, time.Now().Add(15*time.Second)) - bfCtx.SetValue(schemas.BifrostContextKeySkipPluginPipeline, true) - defer bfCtx.Cancel() - - modelData, listModelsErr := s.Client.ListModelsRequest(bfCtx, &schemas.BifrostListModelsRequest{ - Provider: provider, - }) - if modelData != nil && len(modelData.KeyStatuses) > 0 && s.Config.ConfigStore != nil { - s.updateKeyStatus(ctx, modelData.KeyStatuses) - } - if listModelsErr != nil { - if len(listModelsErr.ExtraFields.KeyStatuses) > 0 && s.Config.ConfigStore != nil { - s.updateKeyStatus(ctx, listModelsErr.ExtraFields.KeyStatuses) - } - logger.Error("failed to list models for provider %s: %v: falling back onto the static datasheet", provider, bifrost.GetErrorMessage(listModelsErr)) - } - allowedModels := make([]schemas.Model, 0) - for _, key := range providerConfig.Keys { - if key.Models.IsUnrestricted() { - continue - } - for _, model := range key.Models { - allowedModels = append(allowedModels, schemas.Model{ - ID: string(provider) + "/" + model, - }) - } - } - s.Config.ModelCatalog.UpsertModelDataForProvider(provider, modelData, allowedModels) - unfilteredModelData, listModelsErr := s.Client.ListModelsRequest(bfCtx, &schemas.BifrostListModelsRequest{ - Provider: provider, - Unfiltered: true, - }) - if listModelsErr != nil { - logger.Error("failed to list unfiltered models for provider %s: %v: falling back onto the static datasheet", provider, bifrost.GetErrorMessage(listModelsErr)) - } else { - s.Config.ModelCatalog.UpsertUnfilteredModelDataForProvider(provider, unfilteredModelData) - } - }(provider, providerConfig) + s.RefreshLiveModelsForProvider(ctx, p, keys) + }(provider, providerConfig.Keys) } wg.Wait() }