diff --git a/core/pagination/pagination.go b/core/pagination/pagination.go new file mode 100644 index 00000000..5f268b5a --- /dev/null +++ b/core/pagination/pagination.go @@ -0,0 +1,65 @@ +package pagination + +import ( + "net/http" + + "github.com/Authula/authula/util" +) + +const ( + DefaultPage = 1 + DefaultLimit = 10 + MaxLimit = 100 +) + +type Params struct { + Page int + Limit int +} + +type Pagination struct { + Page int `json:"page" required:"true" nullable:"false"` + Limit int `json:"limit" required:"true" nullable:"false"` + Total int `json:"total" required:"true" nullable:"false"` + TotalPages int `json:"total_pages" required:"true" nullable:"false"` + HasMore bool `json:"has_more" required:"true" nullable:"false"` +} + +func Clamp(params Params) Params { + if params.Page < DefaultPage { + params.Page = DefaultPage + } + if params.Limit <= 0 { + params.Limit = DefaultLimit + } + if params.Limit > MaxLimit { + params.Limit = MaxLimit + } + return params +} + +func New(page, limit, total int) Pagination { + if limit <= 0 { + limit = DefaultLimit + } + + totalPages := total / limit + if total%limit != 0 { + totalPages++ + } + + return Pagination{ + Page: page, + Limit: limit, + Total: total, + TotalPages: totalPages, + HasMore: page < totalPages, + } +} + +func ParseFromRequest(r *http.Request) Params { + return Params{ + Page: util.GetQueryInt(r, "page", DefaultPage), + Limit: util.GetQueryInt(r, "limit", DefaultLimit), + } +} diff --git a/core/pagination/pagination_test.go b/core/pagination/pagination_test.go new file mode 100644 index 00000000..4690b8a1 --- /dev/null +++ b/core/pagination/pagination_test.go @@ -0,0 +1,153 @@ +package pagination_test + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/Authula/authula/core/pagination" +) + +func TestClamp(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + params pagination.Params + expected pagination.Params + }{ + { + name: "valid params are untouched", + params: pagination.Params{Page: 1, Limit: 10}, + expected: pagination.Params{Page: 1, Limit: 10}, + }, + { + name: "page zero becomes the first page", + params: pagination.Params{Page: 0, Limit: 10}, + expected: pagination.Params{Page: 1, Limit: 10}, + }, + { + name: "negative page becomes the first page", + params: pagination.Params{Page: -5, Limit: 10}, + expected: pagination.Params{Page: 1, Limit: 10}, + }, + { + name: "zero limit falls back to the default limit", + params: pagination.Params{Page: 3, Limit: 0}, + expected: pagination.Params{Page: 3, Limit: pagination.DefaultLimit}, + }, + { + name: "negative limit falls back to the default limit", + params: pagination.Params{Page: 3, Limit: -1}, + expected: pagination.Params{Page: 3, Limit: pagination.DefaultLimit}, + }, + { + name: "limit exactly at the maximum is not clamped", + params: pagination.Params{Page: 1, Limit: pagination.MaxLimit}, + expected: pagination.Params{Page: 1, Limit: pagination.MaxLimit}, + }, + { + name: "limit above the maximum is clamped", + params: pagination.Params{Page: 1, Limit: pagination.MaxLimit + 1}, + expected: pagination.Params{Page: 1, Limit: pagination.MaxLimit}, + }, + { + name: "absurdly large limit is clamped", + params: pagination.Params{Page: 1, Limit: 1_000_000}, + expected: pagination.Params{Page: 1, Limit: pagination.MaxLimit}, + }, + { + name: "high page numbers are legal", + params: pagination.Params{Page: 999999, Limit: 10}, + expected: pagination.Params{Page: 999999, Limit: 10}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + require.Equal(t, tt.expected, pagination.Clamp(tt.params)) + }) + } +} + +func TestNew(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + page int + limit int + total int + expectedTotalPages int + expectedHasMore bool + }{ + {name: "empty result set", page: 1, limit: 10, total: 0, expectedTotalPages: 0, expectedHasMore: false}, + {name: "exactly one full page", page: 1, limit: 10, total: 10, expectedTotalPages: 1, expectedHasMore: false}, + {name: "one row spilling onto a second page", page: 1, limit: 10, total: 11, expectedTotalPages: 2, expectedHasMore: true}, + {name: "middle page has more", page: 2, limit: 25, total: 137, expectedTotalPages: 6, expectedHasMore: true}, + {name: "last page has no more", page: 6, limit: 25, total: 137, expectedTotalPages: 6, expectedHasMore: false}, + {name: "page past the end has no more", page: 99, limit: 10, total: 5, expectedTotalPages: 1, expectedHasMore: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + result := pagination.New(tt.page, tt.limit, tt.total) + require.Equal(t, tt.page, result.Page) + require.Equal(t, tt.limit, result.Limit) + require.Equal(t, tt.total, result.Total) + require.Equal(t, tt.expectedTotalPages, result.TotalPages) + require.Equal(t, tt.expectedHasMore, result.HasMore) + }) + } +} + +func TestParseFromRequest(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + target string + expected pagination.Params + }{ + { + name: "absent query parameters use the defaults", + target: "/organizations", + expected: pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}, + }, + { + name: "explicit values are parsed", + target: "/organizations?page=3&limit=50", + expected: pagination.Params{Page: 3, Limit: 50}, + }, + { + name: "unparseable page falls back to the default page", + target: "/organizations?page=abc", + expected: pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}, + }, + { + name: "unparseable limit falls back to the default limit", + target: "/organizations?page=2&limit=xyz", + expected: pagination.Params{Page: 2, Limit: pagination.DefaultLimit}, + }, + { + name: "out of range values are returned unclamped", + target: "/organizations?page=-4&limit=5000", + expected: pagination.Params{Page: -4, Limit: 5000}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + request := httptest.NewRequest(http.MethodGet, tt.target, nil) + require.Equal(t, tt.expected, pagination.ParseFromRequest(request)) + }) + } +} diff --git a/internal/tests/data_helpers.go b/internal/tests/data_helpers.go index c1ef52e2..c6f97779 100644 --- a/internal/tests/data_helpers.go +++ b/internal/tests/data_helpers.go @@ -115,6 +115,28 @@ func AssertErrorMessage(t *testing.T, reqCtx *models.RequestContext, status int, } } +func AssertErrorResponse(t *testing.T, reqCtx *models.RequestContext, status int, code string, message string) { + t.Helper() + + if !reqCtx.Handled { + t.Fatal("expected request to be marked as handled") + } + if reqCtx.ResponseStatus != status { + t.Fatalf("expected status %d, got %d", status, reqCtx.ResponseStatus) + } + + payload := DecodeResponseJSON[struct { + Code string `json:"code"` + Message string `json:"message"` + }](t, reqCtx) + if payload.Code != code { + t.Fatalf("expected code %q, got %q", code, payload.Code) + } + if payload.Message != message { + t.Fatalf("expected message %q, got %q", message, payload.Message) + } +} + func DecodeResponseJSON[T any](t *testing.T, reqCtx *models.RequestContext) T { t.Helper() diff --git a/openapi.json b/openapi.json index 31ef9de3..453f973f 100644 --- a/openapi.json +++ b/openapi.json @@ -2404,21 +2404,31 @@ "Organizations" ], "summary": "List organizations", - "description": "Lists all organizations owned by the authenticated user.", + "description": "Lists every organization the authenticated user can access, both the ones they own and the ones they are a member of, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected.", "operationId": "listOrganizations", + "parameters": [ + { + "name": "page", + "in": "query", + "schema": { + "type": "integer" + } + }, + { + "name": "limit", + "in": "query", + "schema": { + "type": "integer" + } + } + ], "responses": { "200": { "description": "OK", "content": { "application/json": { "schema": { - "items": { - "$ref": "#/components/schemas/Organization" - }, - "type": [ - "null", - "array" - ] + "$ref": "#/components/schemas/ListOrganizationsResponse" } } } @@ -2562,9 +2572,23 @@ "Organization Invitations" ], "summary": "List invitations", - "description": "Lists all invitations for an organization.", + "description": "Lists the invitations for an organization, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected.", "operationId": "listOrganizationInvitations", "parameters": [ + { + "name": "page", + "in": "query", + "schema": { + "type": "integer" + } + }, + { + "name": "limit", + "in": "query", + "schema": { + "type": "integer" + } + }, { "name": "organization_id", "in": "path", @@ -2580,13 +2604,7 @@ "content": { "application/json": { "schema": { - "items": { - "$ref": "#/components/schemas/GetOrganizationInvitationResponse" - }, - "type": [ - "null", - "array" - ] + "$ref": "#/components/schemas/ListOrganizationInvitationsResponse" } } } @@ -2822,7 +2840,7 @@ "Organization Members" ], "summary": "List members", - "description": "Lists all members of an organization with pagination.", + "description": "Lists the members of an organization, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected.", "operationId": "listOrganizationMembers", "parameters": [ { @@ -2854,13 +2872,7 @@ "content": { "application/json": { "schema": { - "items": { - "$ref": "#/components/schemas/OrganizationMemberResponse" - }, - "type": [ - "null", - "array" - ] + "$ref": "#/components/schemas/ListOrganizationMembersResponse" } } } @@ -3078,9 +3090,23 @@ "Organization Teams" ], "summary": "List teams", - "description": "Lists all teams within an organization.", + "description": "Lists the teams within an organization, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected.", "operationId": "listOrganizationTeams", "parameters": [ + { + "name": "page", + "in": "query", + "schema": { + "type": "integer" + } + }, + { + "name": "limit", + "in": "query", + "schema": { + "type": "integer" + } + }, { "name": "organization_id", "in": "path", @@ -3096,13 +3122,7 @@ "content": { "application/json": { "schema": { - "items": { - "$ref": "#/components/schemas/OrganizationTeam" - }, - "type": [ - "null", - "array" - ] + "$ref": "#/components/schemas/ListOrganizationTeamsResponse" } } } @@ -3280,7 +3300,7 @@ "Organization Team Members" ], "summary": "List team members", - "description": "Lists all members of a team with pagination.", + "description": "Lists the members of a team, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected.", "operationId": "listOrganizationTeamMembers", "parameters": [ { @@ -3320,13 +3340,7 @@ "content": { "application/json": { "schema": { - "items": { - "$ref": "#/components/schemas/OrganizationTeamMemberResponse" - }, - "type": [ - "null", - "array" - ] + "$ref": "#/components/schemas/ListOrganizationTeamMembersResponse" } } } @@ -5030,6 +5044,96 @@ ], "type": "object" }, + "ListOrganizationInvitationsResponse": { + "properties": { + "data": { + "items": { + "$ref": "#/components/schemas/GetOrganizationInvitationResponse" + }, + "type": "array" + }, + "pagination": { + "$ref": "#/components/schemas/Pagination" + } + }, + "required": [ + "data", + "pagination" + ], + "type": "object" + }, + "ListOrganizationMembersResponse": { + "properties": { + "data": { + "items": { + "$ref": "#/components/schemas/OrganizationMemberResponse" + }, + "type": "array" + }, + "pagination": { + "$ref": "#/components/schemas/Pagination" + } + }, + "required": [ + "data", + "pagination" + ], + "type": "object" + }, + "ListOrganizationTeamMembersResponse": { + "properties": { + "data": { + "items": { + "$ref": "#/components/schemas/OrganizationTeamMemberResponse" + }, + "type": "array" + }, + "pagination": { + "$ref": "#/components/schemas/Pagination" + } + }, + "required": [ + "data", + "pagination" + ], + "type": "object" + }, + "ListOrganizationTeamsResponse": { + "properties": { + "data": { + "items": { + "$ref": "#/components/schemas/OrganizationTeam" + }, + "type": "array" + }, + "pagination": { + "$ref": "#/components/schemas/Pagination" + } + }, + "required": [ + "data", + "pagination" + ], + "type": "object" + }, + "ListOrganizationsResponse": { + "properties": { + "data": { + "items": { + "$ref": "#/components/schemas/Organization" + }, + "type": "array" + }, + "pagination": { + "$ref": "#/components/schemas/Pagination" + } + }, + "required": [ + "data", + "pagination" + ], + "type": "object" + }, "MagicLinkExchangeRequest": { "properties": { "token": { @@ -5398,6 +5502,33 @@ ], "type": "object" }, + "Pagination": { + "properties": { + "has_more": { + "type": "boolean" + }, + "limit": { + "type": "integer" + }, + "page": { + "type": "integer" + }, + "total": { + "type": "integer" + }, + "total_pages": { + "type": "integer" + } + }, + "required": [ + "page", + "limit", + "total", + "total_pages", + "has_more" + ], + "type": "object" + }, "Permission": { "properties": { "created_at": { diff --git a/plugins/organizations/api.go b/plugins/organizations/api.go index 17ac45b9..8913a8ca 100644 --- a/plugins/organizations/api.go +++ b/plugins/organizations/api.go @@ -4,6 +4,7 @@ import ( "context" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" "github.com/Authula/authula/plugins/organizations/repositories" "github.com/Authula/authula/plugins/organizations/services" @@ -59,8 +60,8 @@ func (a *API) CreateOrganization(ctx context.Context, actor *models.Actor, reque return a.organizationService.CreateOrganization(ctx, actor, request) } -func (a *API) GetAllOrganizationsByOwner(ctx context.Context, actor *models.Actor) ([]types.Organization, error) { - return a.organizationService.GetAllOrganizationsByOwner(ctx, actor) +func (a *API) GetAllOrganizations(ctx context.Context, actor *models.Actor, params pagination.Params) (*types.ListOrganizationsResponse, error) { + return a.organizationService.GetAllOrganizations(ctx, actor, params) } func (a *API) GetOrganizationByID(ctx context.Context, actor *models.Actor, organizationID string) (*types.Organization, error) { @@ -81,8 +82,8 @@ func (a *API) CreateInvitation(ctx context.Context, actor *models.Actor, organiz return a.invitationService.CreateOrganizationInvitation(ctx, actor, organizationID, request, redirectURL) } -func (a *API) GetAllInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { - return a.invitationService.GetAllOrganizationInvitationsByOrgIDWithOrg(ctx, organizationID) +func (a *API) GetAllInvitations(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationInvitationsResponse, error) { + return a.invitationService.GetAllOrganizationInvitationsByOrgIDWithOrg(ctx, organizationID, params) } func (a *API) GetInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.GetOrganizationInvitationResponse, error) { @@ -115,8 +116,8 @@ func (a *API) AddMember(ctx context.Context, actor *models.Actor, organizationID return a.memberService.AddMember(ctx, actor, organizationID, request) } -func (a *API) GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, error) { - return a.memberService.GetAllMembers(ctx, actor, organizationID, page, limit) +func (a *API) GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationMembersResponse, error) { + return a.memberService.GetAllMembers(ctx, actor, organizationID, params) } func (a *API) GetMember(ctx context.Context, actor *models.Actor, organizationID string, memberID string) (*types.OrganizationMemberResponse, error) { @@ -141,8 +142,8 @@ func (a *API) CreateTeam(ctx context.Context, actor *models.Actor, organizationI return a.teamService.CreateTeam(ctx, actor, organizationID, request) } -func (a *API) GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationTeam, error) { - return a.teamService.GetAllTeams(ctx, actor, organizationID) +func (a *API) GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationTeamsResponse, error) { + return a.teamService.GetAllTeams(ctx, actor, organizationID, params) } func (a *API) GetTeam(ctx context.Context, actor *models.Actor, organizationID string, teamID string) (*types.OrganizationTeam, error) { @@ -163,8 +164,8 @@ func (a *API) AddTeamMember(ctx context.Context, actor *models.Actor, organizati return a.teamMemberService.AddTeamMember(ctx, actor, organizationID, teamID, request) } -func (a *API) GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, error) { - return a.teamMemberService.GetAllTeamMembers(ctx, actor, organizationID, teamID, page, limit) +func (a *API) GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, params pagination.Params) (*types.ListOrganizationTeamMembersResponse, error) { + return a.teamMemberService.GetAllTeamMembers(ctx, actor, organizationID, teamID, params) } func (a *API) GetTeamMember(ctx context.Context, actor *models.Actor, organizationID string, teamID string, memberID string) (*types.OrganizationTeamMemberResponse, error) { diff --git a/plugins/organizations/constants/errors.go b/plugins/organizations/constants/errors.go index 1290e07d..bb6506fd 100644 --- a/plugins/organizations/constants/errors.go +++ b/plugins/organizations/constants/errors.go @@ -8,25 +8,50 @@ import ( "github.com/Authula/authula/models" ) -var ( - ErrOrganizationsQuotaExceeded = errors.New("organizations quota exceeded") - ErrMembersQuotaExceeded = errors.New("members quota exceeded") - ErrInvitationsQuotaExceeded = errors.New("invitations quota exceeded") - ErrInvitationEmailMismatch = errors.New("this invitation was sent to a different email address") -) +type Error struct { + Code string + Status int + Message string +} -func HandleError(err error, reqCtx *models.RequestContext) { - var status int +func (e *Error) Error() string { return e.Message } - switch err { - case ErrOrganizationsQuotaExceeded, ErrMembersQuotaExceeded, ErrInvitationsQuotaExceeded: - status = http.StatusTooManyRequests - case ErrInvitationEmailMismatch: - status = http.StatusForbidden +var ( + ErrOrganizationsQuotaExceeded = &Error{ + Code: CodeOrganizationsQuotaExceeded, + Status: http.StatusConflict, + Message: "organizations quota exceeded", + } + ErrMembersQuotaExceeded = &Error{ + Code: CodeMembersQuotaExceeded, + Status: http.StatusConflict, + Message: "members quota exceeded", } + ErrInvitationsQuotaExceeded = &Error{ + Code: CodeInvitationsQuotaExceeded, + Status: http.StatusConflict, + Message: "invitations quota exceeded", + } + ErrInvitationEmailMismatch = &Error{ + Code: CodeInvitationEmailMismatch, + Status: http.StatusForbidden, + Message: "this invitation was sent to a different email address", + } +) - if status != 0 { - reqCtx.SetJSONResponse(status, map[string]any{"message": err.Error()}) +const ( + CodeOrganizationsQuotaExceeded = "organizations_quota_exceeded" + CodeMembersQuotaExceeded = "members_quota_exceeded" + CodeInvitationsQuotaExceeded = "invitations_quota_exceeded" + CodeInvitationEmailMismatch = "invitation_email_mismatch" +) + +func HandleError(err error, reqCtx *models.RequestContext) { + if pluginErr, ok := errors.AsType[*Error](err); ok { + reqCtx.SetJSONResponse(pluginErr.Status, map[string]any{ + "code": pluginErr.Code, + "message": pluginErr.Message, + }) reqCtx.Handled = true return } diff --git a/plugins/organizations/constants/errors_test.go b/plugins/organizations/constants/errors_test.go new file mode 100644 index 00000000..784b42cd --- /dev/null +++ b/plugins/organizations/constants/errors_test.go @@ -0,0 +1,179 @@ +package constants_test + +import ( + "errors" + "fmt" + "net/http" + "testing" + + "github.com/stretchr/testify/require" + + coreerrors "github.com/Authula/authula/core/errors" + internaltests "github.com/Authula/authula/internal/tests" + "github.com/Authula/authula/plugins/organizations/constants" +) + +func TestHandleError(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + expectedStatus int + expectedCode string + expectedMessage string + }{ + { + name: "organizations quota is a conflict, not a rate limit", + err: constants.ErrOrganizationsQuotaExceeded, + expectedStatus: http.StatusConflict, + expectedCode: constants.CodeOrganizationsQuotaExceeded, + expectedMessage: "organizations quota exceeded", + }, + { + name: "members quota is a conflict, not a rate limit", + err: constants.ErrMembersQuotaExceeded, + expectedStatus: http.StatusConflict, + expectedCode: constants.CodeMembersQuotaExceeded, + expectedMessage: "members quota exceeded", + }, + { + name: "invitations quota is a conflict, not a rate limit", + err: constants.ErrInvitationsQuotaExceeded, + expectedStatus: http.StatusConflict, + expectedCode: constants.CodeInvitationsQuotaExceeded, + expectedMessage: "invitations quota exceeded", + }, + { + name: "invitation email mismatch stays forbidden", + err: constants.ErrInvitationEmailMismatch, + expectedStatus: http.StatusForbidden, + expectedCode: constants.CodeInvitationEmailMismatch, + expectedMessage: "this invitation was sent to a different email address", + }, + { + // The regression test for dispatching on errors.AsType rather than + // on error identity: a wrapped sentinel used to fall through to the + // core handler and silently degrade to 400. + name: "a wrapped quota error keeps its status and code", + err: fmt.Errorf("creating organization: %w", constants.ErrMembersQuotaExceeded), + expectedStatus: http.StatusConflict, + expectedCode: constants.CodeMembersQuotaExceeded, + expectedMessage: "members quota exceeded", + }, + { + name: "doubly wrapped quota error keeps its status and code", + err: fmt.Errorf("outer: %w", fmt.Errorf("inner: %w", constants.ErrOrganizationsQuotaExceeded)), + expectedStatus: http.StatusConflict, + expectedCode: constants.CodeOrganizationsQuotaExceeded, + expectedMessage: "organizations quota exceeded", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + reqCtx := internaltests.NewRequestContext(t, http.MethodPost, "/organizations", nil) + + constants.HandleError(tt.err, reqCtx) + + internaltests.AssertErrorResponse(t, reqCtx, tt.expectedStatus, tt.expectedCode, tt.expectedMessage) + }) + } +} + +// Errors the plugin does not own must still reach the core handler unchanged. +func TestHandleErrorDelegatesToCore(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + expectedStatus int + expectedMessage string + }{ + {name: "not found", err: coreerrors.ErrNotFound, expectedStatus: http.StatusNotFound, expectedMessage: "not found"}, + {name: "forbidden", err: coreerrors.ErrForbidden, expectedStatus: http.StatusForbidden, expectedMessage: "forbidden"}, + {name: "conflict", err: coreerrors.ErrConflict, expectedStatus: http.StatusConflict, expectedMessage: "conflict"}, + {name: "unauthorized", err: coreerrors.ErrUnauthorized, expectedStatus: http.StatusUnauthorized, expectedMessage: "unauthorized"}, + {name: "unknown error falls through", err: errors.New("some repository failure"), expectedStatus: http.StatusBadRequest, expectedMessage: "some repository failure"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + reqCtx := internaltests.NewRequestContext(t, http.MethodPost, "/organizations", nil) + + constants.HandleError(tt.err, reqCtx) + + internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + }) + } +} + +// Typing the sentinels must not change how callers compare them. Every service +// returns the bare sentinel and the service tests assert on the value, so +// identity and errors.Is have to keep holding. +func TestSentinelsRemainComparable(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err *constants.Error + expectedMessage string + }{ + {name: "organizations quota", err: constants.ErrOrganizationsQuotaExceeded, expectedMessage: "organizations quota exceeded"}, + {name: "members quota", err: constants.ErrMembersQuotaExceeded, expectedMessage: "members quota exceeded"}, + {name: "invitations quota", err: constants.ErrInvitationsQuotaExceeded, expectedMessage: "invitations quota exceeded"}, + {name: "invitation email mismatch", err: constants.ErrInvitationEmailMismatch, expectedMessage: "this invitation was sent to a different email address"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + var asError error = tt.err + + require.Equal(t, tt.expectedMessage, tt.err.Error(), "message must stay byte-identical") + require.True(t, asError == tt.err, "identity comparison must still hold") //nolint:testifylint // asserting == on purpose + require.ErrorIs(t, asError, tt.err) + require.ErrorIs(t, fmt.Errorf("wrapped: %w", asError), tt.err) + }) + } +} + +// Error codes are a public contract: once shipped, a code is never renamed. +// This golden list makes any addition or rename an explicit, reviewed diff. +func TestErrorCodeRegistry(t *testing.T) { + t.Parallel() + + require.Equal(t, map[string]string{ + "organizations_quota_exceeded": "organizations quota exceeded", + "members_quota_exceeded": "members quota exceeded", + "invitations_quota_exceeded": "invitations quota exceeded", + "invitation_email_mismatch": "this invitation was sent to a different email address", + }, map[string]string{ + constants.ErrOrganizationsQuotaExceeded.Code: constants.ErrOrganizationsQuotaExceeded.Message, + constants.ErrMembersQuotaExceeded.Code: constants.ErrMembersQuotaExceeded.Message, + constants.ErrInvitationsQuotaExceeded.Code: constants.ErrInvitationsQuotaExceeded.Message, + constants.ErrInvitationEmailMismatch.Code: constants.ErrInvitationEmailMismatch.Message, + }) +} + +// No organizations error may use 429: that status means rate limiting in this +// API and is reserved for the rate-limit and api-key plugins, which pair it +// with Retry-After and X-RateLimit-* headers. +func TestNoQuotaErrorUsesTooManyRequests(t *testing.T) { + t.Parallel() + + for _, err := range []*constants.Error{ + constants.ErrOrganizationsQuotaExceeded, + constants.ErrMembersQuotaExceeded, + constants.ErrInvitationsQuotaExceeded, + constants.ErrInvitationEmailMismatch, + } { + require.NotEqual(t, http.StatusTooManyRequests, err.Status, "%s must not use 429", err.Code) + } +} diff --git a/plugins/organizations/handlers/organization_handlers.go b/plugins/organizations/handlers/organization_handlers.go index 2bdea9a3..8d542172 100644 --- a/plugins/organizations/handlers/organization_handlers.go +++ b/plugins/organizations/handlers/organization_handlers.go @@ -3,6 +3,7 @@ package handlers import ( "net/http" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" "github.com/Authula/authula/plugins/organizations/types" @@ -48,16 +49,17 @@ func (h *CreateOrganizationHandler) Handle() http.HandlerFunc { } } -type GetAllOrganizationsByOwnerHandler struct { +type GetAllOrganizationsHandler struct { UseCases *orgusecases.UseCases } -func (h *GetAllOrganizationsByOwnerHandler) Handle() http.HandlerFunc { +func (h *GetAllOrganizationsHandler) Handle() http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() reqCtx, _ := models.GetRequestContext(ctx) actor := reqCtx.Actor - organizations, err := h.UseCases.GetAllOrganizationsByOwner(ctx, actor) + paginationParams := pagination.ParseFromRequest(r) + organizations, err := h.UseCases.GetAllOrganizations(ctx, actor, paginationParams) if err != nil { orgconstants.HandleError(err, reqCtx) return diff --git a/plugins/organizations/handlers/organization_handlers_test.go b/plugins/organizations/handlers/organization_handlers_test.go index dfaadbb4..a78d9ae2 100644 --- a/plugins/organizations/handlers/organization_handlers_test.go +++ b/plugins/organizations/handlers/organization_handlers_test.go @@ -11,6 +11,7 @@ import ( "github.com/stretchr/testify/require" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" @@ -29,6 +30,7 @@ type organizationHandlerCase struct { organizationID string prepare func(*organizationHandlerFixture) expectedStatus int + expectedCode string expectedMessage string checkResponse func(t *testing.T, reqCtx *models.RequestContext) } @@ -97,7 +99,8 @@ func TestCreateOrganizationHandler(t *testing.T) { return request.Name == "Acme Inc" && request.Role == "member" })).Return((*orgtypes.Organization)(nil), orgconstants.ErrOrganizationsQuotaExceeded).Once() }, - expectedStatus: http.StatusTooManyRequests, + expectedStatus: http.StatusConflict, + expectedCode: orgconstants.CodeOrganizationsQuotaExceeded, expectedMessage: "organizations quota exceeded", }, { @@ -163,7 +166,11 @@ func TestCreateOrganizationHandler(t *testing.T) { assert.Equal(t, tt.expectedStatus, reqCtx.ResponseStatus) if tt.expectedMessage != "" { - internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + if tt.expectedCode != "" { + internaltests.AssertErrorResponse(t, reqCtx, tt.expectedStatus, tt.expectedCode, tt.expectedMessage) + } else { + internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + } } if tt.checkResponse != nil { tt.checkResponse(t, reqCtx) @@ -173,7 +180,7 @@ func TestCreateOrganizationHandler(t *testing.T) { } } -func TestGetAllOrganizationsByOwnerHandler(t *testing.T) { +func TestGetAllOrganizationsHandler(t *testing.T) { t.Parallel() tests := []organizationHandlerCase{ @@ -186,7 +193,8 @@ func TestGetAllOrganizationsByOwnerHandler(t *testing.T) { name: "service_error", userID: new("user-1"), prepare: func(f *organizationHandlerFixture) { - f.service.On("GetAllOrganizationsByOwner", mock.Anything, "user-1").Return(([]orgtypes.Organization)(nil), errors.New("some error")).Once() + f.service.On("GetAllOrganizations", mock.Anything, "user-1", pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}). + Return((*orgtypes.ListOrganizationsResponse)(nil), errors.New("some error")).Once() }, expectedStatus: http.StatusBadRequest, expectedMessage: "some error", @@ -195,15 +203,20 @@ func TestGetAllOrganizationsByOwnerHandler(t *testing.T) { name: "success", userID: new("user-1"), prepare: func(f *organizationHandlerFixture) { - f.service.On("GetAllOrganizationsByOwner", mock.Anything, "user-1").Return([]orgtypes.Organization{{ID: "org-1", OwnerID: "user-1", Name: "Acme Inc", Slug: "acme-inc"}}, nil).Once() + f.service.On("GetAllOrganizations", mock.Anything, "user-1", pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}). + Return(&orgtypes.ListOrganizationsResponse{ + Data: []orgtypes.Organization{{ID: "org-1", OwnerID: "user-1", Name: "Acme Inc", Slug: "acme-inc"}}, + Pagination: pagination.New(1, 10, 1), + }, nil).Once() }, expectedStatus: http.StatusOK, checkResponse: func(t *testing.T, reqCtx *models.RequestContext) { - organizations := internaltests.DecodeResponseJSON[[]orgtypes.Organization](t, reqCtx) - require.Len(t, organizations, 1) - assert.Equal(t, "org-1", organizations[0].ID) - assert.Equal(t, "user-1", organizations[0].OwnerID) - assert.Equal(t, "Acme Inc", organizations[0].Name) + resp := internaltests.DecodeResponseJSON[orgtypes.ListOrganizationsResponse](t, reqCtx) + require.Len(t, resp.Data, 1) + assert.Equal(t, "org-1", resp.Data[0].ID) + assert.Equal(t, "user-1", resp.Data[0].OwnerID) + assert.Equal(t, "Acme Inc", resp.Data[0].Name) + assert.Equal(t, pagination.Pagination{Page: 1, Limit: 10, Total: 1, TotalPages: 1, HasMore: false}, resp.Pagination) }, }, } @@ -217,7 +230,7 @@ func TestGetAllOrganizationsByOwnerHandler(t *testing.T) { tt.prepare(fixture) } - handler := &GetAllOrganizationsByOwnerHandler{UseCases: newOrgUseCases(fixture.service)} + handler := &GetAllOrganizationsHandler{UseCases: newOrgUseCases(fixture.service)} req, w, reqCtx := fixture.newRequest(t, http.MethodGet, "/organizations", nil, tt.userID, "") if tt.name == "missing_user" { reqCtx.SetJSONResponse(http.StatusUnauthorized, map[string]any{"message": "Unauthorized"}) @@ -228,7 +241,11 @@ func TestGetAllOrganizationsByOwnerHandler(t *testing.T) { assert.Equal(t, tt.expectedStatus, reqCtx.ResponseStatus) if tt.expectedMessage != "" { - internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + if tt.expectedCode != "" { + internaltests.AssertErrorResponse(t, reqCtx, tt.expectedStatus, tt.expectedCode, tt.expectedMessage) + } else { + internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + } } if tt.checkResponse != nil { tt.checkResponse(t, reqCtx) @@ -322,7 +339,11 @@ func TestGetOrganizationByIDHandler(t *testing.T) { assert.Equal(t, tt.expectedStatus, reqCtx.ResponseStatus) if tt.expectedMessage != "" { - internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + if tt.expectedCode != "" { + internaltests.AssertErrorResponse(t, reqCtx, tt.expectedStatus, tt.expectedCode, tt.expectedMessage) + } else { + internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + } } if tt.checkResponse != nil { tt.checkResponse(t, reqCtx) @@ -471,7 +492,11 @@ func TestUpdateOrganizationHandler(t *testing.T) { assert.Equal(t, tt.expectedStatus, reqCtx.ResponseStatus) if tt.expectedMessage != "" { - internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + if tt.expectedCode != "" { + internaltests.AssertErrorResponse(t, reqCtx, tt.expectedStatus, tt.expectedCode, tt.expectedMessage) + } else { + internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + } } if tt.checkResponse != nil { tt.checkResponse(t, reqCtx) @@ -546,7 +571,11 @@ func TestDeleteOrganizationHandler(t *testing.T) { assert.Equal(t, tt.expectedStatus, reqCtx.ResponseStatus) if tt.expectedMessage != "" { - internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + if tt.expectedCode != "" { + internaltests.AssertErrorResponse(t, reqCtx, tt.expectedStatus, tt.expectedCode, tt.expectedMessage) + } else { + internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + } } if tt.checkResponse != nil { tt.checkResponse(t, reqCtx) diff --git a/plugins/organizations/handlers/organization_invitation_handlers.go b/plugins/organizations/handlers/organization_invitation_handlers.go index a8184fe6..951b2998 100644 --- a/plugins/organizations/handlers/organization_invitation_handlers.go +++ b/plugins/organizations/handlers/organization_invitation_handlers.go @@ -3,6 +3,7 @@ package handlers import ( "net/http" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" "github.com/Authula/authula/plugins/organizations/types" @@ -56,7 +57,8 @@ func (h *GetAllOrganizationInvitationsHandler) Handle() http.HandlerFunc { actor := reqCtx.Actor organizationID := r.PathValue("organization_id") - invitations, err := h.UseCases.GetAllOrganizationInvitations(ctx, actor, organizationID) + paginationParams := pagination.ParseFromRequest(r) + invitations, err := h.UseCases.GetAllOrganizationInvitations(ctx, actor, organizationID, paginationParams) if err != nil { orgconstants.HandleError(err, reqCtx) return diff --git a/plugins/organizations/handlers/organization_invitation_handlers_test.go b/plugins/organizations/handlers/organization_invitation_handlers_test.go index 2d098d16..74932215 100644 --- a/plugins/organizations/handlers/organization_invitation_handlers_test.go +++ b/plugins/organizations/handlers/organization_invitation_handlers_test.go @@ -11,6 +11,7 @@ import ( "github.com/stretchr/testify/require" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" @@ -31,6 +32,7 @@ type organizationInvitationHandlerCase struct { invitationID string prepare func(*organizationInvitationHandlerFixture) expectedStatus int + expectedCode string expectedMessage string checkResponse func(*testing.T, *models.RequestContext) } @@ -82,7 +84,11 @@ func runOrganizationInvitationHandlerCases(t *testing.T, method, path string, bu assert.Equal(t, tt.expectedStatus, reqCtx.ResponseStatus) if tt.expectedMessage != "" { - internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + if tt.expectedCode != "" { + internaltests.AssertErrorResponse(t, reqCtx, tt.expectedStatus, tt.expectedCode, tt.expectedMessage) + } else { + internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + } } if tt.checkResponse != nil { tt.checkResponse(t, reqCtx) @@ -133,7 +139,8 @@ func TestCreateOrganizationInvitationHandler(t *testing.T) { prepare: func(fixture *organizationInvitationHandlerFixture) { fixture.service.On("CreateOrganizationInvitation", mock.Anything, "user-1", "org-1", mock.Anything, mock.Anything).Return((*orgtypes.OrganizationInvitation)(nil), orgconstants.ErrInvitationsQuotaExceeded).Once() }, - expectedStatus: http.StatusTooManyRequests, + expectedStatus: http.StatusConflict, + expectedCode: orgconstants.CodeInvitationsQuotaExceeded, expectedMessage: orgconstants.ErrInvitationsQuotaExceeded.Error(), }, { @@ -173,7 +180,8 @@ func TestGetAllOrganizationInvitationsHandler(t *testing.T) { userID: new("user-1"), organizationID: "org-1", prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("GetAllOrganizationInvitationsByOrgIDWithOrg", mock.Anything, "org-1").Return(([]orgtypes.GetOrganizationInvitationResponse)(nil), errors.New("some error")).Once() + fixture.service.On("GetAllOrganizationInvitationsByOrgIDWithOrg", mock.Anything, "org-1", pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}). + Return((*orgtypes.ListOrganizationInvitationsResponse)(nil), errors.New("some error")).Once() }, expectedStatus: http.StatusBadRequest, expectedMessage: "some error", @@ -183,20 +191,25 @@ func TestGetAllOrganizationInvitationsHandler(t *testing.T) { userID: new("user-1"), organizationID: "org-1", prepare: func(fixture *organizationInvitationHandlerFixture) { - fixture.service.On("GetAllOrganizationInvitationsByOrgIDWithOrg", mock.Anything, "org-1").Return([]orgtypes.GetOrganizationInvitationResponse{ - { - Invitation: &orgtypes.OrganizationInvitation{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}, - Organization: orgtypes.OrganizationSummary{ID: "org-1", Name: "Acme Corp", Slug: "acme"}, - }, - }, nil).Once() + fixture.service.On("GetAllOrganizationInvitationsByOrgIDWithOrg", mock.Anything, "org-1", pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}). + Return(&orgtypes.ListOrganizationInvitationsResponse{ + Data: []orgtypes.GetOrganizationInvitationResponse{ + { + Invitation: &orgtypes.OrganizationInvitation{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: orgtypes.OrganizationInvitationStatusPending}, + Organization: orgtypes.OrganizationSummary{ID: "org-1", Name: "Acme Corp", Slug: "acme"}, + }, + }, + Pagination: pagination.New(1, 10, 1), + }, nil).Once() }, expectedStatus: http.StatusOK, checkResponse: func(t *testing.T, reqCtx *models.RequestContext) { - resp := internaltests.DecodeResponseJSON[[]orgtypes.GetOrganizationInvitationResponse](t, reqCtx) - require.Len(t, resp, 1) - assert.Equal(t, "inv-1", resp[0].Invitation.ID) - assert.Equal(t, "org-1", resp[0].Invitation.OrganizationID) - assert.Equal(t, "Acme Corp", resp[0].Organization.Name) + resp := internaltests.DecodeResponseJSON[orgtypes.ListOrganizationInvitationsResponse](t, reqCtx) + require.Len(t, resp.Data, 1) + assert.Equal(t, "inv-1", resp.Data[0].Invitation.ID) + assert.Equal(t, "org-1", resp.Data[0].Invitation.OrganizationID) + assert.Equal(t, "Acme Corp", resp.Data[0].Organization.Name) + assert.Equal(t, pagination.Pagination{Page: 1, Limit: 10, Total: 1, TotalPages: 1, HasMore: false}, resp.Pagination) }, }, }) diff --git a/plugins/organizations/handlers/organization_member_handlers.go b/plugins/organizations/handlers/organization_member_handlers.go index 8ba7c4c6..c2173e35 100644 --- a/plugins/organizations/handlers/organization_member_handlers.go +++ b/plugins/organizations/handlers/organization_member_handlers.go @@ -3,6 +3,7 @@ package handlers import ( "net/http" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" "github.com/Authula/authula/plugins/organizations/types" @@ -61,9 +62,8 @@ func (h *GetAllOrganizationMembersHandler) Handle() http.HandlerFunc { actor := reqCtx.Actor organizationID := r.PathValue("organization_id") - page := util.GetQueryInt(r, "page", 1) - limit := util.GetQueryInt(r, "limit", 10) - members, err := h.UseCases.GetAllMembers(ctx, actor, organizationID, page, limit) + paginationParams := pagination.ParseFromRequest(r) + members, err := h.UseCases.GetAllMembers(ctx, actor, organizationID, paginationParams) if err != nil { orgconstants.HandleError(err, reqCtx) return diff --git a/plugins/organizations/handlers/organization_member_handlers_test.go b/plugins/organizations/handlers/organization_member_handlers_test.go index 67e86478..782a138b 100644 --- a/plugins/organizations/handlers/organization_member_handlers_test.go +++ b/plugins/organizations/handlers/organization_member_handlers_test.go @@ -10,6 +10,7 @@ import ( "github.com/stretchr/testify/mock" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" @@ -30,6 +31,7 @@ type organizationMemberHandlerCase struct { targetUserID string prepare func(*organizationMemberHandlerFixture) expectedStatus int + expectedCode string expectedMessage string checkResponse func(*testing.T, *models.RequestContext) } @@ -81,7 +83,11 @@ func runOrganizationMemberHandlerCases(t *testing.T, method, path string, buildH assert.Equal(t, tt.expectedStatus, reqCtx.ResponseStatus) if tt.expectedMessage != "" { - internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + if tt.expectedCode != "" { + internaltests.AssertErrorResponse(t, reqCtx, tt.expectedStatus, tt.expectedCode, tt.expectedMessage) + } else { + internaltests.AssertErrorMessage(t, reqCtx, tt.expectedStatus, tt.expectedMessage) + } } if tt.checkResponse != nil { tt.checkResponse(t, reqCtx) @@ -131,7 +137,8 @@ func TestAddOrganizationMemberHandler(t *testing.T) { prepare: func(fixture *organizationMemberHandlerFixture) { fixture.service.On("AddMember", mock.Anything, "user-1", "org-1", mock.Anything).Return((*orgtypes.OrganizationMember)(nil), orgconstants.ErrMembersQuotaExceeded).Once() }, - expectedStatus: http.StatusTooManyRequests, + expectedStatus: http.StatusConflict, + expectedCode: orgconstants.CodeMembersQuotaExceeded, expectedMessage: "members quota exceeded", }, { @@ -156,6 +163,8 @@ func TestAddOrganizationMemberHandler(t *testing.T) { func TestGetAllOrganizationMembersHandler(t *testing.T) { t.Parallel() + defaultParams := pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit} + runOrganizationMemberHandlerCases(t, http.MethodGet, "/organizations/org-1/members", func(fixture *organizationMemberHandlerFixture) http.HandlerFunc { return (&GetAllOrganizationMembersHandler{UseCases: newMemberUseCases(fixture.service)}).Handle() }, []organizationMemberHandlerCase{ @@ -170,7 +179,8 @@ func TestGetAllOrganizationMembersHandler(t *testing.T) { userID: new("user-1"), organizationID: "org-1", prepare: func(fixture *organizationMemberHandlerFixture) { - fixture.service.On("GetAllMembers", mock.Anything, "user-1", "org-1", 1, 10).Return(([]orgtypes.OrganizationMemberResponse)(nil), errors.New("some error")).Once() + fixture.service.On("GetAllMembers", mock.Anything, "user-1", "org-1", defaultParams). + Return((*orgtypes.ListOrganizationMembersResponse)(nil), errors.New("some error")).Once() }, expectedStatus: http.StatusBadRequest, expectedMessage: "some error", @@ -180,18 +190,63 @@ func TestGetAllOrganizationMembersHandler(t *testing.T) { userID: new("user-1"), organizationID: "org-1", prepare: func(fixture *organizationMemberHandlerFixture) { - fixture.service.On("GetAllMembers", mock.Anything, "user-1", "org-1", 1, 10).Return([]orgtypes.OrganizationMemberResponse{{ID: "mem-1", OrganizationID: "org-1", Role: "member"}}, nil).Once() + fixture.service.On("GetAllMembers", mock.Anything, "user-1", "org-1", defaultParams). + Return(&orgtypes.ListOrganizationMembersResponse{ + Data: []orgtypes.OrganizationMemberResponse{{ID: "mem-1", OrganizationID: "org-1", Role: "member"}}, + Pagination: pagination.New(1, 10, 25), + }, nil).Once() }, expectedStatus: http.StatusOK, checkResponse: func(t *testing.T, reqCtx *models.RequestContext) { - members := internaltests.DecodeResponseJSON[[]orgtypes.OrganizationMemberResponse](t, reqCtx) - assert.Len(t, members, 1) - assert.Equal(t, "mem-1", members[0].ID) + resp := internaltests.DecodeResponseJSON[orgtypes.ListOrganizationMembersResponse](t, reqCtx) + assert.Len(t, resp.Data, 1) + assert.Equal(t, "mem-1", resp.Data[0].ID) + assert.Equal(t, pagination.Pagination{Page: 1, Limit: 10, Total: 25, TotalPages: 3, HasMore: true}, resp.Pagination) }, }, }) } +func TestGetAllOrganizationMembersHandlerParsesPagination(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + query string + expectedParams pagination.Params + }{ + {name: "no query string uses the defaults", query: "", expectedParams: pagination.Params{Page: 1, Limit: 10}}, + {name: "explicit values are forwarded", query: "?page=3&limit=50", expectedParams: pagination.Params{Page: 3, Limit: 50}}, + {name: "unparseable values fall back to the defaults", query: "?page=abc&limit=xyz", expectedParams: pagination.Params{Page: 1, Limit: 10}}, + {name: "an absurd limit is forwarded for the service to clamp", query: "?limit=100000", expectedParams: pagination.Params{Page: 1, Limit: 100000}}, + {name: "a negative limit is forwarded for the service to clamp", query: "?limit=-1", expectedParams: pagination.Params{Page: 1, Limit: -1}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + runOrganizationMemberHandlerCases(t, http.MethodGet, "/organizations/org-1/members"+tt.query, func(fixture *organizationMemberHandlerFixture) http.HandlerFunc { + return (&GetAllOrganizationMembersHandler{UseCases: newMemberUseCases(fixture.service)}).Handle() + }, []organizationMemberHandlerCase{ + { + name: "forwards_parsed_params", + userID: new("user-1"), + organizationID: "org-1", + prepare: func(fixture *organizationMemberHandlerFixture) { + fixture.service.On("GetAllMembers", mock.Anything, "user-1", "org-1", tt.expectedParams). + Return(&orgtypes.ListOrganizationMembersResponse{ + Data: []orgtypes.OrganizationMemberResponse{}, + Pagination: pagination.New(1, 10, 0), + }, nil).Once() + }, + expectedStatus: http.StatusOK, + }, + }) + }) + } +} + func TestGetOrganizationMemberHandler(t *testing.T) { t.Parallel() diff --git a/plugins/organizations/handlers/organization_team_handlers.go b/plugins/organizations/handlers/organization_team_handlers.go index 8ca2d5da..02f5a307 100644 --- a/plugins/organizations/handlers/organization_team_handlers.go +++ b/plugins/organizations/handlers/organization_team_handlers.go @@ -3,6 +3,7 @@ package handlers import ( "net/http" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" "github.com/Authula/authula/plugins/organizations/types" @@ -55,7 +56,8 @@ func (h *GetAllOrganizationTeamsHandler) Handle() http.HandlerFunc { actor := reqCtx.Actor organizationID := r.PathValue("organization_id") - teams, err := h.UseCases.GetAllTeams(ctx, actor, organizationID) + paginationParams := pagination.ParseFromRequest(r) + teams, err := h.UseCases.GetAllTeams(ctx, actor, organizationID, paginationParams) if err != nil { orgconstants.HandleError(err, reqCtx) return diff --git a/plugins/organizations/handlers/organization_team_handlers_test.go b/plugins/organizations/handlers/organization_team_handlers_test.go index 91f09db2..9c849727 100644 --- a/plugins/organizations/handlers/organization_team_handlers_test.go +++ b/plugins/organizations/handlers/organization_team_handlers_test.go @@ -10,6 +10,7 @@ import ( "github.com/stretchr/testify/mock" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/models" orgtests "github.com/Authula/authula/plugins/organizations/tests" @@ -154,7 +155,8 @@ func TestGetAllOrganizationTeamsHandler(t *testing.T) { userID: new("user-1"), organizationID: "org-1", prepare: func(fixture *organizationTeamHandlerFixture) { - fixture.service.On("GetAllTeams", mock.Anything, "user-1", "org-1").Return(([]orgtypes.OrganizationTeam)(nil), errors.New("some error")).Once() + fixture.service.On("GetAllTeams", mock.Anything, "user-1", "org-1", pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}). + Return((*orgtypes.ListOrganizationTeamsResponse)(nil), errors.New("some error")).Once() }, expectedStatus: http.StatusBadRequest, expectedMessage: "some error", @@ -164,13 +166,18 @@ func TestGetAllOrganizationTeamsHandler(t *testing.T) { userID: new("user-1"), organizationID: "org-1", prepare: func(fixture *organizationTeamHandlerFixture) { - fixture.service.On("GetAllTeams", mock.Anything, "user-1", "org-1").Return([]orgtypes.OrganizationTeam{{ID: "team-1", OrganizationID: "org-1", Name: "Platform", Slug: "platform"}}, nil).Once() + fixture.service.On("GetAllTeams", mock.Anything, "user-1", "org-1", pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}). + Return(&orgtypes.ListOrganizationTeamsResponse{ + Data: []orgtypes.OrganizationTeam{{ID: "team-1", OrganizationID: "org-1", Name: "Platform", Slug: "platform"}}, + Pagination: pagination.New(1, 10, 1), + }, nil).Once() }, expectedStatus: http.StatusOK, checkResponse: func(t *testing.T, reqCtx *models.RequestContext) { - teams := internaltests.DecodeResponseJSON[[]orgtypes.OrganizationTeam](t, reqCtx) - assert.Len(t, teams, 1) - assert.Equal(t, "team-1", teams[0].ID) + resp := internaltests.DecodeResponseJSON[orgtypes.ListOrganizationTeamsResponse](t, reqCtx) + assert.Len(t, resp.Data, 1) + assert.Equal(t, "team-1", resp.Data[0].ID) + assert.Equal(t, pagination.Pagination{Page: 1, Limit: 10, Total: 1, TotalPages: 1, HasMore: false}, resp.Pagination) }, }, }) diff --git a/plugins/organizations/handlers/organization_team_member_handlers.go b/plugins/organizations/handlers/organization_team_member_handlers.go index de5e493a..c46d1b30 100644 --- a/plugins/organizations/handlers/organization_team_member_handlers.go +++ b/plugins/organizations/handlers/organization_team_member_handlers.go @@ -3,6 +3,7 @@ package handlers import ( "net/http" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" "github.com/Authula/authula/plugins/organizations/types" @@ -57,9 +58,8 @@ func (h *GetAllOrganizationTeamMembersHandler) Handle() http.HandlerFunc { organizationID := r.PathValue("organization_id") teamID := r.PathValue("team_id") - page := util.GetQueryInt(r, "page", 1) - limit := util.GetQueryInt(r, "limit", 10) - teamMembers, err := h.UseCases.GetAllTeamMembers(ctx, actor, organizationID, teamID, page, limit) + paginationParams := pagination.ParseFromRequest(r) + teamMembers, err := h.UseCases.GetAllTeamMembers(ctx, actor, organizationID, teamID, paginationParams) if err != nil { orgconstants.HandleError(err, reqCtx) return diff --git a/plugins/organizations/handlers/organization_team_member_handlers_test.go b/plugins/organizations/handlers/organization_team_member_handlers_test.go index 69ed8259..a6714de3 100644 --- a/plugins/organizations/handlers/organization_team_member_handlers_test.go +++ b/plugins/organizations/handlers/organization_team_member_handlers_test.go @@ -10,6 +10,7 @@ import ( "github.com/stretchr/testify/mock" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/models" orgtests "github.com/Authula/authula/plugins/organizations/tests" @@ -164,7 +165,8 @@ func TestGetAllOrganizationTeamMembersHandler(t *testing.T) { organizationID: "org-1", teamID: "team-1", prepare: func(fixture *organizationTeamMemberHandlerFixture) { - fixture.service.On("GetAllTeamMembers", mock.Anything, "user-1", "org-1", "team-1", 1, 10).Return(([]orgtypes.OrganizationTeamMemberResponse)(nil), errors.New("some error")).Once() + fixture.service.On("GetAllTeamMembers", mock.Anything, "user-1", "org-1", "team-1", pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}). + Return((*orgtypes.ListOrganizationTeamMembersResponse)(nil), errors.New("some error")).Once() }, expectedStatus: http.StatusBadRequest, expectedMessage: "some error", @@ -175,13 +177,18 @@ func TestGetAllOrganizationTeamMembersHandler(t *testing.T) { organizationID: "org-1", teamID: "team-1", prepare: func(fixture *organizationTeamMemberHandlerFixture) { - fixture.service.On("GetAllTeamMembers", mock.Anything, "user-1", "org-1", "team-1", 1, 10).Return([]orgtypes.OrganizationTeamMemberResponse{{ID: "team-mem-1", TeamID: "team-1"}}, nil).Once() + fixture.service.On("GetAllTeamMembers", mock.Anything, "user-1", "org-1", "team-1", pagination.Params{Page: pagination.DefaultPage, Limit: pagination.DefaultLimit}). + Return(&orgtypes.ListOrganizationTeamMembersResponse{ + Data: []orgtypes.OrganizationTeamMemberResponse{{ID: "team-mem-1", TeamID: "team-1"}}, + Pagination: pagination.New(1, 10, 1), + }, nil).Once() }, expectedStatus: http.StatusOK, checkResponse: func(t *testing.T, reqCtx *models.RequestContext) { - teamMembers := internaltests.DecodeResponseJSON[[]orgtypes.OrganizationTeamMemberResponse](t, reqCtx) - assert.Len(t, teamMembers, 1) - assert.Equal(t, "team-mem-1", teamMembers[0].ID) + resp := internaltests.DecodeResponseJSON[orgtypes.ListOrganizationTeamMembersResponse](t, reqCtx) + assert.Len(t, resp.Data, 1) + assert.Equal(t, "team-mem-1", resp.Data[0].ID) + assert.Equal(t, pagination.Pagination{Page: 1, Limit: 10, Total: 1, TotalPages: 1, HasMore: false}, resp.Pagination) }, }, }) diff --git a/plugins/organizations/migrationset/migrations.go b/plugins/organizations/migrationset/migrations.go index ca181793..915d267f 100644 --- a/plugins/organizations/migrationset/migrations.go +++ b/plugins/organizations/migrationset/migrations.go @@ -37,6 +37,7 @@ func organizationsSQLiteInitial() migrations.Migration { FOREIGN KEY (owner_id) REFERENCES users(id) ON DELETE CASCADE );`, `CREATE INDEX IF NOT EXISTS idx_organizations_owner_id ON organizations(owner_id);`, + `CREATE INDEX IF NOT EXISTS idx_organizations_owner_created_id ON organizations(owner_id, created_at, id);`, `DROP TRIGGER IF EXISTS update_organizations_updated_at_trigger;`, `CREATE TRIGGER update_organizations_updated_at_trigger AFTER UPDATE ON organizations @@ -62,6 +63,8 @@ func organizationsSQLiteInitial() migrations.Migration { `CREATE INDEX IF NOT EXISTS idx_organization_invitations_organization_id ON organization_invitations(organization_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_invitations_inviter_id ON organization_invitations(inviter_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_invitations_status_expires_at ON organization_invitations(status, expires_at);`, + `CREATE INDEX IF NOT EXISTS idx_organization_invitations_org_created_id ON organization_invitations(organization_id, created_at, id);`, + `CREATE INDEX IF NOT EXISTS idx_organization_invitations_email_status_expires ON organization_invitations(email, status, expires_at);`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_members ( id TEXT PRIMARY KEY, @@ -76,6 +79,8 @@ func organizationsSQLiteInitial() migrations.Migration { );`, `CREATE INDEX IF NOT EXISTS idx_organization_members_organization_id ON organization_members(organization_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_members_user_id ON organization_members(user_id);`, + `CREATE INDEX IF NOT EXISTS idx_organization_members_org_created_id ON organization_members(organization_id, created_at, id);`, + `CREATE INDEX IF NOT EXISTS idx_organization_members_user_org ON organization_members(user_id, organization_id);`, `DROP TRIGGER IF EXISTS update_organization_members_updated_at_trigger;`, `CREATE TRIGGER update_organization_members_updated_at_trigger AFTER UPDATE ON organization_members @@ -97,6 +102,7 @@ func organizationsSQLiteInitial() migrations.Migration { );`, `CREATE INDEX IF NOT EXISTS idx_organization_teams_organization_id ON organization_teams(organization_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_teams_slug ON organization_teams(slug);`, + `CREATE INDEX IF NOT EXISTS idx_organization_teams_org_created_id ON organization_teams(organization_id, created_at, id);`, `DROP TRIGGER IF EXISTS update_organization_teams_updated_at_trigger;`, `CREATE TRIGGER update_organization_teams_updated_at_trigger AFTER UPDATE ON organization_teams @@ -115,6 +121,7 @@ func organizationsSQLiteInitial() migrations.Migration { );`, `CREATE INDEX IF NOT EXISTS idx_organization_team_members_team_id ON organization_team_members(team_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_team_members_member_id ON organization_team_members(member_id);`, + `CREATE INDEX IF NOT EXISTS idx_organization_team_members_team_created_id ON organization_team_members(team_id, created_at, id);`, // ----------------------------------- ) }, @@ -164,6 +171,7 @@ func organizationsPostgresInitial() migrations.Migration { FOR EACH ROW EXECUTE FUNCTION organizations_set_updated_at_fn();`, `CREATE INDEX IF NOT EXISTS idx_organizations_owner_id ON organizations(owner_id);`, + `CREATE INDEX IF NOT EXISTS idx_organizations_owner_created_id ON organizations(owner_id, created_at, id);`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_invitations ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -188,6 +196,8 @@ func organizationsPostgresInitial() migrations.Migration { `CREATE INDEX IF NOT EXISTS idx_organization_invitations_inviter_id ON organization_invitations(inviter_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_invitations_email ON organization_invitations(email);`, `CREATE INDEX IF NOT EXISTS idx_organization_invitations_status_expires_at ON organization_invitations(status, expires_at);`, + `CREATE INDEX IF NOT EXISTS idx_organization_invitations_org_created_id ON organization_invitations(organization_id, created_at, id);`, + `CREATE INDEX IF NOT EXISTS idx_organization_invitations_email_status_expires ON organization_invitations(email, status, expires_at);`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_members ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -207,6 +217,8 @@ func organizationsPostgresInitial() migrations.Migration { EXECUTE FUNCTION organizations_set_updated_at_fn();`, `CREATE INDEX IF NOT EXISTS idx_organization_members_organization_id ON organization_members(organization_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_members_user_id ON organization_members(user_id);`, + `CREATE INDEX IF NOT EXISTS idx_organization_members_org_created_id ON organization_members(organization_id, created_at, id);`, + `CREATE INDEX IF NOT EXISTS idx_organization_members_user_org ON organization_members(user_id, organization_id);`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_teams ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -227,6 +239,7 @@ func organizationsPostgresInitial() migrations.Migration { EXECUTE FUNCTION organizations_set_updated_at_fn();`, `CREATE INDEX IF NOT EXISTS idx_organization_teams_organization_id ON organization_teams(organization_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_teams_slug ON organization_teams(slug);`, + `CREATE INDEX IF NOT EXISTS idx_organization_teams_org_created_id ON organization_teams(organization_id, created_at, id);`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_team_members ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -239,6 +252,7 @@ func organizationsPostgresInitial() migrations.Migration { );`, `CREATE INDEX IF NOT EXISTS idx_organization_team_members_team_id ON organization_team_members(team_id);`, `CREATE INDEX IF NOT EXISTS idx_organization_team_members_member_id ON organization_team_members(member_id);`, + `CREATE INDEX IF NOT EXISTS idx_organization_team_members_team_created_id ON organization_team_members(team_id, created_at, id);`, // ----------------------------------- ) }, @@ -279,7 +293,8 @@ func organizationsMySQLInitial() migrations.Migration { created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP, CONSTRAINT fk_organizations_owner FOREIGN KEY (owner_id) REFERENCES users(id) ON DELETE CASCADE, - INDEX idx_organizations_owner_id (owner_id) + INDEX idx_organizations_owner_id (owner_id), + INDEX idx_organizations_owner_created_id (owner_id, created_at, id) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_invitations ( @@ -298,7 +313,9 @@ func organizationsMySQLInitial() migrations.Migration { INDEX idx_organization_invitations_organization_id (organization_id), INDEX idx_organization_invitations_inviter_id (inviter_id), INDEX idx_organization_invitations_email (email), - INDEX idx_organization_invitations_status_expires_at (status, expires_at) + INDEX idx_organization_invitations_status_expires_at (status, expires_at), + INDEX idx_organization_invitations_org_created_id (organization_id, created_at, id), + INDEX idx_organization_invitations_email_status_expires (email, status, expires_at) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_members ( @@ -312,7 +329,9 @@ func organizationsMySQLInitial() migrations.Migration { CONSTRAINT fk_organization_members_user FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, CONSTRAINT uq_organization_members_organization_user UNIQUE (organization_id, user_id), INDEX idx_organization_members_organization_id (organization_id), - INDEX idx_organization_members_user_id (user_id) + INDEX idx_organization_members_user_id (user_id), + INDEX idx_organization_members_org_created_id (organization_id, created_at, id), + INDEX idx_organization_members_user_org (user_id, organization_id) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_teams ( @@ -327,7 +346,8 @@ func organizationsMySQLInitial() migrations.Migration { CONSTRAINT fk_organization_teams_organization FOREIGN KEY (organization_id) REFERENCES organizations(id) ON DELETE CASCADE, CONSTRAINT uq_organization_teams_organization_slug UNIQUE (organization_id, slug), INDEX idx_organization_teams_organization_id (organization_id), - INDEX idx_organization_teams_slug (slug) + INDEX idx_organization_teams_slug (slug), + INDEX idx_organization_teams_org_created_id (organization_id, created_at, id) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;`, // ----------------------------------- `CREATE TABLE IF NOT EXISTS organization_team_members ( @@ -339,7 +359,8 @@ func organizationsMySQLInitial() migrations.Migration { CONSTRAINT fk_organization_team_members_member FOREIGN KEY (member_id) REFERENCES organization_members(id) ON DELETE CASCADE, CONSTRAINT uq_organization_team_members_team_member UNIQUE (team_id, member_id), INDEX idx_organization_team_members_team_id (team_id), - INDEX idx_organization_team_members_member_id (member_id) + INDEX idx_organization_team_members_member_id (member_id), + INDEX idx_organization_team_members_team_created_id (team_id, created_at, id) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;`, // ----------------------------------- ) diff --git a/plugins/organizations/openapi/openapi_docs.go b/plugins/organizations/openapi/openapi_docs.go index 7e7a6359..40542590 100644 --- a/plugins/organizations/openapi/openapi_docs.go +++ b/plugins/organizations/openapi/openapi_docs.go @@ -26,9 +26,10 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { "/organizations", openapi.WithOperationID("listOrganizations"), openapi.WithSummary("List organizations"), - openapi.WithDescription("Lists all organizations owned by the authenticated user."), + openapi.WithDescription("Lists every organization the authenticated user can access, both the ones they own and the ones they are a member of, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected."), openapi.WithTags("Organizations"), - openapi.WithResponseStatus(http.StatusOK, &[]types.Organization{}), + openapi.WithRequest(&types.ListOrganizationsRequest{}), + openapi.WithResponseStatus(http.StatusOK, &types.ListOrganizationsResponse{}), ), svc.AddOperation( http.MethodGet, @@ -80,10 +81,10 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { "/organizations/{organization_id}/invitations", openapi.WithOperationID("listOrganizationInvitations"), openapi.WithSummary("List invitations"), - openapi.WithDescription("Lists all invitations for an organization."), + openapi.WithDescription("Lists the invitations for an organization, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected."), openapi.WithTags("Organization Invitations"), - openapi.WithRequest(&types.OrganizationID{}), - openapi.WithResponseStatus(http.StatusOK, &[]types.GetOrganizationInvitationResponse{}), + openapi.WithRequest(&types.ListOrganizationInvitationsRequest{}), + openapi.WithResponseStatus(http.StatusOK, &types.ListOrganizationInvitationsResponse{}), ), svc.AddOperation( http.MethodGet, @@ -143,10 +144,10 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { "/organizations/{organization_id}/members", openapi.WithOperationID("listOrganizationMembers"), openapi.WithSummary("List members"), - openapi.WithDescription("Lists all members of an organization with pagination."), + openapi.WithDescription("Lists the members of an organization, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected."), openapi.WithTags("Organization Members"), openapi.WithRequest(&types.ListOrganizationMembersRequest{}), - openapi.WithResponseStatus(http.StatusOK, &[]types.OrganizationMemberResponse{}), + openapi.WithResponseStatus(http.StatusOK, &types.ListOrganizationMembersResponse{}), ), svc.AddOperation( http.MethodGet, @@ -207,10 +208,10 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { "/organizations/{organization_id}/teams", openapi.WithOperationID("listOrganizationTeams"), openapi.WithSummary("List teams"), - openapi.WithDescription("Lists all teams within an organization."), + openapi.WithDescription("Lists the teams within an organization, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected."), openapi.WithTags("Organization Teams"), - openapi.WithRequest(&types.OrganizationID{}), - openapi.WithResponseStatus(http.StatusOK, &[]types.OrganizationTeam{}), + openapi.WithRequest(&types.ListOrganizationTeamsRequest{}), + openapi.WithResponseStatus(http.StatusOK, &types.ListOrganizationTeamsResponse{}), ), svc.AddOperation( http.MethodGet, @@ -261,10 +262,10 @@ func RegisterOpenAPIDocs(svc openapi.OpenAPIService) error { "/organizations/{organization_id}/teams/{team_id}/members", openapi.WithOperationID("listOrganizationTeamMembers"), openapi.WithSummary("List team members"), - openapi.WithDescription("Lists all members of a team with pagination."), + openapi.WithDescription("Lists the members of a team, newest first. Results are paginated: `page` defaults to 1 and `limit` defaults to 10, with a hard maximum of 100. Out-of-range values are clamped silently rather than rejected."), openapi.WithTags("Organization Team Members"), openapi.WithRequest(&types.ListOrganizationTeamMembersRequest{}), - openapi.WithResponseStatus(http.StatusOK, &[]types.OrganizationTeamMemberResponse{}), + openapi.WithResponseStatus(http.StatusOK, &types.ListOrganizationTeamMembersResponse{}), ), svc.AddOperation( http.MethodGet, diff --git a/plugins/organizations/repositories/bun_organization_invitation_repository.go b/plugins/organizations/repositories/bun_organization_invitation_repository.go index 96c6f7f6..5e02d8af 100644 --- a/plugins/organizations/repositories/bun_organization_invitation_repository.go +++ b/plugins/organizations/repositories/bun_organization_invitation_repository.go @@ -69,22 +69,16 @@ func (r *BunOrganizationInvitationRepository) GetByOrganizationIDAndEmail(ctx co return invitation, err } -func (r *BunOrganizationInvitationRepository) GetAllByOrganizationID(ctx context.Context, organizationID string) ([]types.OrganizationInvitation, error) { - invitations := make([]types.OrganizationInvitation, 0) - err := r.db.NewSelect().Model(&invitations). - Where("organization_id = ?", organizationID). - OrderExpr("created_at DESC"). - Scan(ctx) - if err == sql.ErrNoRows { - return []types.OrganizationInvitation{}, nil +func (r *BunOrganizationInvitationRepository) GetAllPendingByEmail(ctx context.Context, email string, limit int) ([]types.OrganizationInvitation, error) { + if limit <= 0 || limit > MaxPendingInvitationsPerBatch { + limit = MaxPendingInvitationsPerBatch } - return invitations, err -} -func (r *BunOrganizationInvitationRepository) GetAllPendingByEmail(ctx context.Context, email string) ([]types.OrganizationInvitation, error) { invites := make([]types.OrganizationInvitation, 0) err := r.db.NewSelect().Model(&invites). Where("email = ? AND status = ? AND expires_at > ?", email, types.OrganizationInvitationStatusPending, time.Now().UTC()). + OrderExpr("created_at ASC, id ASC"). + Limit(limit). Scan(ctx) if err == sql.ErrNoRows { return []types.OrganizationInvitation{}, nil @@ -142,6 +136,10 @@ const invitationWithOrgColumns = `i.id, i.email, i.inviter_id, i.organization_id ` o.id AS org_id, o.owner_id AS org_owner_id, o.name AS org_name,` + ` o.slug AS org_slug, o.logo AS org_logo, o.metadata AS org_metadata` +const invitationWithOrgByOrganizationFrom = ` FROM organization_invitations i` + + ` INNER JOIN organizations o ON o.id = i.organization_id` + + ` WHERE i.organization_id = ?` + func mapToInvitationWithOrgResponse(row invitationOrgRow) types.GetOrganizationInvitationResponse { return types.GetOrganizationInvitationResponse{ Invitation: &types.OrganizationInvitation{ @@ -183,24 +181,28 @@ func (r *BunOrganizationInvitationRepository) GetByIDWithOrg(ctx context.Context return &result, nil } -func (r *BunOrganizationInvitationRepository) GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { +func (r *BunOrganizationInvitationRepository) GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string, page int, limit int) ([]types.GetOrganizationInvitationResponse, int, error) { + limit = pageLimit(limit) + + var total int + if err := r.db.NewRaw(`SELECT COUNT(*)`+invitationWithOrgByOrganizationFrom, organizationID).Scan(ctx, &total); err != nil { + return nil, 0, err + } + var rows []invitationOrgRow - err := r.db.NewRaw(` - SELECT `+invitationWithOrgColumns+` - FROM organization_invitations i - INNER JOIN organizations o ON o.id = i.organization_id - WHERE i.organization_id = ? - ORDER BY i.created_at DESC - `, organizationID).Scan(ctx, &rows) + err := r.db.NewRaw(`SELECT `+invitationWithOrgColumns+invitationWithOrgByOrganizationFrom+` + ORDER BY i.created_at DESC, i.id DESC + LIMIT ? OFFSET ? + `, organizationID, limit, pageOffset(page, limit)).Scan(ctx, &rows) if err == sql.ErrNoRows { - return []types.GetOrganizationInvitationResponse{}, nil + return []types.GetOrganizationInvitationResponse{}, total, nil } if err != nil { - return nil, err + return nil, 0, err } results := make([]types.GetOrganizationInvitationResponse, len(rows)) for i, row := range rows { results[i] = mapToInvitationWithOrgResponse(row) } - return results, nil + return results, total, nil } diff --git a/plugins/organizations/repositories/bun_organization_invitation_repository_test.go b/plugins/organizations/repositories/bun_organization_invitation_repository_test.go index a93b72d6..e8d891a1 100644 --- a/plugins/organizations/repositories/bun_organization_invitation_repository_test.go +++ b/plugins/organizations/repositories/bun_organization_invitation_repository_test.go @@ -2,6 +2,7 @@ package repositories_test import ( "context" + "fmt" "testing" "time" @@ -200,116 +201,161 @@ func TestBunOrganizationInvitationRepository_GetByOrganizationIDAndEmail(t *test func TestBunOrganizationInvitationRepository_GetAllPendingByEmail(t *testing.T) { t.Parallel() + setup := func(t *testing.T) (repositories.OrganizationInvitationRepository, context.Context) { + t.Helper() + + db := plugintests.SetupRepoDB(t) + plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") + plugintests.SeedOrganization(t, db, "org-2", "user-2", "Beta Inc", "beta-inc") + plugintests.SeedOrganization(t, db, "org-3", "user-2", "Gamma Inc", "gamma-inc") + plugintests.SeedOrganization(t, db, "org-4", "user-2", "Delta Inc", "delta-inc") + plugintests.SeedOrganization(t, db, "org-5", "user-2", "Epsilon Inc", "epsilon-inc") + + repo := repositories.NewBunOrganizationInvitationRepository(db) + ctx := context.Background() + + for i, invitation := range []*types.OrganizationInvitation{ + {ID: "inv-1", OrganizationID: "org-1", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}, + {ID: "inv-2", OrganizationID: "org-2", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}, + {ID: "inv-3", OrganizationID: "org-3", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}, + {ID: "inv-4", OrganizationID: "org-4", Status: types.OrganizationInvitationStatusAccepted, ExpiresAt: time.Now().UTC().Add(time.Hour)}, + {ID: "inv-5", OrganizationID: "org-5", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(-time.Hour)}, + } { + invitation.Email = "user@example.com" + invitation.InviterID = "user-1" + invitation.Role = "member" + _, err := repo.Create(ctx, invitation) + require.NoError(t, err, "seeding invitation %d", i) + } + + return repo, ctx + } + tests := []struct { - name string - email string - setup func(*testing.T) (repositories.OrganizationInvitationRepository, context.Context) - expectPending int + name string + email string + limit int + expectedIDs []string }{ - { - name: "pending only", - email: "user@example.com", - expectPending: 1, - setup: func(t *testing.T) (repositories.OrganizationInvitationRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - plugintests.SeedOrganization(t, db, "org-2", "user-2", "Beta Inc", "beta-inc") - repo := repositories.NewBunOrganizationInvitationRepository(db) - ctx := context.Background() - - _, err := repo.Create(ctx, &types.OrganizationInvitation{ID: "inv-1", Email: "user@example.com", InviterID: "user-1", OrganizationID: "org-1", Role: "member", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationInvitation{ID: "inv-2", Email: "user@example.com", InviterID: "user-1", OrganizationID: "org-2", Role: "member", Status: types.OrganizationInvitationStatusAccepted, ExpiresAt: time.Now().UTC().Add(time.Hour)}) - require.NoError(t, err) - - return repo, ctx - }, - }, - { - name: "missing", - email: "missing@example.com", - expectPending: 0, - setup: func(t *testing.T) (repositories.OrganizationInvitationRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - plugintests.SeedOrganization(t, db, "org-2", "user-2", "Beta Inc", "beta-inc") - return repositories.NewBunOrganizationInvitationRepository(db), context.Background() - }, - }, + {name: "returns every pending invitation", email: "user@example.com", limit: 10, expectedIDs: []string{"inv-1", "inv-2", "inv-3"}}, + {name: "capped batch returns the oldest pending invitations", email: "user@example.com", limit: 2, expectedIDs: []string{"inv-1", "inv-2"}}, + {name: "unknown email has nothing pending", email: "missing@example.com", limit: 10, expectedIDs: []string{}}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - repo, ctx := tt.setup(t) + repo, ctx := setup(t) - pending, err := repo.GetAllPendingByEmail(ctx, tt.email) + pending, err := repo.GetAllPendingByEmail(ctx, tt.email, tt.limit) require.NoError(t, err) - require.Len(t, pending, tt.expectPending) + + ids := make([]string, 0, len(pending)) + for _, invitation := range pending { + ids = append(ids, invitation.ID) + } + require.Equal(t, tt.expectedIDs, ids) }) } } -func TestBunOrganizationInvitationRepository_GetAllByOrganizationID(t *testing.T) { +func TestBunOrganizationInvitationRepository_GetAllByOrganizationIDWithOrg(t *testing.T) { t.Parallel() + setup := func(t *testing.T) (repositories.OrganizationInvitationRepository, context.Context) { + t.Helper() + + db := plugintests.SetupRepoDB(t) + plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") + plugintests.SeedOrganization(t, db, "org-2", "user-2", "Beta Inc", "beta-inc") + + repo := repositories.NewBunOrganizationInvitationRepository(db) + ctx := context.Background() + + for i := 1; i <= 5; i++ { + _, err := repo.Create(ctx, &types.OrganizationInvitation{ + ID: fmt.Sprintf("inv-%d", i), + Email: fmt.Sprintf("user%d@example.com", i), + InviterID: "user-1", + OrganizationID: "org-1", + Role: "member", + Status: types.OrganizationInvitationStatusPending, + ExpiresAt: time.Now().UTC().Add(time.Hour), + }) + require.NoError(t, err) + } + + return repo, ctx + } + tests := []struct { name string organizationID string - setup func(*testing.T) (repositories.OrganizationInvitationRepository, context.Context) + page int + limit int expectCount int - verify func(*testing.T, []types.OrganizationInvitation) + expectTotal int }{ - { - name: "returns invitations for organization", - organizationID: "org-1", - expectCount: 2, - setup: func(t *testing.T) (repositories.OrganizationInvitationRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - plugintests.SeedOrganization(t, db, "org-2", "user-2", "Beta Inc", "beta-inc") - repo := repositories.NewBunOrganizationInvitationRepository(db) - ctx := context.Background() - - _, err := repo.Create(ctx, &types.OrganizationInvitation{ID: "inv-1", Email: "user@example.com", InviterID: "user-1", OrganizationID: "org-1", Role: "member", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationInvitation{ID: "inv-2", Email: "other@example.com", InviterID: "user-1", OrganizationID: "org-1", Role: "member", Status: types.OrganizationInvitationStatusRejected, ExpiresAt: time.Now().UTC().Add(time.Hour)}) - require.NoError(t, err) - - return repo, ctx - }, - verify: func(t *testing.T, invitations []types.OrganizationInvitation) { - t.Helper() - statusByID := map[string]types.OrganizationInvitationStatus{} - for _, invitation := range invitations { - statusByID[invitation.ID] = invitation.Status - } - - require.Equal(t, types.OrganizationInvitationStatusPending, statusByID["inv-1"]) - require.Equal(t, types.OrganizationInvitationStatusRejected, statusByID["inv-2"]) - }, - }, + {name: "first page", organizationID: "org-1", page: 1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "last partial page", organizationID: "org-1", page: 3, limit: 2, expectCount: 1, expectTotal: 5}, + {name: "page past the end keeps the true total", organizationID: "org-1", page: 4, limit: 2, expectCount: 0, expectTotal: 5}, + {name: "page zero does not error", organizationID: "org-1", page: 0, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "zero limit falls back to a bounded page", organizationID: "org-1", page: 1, limit: 0, expectCount: 5, expectTotal: 5}, + {name: "organization without invitations is empty", organizationID: "org-2", page: 1, limit: 10, expectCount: 0, expectTotal: 0}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - repo, ctx := tt.setup(t) - invitations, err := repo.GetAllByOrganizationID(ctx, tt.organizationID) + repo, ctx := setup(t) + + invitations, total, err := repo.GetAllByOrganizationIDWithOrg(ctx, tt.organizationID, tt.page, tt.limit) require.NoError(t, err) require.Len(t, invitations, tt.expectCount) - if tt.verify != nil { - tt.verify(t, invitations) + require.Equal(t, tt.expectTotal, total) + for _, invitation := range invitations { + require.Equal(t, tt.organizationID, invitation.Organization.ID, "joined organization must be populated") } }) } } +func TestBunOrganizationInvitationRepository_GetAllByOrganizationIDWithOrgPagesPartitionCleanly(t *testing.T) { + t.Parallel() + + db := plugintests.SetupRepoDB(t) + plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") + repo := repositories.NewBunOrganizationInvitationRepository(db) + ctx := context.Background() + + for i := 1; i <= 5; i++ { + _, err := repo.Create(ctx, &types.OrganizationInvitation{ + ID: fmt.Sprintf("inv-%d", i), + Email: fmt.Sprintf("user%d@example.com", i), + InviterID: "user-1", + OrganizationID: "org-1", + Role: "member", + Status: types.OrganizationInvitationStatusPending, + ExpiresAt: time.Now().UTC().Add(time.Hour), + }) + require.NoError(t, err) + } + + seen := make([]string, 0, 5) + for page := 1; page <= 3; page++ { + invitations, total, err := repo.GetAllByOrganizationIDWithOrg(ctx, "org-1", page, 2) + require.NoError(t, err) + require.Equal(t, 5, total) + for _, invitation := range invitations { + seen = append(seen, invitation.Invitation.ID) + } + } + + require.ElementsMatch(t, []string{"inv-1", "inv-2", "inv-3", "inv-4", "inv-5"}, seen) +} + func TestBunOrganizationInvitationRepository_GetByID(t *testing.T) { t.Parallel() diff --git a/plugins/organizations/repositories/bun_organization_member_repository.go b/plugins/organizations/repositories/bun_organization_member_repository.go index 2f0528fb..c43da1f5 100644 --- a/plugins/organizations/repositories/bun_organization_member_repository.go +++ b/plugins/organizations/repositories/bun_organization_member_repository.go @@ -54,6 +54,10 @@ const memberWithUserColumns = `m.id, m.organization_id, m.role, m.created_at, m. ` u.metadata AS user_metadata, u.created_at AS user_created_at,` + ` u.updated_at AS user_updated_at` +const memberWithUserByOrganizationFrom = ` FROM organization_members m` + + ` INNER JOIN users u ON u.id = m.user_id` + + ` WHERE m.organization_id = ?` + type BunOrganizationMemberRepository struct { db bun.IDB } @@ -95,37 +99,41 @@ func (r *BunOrganizationMemberRepository) GetByID(ctx context.Context, memberID return member, err } -func (r *BunOrganizationMemberRepository) GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMember, error) { +func (r *BunOrganizationMemberRepository) GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMember, int, error) { members := make([]types.OrganizationMember, 0) - err := r.db.NewSelect().Model(&members). + limit = pageLimit(limit) + total, err := r.db.NewSelect().Model(&members). Where("organization_id = ?", organizationID). OrderExpr("created_at DESC, id DESC"). - Offset((page - 1) * limit).Limit(limit). - Scan(ctx) + Offset(pageOffset(page, limit)).Limit(limit). + ScanAndCount(ctx) if err == sql.ErrNoRows { - return []types.OrganizationMember{}, nil + return []types.OrganizationMember{}, total, nil } - return members, err + return members, total, err } -func (r *BunOrganizationMemberRepository) GetAllByOrganizationIDWithUser(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, error) { +func (r *BunOrganizationMemberRepository) GetAllByOrganizationIDWithUser(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, int, error) { + limit = pageLimit(limit) + + var total int + if err := r.db.NewRaw(`SELECT COUNT(*)`+memberWithUserByOrganizationFrom, organizationID).Scan(ctx, &total); err != nil { + return nil, 0, err + } + var rows []memberUserRow - err := r.db.NewRaw(` - SELECT `+memberWithUserColumns+` - FROM organization_members m - INNER JOIN users u ON u.id = m.user_id - WHERE m.organization_id = ? + err := r.db.NewRaw(`SELECT `+memberWithUserColumns+memberWithUserByOrganizationFrom+` ORDER BY m.created_at DESC, m.id DESC LIMIT ? OFFSET ? - `, organizationID, limit, (page-1)*limit).Scan(ctx, &rows) + `, organizationID, limit, pageOffset(page, limit)).Scan(ctx, &rows) if err != nil { - return nil, err + return nil, 0, err } results := make([]types.OrganizationMemberResponse, len(rows)) for i, row := range rows { results[i] = mapToMemberResponse(row) } - return results, nil + return results, total, nil } func (r *BunOrganizationMemberRepository) GetByIDWithUser(ctx context.Context, memberID string) (*types.OrganizationMemberResponse, error) { @@ -164,17 +172,6 @@ func (r *BunOrganizationMemberRepository) GetByOrganizationIDAndUserIDWithUser(c return &result, nil } -func (r *BunOrganizationMemberRepository) GetAllByUserID(ctx context.Context, userID string) ([]types.OrganizationMember, error) { - members := make([]types.OrganizationMember, 0) - err := r.db.NewSelect().Model(&members). - Where("user_id = ?", userID). - Scan(ctx) - if err == sql.ErrNoRows { - return []types.OrganizationMember{}, nil - } - return members, err -} - func (r *BunOrganizationMemberRepository) GetByOrganizationIDAndUserID(ctx context.Context, organizationID string, userID string) (*types.OrganizationMember, error) { member := new(types.OrganizationMember) err := r.db.NewSelect().Model(member). diff --git a/plugins/organizations/repositories/bun_organization_member_repository_test.go b/plugins/organizations/repositories/bun_organization_member_repository_test.go index f57c619d..e35025a1 100644 --- a/plugins/organizations/repositories/bun_organization_member_repository_test.go +++ b/plugins/organizations/repositories/bun_organization_member_repository_test.go @@ -2,6 +2,7 @@ package repositories_test import ( "context" + "fmt" "testing" "github.com/stretchr/testify/require" @@ -33,9 +34,10 @@ func TestBunOrganizationMemberRepository_CreateGetUpdateDelete(t *testing.T) { name: "list by organization returns created member", run: func(t *testing.T, orgMemberRepo repositories.OrganizationMemberRepository, ctx context.Context, created *types.OrganizationMember) { t.Helper() - members, err := orgMemberRepo.GetAllByOrganizationID(ctx, "org-1", 1, 10) + members, total, err := orgMemberRepo.GetAllByOrganizationID(ctx, "org-1", 1, 10) require.NoError(t, err) require.Len(t, members, 1) + require.Equal(t, 1, total) require.Equal(t, created.ID, members[0].ID) }, }, @@ -143,152 +145,149 @@ func TestBunOrganizationMemberRepository_Create(t *testing.T) { } } +func seedMembers(t *testing.T, count int) (repositories.OrganizationMemberRepository, context.Context) { + t.Helper() + + db := plugintests.SetupRepoDB(t) + plugintests.SeedUsers(t, db, count) + plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") + + repo := repositories.NewBunOrganizationMemberRepository(db) + ctx := context.Background() + for i := 1; i <= count; i++ { + _, err := repo.Create(ctx, &types.OrganizationMember{ + ID: fmt.Sprintf("mem-%d", i), + OrganizationID: "org-1", + UserID: fmt.Sprintf("user-%d", i), + Role: "member", + }) + require.NoError(t, err) + } + + return repo, ctx +} + +func memberIDs(members []types.OrganizationMember) []string { + ids := make([]string, 0, len(members)) + for _, member := range members { + ids = append(ids, member.ID) + } + return ids +} + +func memberResponseIDs(members []types.OrganizationMemberResponse) []string { + ids := make([]string, 0, len(members)) + for _, member := range members { + ids = append(ids, member.ID) + } + return ids +} + func TestBunOrganizationMemberRepository_GetAllByOrganizationID(t *testing.T) { t.Parallel() + tests := []struct { name string organizationID string page int limit int - setup func(*testing.T) (repositories.OrganizationMemberRepository, context.Context) expectCount int + expectTotal int }{ - { - name: "first page", - organizationID: "org-1", - page: 1, - limit: 1, - expectCount: 1, - setup: func(t *testing.T) (repositories.OrganizationMemberRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - repo := repositories.NewBunOrganizationMemberRepository(db) - ctx := context.Background() - - _, err := repo.Create(ctx, &types.OrganizationMember{ID: "mem-1", OrganizationID: "org-1", UserID: "user-1", Role: "member"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationMember{ID: "mem-2", OrganizationID: "org-1", UserID: "user-2", Role: "admin"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationMember{ID: "mem-3", OrganizationID: "org-1", UserID: "user-3", Role: "member"}) - require.NoError(t, err) - - return repo, ctx - }, - }, - { - name: "second page", - organizationID: "org-1", - page: 2, - limit: 1, - expectCount: 1, - setup: func(t *testing.T) (repositories.OrganizationMemberRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - repo := repositories.NewBunOrganizationMemberRepository(db) - ctx := context.Background() - - _, err := repo.Create(ctx, &types.OrganizationMember{ID: "mem-1", OrganizationID: "org-1", UserID: "user-1", Role: "member"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationMember{ID: "mem-2", OrganizationID: "org-1", UserID: "user-2", Role: "admin"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationMember{ID: "mem-3", OrganizationID: "org-1", UserID: "user-3", Role: "member"}) - require.NoError(t, err) - - return repo, ctx - }, - }, - { - name: "empty result", - organizationID: "org-2", - page: 1, - limit: 10, - expectCount: 0, - setup: func(t *testing.T) (repositories.OrganizationMemberRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - repo := repositories.NewBunOrganizationMemberRepository(db) - return repo, context.Background() - }, - }, + {name: "first page", organizationID: "org-1", page: 1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "last partial page", organizationID: "org-1", page: 3, limit: 2, expectCount: 1, expectTotal: 5}, + {name: "page past the end keeps the true total", organizationID: "org-1", page: 4, limit: 2, expectCount: 0, expectTotal: 5}, + {name: "page zero does not error", organizationID: "org-1", page: 0, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "negative page does not error", organizationID: "org-1", page: -1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "zero limit falls back to a bounded page", organizationID: "org-1", page: 1, limit: 0, expectCount: 5, expectTotal: 5}, + {name: "negative limit falls back to a bounded page", organizationID: "org-1", page: 1, limit: -1, expectCount: 5, expectTotal: 5}, + {name: "unknown organization is empty", organizationID: "org-2", page: 1, limit: 10, expectCount: 0, expectTotal: 0}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - repo, ctx := tt.setup(t) + repo, ctx := seedMembers(t, 5) - members, err := repo.GetAllByOrganizationID(ctx, tt.organizationID, tt.page, tt.limit) + members, total, err := repo.GetAllByOrganizationID(ctx, tt.organizationID, tt.page, tt.limit) require.NoError(t, err) require.Len(t, members, tt.expectCount) - if tt.page == 1 && len(members) > 0 { - require.Equal(t, "mem-3", members[0].ID) - } - if tt.page == 2 && len(members) > 0 { - require.Equal(t, "mem-2", members[0].ID) - } + require.Equal(t, tt.expectTotal, total) }) } } -func TestBunOrganizationMemberRepository_GetAllByUserID(t *testing.T) { +func TestBunOrganizationMemberRepository_GetAllByOrganizationIDPagesPartitionCleanly(t *testing.T) { t.Parallel() - tests := []struct { - name string - userID string - setup func(*testing.T) (repositories.OrganizationMemberRepository, context.Context) - expectCount int - }{ - { - name: "found", - userID: "user-1", - expectCount: 2, - setup: func(t *testing.T) (repositories.OrganizationMemberRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - plugintests.SeedOrganization(t, db, "org-2", "user-2", "Platform", "platform") - repo := repositories.NewBunOrganizationMemberRepository(db) - ctx := context.Background() - _, err := repo.Create(ctx, &types.OrganizationMember{ID: "mem-1", OrganizationID: "org-1", UserID: "user-1", Role: "member"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationMember{ID: "mem-2", OrganizationID: "org-2", UserID: "user-1", Role: "admin"}) - require.NoError(t, err) + repo, ctx := seedMembers(t, 5) - return repo, ctx - }, - }, - { - name: "empty", - userID: "user-3", - expectCount: 0, - setup: func(t *testing.T) (repositories.OrganizationMemberRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - plugintests.SeedOrganization(t, db, "org-2", "user-2", "Platform", "platform") - return repositories.NewBunOrganizationMemberRepository(db), context.Background() - }, - }, + seen := make([]string, 0, 5) + for page := 1; page <= 3; page++ { + members, total, err := repo.GetAllByOrganizationID(ctx, "org-1", page, 2) + require.NoError(t, err) + require.Equal(t, 5, total) + seen = append(seen, memberIDs(members)...) + } + + require.ElementsMatch(t, []string{"mem-1", "mem-2", "mem-3", "mem-4", "mem-5"}, seen) +} + +func TestBunOrganizationMemberRepository_GetAllByOrganizationIDWithUser(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + organizationID string + page int + limit int + expectCount int + expectTotal int + }{ + {name: "first page", organizationID: "org-1", page: 1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "last partial page", organizationID: "org-1", page: 3, limit: 2, expectCount: 1, expectTotal: 5}, + {name: "page past the end keeps the true total", organizationID: "org-1", page: 4, limit: 2, expectCount: 0, expectTotal: 5}, + {name: "page zero does not error", organizationID: "org-1", page: 0, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "negative page does not error", organizationID: "org-1", page: -1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "zero limit falls back to a bounded page", organizationID: "org-1", page: 1, limit: 0, expectCount: 5, expectTotal: 5}, + {name: "limit beyond the maximum is capped", organizationID: "org-1", page: 1, limit: 100000, expectCount: 5, expectTotal: 5}, + {name: "unknown organization is empty", organizationID: "org-2", page: 1, limit: 10, expectCount: 0, expectTotal: 0}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - repo, ctx := tt.setup(t) + repo, ctx := seedMembers(t, 5) - members, err := repo.GetAllByUserID(ctx, tt.userID) + members, total, err := repo.GetAllByOrganizationIDWithUser(ctx, tt.organizationID, tt.page, tt.limit) require.NoError(t, err) require.Len(t, members, tt.expectCount) + require.Equal(t, tt.expectTotal, total) + for _, member := range members { + require.NotEmpty(t, member.User.ID, "joined user must be populated") + } }) } } +func TestBunOrganizationMemberRepository_GetAllByOrganizationIDWithUserPagesPartitionCleanly(t *testing.T) { + t.Parallel() + + repo, ctx := seedMembers(t, 5) + + seen := make([]string, 0, 5) + for page := 1; page <= 3; page++ { + members, total, err := repo.GetAllByOrganizationIDWithUser(ctx, "org-1", page, 2) + require.NoError(t, err) + require.Equal(t, 5, total) + seen = append(seen, memberResponseIDs(members)...) + } + + require.ElementsMatch(t, []string{"mem-1", "mem-2", "mem-3", "mem-4", "mem-5"}, seen) +} + func TestBunOrganizationMemberRepository_GetByOrganizationIDAndUserID(t *testing.T) { t.Parallel() tests := []struct { diff --git a/plugins/organizations/repositories/bun_organization_repository.go b/plugins/organizations/repositories/bun_organization_repository.go index 17819cb5..d593148b 100644 --- a/plugins/organizations/repositories/bun_organization_repository.go +++ b/plugins/organizations/repositories/bun_organization_repository.go @@ -55,13 +55,31 @@ func (r *BunOrganizationRepository) GetBySlug(ctx context.Context, slug string) return organization, err } -func (r *BunOrganizationRepository) GetAllByOwnerID(ctx context.Context, ownerID string) ([]types.Organization, error) { +const organizationAccessibleWhere = `o.owner_id = ? OR EXISTS (` + + `SELECT 1 FROM organization_members m WHERE m.organization_id = o.id AND m.user_id = ?)` + +func (r *BunOrganizationRepository) GetAllAccessibleByUserID(ctx context.Context, userID string, page int, limit int) ([]types.Organization, int, error) { organizations := make([]types.Organization, 0) - err := r.db.NewSelect().Model(&organizations).Where("owner_id = ?", ownerID).Order("created_at DESC").Scan(ctx) + limit = pageLimit(limit) + total, err := r.db.NewSelect().Model(&organizations). + ModelTableExpr("organizations AS o"). + ColumnExpr("o.*"). + Where(organizationAccessibleWhere, userID, userID). + OrderExpr("o.created_at DESC, o.id DESC"). + Offset(pageOffset(page, limit)).Limit(limit). + ScanAndCount(ctx) if err == sql.ErrNoRows { - return []types.Organization{}, nil + return []types.Organization{}, total, nil } - return organizations, err + return organizations, total, err +} + +func (r *BunOrganizationRepository) CountAccessibleByUserID(ctx context.Context, userID string) (int, error) { + return r.db.NewSelect(). + Model((*types.Organization)(nil)). + ModelTableExpr("organizations AS o"). + Where(organizationAccessibleWhere, userID, userID). + Count(ctx) } func (r *BunOrganizationRepository) Update(ctx context.Context, organization *types.Organization) (*types.Organization, error) { diff --git a/plugins/organizations/repositories/bun_organization_repository_test.go b/plugins/organizations/repositories/bun_organization_repository_test.go index bf73d39a..d76c2728 100644 --- a/plugins/organizations/repositories/bun_organization_repository_test.go +++ b/plugins/organizations/repositories/bun_organization_repository_test.go @@ -152,41 +152,74 @@ func TestBunOrganizationRepository_GetBySlug(t *testing.T) { } } -func TestBunOrganizationRepository_GetAllByOwnerID(t *testing.T) { +func seedAccessibleOrganizations(t *testing.T) (repositories.OrganizationRepository, context.Context) { + t.Helper() + + db := plugintests.SetupRepoDB(t) + ctx := context.Background() + + plugintests.SeedOrganization(t, db, "org-a", "user-1", "Owned", "owned") + plugintests.SeedOrganization(t, db, "org-b", "user-2", "Member Only", "member-only") + plugintests.SeedOrganization(t, db, "org-c", "user-1", "Owner And Member", "owner-and-member") + plugintests.SeedOrganization(t, db, "org-d", "user-2", "Unrelated", "unrelated") + + plugintests.SeedOrganizationMember(t, db, "mem-b", "org-b", "user-1", "member") + plugintests.SeedOrganizationMember(t, db, "mem-c", "org-c", "user-1", "owner") + plugintests.SeedOrganizationMember(t, db, "mem-d", "org-d", "user-2", "owner") + + return repositories.NewBunOrganizationRepository(db), ctx +} + +func organizationIDs(organizations []types.Organization) []string { + ids := make([]string, 0, len(organizations)) + for _, organization := range organizations { + ids = append(ids, organization.ID) + } + return ids +} + +func TestBunOrganizationRepository_GetAllAccessibleByUserID(t *testing.T) { t.Parallel() tests := []struct { name string - ownerID string - setup func(*testing.T) (repositories.OrganizationRepository, context.Context) - expectCount int + userID string + page int + limit int + expectedIDs []string + expectTotal int }{ { - name: "found", - ownerID: "user-1", - expectCount: 2, - setup: func(t *testing.T) (repositories.OrganizationRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - repo := repositories.NewBunOrganizationRepository(db) - ctx := context.Background() - - _, err := repo.Create(ctx, &types.Organization{ID: "org-1", OwnerID: "user-1", Name: "Acme Inc", Slug: "acme-inc"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.Organization{ID: "org-2", OwnerID: "user-1", Name: "Platform", Slug: "platform"}) - require.NoError(t, err) - - return repo, ctx - }, + name: "returns owned and joined organizations without duplicates", + userID: "user-1", + page: 1, + limit: 10, + expectedIDs: []string{"org-a", "org-b", "org-c"}, + expectTotal: 3, }, { - name: "empty", - ownerID: "user-2", - expectCount: 0, - setup: func(t *testing.T) (repositories.OrganizationRepository, context.Context) { - t.Helper() - return repositories.NewBunOrganizationRepository(plugintests.SetupRepoDB(t)), context.Background() - }, + name: "excludes organizations the user cannot access", + userID: "user-2", + page: 1, + limit: 10, + expectedIDs: []string{"org-b", "org-d"}, + expectTotal: 2, + }, + { + name: "page past the end is empty but keeps the true total", + userID: "user-1", + page: 4, + limit: 2, + expectedIDs: []string{}, + expectTotal: 3, + }, + { + name: "page zero does not produce a negative offset", + userID: "user-1", + page: 0, + limit: 10, + expectedIDs: []string{"org-a", "org-b", "org-c"}, + expectTotal: 3, }, } @@ -194,11 +227,59 @@ func TestBunOrganizationRepository_GetAllByOwnerID(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - repo, ctx := tt.setup(t) + repo, ctx := seedAccessibleOrganizations(t) + + found, total, err := repo.GetAllAccessibleByUserID(ctx, tt.userID, tt.page, tt.limit) + require.NoError(t, err) + require.Equal(t, tt.expectTotal, total) + require.ElementsMatch(t, tt.expectedIDs, organizationIDs(found)) + }) + } +} + +func TestBunOrganizationRepository_GetAllAccessibleByUserIDPagesWithoutLoss(t *testing.T) { + t.Parallel() + + repo, ctx := seedAccessibleOrganizations(t) + + seen := make([]string, 0, 3) + for page := 1; page <= 3; page++ { + found, total, err := repo.GetAllAccessibleByUserID(ctx, "user-1", page, 1) + require.NoError(t, err) + require.Equal(t, 3, total) + require.Len(t, found, 1) + seen = append(seen, found[0].ID) + } + + require.ElementsMatch(t, []string{"org-a", "org-b", "org-c"}, seen) +} + +func TestBunOrganizationRepository_CountAccessibleByUserID(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + userID string + expected int + }{ + {name: "counts owned and joined organizations once each", userID: "user-1", expected: 3}, + {name: "counts only what the user can access", userID: "user-2", expected: 2}, + {name: "unknown user has no organizations", userID: "user-unknown", expected: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + repo, ctx := seedAccessibleOrganizations(t) + + count, err := repo.CountAccessibleByUserID(ctx, tt.userID) + require.NoError(t, err) + require.Equal(t, tt.expected, count) - found, err := repo.GetAllByOwnerID(ctx, tt.ownerID) + _, total, err := repo.GetAllAccessibleByUserID(ctx, tt.userID, 1, 10) require.NoError(t, err) - require.Len(t, found, tt.expectCount) + require.Equal(t, count, total, "count must agree with the list total") }) } } diff --git a/plugins/organizations/repositories/bun_organization_team_member_repository.go b/plugins/organizations/repositories/bun_organization_team_member_repository.go index 2acf00a3..51533ab8 100644 --- a/plugins/organizations/repositories/bun_organization_team_member_repository.go +++ b/plugins/organizations/repositories/bun_organization_team_member_repository.go @@ -65,6 +65,11 @@ const teamMemberWithMemberAndUserColumns = `tm.id, tm.team_id, tm.created_at,` + ` u.metadata AS user_metadata, u.created_at AS user_created_at,` + ` u.updated_at AS user_updated_at` +const teamMemberWithMemberAndUserByTeamFrom = ` FROM organization_team_members tm` + + ` INNER JOIN organization_members m ON m.id = tm.member_id` + + ` INNER JOIN users u ON u.id = m.user_id` + + ` WHERE tm.team_id = ?` + type BunOrganizationTeamMemberRepository struct { db bun.IDB } @@ -113,38 +118,41 @@ func (r *BunOrganizationTeamMemberRepository) GetByTeamIDAndMemberID(ctx context return teamMember, err } -func (r *BunOrganizationTeamMemberRepository) GetAllByTeamID(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMember, error) { +func (r *BunOrganizationTeamMemberRepository) GetAllByTeamID(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMember, int, error) { teamMembers := make([]types.OrganizationTeamMember, 0) - err := r.db.NewSelect().Model(&teamMembers). + limit = pageLimit(limit) + total, err := r.db.NewSelect().Model(&teamMembers). Where("team_id = ?", teamID). OrderExpr("created_at DESC, id DESC"). - Offset((page - 1) * limit).Limit(limit). - Scan(ctx) + Offset(pageOffset(page, limit)).Limit(limit). + ScanAndCount(ctx) if err == sql.ErrNoRows { - return []types.OrganizationTeamMember{}, nil + return []types.OrganizationTeamMember{}, total, nil } - return teamMembers, err + return teamMembers, total, err } -func (r *BunOrganizationTeamMemberRepository) GetAllByTeamIDWithMemberAndUser(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, error) { +func (r *BunOrganizationTeamMemberRepository) GetAllByTeamIDWithMemberAndUser(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, int, error) { + limit = pageLimit(limit) + + var total int + if err := r.db.NewRaw(`SELECT COUNT(*)`+teamMemberWithMemberAndUserByTeamFrom, teamID).Scan(ctx, &total); err != nil { + return nil, 0, err + } + var rows []teamMemberMemberUserRow - err := r.db.NewRaw(` - SELECT `+teamMemberWithMemberAndUserColumns+` - FROM organization_team_members tm - INNER JOIN organization_members m ON m.id = tm.member_id - INNER JOIN users u ON u.id = m.user_id - WHERE tm.team_id = ? + err := r.db.NewRaw(`SELECT `+teamMemberWithMemberAndUserColumns+teamMemberWithMemberAndUserByTeamFrom+` ORDER BY tm.created_at DESC, tm.id DESC LIMIT ? OFFSET ? - `, teamID, limit, (page-1)*limit).Scan(ctx, &rows) + `, teamID, limit, pageOffset(page, limit)).Scan(ctx, &rows) if err != nil { - return nil, err + return nil, 0, err } results := make([]types.OrganizationTeamMemberResponse, len(rows)) for i, row := range rows { results[i] = mapToTeamMemberResponse(row) } - return results, nil + return results, total, nil } func (r *BunOrganizationTeamMemberRepository) GetByIDWithMemberAndUser(ctx context.Context, teamMemberID string) (*types.OrganizationTeamMemberResponse, error) { diff --git a/plugins/organizations/repositories/bun_organization_team_member_repository_test.go b/plugins/organizations/repositories/bun_organization_team_member_repository_test.go index a2befb95..2a20cf10 100644 --- a/plugins/organizations/repositories/bun_organization_team_member_repository_test.go +++ b/plugins/organizations/repositories/bun_organization_team_member_repository_test.go @@ -2,6 +2,7 @@ package repositories_test import ( "context" + "fmt" "testing" "github.com/stretchr/testify/require" @@ -234,6 +235,45 @@ func TestBunOrganizationTeamMemberRepository_GetByTeamIDAndMemberID(t *testing.T } } +func seedTeamMembers(t *testing.T, count int) (repositories.OrganizationTeamMemberRepository, context.Context) { + t.Helper() + + db := plugintests.SetupRepoDB(t) + plugintests.SeedUsers(t, db, count) + plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") + plugintests.SeedOrganizationTeam(t, db, "team-1", "org-1", "Platform", "platform") + + repo := repositories.NewBunOrganizationTeamMemberRepository(db) + ctx := context.Background() + for i := 1; i <= count; i++ { + plugintests.SeedOrganizationMember(t, db, fmt.Sprintf("member-%d", i), "org-1", fmt.Sprintf("user-%d", i), "member") + _, err := repo.Create(ctx, &types.OrganizationTeamMember{ + ID: fmt.Sprintf("team-member-%d", i), + TeamID: "team-1", + MemberID: fmt.Sprintf("member-%d", i), + }) + require.NoError(t, err) + } + + return repo, ctx +} + +func teamMemberIDs(teamMembers []types.OrganizationTeamMember) []string { + ids := make([]string, 0, len(teamMembers)) + for _, teamMember := range teamMembers { + ids = append(ids, teamMember.ID) + } + return ids +} + +func teamMemberResponseIDs(teamMembers []types.OrganizationTeamMemberResponse) []string { + ids := make([]string, 0, len(teamMembers)) + for _, teamMember := range teamMembers { + ids = append(ids, teamMember.ID) + } + return ids +} + func TestBunOrganizationTeamMemberRepository_GetAllByTeamID(t *testing.T) { t.Parallel() @@ -242,104 +282,104 @@ func TestBunOrganizationTeamMemberRepository_GetAllByTeamID(t *testing.T) { teamID string page int limit int - setup func(*testing.T) (repositories.OrganizationTeamMemberRepository, context.Context) expectCount int + expectTotal int }{ - { - name: "first page", - teamID: "team-1", - page: 1, - limit: 2, - expectCount: 2, - setup: func(t *testing.T) (repositories.OrganizationTeamMemberRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - plugintests.SeedOrganizationMember(t, db, "member-1", "org-1", "user-1", "member") - plugintests.SeedOrganizationMember(t, db, "member-2", "org-1", "user-2", "admin") - plugintests.SeedUser(t, db, "user-3") - plugintests.SeedOrganizationMember(t, db, "member-3", "org-1", "user-3", "member") - plugintests.SeedOrganizationTeam(t, db, "team-1", "org-1", "Platform", "platform") - repo := repositories.NewBunOrganizationTeamMemberRepository(db) - ctx := context.Background() - _, err := repo.Create(ctx, &types.OrganizationTeamMember{ID: "team-member-1", TeamID: "team-1", MemberID: "member-1"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationTeamMember{ID: "team-member-2", TeamID: "team-1", MemberID: "member-2"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationTeamMember{ID: "team-member-3", TeamID: "team-1", MemberID: "member-3"}) - require.NoError(t, err) - return repo, ctx - }, - }, - { - name: "second page", - teamID: "team-1", - page: 2, - limit: 2, - expectCount: 1, - setup: func(t *testing.T) (repositories.OrganizationTeamMemberRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - plugintests.SeedOrganizationMember(t, db, "member-1", "org-1", "user-1", "member") - plugintests.SeedOrganizationMember(t, db, "member-2", "org-1", "user-2", "admin") - plugintests.SeedUser(t, db, "user-3") - plugintests.SeedOrganizationMember(t, db, "member-3", "org-1", "user-3", "member") - plugintests.SeedOrganizationTeam(t, db, "team-1", "org-1", "Platform", "platform") - repo := repositories.NewBunOrganizationTeamMemberRepository(db) - ctx := context.Background() - _, err := repo.Create(ctx, &types.OrganizationTeamMember{ID: "team-member-1", TeamID: "team-1", MemberID: "member-1"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationTeamMember{ID: "team-member-2", TeamID: "team-1", MemberID: "member-2"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationTeamMember{ID: "team-member-3", TeamID: "team-1", MemberID: "member-3"}) - require.NoError(t, err) - return repo, ctx - }, - }, - { - name: "empty", - teamID: "team-2", - page: 1, - limit: 10, - expectCount: 0, - setup: func(t *testing.T) (repositories.OrganizationTeamMemberRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - plugintests.SeedOrganizationMember(t, db, "member-1", "org-1", "user-1", "member") - plugintests.SeedOrganizationMember(t, db, "member-2", "org-1", "user-2", "admin") - plugintests.SeedUser(t, db, "user-3") - plugintests.SeedOrganizationMember(t, db, "member-3", "org-1", "user-3", "member") - plugintests.SeedOrganizationTeam(t, db, "team-1", "org-1", "Platform", "platform") - return repositories.NewBunOrganizationTeamMemberRepository(db), context.Background() - }, - }, + {name: "first page", teamID: "team-1", page: 1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "last partial page", teamID: "team-1", page: 3, limit: 2, expectCount: 1, expectTotal: 5}, + {name: "page past the end keeps the true total", teamID: "team-1", page: 4, limit: 2, expectCount: 0, expectTotal: 5}, + {name: "page zero does not error", teamID: "team-1", page: 0, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "negative page does not error", teamID: "team-1", page: -1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "zero limit falls back to a bounded page", teamID: "team-1", page: 1, limit: 0, expectCount: 5, expectTotal: 5}, + {name: "unknown team is empty", teamID: "team-2", page: 1, limit: 10, expectCount: 0, expectTotal: 0}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - repo, ctx := tt.setup(t) + repo, ctx := seedTeamMembers(t, 5) - found, err := repo.GetAllByTeamID(ctx, tt.teamID, tt.page, tt.limit) + found, total, err := repo.GetAllByTeamID(ctx, tt.teamID, tt.page, tt.limit) require.NoError(t, err) require.Len(t, found, tt.expectCount) + require.Equal(t, tt.expectTotal, total) for _, teamMember := range found { - require.Equal(t, "team-1", teamMember.TeamID) - } - if tt.page == 1 && len(found) > 1 { - require.Equal(t, "team-member-3", found[0].ID) - require.Equal(t, "team-member-2", found[1].ID) + require.Equal(t, tt.teamID, teamMember.TeamID) } - if tt.page == 2 && len(found) > 0 { - require.Equal(t, "team-member-1", found[0].ID) + }) + } +} + +func TestBunOrganizationTeamMemberRepository_GetAllByTeamIDPagesPartitionCleanly(t *testing.T) { + t.Parallel() + + repo, ctx := seedTeamMembers(t, 5) + + seen := make([]string, 0, 5) + for page := 1; page <= 3; page++ { + found, total, err := repo.GetAllByTeamID(ctx, "team-1", page, 2) + require.NoError(t, err) + require.Equal(t, 5, total) + seen = append(seen, teamMemberIDs(found)...) + } + + require.ElementsMatch(t, []string{"team-member-1", "team-member-2", "team-member-3", "team-member-4", "team-member-5"}, seen) +} + +func TestBunOrganizationTeamMemberRepository_GetAllByTeamIDWithMemberAndUser(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + teamID string + page int + limit int + expectCount int + expectTotal int + }{ + {name: "first page", teamID: "team-1", page: 1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "last partial page", teamID: "team-1", page: 3, limit: 2, expectCount: 1, expectTotal: 5}, + {name: "page past the end keeps the true total", teamID: "team-1", page: 4, limit: 2, expectCount: 0, expectTotal: 5}, + {name: "page zero does not error", teamID: "team-1", page: 0, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "zero limit falls back to a bounded page", teamID: "team-1", page: 1, limit: 0, expectCount: 5, expectTotal: 5}, + {name: "limit beyond the maximum is capped", teamID: "team-1", page: 1, limit: 100000, expectCount: 5, expectTotal: 5}, + {name: "unknown team is empty", teamID: "team-2", page: 1, limit: 10, expectCount: 0, expectTotal: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + repo, ctx := seedTeamMembers(t, 5) + + found, total, err := repo.GetAllByTeamIDWithMemberAndUser(ctx, tt.teamID, tt.page, tt.limit) + require.NoError(t, err) + require.Len(t, found, tt.expectCount) + require.Equal(t, tt.expectTotal, total) + for _, teamMember := range found { + require.NotEmpty(t, teamMember.Member.User.ID, "joined user must be populated") } }) } } +func TestBunOrganizationTeamMemberRepository_GetAllByTeamIDWithMemberAndUserPagesPartitionCleanly(t *testing.T) { + t.Parallel() + + repo, ctx := seedTeamMembers(t, 5) + + seen := make([]string, 0, 5) + for page := 1; page <= 3; page++ { + found, total, err := repo.GetAllByTeamIDWithMemberAndUser(ctx, "team-1", page, 2) + require.NoError(t, err) + require.Equal(t, 5, total) + seen = append(seen, teamMemberResponseIDs(found)...) + } + + require.ElementsMatch(t, []string{"team-member-1", "team-member-2", "team-member-3", "team-member-4", "team-member-5"}, seen) +} + func TestBunOrganizationTeamMemberRepository_DeleteByTeamIDAndMemberID(t *testing.T) { t.Parallel() diff --git a/plugins/organizations/repositories/bun_organization_team_repository.go b/plugins/organizations/repositories/bun_organization_team_repository.go index 65297fe0..82ec1f10 100644 --- a/plugins/organizations/repositories/bun_organization_team_repository.go +++ b/plugins/organizations/repositories/bun_organization_team_repository.go @@ -57,13 +57,18 @@ func (r *BunOrganizationTeamRepository) GetByOrganizationIDAndSlug(ctx context.C return team, err } -func (r *BunOrganizationTeamRepository) GetAllByOrganizationID(ctx context.Context, organizationID string) ([]types.OrganizationTeam, error) { +func (r *BunOrganizationTeamRepository) GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationTeam, int, error) { teams := make([]types.OrganizationTeam, 0) - err := r.db.NewSelect().Model(&teams).Where("organization_id = ?", organizationID).Scan(ctx) + limit = pageLimit(limit) + total, err := r.db.NewSelect().Model(&teams). + Where("organization_id = ?", organizationID). + OrderExpr("created_at DESC, id DESC"). + Offset(pageOffset(page, limit)).Limit(limit). + ScanAndCount(ctx) if err == sql.ErrNoRows { - return []types.OrganizationTeam{}, nil + return []types.OrganizationTeam{}, total, nil } - return teams, err + return teams, total, err } func (r *BunOrganizationTeamRepository) Update(ctx context.Context, team *types.OrganizationTeam) (*types.OrganizationTeam, error) { diff --git a/plugins/organizations/repositories/bun_organization_team_repository_test.go b/plugins/organizations/repositories/bun_organization_team_repository_test.go index 2612e100..df3c9977 100644 --- a/plugins/organizations/repositories/bun_organization_team_repository_test.go +++ b/plugins/organizations/repositories/bun_organization_team_repository_test.go @@ -2,6 +2,7 @@ package repositories_test import ( "context" + "fmt" "testing" "github.com/stretchr/testify/require" @@ -158,58 +159,85 @@ func TestBunOrganizationTeamRepository_GetByOrganizationIDAndSlug(t *testing.T) } } +func seedTeams(t *testing.T, count int) (repositories.OrganizationTeamRepository, context.Context) { + t.Helper() + + db := plugintests.SetupRepoDB(t) + plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") + + repo := repositories.NewBunOrganizationTeamRepository(db) + ctx := context.Background() + for i := 1; i <= count; i++ { + _, err := repo.Create(ctx, &types.OrganizationTeam{ + ID: fmt.Sprintf("team-%d", i), + OrganizationID: "org-1", + Name: fmt.Sprintf("Team %d", i), + Slug: fmt.Sprintf("team-%d", i), + }) + require.NoError(t, err) + } + + return repo, ctx +} + +func teamIDs(teams []types.OrganizationTeam) []string { + ids := make([]string, 0, len(teams)) + for _, team := range teams { + ids = append(ids, team.ID) + } + return ids +} + func TestBunOrganizationTeamRepository_GetAllByOrganizationID(t *testing.T) { t.Parallel() tests := []struct { name string organizationID string - setup func(*testing.T) (repositories.OrganizationTeamRepository, context.Context) + page int + limit int expectCount int + expectTotal int }{ - { - name: "found", - organizationID: "org-1", - expectCount: 2, - setup: func(t *testing.T) (repositories.OrganizationTeamRepository, context.Context) { - t.Helper() - db := plugintests.SetupRepoDB(t) - plugintests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") - repo := repositories.NewBunOrganizationTeamRepository(db) - ctx := context.Background() - - _, err := repo.Create(ctx, &types.OrganizationTeam{ID: "team-1", OrganizationID: "org-1", Name: "Platform", Slug: "platform"}) - require.NoError(t, err) - _, err = repo.Create(ctx, &types.OrganizationTeam{ID: "team-2", OrganizationID: "org-1", Name: "Core", Slug: "core"}) - require.NoError(t, err) - - return repo, ctx - }, - }, - { - name: "empty", - organizationID: "org-2", - expectCount: 0, - setup: func(t *testing.T) (repositories.OrganizationTeamRepository, context.Context) { - t.Helper() - return repositories.NewBunOrganizationTeamRepository(plugintests.SetupRepoDB(t)), context.Background() - }, - }, + {name: "first page", organizationID: "org-1", page: 1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "last partial page", organizationID: "org-1", page: 3, limit: 2, expectCount: 1, expectTotal: 5}, + {name: "page past the end keeps the true total", organizationID: "org-1", page: 4, limit: 2, expectCount: 0, expectTotal: 5}, + {name: "page zero does not error", organizationID: "org-1", page: 0, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "negative page does not error", organizationID: "org-1", page: -1, limit: 2, expectCount: 2, expectTotal: 5}, + {name: "zero limit falls back to a bounded page", organizationID: "org-1", page: 1, limit: 0, expectCount: 5, expectTotal: 5}, + {name: "unknown organization is empty", organizationID: "org-2", page: 1, limit: 10, expectCount: 0, expectTotal: 0}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - repo, ctx := tt.setup(t) + repo, ctx := seedTeams(t, 5) - found, err := repo.GetAllByOrganizationID(ctx, tt.organizationID) + found, total, err := repo.GetAllByOrganizationID(ctx, tt.organizationID, tt.page, tt.limit) require.NoError(t, err) require.Len(t, found, tt.expectCount) + require.Equal(t, tt.expectTotal, total) }) } } +func TestBunOrganizationTeamRepository_GetAllByOrganizationIDPagesPartitionCleanly(t *testing.T) { + t.Parallel() + + repo, ctx := seedTeams(t, 5) + + seen := make([]string, 0, 5) + for page := 1; page <= 3; page++ { + found, total, err := repo.GetAllByOrganizationID(ctx, "org-1", page, 2) + require.NoError(t, err) + require.Equal(t, 5, total) + seen = append(seen, teamIDs(found)...) + } + + require.ElementsMatch(t, []string{"team-1", "team-2", "team-3", "team-4", "team-5"}, seen) +} + func TestBunOrganizationTeamRepository_Update(t *testing.T) { t.Parallel() tests := []struct { diff --git a/plugins/organizations/repositories/interfaces.go b/plugins/organizations/repositories/interfaces.go index 50e08c6e..5a684e00 100644 --- a/plugins/organizations/repositories/interfaces.go +++ b/plugins/organizations/repositories/interfaces.go @@ -12,7 +12,8 @@ type OrganizationRepository interface { Create(ctx context.Context, organization *types.Organization) (*types.Organization, error) GetByID(ctx context.Context, organizationID string) (*types.Organization, error) GetBySlug(ctx context.Context, slug string) (*types.Organization, error) - GetAllByOwnerID(ctx context.Context, ownerID string) ([]types.Organization, error) + GetAllAccessibleByUserID(ctx context.Context, userID string, page int, limit int) ([]types.Organization, int, error) + CountAccessibleByUserID(ctx context.Context, userID string) (int, error) Update(ctx context.Context, organization *types.Organization) (*types.Organization, error) Delete(ctx context.Context, organizationID string) error WithTx(tx bun.IDB) OrganizationRepository @@ -23,9 +24,8 @@ type OrganizationInvitationRepository interface { GetByID(ctx context.Context, invitationID string) (*types.OrganizationInvitation, error) GetByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) GetByOrganizationIDAndEmail(ctx context.Context, organizationID string, email string, status ...types.OrganizationInvitationStatus) (*types.OrganizationInvitation, error) - GetAllByOrganizationID(ctx context.Context, organizationID string) ([]types.OrganizationInvitation, error) - GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) - GetAllPendingByEmail(ctx context.Context, email string) ([]types.OrganizationInvitation, error) + GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string, page int, limit int) ([]types.GetOrganizationInvitationResponse, int, error) + GetAllPendingByEmail(ctx context.Context, email string, limit int) ([]types.OrganizationInvitation, error) Update(ctx context.Context, invitation *types.OrganizationInvitation) (*types.OrganizationInvitation, error) CountByOrganizationIDAndEmail(ctx context.Context, organizationID string, email string) (int, error) WithTx(tx bun.IDB) OrganizationInvitationRepository @@ -34,13 +34,12 @@ type OrganizationInvitationRepository interface { type OrganizationMemberRepository interface { Create(ctx context.Context, member *types.OrganizationMember) (*types.OrganizationMember, error) CountByOrganizationID(ctx context.Context, organizationID string) (int, error) - GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMember, error) - GetAllByOrganizationIDWithUser(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, error) + GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMember, int, error) + GetAllByOrganizationIDWithUser(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, int, error) GetByID(ctx context.Context, memberID string) (*types.OrganizationMember, error) GetByIDWithUser(ctx context.Context, memberID string) (*types.OrganizationMemberResponse, error) GetByOrganizationIDAndUserID(ctx context.Context, organizationID string, userID string) (*types.OrganizationMember, error) GetByOrganizationIDAndUserIDWithUser(ctx context.Context, organizationID string, userID string) (*types.OrganizationMemberResponse, error) - GetAllByUserID(ctx context.Context, userID string) ([]types.OrganizationMember, error) Update(ctx context.Context, member *types.OrganizationMember) (*types.OrganizationMember, error) Delete(ctx context.Context, memberID string) error WithTx(tx bun.IDB) OrganizationMemberRepository @@ -50,7 +49,7 @@ type OrganizationTeamRepository interface { Create(ctx context.Context, team *types.OrganizationTeam) (*types.OrganizationTeam, error) GetByID(ctx context.Context, teamID string) (*types.OrganizationTeam, error) GetByOrganizationIDAndSlug(ctx context.Context, organizationID, slug string) (*types.OrganizationTeam, error) - GetAllByOrganizationID(ctx context.Context, organizationID string) ([]types.OrganizationTeam, error) + GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationTeam, int, error) Update(ctx context.Context, team *types.OrganizationTeam) (*types.OrganizationTeam, error) Delete(ctx context.Context, teamID string) error WithTx(tx bun.IDB) OrganizationTeamRepository @@ -60,8 +59,8 @@ type OrganizationTeamMemberRepository interface { Create(ctx context.Context, teamMember *types.OrganizationTeamMember) (*types.OrganizationTeamMember, error) GetByID(ctx context.Context, teamMemberID string) (*types.OrganizationTeamMember, error) GetByTeamIDAndMemberID(ctx context.Context, teamID, memberID string) (*types.OrganizationTeamMember, error) - GetAllByTeamID(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMember, error) - GetAllByTeamIDWithMemberAndUser(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, error) + GetAllByTeamID(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMember, int, error) + GetAllByTeamIDWithMemberAndUser(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, int, error) GetByIDWithMemberAndUser(ctx context.Context, teamMemberID string) (*types.OrganizationTeamMemberResponse, error) DeleteByTeamIDAndMemberID(ctx context.Context, teamID, memberID string) error WithTx(tx bun.IDB) OrganizationTeamMemberRepository diff --git a/plugins/organizations/repositories/pagination.go b/plugins/organizations/repositories/pagination.go new file mode 100644 index 00000000..487eb229 --- /dev/null +++ b/plugins/organizations/repositories/pagination.go @@ -0,0 +1,24 @@ +package repositories + +import ( + "github.com/Authula/authula/core/pagination" +) + +const MaxPendingInvitationsPerBatch = 500 + +func pageLimit(limit int) int { + if limit <= 0 { + return pagination.DefaultLimit + } + if limit > pagination.MaxLimit { + return pagination.MaxLimit + } + return limit +} + +func pageOffset(page, limit int) int { + if page < 1 || limit <= 0 { + return 0 + } + return (page - 1) * limit +} diff --git a/plugins/organizations/repositories/pagination_internal_test.go b/plugins/organizations/repositories/pagination_internal_test.go new file mode 100644 index 00000000..725030ab --- /dev/null +++ b/plugins/organizations/repositories/pagination_internal_test.go @@ -0,0 +1,58 @@ +package repositories + +import ( + "testing" + + "github.com/stretchr/testify/require" + + "github.com/Authula/authula/core/pagination" +) + +func TestPageLimit(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + limit int + expected int + }{ + {name: "positive limit is preserved", limit: 25, expected: 25}, + {name: "zero limit falls back to the default", limit: 0, expected: pagination.DefaultLimit}, + {name: "negative limit falls back to the default", limit: -1, expected: pagination.DefaultLimit}, + {name: "limit at the maximum is preserved", limit: pagination.MaxLimit, expected: pagination.MaxLimit}, + {name: "limit above the maximum is capped", limit: 100000, expected: pagination.MaxLimit}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + require.Equal(t, tt.expected, pageLimit(tt.limit)) + }) + } +} + +func TestPageOffset(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + page int + limit int + expected int + }{ + {name: "first page starts at zero", page: 1, limit: 10, expected: 0}, + {name: "third page skips two pages", page: 3, limit: 10, expected: 20}, + {name: "page zero never produces a negative offset", page: 0, limit: 10, expected: 0}, + {name: "negative page never produces a negative offset", page: -5, limit: 10, expected: 0}, + {name: "non positive limit produces no offset", page: 3, limit: 0, expected: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + require.Equal(t, tt.expected, pageOffset(tt.page, tt.limit)) + }) + } +} diff --git a/plugins/organizations/routes.go b/plugins/organizations/routes.go index bf91e023..5d32bd27 100644 --- a/plugins/organizations/routes.go +++ b/plugins/organizations/routes.go @@ -10,7 +10,7 @@ import ( func Routes(plugin *OrganizationsPlugin) []models.Route { createOrganizationHandler := &handlers.CreateOrganizationHandler{UseCases: plugin.useCases} - getAllOrganizationsByOwnerHandler := &handlers.GetAllOrganizationsByOwnerHandler{UseCases: plugin.useCases} + getAllOrganizationsHandler := &handlers.GetAllOrganizationsHandler{UseCases: plugin.useCases} getOrganizationByIDHandler := &handlers.GetOrganizationByIDHandler{UseCases: plugin.useCases} updateOrganizationHandler := &handlers.UpdateOrganizationHandler{UseCases: plugin.useCases} deleteOrganizationHandler := &handlers.DeleteOrganizationHandler{UseCases: plugin.useCases} @@ -56,7 +56,7 @@ func Routes(plugin *OrganizationsPlugin) []models.Route { Middleware: []func(http.Handler) http.Handler{ middleware.RequireAuthenticated(), }, - Handler: getAllOrganizationsByOwnerHandler.Handle(), + Handler: getAllOrganizationsHandler.Handle(), }, { Method: http.MethodGet, diff --git a/plugins/organizations/services/interfaces.go b/plugins/organizations/services/interfaces.go index 5900b632..cea5fb2c 100644 --- a/plugins/organizations/services/interfaces.go +++ b/plugins/organizations/services/interfaces.go @@ -3,13 +3,14 @@ package services import ( "context" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" "github.com/Authula/authula/plugins/organizations/types" ) type OrganizationService interface { CreateOrganization(ctx context.Context, actor *models.Actor, request types.CreateOrganizationRequest) (*types.Organization, error) - GetAllOrganizationsByOwner(ctx context.Context, actor *models.Actor) ([]types.Organization, error) + GetAllOrganizations(ctx context.Context, actor *models.Actor, params pagination.Params) (*types.ListOrganizationsResponse, error) GetOrganizationByID(ctx context.Context, actor *models.Actor, organizationID string) (*types.Organization, error) UpdateOrganization(ctx context.Context, actor *models.Actor, organizationID string, request types.UpdateOrganizationRequest) (*types.Organization, error) DeleteOrganization(ctx context.Context, actor *models.Actor, organizationID string) error @@ -21,8 +22,7 @@ type OrganizationInvitationService interface { GetOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) GetOrganizationInvitationByID(ctx context.Context, invitationID string) (*types.OrganizationInvitation, error) GetOrganizationInvitationByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) - GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, error) - GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) + GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string, params pagination.Params) (*types.ListOrganizationInvitationsResponse, error) RevokeOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) AcceptOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) RejectOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) @@ -30,7 +30,7 @@ type OrganizationInvitationService interface { type OrganizationMemberService interface { AddMember(ctx context.Context, actor *models.Actor, organizationID string, request types.AddOrganizationMemberRequest) (*types.OrganizationMember, error) - GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, error) + GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationMembersResponse, error) GetMember(ctx context.Context, actor *models.Actor, organizationID string, memberID string) (*types.OrganizationMemberResponse, error) GetMemberByUserID(ctx context.Context, actor *models.Actor, organizationID string, userID string) (*types.OrganizationMemberResponse, error) UpdateMember(ctx context.Context, actor *models.Actor, organizationID string, memberID string, request types.UpdateOrganizationMemberRequest) (*types.OrganizationMember, error) @@ -39,7 +39,7 @@ type OrganizationMemberService interface { type OrganizationTeamService interface { CreateTeam(ctx context.Context, actor *models.Actor, organizationID string, request types.CreateOrganizationTeamRequest) (*types.OrganizationTeam, error) - GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationTeam, error) + GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationTeamsResponse, error) GetTeam(ctx context.Context, actor *models.Actor, organizationID string, teamID string) (*types.OrganizationTeam, error) UpdateTeam(ctx context.Context, actor *models.Actor, organizationID string, teamID string, request types.UpdateOrganizationTeamRequest) (*types.OrganizationTeam, error) DeleteTeam(ctx context.Context, actor *models.Actor, organizationID string, teamID string) error @@ -47,7 +47,7 @@ type OrganizationTeamService interface { type OrganizationTeamMemberService interface { AddTeamMember(ctx context.Context, actor *models.Actor, organizationID string, teamID string, request types.AddOrganizationTeamMemberRequest) (*types.OrganizationTeamMember, error) - GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, error) + GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, params pagination.Params) (*types.ListOrganizationTeamMembersResponse, error) GetTeamMember(ctx context.Context, actor *models.Actor, organizationID string, teamID string, memberID string) (*types.OrganizationTeamMemberResponse, error) RemoveTeamMember(ctx context.Context, actor *models.Actor, organizationID string, teamID string, memberID string) error } diff --git a/plugins/organizations/services/organization_invitation_service.go b/plugins/organizations/services/organization_invitation_service.go index de69ce2c..ec8a9c47 100644 --- a/plugins/organizations/services/organization_invitation_service.go +++ b/plugins/organizations/services/organization_invitation_service.go @@ -13,6 +13,7 @@ import ( emailconstants "github.com/Authula/authula/core/email/constants" emailtmpl "github.com/Authula/authula/core/email/template" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" "github.com/Authula/authula/plugins/organizations/repositories" @@ -256,23 +257,6 @@ func (s *organizationInvitationService) publishOrganizationInvitationCreatedEven }) } -func (s *organizationInvitationService) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, error) { - if actor == nil || actor.ID == "" || organizationID == "" { - return nil, coreerrors.ErrUnauthorized - } - - if _, _, err := s.serviceUtils.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { - return nil, err - } - - invitations, err := s.orgInvitationRepo.GetAllByOrganizationID(ctx, organizationID) - if err != nil { - return nil, err - } - - return invitations, nil -} - func (s *organizationInvitationService) GetOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { if actor == nil || actor.ID == "" || organizationID == "" || invitationID == "" { return nil, coreerrors.ErrUnauthorized @@ -325,17 +309,25 @@ func (s *organizationInvitationService) GetOrganizationInvitationByIDWithOrg(ctx return resp, nil } -func (s *organizationInvitationService) GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { +func (s *organizationInvitationService) GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string, params pagination.Params) (*types.ListOrganizationInvitationsResponse, error) { if organizationID == "" { return nil, coreerrors.ErrNotFound } - resp, err := s.orgInvitationRepo.GetAllByOrganizationIDWithOrg(ctx, organizationID) + params = pagination.Clamp(params) + + invitations, total, err := s.orgInvitationRepo.GetAllByOrganizationIDWithOrg(ctx, organizationID, params.Page, params.Limit) if err != nil { return nil, err } + if invitations == nil { + invitations = []types.GetOrganizationInvitationResponse{} + } - return resp, nil + return &types.ListOrganizationInvitationsResponse{ + Data: invitations, + Pagination: pagination.New(params.Page, params.Limit, total), + }, nil } func (s *organizationInvitationService) RevokeOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { @@ -502,13 +494,20 @@ func (s *organizationInvitationService) AcceptPendingOrganizationInvitationsForE } } - pendingInvitations, err := s.orgInvitationRepo.GetAllPendingByEmail(ctx, email) + pendingInvitations, err := s.orgInvitationRepo.GetAllPendingByEmail(ctx, email, repositories.MaxPendingInvitationsPerBatch) if err != nil { return nil, err } if len(pendingInvitations) == 0 { return []types.OrganizationInvitation{}, nil } + if len(pendingInvitations) == repositories.MaxPendingInvitationsPerBatch { + s.logger.Warn( + "pending organization invitation batch hit the per-call cap; remaining invitations will be accepted on the next call", + "email", email, + "limit", repositories.MaxPendingInvitationsPerBatch, + ) + } return s.acceptOrganizationInvitations(ctx, userID, pendingInvitations) } diff --git a/plugins/organizations/services/organization_invitation_service_test.go b/plugins/organizations/services/organization_invitation_service_test.go index 67b19bf1..4b98bff1 100644 --- a/plugins/organizations/services/organization_invitation_service_test.go +++ b/plugins/organizations/services/organization_invitation_service_test.go @@ -13,6 +13,7 @@ import ( emailtmpl "github.com/Authula/authula/core/email/template" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" @@ -720,93 +721,75 @@ func TestOrganizationInvitationService_GetOrganizationInvitation(t *testing.T) { } } -func TestOrganizationInvitationService_GetAllOrganizationInvitations(t *testing.T) { +func TestOrganizationInvitationService_GetAllOrganizationInvitationsByOrgIDWithOrg(t *testing.T) { t.Parallel() repoErr := errors.New("repository error") tests := []struct { - name string - actorUserID string - organizationID string - setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationInvitationRepository, *orgtests.MockOrganizationMemberRepository) - expectErr error - expectLen int + name string + organizationID string + params pagination.Params + setup func(*orgtests.MockOrganizationInvitationRepository) + expectErr error + expectLen int + expectPagination pagination.Pagination }{ { - name: "success", - actorUserID: "user-1", - organizationID: "org-1", - setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { - orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() - invRepo.On("GetAllByOrganizationID", mock.Anything, "org-1").Return([]types.OrganizationInvitation{{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}}, nil).Once() - }, - expectLen: 1, - }, - { - name: "rejected invitation is returned", - actorUserID: "user-1", + name: "returns the requested page with its metadata", organizationID: "org-1", - setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { - orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() - invRepo.On("GetAllByOrganizationID", mock.Anything, "org-1").Return([]types.OrganizationInvitation{{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: types.OrganizationInvitationStatusRejected, ExpiresAt: time.Now().UTC().Add(time.Hour)}}, nil).Once() - }, - expectLen: 1, - }, - { - name: "org member can list", - actorUserID: "user-2", - organizationID: "org-1", - setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { - orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "owner-1"}, nil).Once() - memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-2").Return(&types.OrganizationMember{ID: "mem-1", OrganizationID: "org-1", UserID: "user-2", Role: "member"}, nil).Once() - invRepo.On("GetAllByOrganizationID", mock.Anything, "org-1").Return([]types.OrganizationInvitation{{ID: "inv-1", OrganizationID: "org-1", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}}, nil).Once() + params: pagination.Params{Page: 2, Limit: 2}, + setup: func(invRepo *orgtests.MockOrganizationInvitationRepository) { + invRepo.On("GetAllByOrganizationIDWithOrg", mock.Anything, "org-1", 2, 2). + Return([]types.GetOrganizationInvitationResponse{ + {Invitation: &types.OrganizationInvitation{ID: "inv-3", OrganizationID: "org-1"}}, + {Invitation: &types.OrganizationInvitation{ID: "inv-4", OrganizationID: "org-1"}}, + }, 5, nil).Once() }, - expectLen: 1, + expectLen: 2, + expectPagination: pagination.Pagination{Page: 2, Limit: 2, Total: 5, TotalPages: 3, HasMore: true}, }, - {name: "unauthorized", actorUserID: "", organizationID: "org-1", expectErr: coreerrors.ErrUnauthorized}, { - name: "organization not found", - actorUserID: "user-1", + name: "out of range params are clamped before reaching the repository", organizationID: "org-1", - setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { - orgRepo.On("GetByID", mock.Anything, "org-1").Return(nil, nil).Once() + params: pagination.Params{Page: -4, Limit: 5000}, + setup: func(invRepo *orgtests.MockOrganizationInvitationRepository) { + invRepo.On("GetAllByOrganizationIDWithOrg", mock.Anything, "org-1", 1, pagination.MaxLimit). + Return([]types.GetOrganizationInvitationResponse{}, 0, nil).Once() }, - expectErr: coreerrors.ErrNotFound, + expectLen: 0, + expectPagination: pagination.Pagination{Page: 1, Limit: pagination.MaxLimit, Total: 0, TotalPages: 0, HasMore: false}, }, { - name: "organization lookup error", - actorUserID: "user-1", + name: "nil result is normalised to an empty slice", organizationID: "org-1", - setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { - orgRepo.On("GetByID", mock.Anything, "org-1").Return((*types.Organization)(nil), repoErr).Once() + params: pagination.Params{Page: 1, Limit: 10}, + setup: func(invRepo *orgtests.MockOrganizationInvitationRepository) { + invRepo.On("GetAllByOrganizationIDWithOrg", mock.Anything, "org-1", 1, 10). + Return(([]types.GetOrganizationInvitationResponse)(nil), 0, nil).Once() }, - expectErr: repoErr, + expectLen: 0, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: 0, TotalPages: 0, HasMore: false}, }, { - name: "forbidden", - actorUserID: "user-1", - organizationID: "org-1", - setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { - orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "owner-1"}, nil).Once() - memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-1").Return(nil, nil).Once() - }, - expectErr: coreerrors.ErrForbidden, + name: "missing organization id is not found", + organizationID: "", + params: pagination.Params{Page: 1, Limit: 10}, + expectErr: coreerrors.ErrNotFound, }, { - name: "repo error", - actorUserID: "user-1", + name: "repository error is propagated", organizationID: "org-1", - setup: func(orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { - orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() - invRepo.On("GetAllByOrganizationID", mock.Anything, "org-1").Return(([]types.OrganizationInvitation)(nil), repoErr).Once() + params: pagination.Params{Page: 1, Limit: 10}, + setup: func(invRepo *orgtests.MockOrganizationInvitationRepository) { + invRepo.On("GetAllByOrganizationIDWithOrg", mock.Anything, "org-1", 1, 10). + Return(([]types.GetOrganizationInvitationResponse)(nil), 0, repoErr).Once() }, expectErr: repoErr, }, } for _, tt := range tests { - tt := tt t.Run(tt.name, func(t *testing.T) { t.Parallel() @@ -818,29 +801,27 @@ func TestOrganizationInvitationService_GetAllOrganizationInvitations(t *testing. invRepo := &orgtests.MockOrganizationInvitationRepository{} memberRepo := &orgtests.MockOrganizationMemberRepository{} if tt.setup != nil { - tt.setup(orgRepo, invRepo, memberRepo) + tt.setup(invRepo) } - expectActorMember(memberRepo, tt.organizationID, tt.actorUserID) svc := newTestOrganizationInvitationService(&orgtests.MockOrganizationInvitationTxRunner{}, pluginConfig, &internaltests.MockUserService{}, orgtests.NewAccessControlServiceStub(), orgRepo, invRepo, memberRepo) - invitations, err := svc.GetAllOrganizationInvitations(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID) + resp, err := svc.GetAllOrganizationInvitationsByOrgIDWithOrg(context.Background(), tt.organizationID, tt.params) if tt.expectErr != nil { require.Error(t, err) require.ErrorIs(t, err, tt.expectErr) - require.True(t, orgRepo.AssertExpectations(t)) require.True(t, invRepo.AssertExpectations(t)) - require.True(t, memberRepo.AssertExpectations(t)) return } + require.NoError(t, err) - require.Len(t, invitations, tt.expectLen) - require.True(t, orgRepo.AssertExpectations(t)) + require.NotNil(t, resp) + require.NotNil(t, resp.Data) + require.Len(t, resp.Data, tt.expectLen) + require.Equal(t, tt.expectPagination, resp.Pagination) require.True(t, invRepo.AssertExpectations(t)) - require.True(t, memberRepo.AssertExpectations(t)) }) } } - func TestOrganizationInvitationService_RevokeOrganizationInvitation(t *testing.T) { t.Parallel() @@ -974,7 +955,7 @@ func TestOrganizationInvitationService_AcceptPendingOrganizationInvitationsForEm userID: "user-2", email: "USER@EXAMPLE.COM", setup: func(userSvc *internaltests.MockUserService, orgRepo *orgtests.MockOrganizationRepository, invRepo *orgtests.MockOrganizationInvitationRepository, memberRepo *orgtests.MockOrganizationMemberRepository, hooks *orgtests.MockOrganizationInvitationHooks, memberHooks *orgtests.MockOrganizationMemberHooks) { - invRepo.On("GetAllPendingByEmail", mock.Anything, "user@example.com").Return([]types.OrganizationInvitation{{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}}, nil).Once() + invRepo.On("GetAllPendingByEmail", mock.Anything, "user@example.com", repositories.MaxPendingInvitationsPerBatch).Return([]types.OrganizationInvitation{{ID: "inv-1", OrganizationID: "org-1", Email: "user@example.com", Role: "member", Status: types.OrganizationInvitationStatusPending, ExpiresAt: time.Now().UTC().Add(time.Hour)}}, nil).Once() memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-2").Return(nil, nil).Once() memberRepo.On("Create", mock.Anything, mock.MatchedBy(func(member *types.OrganizationMember) bool { return member != nil && member.OrganizationID == "org-1" && member.UserID == "user-2" && member.Role == "member" diff --git a/plugins/organizations/services/organization_member_service.go b/plugins/organizations/services/organization_member_service.go index d9849fc9..0df6c333 100644 --- a/plugins/organizations/services/organization_member_service.go +++ b/plugins/organizations/services/organization_member_service.go @@ -8,6 +8,7 @@ import ( "github.com/uptrace/bun" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" "github.com/Authula/authula/plugins/organizations/constants" "github.com/Authula/authula/plugins/organizations/repositories" @@ -114,12 +115,25 @@ func (s *organizationMemberService) AddMember(ctx context.Context, actor *models return created, nil } -func (s *organizationMemberService) GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, error) { +func (s *organizationMemberService) GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationMembersResponse, error) { if _, _, err := s.serviceUtils.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { return nil, err } - return s.orgMemberRepo.GetAllByOrganizationIDWithUser(ctx, organizationID, page, limit) + params = pagination.Clamp(params) + + members, total, err := s.orgMemberRepo.GetAllByOrganizationIDWithUser(ctx, organizationID, params.Page, params.Limit) + if err != nil { + return nil, err + } + if members == nil { + members = []types.OrganizationMemberResponse{} + } + + return &types.ListOrganizationMembersResponse{ + Data: members, + Pagination: pagination.New(params.Page, params.Limit, total), + }, nil } func (s *organizationMemberService) GetMember(ctx context.Context, actor *models.Actor, organizationID string, memberID string) (*types.OrganizationMemberResponse, error) { diff --git a/plugins/organizations/services/organization_member_service_test.go b/plugins/organizations/services/organization_member_service_test.go index 06f2eba8..bda1a073 100644 --- a/plugins/organizations/services/organization_member_service_test.go +++ b/plugins/organizations/services/organization_member_service_test.go @@ -10,6 +10,7 @@ import ( "github.com/stretchr/testify/require" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" @@ -309,23 +310,27 @@ func TestOrganizationMemberService_GetAllMembers(t *testing.T) { repoErr := errors.New("repository error") tests := []struct { - name string - actorUserID string - organizationID string - setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationMemberRepository) - expectErr error - expectLen int + name string + actorUserID string + organizationID string + params pagination.Params + setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationMemberRepository) + expectErr error + expectLen int + expectPagination pagination.Pagination }{ { name: "unauthorized", actorUserID: "", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, expectErr: coreerrors.ErrUnauthorized, }, { name: "organization not found", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(nil, nil).Once() }, @@ -335,6 +340,7 @@ func TestOrganizationMemberService_GetAllMembers(t *testing.T) { name: "forbidden", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "owner-1"}, nil).Once() memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-1").Return(nil, nil).Once() @@ -345,9 +351,10 @@ func TestOrganizationMemberService_GetAllMembers(t *testing.T) { name: "repository error", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() - memberRepo.On("GetAllByOrganizationIDWithUser", mock.Anything, "org-1", 1, 10).Return(nil, repoErr).Once() + memberRepo.On("GetAllByOrganizationIDWithUser", mock.Anything, "org-1", 1, 10).Return(nil, 0, repoErr).Once() }, expectErr: repoErr, }, @@ -355,11 +362,40 @@ func TestOrganizationMemberService_GetAllMembers(t *testing.T) { name: "success", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, + setup: func(orgRepo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { + orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() + memberRepo.On("GetAllByOrganizationIDWithUser", mock.Anything, "org-1", 1, 10). + Return([]types.OrganizationMemberResponse{{ID: "mem-1", OrganizationID: "org-1", Role: "member"}}, 25, nil).Once() + }, + expectLen: 1, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: 25, TotalPages: 3, HasMore: true}, + }, + { + name: "out of range params are clamped before reaching the repository", + actorUserID: "user-1", + organizationID: "org-1", + params: pagination.Params{Page: -4, Limit: 5000}, setup: func(orgRepo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() - memberRepo.On("GetAllByOrganizationIDWithUser", mock.Anything, "org-1", 1, 10).Return([]types.OrganizationMemberResponse{{ID: "mem-1", OrganizationID: "org-1", Role: "member"}}, nil).Once() + memberRepo.On("GetAllByOrganizationIDWithUser", mock.Anything, "org-1", 1, pagination.MaxLimit). + Return([]types.OrganizationMemberResponse{}, 0, nil).Once() }, - expectLen: 1, + expectLen: 0, + expectPagination: pagination.Pagination{Page: 1, Limit: pagination.MaxLimit, Total: 0, TotalPages: 0, HasMore: false}, + }, + { + name: "nil result is normalised to an empty slice", + actorUserID: "user-1", + organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, + setup: func(orgRepo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { + orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() + memberRepo.On("GetAllByOrganizationIDWithUser", mock.Anything, "org-1", 1, 10). + Return(([]types.OrganizationMemberResponse)(nil), 0, nil).Once() + }, + expectLen: 0, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: 0, TotalPages: 0, HasMore: false}, }, } @@ -376,7 +412,7 @@ func TestOrganizationMemberService_GetAllMembers(t *testing.T) { expectActorMember(memberRepo, tt.organizationID, tt.actorUserID) svc := newTestOrganizationMemberService(userService, orgtests.NewAccessControlServiceStub(), orgRepo, memberRepo, nil) - members, err := svc.GetAllMembers(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, 1, 10) + resp, err := svc.GetAllMembers(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.params) if tt.expectErr != nil { require.Error(t, err) require.ErrorIs(t, err, tt.expectErr) @@ -386,12 +422,14 @@ func TestOrganizationMemberService_GetAllMembers(t *testing.T) { } return } + require.NoError(t, err) - require.Len(t, members, tt.expectLen) - if tt.setup != nil { - require.True(t, orgRepo.AssertExpectations(t)) - require.True(t, memberRepo.AssertExpectations(t)) - } + require.NotNil(t, resp) + require.NotNil(t, resp.Data) + require.Len(t, resp.Data, tt.expectLen) + require.Equal(t, tt.expectPagination, resp.Pagination) + require.True(t, orgRepo.AssertExpectations(t)) + require.True(t, memberRepo.AssertExpectations(t)) }) } } diff --git a/plugins/organizations/services/organization_service.go b/plugins/organizations/services/organization_service.go index e081f0cd..6ccc8c2f 100644 --- a/plugins/organizations/services/organization_service.go +++ b/plugins/organizations/services/organization_service.go @@ -3,14 +3,13 @@ package services import ( "context" "database/sql" - "sort" "strings" - "sync" "unicode" "github.com/uptrace/bun" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" "github.com/Authula/authula/plugins/organizations/constants" "github.com/Authula/authula/plugins/organizations/repositories" @@ -118,7 +117,7 @@ func (s *organizationService) CreateOrganization(ctx context.Context, actor *mod var created *types.Organization createFn := func(ctx context.Context, orgRepo repositories.OrganizationRepository, memberRepo repositories.OrganizationMemberRepository) error { - if err := s.ensureOrganizationLimit(ctx, actor, orgRepo, memberRepo); err != nil { + if err := s.ensureOrganizationLimit(ctx, actor, orgRepo); err != nil { return err } @@ -159,7 +158,7 @@ func (s *organizationService) CreateOrganization(ctx context.Context, actor *mod return created, nil } -func (s *organizationService) ensureOrganizationLimit(ctx context.Context, actor *models.Actor, orgRepo repositories.OrganizationRepository, memberRepo repositories.OrganizationMemberRepository) error { +func (s *organizationService) ensureOrganizationLimit(ctx context.Context, actor *models.Actor, orgRepo repositories.OrganizationRepository) error { if actor == nil || actor.ID == "" { return coreerrors.ErrUnauthorized } @@ -168,104 +167,37 @@ func (s *organizationService) ensureOrganizationLimit(ctx context.Context, actor return nil } - var ( - ownedOrganizations []types.Organization - memberRecords []types.OrganizationMember - ownedOrganizationsErr error - memberRecordsErr error - ) - - wg := sync.WaitGroup{} - - wg.Go(func() { - ownedOrganizations, ownedOrganizationsErr = orgRepo.GetAllByOwnerID(ctx, actor.ID) - }) - - wg.Go(func() { - memberRecords, memberRecordsErr = memberRepo.GetAllByUserID(ctx, actor.ID) - }) - - wg.Wait() - - if ownedOrganizationsErr != nil { - return ownedOrganizationsErr - } - if memberRecordsErr != nil { - return memberRecordsErr - } - - organizationIDs := make(map[string]struct{}, len(ownedOrganizations)+len(memberRecords)) - for _, organization := range ownedOrganizations { - if organization.ID == "" { - continue - } - organizationIDs[organization.ID] = struct{}{} - } - for _, member := range memberRecords { - if member.OrganizationID == "" { - continue - } - organizationIDs[member.OrganizationID] = struct{}{} + organizationCount, err := orgRepo.CountAccessibleByUserID(ctx, actor.ID) + if err != nil { + return err } - if len(organizationIDs) >= *s.organizationsLimit { + if organizationCount >= *s.organizationsLimit { return constants.ErrOrganizationsQuotaExceeded } return nil } -func (s *organizationService) GetAllOrganizationsByOwner(ctx context.Context, actor *models.Actor) ([]types.Organization, error) { +func (s *organizationService) GetAllOrganizations(ctx context.Context, actor *models.Actor, params pagination.Params) (*types.ListOrganizationsResponse, error) { if actor == nil || actor.ID == "" { return nil, coreerrors.ErrUnauthorized } - ownedOrganizations, err := s.orgRepo.GetAllByOwnerID(ctx, actor.ID) - if err != nil { - return nil, err - } + params = pagination.Clamp(params) - memberRecords, err := s.orgMemberRepo.GetAllByUserID(ctx, actor.ID) + organizations, total, err := s.orgRepo.GetAllAccessibleByUserID(ctx, actor.ID, params.Page, params.Limit) if err != nil { return nil, err } - - organizationMap := make(map[string]types.Organization, len(ownedOrganizations)) - for _, organization := range ownedOrganizations { - organizationMap[organization.ID] = organization - } - - for _, member := range memberRecords { - if member.OrganizationID == "" { - continue - } - if _, exists := organizationMap[member.OrganizationID]; exists { - continue - } - - organization, err := s.orgRepo.GetByID(ctx, member.OrganizationID) - if err != nil { - return nil, err - } - if organization == nil { - continue - } - - organizationMap[organization.ID] = *organization - } - - organizationIDs := make([]string, 0, len(organizationMap)) - for organizationID := range organizationMap { - organizationIDs = append(organizationIDs, organizationID) - } - sort.Strings(organizationIDs) - - organizations := make([]types.Organization, 0, len(organizationIDs)) - for _, organizationID := range organizationIDs { - organizations = append(organizations, organizationMap[organizationID]) + if organizations == nil { + organizations = []types.Organization{} } - return organizations, nil + return &types.ListOrganizationsResponse{ + Data: organizations, + Pagination: pagination.New(params.Page, params.Limit, total), + }, nil } func (s *organizationService) GetOrganizationByID(ctx context.Context, actor *models.Actor, organizationID string) (*types.Organization, error) { diff --git a/plugins/organizations/services/organization_service_test.go b/plugins/organizations/services/organization_service_test.go index 91ac810c..def5f7ab 100644 --- a/plugins/organizations/services/organization_service_test.go +++ b/plugins/organizations/services/organization_service_test.go @@ -2,12 +2,14 @@ package services import ( "context" + "errors" "testing" "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" internaltests "github.com/Authula/authula/internal/tests" "github.com/Authula/authula/plugins/organizations/constants" orgtests "github.com/Authula/authula/plugins/organizations/tests" @@ -33,8 +35,7 @@ func TestOrganizationService_CreateOrganization(t *testing.T) { } limitSuccessSetup := func(repo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository, hooks *orgtests.MockOrganizationHooks, serviceUtils *ServiceUtils) { - repo.On("GetAllByOwnerID", mock.Anything, "user-1").Return([]types.Organization{{ID: "org-owned", OwnerID: "user-1", Name: "Owned Org"}}, nil).Once() - memberRepo.On("GetAllByUserID", mock.Anything, "user-1").Return([]types.OrganizationMember{{ID: "mem-owned", OrganizationID: "org-owned", UserID: "user-1", Role: "member"}, {ID: "mem-2", OrganizationID: "org-member", UserID: "user-1", Role: "member"}}, nil).Once() + repo.On("CountAccessibleByUserID", mock.Anything, "user-1").Return(2, nil).Once() repo.On("Create", mock.Anything, mock.MatchedBy(func(org *types.Organization) bool { return org != nil && org.OwnerID == "user-1" && org.Name == "Acme Labs" && org.Slug == "acme-labs" && string(internaltests.MarshalToJSON(t, org.Metadata)) == "{}" })).Return(&types.Organization{ID: "org-2", OwnerID: "user-1", Name: "Acme Labs", Slug: "acme-labs", Metadata: map[string]any{}}, nil).Once() @@ -44,8 +45,7 @@ func TestOrganizationService_CreateOrganization(t *testing.T) { } quotaExceededSetup := func(repo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository, hooks *orgtests.MockOrganizationHooks, serviceUtils *ServiceUtils) { - repo.On("GetAllByOwnerID", mock.Anything, "user-1").Return([]types.Organization{{ID: "org-owned", OwnerID: "user-1", Name: "Owned Org"}}, nil).Once() - memberRepo.On("GetAllByUserID", mock.Anything, "user-1").Return([]types.OrganizationMember{{ID: "mem-owned", OrganizationID: "org-owned", UserID: "user-1", Role: "member"}, {ID: "mem-2", OrganizationID: "org-member", UserID: "user-1", Role: "member"}}, nil).Once() + repo.On("CountAccessibleByUserID", mock.Anything, "user-1").Return(2, nil).Once() } tests := []struct { @@ -223,30 +223,71 @@ func TestOrganizationService_CreateOrganizationRequiresPrivilegedRole(t *testing } } -func TestOrganizationService_GetAllOrganizationsByOwner(t *testing.T) { +func TestOrganizationService_GetAllOrganizations(t *testing.T) { t.Parallel() + repoErr := errors.New("repository error") + tests := []struct { - name string - actorUserID string - setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationMemberRepository) - expectErr error - expectLen int + name string + actorUserID string + params pagination.Params + setup func(*orgtests.MockOrganizationRepository) + expectErr error + expectLen int + expectPagination pagination.Pagination }{ { name: "unauthorized", actorUserID: "", + params: pagination.Params{Page: 1, Limit: 10}, expectErr: coreerrors.ErrUnauthorized, }, { - name: "success", + name: "returns the requested page with its metadata", actorUserID: "user-1", - setup: func(repo *orgtests.MockOrganizationRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { - repo.On("GetAllByOwnerID", mock.Anything, "user-1").Return([]types.Organization{{ID: "org-1", OwnerID: "user-1", Name: "Acme"}}, nil).Once() - memberRepo.On("GetAllByUserID", mock.Anything, "user-1").Return([]types.OrganizationMember{{ID: "mem-1", OrganizationID: "org-2", UserID: "user-1", Role: "member"}}, nil).Once() - repo.On("GetByID", mock.Anything, "org-2").Return(&types.Organization{ID: "org-2", OwnerID: "owner-2", Name: "Platform"}, nil).Once() + params: pagination.Params{Page: 1, Limit: 10}, + setup: func(repo *orgtests.MockOrganizationRepository) { + repo.On("GetAllAccessibleByUserID", mock.Anything, "user-1", 1, 10). + Return([]types.Organization{ + {ID: "org-1", OwnerID: "user-1", Name: "Acme"}, + {ID: "org-2", OwnerID: "owner-2", Name: "Platform"}, + }, 2, nil).Once() + }, + expectLen: 2, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: 2, TotalPages: 1, HasMore: false}, + }, + { + name: "out of range params are clamped before reaching the repository", + actorUserID: "user-1", + params: pagination.Params{Page: -4, Limit: 5000}, + setup: func(repo *orgtests.MockOrganizationRepository) { + repo.On("GetAllAccessibleByUserID", mock.Anything, "user-1", 1, pagination.MaxLimit). + Return([]types.Organization{}, 0, nil).Once() + }, + expectLen: 0, + expectPagination: pagination.Pagination{Page: 1, Limit: pagination.MaxLimit, Total: 0, TotalPages: 0, HasMore: false}, + }, + { + name: "nil result is normalised to an empty slice", + actorUserID: "user-1", + params: pagination.Params{Page: 1, Limit: 10}, + setup: func(repo *orgtests.MockOrganizationRepository) { + repo.On("GetAllAccessibleByUserID", mock.Anything, "user-1", 1, 10). + Return(([]types.Organization)(nil), 0, nil).Once() + }, + expectLen: 0, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: 0, TotalPages: 0, HasMore: false}, + }, + { + name: "repository error is propagated", + actorUserID: "user-1", + params: pagination.Params{Page: 1, Limit: 10}, + setup: func(repo *orgtests.MockOrganizationRepository) { + repo.On("GetAllAccessibleByUserID", mock.Anything, "user-1", 1, 10). + Return(([]types.Organization)(nil), 0, repoErr).Once() }, - expectLen: 2, + expectErr: repoErr, }, } @@ -257,25 +298,114 @@ func TestOrganizationService_GetAllOrganizationsByOwner(t *testing.T) { repo := &orgtests.MockOrganizationRepository{} memberRepo := &orgtests.MockOrganizationMemberRepository{} if tt.setup != nil { - tt.setup(repo, memberRepo) + tt.setup(repo) } serviceUtils := &ServiceUtils{orgRepo: repo, orgMemberRepo: memberRepo} svc := NewOrganizationService(repo, memberRepo, serviceUtils, nil, nil, nil) - organizations, err := svc.GetAllOrganizationsByOwner(context.Background(), orgtests.Actor(tt.actorUserID)) + resp, err := svc.GetAllOrganizations(context.Background(), orgtests.Actor(tt.actorUserID), tt.params) if tt.expectErr != nil { require.Error(t, err) require.ErrorIs(t, err, tt.expectErr) + require.True(t, repo.AssertExpectations(t)) return } + require.NoError(t, err) - require.Len(t, organizations, tt.expectLen) + require.NotNil(t, resp) + require.NotNil(t, resp.Data) + require.Len(t, resp.Data, tt.expectLen) + require.Equal(t, tt.expectPagination, resp.Pagination) require.True(t, repo.AssertExpectations(t)) require.True(t, memberRepo.AssertExpectations(t)) }) } } +// The organization quota must be enforced from a SQL count: the previous +// implementation counted a fetched list, which a page size would silently cap. +func TestOrganizationService_EnsureOrganizationLimit(t *testing.T) { + t.Parallel() + + countErr := errors.New("count error") + limit := 3 + + tests := []struct { + name string + actorUserID string + organizationLimit *int + setup func(*orgtests.MockOrganizationRepository) + expectErr error + expectNoCount bool + }{ + { + name: "unauthorized actor is rejected before counting", + actorUserID: "", + organizationLimit: &limit, + expectErr: coreerrors.ErrUnauthorized, + expectNoCount: true, + }, + { + name: "no configured limit skips the count entirely", + actorUserID: "user-1", + expectNoCount: true, + }, + { + name: "count one below the limit is allowed", + actorUserID: "user-1", + organizationLimit: &limit, + setup: func(repo *orgtests.MockOrganizationRepository) { + repo.On("CountAccessibleByUserID", mock.Anything, "user-1").Return(limit-1, nil).Once() + }, + }, + { + name: "count at the limit is rejected", + actorUserID: "user-1", + organizationLimit: &limit, + setup: func(repo *orgtests.MockOrganizationRepository) { + repo.On("CountAccessibleByUserID", mock.Anything, "user-1").Return(limit, nil).Once() + }, + expectErr: constants.ErrOrganizationsQuotaExceeded, + }, + { + name: "count error is propagated", + actorUserID: "user-1", + organizationLimit: &limit, + setup: func(repo *orgtests.MockOrganizationRepository) { + repo.On("CountAccessibleByUserID", mock.Anything, "user-1").Return(0, countErr).Once() + }, + expectErr: countErr, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + repo := &orgtests.MockOrganizationRepository{} + memberRepo := &orgtests.MockOrganizationMemberRepository{} + if tt.setup != nil { + tt.setup(repo) + } + + serviceUtils := &ServiceUtils{orgRepo: repo, orgMemberRepo: memberRepo} + svc := NewOrganizationService(repo, memberRepo, serviceUtils, nil, tt.organizationLimit, nil) + + err := svc.ensureOrganizationLimit(context.Background(), orgtests.Actor(tt.actorUserID), repo) + if tt.expectErr != nil { + require.ErrorIs(t, err, tt.expectErr) + } else { + require.NoError(t, err) + } + + if tt.expectNoCount { + repo.AssertNotCalled(t, "CountAccessibleByUserID", mock.Anything, mock.Anything) + } + require.True(t, repo.AssertExpectations(t)) + }) + } +} + func TestOrganizationService_GetOrganizationByID(t *testing.T) { t.Parallel() diff --git a/plugins/organizations/services/organization_team_member_service.go b/plugins/organizations/services/organization_team_member_service.go index 08ecf8eb..31266ba1 100644 --- a/plugins/organizations/services/organization_team_member_service.go +++ b/plugins/organizations/services/organization_team_member_service.go @@ -4,6 +4,7 @@ import ( "context" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" "github.com/Authula/authula/plugins/organizations/repositories" "github.com/Authula/authula/plugins/organizations/types" @@ -96,7 +97,7 @@ func (s *organizationTeamMemberService) AddTeamMember(ctx context.Context, actor return created, nil } -func (s *organizationTeamMemberService) GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, error) { +func (s *organizationTeamMemberService) GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, params pagination.Params) (*types.ListOrganizationTeamMembersResponse, error) { if _, _, err := s.serviceUtils.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { return nil, err } @@ -113,7 +114,20 @@ func (s *organizationTeamMemberService) GetAllTeamMembers(ctx context.Context, a return nil, coreerrors.ErrNotFound } - return s.orgTeamMemberRepo.GetAllByTeamIDWithMemberAndUser(ctx, teamID, page, limit) + params = pagination.Clamp(params) + + teamMembers, total, err := s.orgTeamMemberRepo.GetAllByTeamIDWithMemberAndUser(ctx, teamID, params.Page, params.Limit) + if err != nil { + return nil, err + } + if teamMembers == nil { + teamMembers = []types.OrganizationTeamMemberResponse{} + } + + return &types.ListOrganizationTeamMembersResponse{ + Data: teamMembers, + Pagination: pagination.New(params.Page, params.Limit, total), + }, nil } func (s *organizationTeamMemberService) GetTeamMember(ctx context.Context, actor *models.Actor, organizationID string, teamID string, memberID string) (*types.OrganizationTeamMemberResponse, error) { diff --git a/plugins/organizations/services/organization_team_member_service_test.go b/plugins/organizations/services/organization_team_member_service_test.go index 436b17fc..53aa5021 100644 --- a/plugins/organizations/services/organization_team_member_service_test.go +++ b/plugins/organizations/services/organization_team_member_service_test.go @@ -9,6 +9,7 @@ import ( "github.com/stretchr/testify/require" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" orgtests "github.com/Authula/authula/plugins/organizations/tests" "github.com/Authula/authula/plugins/organizations/types" ) @@ -33,6 +34,7 @@ func TestOrganizationTeamService_GetAllTeamMembers(t *testing.T) { actorUserID string organizationID string teamID string + params pagination.Params setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationTeamRepository, *orgtests.MockOrganizationMemberRepository, *orgtests.MockOrganizationTeamMemberRepository) expectErr error expectLen int @@ -47,7 +49,7 @@ func TestOrganizationTeamService_GetAllTeamMembers(t *testing.T) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-1").Return(&types.OrganizationMember{ID: "owner-member-1", OrganizationID: "org-1", UserID: "user-1", Role: "owner"}, nil).Twice() teamRepo.On("GetByID", mock.Anything, "team-1").Return(&types.OrganizationTeam{ID: "team-1", OrganizationID: "org-1"}, nil).Twice() - teamMemberRepo.On("GetAllByTeamIDWithMemberAndUser", mock.Anything, "team-1", 1, 10).Return([]types.OrganizationTeamMemberResponse{{ID: "tm-1", TeamID: "team-1"}}, nil).Once() + teamMemberRepo.On("GetAllByTeamIDWithMemberAndUser", mock.Anything, "team-1", 1, 10).Return([]types.OrganizationTeamMemberResponse{{ID: "tm-1", TeamID: "team-1"}}, 1, nil).Once() }, expectLen: 1, expectCalled: true, @@ -61,7 +63,7 @@ func TestOrganizationTeamService_GetAllTeamMembers(t *testing.T) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "owner-1"}, nil).Once() memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-2").Return(&types.OrganizationMember{ID: "org-member-1", OrganizationID: "org-1", UserID: "user-2", Role: "member"}, nil).Twice() teamRepo.On("GetByID", mock.Anything, "team-1").Return(&types.OrganizationTeam{ID: "team-1", OrganizationID: "org-1"}, nil).Twice() - teamMemberRepo.On("GetAllByTeamIDWithMemberAndUser", mock.Anything, "team-1", 1, 10).Return([]types.OrganizationTeamMemberResponse{{ID: "tm-1", TeamID: "team-1"}}, nil).Once() + teamMemberRepo.On("GetAllByTeamIDWithMemberAndUser", mock.Anything, "team-1", 1, 10).Return([]types.OrganizationTeamMemberResponse{{ID: "tm-1", TeamID: "team-1"}}, 1, nil).Once() }, expectLen: 1, expectCalled: true, @@ -144,11 +146,26 @@ func TestOrganizationTeamService_GetAllTeamMembers(t *testing.T) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-1").Return(&types.OrganizationMember{ID: "owner-member-1", OrganizationID: "org-1", UserID: "user-1", Role: "owner"}, nil).Twice() teamRepo.On("GetByID", mock.Anything, "team-1").Return(&types.OrganizationTeam{ID: "team-1", OrganizationID: "org-1"}, nil).Twice() - teamMemberRepo.On("GetAllByTeamIDWithMemberAndUser", mock.Anything, "team-1", 1, 10).Return(([]types.OrganizationTeamMemberResponse)(nil), repoErr).Once() + teamMemberRepo.On("GetAllByTeamIDWithMemberAndUser", mock.Anything, "team-1", 1, 10).Return(([]types.OrganizationTeamMemberResponse)(nil), 0, repoErr).Once() }, expectErr: repoErr, expectCalled: true, }, + { + name: "out of range params are clamped before reaching the repository", + actorUserID: "user-1", + organizationID: "org-1", + teamID: "team-1", + params: pagination.Params{Page: -4, Limit: 5000}, + setup: func(orgRepo *orgtests.MockOrganizationRepository, teamRepo *orgtests.MockOrganizationTeamRepository, memberRepo *orgtests.MockOrganizationMemberRepository, teamMemberRepo *orgtests.MockOrganizationTeamMemberRepository) { + orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() + memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-1").Return(&types.OrganizationMember{ID: "owner-member-1", OrganizationID: "org-1", UserID: "user-1", Role: "owner"}, nil).Twice() + teamRepo.On("GetByID", mock.Anything, "team-1").Return(&types.OrganizationTeam{ID: "team-1", OrganizationID: "org-1"}, nil).Twice() + teamMemberRepo.On("GetAllByTeamIDWithMemberAndUser", mock.Anything, "team-1", 1, pagination.MaxLimit).Return(([]types.OrganizationTeamMemberResponse)(nil), 0, nil).Once() + }, + expectLen: 0, + expectCalled: true, + }, } for _, tt := range tests { @@ -163,8 +180,13 @@ func TestOrganizationTeamService_GetAllTeamMembers(t *testing.T) { tt.setup(orgRepo, teamRepo, memberRepo, teamMemberRepo) } + params := tt.params + if params == (pagination.Params{}) { + params = pagination.Params{Page: 1, Limit: 10} + } + svc := newTestOrganizationTeamMemberService(orgRepo, memberRepo, teamRepo, teamMemberRepo) - members, err := svc.GetAllTeamMembers(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.teamID, 1, 10) + resp, err := svc.GetAllTeamMembers(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.teamID, params) if tt.expectErr != nil { require.Error(t, err) require.ErrorIs(t, err, tt.expectErr) @@ -175,7 +197,11 @@ func TestOrganizationTeamService_GetAllTeamMembers(t *testing.T) { return } require.NoError(t, err) - require.Len(t, members, tt.expectLen) + require.NotNil(t, resp) + require.NotNil(t, resp.Data) + require.Len(t, resp.Data, tt.expectLen) + require.Equal(t, pagination.Clamp(params).Page, resp.Pagination.Page) + require.Equal(t, pagination.Clamp(params).Limit, resp.Pagination.Limit) require.True(t, orgRepo.AssertExpectations(t)) require.True(t, memberRepo.AssertExpectations(t)) require.True(t, teamRepo.AssertExpectations(t)) diff --git a/plugins/organizations/services/organization_team_service.go b/plugins/organizations/services/organization_team_service.go index 0b8e281d..b95a7151 100644 --- a/plugins/organizations/services/organization_team_service.go +++ b/plugins/organizations/services/organization_team_service.go @@ -7,6 +7,7 @@ import ( "github.com/uptrace/bun" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" "github.com/Authula/authula/plugins/organizations/repositories" "github.com/Authula/authula/plugins/organizations/types" @@ -150,12 +151,25 @@ func (s *organizationTeamService) CreateTeam(ctx context.Context, actor *models. return created, nil } -func (s *organizationTeamService) GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationTeam, error) { +func (s *organizationTeamService) GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationTeamsResponse, error) { if _, _, err := s.serviceUtils.AuthorizeOrganizationAccess(ctx, actor, organizationID); err != nil { return nil, err } - return s.orgTeamRepo.GetAllByOrganizationID(ctx, organizationID) + params = pagination.Clamp(params) + + teams, total, err := s.orgTeamRepo.GetAllByOrganizationID(ctx, organizationID, params.Page, params.Limit) + if err != nil { + return nil, err + } + if teams == nil { + teams = []types.OrganizationTeam{} + } + + return &types.ListOrganizationTeamsResponse{ + Data: teams, + Pagination: pagination.New(params.Page, params.Limit, total), + }, nil } func (s *organizationTeamService) GetTeam(ctx context.Context, actor *models.Actor, organizationID string, teamID string) (*types.OrganizationTeam, error) { diff --git a/plugins/organizations/services/organization_team_service_test.go b/plugins/organizations/services/organization_team_service_test.go index e763397c..3f8f7106 100644 --- a/plugins/organizations/services/organization_team_service_test.go +++ b/plugins/organizations/services/organization_team_service_test.go @@ -9,6 +9,7 @@ import ( "github.com/stretchr/testify/require" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" orgtests "github.com/Authula/authula/plugins/organizations/tests" "github.com/Authula/authula/plugins/organizations/types" ) @@ -239,50 +240,66 @@ func TestOrganizationTeamService_GetAllTeams(t *testing.T) { repoErr := errors.New("repository error") tests := []struct { - name string - actorUserID string - organizationID string - setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationTeamRepository, *orgtests.MockOrganizationMemberRepository) - expectErr error - expectLen int + name string + actorUserID string + organizationID string + params pagination.Params + setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationTeamRepository, *orgtests.MockOrganizationMemberRepository) + expectErr error + expectLen int + expectPagination pagination.Pagination }{ { name: "success", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, teamRepo *orgtests.MockOrganizationTeamRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-1").Return(&types.OrganizationMember{ID: "mem-1", OrganizationID: "org-1", UserID: "user-1", Role: "owner"}, nil).Once() - teamRepo.On("GetAllByOrganizationID", mock.Anything, "org-1").Return([]types.OrganizationTeam{{ID: "team-1", OrganizationID: "org-1", Name: "Platform", Slug: "platform"}}, nil).Once() + teamRepo.On("GetAllByOrganizationID", mock.Anything, "org-1", 1, 10).Return([]types.OrganizationTeam{{ID: "team-1", OrganizationID: "org-1", Name: "Platform", Slug: "platform"}}, 1, nil).Once() }, - expectLen: 1, + expectLen: 1, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: 1, TotalPages: 1, HasMore: false}, }, { name: "org member can list", actorUserID: "user-2", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, teamRepo *orgtests.MockOrganizationTeamRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "owner-1"}, nil).Once() memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-2").Return(&types.OrganizationMember{ID: "mem-1", OrganizationID: "org-1", UserID: "user-2", Role: "member"}, nil).Once() - teamRepo.On("GetAllByOrganizationID", mock.Anything, "org-1").Return([]types.OrganizationTeam{{ID: "team-1", OrganizationID: "org-1", Name: "Platform", Slug: "platform"}}, nil).Once() + teamRepo.On("GetAllByOrganizationID", mock.Anything, "org-1", 1, 10).Return([]types.OrganizationTeam{{ID: "team-1", OrganizationID: "org-1", Name: "Platform", Slug: "platform"}}, 1, nil).Once() }, - expectLen: 1, + expectLen: 1, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: 1, TotalPages: 1, HasMore: false}, + }, + { + name: "out of range params are clamped before reaching the repository", + actorUserID: "user-1", + organizationID: "org-1", + params: pagination.Params{Page: -4, Limit: 5000}, + setup: func(orgRepo *orgtests.MockOrganizationRepository, teamRepo *orgtests.MockOrganizationTeamRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { + orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() + memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-1").Return(&types.OrganizationMember{ID: "mem-1", OrganizationID: "org-1", UserID: "user-1", Role: "owner"}, nil).Once() + teamRepo.On("GetAllByOrganizationID", mock.Anything, "org-1", 1, pagination.MaxLimit).Return(([]types.OrganizationTeam)(nil), 0, nil).Once() + }, + expectLen: 0, + expectPagination: pagination.Pagination{Page: 1, Limit: pagination.MaxLimit, Total: 0, TotalPages: 0, HasMore: false}, + }, + { + name: "unauthorized", + actorUserID: "", + organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, + expectErr: coreerrors.ErrUnauthorized, }, - } - - tests = append(tests, []struct { - name string - actorUserID string - organizationID string - setup func(*orgtests.MockOrganizationRepository, *orgtests.MockOrganizationTeamRepository, *orgtests.MockOrganizationMemberRepository) - expectErr error - expectLen int - }{ - {name: "unauthorized", actorUserID: "", organizationID: "org-1", expectErr: coreerrors.ErrUnauthorized}, { name: "organization not found", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, teamRepo *orgtests.MockOrganizationTeamRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(nil, nil).Once() }, @@ -292,6 +309,7 @@ func TestOrganizationTeamService_GetAllTeams(t *testing.T) { name: "organization lookup error", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, teamRepo *orgtests.MockOrganizationTeamRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return((*types.Organization)(nil), repoErr).Once() }, @@ -301,6 +319,7 @@ func TestOrganizationTeamService_GetAllTeams(t *testing.T) { name: "forbidden", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, teamRepo *orgtests.MockOrganizationTeamRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "owner-1"}, nil).Once() }, @@ -310,14 +329,15 @@ func TestOrganizationTeamService_GetAllTeams(t *testing.T) { name: "repo error", actorUserID: "user-1", organizationID: "org-1", + params: pagination.Params{Page: 1, Limit: 10}, setup: func(orgRepo *orgtests.MockOrganizationRepository, teamRepo *orgtests.MockOrganizationTeamRepository, memberRepo *orgtests.MockOrganizationMemberRepository) { orgRepo.On("GetByID", mock.Anything, "org-1").Return(&types.Organization{ID: "org-1", OwnerID: "user-1"}, nil).Once() memberRepo.On("GetByOrganizationIDAndUserID", mock.Anything, "org-1", "user-1").Return(&types.OrganizationMember{ID: "mem-1", OrganizationID: "org-1", UserID: "user-1", Role: "owner"}, nil).Once() - teamRepo.On("GetAllByOrganizationID", mock.Anything, "org-1").Return(([]types.OrganizationTeam)(nil), repoErr).Once() + teamRepo.On("GetAllByOrganizationID", mock.Anything, "org-1", 1, 10).Return(([]types.OrganizationTeam)(nil), 0, repoErr).Once() }, expectErr: repoErr, }, - }...) + } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -331,14 +351,21 @@ func TestOrganizationTeamService_GetAllTeams(t *testing.T) { } svc := newTestOrganizationTeamService(orgRepo, memberRepo, teamRepo, &orgtests.MockOrganizationTeamMemberRepository{}) - teams, err := svc.GetAllTeams(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID) + resp, err := svc.GetAllTeams(context.Background(), orgtests.Actor(tt.actorUserID), tt.organizationID, tt.params) if tt.expectErr != nil { require.Error(t, err) require.ErrorIs(t, err, tt.expectErr) return } + require.NoError(t, err) - require.Len(t, teams, tt.expectLen) + require.NotNil(t, resp) + require.NotNil(t, resp.Data) + require.Len(t, resp.Data, tt.expectLen) + require.Equal(t, tt.expectPagination, resp.Pagination) + require.True(t, orgRepo.AssertExpectations(t)) + require.True(t, teamRepo.AssertExpectations(t)) + require.True(t, memberRepo.AssertExpectations(t)) }) } } diff --git a/plugins/organizations/services/pagination_integration_test.go b/plugins/organizations/services/pagination_integration_test.go new file mode 100644 index 00000000..0e1bf4a5 --- /dev/null +++ b/plugins/organizations/services/pagination_integration_test.go @@ -0,0 +1,140 @@ +package services + +import ( + "context" + "fmt" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/Authula/authula/core/pagination" + internaltests "github.com/Authula/authula/internal/tests" + "github.com/Authula/authula/plugins/organizations/repositories" + orgtests "github.com/Authula/authula/plugins/organizations/tests" + "github.com/Authula/authula/plugins/organizations/types" +) + +func TestOrganizationMemberService_GetAllMembersEnforcesLimitsAgainstSQL(t *testing.T) { + t.Parallel() + + const memberCount = 12 + + setup := func(t *testing.T) (*organizationMemberService, context.Context) { + t.Helper() + + db := orgtests.SetupRepoDB(t) + orgtests.SeedUsers(t, db, memberCount) + orgtests.SeedOrganization(t, db, "org-1", "user-1", "Acme Inc", "acme-inc") + for i := 1; i <= memberCount; i++ { + orgtests.SeedOrganizationMember(t, db, fmt.Sprintf("mem-%02d", i), "org-1", fmt.Sprintf("user-%d", i), "member") + } + + orgRepo := repositories.NewBunOrganizationRepository(db) + memberRepo := repositories.NewBunOrganizationMemberRepository(db) + serviceUtils := &ServiceUtils{orgRepo: orgRepo, orgMemberRepo: memberRepo} + svc := NewOrganizationMemberService( + &internaltests.MockUserService{}, + orgtests.NewAccessControlServiceStub(), + orgRepo, + memberRepo, + nil, + &orgtests.MockTxRunner{}, + serviceUtils, + ) + + return svc, context.Background() + } + + tests := []struct { + name string + params pagination.Params + expectLen int + expectPagination pagination.Pagination + }{ + { + name: "no params applies the defaults", + params: pagination.Params{}, + expectLen: pagination.DefaultLimit, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: memberCount, TotalPages: 2, HasMore: true}, + }, + { + name: "an absurd limit is capped at the maximum", + params: pagination.Params{Page: 1, Limit: 100000}, + expectLen: memberCount, + expectPagination: pagination.Pagination{Page: 1, Limit: pagination.MaxLimit, Total: memberCount, TotalPages: 1, HasMore: false}, + }, + { + name: "a negative limit does not read the whole table", + params: pagination.Params{Page: 1, Limit: -1}, + expectLen: pagination.DefaultLimit, + expectPagination: pagination.Pagination{Page: 1, Limit: 10, Total: memberCount, TotalPages: 2, HasMore: true}, + }, + { + name: "page zero returns the first page", + params: pagination.Params{Page: 0, Limit: 5}, + expectLen: 5, + expectPagination: pagination.Pagination{Page: 1, Limit: 5, Total: memberCount, TotalPages: 3, HasMore: true}, + }, + { + name: "a page past the end is empty with a correct total", + params: pagination.Params{Page: 99999, Limit: 10}, + expectLen: 0, + expectPagination: pagination.Pagination{Page: 99999, Limit: 10, Total: memberCount, TotalPages: 2, HasMore: false}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + svc, ctx := setup(t) + + resp, err := svc.GetAllMembers(ctx, orgtests.Actor("user-1"), "org-1", tt.params) + require.NoError(t, err) + require.NotNil(t, resp) + require.Len(t, resp.Data, tt.expectLen) + require.Equal(t, tt.expectPagination, resp.Pagination) + }) + } +} + +func TestOrganizationService_QuotaSurvivesPagination(t *testing.T) { + t.Parallel() + + const limit = 3 + + db := orgtests.SetupRepoDB(t) + orgRepo := repositories.NewBunOrganizationRepository(db) + memberRepo := repositories.NewBunOrganizationMemberRepository(db) + ctx := context.Background() + + // user-1 owns one organization and is a member of two more it does not own. + orgtests.SeedOrganization(t, db, "org-owned", "user-1", "Owned", "owned") + orgtests.SeedOrganizationMember(t, db, "mem-owned", "org-owned", "user-1", "owner") + orgtests.SeedOrganization(t, db, "org-joined-1", "user-2", "Joined One", "joined-one") + orgtests.SeedOrganizationMember(t, db, "mem-joined-1", "org-joined-1", "user-1", "member") + + serviceUtils := &ServiceUtils{orgRepo: orgRepo, orgMemberRepo: memberRepo} + svc := NewOrganizationService(orgRepo, memberRepo, serviceUtils, nil, new(limit), nil) + + require.NoError(t, svc.ensureOrganizationLimit(ctx, orgtests.Actor("user-1"), orgRepo), "two organizations is below the quota") + + orgtests.SeedOrganization(t, db, "org-joined-2", "user-2", "Joined Two", "joined-two") + orgtests.SeedOrganizationMember(t, db, "mem-joined-2", "org-joined-2", "user-1", "member") + + require.Error(t, svc.ensureOrganizationLimit(ctx, orgtests.Actor("user-1"), orgRepo), "the quota must reject at the boundary") + + // The quota counts the same set the list endpoint returns, deduped. + resp, err := svc.GetAllOrganizations(ctx, orgtests.Actor("user-1"), pagination.Params{Page: 1, Limit: 1}) + require.NoError(t, err) + require.Len(t, resp.Data, 1, "a single-row page") + require.Equal(t, limit, resp.Pagination.Total, "the total must not be capped by the page size") + + var seen []types.Organization + for page := 1; page <= limit; page++ { + pageResp, err := svc.GetAllOrganizations(ctx, orgtests.Actor("user-1"), pagination.Params{Page: page, Limit: 1}) + require.NoError(t, err) + seen = append(seen, pageResp.Data...) + } + require.Len(t, seen, limit, "paging must yield every accessible organization exactly once") +} diff --git a/plugins/organizations/tests/repositories.go b/plugins/organizations/tests/repositories.go index 859158b3..0ce10761 100644 --- a/plugins/organizations/tests/repositories.go +++ b/plugins/organizations/tests/repositories.go @@ -53,12 +53,17 @@ func (m *MockOrganizationRepository) GetBySlug(ctx context.Context, slug string) return args.Get(0).(*types.Organization), args.Error(1) } -func (m *MockOrganizationRepository) GetAllByOwnerID(ctx context.Context, ownerID string) ([]types.Organization, error) { - args := m.Called(ctx, ownerID) +func (m *MockOrganizationRepository) GetAllAccessibleByUserID(ctx context.Context, userID string, page int, limit int) ([]types.Organization, int, error) { + args := m.Called(ctx, userID, page, limit) if args.Get(0) == nil { - return nil, args.Error(1) + return nil, args.Int(1), args.Error(2) } - return args.Get(0).([]types.Organization), args.Error(1) + return args.Get(0).([]types.Organization), args.Int(1), args.Error(2) +} + +func (m *MockOrganizationRepository) CountAccessibleByUserID(ctx context.Context, userID string) (int, error) { + args := m.Called(ctx, userID) + return args.Int(0), args.Error(1) } func (m *MockOrganizationRepository) Update(ctx context.Context, organization *types.Organization) (*types.Organization, error) { @@ -107,20 +112,20 @@ func (m *MockOrganizationMemberRepository) GetByOrganizationIDAndUserID(ctx cont return nil, nil } -func (m *MockOrganizationMemberRepository) GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMember, error) { +func (m *MockOrganizationMemberRepository) GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMember, int, error) { args := m.Called(ctx, organizationID, page, limit) if args.Get(0) == nil { - return nil, args.Error(1) + return nil, args.Int(1), args.Error(2) } - return args.Get(0).([]types.OrganizationMember), args.Error(1) + return args.Get(0).([]types.OrganizationMember), args.Int(1), args.Error(2) } -func (m *MockOrganizationMemberRepository) GetAllByOrganizationIDWithUser(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, error) { +func (m *MockOrganizationMemberRepository) GetAllByOrganizationIDWithUser(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, int, error) { args := m.Called(ctx, organizationID, page, limit) if args.Get(0) == nil { - return nil, args.Error(1) + return nil, args.Int(1), args.Error(2) } - return args.Get(0).([]types.OrganizationMemberResponse), args.Error(1) + return args.Get(0).([]types.OrganizationMemberResponse), args.Int(1), args.Error(2) } func (m *MockOrganizationMemberRepository) GetByIDWithUser(ctx context.Context, memberID string) (*types.OrganizationMemberResponse, error) { @@ -139,14 +144,6 @@ func (m *MockOrganizationMemberRepository) GetByOrganizationIDAndUserIDWithUser( return args.Get(0).(*types.OrganizationMemberResponse), args.Error(1) } -func (m *MockOrganizationMemberRepository) GetAllByUserID(ctx context.Context, userID string) ([]types.OrganizationMember, error) { - args := m.Called(ctx, userID) - if args.Get(0) == nil { - return nil, args.Error(1) - } - return args.Get(0).([]types.OrganizationMember), args.Error(1) -} - func (m *MockOrganizationMemberRepository) GetByID(ctx context.Context, memberID string) (*types.OrganizationMember, error) { args := m.Called(ctx, memberID) if args.Get(0) == nil { @@ -203,14 +200,6 @@ func (m *MockOrganizationInvitationRepository) GetByOrganizationIDAndEmail(ctx c return args.Get(0).(*types.OrganizationInvitation), args.Error(1) } -func (m *MockOrganizationInvitationRepository) GetAllByOrganizationID(ctx context.Context, organizationID string) ([]types.OrganizationInvitation, error) { - args := m.Called(ctx, organizationID) - if args.Get(0) == nil { - return nil, args.Error(1) - } - return args.Get(0).([]types.OrganizationInvitation), args.Error(1) -} - func (m *MockOrganizationInvitationRepository) GetByIDWithOrg(ctx context.Context, invitationID string) (*types.GetOrganizationInvitationResponse, error) { args := m.Called(ctx, invitationID) if args.Get(0) == nil { @@ -219,16 +208,16 @@ func (m *MockOrganizationInvitationRepository) GetByIDWithOrg(ctx context.Contex return args.Get(0).(*types.GetOrganizationInvitationResponse), args.Error(1) } -func (m *MockOrganizationInvitationRepository) GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { - args := m.Called(ctx, organizationID) +func (m *MockOrganizationInvitationRepository) GetAllByOrganizationIDWithOrg(ctx context.Context, organizationID string, page int, limit int) ([]types.GetOrganizationInvitationResponse, int, error) { + args := m.Called(ctx, organizationID, page, limit) if args.Get(0) == nil { - return nil, args.Error(1) + return nil, args.Int(1), args.Error(2) } - return args.Get(0).([]types.GetOrganizationInvitationResponse), args.Error(1) + return args.Get(0).([]types.GetOrganizationInvitationResponse), args.Int(1), args.Error(2) } -func (m *MockOrganizationInvitationRepository) GetAllPendingByEmail(ctx context.Context, email string) ([]types.OrganizationInvitation, error) { - args := m.Called(ctx, email) +func (m *MockOrganizationInvitationRepository) GetAllPendingByEmail(ctx context.Context, email string, limit int) ([]types.OrganizationInvitation, error) { + args := m.Called(ctx, email, limit) if args.Get(0) == nil { return nil, args.Error(1) } @@ -280,12 +269,12 @@ func (m *MockOrganizationTeamRepository) GetByOrganizationIDAndSlug(ctx context. return args.Get(0).(*types.OrganizationTeam), args.Error(1) } -func (m *MockOrganizationTeamRepository) GetAllByOrganizationID(ctx context.Context, organizationID string) ([]types.OrganizationTeam, error) { - args := m.Called(ctx, organizationID) +func (m *MockOrganizationTeamRepository) GetAllByOrganizationID(ctx context.Context, organizationID string, page int, limit int) ([]types.OrganizationTeam, int, error) { + args := m.Called(ctx, organizationID, page, limit) if args.Get(0) == nil { - return nil, args.Error(1) + return nil, args.Int(1), args.Error(2) } - return args.Get(0).([]types.OrganizationTeam), args.Error(1) + return args.Get(0).([]types.OrganizationTeam), args.Int(1), args.Error(2) } func (m *MockOrganizationTeamRepository) Update(ctx context.Context, team *types.OrganizationTeam) (*types.OrganizationTeam, error) { @@ -332,20 +321,20 @@ func (m *MockOrganizationTeamMemberRepository) GetByTeamIDAndMemberID(ctx contex return args.Get(0).(*types.OrganizationTeamMember), args.Error(1) } -func (m *MockOrganizationTeamMemberRepository) GetAllByTeamID(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMember, error) { +func (m *MockOrganizationTeamMemberRepository) GetAllByTeamID(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMember, int, error) { args := m.Called(ctx, teamID, page, limit) if args.Get(0) == nil { - return nil, args.Error(1) + return nil, args.Int(1), args.Error(2) } - return args.Get(0).([]types.OrganizationTeamMember), args.Error(1) + return args.Get(0).([]types.OrganizationTeamMember), args.Int(1), args.Error(2) } -func (m *MockOrganizationTeamMemberRepository) GetAllByTeamIDWithMemberAndUser(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, error) { +func (m *MockOrganizationTeamMemberRepository) GetAllByTeamIDWithMemberAndUser(ctx context.Context, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, int, error) { args := m.Called(ctx, teamID, page, limit) if args.Get(0) == nil { - return nil, args.Error(1) + return nil, args.Int(1), args.Error(2) } - return args.Get(0).([]types.OrganizationTeamMemberResponse), args.Error(1) + return args.Get(0).([]types.OrganizationTeamMemberResponse), args.Int(1), args.Error(2) } func (m *MockOrganizationTeamMemberRepository) GetByIDWithMemberAndUser(ctx context.Context, teamMemberID string) (*types.OrganizationTeamMemberResponse, error) { diff --git a/plugins/organizations/tests/services.go b/plugins/organizations/tests/services.go index e42ec080..48c8fe2b 100644 --- a/plugins/organizations/tests/services.go +++ b/plugins/organizations/tests/services.go @@ -5,6 +5,7 @@ import ( "github.com/stretchr/testify/mock" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" "github.com/Authula/authula/plugins/organizations/types" ) @@ -29,12 +30,12 @@ func (m *MockOrganizationService) CreateOrganization(ctx context.Context, actor return args.Get(0).(*types.Organization), args.Error(1) } -func (m *MockOrganizationService) GetAllOrganizationsByOwner(ctx context.Context, actor *models.Actor) ([]types.Organization, error) { - args := m.Called(ctx, actorID(actor)) +func (m *MockOrganizationService) GetAllOrganizations(ctx context.Context, actor *models.Actor, params pagination.Params) (*types.ListOrganizationsResponse, error) { + args := m.Called(ctx, actorID(actor), params) if args.Get(0) == nil { return nil, args.Error(1) } - return args.Get(0).([]types.Organization), args.Error(1) + return args.Get(0).(*types.ListOrganizationsResponse), args.Error(1) } func (m *MockOrganizationService) GetOrganizationByID(ctx context.Context, actor *models.Actor, organizationID string) (*types.Organization, error) { @@ -99,20 +100,12 @@ func (m *MockOrganizationInvitationService) GetOrganizationInvitationByIDWithOrg return args.Get(0).(*types.GetOrganizationInvitationResponse), args.Error(1) } -func (m *MockOrganizationInvitationService) GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { - args := m.Called(ctx, organizationID) - if args.Get(0) == nil { - return nil, args.Error(1) - } - return args.Get(0).([]types.GetOrganizationInvitationResponse), args.Error(1) -} - -func (m *MockOrganizationInvitationService) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationInvitation, error) { - args := m.Called(ctx, actorID(actor), organizationID) +func (m *MockOrganizationInvitationService) GetAllOrganizationInvitationsByOrgIDWithOrg(ctx context.Context, organizationID string, params pagination.Params) (*types.ListOrganizationInvitationsResponse, error) { + args := m.Called(ctx, organizationID, params) if args.Get(0) == nil { return nil, args.Error(1) } - return args.Get(0).([]types.OrganizationInvitation), args.Error(1) + return args.Get(0).(*types.ListOrganizationInvitationsResponse), args.Error(1) } func (m *MockOrganizationInvitationService) RevokeOrganizationInvitation(ctx context.Context, actor *models.Actor, organizationID string, invitationID string) (*types.OrganizationInvitation, error) { @@ -151,12 +144,12 @@ func (m *MockOrganizationMemberService) AddMember(ctx context.Context, actor *mo return args.Get(0).(*types.OrganizationMember), args.Error(1) } -func (m *MockOrganizationMemberService) GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, error) { - args := m.Called(ctx, actorID(actor), organizationID, page, limit) +func (m *MockOrganizationMemberService) GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationMembersResponse, error) { + args := m.Called(ctx, actorID(actor), organizationID, params) if args.Get(0) == nil { return nil, args.Error(1) } - return args.Get(0).([]types.OrganizationMemberResponse), args.Error(1) + return args.Get(0).(*types.ListOrganizationMembersResponse), args.Error(1) } func (m *MockOrganizationMemberService) GetMember(ctx context.Context, actor *models.Actor, organizationID string, memberID string) (*types.OrganizationMemberResponse, error) { @@ -200,12 +193,12 @@ func (m *MockOrganizationTeamService) CreateTeam(ctx context.Context, actor *mod return args.Get(0).(*types.OrganizationTeam), args.Error(1) } -func (m *MockOrganizationTeamService) GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationTeam, error) { - args := m.Called(ctx, actorID(actor), organizationID) +func (m *MockOrganizationTeamService) GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationTeamsResponse, error) { + args := m.Called(ctx, actorID(actor), organizationID, params) if args.Get(0) == nil { return nil, args.Error(1) } - return args.Get(0).([]types.OrganizationTeam), args.Error(1) + return args.Get(0).(*types.ListOrganizationTeamsResponse), args.Error(1) } func (m *MockOrganizationTeamService) GetTeam(ctx context.Context, actor *models.Actor, organizationID string, teamID string) (*types.OrganizationTeam, error) { @@ -241,12 +234,12 @@ func (m *MockOrganizationTeamMemberService) AddTeamMember(ctx context.Context, a return args.Get(0).(*types.OrganizationTeamMember), args.Error(1) } -func (m *MockOrganizationTeamMemberService) GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, error) { - args := m.Called(ctx, actorID(actor), organizationID, teamID, page, limit) +func (m *MockOrganizationTeamMemberService) GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, params pagination.Params) (*types.ListOrganizationTeamMembersResponse, error) { + args := m.Called(ctx, actorID(actor), organizationID, teamID, params) if args.Get(0) == nil { return nil, args.Error(1) } - return args.Get(0).([]types.OrganizationTeamMemberResponse), args.Error(1) + return args.Get(0).(*types.ListOrganizationTeamMembersResponse), args.Error(1) } func (m *MockOrganizationTeamMemberService) GetTeamMember(ctx context.Context, actor *models.Actor, organizationID string, teamID string, memberID string) (*types.OrganizationTeamMemberResponse, error) { diff --git a/plugins/organizations/tests/test_helpers.go b/plugins/organizations/tests/test_helpers.go index 08a4ab43..35d29577 100644 --- a/plugins/organizations/tests/test_helpers.go +++ b/plugins/organizations/tests/test_helpers.go @@ -3,6 +3,7 @@ package tests import ( "context" "database/sql" + "fmt" "testing" "time" @@ -54,6 +55,21 @@ func SetupRepoDB(t *testing.T) *bun.DB { return db } +func SeedUsers(t *testing.T, db bun.IDB, n int) []string { + t.Helper() + + userIDs := make([]string, 0, n) + for i := 1; i <= n; i++ { + userID := fmt.Sprintf("user-%d", i) + if i > 2 { + SeedUser(t, db, userID) + } + userIDs = append(userIDs, userID) + } + + return userIDs +} + func SeedUser(t *testing.T, db bun.IDB, userID string) { t.Helper() diff --git a/plugins/organizations/types/api.go b/plugins/organizations/types/api.go index b62baf8b..7fefac11 100644 --- a/plugins/organizations/types/api.go +++ b/plugins/organizations/types/api.go @@ -5,6 +5,7 @@ import ( "time" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" ) @@ -38,12 +39,29 @@ type TeamMemberID struct { MemberID string `path:"member_id"` } +type ListOrganizationsRequest struct { + Page int `query:"page" json:"page,omitempty" nullable:"false"` + Limit int `query:"limit" json:"limit,omitempty" nullable:"false"` +} + +type ListOrganizationInvitationsRequest struct { + OrganizationID string `path:"organization_id"` + Page int `query:"page" json:"page,omitempty" nullable:"false"` + Limit int `query:"limit" json:"limit,omitempty" nullable:"false"` +} + type ListOrganizationMembersRequest struct { OrganizationID string `path:"organization_id"` Page int `query:"page" json:"page,omitempty" nullable:"false"` Limit int `query:"limit" json:"limit,omitempty" nullable:"false"` } +type ListOrganizationTeamsRequest struct { + OrganizationID string `path:"organization_id"` + Page int `query:"page" json:"page,omitempty" nullable:"false"` + Limit int `query:"limit" json:"limit,omitempty" nullable:"false"` +} + type ListOrganizationTeamMembersRequest struct { OrganizationID string `path:"organization_id"` TeamID string `path:"team_id"` @@ -51,6 +69,31 @@ type ListOrganizationTeamMembersRequest struct { Limit int `query:"limit" json:"limit,omitempty" nullable:"false"` } +type ListOrganizationsResponse struct { + Data []Organization `json:"data" required:"true" nullable:"false"` + Pagination pagination.Pagination `json:"pagination" required:"true" nullable:"false"` +} + +type ListOrganizationInvitationsResponse struct { + Data []GetOrganizationInvitationResponse `json:"data" required:"true" nullable:"false"` + Pagination pagination.Pagination `json:"pagination" required:"true" nullable:"false"` +} + +type ListOrganizationMembersResponse struct { + Data []OrganizationMemberResponse `json:"data" required:"true" nullable:"false"` + Pagination pagination.Pagination `json:"pagination" required:"true" nullable:"false"` +} + +type ListOrganizationTeamsResponse struct { + Data []OrganizationTeam `json:"data" required:"true" nullable:"false"` + Pagination pagination.Pagination `json:"pagination" required:"true" nullable:"false"` +} + +type ListOrganizationTeamMembersResponse struct { + Data []OrganizationTeamMemberResponse `json:"data" required:"true" nullable:"false"` + Pagination pagination.Pagination `json:"pagination" required:"true" nullable:"false"` +} + type AcceptOrganizationInvitationQuery struct { OrganizationID string `path:"organization_id"` InvitationID string `path:"invitation_id"` diff --git a/plugins/organizations/usecases/usecases.go b/plugins/organizations/usecases/usecases.go index 740f439c..4de128b5 100644 --- a/plugins/organizations/usecases/usecases.go +++ b/plugins/organizations/usecases/usecases.go @@ -6,6 +6,7 @@ import ( "strings" coreerrors "github.com/Authula/authula/core/errors" + "github.com/Authula/authula/core/pagination" "github.com/Authula/authula/models" orgconstants "github.com/Authula/authula/plugins/organizations/constants" orgservices "github.com/Authula/authula/plugins/organizations/services" @@ -98,8 +99,8 @@ func (u *UseCases) CreateOrganization(ctx context.Context, actor *models.Actor, return u.orgService.CreateOrganization(ctx, actor, request) } -func (u *UseCases) GetAllOrganizationsByOwner(ctx context.Context, actor *models.Actor) ([]types.Organization, error) { - return u.orgService.GetAllOrganizationsByOwner(ctx, actor) +func (u *UseCases) GetAllOrganizations(ctx context.Context, actor *models.Actor, params pagination.Params) (*types.ListOrganizationsResponse, error) { + return u.orgService.GetAllOrganizations(ctx, actor, params) } func (u *UseCases) GetOrganizationByID(ctx context.Context, actor *models.Actor, organizationID string) (*types.Organization, error) { @@ -130,12 +131,12 @@ func (u *UseCases) CreateOrganizationInvitation(ctx context.Context, actor *mode return u.invitationService.CreateOrganizationInvitation(ctx, actor, organizationID, request, redirectURL) } -func (u *UseCases) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string) ([]types.GetOrganizationInvitationResponse, error) { +func (u *UseCases) GetAllOrganizationInvitations(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationInvitationsResponse, error) { if err := u.authorizeOrgAccess(ctx, actor, organizationID, orgconstants.OrganizationsInvitationsListPermission); err != nil { return nil, err } - resp, err := u.invitationService.GetAllOrganizationInvitationsByOrgIDWithOrg(ctx, organizationID) + resp, err := u.invitationService.GetAllOrganizationInvitationsByOrgIDWithOrg(ctx, organizationID, params) if err != nil { return nil, err } @@ -193,11 +194,11 @@ func (u *UseCases) AddMember(ctx context.Context, actor *models.Actor, organizat return u.memberService.AddMember(ctx, actor, organizationID, request) } -func (u *UseCases) GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, page int, limit int) ([]types.OrganizationMemberResponse, error) { +func (u *UseCases) GetAllMembers(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationMembersResponse, error) { if err := u.authorizeOrgAccess(ctx, actor, organizationID, orgconstants.OrganizationsMembersListPermission); err != nil { return nil, err } - return u.memberService.GetAllMembers(ctx, actor, organizationID, page, limit) + return u.memberService.GetAllMembers(ctx, actor, organizationID, params) } func (u *UseCases) GetMember(ctx context.Context, actor *models.Actor, organizationID string, memberID string) (*types.OrganizationMemberResponse, error) { @@ -237,11 +238,11 @@ func (u *UseCases) CreateTeam(ctx context.Context, actor *models.Actor, organiza return u.teamService.CreateTeam(ctx, actor, organizationID, request) } -func (u *UseCases) GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string) ([]types.OrganizationTeam, error) { +func (u *UseCases) GetAllTeams(ctx context.Context, actor *models.Actor, organizationID string, params pagination.Params) (*types.ListOrganizationTeamsResponse, error) { if err := u.authorizeOrgAccess(ctx, actor, organizationID, orgconstants.OrganizationsTeamsListPermission); err != nil { return nil, err } - return u.teamService.GetAllTeams(ctx, actor, organizationID) + return u.teamService.GetAllTeams(ctx, actor, organizationID, params) } func (u *UseCases) GetTeam(ctx context.Context, actor *models.Actor, organizationID string, teamID string) (*types.OrganizationTeam, error) { @@ -274,11 +275,11 @@ func (u *UseCases) AddTeamMember(ctx context.Context, actor *models.Actor, organ return u.teamMemberService.AddTeamMember(ctx, actor, organizationID, teamID, request) } -func (u *UseCases) GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, page int, limit int) ([]types.OrganizationTeamMemberResponse, error) { +func (u *UseCases) GetAllTeamMembers(ctx context.Context, actor *models.Actor, organizationID string, teamID string, params pagination.Params) (*types.ListOrganizationTeamMembersResponse, error) { if err := u.authorizeOrgAccess(ctx, actor, organizationID, orgconstants.OrganizationsTeamMembersListPermission); err != nil { return nil, err } - return u.teamMemberService.GetAllTeamMembers(ctx, actor, organizationID, teamID, page, limit) + return u.teamMemberService.GetAllTeamMembers(ctx, actor, organizationID, teamID, params) } func (u *UseCases) GetTeamMember(ctx context.Context, actor *models.Actor, organizationID string, teamID string, memberID string) (*types.OrganizationTeamMemberResponse, error) {