diff --git a/backend/ee/authentication/scim/utils.py b/backend/ee/authentication/scim/utils.py index 9119973a0..ea7ef001a 100644 --- a/backend/ee/authentication/scim/utils.py +++ b/backend/ee/authentication/scim/utils.py @@ -16,9 +16,11 @@ role_has_managed_key, ) from api.utils.keys import revoke_team_environment_keys +from ee.authentication.scim.constants import SCIM_DEFAULT_COUNT from ee.authentication.scim.exceptions import ( SCIMDeactivationForbidden, SCIMProvisioningConflict, + scim_bad_request, ) logger = logging.getLogger(__name__) @@ -199,3 +201,29 @@ def reactivate_scim_user(scim_user): if scim_user.org_member and scim_user.org_member.deleted_at is not None: scim_user.org_member.deleted_at = None scim_user.org_member.save(update_fields=["deleted_at"]) + + +def parse_pagination_params(request): + """Parse the SCIM startIndex and count query parameters. + + Returns (start_index, count, None) on success, or (None, None, response) + where response is a SCIM 400. Both parameters are client supplied, so a + non-integer must be reported rather than raised. Per RFC 7644 section 3.4.2.4 + startIndex is 1-based and a negative count is interpreted as zero, which also + keeps the value out of a negative queryset slice. + """ + raw_start = request.GET.get("startIndex", 1) + try: + start_index = max(int(raw_start), 1) + except (TypeError, ValueError): + return None, None, scim_bad_request( + f"startIndex must be an integer, got {raw_start!r}" + ) + + raw_count = request.GET.get("count", SCIM_DEFAULT_COUNT) + try: + count = min(int(raw_count), SCIM_DEFAULT_COUNT) + except (TypeError, ValueError): + return None, None, scim_bad_request(f"count must be an integer, got {raw_count!r}") + + return start_index, max(count, 0), None diff --git a/backend/ee/authentication/scim/views/groups.py b/backend/ee/authentication/scim/views/groups.py index 66f1de1be..ffd0d000b 100644 --- a/backend/ee/authentication/scim/views/groups.py +++ b/backend/ee/authentication/scim/views/groups.py @@ -27,8 +27,7 @@ ) from api.utils.keys import provision_team_environment_keys, revoke_team_environment_keys from ee.authentication.scim.auth import SCIMTokenAuthentication -from ee.authentication.scim.constants import SCIM_DEFAULT_COUNT -from ee.authentication.scim.utils import resolve_external_id +from ee.authentication.scim.utils import parse_pagination_params, resolve_external_id from ee.authentication.scim.exceptions import ( scim_bad_request, scim_conflict, @@ -232,8 +231,9 @@ def groups_detail(request, scim_group_id): def _list_groups(request, org): filter_str = request.GET.get("filter", "") - start_index = max(int(request.GET.get("startIndex", 1)), 1) - count = min(int(request.GET.get("count", SCIM_DEFAULT_COUNT)), SCIM_DEFAULT_COUNT) + start_index, count, error = parse_pagination_params(request) + if error: + return error qs = SCIMGroup.objects.filter(organisation=org).order_by("created_at") if filter_str: diff --git a/backend/ee/authentication/scim/views/users.py b/backend/ee/authentication/scim/views/users.py index 28156c874..b69459684 100644 --- a/backend/ee/authentication/scim/views/users.py +++ b/backend/ee/authentication/scim/views/users.py @@ -38,13 +38,13 @@ SCIM_USER_ATTR_MAP, scim_filter_to_queryset, ) -from ee.authentication.scim.constants import SCIM_DEFAULT_COUNT from ee.authentication.scim.serializers import ( serialize_list_response, serialize_scim_user, ) from ee.authentication.scim.logging import log_scim_event from ee.authentication.scim.utils import ( + parse_pagination_params, deactivate_scim_user, provision_scim_user, reactivate_scim_user, @@ -135,8 +135,9 @@ def users_detail(request, scim_user_id): def _list_users(request, org): filter_str = request.GET.get("filter", "") - start_index = max(int(request.GET.get("startIndex", 1)), 1) - count = min(int(request.GET.get("count", SCIM_DEFAULT_COUNT)), SCIM_DEFAULT_COUNT) + start_index, count, error = parse_pagination_params(request) + if error: + return error qs = SCIMUser.objects.filter(organisation=org).order_by("created_at") if filter_str: diff --git a/backend/tests/ee/authentication/scim/test_pagination_params.py b/backend/tests/ee/authentication/scim/test_pagination_params.py new file mode 100644 index 000000000..f880a4805 --- /dev/null +++ b/backend/tests/ee/authentication/scim/test_pagination_params.py @@ -0,0 +1,68 @@ +"""Tests for SCIM startIndex/count parsing on the list endpoints. + +startIndex and count come straight from the client, so a malformed value has to +come back as a SCIM 400 rather than a 500. +""" + +from unittest.mock import MagicMock + +import pytest + +from ee.authentication.scim.constants import SCIM_DEFAULT_COUNT +from ee.authentication.scim.utils import parse_pagination_params + + +def _request(**params): + request = MagicMock() + request.GET = params + return request + + +def test_defaults_when_absent(): + start_index, count, error = parse_pagination_params(_request()) + assert error is None + assert start_index == 1 + assert count == SCIM_DEFAULT_COUNT + + +def test_valid_values_are_used(): + start_index, count, error = parse_pagination_params( + _request(startIndex="5", count="20") + ) + assert error is None + assert (start_index, count) == (5, 20) + + +@pytest.mark.parametrize("bad", ["abc", "", "1.5", "null", "1,000"]) +def test_non_integer_start_index_is_a_400(bad): + start_index, count, error = parse_pagination_params(_request(startIndex=bad)) + assert (start_index, count) == (None, None) + assert error.status_code == 400 + + +@pytest.mark.parametrize("bad", ["abc", "", "1.5"]) +def test_non_integer_count_is_a_400(bad): + start_index, count, error = parse_pagination_params(_request(count=bad)) + assert (start_index, count) == (None, None) + assert error.status_code == 400 + + +def test_start_index_is_clamped_to_one(): + # RFC 7644 3.4.2.4: a value less than 1 is interpreted as 1. + start_index, _, error = parse_pagination_params(_request(startIndex="-3")) + assert error is None + assert start_index == 1 + + +def test_negative_count_becomes_zero(): + # Left negative this would reach the queryset as qs[0:-5], and Django raises + # ValueError("Negative indexing is not supported."). + _, count, error = parse_pagination_params(_request(count="-5")) + assert error is None + assert count == 0 + + +def test_count_is_capped_at_the_default(): + _, count, error = parse_pagination_params(_request(count="100000")) + assert error is None + assert count == SCIM_DEFAULT_COUNT