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
17 changes: 10 additions & 7 deletions framework/configstore/migrations.go
Original file line number Diff line number Diff line change
Expand Up @@ -840,6 +840,9 @@ func triggerMigrations(ctx context.Context, db *gorm.DB) error {
if err := migrationAddModelConfigCalendarAlignedColumn(ctx, db); err != nil {
return err
}
if err := migrationMigrateVirtualKeyGovernanceToModelConfigs(ctx, db); err != nil {
return err
}
return nil
}

Expand Down Expand Up @@ -4047,9 +4050,9 @@ func migrationAddBudgetModelConfigIDColumn(ctx context.Context, db *gorm.DB) err
return nil
}

// ensureVKWildcardModelConfig returns the ID of the VK-scoped all-models wildcard model
// config, creating it if absent.
func ensureVKWildcardModelConfig(tx *gorm.DB, vkID string, provider *string, calendarAligned bool, now time.Time) (string, error) {
// ensureVKModelConfig returns the ID of the VK-scoped model config for the given
// (vkID, provider) pair, creating it if absent.
func ensureVKModelConfig(tx *gorm.DB, vkID string, provider *string, calendarAligned bool, now time.Time) (string, error) {
q := tx.Model(&tables.TableModelConfig{}).
Where("scope = ? AND scope_id = ? AND model_name = ?",
tables.ModelConfigScopeVirtualKey, vkID, tables.ModelConfigAllModels)
Expand All @@ -4060,7 +4063,7 @@ func ensureVKWildcardModelConfig(tx *gorm.DB, vkID string, provider *string, cal
}
var existing []tables.TableModelConfig
if err := q.Limit(1).Find(&existing).Error; err != nil {
return "", fmt.Errorf("failed to look up VK wildcard model config: %w", err)
return "", fmt.Errorf("failed to look up VK model config: %w", err)
}
if len(existing) > 0 {
return existing[0].ID, nil
Expand All @@ -4076,7 +4079,7 @@ func ensureVKWildcardModelConfig(tx *gorm.DB, vkID string, provider *string, cal
UpdatedAt: now,
}
if err := tx.Create(&mc).Error; err != nil {
return "", fmt.Errorf("failed to create VK wildcard model config: %w", err)
return "", fmt.Errorf("failed to create VK model config: %w", err)
}
return mc.ID, nil
}
Expand Down Expand Up @@ -4110,7 +4113,7 @@ func migrationMigrateVirtualKeyGovernanceToModelConfigs(ctx context.Context, db

// VK top-level governance -> all-providers wildcard.
if len(vk.Budgets) > 0 || vk.RateLimitID != nil {
mcID, err := ensureVKWildcardModelConfig(tx, vk.ID, nil, vk.CalendarAligned, now)
mcID, err := ensureVKModelConfig(tx, vk.ID, nil, vk.CalendarAligned, now)
if err != nil {
return err
}
Expand Down Expand Up @@ -4139,7 +4142,7 @@ func migrationMigrateVirtualKeyGovernanceToModelConfigs(ctx context.Context, db
continue
}
provider := pc.Provider
mcID, err := ensureVKWildcardModelConfig(tx, vk.ID, &provider, vk.CalendarAligned, now)
mcID, err := ensureVKModelConfig(tx, vk.ID, &provider, vk.CalendarAligned, now)
if err != nil {
return err
}
Expand Down
65 changes: 45 additions & 20 deletions framework/configstore/rdb.go
Original file line number Diff line number Diff line change
Expand Up @@ -1163,29 +1163,36 @@ func (s *RDBConfigStore) DeleteProvider(ctx context.Context, provider schemas.Mo
if err := txDB.WithContext(ctx).Preload("Budgets").Where("provider = ?", string(provider)).Find(&providerModelConfigs).Error; err != nil {
return err
}
for _, mc := range providerModelConfigs {
for i := range mc.Budgets {
if err := txDB.WithContext(ctx).Delete(&tables.TableBudget{}, "id = ?", mc.Budgets[i].ID).Error; err != nil {
return err
if len(providerModelConfigs) > 0 {
var mcIDs []string
var budgetIDs []string
var rateLimitIDs []string
for i := range providerModelConfigs {
mcIDs = append(mcIDs, providerModelConfigs[i].ID)
for j := range providerModelConfigs[i].Budgets {
budgetIDs = append(budgetIDs, providerModelConfigs[i].Budgets[j].ID)
}
if providerModelConfigs[i].BudgetID != nil {
budgetIDs = append(budgetIDs, *providerModelConfigs[i].BudgetID)
}
if providerModelConfigs[i].RateLimitID != nil {
rateLimitIDs = append(rateLimitIDs, *providerModelConfigs[i].RateLimitID)
}
}
if err := txDB.WithContext(ctx).Where("id IN ?", mcIDs).Delete(&tables.TableModelConfig{}).Error; err != nil {
return err
}
if mc.BudgetID != nil {
if err := txDB.WithContext(ctx).Delete(&tables.TableBudget{}, "id = ?", *mc.BudgetID).Error; err != nil {
if len(budgetIDs) > 0 {
if err := txDB.WithContext(ctx).Delete(&tables.TableBudget{}, "id IN ?", budgetIDs).Error; err != nil {
return err
}
}
if mc.RateLimitID != nil {
if err := txDB.WithContext(ctx).Delete(&tables.TableRateLimit{}, "id = ?", *mc.RateLimitID).Error; err != nil {
if len(rateLimitIDs) > 0 {
if err := txDB.WithContext(ctx).Delete(&tables.TableRateLimit{}, "id IN ?", rateLimitIDs).Error; err != nil {
return err
}
}
}
if len(providerModelConfigs) > 0 {
var mcIDs []string
if err := txDB.WithContext(ctx).Where("id IN ?", mcIDs).Delete(&tables.TableModelConfig{}).Error; err != nil {
return err
}
}

return nil
}
Expand Down Expand Up @@ -3002,12 +3009,6 @@ func (s *RDBConfigStore) DeleteVirtualKey(ctx context.Context, id string, tx ...
budgetIDs := make([]string, 0, len(scopedModelConfigs))
rateLimitIDs := make([]string, 0, len(scopedModelConfigs))
for _, mc := range scopedModelConfigs {
// Owned budgets via ModelConfigID plus the legacy single BudgetID for safety.
for i := range mc.Budgets {
if err := txDB.WithContext(ctx).Delete(&tables.TableBudget{}, "id = ?", mc.Budgets[i].ID).Error; err != nil {
return err
}
}
if mc.BudgetID != nil {
budgetIDs = append(budgetIDs, *mc.BudgetID)
}
Expand All @@ -3020,6 +3021,16 @@ func (s *RDBConfigStore) DeleteVirtualKey(ctx context.Context, id string, tx ...
Delete(&tables.TableModelConfig{}).Error; err != nil {
return err
}
if len(budgetIDs) > 0 {
if err := txDB.WithContext(ctx).Delete(&tables.TableBudget{}, "id IN ?", budgetIDs).Error; err != nil {
return err
}
}
if len(rateLimitIDs) > 0 {
if err := txDB.WithContext(ctx).Delete(&tables.TableRateLimit{}, "id IN ?", rateLimitIDs).Error; err != nil {
return err
}
}
rateLimitID := virtualKey.RateLimitID
// Delete the virtual key
if err := txDB.WithContext(ctx).Delete(&tables.TableVirtualKey{}, "id = ?", id).Error; err != nil {
Expand Down Expand Up @@ -4261,6 +4272,20 @@ func (s *RDBConfigStore) GetModelConfigs(ctx context.Context) ([]tables.TableMod
return modelConfigs, nil
}

// GetModelConfigsByScopeAndScopeIDs retrieves model configs for a specific scope limited to the given scope IDs.
func (s *RDBConfigStore) GetModelConfigsByScopeAndScopeIDs(ctx context.Context, scope string, scopeIDs []string) ([]tables.TableModelConfig, error) {
if len(scopeIDs) == 0 {
return nil, nil
}
var modelConfigs []tables.TableModelConfig
if err := s.DB().WithContext(ctx).Preload("Budgets").Preload("Budget").Preload("RateLimit").
Where("scope = ? AND scope_id IN ?", scope, scopeIDs).
Find(&modelConfigs).Error; err != nil {
return nil, err
}
return modelConfigs, nil
}

// GetProviderGovernanceModelConfigs retrieves the wildcard "all models on a provider" configs
func (s *RDBConfigStore) GetProviderGovernanceModelConfigs(ctx context.Context) ([]tables.TableModelConfig, error) {
var modelConfigs []tables.TableModelConfig
Expand Down
1 change: 1 addition & 0 deletions framework/configstore/store.go
Original file line number Diff line number Diff line change
Expand Up @@ -256,6 +256,7 @@ type ConfigStore interface {

// Model config CRUD
GetModelConfigs(ctx context.Context) ([]tables.TableModelConfig, error)
GetModelConfigsByScopeAndScopeIDs(ctx context.Context, scope string, scopeIDs []string) ([]tables.TableModelConfig, error)
GetProviderGovernanceModelConfigs(ctx context.Context) ([]tables.TableModelConfig, error)
GetModelConfigsPaginated(ctx context.Context, params ModelConfigsQueryParams) ([]tables.TableModelConfig, int64, error)
GetModelConfig(ctx context.Context, scope string, scopeID *string, modelName string, provider *string) (*tables.TableModelConfig, error)
Expand Down
32 changes: 28 additions & 4 deletions framework/configstore/tables/modelconfig.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package tables
import (
"fmt"
"strings"
"sync"
"time"

"gorm.io/gorm"
Expand All @@ -12,21 +13,44 @@ import (
const (
ModelConfigScopeGlobal = "global"
ModelConfigScopeVirtualKey = "virtual_key"
ModelConfigScopeUser = "user"
)

// ModelConfigAllModels is the model_name sentinel meaning "all models". Combined with a
// specific provider it expresses provider-level governance (all models on that provider);
// with a nil provider it means all models on all providers.
const ModelConfigAllModels = "*"

// validModelConfigScopes is the set of accepted scope values.
var validModelConfigScopes = map[string]bool{
ModelConfigScopeGlobal: true,
ModelConfigScopeVirtualKey: true,
// validModelConfigScopes is the runtime registry of accepted scope values.
// OSS seeds it with global + virtual_key; downstream consumers (e.g. the
// enterprise build registering "user") extend it at startup via
// RegisterModelConfigScope. Guarded by validModelConfigScopesMu.
var (
validModelConfigScopesMu sync.RWMutex
validModelConfigScopes = map[string]bool{
ModelConfigScopeGlobal: true,
ModelConfigScopeVirtualKey: true,
}
)

// RegisterModelConfigScope adds scope to the allow-list consulted by
// IsValidModelConfigScope and TableModelConfig.BeforeSave. Intended to be
// called once at process startup; safe to call concurrently. Whitespace-
// only input is ignored.
func RegisterModelConfigScope(scope string) {
s := strings.TrimSpace(scope)
if s == "" {
return
}
validModelConfigScopesMu.Lock()
validModelConfigScopes[s] = true
validModelConfigScopesMu.Unlock()
}

// IsValidModelConfigScope reports whether scope is a recognized model config scope.
func IsValidModelConfigScope(scope string) bool {
validModelConfigScopesMu.RLock()
defer validModelConfigScopesMu.RUnlock()
return validModelConfigScopes[scope]
}

Expand Down
32 changes: 16 additions & 16 deletions plugins/governance/modelprovidergovernance_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2247,7 +2247,7 @@ func TestStore_CheckVirtualKeyScopedModelBudget_NilVK(t *testing.T) {
store, err := NewLocalGovernanceStore(context.Background(), logger, nil, &configstore.GovernanceConfig{}, nil)
require.NoError(t, err)

decision, err := store.CheckVirtualKeyScopedModelBudget(context.Background(), nil, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
decision, err := store.CheckScopedModelBudget(context.Background(), "", "", &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
assert.NoError(t, err)
assert.Equal(t, DecisionAllow, decision)
}
Expand All @@ -2258,7 +2258,7 @@ func TestStore_CheckVirtualKeyScopedModelBudget_NoConfig(t *testing.T) {
require.NoError(t, err)

vk := buildVirtualKey("vk1", "vk1-value", "vk1", true)
decision, err := store.CheckVirtualKeyScopedModelBudget(context.Background(), vk, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
decision, err := store.CheckScopedModelBudget(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
assert.NoError(t, err)
assert.Equal(t, DecisionAllow, decision)
}
Expand All @@ -2274,7 +2274,7 @@ func TestStore_CheckVirtualKeyScopedModelBudget_WithinLimit(t *testing.T) {
}, nil)
require.NoError(t, err)

_, err = store.CheckVirtualKeyScopedModelBudget(context.Background(), vk, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
_, err = store.CheckScopedModelBudget(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
assert.NoError(t, err, "Should allow when per-VK model budget is within limit")
}

Expand All @@ -2289,7 +2289,7 @@ func TestStore_CheckVirtualKeyScopedModelBudget_Exceeded(t *testing.T) {
}, nil)
require.NoError(t, err)

_, err = store.CheckVirtualKeyScopedModelBudget(context.Background(), vk, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
_, err = store.CheckScopedModelBudget(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
assert.Error(t, err, "Should reject when per-VK model budget is exceeded")
assert.Contains(t, err.Error(), "budget exceeded")
}
Expand All @@ -2307,7 +2307,7 @@ func TestStore_CheckVirtualKeyScopedModelBudget_OnlyAppliesToMatchingVK(t *testi
require.NoError(t, err)

// A request made with a DIFFERENT virtual key must not be affected by vk1's scoped config.
decision, err := store.CheckVirtualKeyScopedModelBudget(context.Background(), otherVK, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
decision, err := store.CheckScopedModelBudget(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, otherVK.ID, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
assert.NoError(t, err)
assert.Equal(t, DecisionAllow, decision)
}
Expand All @@ -2326,7 +2326,7 @@ func TestStore_CheckVirtualKeyScopedModelBudget_IgnoresGlobalConfig(t *testing.T
}, nil)
require.NoError(t, err)

decision, err := store.CheckVirtualKeyScopedModelBudget(context.Background(), vk, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
decision, err := store.CheckScopedModelBudget(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
assert.NoError(t, err, "Scoped check must not pick up the global config")
assert.Equal(t, DecisionAllow, decision)

Expand All @@ -2346,7 +2346,7 @@ func TestStore_CheckVirtualKeyScopedModelRateLimit_TokenLimitExceeded(t *testing
}, nil)
require.NoError(t, err)

decision, err := store.CheckVirtualKeyScopedModelRateLimit(context.Background(), vk, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil, nil)
decision, err := store.CheckScopedModelRateLimit(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil, nil)
assert.Error(t, err, "Should reject when per-VK model token limit is exceeded")
assert.Equal(t, DecisionTokenLimited, decision)
}
Expand All @@ -2362,7 +2362,7 @@ func TestStore_CheckVirtualKeyScopedModelRateLimit_WithinLimit(t *testing.T) {
}, nil)
require.NoError(t, err)

decision, err := store.CheckVirtualKeyScopedModelRateLimit(context.Background(), vk, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil, nil)
decision, err := store.CheckScopedModelRateLimit(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil, nil)
assert.NoError(t, err)
assert.Equal(t, DecisionAllow, decision)
}
Expand All @@ -2384,16 +2384,16 @@ func TestStore_VirtualKeyScopedModel_RecordThenCheck_TokenLimitTrips(t *testing.
req := &EvaluationRequest{Model: "claude-opus-4-7", Provider: schemas.Anthropic}

// Initially within limit.
decision, err := store.CheckVirtualKeyScopedModelRateLimit(context.Background(), vk, req, nil, nil)
decision, err := store.CheckScopedModelRateLimit(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, req, nil, nil)
require.NoError(t, err)
require.Equal(t, DecisionAllow, decision)

// Record usage above the limit (what the tracker does post-response). Provider differs from
// the config's (which is all-providers), exercising the model-only scoped lookup.
require.NoError(t, store.UpdateVirtualKeyScopedModelRateLimitUsageInMemory(context.Background(), vk, "claude-opus-4-7", schemas.Anthropic, 150, true, true))
require.NoError(t, store.UpdateScopedModelRateLimitUsageInMemory(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, "claude-opus-4-7", schemas.Anthropic, 150, true, true))

// Now the scoped check must trip.
decision, err = store.CheckVirtualKeyScopedModelRateLimit(context.Background(), vk, req, nil, nil)
decision, err = store.CheckScopedModelRateLimit(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, req, nil, nil)
assert.Error(t, err)
assert.Equal(t, DecisionTokenLimited, decision)
}
Expand All @@ -2411,13 +2411,13 @@ func TestStore_VirtualKeyScopedModel_RecordThenCheck_BudgetTrips(t *testing.T) {

req := &EvaluationRequest{Model: "claude-opus-4-7", Provider: schemas.Anthropic}

decision, err := store.CheckVirtualKeyScopedModelBudget(context.Background(), vk, req, nil)
decision, err := store.CheckScopedModelBudget(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, req, nil)
require.NoError(t, err)
require.Equal(t, DecisionAllow, decision)

require.NoError(t, store.UpdateVirtualKeyScopedModelBudgetUsageInMemory(context.Background(), vk, "claude-opus-4-7", schemas.Anthropic, 15.0))
require.NoError(t, store.UpdateScopedModelBudgetUsageInMemory(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, "claude-opus-4-7", schemas.Anthropic, 15.0))

_, err = store.CheckVirtualKeyScopedModelBudget(context.Background(), vk, req, nil)
_, err = store.CheckScopedModelBudget(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, req, nil)
assert.Error(t, err, "scoped budget should trip once usage exceeds the cap")
}

Expand All @@ -2440,7 +2440,7 @@ func TestStore_VKGovernanceBudget_NoDoubleCount(t *testing.T) {
require.NoError(t, err)

// Mirror tracker.UpdateUsage: scoped-model path + hierarchy path, same request/cost.
require.NoError(t, store.UpdateVirtualKeyScopedModelBudgetUsageInMemory(context.Background(), vk, "gpt-4", schemas.OpenAI, 10.0))
require.NoError(t, store.UpdateScopedModelBudgetUsageInMemory(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, "gpt-4", schemas.OpenAI, 10.0))
require.NoError(t, store.UpdateVirtualKeyBudgetUsageInMemory(context.Background(), vk, schemas.OpenAI, 10.0))

b := store.LoadBudget(context.Background(), "vkb")
Expand Down Expand Up @@ -2473,6 +2473,6 @@ func TestStore_CheckVirtualKeyScopedModelBudget_MultiBudget_OneExceededBlocks(t
}, nil)
require.NoError(t, err)

_, err = store.CheckVirtualKeyScopedModelBudget(context.Background(), vk, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
_, err = store.CheckScopedModelBudget(context.Background(), configstoreTables.ModelConfigScopeVirtualKey, vk.ID, &EvaluationRequest{Model: "gpt-4", Provider: schemas.OpenAI}, nil)
assert.Error(t, err, "an exceeded budget among several on a VK-scoped config must block")
}
Loading
Loading