From f2676e8a186e526a54a8c78057b4d8a2bfde322b Mon Sep 17 00:00:00 2001 From: Chenyme <118253778+chenyme@users.noreply.github.com> Date: Tue, 28 Jul 2026 14:43:53 +0800 Subject: [PATCH] feat: add client key routing scopes with fail-closed enforcement --- backend/internal/app/application.go | 6 +- .../internal/application/clientkey/cache.go | 6 + .../internal/application/clientkey/service.go | 47 +++- .../application/clientkey/service_test.go | 80 ++++++ backend/internal/application/gateway/image.go | 10 +- .../internal/application/gateway/selector.go | 108 +++++++- .../gateway/selector_layered_test.go | 14 + .../application/gateway/selector_test.go | 87 +++++++ .../internal/application/gateway/service.go | 29 ++- .../application/gateway/service_test.go | 41 ++- backend/internal/application/gateway/video.go | 2 +- .../application/invalidation/service.go | 11 +- .../application/invalidation/service_test.go | 9 + backend/internal/application/model/service.go | 61 ++++- .../internal/domain/clientkey/client_key.go | 229 +++++++++++++++- .../domain/clientkey/client_key_test.go | 59 +++++ .../relational/client_key_repository.go | 36 ++- .../relational/invalidation_test.go | 56 ++++ .../infra/persistence/relational/mapping.go | 6 +- .../relational/model_repository.go | 131 ++++++++++ .../relational/model_scope_filter_test.go | 156 +++++++++++ .../infra/persistence/relational/models.go | 4 +- .../relational/postgres_integration_test.go | 67 +++++ .../infra/persistence/relational/schema.go | 24 ++ ...ma_client_key_account_pool_upgrade_test.go | 91 +++++++ .../runtime/redis/store_integration_test.go | 12 + backend/internal/repository/list.go | 7 +- backend/internal/repository/model.go | 1 + backend/internal/repository/runtime.go | 17 +- .../transport/http/clientkey/handler.go | 84 +++++- .../transport/http/clientkey/handler_test.go | 37 +++ .../transport/http/inference/handler.go | 70 +++-- .../transport/http/inference/handler_test.go | 31 +++ .../http/inference/model_list_test.go | 14 + .../internal/transport/http/model/handler.go | 35 ++- .../transport/http/model/handler_test.go | 19 ++ frontend/src/entities/model/model-api.ts | 6 + .../features/client-keys/client-keys-api.ts | 13 +- .../features/client-keys/client-keys-page.tsx | 246 +++++++++++++++--- .../creative-console-page.tsx | 13 +- frontend/src/shared/i18n/index.ts | 14 +- 41 files changed, 1866 insertions(+), 123 deletions(-) create mode 100644 backend/internal/domain/clientkey/client_key_test.go create mode 100644 backend/internal/infra/persistence/relational/model_scope_filter_test.go create mode 100644 backend/internal/infra/persistence/relational/schema_client_key_account_pool_upgrade_test.go diff --git a/backend/internal/app/application.go b/backend/internal/app/application.go index 0b30da675..17adc5967 100644 --- a/backend/internal/app/application.go +++ b/backend/internal/app/application.go @@ -300,9 +300,13 @@ func New(ctx context.Context, cfg config.Config, logger *slog.Logger) (*Applicat selector := gateway.NewSelector(accountRepo, concurrency, sticky, providers, cfg.Routing.StickyTTL.Value(), cfg.Routing.CooldownBase.Value(), cfg.Routing.CooldownMax.Value(), cfg.Routing.CapacityWait.Value()) selector.UpdatePreferFreeBuild(cfg.Routing.PreferFreeBuild) selector.UpdateSegmentedSelector(cfg.Routing.SegmentedSelectorEnabled, cfg.Routing.SegmentedMinCandidates, cfg.Routing.SegmentedWindowSize) - invalidationService := invalidationapp.NewService(invalidationBus, invalidationSourceInstance(cfg), selector.ApplyInvalidation, logger) + invalidationService := invalidationapp.NewService(invalidationBus, invalidationSourceInstance(cfg), func(event repository.InvalidationEvent) { + selector.ApplyInvalidation(event) + clientKeyService.ApplyInvalidation(event) + }, logger) accountRepo.SetInvalidationObserver(invalidationService.Notify) modelRepo.SetInvalidationObserver(invalidationService.Notify) + clientKeyRepo.SetInvalidationObserver(invalidationService.Notify) gatewayService := gateway.NewService(modelService, auditService, accountService, clientKeyService, providers, selector, responseRepo, cfg.Routing.MaxAttempts) gatewayService.SetLogger(logger) gatewayService.UpdateBuildForbiddenReauthPolicy(cfg.Accounts.MarkBuildForbiddenReauth, cfg.Accounts.BuildForbiddenReauthCodes) diff --git a/backend/internal/application/clientkey/cache.go b/backend/internal/application/clientkey/cache.go index c838e2701..4ed4d71f8 100644 --- a/backend/internal/application/clientkey/cache.go +++ b/backend/internal/application/clientkey/cache.go @@ -94,6 +94,12 @@ func (c *authKeyCache) deleteIDs(ids []uint64) { } } +func (c *authKeyCache) clear() { + c.mu.Lock() + clear(c.byPrefix) + c.mu.Unlock() +} + // touchTracker 合并非关键的最近使用时间写入。 type touchTracker struct { mu sync.Mutex diff --git a/backend/internal/application/clientkey/service.go b/backend/internal/application/clientkey/service.go index 735f0e05f..bfa2931f6 100644 --- a/backend/internal/application/clientkey/service.go +++ b/backend/internal/application/clientkey/service.go @@ -41,6 +41,8 @@ type CreateInput struct { BillingLimitUSDTicks int64 AllowModelAliases bool AllowedModels []uint64 + ProviderScope clientkeydomain.ProviderScope + TierScope clientkeydomain.TierScope } type UpdateInput struct { @@ -53,6 +55,8 @@ type UpdateInput struct { BillingLimitUSDTicks *int64 AllowModelAliases *bool AllowedModels *[]uint64 + ProviderScope *clientkeydomain.ProviderScope + TierScope *clientkeydomain.TierScope } type Created struct { @@ -97,6 +101,19 @@ func (s *Service) UpdateDefaults(defaultRPM, defaultMax int) { s.defaultMax.Store(int64(defaultMax)) } +// ApplyInvalidation removes cached authorization policy after a local or remote +// client-key mutation. A zero ID represents a batch-wide invalidation. +func (s *Service) ApplyInvalidation(event repository.InvalidationEvent) { + if event.Kind != repository.InvalidationClientKeyChanged { + return + } + if event.ClientKeyID == 0 { + s.authCache.clear() + return + } + s.authCache.deleteID(event.ClientKeyID) +} + func (s *Service) List(ctx context.Context, page, pageSize int, search string, filter ListFilter) ([]clientkeydomain.Key, int64, error) { page, pageSize = normalizePage(page, pageSize) if !validListFilter(filter.Status, "", "active", "disabled", "expired") || !validListFilter(filter.ModelScope, "", "all", "restricted") || !repository.IsValidSort(filter.Sort, "name", "prefix", "status", "rpmLimit", "maxConcurrent", "billingLimit", "expiresAt", "lastUsedAt") { @@ -131,6 +148,11 @@ func (s *Service) Create(ctx context.Context, input CreateInput) (Created, error if input.BillingLimitUSDTicks < 0 || input.BillingLimitUSDTicks > clientkeydomain.MaxBillingLimitTicks { return Created{}, invalidInput("billingLimitUsdTicks 超出允许范围") } + providerScope, providerScopeValid := clientkeydomain.NormalizeProviderScope(input.ProviderScope) + tierScope, tierScopeValid := clientkeydomain.NormalizeTierScope(input.TierScope) + if !providerScopeValid || !tierScopeValid { + return Created{}, invalidInput("providerScope 或 tierScope 无效") + } prefix, err := security.NewHexToken(6) if err != nil { return Created{}, err @@ -164,6 +186,7 @@ func (s *Service) Create(ctx context.Context, input CreateInput) (Created, error Name: strings.TrimSpace(input.Name), Prefix: prefix, SecretHash: security.HashToken(raw), EncryptedSecret: encryptedSecret, Enabled: input.Enabled, ExpiresAt: input.ExpiresAt, RPMLimit: input.RPMLimit, MaxConcurrent: input.MaxConcurrent, BillingLimitUSDTicks: input.BillingLimitUSDTicks, AllowModelAliases: input.AllowModelAliases, AllowedModels: input.AllowedModels, + ProviderScope: providerScope, TierScope: tierScope, }) return Created{Key: value, Secret: raw}, mapRepositoryError(err) } @@ -231,6 +254,20 @@ func (s *Service) Update(ctx context.Context, id uint64, input UpdateInput) (cli if input.AllowedModels != nil { value.AllowedModels = *input.AllowedModels } + if input.ProviderScope != nil { + providerScope, valid := clientkeydomain.NormalizeProviderScope(*input.ProviderScope) + if !valid { + return clientkeydomain.Key{}, invalidInput("providerScope 无效") + } + value.ProviderScope = providerScope + } + if input.TierScope != nil { + tierScope, valid := clientkeydomain.NormalizeTierScope(*input.TierScope) + if !valid { + return clientkeydomain.Key{}, invalidInput("tierScope 无效") + } + value.TierScope = tierScope + } updated, err := s.keys.Update(ctx, value) if err == nil { s.authCache.deleteID(id) @@ -336,15 +373,7 @@ func (s *Service) Authenticate(ctx context.Context, raw string) (clientkeydomain // CanUseModel 判断空权限列表代表全部模型,否则要求显式授权。 func (s *Service) CanUseModel(value clientkeydomain.Key, modelID uint64) bool { - if len(value.AllowedModels) == 0 { - return true - } - for _, allowed := range value.AllowedModels { - if allowed == modelID { - return true - } - } - return false + return value.AllowsModel(modelID) } // ReserveBilling 为有限额 Key 原子预留本次请求的预计费用。 diff --git a/backend/internal/application/clientkey/service_test.go b/backend/internal/application/clientkey/service_test.go index 03e09c4a8..9316b3d03 100644 --- a/backend/internal/application/clientkey/service_test.go +++ b/backend/internal/application/clientkey/service_test.go @@ -243,6 +243,86 @@ func TestAuthenticateCachesUnlimitedKeyAndInvalidatesOnDisable(t *testing.T) { } } +func TestAccountScopePersistsAndAuthCacheInvalidatesOnChange(t *testing.T) { + ctx := context.Background() + database, err := relational.OpenSQLite(ctx, filepath.Join(t.TempDir(), "account-pool-auth-cache.db")) + if err != nil { + t.Fatal(err) + } + defer database.Close() + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + base := relational.NewClientKeyRepository(database) + service := NewService(base, successfulRateLimiter{}, successfulConcurrencyLimiter{}, 60, 5, testCipher(t)) + created, err := service.Create(ctx, CreateInput{Name: "scoped", Enabled: true, ProviderScope: clientkeydomain.ProviderScopeBuild | clientkeydomain.ProviderScopeWeb, TierScope: clientkeydomain.TierScopeFree}) + if err != nil { + t.Fatal(err) + } + value, release, err := service.Authenticate(ctx, created.Secret) + if err != nil { + t.Fatal(err) + } + release() + if value.ProviderScope != clientkeydomain.ProviderScopeBuild|clientkeydomain.ProviderScopeWeb || value.TierScope != clientkeydomain.TierScopeFree { + t.Fatalf("authenticated account scope = %+v", value.AccountScope()) + } + consoleScope := clientkeydomain.ProviderScopeConsole + superTier := clientkeydomain.TierScopeSuper + if _, err := service.Update(ctx, created.Key.ID, UpdateInput{ProviderScope: &consoleScope, TierScope: &superTier}); err != nil { + t.Fatal(err) + } + value, release, err = service.Authenticate(ctx, created.Secret) + if err != nil { + t.Fatal(err) + } + release() + if value.ProviderScope != clientkeydomain.ProviderScopeConsole || value.TierScope != clientkeydomain.TierScopeSuper { + t.Fatalf("account scope after cache invalidation = %+v", value.AccountScope()) + } + stored, err := base.Get(ctx, created.Key.ID) + if err != nil { + t.Fatal(err) + } + stored.ProviderScope = clientkeydomain.ProviderScopeWeb + stored.TierScope = clientkeydomain.TierScopeFree + if _, err := base.Update(ctx, stored); err != nil { + t.Fatal(err) + } + value, release, err = service.Authenticate(ctx, created.Secret) + if err != nil { + t.Fatal(err) + } + release() + if value.ProviderScope != clientkeydomain.ProviderScopeConsole || value.TierScope != clientkeydomain.TierScopeSuper { + t.Fatalf("cache unexpectedly changed before remote invalidation = %+v", value.AccountScope()) + } + service.ApplyInvalidation(repository.InvalidationEvent{Kind: repository.InvalidationClientKeyChanged, ClientKeyID: created.Key.ID}) + value, release, err = service.Authenticate(ctx, created.Secret) + if err != nil { + t.Fatal(err) + } + release() + if value.ProviderScope != clientkeydomain.ProviderScopeWeb || value.TierScope != clientkeydomain.TierScopeFree { + t.Fatalf("account scope after remote invalidation = %+v", value.AccountScope()) + } + service.ApplyInvalidation(repository.InvalidationEvent{Kind: repository.InvalidationClientKeyChanged}) + service.authCache.mu.RLock() + cacheEntries := len(service.authCache.byPrefix) + service.authCache.mu.RUnlock() + if cacheEntries != 0 { + t.Fatalf("batch invalidation retained %d auth cache entries", cacheEntries) + } + invalid := clientkeydomain.ProviderScope(8) + if _, err := service.Update(ctx, created.Key.ID, UpdateInput{ProviderScope: &invalid}); !errors.Is(err, ErrInvalidInput) { + t.Fatalf("invalid account scope error = %v", err) + } + all, err := service.Create(ctx, CreateInput{Name: "legacy-default", Enabled: true}) + if err != nil || all.Key.ProviderScope != clientkeydomain.ProviderScopeAll || all.Key.TierScope != clientkeydomain.TierScopeAll { + t.Fatalf("default account scope = %+v, err = %v", all.Key.AccountScope(), err) + } +} + func testCipher(t *testing.T) *security.Cipher { t.Helper() cipher, err := security.NewCipher(base64.StdEncoding.EncodeToString(make([]byte, 32))) diff --git a/backend/internal/application/gateway/image.go b/backend/internal/application/gateway/image.go index 4ea0984eb..f70e619bc 100644 --- a/backend/internal/application/gateway/image.go +++ b/backend/internal/application/gateway/image.go @@ -2,6 +2,7 @@ package gateway import ( "context" + "errors" "fmt" "net/http" "sync" @@ -173,9 +174,14 @@ func (s *Service) executeImage( var lastCredentialFailure *accountdomain.Credential var lastCredentialError error for attempt := 0; attemptPolicy.allows(attempt); attempt++ { - lease, err = s.selector.Acquire(ctx, route.Provider, route.ID, route.UpstreamModel, quotaMode, "", excluded, false) + lease, err = s.selector.AcquireForKey(ctx, route.Provider, route.ID, route.UpstreamModel, quotaMode, "", excluded, false, key.AccountScope()) if err != nil { - writeFailureAudit(http.StatusServiceUnavailable, "upstream_unavailable", lastCredentialFailure) + errorCode := "upstream_unavailable" + var selectionFailure *SelectionUnavailableError + if errors.As(err, &selectionFailure) { + errorCode = selectionFailure.Code() + } + writeFailureAudit(http.StatusServiceUnavailable, errorCode, lastCredentialFailure) return nil, fmt.Errorf("%w: %w", ErrNoAvailableAccount, err) } excluded[lease.Credential.ID] = true diff --git a/backend/internal/application/gateway/selector.go b/backend/internal/application/gateway/selector.go index 0383c1342..143ddd5ee 100644 --- a/backend/internal/application/gateway/selector.go +++ b/backend/internal/application/gateway/selector.go @@ -12,6 +12,7 @@ import ( "time" "github.com/chenyme/grok2api/backend/internal/domain/account" + clientkeydomain "github.com/chenyme/grok2api/backend/internal/domain/clientkey" "github.com/chenyme/grok2api/backend/internal/pkg/resultcache" "github.com/chenyme/grok2api/backend/internal/repository" "golang.org/x/sync/singleflight" @@ -109,24 +110,47 @@ const ( type SelectionUnavailableError struct { Reason SelectionUnavailableReason RetryAfter time.Duration + Scope clientkeydomain.AccountScope } func (e *SelectionUnavailableError) Error() string { if e == nil { return "没有可用上游账号" } + prefix := "" + if e.Scope.IsRestricted() { + prefix = "Client Key 限定范围" + } switch e.Reason { case SelectionUnsupportedModel: + if prefix != "" { + return prefix + "不支持该模型" + } return "当前账号池不支持该模型" case SelectionCooling: + if prefix != "" { + return prefix + "中的可用账号正在冷却" + } return "可用上游账号正在冷却" case SelectionModelCooling: + if prefix != "" { + return prefix + "中可用账号的目标模型正在冷却" + } return "可用上游账号的目标模型正在冷却" case SelectionQuotaExhausted: + if prefix != "" { + return prefix + "中的可用账号额度等待恢复" + } return "可用上游账号额度等待恢复" case SelectionSaturated: + if prefix != "" { + return prefix + "中的可用账号均达到并发上限" + } return "可用上游账号均达到并发上限" default: + if prefix != "" { + return prefix + "当前没有可用上游账号" + } return "没有可用上游账号" } } @@ -156,6 +180,10 @@ func (e *SelectionUnavailableError) Code() string { return "upstream_saturated" case SelectionUnsupportedModel: return "upstream_model_unavailable" + case SelectionNoAccounts: + if e.Scope.IsRestricted() { + return "client_key_account_scope_unavailable" + } } } return "upstream_unavailable" @@ -270,6 +298,19 @@ func (s *Selector) preferFreeBuildEnabled() bool { } func (s *Selector) Acquire(ctx context.Context, provider account.Provider, modelRouteID uint64, upstreamModel, quotaMode, affinityKey string, excluded map[uint64]bool, allowQuotaProbe bool) (*accountLease, error) { + return s.acquire(ctx, provider, modelRouteID, upstreamModel, quotaMode, affinityKey, excluded, allowQuotaProbe, clientkeydomain.AccountScope{}) +} + +func (s *Selector) AcquireForKey(ctx context.Context, provider account.Provider, modelRouteID uint64, upstreamModel, quotaMode, affinityKey string, excluded map[uint64]bool, allowQuotaProbe bool, scope clientkeydomain.AccountScope) (*accountLease, error) { + return s.acquire(ctx, provider, modelRouteID, upstreamModel, quotaMode, affinityKey, excluded, allowQuotaProbe, scope) +} + +func (s *Selector) acquire(ctx context.Context, provider account.Provider, modelRouteID uint64, upstreamModel, quotaMode, affinityKey string, excluded map[uint64]bool, allowQuotaProbe bool, requestedScope clientkeydomain.AccountScope) (lease *accountLease, err error) { + accountScope, scopeValid := clientkeydomain.NormalizeAccountScope(requestedScope) + defer annotateSelectionAccountScope(&err, accountScope) + if !scopeValid || !accountScope.AllowsProvider(provider) { + return nil, &SelectionUnavailableError{Reason: SelectionNoAccounts, Scope: accountScope} + } now := time.Now().UTC() stickyKey := stickySessionKey(affinityKey) values, err := s.loadCandidates(ctx, provider, modelRouteID, upstreamModel, quotaMode, now) @@ -287,6 +328,9 @@ func (s *Selector) Acquire(ctx context.Context, provider account.Provider, model var earliestRetry time.Time for index, candidate := range values { value := candidate.Credential + if !accountScopeAllowsCandidate(provider, accountScope, candidate) { + continue + } if excluded[value.ID] || value.AuthStatus != account.AuthStatusActive { continue } @@ -552,6 +596,19 @@ func isSelectionUnavailable(err error, reason SelectionUnavailableReason) bool { // AcquirePinned 为 previous_response_id 等账号归属请求获取指定账号租约。 func (s *Selector) AcquirePinned(ctx context.Context, provider account.Provider, accountID, modelRouteID uint64, upstreamModel, quotaMode string, inference bool) (*accountLease, error) { + return s.acquirePinned(ctx, provider, accountID, modelRouteID, upstreamModel, quotaMode, inference, clientkeydomain.AccountScope{}) +} + +func (s *Selector) AcquirePinnedForKey(ctx context.Context, provider account.Provider, accountID, modelRouteID uint64, upstreamModel, quotaMode string, inference bool, scope clientkeydomain.AccountScope) (*accountLease, error) { + return s.acquirePinned(ctx, provider, accountID, modelRouteID, upstreamModel, quotaMode, inference, scope) +} + +func (s *Selector) acquirePinned(ctx context.Context, provider account.Provider, accountID, modelRouteID uint64, upstreamModel, quotaMode string, inference bool, requestedScope clientkeydomain.AccountScope) (lease *accountLease, err error) { + accountScope, scopeValid := clientkeydomain.NormalizeAccountScope(requestedScope) + defer annotateSelectionAccountScope(&err, accountScope) + if !scopeValid || !accountScope.AllowsProvider(provider) { + return nil, &SelectionUnavailableError{Reason: SelectionNoAccounts, Scope: accountScope} + } now := time.Now().UTC() values, err := s.loadCandidates(ctx, provider, modelRouteID, upstreamModel, quotaMode, now) if err != nil { @@ -562,6 +619,9 @@ func (s *Selector) AcquirePinned(ctx context.Context, provider account.Provider, if value.ID != accountID { continue } + if !accountScopeAllowsCandidate(provider, accountScope, candidate) { + return nil, &SelectionUnavailableError{Reason: SelectionNoAccounts} + } if !value.Enabled || value.AuthStatus != account.AuthStatusActive { return nil, &SelectionUnavailableError{Reason: SelectionNoAccounts} } @@ -628,6 +688,46 @@ func (s *Selector) AcquirePinned(ctx context.Context, provider account.Provider, return nil, &SelectionUnavailableError{Reason: SelectionNoAccounts} } +func accountScopeAllowsCandidate(provider account.Provider, scope clientkeydomain.AccountScope, candidate account.RoutingCandidate) bool { + if provider == account.ProviderConsole { + return true + } + tier := clientkeydomain.AccountTierUnknown + switch provider { + case account.ProviderBuild: + if candidate.IsKnownFreeBuild() { + tier = clientkeydomain.AccountTierFree + } else if account.IsBuildSuper(candidate.Credential, candidate.Billing) { + tier = clientkeydomain.AccountTierSuper + } + case account.ProviderWeb: + switch candidate.Credential.WebTier { + case account.WebTierBasic: + tier = clientkeydomain.AccountTierFree + case account.WebTierSuper, account.WebTierHeavy: + tier = clientkeydomain.AccountTierSuper + } + } + switch tier { + case clientkeydomain.AccountTierFree: + return scope.Tiers&clientkeydomain.TierScopeFree != 0 + case clientkeydomain.AccountTierSuper: + return scope.Tiers&clientkeydomain.TierScopeSuper != 0 + default: + return scope.Tiers&clientkeydomain.TierScopeUnknown != 0 + } +} + +func annotateSelectionAccountScope(err *error, scope clientkeydomain.AccountScope) { + if err == nil || *err == nil || !scope.IsRestricted() { + return + } + var unavailable *SelectionUnavailableError + if errors.As(*err, &unavailable) { + unavailable.Scope = scope + } +} + func effectiveQuotaMode(candidate account.RoutingCandidate, fallback string) string { if candidate.QuotaWindow != nil && candidate.QuotaWindow.Mode == "weekly" { return "weekly" @@ -1031,6 +1131,10 @@ func (s *Selector) ApplyInvalidation(event repository.InvalidationEvent) { if !event.Valid() { return } + layer := event.Layer() + if layer != repository.InvalidationLayerRoute && layer != repository.InvalidationLayerBase && layer != repository.InvalidationLayerOverlay { + return + } s.candidateMu.Lock() provider := event.Provider if provider == "" && event.AccountID != 0 { @@ -1044,8 +1148,8 @@ func (s *Selector) ApplyInvalidation(event repository.InvalidationEvent) { } } } - base := event.Layer() == repository.InvalidationLayerBase - overlay := event.Layer() == repository.InvalidationLayerOverlay || event.Layer() == repository.InvalidationLayerRoute + base := layer == repository.InvalidationLayerBase + overlay := layer == repository.InvalidationLayerOverlay || layer == repository.InvalidationLayerRoute if base { if provider == "" { s.baseGlobalVersion++ diff --git a/backend/internal/application/gateway/selector_layered_test.go b/backend/internal/application/gateway/selector_layered_test.go index 8a341cd62..a5bf4df73 100644 --- a/backend/internal/application/gateway/selector_layered_test.go +++ b/backend/internal/application/gateway/selector_layered_test.go @@ -197,6 +197,20 @@ func TestSelectorAppliesOutOfOrderInvalidationsSafely(t *testing.T) { } } +func TestSelectorIgnoresClientKeyInvalidation(t *testing.T) { + selector := NewSelector(nil, nil, nil, nil, time.Hour, time.Second, time.Minute) + expiresAt := time.Now().Add(time.Hour) + selector.routingBases[routingBaseCacheKey{provider: account.ProviderBuild}] = routingBaseSnapshot{expiresAt: expiresAt} + selector.routingOverlays[routingOverlayCacheKey{provider: account.ProviderBuild}] = routingOverlaySnapshot{expiresAt: expiresAt} + selector.candidates[candidateCacheKey{provider: account.ProviderBuild}] = candidateSnapshot{expiresAt: expiresAt} + + selector.ApplyInvalidation(repository.InvalidationEvent{Kind: repository.InvalidationClientKeyChanged, ClientKeyID: 42}) + + if len(selector.routingBases) != 1 || len(selector.routingOverlays) != 1 || len(selector.candidates) != 1 { + t.Fatalf("client-key event invalidated selector caches: bases=%d overlays=%d candidates=%d", len(selector.routingBases), len(selector.routingOverlays), len(selector.candidates)) + } +} + func TestSelectorScopesAccountInvalidationToCachedProvider(t *testing.T) { selector := NewSelector(nil, nil, nil, nil, time.Hour, time.Second, time.Minute) expiresAt := time.Now().Add(time.Hour) diff --git a/backend/internal/application/gateway/selector_test.go b/backend/internal/application/gateway/selector_test.go index 155b42622..e735509e6 100644 --- a/backend/internal/application/gateway/selector_test.go +++ b/backend/internal/application/gateway/selector_test.go @@ -11,6 +11,7 @@ import ( "time" "github.com/chenyme/grok2api/backend/internal/domain/account" + clientkeydomain "github.com/chenyme/grok2api/backend/internal/domain/clientkey" "github.com/chenyme/grok2api/backend/internal/infra/persistence/relational" "github.com/chenyme/grok2api/backend/internal/infra/runtime/memory" "github.com/chenyme/grok2api/backend/internal/repository" @@ -535,6 +536,92 @@ func TestSelectorHonorsWebTierPoolOrderBeforeAccountPriority(t *testing.T) { } } +func TestSelectorEnforcesClientKeyAccountScopeAcrossProvidersAndTiers(t *testing.T) { + ctx := context.Background() + database, err := relational.OpenSQLite(ctx, filepath.Join(t.TempDir(), "selector-client-key-pool.db")) + if err != nil { + t.Fatal(err) + } + defer database.Close() + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + accounts := relational.NewAccountRepository(database) + create := func(value account.Credential) account.Credential { + t.Helper() + created, _, createErr := accounts.UpsertByIdentity(ctx, value) + if createErr != nil { + t.Fatal(createErr) + } + return created + } + buildFree := create(account.Credential{Provider: account.ProviderBuild, Name: "build-free", SourceKey: "build-free", EncryptedAccessToken: "encrypted", Enabled: true, AuthStatus: account.AuthStatusActive, Priority: 10, MaxConcurrent: 2}) + buildSuper := create(account.Credential{Provider: account.ProviderBuild, Name: "build-super", SourceKey: "build-super", EncryptedAccessToken: "encrypted", Enabled: true, AuthStatus: account.AuthStatusActive, Priority: 20, MaxConcurrent: 2}) + buildUnknown := create(account.Credential{Provider: account.ProviderBuild, Name: "build-unknown", SourceKey: "build-unknown", EncryptedAccessToken: "encrypted", Enabled: true, AuthStatus: account.AuthStatusActive, Priority: 100, MaxConcurrent: 2}) + now := time.Now().UTC() + if err := accounts.SaveBilling(ctx, account.Billing{AccountID: buildFree.ID, PlanName: "Free", SyncedAt: now}); err != nil { + t.Fatal(err) + } + if err := accounts.SaveBilling(ctx, account.Billing{AccountID: buildSuper.ID, PlanName: "SuperGrok", SyncedAt: now}); err != nil { + t.Fatal(err) + } + + webFree := create(account.Credential{Provider: account.ProviderWeb, AuthType: account.AuthTypeSSO, WebTier: account.WebTierBasic, Name: "web-free", SourceKey: "web-free", EncryptedAccessToken: "encrypted", Enabled: true, AuthStatus: account.AuthStatusActive, Priority: 100, MaxConcurrent: 2}) + webSuper := create(account.Credential{Provider: account.ProviderWeb, AuthType: account.AuthTypeSSO, WebTier: account.WebTierSuper, Name: "web-super", SourceKey: "web-super", EncryptedAccessToken: "encrypted", Enabled: true, AuthStatus: account.AuthStatusActive, Priority: 20, MaxConcurrent: 2}) + webHeavy := create(account.Credential{Provider: account.ProviderWeb, AuthType: account.AuthTypeSSO, WebTier: account.WebTierHeavy, Name: "web-heavy", SourceKey: "web-heavy", EncryptedAccessToken: "encrypted", Enabled: true, AuthStatus: account.AuthStatusActive, Priority: 10, MaxConcurrent: 2}) + _ = create(account.Credential{Provider: account.ProviderWeb, AuthType: account.AuthTypeSSO, WebTier: account.WebTierAuto, Name: "web-unknown", SourceKey: "web-unknown", EncryptedAccessToken: "encrypted", Enabled: true, AuthStatus: account.AuthStatusActive, Priority: 200, MaxConcurrent: 2}) + console := create(account.Credential{Provider: account.ProviderConsole, AuthType: account.AuthTypeSSO, Name: "console", SourceKey: "console", EncryptedAccessToken: "encrypted", Enabled: true, AuthStatus: account.AuthStatusActive, Priority: 10, MaxConcurrent: 2}) + + selector := NewSelector(accounts, memory.NewConcurrencyLimiter(), memory.NewStickyStore(), staticTierOrder{order: []account.WebTier{account.WebTierBasic, account.WebTierSuper, account.WebTierHeavy}}, time.Hour, time.Second, time.Minute) + providerScope := func(provider account.Provider) clientkeydomain.ProviderScope { + switch provider { + case account.ProviderBuild: + return clientkeydomain.ProviderScopeBuild + case account.ProviderWeb: + return clientkeydomain.ProviderScopeWeb + default: + return clientkeydomain.ProviderScopeConsole + } + } + assertSelected := func(provider account.Provider, tiers clientkeydomain.TierScope, excluded map[uint64]bool, want uint64) { + t.Helper() + scope := clientkeydomain.AccountScope{Providers: providerScope(provider), Tiers: tiers} + lease, acquireErr := selector.AcquireForKey(ctx, provider, 0, "", "", "", excluded, false, scope) + if acquireErr != nil { + t.Fatal(acquireErr) + } + defer lease.Release() + if lease.Credential.ID != want { + t.Fatalf("provider %s tiers %d selected %d, want %d", provider, tiers, lease.Credential.ID, want) + } + } + assertSelected(account.ProviderBuild, clientkeydomain.TierScopeFree, nil, buildFree.ID) + assertSelected(account.ProviderBuild, clientkeydomain.TierScopeSuper, nil, buildSuper.ID) + assertSelected(account.ProviderWeb, clientkeydomain.TierScopeFree, nil, webFree.ID) + assertSelected(account.ProviderWeb, clientkeydomain.TierScopeSuper, nil, webSuper.ID) + assertSelected(account.ProviderWeb, clientkeydomain.TierScopeSuper, map[uint64]bool{webSuper.ID: true}, webHeavy.ID) + assertSelected(account.ProviderConsole, clientkeydomain.TierScopeFree, nil, console.ID) + + freeBuildScope := clientkeydomain.AccountScope{Providers: clientkeydomain.ProviderScopeBuild, Tiers: clientkeydomain.TierScopeFree} + _, err = selector.AcquireForKey(ctx, account.ProviderBuild, 0, "", "", "", map[uint64]bool{buildFree.ID: true}, false, freeBuildScope) + var unavailable *SelectionUnavailableError + if !errors.As(err, &unavailable) || unavailable.Code() != "client_key_account_scope_unavailable" || unavailable.Scope != freeBuildScope { + t.Fatalf("scoped exhaustion error = %#v, err = %v", unavailable, err) + } + superBuildScope := clientkeydomain.AccountScope{Providers: clientkeydomain.ProviderScopeBuild, Tiers: clientkeydomain.TierScopeSuper} + if _, err := selector.AcquirePinnedForKey(ctx, account.ProviderBuild, buildUnknown.ID, 0, "", "", true, superBuildScope); !errors.As(err, &unavailable) || unavailable.Code() != "client_key_account_scope_unavailable" { + t.Fatalf("out-of-pool pinned error = %#v, err = %v", unavailable, err) + } + allKnownBuildTiers := clientkeydomain.AccountScope{Providers: clientkeydomain.ProviderScopeBuild, Tiers: clientkeydomain.TierScopeFree | clientkeydomain.TierScopeSuper} + if _, err := selector.AcquirePinnedForKey(ctx, account.ProviderBuild, buildUnknown.ID, 0, "", "", true, allKnownBuildTiers); !errors.As(err, &unavailable) { + t.Fatalf("unknown Build tier should be excluded: %v", err) + } + buildOnlyScope := clientkeydomain.AccountScope{Providers: clientkeydomain.ProviderScopeBuild, Tiers: clientkeydomain.TierScopeAll} + if _, err := selector.AcquireForKey(ctx, account.ProviderWeb, 0, "", "", "", nil, false, buildOnlyScope); !errors.As(err, &unavailable) || unavailable.Code() != "client_key_account_scope_unavailable" { + t.Fatalf("provider scope should fail closed: %#v, err = %v", unavailable, err) + } +} + func TestSelectorPropagatesConcurrencyStoreFailure(t *testing.T) { ctx := context.Background() database, err := relational.OpenSQLite(ctx, filepath.Join(t.TempDir(), "selector-runtime-error.db")) diff --git a/backend/internal/application/gateway/service.go b/backend/internal/application/gateway/service.go index b8c4eb0b7..b530ba5ce 100644 --- a/backend/internal/application/gateway/service.go +++ b/backend/internal/application/gateway/service.go @@ -510,7 +510,9 @@ func (s *Service) selectConversationRoute(routes []modeldomain.Route, key client return modeldomain.Route{}, ErrModelNotFound } fallback := routes[0] + accountScope := key.AccountScope() matchedOwnership := ownership == nil + scopeMatched := false allowed := false conversationSupported := false storedResponseUnsupported := false @@ -520,6 +522,10 @@ func (s *Service) selectConversationRoute(routes []modeldomain.Route, key client } matchedOwnership = true fallback = route + if !accountScope.AllowsProvider(route.Provider) { + continue + } + scopeMatched = true if !s.clientKeys.CanUseModel(key, route.ID) { continue } @@ -540,6 +546,9 @@ func (s *Service) selectConversationRoute(routes []modeldomain.Route, key client if !matchedOwnership { return fallback, ErrResponseAccountUnavailable } + if !scopeMatched { + return fallback, &SelectionUnavailableError{Reason: SelectionNoAccounts, Scope: accountScope} + } if !allowed { return fallback, clientkeyapp.ErrModelNotAllowed } @@ -558,7 +567,9 @@ func (s *Service) selectMediaRoute(routes []modeldomain.Route, key clientkey.Key return modeldomain.Route{}, ErrModelNotFound } fallback := routes[0] + accountScope := key.AccountScope() capabilityMatched := false + scopeMatched := false allowed := false for _, route := range routes { if route.Capability != capability { @@ -566,6 +577,10 @@ func (s *Service) selectMediaRoute(routes []modeldomain.Route, key clientkey.Key } fallback = route capabilityMatched = true + if !accountScope.AllowsProvider(route.Provider) { + continue + } + scopeMatched = true if !s.clientKeys.CanUseModel(key, route.ID) { continue } @@ -577,6 +592,9 @@ func (s *Service) selectMediaRoute(routes []modeldomain.Route, key clientkey.Key if !capabilityMatched { return fallback, ErrModelNotFound } + if !scopeMatched { + return fallback, &SelectionUnavailableError{Reason: SelectionNoAccounts, Scope: accountScope} + } if !allowed { return fallback, clientkeyapp.ErrModelNotAllowed } @@ -719,6 +737,7 @@ func (s *Service) createResponseAt(ctx context.Context, input Input, path string failureFingerprints := make(map[string]int) authRecoveryAttempted := make(map[uint64]bool) quotaMode := s.providers.QuotaMode(route.Provider, route.UpstreamModel) + accountScope := input.ClientKey.AccountScope() quotaProbeAttempted := false var lastErr error var lastFailure *UpstreamFailure @@ -746,9 +765,9 @@ attemptLoop: var err error selectionStarted := time.Now() if ownership != nil { - lease, err = s.selector.AcquirePinned(ctx, route.Provider, ownership.AccountID, route.ID, route.UpstreamModel, quotaMode, true) + lease, err = s.selector.AcquirePinnedForKey(ctx, route.Provider, ownership.AccountID, route.ID, route.UpstreamModel, quotaMode, true, accountScope) } else { - lease, err = s.selector.Acquire(ctx, route.Provider, route.ID, route.UpstreamModel, quotaMode, affinityKey, excluded, !quotaProbeAttempted) + lease, err = s.selector.AcquireForKey(ctx, route.Provider, route.ID, route.UpstreamModel, quotaMode, affinityKey, excluded, !quotaProbeAttempted, accountScope) } timing.markSelection(time.Since(selectionStarted)) if err != nil { @@ -1367,6 +1386,10 @@ func (s *Service) forwardOwnedResponse(ctx context.Context, input ResourceInput, _ = s.responses.Delete(ctx, input.ResponseID, input.ClientKey.ID) return nil, ErrResponseNotFound } + accountScope := input.ClientKey.AccountScope() + if !accountScope.AllowsProvider(ownership.Provider) { + return nil, &SelectionUnavailableError{Reason: SelectionNoAccounts, Scope: accountScope} + } adapter, ok := s.providers.Responses(ownership.Provider) if !ok { return nil, ErrResponseAccountUnavailable @@ -1376,7 +1399,7 @@ func (s *Service) forwardOwnedResponse(ctx context.Context, input ResourceInput, operation = "response_delete" } physicalCallCtx := infraegress.WithPhysicalCallTrace(ctx, string(ownership.Provider), operation) - lease, err := s.selector.AcquirePinned(ctx, ownership.Provider, ownership.AccountID, 0, "", "", false) + lease, err := s.selector.AcquirePinnedForKey(ctx, ownership.Provider, ownership.AccountID, 0, "", "", false, accountScope) if err != nil { return nil, fmt.Errorf("%w: %w", ErrResponseAccountUnavailable, err) } diff --git a/backend/internal/application/gateway/service_test.go b/backend/internal/application/gateway/service_test.go index 7fec6c941..aee1d6999 100644 --- a/backend/internal/application/gateway/service_test.go +++ b/backend/internal/application/gateway/service_test.go @@ -218,6 +218,16 @@ func TestGatewayFailsOverBeforeReturningBody(t *testing.T) { } adapter.resetAttempts() + blockedKey := clientKey + blockedKey.ProviderScope = clientkey.ProviderScopeConsole + if _, err := service.GetResponse(ctx, ResourceInput{ClientKey: blockedKey, ResponseID: "resp-test"}); err == nil { + t.Fatal("owned response should be rejected after its provider leaves the key scope") + } else { + var unavailable *SelectionUnavailableError + if !errors.As(err, &unavailable) || unavailable.Code() != "client_key_account_scope_unavailable" || len(adapter.attempts) != 0 { + t.Fatalf("scoped owned response error = %#v, attempts = %#v, err = %v", unavailable, adapter.attempts, err) + } + } resource, err := service.GetResponse(ctx, ResourceInput{ClientKey: clientKey, ResponseID: "resp-test", RawQuery: "include=reasoning.encrypted_content"}) if err != nil { t.Fatal(err) @@ -707,19 +717,38 @@ func TestGatewayTeamModelRateLimitOnlySkipsMatchingTeam(t *testing.T) { } func TestSelectConversationRouteRespectsClientKeyAcrossSharedPublicModel(t *testing.T) { - registry := provider.NewRegistry(&failoverAdapter{}, statelessConsoleAdapter{}) + registry := provider.NewRegistry(&failoverAdapter{}, webStoredResponseAdapter{}, statelessConsoleAdapter{}) service := &Service{ clientKeys: clientkeyapp.NewService(nil, nil, nil, 60, 4, nil), providers: registry, } routes := []modeldomain.Route{ {ID: 10, PublicID: "Build/grok-shared", Provider: account.ProviderBuild, UpstreamModel: "grok-shared"}, + {ID: 15, PublicID: "Web/grok-shared", Provider: account.ProviderWeb, UpstreamModel: "grok-shared"}, {ID: 20, PublicID: "Console/grok-shared", Provider: account.ProviderConsole, UpstreamModel: "grok-shared"}, } - selected, err := service.selectConversationRoute(routes, clientkey.Key{AllowedModels: []uint64{20}}, audit.OperationResponses, "/responses", false, nil) + selected, err := service.selectConversationRoute(routes, clientkey.Key{ProviderScope: clientkey.ProviderScopeWeb | clientkey.ProviderScopeConsole}, audit.OperationResponses, "/responses", false, nil) + if err != nil || selected.ID != 15 { + t.Fatalf("provider-scoped route = %#v, err = %v", selected, err) + } + selected, err = service.selectConversationRoute(routes, clientkey.Key{ProviderScope: clientkey.ProviderScopeConsole, AllowedModels: []uint64{20}}, audit.OperationResponses, "/responses", false, nil) if err != nil || selected.ID != 20 { t.Fatalf("selected route = %#v, err = %v", selected, err) } + _, err = service.selectConversationRoute(routes, clientkey.Key{ProviderScope: clientkey.ProviderScopeBuild, AllowedModels: []uint64{20}}, audit.OperationResponses, "/responses", false, nil) + if !errors.Is(err, clientkeyapp.ErrModelNotAllowed) { + t.Fatalf("scope and model intersection should reject the request: %v", err) + } + _, err = service.selectConversationRoute(routes[1:], clientkey.Key{ProviderScope: clientkey.ProviderScopeBuild}, audit.OperationResponses, "/responses", false, nil) + var unavailable *SelectionUnavailableError + if !errors.As(err, &unavailable) || unavailable.Code() != "client_key_account_scope_unavailable" { + t.Fatalf("provider scope must not fall back: %#v, err = %v", unavailable, err) + } + ownership := &inferencedomain.ResponseOwnership{Provider: account.ProviderWeb} + _, err = service.selectConversationRoute(routes, clientkey.Key{ProviderScope: clientkey.ProviderScopeBuild}, audit.OperationResponses, "/responses", true, ownership) + if !errors.As(err, &unavailable) || unavailable.Code() != "client_key_account_scope_unavailable" { + t.Fatalf("owned response must remain inside the updated provider scope: %#v, err = %v", unavailable, err) + } } func TestSelectMediaRouteSkipsSameNamedConversationRoute(t *testing.T) { @@ -739,6 +768,14 @@ func TestSelectMediaRouteSkipsSameNamedConversationRoute(t *testing.T) { if err != nil || selected.ID != 20 { t.Fatalf("selected route = %#v, err = %v", selected, err) } + _, err = service.selectMediaRoute(routes, clientkey.Key{ProviderScope: clientkey.ProviderScopeBuild}, modeldomain.CapabilityImage, func(providerValue account.Provider) bool { + _, ok := registry.ImageGeneration(providerValue) + return ok + }) + var unavailable *SelectionUnavailableError + if !errors.As(err, &unavailable) || unavailable.Code() != "client_key_account_scope_unavailable" { + t.Fatalf("media route must not leave provider scope: %#v, err = %v", unavailable, err) + } } func TestGenerateImageReturnsWhenEveryCredentialRefreshFails(t *testing.T) { diff --git a/backend/internal/application/gateway/video.go b/backend/internal/application/gateway/video.go index 7cfe6c314..3fadffbbe 100644 --- a/backend/internal/application/gateway/video.go +++ b/backend/internal/application/gateway/video.go @@ -67,7 +67,7 @@ func (s *Service) CreateVideo(ctx context.Context, input VideoInput) (media.Job, } externalModel := model.ExternalPublicID(route.Provider, route.PublicID) quotaMode := s.providers.QuotaMode(route.Provider, route.UpstreamModel) - lease, err := s.selector.Acquire(ctx, route.Provider, route.ID, route.UpstreamModel, quotaMode, "", nil, false) + lease, err := s.selector.AcquireForKey(ctx, route.Provider, route.ID, route.UpstreamModel, quotaMode, "", nil, false, input.ClientKey.AccountScope()) if err != nil { return media.Job{}, fmt.Errorf("%w: %w", ErrNoAvailableAccount, err) } diff --git a/backend/internal/application/invalidation/service.go b/backend/internal/application/invalidation/service.go index b6035cd00..e711ce69b 100644 --- a/backend/internal/application/invalidation/service.go +++ b/backend/internal/application/invalidation/service.go @@ -99,12 +99,17 @@ func (s *Service) RunPublisher(ctx context.Context) error { } type invalidationKey struct { - layer repository.InvalidationLayer - provider string + layer repository.InvalidationLayer + provider string + clientKeyID uint64 } func eventKey(event repository.InvalidationEvent) invalidationKey { - return invalidationKey{layer: event.Layer(), provider: string(event.Provider)} + key := invalidationKey{layer: event.Layer(), provider: string(event.Provider)} + if key.layer == repository.InvalidationLayerClientKey { + key.clientKeyID = event.ClientKeyID + } + return key } // RunSubscriber consumes remote events. Pub/Sub delivery is best effort; the diff --git a/backend/internal/application/invalidation/service_test.go b/backend/internal/application/invalidation/service_test.go index dfdb2312b..861dbd999 100644 --- a/backend/internal/application/invalidation/service_test.go +++ b/backend/internal/application/invalidation/service_test.go @@ -97,3 +97,12 @@ func TestRunPublisherDoesNotStopAfterPublishFailure(t *testing.T) { t.Fatalf("publish failures = %d", service.failures.Load()) } } + +func TestEventKeyKeepsClientKeyInvalidationsDistinct(t *testing.T) { + first := eventKey(repository.InvalidationEvent{Kind: repository.InvalidationClientKeyChanged, ClientKeyID: 1}) + second := eventKey(repository.InvalidationEvent{Kind: repository.InvalidationClientKeyChanged, ClientKeyID: 2}) + global := eventKey(repository.InvalidationEvent{Kind: repository.InvalidationClientKeyChanged}) + if first == second || first == global || second == global { + t.Fatalf("client-key invalidation keys were coalesced: first=%#v second=%#v global=%#v", first, second, global) + } +} diff --git a/backend/internal/application/model/service.go b/backend/internal/application/model/service.go index f7f68de8e..2e5394f72 100644 --- a/backend/internal/application/model/service.go +++ b/backend/internal/application/model/service.go @@ -10,6 +10,7 @@ import ( accountapp "github.com/chenyme/grok2api/backend/internal/application/account" "github.com/chenyme/grok2api/backend/internal/domain/account" + clientkeydomain "github.com/chenyme/grok2api/backend/internal/domain/clientkey" modeldomain "github.com/chenyme/grok2api/backend/internal/domain/model" "github.com/chenyme/grok2api/backend/internal/infra/provider" "github.com/chenyme/grok2api/backend/internal/pkg/batch" @@ -48,9 +49,12 @@ type AccountOption struct { } type ListFilter struct { - Provider string - Status string - Sort repository.SortQuery + Provider string + Providers []string + Tiers []string + Status string + ActiveScope bool + Sort repository.SortQuery } // Service 负责上游模型发现、内部来源路由与对外模型名称维护。 @@ -82,7 +86,7 @@ func (s *Service) SetLogger(logger *slog.Logger) { func (s *Service) List(ctx context.Context, page, pageSize int, search string, filter ListFilter) ([]modeldomain.Route, int64, error) { page, pageSize = normalizePage(page, pageSize) - if !validProviderFilter(filter.Provider) || !validModelFilter(filter.Status, "", "enabled", "disabled") || !repository.IsValidSort(filter.Sort, "publicId", "upstreamModel", "status", "provider", "accountSupport", "lastSyncedAt") { + if !validProviderFilter(filter.Provider) || !validProviderFilters(filter.Providers) || !validTierFilters(filter.Tiers) || !validModelFilter(filter.Status, "", "enabled", "disabled") || !repository.IsValidSort(filter.Sort, "publicId", "upstreamModel", "status", "provider", "accountSupport", "lastSyncedAt") { return nil, 0, ErrInvalidFilter } var enabled *bool @@ -90,13 +94,41 @@ func (s *Service) List(ctx context.Context, page, pageSize int, search string, f value := filter.Status == "enabled" enabled = &value } - return s.models.List(ctx, repository.ModelListQuery{Page: repository.PageQuery{Offset: (page - 1) * pageSize, Limit: pageSize, Search: search, Sort: filter.Sort}, Filter: repository.ModelListFilter{Provider: filter.Provider, Enabled: enabled}}) + return s.models.List(ctx, repository.ModelListQuery{Page: repository.PageQuery{Offset: (page - 1) * pageSize, Limit: pageSize, Search: search, Sort: filter.Sort}, Filter: repository.ModelListFilter{Provider: filter.Provider, Providers: filter.Providers, Tiers: filter.Tiers, Enabled: enabled, ActiveScope: filter.ActiveScope}}) } func validProviderFilter(value string) bool { return value == "" || account.Provider(value).IsValid() } +func validProviderFilters(values []string) bool { + seen := make(map[string]struct{}, len(values)) + for _, value := range values { + if !account.Provider(value).IsValid() { + return false + } + if _, exists := seen[value]; exists { + return false + } + seen[value] = struct{}{} + } + return true +} + +func validTierFilters(values []string) bool { + seen := make(map[string]struct{}, len(values)) + for _, value := range values { + if value != "free" && value != "super" { + return false + } + if _, exists := seen[value]; exists { + return false + } + seen[value] = struct{}{} + } + return true +} + func validModelFilter(value string, allowed ...string) bool { for _, candidate := range allowed { if value == candidate { @@ -110,6 +142,25 @@ func (s *Service) ListEnabled(ctx context.Context) ([]modeldomain.Route, error) return s.models.ListEnabled(ctx) } +func (s *Service) ListEnabledForClientKey(ctx context.Context, key clientkeydomain.Key) ([]modeldomain.Route, error) { + scope, valid := clientkeydomain.NormalizeAccountScope(clientkeydomain.AccountScope{Providers: key.ProviderScope, Tiers: key.TierScope}) + if !valid { + return nil, ErrInvalidFilter + } + if !scope.IsRestricted() { + return s.models.ListEnabled(ctx) + } + providers := scope.Providers.Values() + if len(providers) == 1 && providers[0] == "all" { + providers = nil + } + tiers := scope.Tiers.Values() + if len(tiers) == 1 && tiers[0] == "all" { + tiers = nil + } + return s.models.ListEnabledForScope(ctx, repository.ModelListFilter{Providers: providers, Tiers: tiers}) +} + func (s *Service) Get(ctx context.Context, id uint64) (modeldomain.Route, error) { return s.models.Get(ctx, id) } diff --git a/backend/internal/domain/clientkey/client_key.go b/backend/internal/domain/clientkey/client_key.go index 337fda283..cd1819f66 100644 --- a/backend/internal/domain/clientkey/client_key.go +++ b/backend/internal/domain/clientkey/client_key.go @@ -1,6 +1,10 @@ package clientkey -import "time" +import ( + "time" + + "github.com/chenyme/grok2api/backend/internal/domain/account" +) const ( DefaultRPMLimit = 120 @@ -10,6 +14,208 @@ const ( MaxBillingLimitTicks = 9_000_000_000_000_000 ) +type ProviderScope uint8 + +const ( + ProviderScopeBuild ProviderScope = 1 << iota + ProviderScopeWeb + ProviderScopeConsole + ProviderScopeAll = ProviderScopeBuild | ProviderScopeWeb | ProviderScopeConsole +) + +type TierScope uint8 + +const ( + TierScopeFree TierScope = 1 << iota + TierScopeSuper + TierScopeUnknown + TierScopeAll = TierScopeFree | TierScopeSuper | TierScopeUnknown +) + +type AccountTier uint8 + +const ( + AccountTierFree AccountTier = iota + 1 + AccountTierSuper + AccountTierUnknown +) + +type AccountScope struct { + Providers ProviderScope + Tiers TierScope +} + +func ParseProviderScopeValues(values []string) (ProviderScope, bool) { + if len(values) == 0 { + return 0, false + } + var scope ProviderScope + for _, value := range values { + switch value { + case "all": + if len(values) != 1 { + return 0, false + } + return ProviderScopeAll, true + case string(account.ProviderBuild): + scope |= ProviderScopeBuild + case string(account.ProviderWeb): + scope |= ProviderScopeWeb + case string(account.ProviderConsole): + scope |= ProviderScopeConsole + default: + return 0, false + } + } + return NormalizeProviderScope(scope) +} + +func ParseTierScopeValues(values []string) (TierScope, bool) { + if len(values) == 0 { + return 0, false + } + var scope TierScope + for _, value := range values { + switch value { + case "all": + if len(values) != 1 { + return 0, false + } + return TierScopeAll, true + case "free": + scope |= TierScopeFree + case "super": + scope |= TierScopeSuper + default: + return 0, false + } + } + return NormalizeTierScope(scope) +} + +func (s ProviderScope) Values() []string { + value, valid := NormalizeProviderScope(s) + if !valid || value == ProviderScopeAll { + return []string{"all"} + } + values := make([]string, 0, 3) + if value&ProviderScopeBuild != 0 { + values = append(values, string(account.ProviderBuild)) + } + if value&ProviderScopeWeb != 0 { + values = append(values, string(account.ProviderWeb)) + } + if value&ProviderScopeConsole != 0 { + values = append(values, string(account.ProviderConsole)) + } + return values +} + +func (s TierScope) Values() []string { + value, valid := NormalizeTierScope(s) + if !valid || value == TierScopeAll { + return []string{"all"} + } + values := make([]string, 0, 2) + if value&TierScopeFree != 0 { + values = append(values, "free") + } + if value&TierScopeSuper != 0 { + values = append(values, "super") + } + return values +} + +func NormalizeProviderScope(value ProviderScope) (ProviderScope, bool) { + if value == 0 { + return ProviderScopeAll, true + } + if value&^ProviderScopeAll != 0 { + return value, false + } + return value, true +} + +func NormalizeTierScope(value TierScope) (TierScope, bool) { + if value == 0 { + return TierScopeAll, true + } + switch value { + case TierScopeFree, TierScopeSuper, TierScopeFree | TierScopeSuper, TierScopeAll: + return value, true + default: + return value, false + } +} + +func NormalizeAccountScope(value AccountScope) (AccountScope, bool) { + providers, providersValid := NormalizeProviderScope(value.Providers) + tiers, tiersValid := NormalizeTierScope(value.Tiers) + return AccountScope{Providers: providers, Tiers: tiers}, providersValid && tiersValid +} + +func (s ProviderScope) Allows(provider account.Provider) bool { + value, valid := NormalizeProviderScope(s) + if !valid { + return false + } + switch provider { + case account.ProviderBuild: + return value&ProviderScopeBuild != 0 + case account.ProviderWeb: + return value&ProviderScopeWeb != 0 + case account.ProviderConsole: + return value&ProviderScopeConsole != 0 + default: + return false + } +} + +func (s TierScope) Allows(tier AccountTier) bool { + value, valid := NormalizeTierScope(s) + if !valid { + return false + } + switch tier { + case AccountTierFree: + return value&TierScopeFree != 0 + case AccountTierSuper: + return value&TierScopeSuper != 0 + case AccountTierUnknown: + return value&TierScopeUnknown != 0 + default: + return false + } +} + +func (s AccountScope) AllowsProvider(provider account.Provider) bool { + value, valid := NormalizeAccountScope(s) + return valid && value.Providers.Allows(provider) +} + +// AllowsAccount applies provider restrictions to every channel while tier +// restrictions apply only to channels with a reliable tier classification. +func (s AccountScope) AllowsAccount(provider account.Provider, tier AccountTier) bool { + value, valid := NormalizeAccountScope(s) + if !valid || !value.Providers.Allows(provider) { + return false + } + if provider == account.ProviderConsole { + return true + } + return value.Tiers.Allows(tier) +} + +func (s AccountScope) IsRestricted() bool { + value, valid := NormalizeAccountScope(s) + return valid && (value.Providers != ProviderScopeAll || value.Tiers != TierScopeAll) +} + +func (k Key) AccountScope() AccountScope { + value, _ := NormalizeAccountScope(AccountScope{Providers: k.ProviderScope, Tiers: k.TierScope}) + return value +} + // Key 表示下游客户端调用凭据及其限制。 type Key struct { ID uint64 @@ -28,9 +234,12 @@ type Key struct { // Registered compatibility aliases remain available to avoid breaking existing clients. AllowModelAliases bool AllowedModels []uint64 - LastUsedAt *time.Time - CreatedAt time.Time - UpdatedAt time.Time + // ProviderScope and TierScope narrow routing without adding request-time storage lookups. + ProviderScope ProviderScope + TierScope TierScope + LastUsedAt *time.Time + CreatedAt time.Time + UpdatedAt time.Time } // IsAvailable 判断客户端 Key 当前是否可用。 @@ -40,3 +249,15 @@ func (k Key) IsAvailable(now time.Time) bool { } return k.ExpiresAt == nil || now.Before(*k.ExpiresAt) } + +func (k Key) AllowsModel(modelID uint64) bool { + if len(k.AllowedModels) == 0 { + return true + } + for _, allowed := range k.AllowedModels { + if allowed == modelID { + return true + } + } + return false +} diff --git a/backend/internal/domain/clientkey/client_key_test.go b/backend/internal/domain/clientkey/client_key_test.go new file mode 100644 index 000000000..d4e27246b --- /dev/null +++ b/backend/internal/domain/clientkey/client_key_test.go @@ -0,0 +1,59 @@ +package clientkey + +import ( + "reflect" + "testing" + + "github.com/chenyme/grok2api/backend/internal/domain/account" +) + +func TestAccountScopeParsingAndNormalization(t *testing.T) { + if _, valid := ParseProviderScopeValues(nil); valid { + t.Fatal("an explicitly empty provider scope must be rejected") + } + if _, valid := ParseTierScopeValues(nil); valid { + t.Fatal("an explicitly empty tier scope must be rejected") + } + providers, valid := ParseProviderScopeValues([]string{"grok_build", "grok_web"}) + if !valid || providers != ProviderScopeBuild|ProviderScopeWeb { + t.Fatalf("provider scope = %d, valid = %v", providers, valid) + } + if _, valid := ParseProviderScopeValues([]string{"all", "grok_web"}); valid { + t.Fatal("all mixed with a provider must be rejected") + } + tiers, valid := ParseTierScopeValues([]string{"free", "super"}) + if !valid || tiers != TierScopeFree|TierScopeSuper { + t.Fatalf("tier scope = %d, valid = %v", tiers, valid) + } + if tiers == TierScopeAll { + t.Fatal("free + super must continue excluding unknown tiers") + } + if _, valid := ParseTierScopeValues([]string{"all", "free"}); valid { + t.Fatal("all mixed with a tier must be rejected") + } + if values := ProviderScopeAll.Values(); !reflect.DeepEqual(values, []string{"all"}) { + t.Fatalf("all provider values = %#v", values) + } + if values := tiers.Values(); !reflect.DeepEqual(values, []string{"free", "super"}) { + t.Fatalf("tier values = %#v", values) + } +} + +func TestAccountScopeAllowsProviderAndTierIntersection(t *testing.T) { + scope := AccountScope{Providers: ProviderScopeBuild | ProviderScopeConsole, Tiers: TierScopeFree} + if !scope.AllowsAccount(account.ProviderBuild, AccountTierFree) { + t.Fatal("Build Free should be allowed") + } + if scope.AllowsAccount(account.ProviderBuild, AccountTierSuper) { + t.Fatal("Build Super should be excluded") + } + if scope.AllowsAccount(account.ProviderWeb, AccountTierFree) { + t.Fatal("Web should be excluded by provider scope") + } + if !scope.AllowsAccount(account.ProviderConsole, AccountTierUnknown) { + t.Fatal("Console should ignore tier restrictions") + } + if (AccountScope{Providers: ProviderScopeBuild, Tiers: TierScopeFree | TierScopeSuper}).AllowsAccount(account.ProviderBuild, AccountTierUnknown) { + t.Fatal("Free + Super should exclude unknown Build accounts") + } +} diff --git a/backend/internal/infra/persistence/relational/client_key_repository.go b/backend/internal/infra/persistence/relational/client_key_repository.go index 2656a6a19..d050edc0b 100644 --- a/backend/internal/infra/persistence/relational/client_key_repository.go +++ b/backend/internal/infra/persistence/relational/client_key_repository.go @@ -13,10 +13,23 @@ import ( "gorm.io/gorm/clause" ) -type ClientKeyRepository struct{ db *Database } +type ClientKeyRepository struct { + db *Database + observer repository.InvalidationObserver +} func NewClientKeyRepository(db *Database) *ClientKeyRepository { return &ClientKeyRepository{db: db} } +func (r *ClientKeyRepository) SetInvalidationObserver(observer repository.InvalidationObserver) { + r.observer = observer +} + +func (r *ClientKeyRepository) notifyInvalidation(ctx context.Context, clientKeyID uint64) { + if r.observer != nil { + r.observer(ctx, repository.InvalidationEvent{Kind: repository.InvalidationClientKeyChanged, ClientKeyID: clientKeyID}) + } +} + func (r *ClientKeyRepository) List(ctx context.Context, input repository.ClientKeyListQuery) ([]clientkey.Key, int64, error) { var total int64 query := r.db.db.WithContext(ctx).Model(&clientKeyModel{}) @@ -53,7 +66,7 @@ func (r *ClientKeyRepository) List(ctx context.Context, input repository.ClientK "expiresAt": {expression: "client_keys.expires_at", nullsLast: true, defaultDirection: repository.SortDescending}, "lastUsedAt": {expression: "client_keys.last_used_at", nullsLast: true, defaultDirection: repository.SortDescending}, }, sortSpec{expression: "client_keys.created_at", defaultDirection: repository.SortDescending}, "client_keys.id") - if err := query.Select("id", "name", "prefix", "enabled", "expires_at", "rpm_limit", "max_concurrent", "billing_limit_usd_ticks", "billed_usage_usd_ticks", "reserved_usage_usd_ticks", "last_used_at", "created_at", "updated_at").Offset(input.Page.Offset).Limit(input.Page.Limit).Find(&rows).Error; err != nil { + if err := query.Select("id", "name", "prefix", "enabled", "expires_at", "rpm_limit", "max_concurrent", "billing_limit_usd_ticks", "billed_usage_usd_ticks", "reserved_usage_usd_ticks", "allow_model_aliases", "provider_scope_mask", "tier_scope_mask", "last_used_at", "created_at", "updated_at").Offset(input.Page.Offset).Limit(input.Page.Limit).Find(&rows).Error; err != nil { return nil, 0, err } ids := make([]uint64, 0, len(rows)) @@ -76,11 +89,18 @@ func (r *ClientKeyRepository) UpdateManyEnabled(ctx context.Context, ids []uint6 return 0, nil } result := r.db.db.WithContext(ctx).Model(&clientKeyModel{}).Where("id IN ?", ids).Update("enabled", enabled) + if result.Error == nil && result.RowsAffected > 0 { + r.notifyInvalidation(ctx, 0) + } return result.RowsAffected, result.Error } func (r *ClientKeyRepository) Create(ctx context.Context, value clientkey.Key) (clientkey.Key, error) { - row := clientKeyModel{Name: value.Name, Prefix: value.Prefix, SecretHash: value.SecretHash, EncryptedSecret: value.EncryptedSecret, Enabled: value.Enabled, ExpiresAt: value.ExpiresAt, RPMLimit: value.RPMLimit, MaxConcurrent: value.MaxConcurrent, BillingLimitUSDTicks: value.BillingLimitUSDTicks, BilledUsageUSDTicks: value.BilledUsageUSDTicks, ReservedUsageUSDTicks: value.ReservedUsageUSDTicks, AllowModelAliases: value.AllowModelAliases} + scope, valid := clientkey.NormalizeAccountScope(clientkey.AccountScope{Providers: value.ProviderScope, Tiers: value.TierScope}) + if !valid { + return clientkey.Key{}, repository.ErrConflict + } + row := clientKeyModel{Name: value.Name, Prefix: value.Prefix, SecretHash: value.SecretHash, EncryptedSecret: value.EncryptedSecret, Enabled: value.Enabled, ExpiresAt: value.ExpiresAt, RPMLimit: value.RPMLimit, MaxConcurrent: value.MaxConcurrent, BillingLimitUSDTicks: value.BillingLimitUSDTicks, BilledUsageUSDTicks: value.BilledUsageUSDTicks, ReservedUsageUSDTicks: value.ReservedUsageUSDTicks, AllowModelAliases: value.AllowModelAliases, ProviderScopeMask: uint8(scope.Providers), TierScopeMask: uint8(scope.Tiers)} err := r.db.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { if err := tx.Create(&row).Error; err != nil { return err @@ -134,11 +154,16 @@ func (r *ClientKeyRepository) GetByPrefix(ctx context.Context, prefix string) (c } func (r *ClientKeyRepository) Update(ctx context.Context, value clientkey.Key) (clientkey.Key, error) { + scope, valid := clientkey.NormalizeAccountScope(clientkey.AccountScope{Providers: value.ProviderScope, Tiers: value.TierScope}) + if !valid { + return clientkey.Key{}, repository.ErrConflict + } err := r.db.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { result := tx.Model(&clientKeyModel{}).Where("id = ?", value.ID).Updates(map[string]any{ "name": value.Name, "enabled": value.Enabled, "expires_at": value.ExpiresAt, "rpm_limit": value.RPMLimit, "max_concurrent": value.MaxConcurrent, "billing_limit_usd_ticks": value.BillingLimitUSDTicks, "allow_model_aliases": value.AllowModelAliases, + "provider_scope_mask": uint8(scope.Providers), "tier_scope_mask": uint8(scope.Tiers), "updated_at": time.Now().UTC(), }) if result.Error != nil { @@ -152,6 +177,7 @@ func (r *ClientKeyRepository) Update(ctx context.Context, value clientkey.Key) ( if err != nil { return clientkey.Key{}, mapError(err) } + r.notifyInvalidation(ctx, value.ID) return r.Get(ctx, value.ID) } @@ -163,6 +189,7 @@ func (r *ClientKeyRepository) Delete(ctx context.Context, id uint64) error { if result.RowsAffected == 0 { return repository.ErrNotFound } + r.notifyInvalidation(ctx, id) return nil } @@ -171,6 +198,9 @@ func (r *ClientKeyRepository) DeleteMany(ctx context.Context, ids []uint64) (int return 0, nil } result := r.db.db.WithContext(ctx).Where("id IN ?", ids).Delete(&clientKeyModel{}) + if result.Error == nil && result.RowsAffected > 0 { + r.notifyInvalidation(ctx, 0) + } return result.RowsAffected, result.Error } diff --git a/backend/internal/infra/persistence/relational/invalidation_test.go b/backend/internal/infra/persistence/relational/invalidation_test.go index 6a0e5191c..9ce7f97c6 100644 --- a/backend/internal/infra/persistence/relational/invalidation_test.go +++ b/backend/internal/infra/persistence/relational/invalidation_test.go @@ -7,6 +7,7 @@ import ( "time" "github.com/chenyme/grok2api/backend/internal/domain/account" + clientkeydomain "github.com/chenyme/grok2api/backend/internal/domain/clientkey" "github.com/chenyme/grok2api/backend/internal/domain/model" "github.com/chenyme/grok2api/backend/internal/repository" ) @@ -81,6 +82,61 @@ func TestRoutingMutationsNotifyAfterCommit(t *testing.T) { } } +func TestClientKeyMutationsNotifyAfterCommit(t *testing.T) { + ctx := context.Background() + database, err := OpenSQLite(ctx, filepath.Join(t.TempDir(), "client-key-invalidation.db")) + if err != nil { + t.Fatal(err) + } + defer database.Close() + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + keys := NewClientKeyRepository(database) + created, err := keys.Create(ctx, clientkeydomain.Key{ + Name: "scoped", Prefix: "scoped", SecretHash: testSecretHash, EncryptedSecret: testEncryptedToken, + Enabled: true, RPMLimit: 60, MaxConcurrent: 4, ProviderScope: clientkeydomain.ProviderScopeBuild, TierScope: clientkeydomain.TierScopeFree, + }) + if err != nil { + t.Fatal(err) + } + var events []repository.InvalidationEvent + keys.SetInvalidationObserver(func(_ context.Context, event repository.InvalidationEvent) { + events = append(events, event) + }) + + created.ProviderScope = clientkeydomain.ProviderScopeWeb + created.TierScope = clientkeydomain.TierScopeSuper + if _, err := keys.Update(ctx, created); err != nil { + t.Fatal(err) + } + if len(events) != 1 || events[0].Kind != repository.InvalidationClientKeyChanged || events[0].ClientKeyID != created.ID || !events[0].Valid() { + t.Fatalf("client-key update events = %#v", events) + } + + missing := created + missing.ID = created.ID + 999 + if _, err := keys.Update(ctx, missing); err == nil { + t.Fatal("missing client-key update should fail") + } + if len(events) != 1 { + t.Fatalf("failed update emitted invalidation: %#v", events[1:]) + } + + if _, err := keys.UpdateManyEnabled(ctx, []uint64{created.ID}, false); err != nil { + t.Fatal(err) + } + if len(events) != 2 || events[1].ClientKeyID != 0 || !events[1].Valid() { + t.Fatalf("client-key batch events = %#v", events) + } + if err := keys.Delete(ctx, created.ID); err != nil { + t.Fatal(err) + } + if len(events) != 3 || events[2].ClientKeyID != created.ID || !events[2].Valid() { + t.Fatalf("client-key delete events = %#v", events) + } +} + func TestModelUpdateInvalidationUsesStoredRouteIdentity(t *testing.T) { ctx := context.Background() database, err := OpenSQLite(ctx, filepath.Join(t.TempDir(), "model-update-invalidation.db")) diff --git a/backend/internal/infra/persistence/relational/mapping.go b/backend/internal/infra/persistence/relational/mapping.go index 722a5caf7..cc1ca2ea4 100644 --- a/backend/internal/infra/persistence/relational/mapping.go +++ b/backend/internal/infra/persistence/relational/mapping.go @@ -210,11 +210,15 @@ func toModelDomain(value modelRouteModel) model.Route { } func toClientKeyDomain(value clientKeyModel, allowedModels []uint64) clientkey.Key { + providerScope, _ := clientkey.NormalizeProviderScope(clientkey.ProviderScope(value.ProviderScopeMask)) + tierScope, _ := clientkey.NormalizeTierScope(clientkey.TierScope(value.TierScopeMask)) return clientkey.Key{ ID: value.ID, Name: value.Name, Prefix: value.Prefix, SecretHash: value.SecretHash, EncryptedSecret: value.EncryptedSecret, Enabled: value.Enabled, ExpiresAt: value.ExpiresAt, RPMLimit: value.RPMLimit, MaxConcurrent: value.MaxConcurrent, BillingLimitUSDTicks: value.BillingLimitUSDTicks, BilledUsageUSDTicks: value.BilledUsageUSDTicks, ReservedUsageUSDTicks: value.ReservedUsageUSDTicks, - AllowModelAliases: value.AllowModelAliases, AllowedModels: allowedModels, LastUsedAt: value.LastUsedAt, CreatedAt: value.CreatedAt, UpdatedAt: value.UpdatedAt, + AllowModelAliases: value.AllowModelAliases, AllowedModels: allowedModels, + ProviderScope: providerScope, TierScope: tierScope, + LastUsedAt: value.LastUsedAt, CreatedAt: value.CreatedAt, UpdatedAt: value.UpdatedAt, } } diff --git a/backend/internal/infra/persistence/relational/model_repository.go b/backend/internal/infra/persistence/relational/model_repository.go index c9e5aa396..bef044816 100644 --- a/backend/internal/infra/persistence/relational/model_repository.go +++ b/backend/internal/infra/persistence/relational/model_repository.go @@ -72,6 +72,105 @@ const modelSharedPaidBuildSupportAvailabilityExpression = `(route.provider = 'gr AND ` + modelPeerBuildSuperPredicate + ` ))` +// These predicates mirror the gateway's client-key scope classification. They +// are used only by the admin model picker to avoid presenting routes that the +// selected account scope can never serve. +const modelAccountBuildFreePredicate = `(account.provider = 'grok_build' + AND NOT ` + modelAccountBuildSuperPredicate + ` + AND ( + EXISTS (SELECT 1 FROM account_quota_recovery recovery WHERE recovery.account_id = account.id AND recovery.kind = 'free') + OR LOWER(TRIM(account.observed_model)) LIKE '%-build-free' + OR EXISTS (SELECT 1 FROM account_billing_snapshots billing WHERE billing.account_id = account.id AND ` + accountFreeBillingSignal + `) + ))` + +const modelAccountBuildSuperTierPredicate = `(account.provider = 'grok_build' AND ` + modelAccountBuildSuperPredicate + `)` + +const modelAccountWebFreePredicate = `(account.provider = 'grok_web' + AND EXISTS (SELECT 1 FROM web_account_profiles profile WHERE profile.account_id = account.id AND profile.tier = 'basic'))` + +const modelAccountWebSuperPredicate = `(account.provider = 'grok_web' + AND EXISTS (SELECT 1 FROM web_account_profiles profile WHERE profile.account_id = account.id AND profile.tier IN ('super', 'heavy')))` + +const modelSharedPaidBuildScopeExpression = `(model_routes.provider = 'grok_build' + AND ` + modelAccountBuildSuperPredicate + ` + AND EXISTS ( + SELECT 1 + FROM provider_accounts peer + JOIN account_model_capabilities peer_capability ON peer_capability.account_id = peer.id AND peer_capability.upstream_model = model_routes.upstream_model + WHERE peer.provider = model_routes.provider + AND ` + modelPeerBuildSuperPredicate + ` + ))` + +const modelRouteAccountCapabilityPredicate = `( + EXISTS ( + SELECT 1 FROM model_route_accounts binding + WHERE binding.model_route_id = model_routes.id + AND binding.account_id = account.id + ) + OR ( + NOT EXISTS (SELECT 1 FROM model_route_accounts binding WHERE binding.model_route_id = model_routes.id) + AND ( + EXISTS ( + SELECT 1 FROM account_model_capabilities capability + WHERE capability.account_id = account.id + AND capability.upstream_model = model_routes.upstream_model + ) + OR ` + modelSharedPaidBuildScopeExpression + ` + ) + ) +)` + +const modelAvailableRouteAccountCapabilityPredicate = `( + EXISTS ( + SELECT 1 FROM model_route_accounts binding + WHERE binding.model_route_id = model_routes.id + AND binding.account_id = account.id + ) + OR ( + NOT EXISTS (SELECT 1 FROM model_route_accounts binding WHERE binding.model_route_id = model_routes.id) + AND ( + EXISTS ( + SELECT 1 FROM account_model_capabilities capability + WHERE capability.account_id = account.id + AND capability.upstream_model = model_routes.upstream_model + ) + OR ` + modelSharedPaidBuildSupportSortExpression + ` + ) + ) +)` + +func modelTierAvailabilityPredicate(tiers []string) string { + return modelTierAvailabilityPredicateWithAvailability(tiers, false) +} + +func modelTierAvailabilityPredicateWithAvailability(tiers []string, activeOnly bool) string { + parts := make([]string, 0, len(tiers)) + for _, tier := range tiers { + switch tier { + case "free": + parts = append(parts, modelAccountBuildFreePredicate, modelAccountWebFreePredicate) + case "super": + parts = append(parts, modelAccountBuildSuperTierPredicate, modelAccountWebSuperPredicate) + } + } + if len(parts) == 0 { + return "" + } + accountPredicate := "" + capabilityPredicate := modelRouteAccountCapabilityPredicate + if activeOnly { + accountPredicate = " AND account.enabled = TRUE AND account.auth_status = 'active'" + capabilityPredicate = modelAvailableRouteAccountCapabilityPredicate + } + return `(model_routes.provider = 'grok_console' OR EXISTS ( + SELECT 1 FROM provider_accounts account + WHERE account.provider = model_routes.provider + ` + accountPredicate + ` + AND (` + strings.Join(parts, " OR ") + `) + AND ` + capabilityPredicate + ` + ))` +} + const ( modelProviderPriorityExpression = "CASE model_routes.provider WHEN 'grok_build' THEN 0 WHEN 'grok_web' THEN 1 WHEN 'grok_console' THEN 2 ELSE 3 END" modelSupportSortExpression = `(SELECT COUNT(*) FROM provider_accounts account WHERE account.provider = model_routes.provider AND account.enabled = TRUE AND account.auth_status = 'active' AND (EXISTS (SELECT 1 FROM model_route_accounts binding WHERE binding.model_route_id = model_routes.id AND binding.account_id = account.id) OR (NOT EXISTS (SELECT 1 FROM model_route_accounts binding WHERE binding.model_route_id = model_routes.id) AND (EXISTS (SELECT 1 FROM account_model_capabilities capability WHERE capability.account_id = account.id AND capability.upstream_model = model_routes.upstream_model) OR ` + modelSharedPaidBuildSupportSortExpression + `))))` @@ -93,6 +192,9 @@ func (r *ModelRepository) notifyInvalidation(ctx context.Context, event reposito func (r *ModelRepository) List(ctx context.Context, input repository.ModelListQuery) ([]model.Route, int64, error) { var total int64 query := r.db.db.WithContext(ctx).Model(&modelRouteModel{}) + if input.Filter.ActiveScope { + query = r.availableRoutes(query) + } if search := strings.TrimSpace(input.Page.Search); search != "" { pattern := "%" + strings.ToLower(search) + "%" query = query.Where("LOWER(public_id) LIKE ? OR LOWER(upstream_model) LIKE ?", pattern, pattern) @@ -100,6 +202,16 @@ func (r *ModelRepository) List(ctx context.Context, input repository.ModelListQu if input.Filter.Provider != "" { query = query.Where("provider = ?", input.Filter.Provider) } + if len(input.Filter.Providers) > 0 { + query = query.Where("provider IN ?", input.Filter.Providers) + } + tierPredicate := modelTierAvailabilityPredicate(input.Filter.Tiers) + if input.Filter.ActiveScope { + tierPredicate = modelTierAvailabilityPredicateWithAvailability(input.Filter.Tiers, true) + } + if tierPredicate != "" { + query = query.Where(tierPredicate) + } if input.Filter.Enabled != nil { query = query.Where("enabled = ?", *input.Filter.Enabled) } @@ -137,6 +249,25 @@ func (r *ModelRepository) ListEnabled(ctx context.Context) ([]model.Route, error return values, nil } +func (r *ModelRepository) ListEnabledForScope(ctx context.Context, filter repository.ModelListFilter) ([]model.Route, error) { + query := r.availableRoutes(r.db.db.WithContext(ctx)).Where("enabled = ?", true) + if len(filter.Providers) > 0 { + query = query.Where("provider IN ?", filter.Providers) + } + if tierPredicate := modelTierAvailabilityPredicateWithAvailability(filter.Tiers, true); tierPredicate != "" { + query = query.Where(tierPredicate) + } + var rows []modelRouteModel + if err := query.Order("public_id ASC, id ASC").Find(&rows).Error; err != nil { + return nil, err + } + values := mapModelRows(rows) + if err := r.annotateAvailability(ctx, values); err != nil { + return nil, err + } + return values, nil +} + // ListConfiguredEnabled 返回所有已启用配置,包括暂时没有可用账号的路由,供 readiness 展示部分故障。 func (r *ModelRepository) ListConfiguredEnabled(ctx context.Context) ([]model.Route, error) { var rows []modelRouteModel diff --git a/backend/internal/infra/persistence/relational/model_scope_filter_test.go b/backend/internal/infra/persistence/relational/model_scope_filter_test.go new file mode 100644 index 000000000..05941a191 --- /dev/null +++ b/backend/internal/infra/persistence/relational/model_scope_filter_test.go @@ -0,0 +1,156 @@ +package relational + +import ( + "context" + "os" + "testing" + "time" + + "github.com/chenyme/grok2api/backend/internal/domain/account" + modeldomain "github.com/chenyme/grok2api/backend/internal/domain/model" + "github.com/chenyme/grok2api/backend/internal/repository" +) + +func TestModelListFiltersByClientKeyProviderAndTierScope(t *testing.T) { + database := openTestDatabase(t) + assertModelListFiltersByClientKeyProviderAndTierScope(t, database) +} + +func TestPostgresModelListFiltersByClientKeyProviderAndTierScope(t *testing.T) { + dsn := os.Getenv("TEST_POSTGRES_DSN") + if dsn == "" { + t.Skip("TEST_POSTGRES_DSN is not configured") + } + ctx, cancel := context.WithTimeout(context.Background(), time.Minute) + defer cancel() + database, err := OpenPostgres(ctx, dsn, 5, 2) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := database.Close(); err != nil { + t.Errorf("close PostgreSQL model scope database: %v", err) + } + }) + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + assertModelListFiltersByClientKeyProviderAndTierScope(t, database) +} + +func assertModelListFiltersByClientKeyProviderAndTierScope(t *testing.T, database *Database) { + t.Helper() + ctx := context.Background() + repo := NewModelRepository(database) + + now := time.Now().UTC() + accounts := []accountModel{ + {IdentityKey: testIdentityKey("scope-build-free"), Provider: string(account.ProviderBuild), Name: "build-free", SourceKey: "build-free", ObservedModel: "grok-build-free", Enabled: true, AuthStatus: string(account.AuthStatusActive), Priority: 1}, + {IdentityKey: testIdentityKey("scope-build-super"), Provider: string(account.ProviderBuild), Name: "build-super", SourceKey: "build-super", Enabled: true, AuthStatus: string(account.AuthStatusActive), Priority: 1}, + {IdentityKey: testIdentityKey("scope-web-free"), Provider: string(account.ProviderWeb), Name: "web-free", SourceKey: "web-free", Enabled: true, AuthStatus: string(account.AuthStatusActive), Priority: 1}, + {IdentityKey: testIdentityKey("scope-web-super"), Provider: string(account.ProviderWeb), Name: "web-super", SourceKey: "web-super", Enabled: true, AuthStatus: string(account.AuthStatusActive), Priority: 1}, + {IdentityKey: testIdentityKey("scope-console"), Provider: string(account.ProviderConsole), Name: "console", SourceKey: "console", Enabled: true, AuthStatus: string(account.AuthStatusActive), Priority: 1}, + } + for index := range accounts { + if err := database.db.WithContext(ctx).Create(&accounts[index]).Error; err != nil { + t.Fatal(err) + } + } + accountIDs := make([]uint64, 0, len(accounts)) + for _, value := range accounts { + accountIDs = append(accountIDs, value.ID) + } + if err := database.db.WithContext(ctx).Create(&billingModel{AccountID: accounts[1].ID, PlanName: "SuperGrokPro", SyncedAt: now}).Error; err != nil { + t.Fatal(err) + } + if err := database.db.WithContext(ctx).Create(&webAccountProfileModel{AccountID: accounts[2].ID, Tier: "basic", SyncedAt: &now}).Error; err != nil { + t.Fatal(err) + } + if err := database.db.WithContext(ctx).Create(&webAccountProfileModel{AccountID: accounts[3].ID, Tier: "super", SyncedAt: &now}).Error; err != nil { + t.Fatal(err) + } + + routes := []modeldomain.Route{ + {PublicID: "Build/scope-free", Provider: account.ProviderBuild, UpstreamModel: "scope-free", Capability: modeldomain.CapabilityResponses, Enabled: true}, + {PublicID: "Build/scope-super", Provider: account.ProviderBuild, UpstreamModel: "scope-super", Capability: modeldomain.CapabilityResponses, Enabled: true}, + {PublicID: "scope-web-free", Provider: account.ProviderWeb, UpstreamModel: "scope-web-free", Capability: modeldomain.CapabilityChat, Enabled: true}, + {PublicID: "scope-web-super", Provider: account.ProviderWeb, UpstreamModel: "scope-web-super", Capability: modeldomain.CapabilityChat, Enabled: true}, + {PublicID: "scope-console", Provider: account.ProviderConsole, UpstreamModel: "scope-console", Capability: modeldomain.CapabilityResponses, Enabled: true}, + } + routeIDs := make([]uint64, 0, len(routes)) + for index := range routes { + created, err := repo.Create(ctx, routes[index], nil) + if err != nil { + t.Fatal(err) + } + routeIDs = append(routeIDs, created.ID) + } + t.Cleanup(func() { + cleanupCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if _, err := repo.DeleteMany(cleanupCtx, routeIDs); err != nil { + t.Errorf("delete scoped model routes: %v", err) + } + if err := database.db.WithContext(cleanupCtx).Where("id IN ?", accountIDs).Delete(&accountModel{}).Error; err != nil { + t.Errorf("delete scoped model accounts: %v", err) + } + }) + capabilities := []accountModelCapabilityModel{ + {AccountID: accounts[0].ID, UpstreamModel: "scope-free"}, + {AccountID: accounts[1].ID, UpstreamModel: "scope-super"}, + {AccountID: accounts[2].ID, UpstreamModel: "scope-web-free"}, + {AccountID: accounts[3].ID, UpstreamModel: "scope-web-super"}, + {AccountID: accounts[4].ID, UpstreamModel: "scope-console"}, + } + if err := database.db.WithContext(ctx).Create(&capabilities).Error; err != nil { + t.Fatal(err) + } + + list := func(tiers []string, providers []string) map[string]bool { + values, _, err := repo.List(ctx, repository.ModelListQuery{Page: repository.PageQuery{Limit: 20}, Filter: repository.ModelListFilter{Providers: providers, Tiers: tiers}}) + if err != nil { + t.Fatal(err) + } + result := make(map[string]bool, len(values)) + for _, value := range values { + result[value.PublicID] = true + } + return result + } + + free := list([]string{"free"}, nil) + if len(free) != 3 || !free["Build/scope-free"] || !free["Web/scope-web-free"] || !free["Console/scope-console"] || free["Build/scope-super"] || free["Web/scope-web-super"] { + t.Fatalf("free scope models = %#v", free) + } + super := list([]string{"super"}, nil) + if len(super) != 3 || !super["Build/scope-super"] || !super["Web/scope-web-super"] || !super["Console/scope-console"] || super["Build/scope-free"] || super["Web/scope-web-free"] { + t.Fatalf("super scope models = %#v", super) + } + buildFree := list([]string{"free"}, []string{"grok_build"}) + if len(buildFree) != 1 || !buildFree["Build/scope-free"] { + t.Fatalf("build/free scope models = %#v", buildFree) + } + + enabledForScope := func(tiers []string, providers []string) map[string]bool { + values, err := repo.ListEnabledForScope(ctx, repository.ModelListFilter{Providers: providers, Tiers: tiers}) + if err != nil { + t.Fatal(err) + } + result := make(map[string]bool, len(values)) + for _, value := range values { + result[value.PublicID] = true + } + return result + } + enabledBuildFree := enabledForScope([]string{"free"}, []string{"grok_build"}) + if len(enabledBuildFree) != 1 || !enabledBuildFree["Build/scope-free"] { + t.Fatalf("enabled build/free scope models = %#v", enabledBuildFree) + } + if err := database.db.WithContext(ctx).Model(&accountModel{}).Where("id IN ?", []uint64{accounts[1].ID, accounts[3].ID}).Update("enabled", false).Error; err != nil { + t.Fatal(err) + } + enabledSuper := enabledForScope([]string{"super"}, nil) + if len(enabledSuper) != 1 || !enabledSuper["Console/scope-console"] { + t.Fatalf("inactive Super accounts must not advertise routes: %#v", enabledSuper) + } +} diff --git a/backend/internal/infra/persistence/relational/models.go b/backend/internal/infra/persistence/relational/models.go index bf01be2c7..dfacdcc3f 100644 --- a/backend/internal/infra/persistence/relational/models.go +++ b/backend/internal/infra/persistence/relational/models.go @@ -252,7 +252,9 @@ type clientKeyModel struct { BilledUsageUSDTicks int64 `gorm:"not null;default:0;check:chk_client_keys_billed_usage,billed_usage_usd_ticks >= 0"` ReservedUsageUSDTicks int64 `gorm:"not null;default:0;check:chk_client_keys_reserved_usage,reserved_usage_usd_ticks >= 0"` // AllowModelAliases defaults false so existing keys keep a clean base-model list. - AllowModelAliases bool `gorm:"not null;default:false"` + AllowModelAliases bool `gorm:"not null;default:false"` + ProviderScopeMask uint8 `gorm:"not null;default:7;check:chk_client_keys_provider_scope,provider_scope_mask BETWEEN 1 AND 7"` + TierScopeMask uint8 `gorm:"not null;default:7;check:chk_client_keys_tier_scope,tier_scope_mask IN (1,2,3,7)"` LastUsedAt *time.Time CreatedAt time.Time `gorm:"not null"` UpdatedAt time.Time `gorm:"not null"` diff --git a/backend/internal/infra/persistence/relational/postgres_integration_test.go b/backend/internal/infra/persistence/relational/postgres_integration_test.go index bdf48f0d8..ee1c8f62e 100644 --- a/backend/internal/infra/persistence/relational/postgres_integration_test.go +++ b/backend/internal/infra/persistence/relational/postgres_integration_test.go @@ -279,6 +279,33 @@ func TestPostgresRepositoriesIntegration(t *testing.T) { if err := database.InitializeSchema(ctx); err != nil { t.Fatal(err) } + keyRepository := NewClientKeyRepository(database) + poolKey, err := keyRepository.Create(ctx, clientkey.Key{Name: "postgres-account-scope", Prefix: "postgres-account-scope", SecretHash: testSecretHash, EncryptedSecret: testEncryptedToken, Enabled: true}) + if err != nil { + t.Fatal(err) + } + if err := database.db.WithContext(ctx).Migrator().DropColumn(&clientKeyModel{}, "ProviderScopeMask"); err != nil { + t.Fatal(err) + } + if err := database.db.WithContext(ctx).Migrator().DropColumn(&clientKeyModel{}, "TierScopeMask"); err != nil { + t.Fatal(err) + } + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + loadedPoolKey, err := keyRepository.Get(ctx, poolKey.ID) + if err != nil || loadedPoolKey.ProviderScope != clientkey.ProviderScopeAll || loadedPoolKey.TierScope != clientkey.TierScopeAll { + t.Fatalf("postgres migrated client key account scope = %+v, err = %v", loadedPoolKey.AccountScope(), err) + } + loadedPoolKey.ProviderScope = clientkey.ProviderScopeBuild | clientkey.ProviderScopeWeb + loadedPoolKey.TierScope = clientkey.TierScopeSuper + loadedPoolKey, err = keyRepository.Update(ctx, loadedPoolKey) + if err != nil || loadedPoolKey.ProviderScope != clientkey.ProviderScopeBuild|clientkey.ProviderScopeWeb || loadedPoolKey.TierScope != clientkey.TierScopeSuper { + t.Fatalf("postgres updated client key account scope = %+v, err = %v", loadedPoolKey.AccountScope(), err) + } + if err := keyRepository.Delete(ctx, poolKey.ID); err != nil { + t.Fatal(err) + } verifyPostgresMediaJobInputConstraintUpgrade(t, ctx, database) repository := NewAccountRepository(database) created, wasCreated, err := repository.UpsertByIdentity(ctx, account.Credential{ @@ -388,6 +415,46 @@ func TestPostgresRepositoriesIntegration(t *testing.T) { } } +func TestPostgresMigratesLegacyClientKeyAccountPoolToScopes(t *testing.T) { + dsn := os.Getenv("TEST_POSTGRES_DSN") + if dsn == "" { + t.Skip("TEST_POSTGRES_DSN is not configured") + } + ctx := context.Background() + database, err := OpenPostgres(ctx, dsn, 10, 2) + if err != nil { + t.Fatal(err) + } + defer database.Close() + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + repository := NewClientKeyRepository(database) + value, err := repository.Create(ctx, clientkey.Key{Name: "postgres-legacy-pool", Prefix: "postgres-legacy-pool", SecretHash: testSecretHash, EncryptedSecret: testEncryptedToken, Enabled: true}) + if err != nil { + t.Fatal(err) + } + if err := database.db.WithContext(ctx).Exec("ALTER TABLE client_keys DROP COLUMN provider_scope_mask, DROP COLUMN tier_scope_mask").Error; err != nil { + t.Fatal(err) + } + if err := database.db.WithContext(ctx).Exec("ALTER TABLE client_keys ADD COLUMN account_pool text NOT NULL DEFAULT 'all' CHECK (account_pool IN ('all','free','super'))").Error; err != nil { + t.Fatal(err) + } + if err := database.db.WithContext(ctx).Table("client_keys").Where("id = ?", value.ID).Update("account_pool", "free").Error; err != nil { + t.Fatal(err) + } + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + stored, err := repository.Get(ctx, value.ID) + if err != nil || stored.ProviderScope != clientkey.ProviderScopeAll || stored.TierScope != clientkey.TierScopeFree { + t.Fatalf("migrated PostgreSQL account scope = %+v, err = %v", stored.AccountScope(), err) + } + if err := repository.Delete(ctx, value.ID); err != nil { + t.Fatal(err) + } +} + func TestPostgresUnhealthyEgressCleanupUsesBothAddressFamilies(t *testing.T) { dsn := os.Getenv("TEST_POSTGRES_DSN") if dsn == "" { diff --git a/backend/internal/infra/persistence/relational/schema.go b/backend/internal/infra/persistence/relational/schema.go index 4b5a581a5..38b079036 100644 --- a/backend/internal/infra/persistence/relational/schema.go +++ b/backend/internal/infra/persistence/relational/schema.go @@ -8,6 +8,7 @@ import ( "strconv" "strings" + clientkeydomain "github.com/chenyme/grok2api/backend/internal/domain/clientkey" "github.com/chenyme/grok2api/backend/internal/domain/media" settingsdomain "github.com/chenyme/grok2api/backend/internal/domain/settings" "gorm.io/gorm" @@ -114,6 +115,10 @@ func (d *Database) InitializeSchema(ctx context.Context) error { func (d *Database) initializeSchema(ctx context.Context) error { db := d.db.WithContext(ctx) + hadClientKeys := db.Migrator().HasTable(&clientKeyModel{}) + hadProviderScope := hadClientKeys && db.Migrator().HasColumn(&clientKeyModel{}, "ProviderScopeMask") + hadTierScope := hadClientKeys && db.Migrator().HasColumn(&clientKeyModel{}, "TierScopeMask") + hadLegacyAccountPool := hadClientKeys && db.Migrator().HasColumn("client_keys", "account_pool") // all 作用域会让 Build 与 Web 共用 UA、健康度和冷却状态,升级时直接移除旧节点。 if db.Migrator().HasTable(&egressNodeModel{}) { if err := db.Where("scope = ?", "all").Delete(&egressNodeModel{}).Error; err != nil { @@ -134,6 +139,9 @@ func (d *Database) initializeSchema(ctx context.Context) error { if migrateErr != nil { return fmt.Errorf("初始化数据库表: %w", migrateErr) } + if err := d.migrateClientKeyAccountScopes(ctx, hadLegacyAccountPool, !hadProviderScope, !hadTierScope); err != nil { + return fmt.Errorf("迁移客户端 Key 调用范围: %w", err) + } if err := d.migrateBuildResponseHeaderTimeout(ctx); err != nil { return fmt.Errorf("迁移 Grok Build 响应头超时: %w", err) } @@ -189,6 +197,22 @@ func (d *Database) initializeSchema(ctx context.Context) error { return nil } +// migrateClientKeyAccountScopes translates the short-lived account_pool +// representation only when the corresponding scope columns are first added. +func (d *Database) migrateClientKeyAccountScopes(ctx context.Context, hadLegacyAccountPool, providerScopeAdded, tierScopeAdded bool) error { + if !hadLegacyAccountPool || (!providerScopeAdded && !tierScopeAdded) { + return nil + } + updates := make(map[string]any, 2) + if providerScopeAdded { + updates["provider_scope_mask"] = uint8(clientkeydomain.ProviderScopeAll) + } + if tierScopeAdded { + updates["tier_scope_mask"] = gorm.Expr("CASE account_pool WHEN 'free' THEN 1 WHEN 'super' THEN 2 ELSE 7 END") + } + return d.db.WithContext(ctx).Table("client_keys").Where("1 = 1").Updates(updates).Error +} + // migrateBuildResponseHeaderTimeout persists the runtime default for settings // rows created before the Build response-header timeout became configurable. func (d *Database) migrateBuildResponseHeaderTimeout(ctx context.Context) error { diff --git a/backend/internal/infra/persistence/relational/schema_client_key_account_pool_upgrade_test.go b/backend/internal/infra/persistence/relational/schema_client_key_account_pool_upgrade_test.go new file mode 100644 index 000000000..60039cbf6 --- /dev/null +++ b/backend/internal/infra/persistence/relational/schema_client_key_account_pool_upgrade_test.go @@ -0,0 +1,91 @@ +package relational + +import ( + "context" + "fmt" + "path/filepath" + "strings" + "testing" + + clientkeydomain "github.com/chenyme/grok2api/backend/internal/domain/clientkey" +) + +func TestInitializeSchemaMigratesLegacyClientKeyAccountPoolToScopes(t *testing.T) { + ctx := context.Background() + database, err := OpenSQLite(ctx, filepath.Join(t.TempDir(), "legacy-client-key-account-pool.db")) + if err != nil { + t.Fatal(err) + } + defer database.Close() + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + repository := NewClientKeyRepository(database) + create := func(prefix string) clientkeydomain.Key { + value, createErr := repository.Create(ctx, clientkeydomain.Key{Name: prefix, Prefix: prefix, SecretHash: testSecretHash, EncryptedSecret: testEncryptedToken, Enabled: true}) + if createErr != nil { + t.Fatal(createErr) + } + return value + } + freeKey := create("legacy-free") + superKey := create("legacy-super") + allKey := create("legacy-all") + if err := recreateClientKeysWithLegacyAccountPool(ctx, database); err != nil { + t.Fatal(err) + } + if err := database.withSQLiteForeignKeysDisabled(ctx, func() error { + db := database.db.WithContext(ctx) + if err := db.Table("client_keys").Where("id = ?", freeKey.ID).Update("account_pool", "free").Error; err != nil { + return err + } + return db.Table("client_keys").Where("id = ?", superKey.ID).Update("account_pool", "super").Error + }); err != nil { + t.Fatal(err) + } + if err := database.InitializeSchema(ctx); err != nil { + t.Fatal(err) + } + assertScope := func(id uint64, wantTier clientkeydomain.TierScope) { + stored, getErr := repository.Get(ctx, id) + if getErr != nil { + t.Fatal(getErr) + } + if stored.ProviderScope != clientkeydomain.ProviderScopeAll || stored.TierScope != wantTier { + t.Fatalf("migrated scope for %d = %+v", id, stored.AccountScope()) + } + } + assertScope(freeKey.ID, clientkeydomain.TierScopeFree) + assertScope(superKey.ID, clientkeydomain.TierScopeSuper) + assertScope(allKey.ID, clientkeydomain.TierScopeAll) + if err := database.db.WithContext(ctx).Model(&clientKeyModel{}).Where("id = ?", allKey.ID).Update("tier_scope_mask", 4).Error; err == nil { + t.Fatal("tier scope constraint accepted an unsupported unknown-only value") + } +} + +func recreateClientKeysWithLegacyAccountPool(ctx context.Context, database *Database) error { + return database.withSQLiteForeignKeysDisabled(ctx, func() error { + db := database.db.WithContext(ctx) + var tableSQL string + if err := db.Raw("SELECT sql FROM sqlite_master WHERE type = 'table' AND name = 'client_keys'").Scan(&tableSQL).Error; err != nil { + return err + } + legacySQL := strings.Replace(tableSQL, ",`provider_scope_mask` integer NOT NULL DEFAULT 7", "", 1) + legacySQL = strings.Replace(legacySQL, ",`tier_scope_mask` integer NOT NULL DEFAULT 7", "", 1) + legacySQL = strings.Replace(legacySQL, ",CONSTRAINT `chk_client_keys_provider_scope` CHECK (provider_scope_mask BETWEEN 1 AND 7)", "", 1) + legacySQL = strings.Replace(legacySQL, ",CONSTRAINT `chk_client_keys_tier_scope` CHECK (tier_scope_mask IN (1,2,3,7))", "", 1) + legacySQL = strings.Replace(legacySQL, ",CONSTRAINT", ",`account_pool` text NOT NULL DEFAULT 'all' CHECK (account_pool IN ('all','free','super')),CONSTRAINT", 1) + legacySQL = strings.Replace(legacySQL, "client_keys", "client_keys_legacy_scope", 1) + if err := db.Exec(legacySQL).Error; err != nil { + return fmt.Errorf("create legacy table with %s: %w", legacySQL, err) + } + const columns = "id,name,prefix,secret_hash,encrypted_secret,enabled,expires_at,rpm_limit,max_concurrent,billing_limit_usd_ticks,billed_usage_usd_ticks,reserved_usage_usd_ticks,allow_model_aliases,last_used_at,created_at,updated_at" + if err := db.Exec("INSERT INTO client_keys_legacy_scope (" + columns + ") SELECT " + columns + " FROM client_keys").Error; err != nil { + return err + } + if err := db.Exec("DROP TABLE client_keys").Error; err != nil { + return err + } + return db.Exec("ALTER TABLE client_keys_legacy_scope RENAME TO client_keys").Error + }) +} diff --git a/backend/internal/infra/runtime/redis/store_integration_test.go b/backend/internal/infra/runtime/redis/store_integration_test.go index fbbc997bc..b0eac3b20 100644 --- a/backend/internal/infra/runtime/redis/store_integration_test.go +++ b/backend/internal/infra/runtime/redis/store_integration_test.go @@ -452,6 +452,18 @@ func TestRedisInvalidationBusIntegration(t *testing.T) { case <-time.After(time.Second): t.Fatal("second invalidation notification was not delivered") } + clientKeyEvent := repository.InvalidationEvent{Kind: repository.InvalidationClientKeyChanged, ClientKeyID: 42, SourceInstance: "instance-a"} + if err := store.PublishInvalidation(ctx, clientKeyEvent); err != nil { + t.Fatal(err) + } + select { + case clientKeyInvalidation := <-received: + if clientKeyInvalidation.Layer() != repository.InvalidationLayerClientKey || clientKeyInvalidation.ClientKeyID != 42 || clientKeyInvalidation.Revision == 0 { + t.Fatalf("client-key invalidation = %#v", clientKeyInvalidation) + } + case <-time.After(time.Second): + t.Fatal("client-key invalidation notification was not delivered") + } cancel() if err := <-done; err != nil { t.Fatal(err) diff --git a/backend/internal/repository/list.go b/backend/internal/repository/list.go index feb9b506e..e7ac0f3f3 100644 --- a/backend/internal/repository/list.go +++ b/backend/internal/repository/list.go @@ -103,8 +103,11 @@ type AccountSummary struct { } type ModelListFilter struct { - Provider string - Enabled *bool + Provider string + Providers []string + Tiers []string + Enabled *bool + ActiveScope bool } type ModelListQuery struct { diff --git a/backend/internal/repository/model.go b/backend/internal/repository/model.go index ac6a516fa..63bef5da8 100644 --- a/backend/internal/repository/model.go +++ b/backend/internal/repository/model.go @@ -12,6 +12,7 @@ import ( type ModelRepository interface { List(ctx context.Context, query ModelListQuery) ([]model.Route, int64, error) ListEnabled(ctx context.Context) ([]model.Route, error) + ListEnabledForScope(ctx context.Context, filter ModelListFilter) ([]model.Route, error) ListConfiguredEnabled(ctx context.Context) ([]model.Route, error) Get(ctx context.Context, id uint64) (model.Route, error) GetByPublicID(ctx context.Context, publicID string) (model.Route, error) diff --git a/backend/internal/repository/runtime.go b/backend/internal/repository/runtime.go index 98adc59a6..56ab7c8f6 100644 --- a/backend/internal/repository/runtime.go +++ b/backend/internal/repository/runtime.go @@ -98,20 +98,23 @@ const ( InvalidationAccountQuotaChanged InvalidationKind = "account_quota_changed" InvalidationAccountRecoveryChanged InvalidationKind = "account_recovery_changed" InvalidationAccountModelQuotaChanged InvalidationKind = "account_model_quota_changed" + InvalidationClientKeyChanged InvalidationKind = "client_key_changed" ) type InvalidationLayer string const ( - InvalidationLayerRoute InvalidationLayer = "route" - InvalidationLayerBase InvalidationLayer = "account_base" - InvalidationLayerOverlay InvalidationLayer = "account_overlay" + InvalidationLayerRoute InvalidationLayer = "route" + InvalidationLayerBase InvalidationLayer = "account_base" + InvalidationLayerOverlay InvalidationLayer = "account_overlay" + InvalidationLayerClientKey InvalidationLayer = "client_key" ) type InvalidationEvent struct { Kind InvalidationKind `json:"kind"` Provider account.Provider `json:"provider,omitempty"` AccountID uint64 `json:"accountId,omitempty"` + ClientKeyID uint64 `json:"clientKeyId,omitempty"` UpstreamModel string `json:"upstreamModel,omitempty"` Revision uint64 `json:"revision,omitempty"` SourceInstance string `json:"sourceInstance,omitempty"` @@ -126,15 +129,21 @@ func (e InvalidationEvent) Layer() InvalidationLayer { return InvalidationLayerOverlay case InvalidationAccountStateChanged, InvalidationAccountCredentialChanged, InvalidationAccountBillingChanged, InvalidationAccountQuotaChanged, InvalidationAccountRecoveryChanged: return InvalidationLayerBase + case InvalidationClientKeyChanged: + return InvalidationLayerClientKey default: return "" } } func (e InvalidationEvent) Valid() bool { - if e.Layer() == "" { + layer := e.Layer() + if layer == "" { return false } + if layer == InvalidationLayerClientKey { + return e.Provider == "" && e.AccountID == 0 && e.UpstreamModel == "" + } switch e.Provider { case "", account.ProviderBuild, account.ProviderWeb, account.ProviderConsole: return true diff --git a/backend/internal/transport/http/clientkey/handler.go b/backend/internal/transport/http/clientkey/handler.go index 791f997f9..1a6d160e9 100644 --- a/backend/internal/transport/http/clientkey/handler.go +++ b/backend/internal/transport/http/clientkey/handler.go @@ -5,6 +5,7 @@ import ( "fmt" "net/http" "strconv" + "strings" "time" clientkeyapp "github.com/chenyme/grok2api/backend/internal/application/clientkey" @@ -29,14 +30,17 @@ func (h *Handler) Register(router *gin.RouterGroup) { } type createRequest struct { - Name string `json:"name" binding:"required"` - Enabled *bool `json:"enabled"` - ExpiresAt string `json:"expiresAt"` - RPMLimit *int `json:"rpmLimit"` - MaxConcurrent *int `json:"maxConcurrent"` - BillingLimitUSDTicks int64 `json:"billingLimitUsdTicks"` - AllowModelAliases *bool `json:"allowModelAliases"` - AllowedModelIDs []string `json:"allowedModelIds"` + Name string `json:"name" binding:"required"` + Enabled *bool `json:"enabled"` + ExpiresAt string `json:"expiresAt"` + RPMLimit *int `json:"rpmLimit"` + MaxConcurrent *int `json:"maxConcurrent"` + BillingLimitUSDTicks int64 `json:"billingLimitUsdTicks"` + AllowModelAliases *bool `json:"allowModelAliases"` + AllowedModelIDs []string `json:"allowedModelIds"` + ProviderScope *[]string `json:"providerScope"` + TierScope *[]string `json:"tierScope"` + AccountPool *string `json:"accountPool"` } type updateRequest struct { @@ -48,6 +52,9 @@ type updateRequest struct { BillingLimitUSDTicks *int64 `json:"billingLimitUsdTicks"` AllowModelAliases *bool `json:"allowModelAliases"` AllowedModelIDs *[]string `json:"allowedModelIds"` + ProviderScope *[]string `json:"providerScope"` + TierScope *[]string `json:"tierScope"` + AccountPool *string `json:"accountPool"` } type batchUpdateRequest struct { @@ -71,6 +78,8 @@ type keyResponse struct { BilledUsageUSDTicks int64 `json:"billedUsageUsdTicks"` AllowModelAliases bool `json:"allowModelAliases"` AllowedModelIDs []string `json:"allowedModelIds"` + ProviderScope []string `json:"providerScope"` + TierScope []string `json:"tierScope"` LastUsedAt *time.Time `json:"lastUsedAt,omitempty"` } @@ -150,7 +159,18 @@ func (h *Handler) create(c *gin.Context) { if request.Enabled != nil { enabled = *request.Enabled } + providerScope, tierScope, err := parseRequestedScopes(request.ProviderScope, request.TierScope, request.AccountPool) + if err != nil { + response.Error(c, http.StatusBadRequest, "invalidAccountScope", err.Error()) + return + } input := clientkeyapp.CreateInput{Name: request.Name, Enabled: enabled, ExpiresAt: expiresAt, BillingLimitUSDTicks: request.BillingLimitUSDTicks, AllowedModels: modelIDs} + if providerScope != nil { + input.ProviderScope = *providerScope + } + if tierScope != nil { + input.TierScope = *tierScope + } if request.AllowModelAliases != nil { input.AllowModelAliases = *request.AllowModelAliases } @@ -180,7 +200,12 @@ func (h *Handler) update(c *gin.Context) { response.Error(c, http.StatusBadRequest, "invalidRequest", "请求参数无效") return } - input := clientkeyapp.UpdateInput{Name: request.Name, Enabled: request.Enabled, RPMLimit: request.RPMLimit, MaxConcurrent: request.MaxConcurrent, BillingLimitUSDTicks: request.BillingLimitUSDTicks, AllowModelAliases: request.AllowModelAliases} + providerScope, tierScope, err := parseRequestedScopes(request.ProviderScope, request.TierScope, request.AccountPool) + if err != nil { + response.Error(c, http.StatusBadRequest, "invalidAccountScope", err.Error()) + return + } + input := clientkeyapp.UpdateInput{Name: request.Name, Enabled: request.Enabled, RPMLimit: request.RPMLimit, MaxConcurrent: request.MaxConcurrent, BillingLimitUSDTicks: request.BillingLimitUSDTicks, AllowModelAliases: request.AllowModelAliases, ProviderScope: providerScope, TierScope: tierScope} if request.ExpiresAt != nil { if *request.ExpiresAt == "" { input.ClearExpiresAt = true @@ -260,8 +285,47 @@ func newKeyResponse(value clientkeydomain.Key) keyResponse { return keyResponse{ ID: value.ID, Name: value.Name, Prefix: value.Prefix, Enabled: value.Enabled, ExpiresAt: value.ExpiresAt, RPMLimit: value.RPMLimit, MaxConcurrent: value.MaxConcurrent, BillingLimitUSDTicks: value.BillingLimitUSDTicks, - BilledUsageUSDTicks: value.BilledUsageUSDTicks, AllowModelAliases: value.AllowModelAliases, AllowedModelIDs: ids, LastUsedAt: value.LastUsedAt, + BilledUsageUSDTicks: value.BilledUsageUSDTicks, AllowModelAliases: value.AllowModelAliases, AllowedModelIDs: ids, + ProviderScope: value.ProviderScope.Values(), TierScope: value.TierScope.Values(), LastUsedAt: value.LastUsedAt, + } +} + +func parseRequestedScopes(providerValues, tierValues *[]string, legacyPool *string) (*clientkeydomain.ProviderScope, *clientkeydomain.TierScope, error) { + if legacyPool != nil && (providerValues != nil || tierValues != nil) { + return nil, nil, errors.New("accountPool 不能与 providerScope 或 tierScope 同时设置") + } + if legacyPool != nil { + providers := clientkeydomain.ProviderScopeAll + var tiers clientkeydomain.TierScope + switch strings.TrimSpace(*legacyPool) { + case "all": + tiers = clientkeydomain.TierScopeAll + case "free": + tiers = clientkeydomain.TierScopeFree + case "super": + tiers = clientkeydomain.TierScopeSuper + default: + return nil, nil, errors.New("accountPool 必须是 all、free 或 super") + } + return &providers, &tiers, nil + } + var providers *clientkeydomain.ProviderScope + if providerValues != nil { + value, valid := clientkeydomain.ParseProviderScopeValues(*providerValues) + if !valid { + return nil, nil, errors.New("providerScope 必须是 all,或 grok_build、grok_web、grok_console 的组合") + } + providers = &value + } + var tiers *clientkeydomain.TierScope + if tierValues != nil { + value, valid := clientkeydomain.ParseTierScopeValues(*tierValues) + if !valid { + return nil, nil, errors.New("tierScope 必须是 all,或 free、super 的组合") + } + tiers = &value } + return providers, tiers, nil } func parseTime(value string) (*time.Time, error) { diff --git a/backend/internal/transport/http/clientkey/handler_test.go b/backend/internal/transport/http/clientkey/handler_test.go index c2c8668f9..7403d37da 100644 --- a/backend/internal/transport/http/clientkey/handler_test.go +++ b/backend/internal/transport/http/clientkey/handler_test.go @@ -46,6 +46,8 @@ func TestCreateDistinguishesOmittedLimitsFromExplicitZero(t *testing.T) { } assertCreate(`{"name":"defaults"}`) assertCreate(`{"name":"unlimited","rpmLimit":0,"maxConcurrent":0}`) + assertCreate(`{"name":"free-pool","accountPool":"free"}`) + assertCreate(`{"name":"mixed-scope","providerScope":["grok_build","grok_web"],"tierScope":["free","super"]}`) defaults, total, err := service.List(ctx, 1, 20, "defaults", clientkeyapp.ListFilter{}) if err != nil || total != 1 || len(defaults) != 1 { @@ -61,4 +63,39 @@ func TestCreateDistinguishesOmittedLimitsFromExplicitZero(t *testing.T) { if unlimited[0].RPMLimit != 0 || unlimited[0].MaxConcurrent != 0 { t.Fatalf("explicit zero limits = rpm %d, concurrency %d", unlimited[0].RPMLimit, unlimited[0].MaxConcurrent) } + freePool, total, err := service.List(ctx, 1, 20, "free-pool", clientkeyapp.ListFilter{}) + if err != nil || total != 1 || len(freePool) != 1 || freePool[0].ProviderScope != 7 || freePool[0].TierScope != 1 { + t.Fatalf("free-pool key list = %#v, total = %d, err = %v", freePool, total, err) + } + mixedScope, total, err := service.List(ctx, 1, 20, "mixed-scope", clientkeyapp.ListFilter{}) + if err != nil || total != 1 || len(mixedScope) != 1 || mixedScope[0].ProviderScope != 3 || mixedScope[0].TierScope != 3 { + t.Fatalf("mixed-scope key list = %#v, total = %d, err = %v", mixedScope, total, err) + } + request := httptest.NewRequest(http.MethodPost, "/api/client-keys", bytes.NewBufferString(`{"name":"invalid-pool","accountPool":"unknown"}`)) + request.Header.Set("Content-Type", "application/json") + response := httptest.NewRecorder() + router.ServeHTTP(response, request) + if response.Code != http.StatusBadRequest { + t.Fatalf("invalid pool response = %d %s", response.Code, response.Body.String()) + } + request = httptest.NewRequest(http.MethodPost, "/api/client-keys", bytes.NewBufferString(`{"name":"ambiguous","accountPool":"free","tierScope":["super"]}`)) + request.Header.Set("Content-Type", "application/json") + response = httptest.NewRecorder() + router.ServeHTTP(response, request) + if response.Code != http.StatusBadRequest { + t.Fatalf("ambiguous scope response = %d %s", response.Code, response.Body.String()) + } + for _, body := range []string{ + `{"name":"empty-providers","providerScope":[]}`, + `{"name":"empty-tiers","tierScope":[]}`, + `{"name":"empty-legacy-pool","accountPool":""}`, + } { + request = httptest.NewRequest(http.MethodPost, "/api/client-keys", bytes.NewBufferString(body)) + request.Header.Set("Content-Type", "application/json") + response = httptest.NewRecorder() + router.ServeHTTP(response, request) + if response.Code != http.StatusBadRequest { + t.Fatalf("empty scope response = %d %s", response.Code, response.Body.String()) + } + } } diff --git a/backend/internal/transport/http/inference/handler.go b/backend/internal/transport/http/inference/handler.go index c3f9e7ff6..db334fe34 100644 --- a/backend/internal/transport/http/inference/handler.go +++ b/backend/internal/transport/http/inference/handler.go @@ -176,17 +176,30 @@ type modelListItem struct { } func (h *Handler) listModels(c *gin.Context) { - values, err := h.models.ListEnabled(c.Request.Context()) - if err != nil { - writeOpenAIError(c, http.StatusInternalServerError, "model_list_failed", "读取模型列表失败") - return - } allowAliases := false + var clientKey clientkeydomain.Key + hasClientKey := false if clientValue, exists := c.Get(middleware.ClientKey); exists { - if clientKey, ok := clientValue.(clientkeydomain.Key); ok { + if value, ok := clientValue.(clientkeydomain.Key); ok { + clientKey = value + hasClientKey = true allowAliases = clientKey.AllowModelAliases } } + var values []modeldomain.Route + var err error + if hasClientKey { + values, err = h.models.ListEnabledForClientKey(c.Request.Context(), clientKey) + } else { + values, err = h.models.ListEnabled(c.Request.Context()) + } + if err != nil { + writeOpenAIError(c, http.StatusInternalServerError, "model_list_failed", "读取模型列表失败") + return + } + if hasClientKey { + values = filterModelRoutesForClientKey(values, clientKey) + } items := newModelListItems(values) if allowAliases { items = appendReasoningModelAliases(items) @@ -198,6 +211,17 @@ func (h *Handler) listModels(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"object": "list", "data": items}) } +func filterModelRoutesForClientKey(values []modeldomain.Route, key clientkeydomain.Key) []modeldomain.Route { + filtered := make([]modeldomain.Route, 0, len(values)) + scope := key.AccountScope() + for _, value := range values { + if scope.AllowsProvider(value.Provider) && key.AllowsModel(value.ID) { + filtered = append(filtered, value) + } + } + return filtered +} + // newModelListItems deduplicates by downstream public name and hides Provider prefixes used only for internal routing. func newModelListItems(values []modeldomain.Route) []modelListItem { data := make([]modelListItem, 0, len(values)) @@ -1636,6 +1660,9 @@ func writeGatewayError(c *gin.Context, err error) { case errors.Is(err, clientkeyapp.ErrBillingLimit): status, code = http.StatusTooManyRequests, "billing_limit_exceeded" message = clientkeyapp.ErrBillingLimit.Error() + case errors.Is(err, clientkeyapp.ErrModelNotAllowed): + status, code = http.StatusForbidden, "model_not_allowed" + message = clientkeyapp.ErrModelNotAllowed.Error() case errors.Is(err, gateway.ErrModelNotFound): status, code = http.StatusNotFound, "model_not_found" message = "模型不存在" @@ -1684,6 +1711,9 @@ func writeGatewayAnthropicError(c *gin.Context, err error) { case errors.Is(err, clientkeyapp.ErrBillingLimit): status, errorType = http.StatusTooManyRequests, "rate_limit_error" message = clientkeyapp.ErrBillingLimit.Error() + case errors.Is(err, clientkeyapp.ErrModelNotAllowed): + status, errorType, clientCode = http.StatusForbidden, "permission_error", "model_not_allowed" + message = clientkeyapp.ErrModelNotAllowed.Error() case errors.Is(err, gateway.ErrModelNotFound): status, errorType = http.StatusNotFound, "not_found_error" message = "模型不存在" @@ -1710,7 +1740,7 @@ func writeGatewayAnthropicError(c *gin.Context, err error) { errorType = "rate_limit_error" } case errors.As(err, &selectionFailure): - status, _, message = selectionErrorResponse(c, selectionFailure) + status, clientCode, message = selectionErrorResponse(c, selectionFailure) if status == http.StatusTooManyRequests { errorType = "rate_limit_error" } else { @@ -1738,17 +1768,21 @@ func selectionErrorResponse(c *gin.Context, failure *gateway.SelectionUnavailabl return status, code, message } status, code = failure.HTTPStatus(), failure.Code() - switch failure.Reason { - case gateway.SelectionCooling: - message = "上游账号正在冷却" - case gateway.SelectionModelCooling: - message = "上游账号的目标模型正在冷却" - case gateway.SelectionQuotaExhausted: - message = "上游账号额度等待恢复" - case gateway.SelectionSaturated: - message = "上游账号当前均达到并发上限" - case gateway.SelectionUnsupportedModel: - message = "当前账号池不支持该模型" + if failure.Scope.IsRestricted() { + message = failure.Error() + } else { + switch failure.Reason { + case gateway.SelectionCooling: + message = "上游账号正在冷却" + case gateway.SelectionModelCooling: + message = "上游账号的目标模型正在冷却" + case gateway.SelectionQuotaExhausted: + message = "上游账号额度等待恢复" + case gateway.SelectionSaturated: + message = "上游账号当前均达到并发上限" + case gateway.SelectionUnsupportedModel: + message = "当前账号池不支持该模型" + } } if failure.RetryAfter > 0 { seconds := max(int64(1), int64((failure.RetryAfter+time.Second-1)/time.Second)) diff --git a/backend/internal/transport/http/inference/handler_test.go b/backend/internal/transport/http/inference/handler_test.go index 71e96efeb..bc5d7e39d 100644 --- a/backend/internal/transport/http/inference/handler_test.go +++ b/backend/internal/transport/http/inference/handler_test.go @@ -13,7 +13,9 @@ import ( "time" "unicode/utf8" + clientkeyapp "github.com/chenyme/grok2api/backend/internal/application/clientkey" "github.com/chenyme/grok2api/backend/internal/application/gateway" + clientkeydomain "github.com/chenyme/grok2api/backend/internal/domain/clientkey" mediadomain "github.com/chenyme/grok2api/backend/internal/domain/media" "github.com/gin-gonic/gin" ) @@ -165,6 +167,34 @@ func TestGatewayErrorMapsLedgerUnavailableToServiceUnavailable(t *testing.T) { } } +func TestGatewayErrorMapsDisallowedModelWithoutCallingItUpstreamUnavailable(t *testing.T) { + gin.SetMode(gin.TestMode) + for _, test := range []struct { + name string + anthropic bool + wantType string + }{ + {name: "openai", wantType: `"code":"model_not_allowed"`}, + {name: "anthropic", anthropic: true, wantType: `"type":"permission_error"`}, + } { + t.Run(test.name, func(t *testing.T) { + router := gin.New() + router.GET("/", func(c *gin.Context) { + if test.anthropic { + writeGatewayAnthropicError(c, clientkeyapp.ErrModelNotAllowed) + return + } + writeGatewayError(c, clientkeyapp.ErrModelNotAllowed) + }) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/", nil)) + if recorder.Code != http.StatusForbidden || !strings.Contains(recorder.Body.String(), test.wantType) || strings.Contains(recorder.Body.String(), "upstream_unavailable") { + t.Fatalf("status=%d body=%s", recorder.Code, recorder.Body.String()) + } + }) + } +} + func TestGatewayErrorMapsResponseHeaderTimeout(t *testing.T) { gin.SetMode(gin.TestMode) openAIRouter := gin.New() @@ -954,6 +984,7 @@ func TestSelectionErrorResponseDistinguishesCoolingAndSaturation(t *testing.T) { {name: "cooling", failure: &gateway.SelectionUnavailableError{Reason: gateway.SelectionCooling, RetryAfter: 1500 * time.Millisecond}, status: http.StatusTooManyRequests, code: "upstream_cooling", retryAfter: "2"}, {name: "model cooling", failure: &gateway.SelectionUnavailableError{Reason: gateway.SelectionModelCooling, RetryAfter: time.Second}, status: http.StatusTooManyRequests, code: "upstream_model_cooling", retryAfter: "1"}, {name: "saturated", failure: &gateway.SelectionUnavailableError{Reason: gateway.SelectionSaturated, RetryAfter: time.Second}, status: http.StatusServiceUnavailable, code: "upstream_saturated", retryAfter: "1"}, + {name: "scoped account range", failure: &gateway.SelectionUnavailableError{Reason: gateway.SelectionNoAccounts, Scope: clientkeydomain.AccountScope{Providers: clientkeydomain.ProviderScopeBuild, Tiers: clientkeydomain.TierScopeFree}}, status: http.StatusServiceUnavailable, code: "client_key_account_scope_unavailable"}, } { t.Run(test.name, func(t *testing.T) { recorder := httptest.NewRecorder() diff --git a/backend/internal/transport/http/inference/model_list_test.go b/backend/internal/transport/http/inference/model_list_test.go index 371cc388c..639739e85 100644 --- a/backend/internal/transport/http/inference/model_list_test.go +++ b/backend/internal/transport/http/inference/model_list_test.go @@ -8,6 +8,7 @@ import ( "time" "github.com/chenyme/grok2api/backend/internal/domain/account" + clientkeydomain "github.com/chenyme/grok2api/backend/internal/domain/clientkey" modeldomain "github.com/chenyme/grok2api/backend/internal/domain/model" "github.com/gin-gonic/gin" ) @@ -24,6 +25,19 @@ func TestNewModelListItemsDeduplicatesSharedPublicName(t *testing.T) { } } +func TestFilterModelRoutesForClientKeyUsesProviderAndModelIntersection(t *testing.T) { + routes := []modeldomain.Route{ + {ID: 1, Provider: account.ProviderBuild, PublicID: "Build/grok-shared"}, + {ID: 2, Provider: account.ProviderWeb, PublicID: "Web/grok-shared"}, + {ID: 3, Provider: account.ProviderConsole, PublicID: "Console/grok-shared"}, + } + key := clientkeydomain.Key{ProviderScope: clientkeydomain.ProviderScopeWeb | clientkeydomain.ProviderScopeConsole, AllowedModels: []uint64{2, 3}} + filtered := filterModelRoutesForClientKey(routes, key) + if len(filtered) != 2 || filtered[0].ID != 2 || filtered[1].ID != 3 { + t.Fatalf("filtered routes = %#v", filtered) + } +} + func TestAppendReasoningModelAliasesUsesRealSupportedLevels(t *testing.T) { now := time.Unix(100, 0).UTC() base := newModelListItems([]modeldomain.Route{ diff --git a/backend/internal/transport/http/model/handler.go b/backend/internal/transport/http/model/handler.go index 726c21821..6a1852812 100644 --- a/backend/internal/transport/http/model/handler.go +++ b/backend/internal/transport/http/model/handler.go @@ -5,6 +5,7 @@ import ( "fmt" "net/http" "strconv" + "strings" "time" modelapp "github.com/chenyme/grok2api/backend/internal/application/model" @@ -79,7 +80,12 @@ type accountOptionResponse struct { func (h *Handler) list(c *gin.Context) { page, pageSize := pagination(c) - values, total, err := h.service.List(c.Request.Context(), page, pageSize, c.Query("search"), modelapp.ListFilter{Provider: c.Query("provider"), Status: c.Query("status"), Sort: repository.SortQuery{Field: c.Query("sortBy"), Direction: repository.SortDirection(c.Query("sortOrder"))}}) + activeScope, ok := parseOptionalBool(c.Query("activeScope")) + if !ok { + response.Error(c, http.StatusBadRequest, "invalidFilter", "activeScope 必须是 true 或 false") + return + } + values, total, err := h.service.List(c.Request.Context(), page, pageSize, c.Query("search"), modelapp.ListFilter{Provider: c.Query("provider"), Providers: splitScopeQuery(c.QueryArray("providerScope")), Tiers: splitScopeQuery(c.QueryArray("tierScope")), Status: c.Query("status"), ActiveScope: activeScope, Sort: repository.SortQuery{Field: c.Query("sortBy"), Direction: repository.SortDirection(c.Query("sortOrder"))}}) if errors.Is(err, modelapp.ErrInvalidFilter) { response.Error(c, http.StatusBadRequest, "invalidFilter", err.Error()) return @@ -95,6 +101,33 @@ func (h *Handler) list(c *gin.Context) { response.Success(c, http.StatusOK, gin.H{"items": items, "page": page, "pageSize": pageSize, "total": total}) } +func splitScopeQuery(values []string) []string { + if len(values) == 0 { + return nil + } + result := make([]string, 0, len(values)) + for _, value := range values { + for _, item := range strings.Split(value, ",") { + item = strings.TrimSpace(item) + if item != "" { + result = append(result, item) + } + } + } + return result +} + +func parseOptionalBool(value string) (bool, bool) { + switch strings.TrimSpace(value) { + case "", "false": + return false, true + case "true": + return true, true + default: + return false, false + } +} + func (h *Handler) listAccounts(c *gin.Context) { values, err := h.service.ListBindableAccounts(c.Request.Context(), account.Provider(c.Query("provider"))) if err != nil { diff --git a/backend/internal/transport/http/model/handler_test.go b/backend/internal/transport/http/model/handler_test.go index 8d1e8135d..abb6713f7 100644 --- a/backend/internal/transport/http/model/handler_test.go +++ b/backend/internal/transport/http/model/handler_test.go @@ -16,3 +16,22 @@ func TestNewModelResponseSeparatesPublicAndUpstreamNames(t *testing.T) { t.Fatalf("model response = %#v", response) } } + +func TestParseOptionalBoolRejectsAmbiguousValues(t *testing.T) { + for _, test := range []struct { + input string + value bool + valid bool + }{ + {input: "", valid: true}, + {input: "false", valid: true}, + {input: "true", value: true, valid: true}, + {input: "1", valid: false}, + {input: "yes", valid: false}, + } { + value, valid := parseOptionalBool(test.input) + if value != test.value || valid != test.valid { + t.Fatalf("parseOptionalBool(%q) = (%v, %v), want (%v, %v)", test.input, value, valid, test.value, test.valid) + } + } +} diff --git a/frontend/src/entities/model/model-api.ts b/frontend/src/entities/model/model-api.ts index 7d0957c49..049126c95 100644 --- a/frontend/src/entities/model/model-api.ts +++ b/frontend/src/entities/model/model-api.ts @@ -9,6 +9,9 @@ type ListModelsInput = { search?: string; status?: string; provider?: "grok_build" | "grok_web" | "grok_console" | ""; + providerScope?: Array<"grok_build" | "grok_web" | "grok_console">; + tierScope?: Array<"free" | "super">; + activeScope?: boolean; sortBy?: string; sortOrder?: SortOrder; }; @@ -45,6 +48,9 @@ export function listModels(input: ListModelsInput): Promise("client key", { id: isString, name: isString, prefix: isString, enabled: isBoolean, expiresAt: isOptional(isString), rpmLimit: isNumber, maxConcurrent: isNumber, billingLimitUsdTicks: isNumber, billedUsageUsdTicks: isNumber, - allowModelAliases: isBoolean, allowedModelIds: isArrayOf(isString), lastUsedAt: isOptional(isString), + allowModelAliases: isBoolean, allowedModelIds: isArrayOf(isString), providerScope: isOptional(isArrayOf(isOneOf("all", "grok_build", "grok_web", "grok_console"))), tierScope: isOptional(isArrayOf(isOneOf("all", "free", "super"))), lastUsedAt: isOptional(isString), }); const decodeClientKeyPage = createPaginatedDecoder(clientKeyValidator); const decodeCreatedClientKey = createObjectDecoder("created client key", { key: clientKeyValidator, secret: isString }); diff --git a/frontend/src/features/client-keys/client-keys-page.tsx b/frontend/src/features/client-keys/client-keys-page.tsx index 22edf90dd..7013f734b 100644 --- a/frontend/src/features/client-keys/client-keys-page.tsx +++ b/frontend/src/features/client-keys/client-keys-page.tsx @@ -1,6 +1,6 @@ import { zodResolver } from "@hookform/resolvers/zod"; import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query"; -import { ChevronLeft, ChevronRight, CircleHelp, Copy, MoreHorizontal, Pencil, Plus, Search, Trash2 } from "lucide-react"; +import { ChevronDown, ChevronLeft, ChevronRight, CircleHelp, Copy, MoreHorizontal, Pencil, Plus, Search, Trash2 } from "lucide-react"; import { useState } from "react"; import { Controller, useForm, useWatch } from "react-hook-form"; import { useTranslation } from "react-i18next"; @@ -14,7 +14,7 @@ import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Checkbox } from "@/components/ui/checkbox"; import { Dialog, DialogContent, DialogDescription, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog"; -import { DropdownMenu, DropdownMenuContent, DropdownMenuItem, DropdownMenuSeparator, DropdownMenuTrigger } from "@/components/ui/dropdown-menu"; +import { DropdownMenu, DropdownMenuCheckboxItem, DropdownMenuContent, DropdownMenuItem, DropdownMenuRadioGroup, DropdownMenuRadioItem, DropdownMenuSeparator, DropdownMenuTrigger } from "@/components/ui/dropdown-menu"; import { Input } from "@/components/ui/input"; import { Label } from "@/components/ui/label"; import { Switch } from "@/components/ui/switch"; @@ -22,7 +22,7 @@ import { Spinner } from "@/components/ui/spinner"; import { Table, TableActionCell, TableActionHead, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; import { listModels } from "@/entities/model/model-api"; -import { createClientKey, deleteClientKey, deleteClientKeys, getClientKeySecret, listClientKeys, updateClientKey, updateClientKeysEnabled, type ClientKeyDTO, type CreateKeyResponseDTO } from "@/features/client-keys/client-keys-api"; +import { createClientKey, deleteClientKey, deleteClientKeys, getClientKeySecret, listClientKeys, updateClientKey, updateClientKeysEnabled, type ClientKeyDTO, type CreateKeyResponseDTO, type ProviderScopeValue, type TierScopeValue } from "@/features/client-keys/client-keys-api"; import { EmptyState, ErrorState, LoadingState, TableLoadingRow } from "@/shared/components/data-state"; import { DataTableShell } from "@/shared/components/data-table-shell"; import { DataTableFilters } from "@/shared/components/data-table-filters"; @@ -74,33 +74,47 @@ export function ClientKeysPage() { billingUnlimited: z.boolean(), billingLimitUsd: z.number().min(0.01, t("errors.positive")).max(MAX_BILLING_LIMIT_USD), allowModelAliases: z.boolean(), + modelScopeMode: z.enum(["all", "restricted"]), allowedModelIds: z.array(z.string()), + providerScope: z.array(z.enum(["all", "grok_build", "grok_web", "grok_console"])).min(1), + tierScope: z.array(z.enum(["all", "free", "super"])).min(1), }).superRefine((value, context) => { if (!value.expiryUnlimited && !value.expiresAt) { context.addIssue({ code: "custom", path: ["expiresAt"], message: t("errors.required") }); } + if (value.modelScopeMode === "restricted" && value.allowedModelIds.length === 0) { + context.addIssue({ code: "custom", path: ["allowedModelIds"], message: t("keys.selectModelRequired") }); + } }); type KeyForm = z.infer; const form = useForm({ resolver: zodResolver(schema), - defaultValues: { name: "", enabled: true, expiryUnlimited: true, expiresAt: "", rpmUnlimited: false, rpmLimit: 120, concurrencyUnlimited: false, maxConcurrent: 8, billingUnlimited: true, billingLimitUsd: 10, allowModelAliases: false, allowedModelIds: [] }, + defaultValues: { name: "", enabled: true, expiryUnlimited: true, expiresAt: "", rpmUnlimited: false, rpmLimit: 120, concurrencyUnlimited: false, maxConcurrent: 8, billingUnlimited: true, billingLimitUsd: 10, allowModelAliases: false, modelScopeMode: "all", allowedModelIds: [], providerScope: ["all"], tierScope: ["all"] }, }); const keyEnabled = useWatch({ control: form.control, name: "enabled" }); const allowModelAliases = useWatch({ control: form.control, name: "allowModelAliases" }); + const modelScopeMode = useWatch({ control: form.control, name: "modelScopeMode" }); const selectedModels = useWatch({ control: form.control, name: "allowedModelIds" }); + const providerScope = useWatch({ control: form.control, name: "providerScope" }); + const tierScope = useWatch({ control: form.control, name: "tierScope" }); const expiryUnlimited = useWatch({ control: form.control, name: "expiryUnlimited" }); const rpmUnlimited = useWatch({ control: form.control, name: "rpmUnlimited" }); const concurrencyUnlimited = useWatch({ control: form.control, name: "concurrencyUnlimited" }); const billingUnlimited = useWatch({ control: form.control, name: "billingUnlimited" }); + const modelProviderScope = providerScope.filter((value): value is Exclude => value !== "all"); + const modelTierScope = tierScope.filter((value): value is Exclude => value !== "all"); + const providerScopeSummary = providerScope.includes("all") ? t("keys.allProviders") : modelProviderScope.map((value) => ({ grok_build: "Build", grok_web: "Web", grok_console: "Console" })[value]).join(" · "); + const tierScopeSummary = tierScope.includes("all") ? t("keys.allTiers") : modelTierScope.map((value) => value === "free" ? "Free" : "Super").join(" · "); + const modelScopeSummary = modelScopeMode === "all" ? t("keys.allModels") : t("keys.selectedModels", { count: selectedModels.length }); const keysQuery = useQuery({ queryKey: ["client-keys", page, pageSize, debouncedSearch, statusFilter, modelScopeFilter, sort.field, sort.order], queryFn: () => listClientKeys({ page, pageSize, search: debouncedSearch, status: statusFilter, modelScope: modelScopeFilter, sortBy: sort.field || undefined, sortOrder: sort.field ? sort.order : undefined }), }); const modelsQuery = useQuery({ - queryKey: ["models", "options", modelOptionsPage, debouncedModelOptionsSearch], - queryFn: () => listModels({ page: modelOptionsPage, pageSize: 50, search: debouncedModelOptionsSearch }), - enabled: editing !== null, + queryKey: ["models", "options", modelOptionsPage, debouncedModelOptionsSearch, modelProviderScope.join(","), modelTierScope.join(",")], + queryFn: () => listModels({ page: modelOptionsPage, pageSize: 50, search: debouncedModelOptionsSearch, providerScope: modelProviderScope, tierScope: modelTierScope }), + enabled: editing !== null && modelScopeMode === "restricted", }); const saveMutation = useMutation({ @@ -113,6 +127,8 @@ export function ClientKeysPage() { billingLimitUsdTicks: values.billingUnlimited ? 0 : Math.round(values.billingLimitUsd * USD_TICKS), allowModelAliases: values.allowModelAliases, allowedModelIds: values.allowedModelIds, + providerScope: values.providerScope, + tierScope: values.tierScope, expiresAt: values.expiryUnlimited ? "" : new Date(values.expiresAt).toISOString(), }; if (editing === "new") { @@ -179,7 +195,7 @@ export function ClientKeysPage() { setEditing("new"); setModelOptionsPage(1); setModelOptionsSearch(""); - form.reset({ name: "", enabled: true, expiryUnlimited: true, expiresAt: "", rpmUnlimited: false, rpmLimit: 120, concurrencyUnlimited: false, maxConcurrent: 8, billingUnlimited: true, billingLimitUsd: 10, allowModelAliases: false, allowedModelIds: [] }); + form.reset({ name: "", enabled: true, expiryUnlimited: true, expiresAt: "", rpmUnlimited: false, rpmLimit: 120, concurrencyUnlimited: false, maxConcurrent: 8, billingUnlimited: true, billingLimitUsd: 10, allowModelAliases: false, modelScopeMode: "all", allowedModelIds: [], providerScope: ["all"], tierScope: ["all"] }); } function beginEdit(key: ClientKeyDTO): void { @@ -198,13 +214,16 @@ export function ClientKeysPage() { billingUnlimited: key.billingLimitUsdTicks === 0, billingLimitUsd: key.billingLimitUsdTicks > 0 ? key.billingLimitUsdTicks / USD_TICKS : 10, allowModelAliases: key.allowModelAliases, + modelScopeMode: key.allowedModelIds.length > 0 ? "restricted" : "all", allowedModelIds: key.allowedModelIds, + providerScope: key.providerScope ?? ["all"], + tierScope: key.tierScope ?? ["all"], }); } function toggleModel(id: string): void { const current = form.getValues("allowedModelIds"); - form.setValue("allowedModelIds", current.includes(id) ? current.filter((value) => value !== id) : [...current, id], { shouldDirty: true }); + form.setValue("allowedModelIds", current.includes(id) ? current.filter((value) => value !== id) : [...current, id], { shouldDirty: true, shouldValidate: true }); } const result = keysQuery.data; @@ -257,7 +276,7 @@ export function ClientKeysPage() { { value: "disabled", label: t("common.disabled") }, { value: "expired", label: t("keys.statusExpired") }, ] }, - { id: "modelScope", label: t("keys.models"), value: modelScopeFilter, onChange: (value) => { setModelScopeFilter(value); setPage(1); }, options: [ + { id: "modelScope", label: t("keys.modelScope"), value: modelScopeFilter, onChange: (value) => { setModelScopeFilter(value); setPage(1); }, options: [ { value: "all", label: t("keys.allModels") }, { value: "restricted", label: t("keys.restrictedModels") }, ] }, @@ -278,12 +297,13 @@ export function ClientKeysPage() { {keysQuery.isError ? void keysQuery.refetch()} /> : null} {result && result.items.length === 0 ? : null} {keysQuery.isPending || (result && result.items.length > 0) ? ( - +
+ @@ -297,6 +317,7 @@ export function ClientKeysPage() { {t("keys.name")}{t("keys.prefix")}{t("keys.status")} + {t("keys.accountScope")}{t("keys.rpmShort")}{t("keys.concurrencyShort")}{t("keys.billingLimit")} @@ -306,9 +327,9 @@ export function ClientKeysPage() { {keysQuery.isPending ? ( - + ) : ( - ( + ( toggleKey(key.id, checked === true)} aria-label={t("common.selectItem", { name: key.name })} /> @@ -332,6 +353,7 @@ export function ClientKeysPage() { + {key.rpmLimit > 0 ? key.rpmLimit : t("keys.unlimited")} {key.maxConcurrent > 0 ? key.maxConcurrent : t("keys.unlimited")} @@ -424,34 +446,6 @@ export function ClientKeysPage() { {form.formState.errors.expiresAt ?

{form.formState.errors.expiresAt.message}

: null} -
-
- {t("keys.models")} - {selectedModels.length === 0 ? t("keys.allModels") : t("keys.selectedModels", { count: selectedModels.length })} -
-
-
- - { setModelOptionsSearch(event.target.value); setModelOptionsPage(1); }} placeholder={t("keys.modelSearch")} aria-label={t("keys.modelSearch")} /> -
-
- {modelsQuery.isPending ? : modelsQuery.data?.items.map((model) => { - const checked = selectedModels.includes(model.id); - const controlId = `allowed-model-${model.id}`; - return ( - - ); - })} - {modelsQuery.data?.items.length === 0 ?

{t("common.noData")}

: null} -
- {modelsQuery.data && modelsQuery.data.total > modelsQuery.data.pageSize ? : null} -
-
@@ -483,6 +477,120 @@ export function ClientKeysPage() {
form.setValue("enabled", checked)} />
+
+
+
+ + ( + { + field.onChange(value); + setModelOptionsPage(1); + if (modelScopeMode === "restricted" && selectedModels.length > 0) { + form.setValue("allowedModelIds", [], { shouldDirty: true, shouldValidate: true }); + toast.info(t("keys.modelsClearedForScopeChange")); + } + }} + normalizeAllWhenComplete + options={[ + { value: "grok_build", label: "Build" }, + { value: "grok_web", label: "Web" }, + { value: "grok_console", label: "Console" }, + ]} + /> + )} /> +
+
+
+ + + + + + {t("keys.accountScopeDescription")} + +
+ ( + { + field.onChange(value); + setModelOptionsPage(1); + if (modelScopeMode === "restricted" && selectedModels.length > 0) { + form.setValue("allowedModelIds", [], { shouldDirty: true, shouldValidate: true }); + toast.info(t("keys.modelsClearedForScopeChange")); + } + }} + allLabel={t("keys.allTiersIncludingUnknown")} + options={[ + { value: "free", label: "Free" }, + { value: "super", label: "Super" }, + ]} + /> + )} /> +
+
+ + ( + + + + + + { + field.onChange(value); + setModelOptionsPage(1); + if (value === "all") { + form.setValue("allowedModelIds", [], { shouldDirty: true, shouldValidate: true }); + form.clearErrors("allowedModelIds"); + } + }}> + {t("keys.allModels")} + {t("keys.restrictedModels")} + + + + )} /> +
+
+ {modelScopeMode === "restricted" ? ( +
+
+
+ + { setModelOptionsSearch(event.target.value); setModelOptionsPage(1); }} placeholder={t("keys.modelSearch")} aria-label={t("keys.modelSearch")} /> +
+
+ {modelsQuery.isPending ? : modelsQuery.data?.items.map((model) => { + const checked = selectedModels.includes(model.id); + const controlId = `allowed-model-${model.id}`; + return ( + + ); + })} + {modelsQuery.data?.items.length === 0 ?

{t("common.noData")}

: null} +
+ {modelsQuery.data && modelsQuery.data.total > modelsQuery.data.pageSize ? : null} +
+ {form.formState.errors.allowedModelIds ?

{form.formState.errors.allowedModelIds.message}

: null} +
+ ) : null} +
@@ -574,3 +682,59 @@ function ClientKeyStatus({ value, referenceTime }: { value: ClientKeyDTO; refere } return {t("keys.statusActive")}; } + +function ScopeDropdown({ allLabel, ariaLabel, summary, value, onChange, options, normalizeAllWhenComplete = false }: { allLabel?: string; ariaLabel: string; summary: string; value: string[]; onChange: (value: string[]) => void; options: Array<{ value: string; label: string }>; normalizeAllWhenComplete?: boolean }) { + function toggle(nextValue: string): void { + if (nextValue === "all") { + onChange(["all"]); + return; + } + const current = value.includes("all") ? (allLabel ? [] : options.map((option) => option.value)) : value; + const next = current.includes(nextValue) ? current.filter((item) => item !== nextValue) : [...current, nextValue]; + if (next.length === 0) { + return; + } + if (normalizeAllWhenComplete && next.length === options.length) { + onChange(["all"]); + return; + } + onChange(next); + } + return ( + + + + + + {allLabel ? ( + <> + toggle("all")} onSelect={(event) => event.preventDefault()}> + {allLabel} + + + + ) : null} + {options.map((option) => ( + toggle(option.value)} onSelect={(event) => event.preventDefault()}> + {option.label} + + ))} + + + ); +} + +function AccountScopeSummary({ providerScope, tierScope }: { providerScope: ProviderScopeValue[]; tierScope: TierScopeValue[] }) { + const { t } = useTranslation(); + const providerLabels: Record = { all: t("keys.allProviders"), grok_build: "Build", grok_web: "Web", grok_console: "Console" }; + const tierLabels: Record = { all: t("keys.allTiers"), free: "Free", super: "Super" }; + return ( + + {providerScope.map((value) => providerLabels[value]).join(" · ")} + {tierScope.map((value) => tierLabels[value]).join(" · ")} + + ); +} diff --git a/frontend/src/features/creative-console/creative-console-page.tsx b/frontend/src/features/creative-console/creative-console-page.tsx index f57be4da1..3420c67a0 100644 --- a/frontend/src/features/creative-console/creative-console-page.tsx +++ b/frontend/src/features/creative-console/creative-console-page.tsx @@ -102,14 +102,17 @@ export function CreativeConsolePage() { queryFn: () => listAllPaginatedItems((page, pageSize) => listClientKeys({ page, pageSize, status: "active" })), staleTime: 30_000, }); - const modelsQuery = useQuery({ - queryKey: ["creative-console", "models"], - queryFn: () => listAllPaginatedItems((page, pageSize) => listModels({ page, pageSize, status: "enabled" })), - staleTime: 30_000, - }); const activeKeys = useMemo(() => (keysQuery.data ?? []).filter(isUsableKey), [keysQuery.data]); const effectiveKeyId = activeKeys.some((key) => key.id === selectedKeyId) ? selectedKeyId : activeKeys[0]?.id ?? ""; const selectedKey = activeKeys.find((key) => key.id === effectiveKeyId); + const modelProviderScope = (selectedKey?.providerScope ?? ["all"]).filter((value) => value !== "all"); + const modelTierScope = (selectedKey?.tierScope ?? ["all"]).filter((value) => value !== "all"); + const modelsQuery = useQuery({ + queryKey: ["creative-console", "models", modelProviderScope.join(","), modelTierScope.join(",")], + queryFn: () => listAllPaginatedItems((page, pageSize) => listModels({ page, pageSize, status: "enabled", providerScope: modelProviderScope, tierScope: modelTierScope, activeScope: true })), + enabled: Boolean(selectedKey), + staleTime: 30_000, + }); const availableModels = useMemo(() => (modelsQuery.data ?? []).filter((model) => model.enabled && model.available), [modelsQuery.data]); const permittedModels = useMemo(() => { if (!selectedKey || selectedKey.allowedModelIds.length === 0) return availableModels; diff --git a/frontend/src/shared/i18n/index.ts b/frontend/src/shared/i18n/index.ts index 448cf1e78..061702129 100644 --- a/frontend/src/shared/i18n/index.ts +++ b/frontend/src/shared/i18n/index.ts @@ -692,6 +692,14 @@ const resources = { lastUsed: "最近使用", rpm: "每分钟请求", maxConcurrent: "最大并发数", + accountScope: "调用范围", + providerScope: "渠道范围", + modelScope: "模型范围", + tierScope: "订阅范围", + allProviders: "全部渠道", + allTiers: "全部订阅", + allTiersIncludingUnknown: "全部(含待识别)", + accountScopeDescription: "模型列表会随渠道和订阅范围联动;调用仅在所选范围内选号,不会越界回退。全部订阅包含待识别,Free + Super 会排除待识别;Console 不受订阅限制。", rpmValue: "{{value}} RPM", concurrentValue: "{{value}} 并发", unlimited: "无限制", @@ -706,8 +714,10 @@ const resources = { modelAliasesOn: "已开启:可发现并调用当前模型支持的推理等级别名。", modelAliasesOff: "已关闭:不生成动态推理别名;既有 Console 兼容别名不受影响。", allModels: "全部模型", - modelSearch: "搜索模型,留空默认允许全部模型。", + modelSearch: "搜索模型", selectedModels: "已选择 {{count}} 个模型", + selectModelRequired: "请至少选择一个模型,或切换为全部模型。", + modelsClearedForScopeChange: "调用范围已变化,请重新选择模型。", neverExpires: "永不过期", selectExpiry: "选择截止日期", clearExpiry: "清除有效期", @@ -1297,7 +1307,7 @@ const resources = { dashboard: { title: "Dashboard", subtitle: "Welcome, {{name}}", lastUpdated: "Updated {{time}}", welcome: "Welcome, {{name}}", gatewayOnline: "Gateway healthy", gatewayOffline: "Unavailable", gatewayReady: "The gateway is ready to process API requests.", gatewayNeedsAccount: "Connect an available account to process API requests.", manageAccounts: "Manage accounts", usage: "Usage overview", seeAll: "View audits", activeAccounts: "Available accounts", accountCapacity: "{{total}} accounts · {{rate}}% success", manage: "Manage", tokens: "Total tokens", requests: "Requests", successRate: "Success rate", accountCount: "Accounts", accountDistribution: "Build {{build}} · Web {{web}} · Console {{console}}", requestSuccessRate: "{{rate}}% success", averageRequestCost: "{{cost}} per request", availableSummary: "{{active}} / {{total}} available", requestBreakdown: "{{success}} succeeded · {{failed}} failed", requestQualitySummary: "{{success}}% success · {{failed}} failed", periodSummary: "Current {{period}} range", successSummary: "{{success}} / {{total}} requests succeeded", trend: "Usage trend", trendRequests: "Requests", trendTokens: "Tokens", trendSummary: "{{requests}} requests · {{rate}}% success", resourcesTitle: "Resource availability", availability: "Account availability", unavailableAccounts: "Unavailable accounts", unavailableSummary: "{{unavailable}} / {{total}} unavailable", enabledModels: "Enabled models", modelsAvailableSummary: "{{enabled}} / {{total}} enabled", activeClientKeys: "Available keys", keysAvailableSummary: "{{active}} / {{total}} available", allTimeRequests: "All-time requests", allTimeRequestsSummary: "Since the first recorded request", qualityTitle: "Request quality", successfulRequests: "Successful requests", failedRequests: "Failed requests", topModels: "Top 10 model billing", model: "Model", share: "Call share", noTopModels: "No model calls in this period", noTrendData: "No usage data in this period", billing: "Billing", billingSummary: "Accumulated in the current {{period}} range", inputTokens: "Input", cachedTokens: "Cached", outputTokens: "Output", reasoningTokens: "Reasoning", uncachedInput: "Uncached input", visibleOutput: "Visible output", otherModels: "Other models", providerDistribution: "Provider distribution", providerDetail: "{{rate}}% success · {{tokens}} tokens", providerStripeDetail: "{{provider}} · {{requests}} requests · {{share}}%", providerNoRequests: "No provider requests", activityTitle: "Request activity", lastDays: "Last {{count}} days", activityDay: "{{date}} · {{requests}} requests", activityLess: "Less", activityMore: "More", tokenBreakdown: "Input {{input}} · Cached {{cached}} · Output {{output}} · Reasoning {{reasoning}}", tokenEfficiency: "Cache hit rate {{rate}}%", explore: "Manage resources", accountsProduct: "Upstream accounts", accountsProductDescription: "Manage OAuth credentials, health, quota, and concurrency.", keysProduct: "Client keys", keysProductDescription: "Create downstream API keys and control models, RPM, and concurrency.", modelsProduct: "Model routes", modelsProductDescription: "Manage public model IDs, upstream mappings, and availability.", auditsProduct: "Request audits", auditsProductDescription: "Inspect status, latency, and token usage records.", gatewayFooter: "A lightweight API gateway", learnMore: "View models", unknown: "Unknown" }, accounts: { title: "Upstream accounts", description: "Manage separate Grok Build OAuth and Grok Web SSO pools, health, concurrency, and quotas.", add: "Connect account", connectAccount: "Connect account", convertToBuild: "Convert to Build", convertToBuildTitle: "Convert Grok Web accounts to Build?", convertToBuildDescription: "Choose the conversion scope for the selected {{count}} Web accounts. Conversion may take a moment.", convertingProgress: "Converting {{completed}} / {{total}}", syncingProgress: "Syncing {{completed}} / {{total}}", conversionCompleted: "Conversion complete: {{created}} created, {{linked}} linked, {{skipped}} skipped, {{failed}} failed", linkedAccountTooltip: "Linked account: {{name}}", syncAll: "Sync all", renewAll: "Renew all", syncAllTitle: "Sync all accounts?", syncAllDescription: "Refresh quota data for every enabled Grok Build account.", syncAllWebDescription: "Refresh real quota data for every enabled Grok Web account.", renewAllTitle: "Renew all accounts?", renewAllDescription: "Refresh every renewable Grok Build credential. Accounts without automatic renewal will be skipped.", search: "Search accounts", deviceLogin: "Device OAuth", importAuth: "Import account files", importSSO: "Import SSO JSON", quickImportSSO: "Quick import SSO", importWebFile: "Import account files", quickImportTitle: "Quick import Grok Web accounts", quickImportDescription: "Paste multiple SSO tokens or upload a TXT file with one token per line. Duplicate entries are ignored.", ssoTokens: "SSO tokens (one per line)", uploadTXT: "Upload TXT", ssoTokenPlaceholder: "sso_token_1\nsso_token_2\nsso=token_3", importAction: "Import", exportAuth: "Export account file", exported: "Account credentials exported", exportTitle: "Export {{provider}} accounts?", exportDescription: "The JSON contains credentials for the current provider and any configured Cloudflare cookies, and can be imported back into this project. Store it only in a trusted location.", exportCount: "Export count", exportCountDescription: "Exports newest accounts first, up to 10,000 per batch.", refreshAllBilling: "Sync all quotas", account: "Account", status: "Status", type: "Type", quota: "Quota", routing: "Routing", credentialExpiry: "Credential expiry", credentialRenewal: "Renewal capability", autoRefresh: "Auto-renewal", noAutoRefresh: "No auto-renewal", buildAccountCount: "Grok Build accounts", webAccountCount: "Grok Web accounts", consoleAccountCount: "Grok Console accounts", abnormalAccountCount: "Abnormal accounts", routableAccountCount: "{{count}} routable", abnormalAccountBreakdown: "Recovering {{recovering}} · Needs attention {{attention}}", riskAccountCount: "Risk {{count}}", riskFilter: "Risk", riskNormal: "Normal", lastUsed: "Last used", createdAt: "Created at", name: "Name", priority: "Priority", maxConcurrent: "Max concurrency", minimumRemaining: "Minimum remaining threshold", routingSummary: "Priority {{priority}} · {{count}} concurrent", minimumRemainingSummary: "Minimum remaining {{value}}", minimumRemainingNotApplicable: "Minimum remaining does not apply to Free", refreshToken: "Refresh credential", refreshBilling: "Sync quota", refreshModeQuota: "Sync quota", quotaNotSynced: "Not synced", quotaResetAt: "Next reset {{time}}", quotaResetUnknown: "The upstream did not return a reset time", webQuotaUsage: "Used {{used}} / {{total}} total · {{remaining}} remaining", webModeQuotaRemaining: "{{mode}} has {{remaining}} calls remaining", webWeeklyQuotaUsage: "Weekly quota {{remaining}}% remaining", paidQuotaDetails: "{{remaining}} credits remaining", created: "Account connected", createdWithSyncFailure: "Account connected, but initial quota or model sync failed", updated: "Account updated", updatedWithModelSyncFailure: "Account updated, but model capability sync failed. Run model sync again later.", deleted: "Account deleted", imported: "Import complete: {{created}} created, {{updated}} updated", importedWithSyncFailures: "Import complete: {{created}} created, {{updated}} updated; {{synced}} synced, {{syncFailed}} failed initial sync", billingRefreshed: "Quota synced", authRefreshed: "Credential refreshed", allBillingRefreshed: "Sync complete: {{succeeded}} succeeded, {{failed}} failed", allTokensRefreshed: "Renewal complete: {{succeeded}} succeeded, {{failed}} failed, {{skipped}} skipped", statusDisabled: "Disabled", statusReauthRequired: "Invalid", statusCooldown: "Cooling", statusActive: "Normal", botRisk: "Risk", botRiskTooltip: "Bot risk control. This may affect video generation and other permissions.", nsfwEnabledMark: "NSFW", nsfwEnabledTooltip: "NSFW enabled · {{time}}", agreementFilter: "Account agreements", agreementNsfwEnabled: "NSFW enabled", agreementNsfwDisabled: "NSFW not enabled", agreementTermsAccepted: "Terms accepted", agreementTermsNotAccepted: "Terms not accepted", agreementAllAccepted: "All accepted", agreementAllNotAccepted: "All not accepted", associationFilter: "Account links", associationWebLinked: "Web linked", associationWebUnlinked: "Web not linked", associationBuildLinked: "Build linked", associationBuildUnlinked: "Build not linked", associationConsoleLinked: "Console linked", associationConsoleUnlinked: "Console not linked", associationAllLinked: "All linked", associationAllUnlinked: "All not linked", buildRouteMode: { label: "Upstream endpoint", auto: "Automatic", build: "Build", xai: "XAI", autoDescription: "Selects an upstream based on account tier and risk status.", buildDescription: "Routes all requests to Build with XAI fallback disabled.", xaiDescription: "Routes all inference and video requests to XAI.", xaiUnconfirmedWarning: "This account has not been confirmed for XAI access; the upstream may reject requests." }, buildSuperEntitled: { label: "Mark as Super manually", description: "Use only when automatic tier detection is inaccurate. Enabled accounts sync models with Super access." }, quotaFree: "Free", quotaSuper: "Super", quotaPaid: "Super", quotaUnknown: "Unknown", paidQuotaUsage: "Upstream quota usage", weeklyQuota: "Weekly", monthlyQuota: "Monthly", weeklyLimit: "Weekly limit {{percent}}% remaining", freeObservedUsage: "{{used}} tokens observed in the last 24 hours", freeEstimatedUsage: "About {{used}} / {{limit}} tokens", freeEstimatedDescription: "Estimated from observed Grok Build Free characteristics; upstream exhaustion data takes precedence", observedRolling: "Free accounts use this proxy's rolling observation and may differ from upstream billing data", upstreamConfirmed: "Exhaustion confirmed by upstream", waitingReset: "Waiting reset", waitingResetUntil: "Waiting for reset · probe after {{time}} on real traffic", paidWaitingResetUntil: "Waiting for reset · check the upstream Billing period after {{time}}", probing: "Probing", probingQuota: "Checking quota recovery with real traffic", paidProbingQuota: "Checking whether the upstream Billing period has recovered", upstreamBilling: "Upstream billing data", deviceTitle: "Connect Grok Build", deviceDescription: "Complete authorization in a new window. This page will check status automatically.", userCode: "Device code", openVerification: "Continue to authorization", waiting: "Waiting for authorization", expiresAt: "Expires {{time}}", deleteTitle: "Delete account?", deleteDescription: "Credentials and quota snapshots will be removed. This cannot be undone.", deleteConfirm: "Delete", linkedDeleteTitle: "Also delete linked accounts?", linkedDeleteSelectAll: "Select all linked", linkedDeleteClearAll: "Clear linked selection", linkedDeleteExtra: "+{{count}}", linkedDeleteExtraPending: "+—", linkedDeleteExtraFailed: "!", linkedDeleteHint: "Only linked accounts are deleted; unlinked accounts remain unchanged.", linkedDeletePreviewFailed: "Failed to load linked counts. Uncheck and retry, or delete later.", batchUpdated: "Selected accounts updated", batchBillingRefreshed: "Quota sync complete: {{succeeded}} succeeded, {{failed}} failed", batchDeleteTitle: "Delete {{count}} selected accounts?", cleanupAction: "Clean up accounts", cleanupTitle: "Clean up {{provider}} accounts", cleanupDescription: "Delete accounts in the selected states from this provider, including their quota data. Normal accounts are never affected. This cannot be undone.", cleanupCompleted: "Deleted {{deleted}} accounts", deletedWithSkipped: "Deleted {{deleted}} accounts, skipped {{skipped}} groups with active video jobs", cleanupCompletedDetailed: "Deleted {{deleted}} accounts ({{linked}} linked), skipped {{skipped}} groups", cleanupLinkedWarning: "Warning: linked accounts are not limited to the selected states and may include healthy accounts; they will be deleted too.", cleanupPreviewFailed: "Failed to load estimated counts. Adjust selection and retry, or clean up later.", cleanupPreviewTotal: "About {{total}} accounts will be deleted", cleanupStart: "Start cleanup" }, models: { title: "Model routes", description: "Map one client-facing model name to provider-specific Build, Web, and Console routes. Multiple public names may also target the same upstream model.", sync: "Sync models", search: "Search models", model: "Public model name", publicId: "Public model name", upstream: "Upstream model", provider: "Model source", providerGrokBuild: "Grok Build", providerGrokWeb: "Grok Web", accountSupport: "Supported accounts", lastSyncedAt: "Last synced", capability: "Account model capabilities", supportSummary: "{{supported}} available / {{total}} total", awaitingCapabilitySync: "Awaiting model capability sync", unknownCapability: "Pending sync", partialCapability: "Awaiting remaining accounts", available: "Available", unavailable: "Unavailable", status: "Status", create: "Add model", createTitle: "Add model route", createDescription: "Add a custom Grok Build route. The public model name has no source prefix; the upstream model uses Build/ to identify its actual channel.", editTitle: "Edit model route", bindAccounts: "Bind specific accounts", boundAccounts: "Bound accounts", bindAccountsDescription: "Route only through selected accounts; when disabled, use synchronized model capabilities.", enabledDescription: "Controls whether this route participates in model discovery and request scheduling.", searchAccounts: "Search account name or ID", selectedAccounts: "{{count}} accounts selected", selectAccountRequired: "Select at least one account", noBindableAccounts: "No accounts from this provider can be bound", deleteTitle: "Delete model route?", deleteDescription: "This removes {{name}} and its client-key permissions. A later sync recreates it if upstream accounts still support the model.", batchDeleteTitle: "Delete {{count}} selected model routes?", batchDeleteDescription: "This removes the selected routes and their client-key permissions. Models still supported by upstream accounts are recreated on a later sync.", synced: "Synced {{count}} models", created: "Model route created", updated: "Model route updated", deleted: "Model route deleted", batchDeleted: "Deleted {{count}} model routes", batchUpdated: "Selected models updated" }, - keys: { title: "Client keys", description: "Control downstream model access, RPM, concurrency, billed usage, and expiration.", search: "Search keys", create: "Create key", name: "Name", prefix: "Key", limits: "Limits", rpmLimit: "RPM limit", concurrencyLimit: "Concurrency limit", rpmShort: "RPM", concurrencyShort: "Concurrency", billingLimit: "Usage limit", models: "Allowed models", status: "Status", statusActive: "Active", statusExpired: "Expired", modelScope: "Model scope", restrictedModels: "Selected models", expires: "Expires", lastUsed: "Last used", rpm: "Requests per minute", maxConcurrent: "Max concurrency", rpmValue: "{{value}} RPM", concurrentValue: "{{value}} concurrent", unlimited: "Unlimited", billedUsage: "{{value}} billed", billingLimitDescription: "Accumulates billed request audits, preferring upstream-reported cost and otherwise using the official price estimate.", enabledDescription: "Controls whether this key can make API requests.", modelAliases: "Model aliases", modelAliasesDescription: "Suffix aliases are generated only for available models with multiple supported reasoning levels.", modelAliasesBuildSupport: "Grok Build: grok-4.5, grok-3-mini, and grok-3-mini-fast support low / medium / high.", modelAliasesConsoleSupport: "Grok Console: grok-4.3 supports none / low / medium / high; grok-4.20-0309-reasoning supports low / medium / high; grok-4.20-multi-agent-0309 also supports xhigh.", modelAliasesUnsupported: "Grok Web, image models, and video models do not generate reasoning-effort suffixes.", modelAliasesOn: "On: discover and call the reasoning aliases supported by available models.", modelAliasesOff: "Off: dynamic reasoning aliases are disabled; existing Console compatibility aliases remain available.", allModels: "All models", modelSearch: "Search models. Leave blank to allow all models.", selectedModels: "{{count}} models selected", neverExpires: "Never expires", selectExpiry: "Select expiration date", clearExpiry: "Clear expiration", expiryTime: "Expiration time", createTitle: "Create client key", editTitle: "Edit client key", secretTitle: "Client key created", secretDescription: "Store this key securely.", secretLabel: "Key", copySecret: "Copy", copySecretTitle: "Copy key", copySecretDescription: "Use the button below to copy your key.", secretReady: "Ready", created: "Client key created", updated: "Client key updated", deleted: "Client key deleted", deleteTitle: "Delete client key?", deleteDescription: "All downstream requests using this key will stop immediately.", batchUpdated: "Selected keys updated", batchDeleteTitle: "Delete {{count}} selected keys?" }, + keys: { title: "Client keys", description: "Control downstream model access, RPM, concurrency, billed usage, and expiration.", search: "Search keys", create: "Create key", name: "Name", prefix: "Key", limits: "Limits", rpmLimit: "RPM limit", concurrencyLimit: "Concurrency limit", rpmShort: "RPM", concurrencyShort: "Concurrency", billingLimit: "Usage limit", models: "Allowed models", status: "Status", statusActive: "Active", statusExpired: "Expired", modelScope: "Model scope", restrictedModels: "Selected models", expires: "Expires", lastUsed: "Last used", rpm: "Requests per minute", maxConcurrent: "Max concurrency", accountScope: "Routing scope", providerScope: "Provider scope", tierScope: "Subscription scope", allProviders: "All providers", allTiers: "All subscriptions", allTiersIncludingUnknown: "All, including unknown", accountScopeDescription: "The model list follows the provider and subscription scope, and routing never falls back outside it. All subscriptions include unknown accounts; Free + Super excludes unknown accounts. Console ignores subscription restrictions.", rpmValue: "{{value}} RPM", concurrentValue: "{{value}} concurrent", unlimited: "Unlimited", billedUsage: "{{value}} billed", billingLimitDescription: "Accumulates billed request audits, preferring upstream-reported cost and otherwise using the official price estimate.", enabledDescription: "Controls whether this key can make API requests.", modelAliases: "Model aliases", modelAliasesDescription: "Suffix aliases are generated only for available models with multiple supported reasoning levels.", modelAliasesBuildSupport: "Grok Build: grok-4.5, grok-3-mini, and grok-3-mini-fast support low / medium / high.", modelAliasesConsoleSupport: "Grok Console: grok-4.3 supports none / low / medium / high; grok-4.20-0309-reasoning supports low / medium / high; grok-4.20-multi-agent-0309 also supports xhigh.", modelAliasesUnsupported: "Grok Web, image models, and video models do not generate reasoning-effort suffixes.", modelAliasesOn: "On: discover and call the reasoning aliases supported by available models.", modelAliasesOff: "Off: dynamic reasoning aliases are disabled; existing Console compatibility aliases remain available.", allModels: "All models", modelSearch: "Search models", selectedModels: "{{count}} models selected", selectModelRequired: "Select at least one model, or switch to All models.", modelsClearedForScopeChange: "The routing scope changed. Select the allowed models again.", neverExpires: "Never expires", selectExpiry: "Select expiration date", clearExpiry: "Clear expiration", expiryTime: "Expiration time", createTitle: "Create client key", editTitle: "Edit client key", secretTitle: "Client key created", secretDescription: "Store this key securely.", secretLabel: "Key", copySecret: "Copy", copySecretTitle: "Copy key", copySecretDescription: "Use the button below to copy your key.", secretReady: "Ready", created: "Client key created", updated: "Client key updated", deleted: "Client key deleted", deleteTitle: "Delete client key?", deleteDescription: "All downstream requests using this key will stop immediately.", batchUpdated: "Selected keys updated", batchDeleteTitle: "Delete {{count}} selected keys?" }, audits: { title: "Request audits", description: "Successful requests store metadata only. Failed requests retain a redacted and size-limited diagnostic snapshot.",