Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions internal/tests/mock_plugin_services.go
Original file line number Diff line number Diff line change
Expand Up @@ -49,3 +49,11 @@ func (m *MockOrganizationService) ExistsByID(ctx context.Context, organizationID
args := m.Called(ctx, organizationID)
return args.Bool(0), args.Error(1)
}

func (m *MockOrganizationService) GetUserPermissionsInOrganization(ctx context.Context, userID string, organizationID string) ([]string, error) {
args := m.Called(ctx, userID, organizationID)
if args.Get(0) == nil {
return nil, args.Error(1)
}
return args.Get(0).([]string), args.Error(1)
}
35 changes: 19 additions & 16 deletions plugins/access-control/services/access_control_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@ package services

import (
"context"
"time"

coreerrors "github.com/Authula/authula/core/errors"
"github.com/Authula/authula/plugins/access-control/repositories"
Expand Down Expand Up @@ -30,34 +29,38 @@ func (s *AccessControlService) RoleExists(ctx context.Context, roleName string)
return role != nil && role.ID != "", nil
}

func (s *AccessControlService) ValidateRoleAssignment(ctx context.Context, roleName string, assignerUserID *string) (bool, error) {
func (s *AccessControlService) GetRolePermissionsByName(ctx context.Context, roleName string) ([]string, error) {
role, err := s.rolesService.GetRoleByName(ctx, nil, roleName)
if err != nil {
return false, err
return nil, err
}
if role == nil || role.ID == "" {
return false, coreerrors.ErrNotFound
}

if assignerUserID == nil || *assignerUserID == "" {
return false, nil
return nil, coreerrors.ErrNotFound
}

assignerRoles, err := s.userRolesService.GetUserRoles(ctx, nil, *assignerUserID)
details, err := s.rolesService.GetRoleByID(ctx, nil, role.ID)
if err != nil {
return false, err
return nil, err
}

highestWeight, activeCount := determineHighestActiveRoleWeight(assignerRoles, time.Now().UTC())
if activeCount == 0 {
return false, coreerrors.ErrForbidden
permissions := make([]string, 0, len(details.Permissions))
for _, permission := range details.Permissions {
permissions = append(permissions, permission.PermissionKey)
}

if role.Weight > highestWeight {
return false, coreerrors.ErrForbidden
return permissions, nil
}

func (s *AccessControlService) GetRoleWeightByName(ctx context.Context, roleName string) (int, error) {
role, err := s.rolesService.GetRoleByName(ctx, nil, roleName)
if err != nil {
return 0, err
}
if role == nil || role.ID == "" {
return 0, coreerrors.ErrNotFound
}

return true, nil
return role.Weight, nil
}

func (s *AccessControlService) ValidatePermissionKeys(ctx context.Context, permissionKeys []string) error {
Expand Down
143 changes: 100 additions & 43 deletions plugins/access-control/services/access_control_service_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,8 @@ package services

import (
"context"
"slices"
"testing"
"time"

"github.com/stretchr/testify/mock"

Expand Down Expand Up @@ -142,66 +142,125 @@ func TestAccessControlServiceValidatePermissionKeys(t *testing.T) {
}
}

func TestAccessControlServiceValidateRoleAssignment(t *testing.T) {
func TestAccessControlServiceGetRolePermissionsByName(t *testing.T) {
t.Parallel()

assignerUserID := func() *string { value := "assigner-user-1"; return &value }()

tests := []struct {
name string
roleName string
assigner *string
setup func(*accesscontroltests.MockRolesRepository, *accesscontroltests.MockUserRolesRepository)
wantErr error
wantOK bool
name string
roleName string
setup func(*accesscontroltests.MockRolesRepository, *accesscontroltests.MockRolePermissionsRepository)
wantErr error
wantPerms []string
}{
{
name: "role not found",
roleName: "missing",
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, userRolesRepo *accesscontroltests.MockUserRolesRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "missing").Return((*types.Role)(nil), nil).Once()
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, rolePermRepo *accesscontroltests.MockRolePermissionsRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "missing").Return((*types.Role)(nil), coreerrors.ErrNotFound).Once()
},
wantErr: coreerrors.ErrNotFound,
},
{
name: "nil role treated as not found",
roleName: "ghost",
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, rolePermRepo *accesscontroltests.MockRolePermissionsRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "ghost").Return((*types.Role)(nil), nil).Once()
},
wantErr: coreerrors.ErrNotFound,
},
{
name: "nil assigner is not allowed",
name: "role permissions resolved",
roleName: "editor",
assigner: nil,
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, userRolesRepo *accesscontroltests.MockUserRolesRepository) {
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, rolePermRepo *accesscontroltests.MockRolePermissionsRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "editor").Return(&types.Role{ID: "role-1", Name: "editor", Weight: 10}, nil).Once()
rolesRepo.On("GetRoleByID", mock.Anything, "role-1").Return(&types.Role{ID: "role-1", Name: "editor", Weight: 10}, nil).Once()
rolePermRepo.On("GetRolePermissions", mock.Anything, "role-1").Return([]types.UserPermissionInfo{
{PermissionKey: "organizations:members:list"},
{PermissionKey: "organizations:members:read"},
}, nil).Once()
},
wantOK: false,
wantPerms: []string{"organizations:members:list", "organizations:members:read"},
},
{
name: "forbidden when assigner has no active roles",
name: "repository error propagates",
roleName: "editor",
assigner: assignerUserID,
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, userRolesRepo *accesscontroltests.MockUserRolesRepository) {
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, rolePermRepo *accesscontroltests.MockRolePermissionsRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "editor").Return(&types.Role{ID: "role-1", Name: "editor", Weight: 10}, nil).Once()
userRolesRepo.On("GetUserRoles", mock.Anything, "assigner-user-1").Return([]types.UserRoleInfo{{RoleID: "role-old", RoleName: "old", RoleWeight: 100, ExpiresAt: func() *time.Time { value := time.Now().UTC().Add(-time.Hour); return &value }()}}, nil).Once()
rolesRepo.On("GetRoleByID", mock.Anything, "role-1").Return(&types.Role{ID: "role-1", Name: "editor", Weight: 10}, nil).Once()
rolePermRepo.On("GetRolePermissions", mock.Anything, "role-1").Return(([]types.UserPermissionInfo)(nil), coreerrors.ErrForbidden).Once()
},
wantErr: coreerrors.ErrForbidden,
},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

