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
26 changes: 16 additions & 10 deletions admin/database/database.go
Original file line number Diff line number Diff line change
Expand Up @@ -155,6 +155,7 @@ type DB interface {
FindUsergroupsForUser(ctx context.Context, userID, orgID string) ([]*Usergroup, error)
FindUsergroupMemberUsers(ctx context.Context, groupID, afterEmail string, limit int) ([]*UsergroupMemberUser, error)
InsertUsergroupMemberUser(ctx context.Context, groupID, userID string) error
InsertUsergroupsMemberUser(ctx context.Context, userID string, groupIDs []string) error
DeleteUsergroupMemberUser(ctx context.Context, groupID, userID string) error
DeleteUsergroupsMemberUser(ctx context.Context, orgID, userID string) error
InsertManagedUsergroupsMemberUser(ctx context.Context, orgID, userID, roleID string) error
Expand Down Expand Up @@ -1050,11 +1051,14 @@ type ProjectMemberUser struct {
type UsergroupMemberUser struct {
ID string
Email string
DisplayName string `db:"display_name"`
PhotoURL string `db:"photo_url"`
RoleName string `db:"name"`
CreatedOn time.Time `db:"created_on"`
UpdatedOn time.Time `db:"updated_on"`
DisplayName string `db:"display_name"`
PhotoURL string `db:"photo_url"`
// PendingAcceptance is true for users who have been invited to the group but have not signed up yet.
// For pending members, ID, DisplayName and PhotoURL are empty.
PendingAcceptance bool `db:"pending_acceptance"`
RoleName string `db:"name"`
CreatedOn time.Time `db:"created_on"`
UpdatedOn time.Time `db:"updated_on"`
}

// MemberUsergroup is a convenience type used for display-friendly representation of an org or project member that is a usergroup.
Expand Down Expand Up @@ -1087,6 +1091,7 @@ type OrganizationInviteWithRole struct {
ID string
Email string
RoleName string `db:"role_name"`
Usergroups []string `db:"usergroups"` // Names of the user groups the user will be added to on acceptance
Attributes map[string]any `db:"attributes"`
InvitedBy *string `db:"invited_by"`
}
Expand Down Expand Up @@ -1164,11 +1169,12 @@ type ProjectWhitelistedDomainWithJoinedRoleNames struct {
}

type InsertOrganizationInviteOptions struct {
Email string `validate:"email"`
InviterID string
OrgID string `validate:"required"`
RoleID string `validate:"required"`
Attributes map[string]any
Email string `validate:"email"`
InviterID string
OrgID string `validate:"required"`
RoleID string `validate:"required"`
UsergroupIDs []string
Attributes map[string]any
}

type InsertProjectInviteOptions struct {
Expand Down
82 changes: 74 additions & 8 deletions admin/database/postgres/postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -1086,8 +1086,26 @@ func (c *connection) UpdateUsergroupDescription(ctx context.Context, description
}

func (c *connection) DeleteUsergroup(ctx context.Context, groupID string) error {
// Pending org invites reference usergroups by ID without a foreign key, so scrub the group from them first.
// The scrub and the delete must land together, otherwise a failed delete would silently drop the invitees' group assignment.
ctx, tx, err := c.NewTx(ctx, true)
if err != nil {
return err
}
defer func() { _ = tx.Rollback() }()

_, err = c.getDB(ctx).ExecContext(ctx, "UPDATE org_invites SET usergroup_ids = array_remove(usergroup_ids, $1::text) WHERE $1::text = ANY(usergroup_ids)", groupID)
if err != nil {
return parseErr("org invites", err)
}

res, err := c.getDB(ctx).ExecContext(ctx, "DELETE FROM usergroups WHERE id=$1", groupID)
return checkDeleteRow("usergroup", res, err)
err = checkDeleteRow("usergroup", res, err)
if err != nil {
return err
}

return tx.Commit()
}

func (c *connection) FindUsergroupsForUser(ctx context.Context, userID, orgID string) ([]*database.Usergroup, error) {
Expand All @@ -1102,13 +1120,25 @@ func (c *connection) FindUsergroupsForUser(ctx context.Context, userID, orgID st
return res, nil
}

// FindUsergroupMemberUsers returns the group's members, including users with a pending org invite that will add them to the group on acceptance.
// Both are keyed by email, so pagination by afterEmail works across the union.
// An email cannot appear in both sets since org invites are deleted when the user signs up.
func (c *connection) FindUsergroupMemberUsers(ctx context.Context, groupID, afterEmail string, limit int) ([]*database.UsergroupMemberUser, error) {
var res []*database.UsergroupMemberUser
err := c.getDB(ctx).SelectContext(ctx, &res, `
SELECT uug.user_id as "id", u.email, u.display_name, u.photo_url FROM usergroups_users uug
JOIN users u ON uug.user_id = u.id
WHERE uug.usergroup_id = $1 AND lower(u.email) > lower($2)
ORDER BY lower(u.email) LIMIT $3
SELECT m.id, m.email, m.display_name, m.photo_url, m.pending_acceptance FROM (
SELECT uug.user_id::text AS id, u.email, u.display_name, u.photo_url, false AS pending_acceptance
FROM usergroups_users uug
JOIN users u ON uug.user_id = u.id
WHERE uug.usergroup_id = $1
UNION ALL
SELECT '' AS id, oi.email, '' AS display_name, '' AS photo_url, true AS pending_acceptance
FROM org_invites oi
JOIN usergroups ug ON ug.org_id = oi.org_id
WHERE ug.id = $1 AND ug.id::text = ANY(oi.usergroup_ids)
) m
WHERE lower(m.email) > lower($2)
ORDER BY lower(m.email) LIMIT $3
`, groupID, afterEmail, limit)
if err != nil {
return nil, parseErr("usergroup member", err)
Expand All @@ -1124,6 +1154,23 @@ func (c *connection) InsertUsergroupMemberUser(ctx context.Context, groupID, use
return nil
}

// InsertUsergroupsMemberUser adds the user to each of the groups, skipping groups the user is already a member of.
// It is safe to call inside a transaction since duplicates do not raise a unique violation.
func (c *connection) InsertUsergroupsMemberUser(ctx context.Context, userID string, groupIDs []string) error {
if len(groupIDs) == 0 {
return nil
}
_, err := c.getDB(ctx).ExecContext(ctx, `
INSERT INTO usergroups_users (user_id, usergroup_id)
SELECT $1::uuid, unnest($2::text[])::uuid
ON CONFLICT DO NOTHING
`, userID, groupIDs)
if err != nil {
return parseErr("usergroup member", err)
}
return nil
}

func (c *connection) DeleteUsergroupMemberUser(ctx context.Context, groupID, userID string) error {
res, err := c.getDB(ctx).ExecContext(ctx, "DELETE FROM usergroups_users WHERE user_id = $1 AND usergroup_id = $2", userID, groupID)
return checkDeleteRow("usergroup member", res, err)
Expand Down Expand Up @@ -2383,9 +2430,12 @@ func (c *connection) FindOrganizationMemberUsergroups(ctx context.Context, orgID
var qry strings.Builder
qry.WriteString("SELECT ug.id, ug.name, ug.managed, ug.created_on, ug.updated_on, COALESCE(r.name, '') as role_name")
if withCounts {
// Counts pending invitees as well, to match the members listed by FindUsergroupMemberUsers.
qry.WriteString(`,
(
SELECT COUNT(*) FROM usergroups_users uug WHERE uug.usergroup_id = ug.id
) + (
SELECT COUNT(*) FROM org_invites oi WHERE oi.org_id = ug.org_id AND ug.id::text = ANY(oi.usergroup_ids)
) as users_count
`)
}
Expand Down Expand Up @@ -2449,9 +2499,12 @@ func (c *connection) FindProjectMemberUsergroups(ctx context.Context, projectID,
var qry strings.Builder
qry.WriteString(`SELECT ug.id, ug.name, ug.managed, ug.created_on, ug.updated_on, r.name as "role_name", upr.resources, upr.restrict_resources`)
if withCounts {
// Counts pending invitees as well, to match the members listed by FindUsergroupMemberUsers.
qry.WriteString(`,
(
SELECT COUNT(*) FROM usergroups_users uug WHERE uug.usergroup_id = ug.id
) + (
SELECT COUNT(*) FROM org_invites oi WHERE oi.org_id = ug.org_id AND ug.id::text = ANY(oi.usergroup_ids)
) as users_count
`)
}
Expand Down Expand Up @@ -2604,7 +2657,8 @@ func (c *connection) DeleteProjectMemberService(ctx context.Context, serviceID,
func (c *connection) FindOrganizationInvites(ctx context.Context, orgID, afterEmail string, limit int) ([]*database.OrganizationInviteWithRole, error) {
var dtos []*organizationInviteWithRoleDTO
err := c.getDB(ctx).SelectContext(ctx, &dtos, `
SELECT uoi.id, uoi.email, ur.name as role_name, uoi.attributes, u.email as invited_by
SELECT uoi.id, uoi.email, ur.name as role_name, uoi.attributes, u.email as invited_by,
(SELECT COALESCE(array_agg(ug.name ORDER BY lower(ug.name)), '{}') FROM usergroups ug WHERE ug.id::text = ANY(uoi.usergroup_ids)) as usergroups
FROM org_invites uoi
JOIN org_roles ur ON uoi.org_role_id = ur.id
LEFT JOIN users u ON uoi.invited_by_user_id = u.id
Expand Down Expand Up @@ -2675,7 +2729,13 @@ func (c *connection) InsertOrganizationInvite(ctx context.Context, opts *databas
inviterID = opts.InviterID
}

_, err = c.getDB(ctx).ExecContext(ctx, "INSERT INTO org_invites (email, invited_by_user_id, org_id, org_role_id, attributes) VALUES ($1, $2, $3, $4, $5)", opts.Email, inviterID, opts.OrgID, opts.RoleID, attrs)
// usergroup_ids is NOT NULL, so a nil slice must be inserted as an empty array
groupIDs := opts.UsergroupIDs
if groupIDs == nil {
groupIDs = []string{}
}

_, err = c.getDB(ctx).ExecContext(ctx, "INSERT INTO org_invites (email, invited_by_user_id, org_id, org_role_id, attributes, usergroup_ids) VALUES ($1, $2, $3, $4, $5, $6)", opts.Email, inviterID, opts.OrgID, opts.RoleID, attrs, groupIDs)
if err != nil {
return parseErr("org invite", err)
}
Expand Down Expand Up @@ -3760,10 +3820,16 @@ func (o *organizationInviteDTO) AsModel() (*database.OrganizationInvite, error)

type organizationInviteWithRoleDTO struct {
*database.OrganizationInviteWithRole
Attributes pgtype.JSON `db:"attributes"`
Usergroups pgtype.TextArray `db:"usergroups"`
Attributes pgtype.JSON `db:"attributes"`
}

func (o *organizationInviteWithRoleDTO) AsModel() (*database.OrganizationInviteWithRole, error) {
err := o.Usergroups.AssignTo(&o.OrganizationInviteWithRole.Usergroups)
if err != nil {
return nil, err
}

// Handle Attributes: Normalize NULL JSONB to empty map
var attrs map[string]any
if err := o.Attributes.AssignTo(&attrs); err != nil {
Expand Down
102 changes: 102 additions & 0 deletions admin/database/postgres/postgres_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@ func TestPostgres(t *testing.T) {
t.Run("TestManagedGitRepos", func(t *testing.T) { testManagedGitRepos(t, db) })
t.Run("TestOrganizationMemberUserAttributes", func(t *testing.T) { testOrganizationMemberUserAttributes(t, db) })
t.Run("TestOrganizationInviteAttributes", func(t *testing.T) { testOrganizationInviteAttributes(t, db) })
t.Run("TestOrganizationInviteUsergroups", func(t *testing.T) { testOrganizationInviteUsergroups(t, db) })
t.Run("TestAttributeValidation", func(t *testing.T) { testAttributeValidation(t, db) })

t.Run("TestOrgNameValidation", func(t *testing.T) {
Expand Down Expand Up @@ -825,6 +826,107 @@ func testOrganizationInviteAttributes(t *testing.T, db database.DB) {
require.NoError(t, db.DeleteOrganization(ctx, org.Name))
}

func testOrganizationInviteUsergroups(t *testing.T, db database.DB) {
ctx := context.Background()

org, err := db.InsertOrganization(ctx, &database.InsertOrganizationOptions{Name: "test-invite-groups-org"})
require.NoError(t, err)

role, err := db.FindOrganizationRole(ctx, database.OrganizationRoleNameViewer)
require.NoError(t, err)

g1, err := db.InsertUsergroup(ctx, &database.InsertUsergroupOptions{OrgID: org.ID, Name: "g1"})
require.NoError(t, err)
g2, err := db.InsertUsergroup(ctx, &database.InsertUsergroupOptions{OrgID: org.ID, Name: "g2"})
require.NoError(t, err)

email := "invitee-groups@rilldata.com"

t.Run("InsertOrganizationInvite with usergroups", func(t *testing.T) {
err := db.InsertOrganizationInvite(ctx, &database.InsertOrganizationInviteOptions{
Email: email,
OrgID: org.ID,
RoleID: role.ID,
UsergroupIDs: []string{g2.ID, g1.ID},
})
require.NoError(t, err)

invite, err := db.FindOrganizationInvite(ctx, org.ID, email)
require.NoError(t, err)
require.Equal(t, []string{g2.ID, g1.ID}, invite.UsergroupIDs)

// The display-friendly listing resolves names, sorted by name
invitesWithRole, err := db.FindOrganizationInvites(ctx, org.ID, "", 10)
require.NoError(t, err)
require.Len(t, invitesWithRole, 1)
require.Equal(t, []string{"g1", "g2"}, invitesWithRole[0].Usergroups)
})

t.Run("InsertOrganizationInvite without usergroups normalizes to empty slice", func(t *testing.T) {
email2 := "invitee-no-groups@rilldata.com"
err := db.InsertOrganizationInvite(ctx, &database.InsertOrganizationInviteOptions{
Email: email2,
OrgID: org.ID,
RoleID: role.ID,
})
require.NoError(t, err)

invite, err := db.FindOrganizationInvite(ctx, org.ID, email2)
require.NoError(t, err)
require.Empty(t, invite.UsergroupIDs)

invitesWithRole, err := db.FindOrganizationInvites(ctx, org.ID, "", 10)
require.NoError(t, err)
require.Len(t, invitesWithRole, 2)
for _, inv := range invitesWithRole {
if inv.Email == email2 {
require.Empty(t, inv.Usergroups)
}
}
require.NoError(t, db.DeleteOrganizationInvite(ctx, invite.ID))
})

t.Run("FindUsergroupMemberUsers includes pending invitees", func(t *testing.T) {
user, err := db.InsertUser(ctx, &database.InsertUserOptions{Email: "member-groups@rilldata.com"})
require.NoError(t, err)
_, err = db.InsertOrganizationMemberUser(ctx, org.ID, user.ID, role.ID, nil, false)
require.NoError(t, err)

// Batch insert is idempotent
require.NoError(t, db.InsertUsergroupsMemberUser(ctx, user.ID, []string{g1.ID, g2.ID}))
require.NoError(t, db.InsertUsergroupsMemberUser(ctx, user.ID, []string{g1.ID}))
require.NoError(t, db.InsertUsergroupsMemberUser(ctx, user.ID, nil))

members, err := db.FindUsergroupMemberUsers(ctx, g1.ID, "", 10)
require.NoError(t, err)
require.Len(t, members, 2)
// Ordered by email: invitee-groups@ sorts before member-groups@
require.Equal(t, email, members[0].Email)
require.True(t, members[0].PendingAcceptance)
require.Empty(t, members[0].ID)
require.Equal(t, user.Email, members[1].Email)
require.False(t, members[1].PendingAcceptance)
require.Equal(t, user.ID, members[1].ID)

// Pagination by email spans both real and pending members
page, err := db.FindUsergroupMemberUsers(ctx, g1.ID, email, 10)
require.NoError(t, err)
require.Len(t, page, 1)
require.Equal(t, user.Email, page[0].Email)
})

t.Run("DeleteUsergroup scrubs pending invites", func(t *testing.T) {
require.NoError(t, db.DeleteUsergroup(ctx, g2.ID))

invite, err := db.FindOrganizationInvite(ctx, org.ID, email)
require.NoError(t, err)
require.Equal(t, []string{g1.ID}, invite.UsergroupIDs)
})

// Cleanup
require.NoError(t, db.DeleteOrganization(ctx, org.Name))
}

func testAttributeValidation(t *testing.T, db database.DB) {
ctx := context.Background()

Expand Down
Loading
Loading