Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion framework/modelcatalog/models.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
12 changes: 12 additions & 0 deletions framework/modelcatalog/pool.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.

// InvalidateLive drops both filtered + unfiltered live entries for one key.
func (mc *ModelCatalog) InvalidateLive(provider schemas.ModelProvider, keyID string) {
mc.live.Invalidate(provider, keyID)
Expand Down
169 changes: 169 additions & 0 deletions framework/modelcatalog/pool_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
43 changes: 0 additions & 43 deletions framework/modelcatalog/shims.go

This file was deleted.

4 changes: 2 additions & 2 deletions plugins/governance/httptransportprehook_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -25,11 +25,11 @@ func TestHTTPTransportPreHook_VirtualKeyReplicateRefinesNestedModel(t *testing.T
mc := modelcatalog.NewTestCatalog(map[string]string{
"openai/gpt-5-nano": "gpt-5-nano",
})
mc.UpsertModelDataForProvider(schemas.Replicate, &schemas.BifrostListModelsResponse{
mc.UpsertLiveFromResponse(schemas.Replicate, "", false, &schemas.BifrostListModelsResponse{
Data: []schemas.Model{
{ID: "replicate/openai/gpt-5-nano"},
},
}, nil)
})

virtualKey := buildVirtualKeyWithProviders(
"vk1",
Expand Down
82 changes: 0 additions & 82 deletions plugins/governance/resolver_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand All @@ -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()
Expand Down
16 changes: 10 additions & 6 deletions transports/bifrost-http/handlers/provider_keys.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
3 changes: 3 additions & 0 deletions transports/bifrost-http/handlers/providers.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
Loading
Loading