This commit is contained in:
+110
-3
@@ -438,7 +438,7 @@ func TestInvitationAccessLifecycleAndLastOwnerProtection(t *testing.T) {
|
||||
if _, err = accessService.Grant(t.Context(), access.Grant{SubjectKind: access.User, SubjectID: owner.ID, Role: "organization.owner", Scope: access.Scope{OrganizationID: organization.ID}, GrantedBy: owner.ID}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = organizationService.SetMembershipStatus(t.Context(), organization.ID, owner.ID, "suspended", owner.ID, "request-last-owner"); !errors.Is(err, organizations.ErrLastOwner) {
|
||||
if err = organizationService.ChangeMembershipStatus(t.Context(), organizations.MembershipStatusChange{OrganizationID: organization.ID, UserID: owner.ID, ExpectedStatus: "active", Status: "suspended", ActorUserID: owner.ID, RequestID: "request-last-owner"}); !errors.Is(err, organizations.ErrLastOwner) {
|
||||
t.Fatalf("last-owner suspension err=%v", err)
|
||||
}
|
||||
team, err := organizationService.CreateTeam(t.Context(), organizations.CreateTeam{OrganizationID: organization.ID, Slug: "operators", Name: "Operators", ActorUserID: owner.ID})
|
||||
@@ -463,10 +463,10 @@ func TestInvitationAccessLifecycleAndLastOwnerProtection(t *testing.T) {
|
||||
if err != nil || len(teams) != 1 || teams[0].ID != team.ID {
|
||||
t.Fatalf("member teams=%+v err=%v", teams, err)
|
||||
}
|
||||
if err = organizationService.SetMembershipStatus(t.Context(), organization.ID, owner.ID, "suspended", owner.ID, "request-suspend-owner"); err != nil {
|
||||
if err = organizationService.ChangeMembershipStatus(t.Context(), organizations.MembershipStatusChange{OrganizationID: organization.ID, UserID: owner.ID, ExpectedStatus: "active", Status: "suspended", ActorUserID: owner.ID, RequestID: "request-suspend-owner"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = organizationService.RemoveMembership(t.Context(), organization.ID, member.ID, member.ID, "request-last-member"); !errors.Is(err, organizations.ErrLastOwner) {
|
||||
if err = organizationService.RemoveMembershipIfCurrent(t.Context(), organizations.MembershipRemoval{OrganizationID: organization.ID, UserID: member.ID, ExpectedStatus: "active", ActorUserID: member.ID, RequestID: "request-last-member"}); !errors.Is(err, organizations.ErrLastOwner) {
|
||||
t.Fatalf("sole active owner removal err=%v", err)
|
||||
}
|
||||
if _, err = organizationService.SetOrganizationStatus(t.Context(), organizations.SetOrganizationStatus{ID: organization.ID, Status: "archived", ActorUserID: member.ID, ExpectedRevision: organization.Revision, RequestID: "request-archive"}); err != nil {
|
||||
@@ -635,6 +635,113 @@ func TestOrganizationRoleAdministrationIsAtomicAndProtectsOwners(t *testing.T) {
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_access_bindings WHERE id=?`, replacement.ID, 0)
|
||||
}
|
||||
|
||||
func TestOptimisticMembershipLifecycleIsSerializedAndAtomic(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, 4, 9, 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: "lifecycle.owner", Email: "lifecycle-owner@example.test", DisplayName: "Lifecycle Owner", Password: "correct horse battery staple"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
member, err := authService.CreateUser(t.Context(), auth.CreateUser{Username: "lifecycle.member", Email: "lifecycle-member@example.test", DisplayName: "Lifecycle 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: "optimistic-lifecycle", Name: "Optimistic Lifecycle", OwnerUserID: owner.ID})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
policy := access.Policy{Roles: map[string]string{"owner": "Owner", "viewer": "Viewer"}, Permissions: map[string]string{"telemetry.read": "Read"}, Grants: map[string][]string{"owner": {"telemetry.read"}, "viewer": {"telemetry.read"}}}
|
||||
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)
|
||||
}
|
||||
if _, err = accessService.Grant(t.Context(), access.Grant{SubjectKind: access.User, SubjectID: owner.ID, Role: "owner", Scope: access.Scope{OrganizationID: organization.ID}, GrantedBy: owner.ID}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
team, err := organizationService.CreateTeam(t.Context(), organizations.CreateTeam{OrganizationID: organization.ID, Slug: "operators", Name: "Operators", ActorUserID: owner.ID})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, _, err := organizationService.InviteWithAccess(t.Context(), organizations.InviteWithAccess{OrganizationID: organization.ID, Email: member.Email, InvitedByUserID: owner.ID, DirectRole: "viewer", TeamIDs: []string{team.ID}, Lifetime: 24 * time.Hour})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = organizationService.AcceptInvitation(t.Context(), raw, member.ID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
start := make(chan struct{})
|
||||
results := make(chan error, 2)
|
||||
for _, requestID := range []string{"request-suspend-one", "request-suspend-two"} {
|
||||
requestID := requestID
|
||||
go func() {
|
||||
<-start
|
||||
results <- organizationService.ChangeMembershipStatus(t.Context(), organizations.MembershipStatusChange{OrganizationID: organization.ID, UserID: member.ID, ExpectedStatus: "active", Status: "suspended", ActorUserID: owner.ID, RequestID: requestID})
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
var successful, conflicted int
|
||||
for range 2 {
|
||||
switch lifecycleErr := <-results; {
|
||||
case lifecycleErr == nil:
|
||||
successful++
|
||||
case errors.Is(lifecycleErr, organizations.ErrRevisionConflict):
|
||||
conflicted++
|
||||
default:
|
||||
t.Fatalf("concurrent membership suspension err=%v", lifecycleErr)
|
||||
}
|
||||
}
|
||||
if successful != 1 || conflicted != 1 {
|
||||
t.Fatalf("concurrent membership suspension success=%d conflict=%d", successful, conflicted)
|
||||
}
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_access_audit_events WHERE action='membership.suspended' AND resource_id=?`, member.ID, 1)
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_team_members WHERE user_id=?`, member.ID, 0)
|
||||
decision, err := accessService.Authorize(t.Context(), member.ID, access.Scope{OrganizationID: organization.ID}, "telemetry.read")
|
||||
if err != nil || decision.Allowed {
|
||||
t.Fatalf("suspended member decision=%+v err=%v", decision, err)
|
||||
}
|
||||
if err = organizationService.RemoveMembershipIfCurrent(t.Context(), organizations.MembershipRemoval{OrganizationID: organization.ID, UserID: member.ID, ExpectedStatus: "active", ActorUserID: owner.ID, RequestID: "request-stale-remove"}); !errors.Is(err, organizations.ErrRevisionConflict) {
|
||||
t.Fatalf("stale membership removal err=%v", err)
|
||||
}
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_access_audit_events WHERE request_id=?`, "request-stale-remove", 0)
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_organization_memberships WHERE user_id=?`, member.ID, 1)
|
||||
|
||||
if err = organizationService.ChangeMembershipStatus(t.Context(), organizations.MembershipStatusChange{OrganizationID: organization.ID, UserID: member.ID, ExpectedStatus: "suspended", Status: "active", ActorUserID: owner.ID, RequestID: "request-reactivate"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_team_members WHERE user_id=?`, member.ID, 0)
|
||||
if err = organizationService.ChangeMembershipStatus(t.Context(), organizations.MembershipStatusChange{OrganizationID: organization.ID, UserID: member.ID, ExpectedStatus: "suspended", Status: "active", ActorUserID: owner.ID, RequestID: "request-stale-reactivate"}); !errors.Is(err, organizations.ErrRevisionConflict) {
|
||||
t.Fatalf("stale membership reactivation err=%v", err)
|
||||
}
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_access_audit_events WHERE request_id=?`, "request-stale-reactivate", 0)
|
||||
|
||||
if err = organizationService.RemoveMembershipIfCurrent(t.Context(), organizations.MembershipRemoval{OrganizationID: organization.ID, UserID: member.ID, ExpectedStatus: "active", ActorUserID: owner.ID, RequestID: "request-remove-member"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_organization_memberships WHERE user_id=?`, member.ID, 0)
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_access_bindings WHERE subject_id=? AND revoked_at IS NOT NULL`, member.ID, 1)
|
||||
assertCount(t, store, `SELECT COUNT(*) FROM gwf_access_audit_events WHERE request_id=?`, "request-remove-member", 1)
|
||||
decision, err = accessService.Authorize(t.Context(), member.ID, access.Scope{OrganizationID: organization.ID}, "telemetry.read")
|
||||
if err != nil || decision.Allowed {
|
||||
t.Fatalf("removed member decision=%+v err=%v", decision, err)
|
||||
}
|
||||
}
|
||||
|
||||
func membershipPresent(values []organizations.Membership, userID, status string) bool {
|
||||
for _, value := range values {
|
||||
if value.UserID == userID && value.Status == status {
|
||||
|
||||
Reference in New Issue
Block a user