rolesRepo := &accesscontroltests.MockRolesRepository{}
rolePermRepo := &accesscontroltests.MockRolePermissionsRepository{}
if tc.setup != nil {
tc.setup(rolesRepo, rolePermRepo)
}

service := NewAccessControlService(NewRolesService(rolesRepo, rolePermRepo, nil), NewUserRolesService(nil, nil), nil)
perms, err := service.GetRolePermissionsByName(context.Background(), tc.roleName)
if err != tc.wantErr {
t.Fatalf("expected err %v, got %v", tc.wantErr, err)
}
if tc.wantErr != nil {
if perms != nil {
t.Fatalf("expected nil permissions, got %v", perms)
}
} else if !slices.Equal(perms, tc.wantPerms) {
t.Fatalf("expected permissions %v, got %v", tc.wantPerms, perms)
}

rolesRepo.AssertExpectations(t)
rolePermRepo.AssertExpectations(t)
})
}
}

func TestAccessControlServiceGetRoleWeightByName(t *testing.T) {
t.Parallel()

tests := []struct {
name string
roleName string
setup func(*accesscontroltests.MockRolesRepository)
wantErr error
wantWeight int
}{
{
name: "expired roles are ignored",
roleName: "editor",
assigner: assignerUserID,
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, userRolesRepo *accesscontroltests.MockUserRolesRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "editor").Return(&types.Role{ID: "role-1", Name: "editor", Weight: 20}, nil).Once()
userRolesRepo.On("GetUserRoles", mock.Anything, "assigner-user-1").Return([]types.UserRoleInfo{
{RoleID: "role-expired", RoleName: "expired", RoleWeight: 100, ExpiresAt: func() *time.Time { value := time.Now().UTC().Add(-time.Hour); return &value }()},
{RoleID: "role-active", RoleName: "active", RoleWeight: 30},
}, nil).Once()
name: "role not found",
roleName: "missing",
setup: func(rolesRepo *accesscontroltests.MockRolesRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "missing").Return((*types.Role)(nil), coreerrors.ErrNotFound).Once()
},
wantOK: true,
wantErr: coreerrors.ErrNotFound,
},
{
name: "nil role treated as not found",
roleName: "ghost",
setup: func(rolesRepo *accesscontroltests.MockRolesRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "ghost").Return((*types.Role)(nil), nil).Once()
},
wantErr: coreerrors.ErrNotFound,
},
{
name: "forbidden when target exceeds assigner weight",
name: "role weight resolved",
roleName: "admin",
assigner: assignerUserID,
setup: func(rolesRepo *accesscontroltests.MockRolesRepository, userRolesRepo *accesscontroltests.MockUserRolesRepository) {
setup: func(rolesRepo *accesscontroltests.MockRolesRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "admin").Return(&types.Role{ID: "role-2", Name: "admin", Weight: 80}, nil).Once()
userRolesRepo.On("GetUserRoles", mock.Anything, "assigner-user-1").Return([]types.UserRoleInfo{{RoleID: "role-member", RoleName: "member", RoleWeight: 10}}, nil).Once()
},
wantWeight: 80,
},
{
name: "repository error propagates",
roleName: "admin",
setup: func(rolesRepo *accesscontroltests.MockRolesRepository) {
rolesRepo.On("GetRoleByName", mock.Anything, "admin").Return((*types.Role)(nil), coreerrors.ErrForbidden).Once()
},
wantErr: coreerrors.ErrForbidden,
},
Expand All @@ -212,26 +271,24 @@ func TestAccessControlServiceValidateRoleAssignment(t *testing.T) {
t.Parallel()

rolesRepo := &accesscontroltests.MockRolesRepository{}
userRolesRepo := &accesscontroltests.MockUserRolesRepository{}
if tc.setup != nil {
tc.setup(rolesRepo, userRolesRepo)
tc.setup(rolesRepo)
}

service := NewAccessControlService(NewRolesService(rolesRepo, nil, userRolesRepo), NewUserRolesService(userRolesRepo, rolesRepo), nil)
ok, err := service.ValidateRoleAssignment(context.Background(), tc.roleName, tc.assigner)
service := NewAccessControlService(NewRolesService(rolesRepo, nil, nil), NewUserRolesService(nil, nil), nil)
weight, err := service.GetRoleWeightByName(context.Background(), tc.roleName)
if err != tc.wantErr {
t.Fatalf("expected err %v, got %v", tc.wantErr, err)
}
if tc.wantErr != nil {
if ok {
t.Fatalf("expected false, got true")
if weight != 0 {
t.Fatalf("expected 0 weight, got %d", weight)
}
} else if ok != tc.wantOK {
t.Fatalf("unexpected result %v", ok)
} else if weight != tc.wantWeight {
t.Fatalf("expected weight %d, got %d", tc.wantWeight, weight)
}

rolesRepo.AssertExpectations(t)
userRolesRepo.AssertExpectations(t)
})
}
}
2 changes: 1 addition & 1 deletion plugins/api-key/plugin.go
Original file line number Diff line number Diff line change
Expand Up @@ -87,7 +87,7 @@ func (p *ApiKeyPlugin) Init(ctx *models.PluginContext) error {
p.rateLimiterService = rateLimiterService
}

