This commit is contained in:
@@ -477,3 +477,169 @@ func TestInvitationAccessLifecycleAndLastOwnerProtection(t *testing.T) {
|
||||
t.Fatalf("archived organization decision=%+v err=%v", decision, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOrganizationRoleAdministrationIsAtomicAndProtectsOwners(t *testing.T) {
|
||||
store, err := Open(filepath.Join(t.TempDir(), "accounts.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer store.Close()
|
||||
now := time.Date(2026, 9, 3, 16, 0, 0, 0, time.UTC)
|
||||
authService, err := auth.New(store, auth.Options{Now: func() time.Time { return now }})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
owner, err := authService.CreateUser(t.Context(), auth.CreateUser{Username: "access.owner", Email: "access-owner@example.test", DisplayName: "Access Owner", Password: "correct horse battery staple"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
member, err := authService.CreateUser(t.Context(), auth.CreateUser{Username: "access.member", Email: "access-member@example.test", DisplayName: "Access Member", Password: "correct horse battery staple"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
organizationService, err := organizations.New(store, organizations.Options{Now: func() time.Time { return now }, OwnerRole: "owner"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
organization, err := organizationService.CreateOrganization(t.Context(), organizations.CreateOrganization{Slug: "access-admin", Name: "Access Admin", OwnerUserID: owner.ID})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, _, err := organizationService.Invite(t.Context(), organization.ID, member.Email, owner.ID, time.Hour)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = organizationService.AcceptInvitation(t.Context(), raw, member.ID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
policy := access.Policy{
|
||||
Roles: map[string]string{"owner": "Owner", "viewer": "Viewer"},
|
||||
Permissions: map[string]string{"site.view": "View site"},
|
||||
Grants: map[string][]string{"owner": {"site.view"}, "viewer": {"site.view"}},
|
||||
}
|
||||
accessService, err := access.New(store, policy, access.Options{Now: func() time.Time { return now }, OwnerRole: "owner"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = accessService.Seed(t.Context()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ownerBinding, err := accessService.Grant(t.Context(), access.Grant{SubjectKind: access.User, SubjectID: owner.ID, Role: "owner", Scope: access.Scope{OrganizationID: organization.ID}, GrantedBy: owner.ID})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
memberBinding, err := accessService.Grant(t.Context(), access.Grant{SubjectKind: access.User, SubjectID: member.ID, Role: "viewer", Scope: access.Scope{OrganizationID: organization.ID}, GrantedBy: owner.ID})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
project, err := organizationService.CreateProject(t.Context(), organizations.CreateProject{OrganizationID: organization.ID, Slug: "narrow", Name: "Narrow"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = accessService.Grant(t.Context(), access.Grant{SubjectKind: access.User, SubjectID: member.ID, Role: "viewer", Scope: access.Scope{OrganizationID: organization.ID, ProjectID: project.ID}, GrantedBy: owner.ID}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
members, err := organizationService.Members(t.Context(), organization.ID, 10)
|
||||
if err != nil || len(members) != 2 || !membershipPresent(members, owner.ID, "active") || !membershipPresent(members, member.ID, "active") {
|
||||
t.Fatalf("members=%+v err=%v", members, err)
|
||||
}
|
||||
direct, err := accessService.OrganizationUserBindings(t.Context(), organization.ID, 10)
|
||||
if err != nil || len(direct) != 2 {
|
||||
t.Fatalf("direct=%+v err=%v", direct, err)
|
||||
}
|
||||
|
||||
type replacementResult struct {
|
||||
binding access.Binding
|
||||
err error
|
||||
}
|
||||
start := make(chan struct{})
|
||||
results := make(chan replacementResult, 2)
|
||||
for _, requestID := range []string{"request-member-owner-one", "request-member-owner-two"} {
|
||||
requestID := requestID
|
||||
go func() {
|
||||
<-start
|
||||
binding, replaceErr := accessService.ReplaceOrganizationUserRole(t.Context(), access.OrganizationUserRoleChange{OrganizationID: organization.ID, UserID: member.ID, Role: "owner", ActorUserID: owner.ID, RequestID: requestID, ExpectedBindingIDs: []string{memberBinding.ID}})
|
||||
results <- replacementResult{binding: binding, err: replaceErr}
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
var memberOwner access.Binding
|
||||
var successful, conflicted int
|
||||
for range 2 {
|
||||
result := <-results
|
||||
switch {
|
||||
case result.err == nil:
|
||||
successful++
|
||||
memberOwner = result.binding
|
||||
case errors.Is(result.err, access.ErrRoleChangeConflict):
|
||||
conflicted++
|
||||
default:
|
||||
t.Fatalf("concurrent replacement err=%v", result.err)
|
||||
}
|
||||
}
|
||||
if successful != 1 || conflicted != 1 {
|
||||
t.Fatalf("concurrent replacements success=%d conflict=%d", successful, conflicted)
|
||||
}
|
||||
if _, err = accessService.ReplaceOrganizationUserRole(t.Context(), access.OrganizationUserRoleChange{OrganizationID: organization.ID, UserID: member.ID, Role: "viewer", ActorUserID: owner.ID, RequestID: "request-stale", ExpectedBindingIDs: []string{memberBinding.ID}}); !errors.Is(err, access.ErrRoleChangeConflict) {
|
||||
t.Fatalf("stale replacement err=%v", err)
|
||||
}
|
||||
if _, err = accessService.ReplaceOrganizationUserRole(t.Context(), access.OrganizationUserRoleChange{OrganizationID: organization.ID, UserID: member.ID, Role: "owner", ActorUserID: owner.ID, RequestID: "request-unchanged", ExpectedBindingIDs: []string{memberOwner.ID}}); !errors.Is(err, access.ErrRoleUnchanged) {
|
||||
t.Fatalf("unchanged replacement err=%v", err)
|
||||
}
|
||||
ownerViewer, err := accessService.ReplaceOrganizationUserRole(t.Context(), access.OrganizationUserRoleChange{OrganizationID: organization.ID, UserID: owner.ID, Role: "viewer", ActorUserID: member.ID, RequestID: "request-owner-viewer", ExpectedBindingIDs: []string{ownerBinding.ID}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = accessService.ReplaceOrganizationUserRole(t.Context(), access.OrganizationUserRoleChange{OrganizationID: organization.ID, UserID: member.ID, Role: "viewer", ActorUserID: member.ID, RequestID: "request-last-owner", ExpectedBindingIDs: []string{memberOwner.ID}}); !errors.Is(err, access.ErrLastOwner) {
|
||||
t.Fatalf("last-owner demotion err=%v", err)
|
||||
}
|
||||
if _, err = accessService.ReplaceOrganizationUserRole(t.Context(), access.OrganizationUserRoleChange{OrganizationID: organization.ID, UserID: owner.ID, Role: "owner", ActorUserID: member.ID, RequestID: "request-restore-owner", ExpectedBindingIDs: []string{ownerViewer.ID}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = organizationService.SetMembershipStatus(t.Context(), organization.ID, member.ID, "suspended", owner.ID, "request-suspend"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
members, err = organizationService.Members(t.Context(), organization.ID, 10)
|
||||
if err != nil || len(members) != 2 || !membershipPresent(members, member.ID, "suspended") {
|
||||
t.Fatalf("suspended members=%+v err=%v", members, err)
|
||||
}
|
||||
if _, err = accessService.ReplaceOrganizationUserRole(t.Context(), access.OrganizationUserRoleChange{OrganizationID: organization.ID, UserID: member.ID, Role: "viewer", ActorUserID: owner.ID, RequestID: "request-suspended", ExpectedBindingIDs: []string{memberOwner.ID}}); err == nil {
|
||||
t.Fatal("suspended member role was replaced")
|
||||
}
|
||||
if err = organizationService.SetMembershipStatus(t.Context(), organization.ID, member.ID, "active", owner.ID, "request-reactivate"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
duplicateAudit := access.AuditEvent{ID: "audit-duplicate-1234", OrganizationID: organization.ID, ActorUserID: owner.ID, Action: "access.role.replace", ResourceType: "user", ResourceID: member.ID, RequestID: "request-rollback", Summary: "Direct organization role replaced", CreatedAt: now}
|
||||
if err = store.AppendAccessAudit(t.Context(), duplicateAudit); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
replacement := access.Binding{ID: "binding-rollback-1234", SubjectKind: access.User, SubjectID: member.ID, Role: "viewer", Scope: access.Scope{OrganizationID: organization.ID}, GrantedBy: owner.ID, GrantedAt: now}
|
||||
if err = store.ReplaceOrganizationUserRole(t.Context(), []string{memberOwner.ID}, replacement, "owner", duplicateAudit); err == nil {
|
||||
t.Fatal("audit failure did not roll back role replacement")
|
||||
}
|
||||
direct, err = store.OrganizationUserBindings(t.Context(), organization.ID, 10)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var memberRoles []string
|
||||
for _, binding := range direct {
|
||||
if binding.SubjectID == member.ID {
|
||||
memberRoles = append(memberRoles, binding.ID+":"+binding.Role)
|
||||
}
|
||||
}
|
||||
if len(memberRoles) != 1 || memberRoles[0] != memberOwner.ID+":owner" {
|
||||
t.Fatalf("rollback member roles=%v", memberRoles)
|
||||
}
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_access_bindings WHERE id=?`, replacement.ID, 0)
|
||||
}
|
||||
|
||||
func membershipPresent(values []organizations.Membership, userID, status string) bool {
|
||||
for _, value := range values {
|
||||
if value.UserID == userID && value.Status == status {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user