Files
easyai-ai-gateway/apps/api/internal/store/access_policy.go
T
wangbo 7376d6fab6 refactor(access): 统一分层白名单权限语义
取消跨主体专属占用,按租户、用户组、用户、当前 API Key 和 scope 分层求交,并在任务落库前统一校验候选。\n\n增加旧 allow 规则归档清理迁移、脱敏审计工具和回滚运行手册,补齐主体隔离、deny 优先及列表与运行时一致性测试。
2026-08-03 15:43:49 +08:00

472 lines
15 KiB
Go

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
}
}