var organizationService rootservices.OrganizationService
var organizationService apiservices.OrganizationLookupService
if p.config.AllowOrgKeys {
orgSvc, ok := ctx.ServiceRegistry.Get(models.ServiceOrganization.String()).(orgplugins.OrganizationLookupService)
if !ok {
Expand Down
23 changes: 17 additions & 6 deletions plugins/api-key/services/api_key_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ type apiKeyService struct {
tokenService rootservices.TokenService
accessControlService rootservices.AccessControlService
rateLimiterService rootservices.RateLimiterService
organizationService rootservices.OrganizationService
organizationService OrganizationLookupService
apiKeyRepo repositories.ApiKeyRepository
}

Expand All @@ -36,7 +36,7 @@ func NewApiKeyService(
tokenService rootservices.TokenService,
accessControlService rootservices.AccessControlService,
rateLimiterService rootservices.RateLimiterService,
organizationService rootservices.OrganizationService,
organizationService OrganizationLookupService,
apiKeyRepo repositories.ApiKeyRepository,
) ApiKeyService {
return &apiKeyService{
Expand All @@ -54,7 +54,7 @@ func NewApiKeyService(
}

func (s *apiKeyService) Create(ctx context.Context, actor *models.Actor, req types.CreateApiKeyRequest) (*types.CreateApiKeyResponse, error) {
if err := s.authorizeCreate(actor, req); err != nil {
if err := s.authorizeCreate(ctx, actor, req); err != nil {
return nil, err
}

Expand Down Expand Up @@ -169,7 +169,7 @@ func (s *apiKeyService) Create(ctx context.Context, actor *models.Actor, req typ
}, nil
}

func (s *apiKeyService) authorizeCreate(actor *models.Actor, req types.CreateApiKeyRequest) error {
func (s *apiKeyService) authorizeCreate(ctx context.Context, actor *models.Actor, req types.CreateApiKeyRequest) error {
switch req.OwnerType {
case types.OwnerTypeUser:
if req.OwnerID != "" && req.OwnerID != actor.ID {
Expand All @@ -185,7 +185,11 @@ func (s *apiKeyService) authorizeCreate(actor *models.Actor, req types.CreateApi
if s.organizationService == nil {
return fmt.Errorf("%w: organization service is not available", coreerrors.ErrUnprocessableEntity)
}
if err := s.validatePermissionsSubset(actor.Scopes, req.Permissions); err != nil {
perms, err := s.organizationService.GetUserPermissionsInOrganization(ctx, actor.ID, req.OwnerID)
if err != nil {
return err
}
if err := s.validatePermissionsSubset(perms, req.Permissions); err != nil {
return err
}
}
Expand Down Expand Up @@ -270,7 +274,14 @@ func (s *apiKeyService) Update(ctx context.Context, actor *models.Actor, id stri
}
case types.OwnerTypeOrganization:
if len(req.Permissions) > 0 {
if err := s.validatePermissionsSubset(actor.Scopes, req.Permissions); err != nil {
if s.organizationService == nil {
return nil, fmt.Errorf("%w: organization service is not available", coreerrors.ErrUnprocessableEntity)
}
perms, err := s.organizationService.GetUserPermissionsInOrganization(ctx, actor.ID, apiKey.OwnerID)
if err != nil {
return nil, err
}
if err := s.validatePermissionsSubset(perms, req.Permissions); err != nil {
return nil, err
}
}
Expand Down
Loading