feat: refine api key permissions and admin routes
This commit is contained in:
@@ -3,6 +3,7 @@ package store
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||
@@ -66,6 +67,35 @@ ORDER BY resource_type ASC, priority ASC, subject_type ASC, created_at DESC`)
|
||||
return items, rows.Err()
|
||||
}
|
||||
|
||||
func (s *Store) ListAPIKeyAccessRules(ctx context.Context, user *auth.User) ([]AccessRule, error) {
|
||||
gatewayUserID := localGatewayUserID(user)
|
||||
if gatewayUserID == "" {
|
||||
return nil, ErrLocalUserRequired
|
||||
}
|
||||
rows, err := s.pool.Query(ctx, `
|
||||
SELECT `+apiKeyAccessRuleColumns+`
|
||||
FROM gateway_access_rules ar
|
||||
JOIN gateway_api_keys k ON k.id = ar.subject_id
|
||||
WHERE ar.subject_type = 'api_key'
|
||||
AND k.gateway_user_id = $1::uuid
|
||||
AND k.deleted_at IS NULL
|
||||
ORDER BY ar.resource_type ASC, ar.priority ASC, ar.created_at DESC`, gatewayUserID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
items := make([]AccessRule, 0)
|
||||
for rows.Next() {
|
||||
item, err := scanAccessRule(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
return items, rows.Err()
|
||||
}
|
||||
|
||||
func (s *Store) CreateAccessRule(ctx context.Context, input AccessRuleInput) (AccessRule, error) {
|
||||
input = normalizeAccessRuleInput(input)
|
||||
conditions, _ := json.Marshal(emptyObjectIfNil(input.Conditions))
|
||||
@@ -191,6 +221,38 @@ DO UPDATE SET priority = EXCLUDED.priority,
|
||||
return s.ListAccessRules(ctx)
|
||||
}
|
||||
|
||||
func (s *Store) BatchAPIKeyAccessRules(ctx context.Context, input AccessRuleBatchInput, user *auth.User) ([]AccessRule, error) {
|
||||
gatewayUserID := localGatewayUserID(user)
|
||||
if gatewayUserID == "" {
|
||||
return nil, ErrLocalUserRequired
|
||||
}
|
||||
input = normalizeAccessRuleBatchInput(input)
|
||||
if input.SubjectType != "api_key" {
|
||||
return nil, pgx.ErrNoRows
|
||||
}
|
||||
var exists bool
|
||||
if err := s.pool.QueryRow(ctx, `
|
||||
SELECT EXISTS (
|
||||
SELECT 1
|
||||
FROM gateway_api_keys
|
||||
WHERE id = $1::uuid
|
||||
AND gateway_user_id = $2::uuid
|
||||
AND deleted_at IS NULL
|
||||
)`, input.SubjectID, gatewayUserID).Scan(&exists); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !exists {
|
||||
return nil, pgx.ErrNoRows
|
||||
}
|
||||
if err := s.ensureAPIKeyAccessRuleResourcesAllowed(ctx, user, input.UpsertResources); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err := s.BatchAccessRules(ctx, input); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.ListAPIKeyAccessRules(ctx, user)
|
||||
}
|
||||
|
||||
func (s *Store) filterCandidatesByAccessRules(ctx context.Context, user *auth.User, candidates []RuntimeModelCandidate) ([]RuntimeModelCandidate, error) {
|
||||
if len(candidates) == 0 {
|
||||
return candidates, nil
|
||||
@@ -221,6 +283,10 @@ func (s *Store) filterCandidatesByAccessRules(ctx context.Context, user *auth.Us
|
||||
}
|
||||
|
||||
func (s *Store) ListAccessiblePlatformModels(ctx context.Context, user *auth.User) ([]PlatformModel, error) {
|
||||
accessUser, err := s.resolveCurrentAccessUser(ctx, user)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
models, err := s.ListModels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -241,7 +307,80 @@ func (s *Store) ListAccessiblePlatformModels(ctx context.Context, user *auth.Use
|
||||
enabled = append(enabled, model)
|
||||
}
|
||||
}
|
||||
return s.filterPlatformModelsByAccessRules(ctx, user, enabled)
|
||||
return s.filterPlatformModelsByAccessRules(ctx, accessUser, enabled)
|
||||
}
|
||||
|
||||
func (s *Store) ensureAPIKeyAccessRuleResourcesAllowed(ctx context.Context, user *auth.User, resources []AccessRuleResourceInput) error {
|
||||
resources = dedupeAccessRuleResources(resources)
|
||||
if len(resources) == 0 {
|
||||
return nil
|
||||
}
|
||||
allowed, err := s.accessibleAccessRuleResources(ctx, user)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, resource := range resources {
|
||||
if !allowed[resource.ResourceType+":"+resource.ResourceID] {
|
||||
return ErrAccessRuleResourceDenied
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Store) accessibleAccessRuleResources(ctx context.Context, user *auth.User) (map[string]bool, error) {
|
||||
models, err := s.ListAccessiblePlatformModels(ctx, user)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
allowed := map[string]bool{}
|
||||
for _, model := range models {
|
||||
allowed["platform:"+model.PlatformID] = true
|
||||
allowed["platform_model:"+model.ID] = true
|
||||
if model.BaseModelID != "" {
|
||||
allowed["base_model:"+model.BaseModelID] = true
|
||||
}
|
||||
}
|
||||
return allowed, nil
|
||||
}
|
||||
|
||||
func (s *Store) resolveCurrentAccessUser(ctx context.Context, user *auth.User) (*auth.User, error) {
|
||||
if user == nil {
|
||||
return nil, nil
|
||||
}
|
||||
gatewayUserID := localGatewayUserID(user)
|
||||
if gatewayUserID == "" {
|
||||
return user, nil
|
||||
}
|
||||
next := *user
|
||||
var userGroupID 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, '')
|
||||
FROM gateway_users u
|
||||
JOIN gateway_api_keys k ON k.gateway_user_id = u.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)
|
||||
} 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)
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return &next, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
next.UserGroupID = userGroupID
|
||||
return &next, nil
|
||||
}
|
||||
|
||||
func (s *Store) filterPlatformModelsByAccessRules(ctx context.Context, user *auth.User, models []PlatformModel) ([]PlatformModel, error) {
|
||||
@@ -432,6 +571,10 @@ const accessRuleColumns = `
|
||||
id::text, subject_type, subject_id::text, resource_type, resource_id::text, effect,
|
||||
priority, min_permission_level, conditions, metadata, status, created_at, updated_at`
|
||||
|
||||
const apiKeyAccessRuleColumns = `
|
||||
ar.id::text, ar.subject_type, ar.subject_id::text, ar.resource_type, ar.resource_id::text, ar.effect,
|
||||
ar.priority, ar.min_permission_level, ar.conditions, ar.metadata, ar.status, ar.created_at, ar.updated_at`
|
||||
|
||||
func scanAccessRule(row scanner) (AccessRule, error) {
|
||||
var item AccessRule
|
||||
var conditions []byte
|
||||
|
||||
Reference in New Issue
Block a user