Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions backend/ee/authentication/scim/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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
8 changes: 4 additions & 4 deletions backend/ee/authentication/scim/views/groups.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
7 changes: 4 additions & 3 deletions backend/ee/authentication/scim/views/users.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
68 changes: 68 additions & 0 deletions backend/tests/ee/authentication/scim/test_pagination_params.py
Original file line number Diff line number Diff line change
@@ -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
Loading