fix(access): 统一 API Key 模型权限与列表契约
将全局启用、用户组基线、API Key 专属或排除规则及 scope 按固定顺序求值,避免 Key 越过所属用户组权限,并让运行时候选与模型列表共用同一权限链。 新增 Key 级可分配模型与失效规则诊断接口、OpenAI 兼容 /v1/models 及 rich 列表迁移路径;前端权限弹窗改为按当前 Key 实时加载并支持清理失效规则。 验证:Go 全量测试与 go vet 通过;Web 22 个测试文件共 142 项通过;pnpm lint、pnpm openapi、pnpm build、Compose 配置、gofmt、ShellCheck 和 git diff --check 通过;独立 PostgreSQL 真实配置验收通过。
This commit is contained in:
@@ -0,0 +1,507 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/modelaccess"
|
||||
)
|
||||
|
||||
type APIKeyAccessRuleDiagnostic struct {
|
||||
RuleID string `json:"ruleId"`
|
||||
ResourceType string `json:"resourceType"`
|
||||
ResourceID string `json:"resourceId"`
|
||||
ResourceName string `json:"resourceName,omitempty"`
|
||||
Effect string `json:"effect"`
|
||||
Effective bool `json:"effective"`
|
||||
Reason string `json:"reason,omitempty" enums:"resource_unavailable,owner_access_revoked,scope_not_allowed"`
|
||||
}
|
||||
|
||||
func (s *Store) enabledPlatformModels(ctx context.Context) ([]PlatformModel, []Platform, error) {
|
||||
models, err := s.ListModels(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
platforms, err := s.ListPlatforms(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
enabledPlatforms := map[string]bool{}
|
||||
for _, platform := range platforms {
|
||||
if platform.Status == "enabled" {
|
||||
enabledPlatforms[platform.ID] = true
|
||||
}
|
||||
}
|
||||
enabled := make([]PlatformModel, 0, len(models))
|
||||
for _, model := range models {
|
||||
if model.Enabled && enabledPlatforms[model.PlatformID] {
|
||||
enabled = append(enabled, model)
|
||||
}
|
||||
}
|
||||
return enabled, platforms, nil
|
||||
}
|
||||
|
||||
func (s *Store) filterPlatformModelsByLayeredAccess(ctx context.Context, user *auth.User, models []PlatformModel) ([]PlatformModel, error) {
|
||||
rules, err := s.listActiveAccessRulesForResources(ctx, platformModelAccessResources(models))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
baselineRules, apiKeyRules := splitLayeredAccessRules(rules)
|
||||
baseline := filterPlatformModelsByRuleSet(models, baselineRules, baselineAccessRuleSubjects(user), permissionLevel(user))
|
||||
if user == nil || strings.TrimSpace(user.APIKeyID) == "" {
|
||||
return baseline, nil
|
||||
}
|
||||
apiKeyUsers, err := s.apiKeyAccessRuleUsers(ctx, apiKeyRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return filterPlatformModelsByLayeredRuleSet(user, baseline, baselineRules, apiKeyRules, apiKeyUsers), nil
|
||||
}
|
||||
|
||||
func filterPlatformModelsByLayeredRuleSet(user *auth.User, baseline []PlatformModel, baselineRules []AccessRule, apiKeyRules []AccessRule, apiKeyUsers map[string]*auth.User) []PlatformModel {
|
||||
if user == nil || strings.TrimSpace(user.APIKeyID) == "" {
|
||||
return baseline
|
||||
}
|
||||
keyFiltered := make([]PlatformModel, 0, len(baseline))
|
||||
for _, model := range baseline {
|
||||
effectiveRules := effectiveAPIKeyRulesForPlatformModel(apiKeyRules, baselineRules, apiKeyUsers, model)
|
||||
if platformModelAllowedByAccessRules(model, effectiveRules, apiKeyAccessRuleSubjects(user), permissionLevel(user)) {
|
||||
keyFiltered = append(keyFiltered, model)
|
||||
}
|
||||
}
|
||||
return filterPlatformModelsByAPIKeyScopes(keyFiltered, user.APIKeyScopes)
|
||||
}
|
||||
|
||||
func (s *Store) filterPlatformModelsByBaselineAccess(ctx context.Context, user *auth.User, models []PlatformModel) ([]PlatformModel, error) {
|
||||
rules, err := s.listActiveAccessRulesForResources(ctx, platformModelAccessResources(models))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
baselineRules, _ := splitLayeredAccessRules(rules)
|
||||
return filterPlatformModelsByRuleSet(models, baselineRules, baselineAccessRuleSubjects(user), permissionLevel(user)), nil
|
||||
}
|
||||
|
||||
func (s *Store) filterRuntimeCandidatesByLayeredAccess(ctx context.Context, user *auth.User, candidates []RuntimeModelCandidate) ([]RuntimeModelCandidate, error) {
|
||||
if len(candidates) == 0 {
|
||||
return candidates, nil
|
||||
}
|
||||
accessUser, err := s.resolveCurrentAccessUser(ctx, user)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rules, err := s.listActiveAccessRulesForResources(ctx, candidateAccessResources(candidates))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
baselineRules, apiKeyRules := splitLayeredAccessRules(rules)
|
||||
baseline := filterCandidatesByRuleSet(candidates, baselineRules, baselineAccessRuleSubjects(accessUser), permissionLevel(accessUser))
|
||||
if accessUser == nil || strings.TrimSpace(accessUser.APIKeyID) == "" {
|
||||
return baseline, nil
|
||||
}
|
||||
apiKeyUsers, err := s.apiKeyAccessRuleUsers(ctx, apiKeyRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keyFiltered := make([]RuntimeModelCandidate, 0, len(baseline))
|
||||
for _, candidate := range baseline {
|
||||
effectiveRules := effectiveAPIKeyRulesForCandidate(apiKeyRules, baselineRules, apiKeyUsers, candidate)
|
||||
if candidateAllowedByAccessRules(candidate, effectiveRules, apiKeyAccessRuleSubjects(accessUser), permissionLevel(accessUser)) {
|
||||
keyFiltered = append(keyFiltered, candidate)
|
||||
}
|
||||
}
|
||||
filtered := make([]RuntimeModelCandidate, 0, len(keyFiltered))
|
||||
for _, candidate := range keyFiltered {
|
||||
if modelaccess.ScopeAllowsModelType(accessUser.APIKeyScopes, candidate.ModelType) {
|
||||
filtered = append(filtered, candidate)
|
||||
}
|
||||
}
|
||||
return filtered, nil
|
||||
}
|
||||
|
||||
func (s *Store) ListAPIKeyAssignablePlatformModelsForKey(ctx context.Context, user *auth.User, apiKeyID string) ([]PlatformModel, []APIKeyAccessRuleDiagnostic, error) {
|
||||
if localGatewayUserID(user) == "" {
|
||||
return nil, nil, ErrLocalUserRequired
|
||||
}
|
||||
accessUser, err := s.resolveOwnedAPIKeyAccessUser(ctx, user, apiKeyID)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
allModels, err := s.ListModels(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
enabledModels, platforms, err := s.enabledPlatformModels(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
rules, err := s.listActiveAccessRulesForResources(ctx, platformModelAccessResources(enabledModels))
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
baselineRules, _ := splitLayeredAccessRules(rules)
|
||||
baseline := filterPlatformModelsByRuleSet(enabledModels, baselineRules, baselineAccessRuleSubjects(accessUser), permissionLevel(accessUser))
|
||||
scoped := filterPlatformModelsByAPIKeyScopes(baseline, accessUser.APIKeyScopes)
|
||||
ownedRules, err := s.ListAPIKeyAccessRules(ctx, user)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
diagnostics := diagnoseAPIKeyRules(apiKeyID, ownedRules, allModels, enabledModels, baseline, scoped, platforms)
|
||||
return scoped, diagnostics, nil
|
||||
}
|
||||
|
||||
func (s *Store) resolveOwnedAPIKeyAccessUser(ctx context.Context, user *auth.User, apiKeyID string) (*auth.User, error) {
|
||||
if user == nil {
|
||||
return nil, ErrLocalUserRequired
|
||||
}
|
||||
gatewayUserID := localGatewayUserID(user)
|
||||
if gatewayUserID == "" {
|
||||
return nil, ErrLocalUserRequired
|
||||
}
|
||||
var scopesBytes []byte
|
||||
var userGroupID string
|
||||
var userGroupKey string
|
||||
err := s.pool.QueryRow(ctx, `
|
||||
SELECT k.scopes, COALESCE(k.user_group_id::text, u.default_user_group_id::text, ''), COALESCE(g.group_key, '')
|
||||
FROM gateway_api_keys k
|
||||
JOIN gateway_users u ON u.id = k.gateway_user_id
|
||||
LEFT JOIN gateway_user_groups g ON g.id = COALESCE(k.user_group_id, u.default_user_group_id)
|
||||
WHERE k.id = $1::uuid
|
||||
AND k.gateway_user_id = $2::uuid
|
||||
AND k.deleted_at IS NULL
|
||||
AND u.deleted_at IS NULL`, strings.TrimSpace(apiKeyID), gatewayUserID).Scan(&scopesBytes, &userGroupID, &userGroupKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
next := *user
|
||||
next.APIKeyID = strings.TrimSpace(apiKeyID)
|
||||
next.APIKeyScopes = decodeStringArray(scopesBytes)
|
||||
next.UserGroupID = userGroupID
|
||||
next.UserGroupKey = userGroupKey
|
||||
next.UserGroupKeys = nil
|
||||
if userGroupKey != "" {
|
||||
next.UserGroupKeys = []string{userGroupKey}
|
||||
}
|
||||
return &next, nil
|
||||
}
|
||||
|
||||
func effectiveAPIKeyRulesForPlatformModel(rules []AccessRule, baselineRules []AccessRule, users map[string]*auth.User, model PlatformModel) []AccessRule {
|
||||
effective := make([]AccessRule, 0, len(rules))
|
||||
for _, rule := range rules {
|
||||
accessUser := users[rule.SubjectID]
|
||||
if accessUser == nil || !accessRuleMatchesPlatformModel(rule, model) ||
|
||||
!platformModelAllowedByAccessRules(model, baselineRules, baselineAccessRuleSubjects(accessUser), permissionLevel(accessUser)) ||
|
||||
len(modelaccess.FilterModelTypes(accessUser.APIKeyScopes, model.ModelType)) == 0 {
|
||||
continue
|
||||
}
|
||||
effective = append(effective, rule)
|
||||
}
|
||||
return effective
|
||||
}
|
||||
|
||||
func effectiveAPIKeyRulesForCandidate(rules []AccessRule, baselineRules []AccessRule, users map[string]*auth.User, candidate RuntimeModelCandidate) []AccessRule {
|
||||
effective := make([]AccessRule, 0, len(rules))
|
||||
for _, rule := range rules {
|
||||
accessUser := users[rule.SubjectID]
|
||||
if accessUser == nil || !accessRuleMatchesCandidate(rule, candidate) ||
|
||||
!candidateAllowedByAccessRules(candidate, baselineRules, baselineAccessRuleSubjects(accessUser), permissionLevel(accessUser)) ||
|
||||
!modelaccess.ScopeAllowsModelType(accessUser.APIKeyScopes, candidate.ModelType) {
|
||||
continue
|
||||
}
|
||||
effective = append(effective, rule)
|
||||
}
|
||||
return effective
|
||||
}
|
||||
|
||||
func (s *Store) apiKeyAccessRuleUsers(ctx context.Context, rules []AccessRule) (map[string]*auth.User, error) {
|
||||
ids := make([]string, 0, len(rules))
|
||||
seen := map[string]bool{}
|
||||
for _, rule := range rules {
|
||||
if rule.SubjectType != "api_key" || rule.SubjectID == "" || seen[rule.SubjectID] {
|
||||
continue
|
||||
}
|
||||
seen[rule.SubjectID] = true
|
||||
ids = append(ids, rule.SubjectID)
|
||||
}
|
||||
if len(ids) == 0 {
|
||||
return map[string]*auth.User{}, nil
|
||||
}
|
||||
rows, err := s.pool.Query(ctx, `
|
||||
SELECT k.id::text, k.scopes, u.id::text,
|
||||
COALESCE(u.gateway_tenant_id::text, ''), COALESCE(u.tenant_id, ''), COALESCE(u.tenant_key, ''),
|
||||
u.roles, COALESCE(k.user_group_id::text, u.default_user_group_id::text, ''), COALESCE(g.group_key, '')
|
||||
FROM gateway_api_keys k
|
||||
JOIN gateway_users u ON u.id = k.gateway_user_id
|
||||
LEFT JOIN gateway_user_groups g ON g.id = COALESCE(k.user_group_id, u.default_user_group_id)
|
||||
WHERE k.id = ANY($1::uuid[])
|
||||
AND k.status = 'active'
|
||||
AND k.deleted_at IS NULL
|
||||
AND (k.expires_at IS NULL OR k.expires_at > now())
|
||||
AND u.status = 'active'
|
||||
AND u.deleted_at IS NULL`, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
users := map[string]*auth.User{}
|
||||
for rows.Next() {
|
||||
var apiKeyID string
|
||||
var scopesBytes []byte
|
||||
var rolesBytes []byte
|
||||
var gatewayUserID string
|
||||
var gatewayTenantID string
|
||||
var tenantID string
|
||||
var tenantKey string
|
||||
var userGroupID string
|
||||
var userGroupKey string
|
||||
if err := rows.Scan(&apiKeyID, &scopesBytes, &gatewayUserID, &gatewayTenantID, &tenantID, &tenantKey, &rolesBytes, &userGroupID, &userGroupKey); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groupKeys := []string(nil)
|
||||
if userGroupKey != "" {
|
||||
groupKeys = []string{userGroupKey}
|
||||
}
|
||||
users[apiKeyID] = &auth.User{
|
||||
GatewayUserID: gatewayUserID,
|
||||
GatewayTenantID: gatewayTenantID,
|
||||
TenantID: tenantID,
|
||||
TenantKey: tenantKey,
|
||||
Roles: decodeStringArray(rolesBytes),
|
||||
UserGroupID: userGroupID,
|
||||
UserGroupKey: userGroupKey,
|
||||
UserGroupKeys: groupKeys,
|
||||
APIKeyID: apiKeyID,
|
||||
APIKeyScopes: decodeStringArray(scopesBytes),
|
||||
}
|
||||
}
|
||||
return users, rows.Err()
|
||||
}
|
||||
|
||||
func splitLayeredAccessRules(rules []AccessRule) ([]AccessRule, []AccessRule) {
|
||||
baseline := make([]AccessRule, 0, len(rules))
|
||||
apiKeys := make([]AccessRule, 0, len(rules))
|
||||
for _, rule := range rules {
|
||||
if rule.SubjectType == "api_key" {
|
||||
apiKeys = append(apiKeys, rule)
|
||||
} else {
|
||||
baseline = append(baseline, rule)
|
||||
}
|
||||
}
|
||||
return baseline, apiKeys
|
||||
}
|
||||
|
||||
func filterPlatformModelsByRuleSet(models []PlatformModel, rules []AccessRule, subjects map[string]bool, level int) []PlatformModel {
|
||||
filtered := make([]PlatformModel, 0, len(models))
|
||||
for _, model := range models {
|
||||
if platformModelAllowedByAccessRules(model, rules, subjects, level) {
|
||||
filtered = append(filtered, model)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func filterCandidatesByRuleSet(candidates []RuntimeModelCandidate, rules []AccessRule, subjects map[string]bool, level int) []RuntimeModelCandidate {
|
||||
filtered := make([]RuntimeModelCandidate, 0, len(candidates))
|
||||
for _, candidate := range candidates {
|
||||
if candidateAllowedByAccessRules(candidate, rules, subjects, level) {
|
||||
filtered = append(filtered, candidate)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func filterPlatformModelsByAPIKeyScopes(models []PlatformModel, scopes []string) []PlatformModel {
|
||||
filtered := make([]PlatformModel, 0, len(models))
|
||||
for _, model := range models {
|
||||
declaredTypes := append(StringList(nil), model.ModelType...)
|
||||
allowedTypes := modelaccess.FilterModelTypes(scopes, model.ModelType)
|
||||
if len(allowedTypes) == 0 {
|
||||
continue
|
||||
}
|
||||
model.ModelType = StringList(allowedTypes)
|
||||
model.Capabilities = filterPlatformModelTypeConfig(model.Capabilities, allowedTypes, declaredTypes)
|
||||
model.BaseCapabilities = filterPlatformModelTypeConfig(model.BaseCapabilities, allowedTypes, declaredTypes)
|
||||
model.CapabilityOverride = filterPlatformModelTypeConfig(model.CapabilityOverride, allowedTypes, declaredTypes)
|
||||
filtered = append(filtered, model)
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func filterPlatformModelTypeConfig(config map[string]any, allowedTypes []string, declaredTypes []string) map[string]any {
|
||||
if len(config) == 0 {
|
||||
return config
|
||||
}
|
||||
allowed := map[string]bool{}
|
||||
for _, modelType := range allowedTypes {
|
||||
allowed[modelType] = true
|
||||
}
|
||||
declared := map[string]bool{}
|
||||
for _, modelType := range declaredTypes {
|
||||
declared[modelType] = true
|
||||
}
|
||||
out := make(map[string]any, len(config))
|
||||
for key, value := range config {
|
||||
if key == "originalTypes" {
|
||||
original := stringValues(value)
|
||||
kept := make([]string, 0, len(original))
|
||||
for _, modelType := range original {
|
||||
if allowed[modelType] {
|
||||
kept = append(kept, modelType)
|
||||
}
|
||||
}
|
||||
if len(kept) > 0 {
|
||||
out[key] = kept
|
||||
}
|
||||
continue
|
||||
}
|
||||
if (declared[key] || knownPlatformModelType(key)) && !allowed[key] {
|
||||
continue
|
||||
}
|
||||
out[key] = value
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func knownPlatformModelType(value string) bool {
|
||||
switch value {
|
||||
case "text_generate", "text_embedding", "text_rerank", "tools_call",
|
||||
"image_generate", "image_edit", "image_analysis", "image_vectorize",
|
||||
"video_generate", "video_enhance", "image_to_video", "text_to_video", "video_edit", "video_reference", "video_first_last_frame", "video_understanding", "omni_video", "omni",
|
||||
"audio_generate", "music_generate", "audio_understanding", "text_to_speech", "voice_clone":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func stringValues(value any) []string {
|
||||
switch values := value.(type) {
|
||||
case []string:
|
||||
return values
|
||||
case StringList:
|
||||
return []string(values)
|
||||
case []any:
|
||||
out := make([]string, 0, len(values))
|
||||
for _, value := range values {
|
||||
if text, ok := value.(string); ok && text != "" {
|
||||
out = append(out, text)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func baselineAccessRuleSubjects(user *auth.User) map[string]bool {
|
||||
subjects := accessRuleSubjects(user)
|
||||
if user != nil && user.APIKeyID != "" {
|
||||
delete(subjects, "api_key:"+user.APIKeyID)
|
||||
}
|
||||
return subjects
|
||||
}
|
||||
|
||||
func apiKeyAccessRuleSubjects(user *auth.User) map[string]bool {
|
||||
subjects := map[string]bool{}
|
||||
if user != nil && strings.TrimSpace(user.APIKeyID) != "" {
|
||||
subjects["api_key:"+strings.TrimSpace(user.APIKeyID)] = true
|
||||
}
|
||||
return subjects
|
||||
}
|
||||
|
||||
func permissionLevel(user *auth.User) int {
|
||||
if user == nil {
|
||||
return 0
|
||||
}
|
||||
return auth.PermissionLevel(user.Roles)
|
||||
}
|
||||
|
||||
func accessRuleMatchesPlatformModel(rule AccessRule, model PlatformModel) bool {
|
||||
switch rule.ResourceType {
|
||||
case "platform":
|
||||
return rule.ResourceID == model.PlatformID
|
||||
case "platform_model":
|
||||
return rule.ResourceID == model.ID
|
||||
case "base_model":
|
||||
return rule.ResourceID != "" && rule.ResourceID == model.BaseModelID
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func accessRuleMatchesCandidate(rule AccessRule, candidate RuntimeModelCandidate) bool {
|
||||
switch rule.ResourceType {
|
||||
case "platform":
|
||||
return rule.ResourceID == candidate.PlatformID
|
||||
case "platform_model":
|
||||
return rule.ResourceID == candidate.PlatformModelID
|
||||
case "base_model":
|
||||
return rule.ResourceID != "" && rule.ResourceID == candidate.BaseModelID
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func diagnoseAPIKeyRules(apiKeyID string, rules []AccessRule, allModels []PlatformModel, enabledModels []PlatformModel, baselineModels []PlatformModel, scopedModels []PlatformModel, platforms []Platform) []APIKeyAccessRuleDiagnostic {
|
||||
platformNames := map[string]string{}
|
||||
for _, platform := range platforms {
|
||||
platformNames[platform.ID] = firstNonEmpty(platform.InternalName, platform.Name, platform.PlatformKey)
|
||||
}
|
||||
modelNames := map[string]string{}
|
||||
baseModelNames := map[string]string{}
|
||||
for _, model := range allModels {
|
||||
modelNames[model.ID] = firstNonEmpty(model.DisplayName, model.ModelName, model.ID)
|
||||
if model.BaseModelID != "" && baseModelNames[model.BaseModelID] == "" {
|
||||
baseModelNames[model.BaseModelID] = firstNonEmpty(model.DisplayName, model.ModelName, model.BaseModelID)
|
||||
}
|
||||
}
|
||||
diagnostics := make([]APIKeyAccessRuleDiagnostic, 0)
|
||||
for _, rule := range rules {
|
||||
if rule.SubjectType != "api_key" || rule.SubjectID != apiKeyID || rule.Status != "active" {
|
||||
continue
|
||||
}
|
||||
diagnostic := APIKeyAccessRuleDiagnostic{
|
||||
RuleID: rule.ID,
|
||||
ResourceType: rule.ResourceType,
|
||||
ResourceID: rule.ResourceID,
|
||||
ResourceName: accessRuleResourceName(rule, platformNames, modelNames, baseModelNames),
|
||||
Effect: rule.Effect,
|
||||
Effective: true,
|
||||
}
|
||||
switch {
|
||||
case !accessRuleMatchesAnyPlatformModel(rule, allModels) || !accessRuleMatchesAnyPlatformModel(rule, enabledModels):
|
||||
diagnostic.Effective = false
|
||||
diagnostic.Reason = "resource_unavailable"
|
||||
case !accessRuleMatchesAnyPlatformModel(rule, baselineModels):
|
||||
diagnostic.Effective = false
|
||||
diagnostic.Reason = "owner_access_revoked"
|
||||
case !accessRuleMatchesAnyPlatformModel(rule, scopedModels):
|
||||
diagnostic.Effective = false
|
||||
diagnostic.Reason = "scope_not_allowed"
|
||||
}
|
||||
diagnostics = append(diagnostics, diagnostic)
|
||||
}
|
||||
return diagnostics
|
||||
}
|
||||
|
||||
func accessRuleMatchesAnyPlatformModel(rule AccessRule, models []PlatformModel) bool {
|
||||
for _, model := range models {
|
||||
if accessRuleMatchesPlatformModel(rule, model) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func accessRuleResourceName(rule AccessRule, platformNames map[string]string, modelNames map[string]string, baseModelNames map[string]string) string {
|
||||
switch rule.ResourceType {
|
||||
case "platform":
|
||||
return firstNonEmpty(platformNames[rule.ResourceID], rule.ResourceID)
|
||||
case "platform_model":
|
||||
return firstNonEmpty(modelNames[rule.ResourceID], rule.ResourceID)
|
||||
case "base_model":
|
||||
return firstNonEmpty(baseModelNames[rule.ResourceID], rule.ResourceID)
|
||||
default:
|
||||
return rule.ResourceID
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||
)
|
||||
|
||||
func TestLayeredAccessDoesNotLetAPIKeyExpandBaseline(t *testing.T) {
|
||||
model := PlatformModel{ID: "model-1", PlatformID: "platform-1", BaseModelID: "base-1"}
|
||||
groupUser := &auth.User{GatewayUserID: "user-1", UserGroupID: "group-1", APIKeyID: "key-1"}
|
||||
baselineRules := []AccessRule{{
|
||||
SubjectType: "user_group", SubjectID: "group-1", ResourceType: "platform_model", ResourceID: "model-1", Effect: "deny",
|
||||
}}
|
||||
keyRules := []AccessRule{{
|
||||
SubjectType: "api_key", SubjectID: "key-1", ResourceType: "platform_model", ResourceID: "model-1", Effect: "allow",
|
||||
}}
|
||||
baseline := filterPlatformModelsByRuleSet([]PlatformModel{model}, baselineRules, baselineAccessRuleSubjects(groupUser), 0)
|
||||
actual := filterPlatformModelsByRuleSet(baseline, keyRules, apiKeyAccessRuleSubjects(groupUser), 0)
|
||||
if len(actual) != 0 {
|
||||
t.Fatalf("api key allow expanded denied baseline: %+v", actual)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAPIKeyAllowControlsOnlyMatchingKeys(t *testing.T) {
|
||||
model := PlatformModel{ID: "model-1", PlatformID: "platform-1"}
|
||||
rules := []AccessRule{{
|
||||
SubjectType: "api_key", SubjectID: "key-a", ResourceType: "platform_model", ResourceID: "model-1", Effect: "allow",
|
||||
}}
|
||||
for _, test := range []struct {
|
||||
keyID string
|
||||
want int
|
||||
}{{"key-a", 1}, {"key-b", 0}} {
|
||||
user := &auth.User{APIKeyID: test.keyID}
|
||||
actual := filterPlatformModelsByRuleSet([]PlatformModel{model}, rules, apiKeyAccessRuleSubjects(user), 0)
|
||||
if len(actual) != test.want {
|
||||
t.Fatalf("key %s received %d models, want %d", test.keyID, len(actual), test.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLayeredAccessDenyWinsAndNoRulesInherit(t *testing.T) {
|
||||
model := PlatformModel{ID: "model-1", PlatformID: "platform-1", ModelType: StringList{"text_generate"}}
|
||||
keyUser := &auth.User{APIKeyID: "key-a", APIKeyScopes: []string{"chat"}}
|
||||
keyUsers := map[string]*auth.User{"key-a": keyUser}
|
||||
if got := filterPlatformModelsByLayeredRuleSet(keyUser, []PlatformModel{model}, nil, nil, keyUsers); len(got) != 1 {
|
||||
t.Fatalf("key without rules did not inherit baseline: %+v", got)
|
||||
}
|
||||
rules := []AccessRule{
|
||||
{SubjectType: "api_key", SubjectID: "key-a", ResourceType: "platform_model", ResourceID: "model-1", Effect: "allow"},
|
||||
{SubjectType: "api_key", SubjectID: "key-a", ResourceType: "platform_model", ResourceID: "model-1", Effect: "deny"},
|
||||
}
|
||||
if got := filterPlatformModelsByLayeredRuleSet(keyUser, []PlatformModel{model}, nil, rules, keyUsers); len(got) != 0 {
|
||||
t.Fatalf("matching deny did not override allow: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectUserIgnoresAPIKeyRules(t *testing.T) {
|
||||
model := PlatformModel{ID: "model-1", PlatformID: "platform-1", ModelType: StringList{"text_generate"}}
|
||||
rules := []AccessRule{{
|
||||
SubjectType: "api_key", SubjectID: "key-a", ResourceType: "platform_model", ResourceID: "model-1", Effect: "allow",
|
||||
}}
|
||||
keyUsers := map[string]*auth.User{"key-a": {APIKeyID: "key-a", APIKeyScopes: []string{"chat"}}}
|
||||
if got := filterPlatformModelsByLayeredRuleSet(&auth.User{GatewayUserID: "user-1"}, []PlatformModel{model}, nil, rules, keyUsers); len(got) != 1 {
|
||||
t.Fatalf("direct user was constrained by API key exclusive rule: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIneffectiveAPIKeyRuleDoesNotControlAnotherKey(t *testing.T) {
|
||||
model := PlatformModel{ID: "model-1", PlatformID: "platform-1", ModelType: StringList{"text_generate"}}
|
||||
staleRule := AccessRule{SubjectType: "api_key", SubjectID: "key-a", ResourceType: "platform_model", ResourceID: "model-1", Effect: "allow"}
|
||||
groupDeny := AccessRule{SubjectType: "user_group", SubjectID: "group-a", ResourceType: "platform_model", ResourceID: "model-1", Effect: "deny"}
|
||||
keyUsers := map[string]*auth.User{
|
||||
"key-a": {APIKeyID: "key-a", UserGroupID: "group-a", APIKeyScopes: []string{"chat"}},
|
||||
}
|
||||
keyB := &auth.User{APIKeyID: "key-b", UserGroupID: "group-b", APIKeyScopes: []string{"chat"}}
|
||||
if got := filterPlatformModelsByLayeredRuleSet(keyB, []PlatformModel{model}, []AccessRule{groupDeny}, []AccessRule{staleRule}, keyUsers); len(got) != 1 {
|
||||
t.Fatalf("stale key-a rule blocked authorized key-b: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlatformRuleOnlyControlsModelsItsKeyCanAccess(t *testing.T) {
|
||||
allowed := PlatformModel{ID: "allowed", PlatformID: "platform-1", ModelType: StringList{"image_generate"}}
|
||||
denied := PlatformModel{ID: "denied", PlatformID: "platform-1", ModelType: StringList{"text_generate"}}
|
||||
rule := AccessRule{SubjectType: "api_key", SubjectID: "key-a", ResourceType: "platform", ResourceID: "platform-1", Effect: "allow"}
|
||||
groupDeny := AccessRule{SubjectType: "user_group", SubjectID: "group-a", ResourceType: "platform_model", ResourceID: "denied", Effect: "deny"}
|
||||
ruleUser := &auth.User{APIKeyID: "key-a", UserGroupID: "group-a", APIKeyScopes: []string{"image"}}
|
||||
users := map[string]*auth.User{"key-a": ruleUser}
|
||||
if got := effectiveAPIKeyRulesForPlatformModel([]AccessRule{rule}, []AccessRule{groupDeny}, users, allowed); len(got) != 1 {
|
||||
t.Fatalf("platform rule should control the allowed image model: %+v", got)
|
||||
}
|
||||
if got := effectiveAPIKeyRulesForPlatformModel([]AccessRule{rule}, []AccessRule{groupDeny}, users, denied); len(got) != 0 {
|
||||
t.Fatalf("platform rule controlled a group-denied or scope-denied model: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilterPlatformModelsByAPIKeyScopesPrunesCapabilities(t *testing.T) {
|
||||
models := []PlatformModel{{
|
||||
ID: "model-1",
|
||||
ModelType: StringList{"text_generate", "image_generate"},
|
||||
Capabilities: map[string]any{
|
||||
"text_generate": map[string]any{"max_context_tokens": 128000},
|
||||
"image_generate": map[string]any{"aspect_ratio_allowed": []any{"1:1"}},
|
||||
"originalTypes": []any{"text_generate", "image_generate"},
|
||||
"shared": true,
|
||||
},
|
||||
}}
|
||||
actual := filterPlatformModelsByAPIKeyScopes(models, []string{"image"})
|
||||
if len(actual) != 1 || !reflect.DeepEqual(actual[0].ModelType, StringList{"image_generate"}) {
|
||||
t.Fatalf("scope-filtered models = %+v", actual)
|
||||
}
|
||||
if _, exists := actual[0].Capabilities["text_generate"]; exists {
|
||||
t.Fatalf("text capability leaked into image scope: %+v", actual[0].Capabilities)
|
||||
}
|
||||
if _, exists := actual[0].Capabilities["image_generate"]; !exists {
|
||||
t.Fatalf("image capability was removed: %+v", actual[0].Capabilities)
|
||||
}
|
||||
if !reflect.DeepEqual(actual[0].Capabilities["originalTypes"], []string{"image_generate"}) {
|
||||
t.Fatalf("originalTypes not pruned: %+v", actual[0].Capabilities["originalTypes"])
|
||||
}
|
||||
if actual[0].Capabilities["shared"] != true {
|
||||
t.Fatalf("shared capability metadata was removed: %+v", actual[0].Capabilities)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiagnoseAPIKeyRulesExplainsEachInactiveLayer(t *testing.T) {
|
||||
rules := []AccessRule{
|
||||
{ID: "gone", SubjectType: "api_key", SubjectID: "key-1", ResourceType: "platform_model", ResourceID: "gone", Effect: "allow", Status: "active"},
|
||||
{ID: "revoked", SubjectType: "api_key", SubjectID: "key-1", ResourceType: "platform_model", ResourceID: "revoked", Effect: "allow", Status: "active"},
|
||||
{ID: "scope", SubjectType: "api_key", SubjectID: "key-1", ResourceType: "platform_model", ResourceID: "scope", Effect: "deny", Status: "active"},
|
||||
{ID: "effective", SubjectType: "api_key", SubjectID: "key-1", ResourceType: "platform_model", ResourceID: "effective", Effect: "allow", Status: "active"},
|
||||
}
|
||||
all := []PlatformModel{{ID: "revoked"}, {ID: "scope"}, {ID: "effective"}}
|
||||
enabled := append([]PlatformModel(nil), all...)
|
||||
baseline := []PlatformModel{{ID: "scope"}, {ID: "effective"}}
|
||||
scoped := []PlatformModel{{ID: "effective"}}
|
||||
diagnostics := diagnoseAPIKeyRules("key-1", rules, all, enabled, baseline, scoped, nil)
|
||||
want := map[string]string{
|
||||
"gone": "resource_unavailable", "revoked": "owner_access_revoked", "scope": "scope_not_allowed", "effective": "",
|
||||
}
|
||||
for _, diagnostic := range diagnostics {
|
||||
if diagnostic.Reason != want[diagnostic.RuleID] {
|
||||
t.Fatalf("diagnostic %s reason = %q, want %q", diagnostic.RuleID, diagnostic.Reason, want[diagnostic.RuleID])
|
||||
}
|
||||
if diagnostic.Effective != (diagnostic.RuleID == "effective") {
|
||||
t.Fatalf("diagnostic %s effective = %v", diagnostic.RuleID, diagnostic.Effective)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -244,7 +244,7 @@ SELECT EXISTS (
|
||||
if !exists {
|
||||
return nil, pgx.ErrNoRows
|
||||
}
|
||||
if err := s.ensureAPIKeyAccessRuleResourcesAllowed(ctx, user, input.UpsertResources); err != nil {
|
||||
if err := s.ensureAPIKeyAccessRuleResourcesAllowed(ctx, user, input.SubjectID, input.UpsertResources); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err := s.BatchAccessRules(ctx, input); err != nil {
|
||||
@@ -254,32 +254,7 @@ SELECT EXISTS (
|
||||
}
|
||||
|
||||
func (s *Store) filterCandidatesByAccessRules(ctx context.Context, user *auth.User, candidates []RuntimeModelCandidate) ([]RuntimeModelCandidate, error) {
|
||||
if len(candidates) == 0 {
|
||||
return candidates, nil
|
||||
}
|
||||
resources := candidateAccessResources(candidates)
|
||||
if len(resources) == 0 {
|
||||
return candidates, nil
|
||||
}
|
||||
rules, err := s.listActiveAccessRulesForResources(ctx, resources)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(rules) == 0 {
|
||||
return candidates, nil
|
||||
}
|
||||
subjects := accessRuleSubjects(user)
|
||||
level := 0
|
||||
if user != nil {
|
||||
level = auth.PermissionLevel(user.Roles)
|
||||
}
|
||||
filtered := candidates[:0]
|
||||
for _, candidate := range candidates {
|
||||
if candidateAllowedByAccessRules(candidate, rules, subjects, level) {
|
||||
filtered = append(filtered, candidate)
|
||||
}
|
||||
}
|
||||
return filtered, nil
|
||||
return s.filterRuntimeCandidatesByLayeredAccess(ctx, user, candidates)
|
||||
}
|
||||
|
||||
func (s *Store) ListAccessiblePlatformModels(ctx context.Context, user *auth.User) ([]PlatformModel, error) {
|
||||
@@ -302,35 +277,22 @@ func (s *Store) listPlatformModelsForAccessRules(ctx context.Context, user *auth
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
models, err := s.ListModels(ctx)
|
||||
models, _, err := s.enabledPlatformModels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
platforms, err := s.ListPlatforms(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
if excludedSubjectTypes["api_key"] {
|
||||
return s.filterPlatformModelsByBaselineAccess(ctx, accessUser, models)
|
||||
}
|
||||
enabledPlatforms := map[string]bool{}
|
||||
for _, platform := range platforms {
|
||||
if platform.Status == "enabled" {
|
||||
enabledPlatforms[platform.ID] = true
|
||||
}
|
||||
}
|
||||
enabled := make([]PlatformModel, 0, len(models))
|
||||
for _, model := range models {
|
||||
if model.Enabled && enabledPlatforms[model.PlatformID] {
|
||||
enabled = append(enabled, model)
|
||||
}
|
||||
}
|
||||
return s.filterPlatformModelsByAccessRules(ctx, accessUser, enabled, excludedSubjectTypes)
|
||||
return s.filterPlatformModelsByLayeredAccess(ctx, accessUser, models)
|
||||
}
|
||||
|
||||
func (s *Store) ensureAPIKeyAccessRuleResourcesAllowed(ctx context.Context, user *auth.User, resources []AccessRuleResourceInput) error {
|
||||
func (s *Store) ensureAPIKeyAccessRuleResourcesAllowed(ctx context.Context, user *auth.User, apiKeyID string, resources []AccessRuleResourceInput) error {
|
||||
resources = dedupeAccessRuleResources(resources)
|
||||
if len(resources) == 0 {
|
||||
return nil
|
||||
}
|
||||
allowed, err := s.accessibleAccessRuleResources(ctx, user)
|
||||
allowed, err := s.accessibleAccessRuleResources(ctx, user, apiKeyID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -342,8 +304,8 @@ func (s *Store) ensureAPIKeyAccessRuleResourcesAllowed(ctx context.Context, user
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Store) accessibleAccessRuleResources(ctx context.Context, user *auth.User) (map[string]bool, error) {
|
||||
models, err := s.ListAPIKeyAssignablePlatformModels(ctx, user)
|
||||
func (s *Store) accessibleAccessRuleResources(ctx context.Context, user *auth.User, apiKeyID string) (map[string]bool, error) {
|
||||
models, _, err := s.ListAPIKeyAssignablePlatformModelsForKey(ctx, user, apiKeyID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -368,25 +330,28 @@ func (s *Store) resolveCurrentAccessUser(ctx context.Context, user *auth.User) (
|
||||
}
|
||||
next := *user
|
||||
var userGroupID string
|
||||
var userGroupKey string
|
||||
var err error
|
||||
if strings.TrimSpace(user.APIKeyID) != "" {
|
||||
err = s.pool.QueryRow(ctx, `
|
||||
SELECT COALESCE(k.user_group_id::text, u.default_user_group_id::text, '')
|
||||
SELECT COALESCE(k.user_group_id::text, u.default_user_group_id::text, ''), COALESCE(g.group_key, '')
|
||||
FROM gateway_users u
|
||||
JOIN gateway_api_keys k ON k.gateway_user_id = u.id
|
||||
LEFT JOIN gateway_user_groups g ON g.id = COALESCE(k.user_group_id, u.default_user_group_id)
|
||||
WHERE u.id = $1::uuid
|
||||
AND k.id = $2::uuid
|
||||
AND u.status = 'active'
|
||||
AND u.deleted_at IS NULL
|
||||
AND k.status = 'active'
|
||||
AND k.deleted_at IS NULL`, gatewayUserID, user.APIKeyID).Scan(&userGroupID)
|
||||
AND k.deleted_at IS NULL`, gatewayUserID, user.APIKeyID).Scan(&userGroupID, &userGroupKey)
|
||||
} else {
|
||||
err = s.pool.QueryRow(ctx, `
|
||||
SELECT COALESCE(default_user_group_id::text, '')
|
||||
FROM gateway_users
|
||||
WHERE id = $1::uuid
|
||||
AND status = 'active'
|
||||
AND deleted_at IS NULL`, gatewayUserID).Scan(&userGroupID)
|
||||
SELECT COALESCE(u.default_user_group_id::text, ''), COALESCE(g.group_key, '')
|
||||
FROM gateway_users u
|
||||
LEFT JOIN gateway_user_groups g ON g.id = u.default_user_group_id
|
||||
WHERE u.id = $1::uuid
|
||||
AND u.status = 'active'
|
||||
AND u.deleted_at IS NULL`, gatewayUserID).Scan(&userGroupID, &userGroupKey)
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
@@ -395,6 +360,11 @@ WHERE id = $1::uuid
|
||||
return nil, err
|
||||
}
|
||||
next.UserGroupID = userGroupID
|
||||
next.UserGroupKey = userGroupKey
|
||||
next.UserGroupKeys = nil
|
||||
if userGroupKey != "" {
|
||||
next.UserGroupKeys = []string{userGroupKey}
|
||||
}
|
||||
return &next, nil
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user