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) { layers := accessRuleLayers(user, true) rules, err := s.listActiveAccessRulesForLayers(ctx, layers) if err != nil { return nil, err } filtered := filterPlatformModelsByAccessLayers(models, rules, layers, permissionLevel(user)) if user != nil && strings.TrimSpace(user.APIKeyID) != "" { filtered = filterPlatformModelsByAPIKeyScopes(filtered, user.APIKeyScopes) } return filtered, nil } func (s *Store) filterPlatformModelsByBaselineAccess(ctx context.Context, user *auth.User, models []PlatformModel) ([]PlatformModel, error) { layers := accessRuleLayers(user, false) rules, err := s.listActiveAccessRulesForLayers(ctx, layers) if err != nil { return nil, err } return filterPlatformModelsByAccessLayers(models, rules, layers, 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 } layers := accessRuleLayers(accessUser, true) rules, err := s.listActiveAccessRulesForLayers(ctx, layers) if err != nil { return nil, err } filtered := filterCandidatesByAccessLayers(candidates, rules, layers, permissionLevel(accessUser)) if accessUser != nil && strings.TrimSpace(accessUser.APIKeyID) != "" { scoped := make([]RuntimeModelCandidate, 0, len(filtered)) for _, candidate := range filtered { if modelaccess.ScopeAllowsModelType(accessUser.APIKeyScopes, candidate.ModelType) { scoped = append(scoped, candidate) } } filtered = scoped } 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 } baseline, err := s.filterPlatformModelsByBaselineAccess(ctx, accessUser, enabledModels) if err != nil { return nil, nil, err } 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 } type accessRuleLayer struct { subjectType string subjectIDs map[string]bool } var accessRuleLayerOrder = []string{"tenant", "user_group", "user", "api_key"} func accessRuleLayers(user *auth.User, includeAPIKey bool) []accessRuleLayer { subjects := accessRuleSubjects(user) layers := make([]accessRuleLayer, 0, len(accessRuleLayerOrder)) for _, subjectType := range accessRuleLayerOrder { if subjectType == "api_key" && !includeAPIKey { continue } prefix := subjectType + ":" ids := map[string]bool{} for subject := range subjects { if strings.HasPrefix(subject, prefix) { ids[strings.TrimPrefix(subject, prefix)] = true } } if len(ids) > 0 { layers = append(layers, accessRuleLayer{subjectType: subjectType, subjectIDs: ids}) } } return layers } func (s *Store) listActiveAccessRulesForLayers(ctx context.Context, layers []accessRuleLayer) ([]AccessRule, error) { subjects := make([]string, 0) for _, layer := range layers { for id := range layer.subjectIDs { subjects = append(subjects, layer.subjectType+":"+id) } } if len(subjects) == 0 { return nil, nil } rows, err := s.pool.Query(ctx, ` SELECT `+accessRuleColumns+` FROM gateway_access_rules WHERE status = 'active' AND (subject_type || ':' || subject_id::text) = ANY($1) ORDER BY subject_type ASC, priority ASC, created_at ASC`, subjects) if err != nil { return nil, err } defer rows.Close() rules := make([]AccessRule, 0) for rows.Next() { item, err := scanAccessRule(rows) if err != nil { return nil, err } rules = append(rules, item) } return rules, rows.Err() } func filterPlatformModelsByAccessLayers(models []PlatformModel, rules []AccessRule, layers []accessRuleLayer, level int) []PlatformModel { filtered := append([]PlatformModel(nil), models...) for _, layer := range layers { layerRules := accessRulesForLayer(rules, layer) if len(layerRules) == 0 { continue } next := make([]PlatformModel, 0, len(filtered)) for _, model := range filtered { if platformModelAllowedBySubjectLayer(model, layerRules, level) { next = append(next, model) } } filtered = next } return filtered } func filterCandidatesByAccessLayers(candidates []RuntimeModelCandidate, rules []AccessRule, layers []accessRuleLayer, level int) []RuntimeModelCandidate { filtered := append([]RuntimeModelCandidate(nil), candidates...) for _, layer := range layers { layerRules := accessRulesForLayer(rules, layer) if len(layerRules) == 0 { continue } next := make([]RuntimeModelCandidate, 0, len(filtered)) for _, candidate := range filtered { if candidateAllowedBySubjectLayer(candidate, layerRules, level) { next = append(next, candidate) } } filtered = next } return filtered } func accessRulesForLayer(rules []AccessRule, layer accessRuleLayer) []AccessRule { filtered := make([]AccessRule, 0) for _, rule := range rules { if rule.SubjectType == layer.subjectType && layer.subjectIDs[rule.SubjectID] { filtered = append(filtered, rule) } } return filtered } func platformModelAllowedBySubjectLayer(model PlatformModel, rules []AccessRule, level int) bool { return allowedBySubjectLayer(rules, level, func(rule AccessRule) bool { return accessRuleMatchesPlatformModel(rule, model) }) } func candidateAllowedBySubjectLayer(candidate RuntimeModelCandidate, rules []AccessRule, level int) bool { return allowedBySubjectLayer(rules, level, func(rule AccessRule) bool { return accessRuleMatchesCandidate(rule, candidate) }) } func allowedBySubjectLayer(rules []AccessRule, level int, matches func(AccessRule) bool) bool { hasAllow := false matchedAllow := false for _, rule := range rules { switch rule.Effect { case "allow": hasAllow = true if level >= rule.MinPermissionLevel && matches(rule) { matchedAllow = true } case "deny": if matches(rule) { return false } } } return !hasAllow || matchedAllow } 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 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 } }