From 3a5f36a8e88d1b335a2a5934f5a8ecc18bf53b32 Mon Sep 17 00:00:00 2001 From: wadii Date: Fri, 14 Aug 2026 16:26:48 +0200 Subject: [PATCH 1/2] feat: cohort CSV sync endpoint and cohort summary on segments --- api/cohorts/constants.py | 4 + api/cohorts/dataclasses.py | 25 ++ api/cohorts/exceptions.py | 13 + api/cohorts/metrics.py | 12 + api/cohorts/serializers.py | 58 +++- api/cohorts/services.py | 144 +++++++++- api/cohorts/views.py | 41 ++- api/segments/serializers.py | 21 ++ api/segments/views.py | 1 + api/tests/unit/cohorts/conftest.py | 23 ++ api/tests/unit/cohorts/test_services.py | 239 ++++++++++++++++ api/tests/unit/cohorts/test_views.py | 256 +++++++++++++++++- .../unit/segments/test_unit_segments_views.py | 39 ++- .../observability/_events-catalogue.md | 23 +- .../observability/_metrics-catalogue.md | 16 ++ 15 files changed, 901 insertions(+), 14 deletions(-) create mode 100644 api/cohorts/dataclasses.py create mode 100644 api/cohorts/exceptions.py diff --git a/api/cohorts/constants.py b/api/cohorts/constants.py index 79aa289ab9f5..a1eb9e4b880e 100644 --- a/api/cohorts/constants.py +++ b/api/cohorts/constants.py @@ -1,6 +1,10 @@ COHORT_SYSTEM_TRAIT_KEY_PREFIX = "flagsmith_cohort_" +# Edge identifiers are DynamoDB sort keys, capped at 1024 bytes. +COHORT_IDENTIFIER_MAX_BYTES = 1024 COHORT_MEMBERSHIP_APPLY_BATCH_SIZE = 100 COHORT_MEMBERSHIP_APPLY_MAX_BATCHES_PER_RUN = 10 +COHORT_CSV_MAX_FILE_SIZE_BYTES = 10 * 1024 * 1024 +COHORT_CSV_MEMBERSHIP_CREATE_BATCH_SIZE = 1000 DYNAMODB_THROTTLING_ERROR_CODES = frozenset( { "ProvisionedThroughputExceededException", diff --git a/api/cohorts/dataclasses.py b/api/cohorts/dataclasses.py new file mode 100644 index 000000000000..ad73d27b5669 --- /dev/null +++ b/api/cohorts/dataclasses.py @@ -0,0 +1,25 @@ +from dataclasses import dataclass + + +@dataclass +class CsvIdentifierExtraction: + identifiers: list[str] + empty_count: int + duplicate_count: int + too_long_count: int + + +@dataclass +class CohortCsvIgnoredRows: + empty: int + duplicates: int + too_long: int + + +@dataclass +class CohortCsvSyncResult: + version: int + added: int + removed: int + unchanged: int + ignored: CohortCsvIgnoredRows diff --git a/api/cohorts/exceptions.py b/api/cohorts/exceptions.py new file mode 100644 index 000000000000..6286d4beed6b --- /dev/null +++ b/api/cohorts/exceptions.py @@ -0,0 +1,13 @@ +from rest_framework import status +from rest_framework.exceptions import APIException + +from cohorts.constants import COHORT_CSV_MAX_FILE_SIZE_BYTES + + +class CsvFileTooLargeError(APIException): + status_code = status.HTTP_413_REQUEST_ENTITY_TOO_LARGE + default_detail = ( + "CSV file exceeds the " + f"{COHORT_CSV_MAX_FILE_SIZE_BYTES // (1024 * 1024)}MB size limit." + ) + default_code = "csv_file_too_large" diff --git a/api/cohorts/metrics.py b/api/cohorts/metrics.py index cb9236f69064..22047c3559aa 100644 --- a/api/cohorts/metrics.py +++ b/api/cohorts/metrics.py @@ -7,3 +7,15 @@ "The `operation` label is either `add` or `remove`.", ["operation"], ) + +flagsmith_cohorts_csv_syncs_total = prometheus_client.Counter( + "flagsmith_cohorts_csv_syncs_total", + "Total number of accepted cohort CSV synchronisations, i.e. uploads that " + "yielded at least one valid identifier and enqueued a membership sync.", +) + +flagsmith_cohorts_csv_sync_identifiers = prometheus_client.Histogram( + "flagsmith_cohorts_csv_sync_identifiers", + "Number of unique identifiers extracted per accepted cohort CSV synchronisation.", + buckets=(10, 100, 1_000, 10_000, 100_000, 1_000_000), +) diff --git a/api/cohorts/serializers.py b/api/cohorts/serializers.py index 780216bb9f4a..9e0da4bf8072 100644 --- a/api/cohorts/serializers.py +++ b/api/cohorts/serializers.py @@ -1,9 +1,22 @@ import typing +from django.core.files.uploadedfile import UploadedFile from rest_framework import serializers +from cohorts.constants import COHORT_CSV_MAX_FILE_SIZE_BYTES +from cohorts.exceptions import CsvFileTooLargeError from cohorts.models import Cohort from cohorts.services import create_cohort +from environments.models import Environment +from metadata.serializers import MetadataSerializer, MetadataSerializerMixin +from segments.models import Segment + + +class _SegmentMetadataHandler(MetadataSerializerMixin): + # The mixin derives the metadata content type from Meta.model; cohort + # metadata lives on the managed segment, not on the cohort itself. + class Meta: + model = Segment class CohortSerializer(serializers.ModelSerializer[Cohort]): @@ -11,6 +24,7 @@ class CohortSerializer(serializers.ModelSerializer[Cohort]): description = serializers.CharField( source="segment.description", required=False, allow_null=True ) + metadata = MetadataSerializer(required=False, many=True, write_only=True) class Meta: model = Cohort @@ -19,6 +33,7 @@ class Meta: "uuid", "name", "description", + "metadata", "segment", "source_type", "version", @@ -26,10 +41,51 @@ class Meta: ) read_only_fields = ("segment", "source_type", "version", "created_at") + def validate(self, attrs: dict[str, typing.Any]) -> dict[str, typing.Any]: + attrs = super().validate(attrs) + environment = Environment.objects.get( + api_key=self.context["view"].kwargs["environment_api_key"] + ) + project = environment.project + _SegmentMetadataHandler()._validate_required_metadata( + project.organisation, attrs.get("metadata", []), project + ) + return attrs + def create(self, validated_data: dict[str, typing.Any]) -> Cohort: segment_data = validated_data["segment"] - return create_cohort( + metadata_data = validated_data.pop("metadata", []) + cohort = create_cohort( environment=validated_data["environment"], name=segment_data["name"], description=segment_data.get("description"), ) + if metadata_data: + _SegmentMetadataHandler()._update_metadata(cohort.segment, metadata_data) + return cohort + + +class CohortCsvSyncSerializer(serializers.Serializer): # type: ignore[type-arg] + file = serializers.FileField() + identifier_column = serializers.IntegerField(required=False, default=0, min_value=0) + has_header = serializers.BooleanField(required=False, default=True) + + def validate_file(self, file: UploadedFile) -> UploadedFile: + if file.size and file.size > COHORT_CSV_MAX_FILE_SIZE_BYTES: + # Deliberately not a ValidationError: propagates as a 413. + raise CsvFileTooLargeError() + return file + + +class CohortCsvSyncIgnoredRowsSerializer(serializers.Serializer): # type: ignore[type-arg] + empty = serializers.IntegerField(min_value=0) + duplicates = serializers.IntegerField(min_value=0) + too_long = serializers.IntegerField(min_value=0) + + +class CohortCsvSyncResultSerializer(serializers.Serializer): # type: ignore[type-arg] + version = serializers.IntegerField(min_value=0) + added = serializers.IntegerField(min_value=0) + removed = serializers.IntegerField(min_value=0) + unchanged = serializers.IntegerField(min_value=0) + ignored = CohortCsvSyncIgnoredRowsSerializer() diff --git a/api/cohorts/services.py b/api/cohorts/services.py index c2ba9f3bd111..2e7bdcfe4aa7 100644 --- a/api/cohorts/services.py +++ b/api/cohorts/services.py @@ -1,3 +1,5 @@ +import csv +import io import typing import structlog @@ -5,9 +7,23 @@ from django.db.models import QuerySet from django.utils import timezone from flag_engine.segments.constants import IS_SET +from rest_framework.exceptions import ValidationError -from cohorts.constants import COHORT_MEMBERSHIP_APPLY_BATCH_SIZE -from cohorts.metrics import flagsmith_cohorts_membership_deltas_applied_total +from cohorts.constants import ( + COHORT_CSV_MEMBERSHIP_CREATE_BATCH_SIZE, + COHORT_IDENTIFIER_MAX_BYTES, + COHORT_MEMBERSHIP_APPLY_BATCH_SIZE, +) +from cohorts.dataclasses import ( + CohortCsvIgnoredRows, + CohortCsvSyncResult, + CsvIdentifierExtraction, +) +from cohorts.metrics import ( + flagsmith_cohorts_csv_sync_identifiers, + flagsmith_cohorts_csv_syncs_total, + flagsmith_cohorts_membership_deltas_applied_total, +) from cohorts.models import Cohort, CohortMembership, CohortMembershipState from core.dataclasses import AuthorData from environments.dynamodb import DynamoIdentityWrapper @@ -111,6 +127,130 @@ def create_cohort( return cohort +def extract_identifiers_from_csv( + file: typing.IO[bytes], + *, + identifier_column: int = 0, + has_header: bool = True, +) -> CsvIdentifierExtraction: + # The upload size cap keeps a full read cheap. + text = io.StringIO(file.read().decode("utf-8-sig", errors="replace"), newline="") + reader = csv.reader(text) + seen: set[str] = set() + identifiers: list[str] = [] + empty_count = duplicate_count = too_long_count = 0 + try: + for row_number, row in enumerate(reader): + if has_header and row_number == 0: + continue + if not row: + continue + value = ( + row[identifier_column].strip() if identifier_column < len(row) else "" + ) + if not value: + empty_count += 1 + elif len(value.encode()) > COHORT_IDENTIFIER_MAX_BYTES: + too_long_count += 1 + elif value in seen: + duplicate_count += 1 + else: + seen.add(value) + identifiers.append(value) + except csv.Error as exc: + raise ValidationError({"file": "Could not parse the CSV file."}) from exc + return CsvIdentifierExtraction( + identifiers=identifiers, + empty_count=empty_count, + duplicate_count=duplicate_count, + too_long_count=too_long_count, + ) + + +def sync_cohort_memberships_from_csv( + *, + cohort: Cohort, + file: typing.IO[bytes], + identifier_column: int = 0, + has_header: bool = True, +) -> CohortCsvSyncResult: + from cohorts.tasks import apply_cohort_membership_deltas + + extraction = extract_identifiers_from_csv( + file, identifier_column=identifier_column, has_header=has_header + ) + if not extraction.identifiers: + raise ValidationError({"file": "No valid identifiers found in the CSV file."}) + + incoming = set(extraction.identifiers) + added = removed = unchanged = 0 + with transaction.atomic(): + # Serialise concurrent syncs of the same cohort. + locked_cohort = Cohort.objects.select_for_update().get(id=cohort.id) + existing = { + membership.identifier: membership + for membership in CohortMembership.objects.filter(cohort=cohort).only( + "id", "identifier", "state" + ) + } + CohortMembership.objects.bulk_create( + [ + CohortMembership(cohort=cohort, identifier=identifier) + for identifier in extraction.identifiers + if identifier not in existing + ], + batch_size=COHORT_CSV_MEMBERSHIP_CREATE_BATCH_SIZE, + ) + added += len(incoming - existing.keys()) + + readd_ids: list[int] = [] + remove_ids: list[int] = [] + for identifier, membership in existing.items(): + if identifier in incoming: + if membership.state == CohortMembershipState.PENDING_REMOVE: + readd_ids.append(membership.id) + else: + unchanged += 1 + elif membership.state != CohortMembershipState.PENDING_REMOVE: + # A pending add may have had its trait written by a concurrent + # applier run, so drain it via pending remove, never delete. + remove_ids.append(membership.id) + + added += CohortMembership.objects.filter(id__in=readd_ids).update( + state=CohortMembershipState.PENDING_ADD, updated_at=timezone.now() + ) + removed += CohortMembership.objects.filter(id__in=remove_ids).update( + state=CohortMembershipState.PENDING_REMOVE, updated_at=timezone.now() + ) + + locked_cohort.version += 1 + locked_cohort.save(update_fields=["version"]) + apply_cohort_membership_deltas.delay(kwargs={"cohort_id": cohort.id}) + + flagsmith_cohorts_csv_syncs_total.inc() + flagsmith_cohorts_csv_sync_identifiers.observe(len(incoming)) + logger.info( + "csv.synced", + cohort__id=cohort.id, + environment__id=cohort.environment_id, + cohort__version=locked_cohort.version, + adds__count=added, + removes__count=removed, + unchanged__count=unchanged, + ) + return CohortCsvSyncResult( + version=locked_cohort.version, + added=added, + removed=removed, + unchanged=unchanged, + ignored=CohortCsvIgnoredRows( + empty=extraction.empty_count, + duplicates=extraction.duplicate_count, + too_long=extraction.too_long_count, + ), + ) + + def edge_sync_enabled(project: "Project") -> bool: return bool(project.enable_dynamo_db and DynamoIdentityWrapper().is_enabled) diff --git a/api/cohorts/views.py b/api/cohorts/views.py index eda7c209a3d9..c86454fcf054 100644 --- a/api/cohorts/views.py +++ b/api/cohorts/views.py @@ -1,14 +1,21 @@ from django.db.models import QuerySet from drf_spectacular.utils import extend_schema, extend_schema_view from rest_framework import mixins, status +from rest_framework.decorators import action +from rest_framework.parsers import MultiPartParser from rest_framework.permissions import IsAuthenticated from rest_framework.request import Request from rest_framework.response import Response +from api.serializers import ErrorSerializer from cohorts import services from cohorts.models import Cohort from cohorts.permissions import CohortPermission, CohortPlanPermission -from cohorts.serializers import CohortSerializer +from cohorts.serializers import ( + CohortCsvSyncResultSerializer, + CohortCsvSyncSerializer, + CohortSerializer, +) from environments.views import NestedEnvironmentViewSet from projects.exceptions import DynamoNotEnabledError @@ -62,3 +69,35 @@ def get_queryset(self) -> QuerySet[Cohort]: def destroy(self, request: Request, *args: object, **kwargs: object) -> Response: services.delete_cohort(self.get_object()) return Response(status=status.HTTP_202_ACCEPTED) + + @extend_schema( + description=( + "Replace the cohort's members with the identifiers found in the " + "uploaded CSV file and trigger a sync to identity data. " + "`identifier_column` is the 0-based index of the column holding " + "the identifiers; `has_header` skips the first row when true." + ), + request=CohortCsvSyncSerializer, + responses={ + 202: CohortCsvSyncResultSerializer, + 400: ErrorSerializer, + 413: ErrorSerializer, + }, + ) + @action( + detail=True, + methods=["post"], + url_path="sync-csv", + parser_classes=[MultiPartParser], + ) + def sync_csv(self, request: Request, *args: object, **kwargs: object) -> Response: + cohort = self.get_object() + serializer = CohortCsvSyncSerializer(data=request.data) + serializer.is_valid(raise_exception=True) + result = services.sync_cohort_memberships_from_csv( + cohort=cohort, **serializer.validated_data + ) + return Response( + CohortCsvSyncResultSerializer(result).data, + status=status.HTTP_202_ACCEPTED, + ) diff --git a/api/segments/serializers.py b/api/segments/serializers.py index 9e9c36fd4bb6..1d03940938f9 100644 --- a/api/segments/serializers.py +++ b/api/segments/serializers.py @@ -8,6 +8,7 @@ from rest_framework import serializers from rest_framework.exceptions import ValidationError +from cohorts.models import Cohort from edge_api.utils import is_edge_enabled from metadata.serializers import MetadataSerializer, MetadataSerializerMixin from projects.models import Project @@ -102,10 +103,23 @@ class Meta: ] +class _SegmentCohortSerializer(serializers.ModelSerializer[Cohort]): + class Meta: + model = Cohort + fields = [ + "id", + "environment", + "source_type", + "version", + "deletion_requested_at", + ] + + class SegmentSerializer(MetadataSerializerMixin, WritableNestedModelSerializer): rules = SegmentRuleSerializer(many=True, required=True, allow_empty=False) metadata = MetadataSerializer(required=False, many=True) membership_counts = SegmentMembershipCountSerializer(many=True, read_only=True) + cohort = serializers.SerializerMethodField() def __init__(self, *args: Any, **kwargs: Any) -> None: """ @@ -140,6 +154,7 @@ class Meta: "metadata", "membership_counts", "managed_by", + "cohort", ] read_only_fields = [ "managed_by", @@ -147,6 +162,12 @@ class Meta: "project", ] + @extend_schema_field(_SegmentCohortSerializer(allow_null=True)) + def get_cohort(self, segment: Segment) -> dict[str, Any] | None: + # next() over the relation keeps the list view's prefetch cache warm. + cohort = next(iter(segment.cohorts.all()), None) + return _SegmentCohortSerializer(cohort).data if cohort else None + def to_internal_value(self, data: dict[str, Any]) -> Any: self._validate_rules_depth(data.get("rules", [])) return super().to_internal_value(data) diff --git a/api/segments/views.py b/api/segments/views.py index 51a80de92eb7..116b34ce2aac 100644 --- a/api/segments/views.py +++ b/api/segments/views.py @@ -104,6 +104,7 @@ def get_queryset(self): # type: ignore[no-untyped-def] # TODO: at the moment, the UI only shows the name and description of the segment in the list view. # we shouldn't return all of the rules and conditions in the list view. queryset = queryset.prefetch_related( + "cohorts", "membership_counts", "rules", "rules__conditions", diff --git a/api/tests/unit/cohorts/conftest.py b/api/tests/unit/cohorts/conftest.py index adf99e3cbeb7..0e0859c503a3 100644 --- a/api/tests/unit/cohorts/conftest.py +++ b/api/tests/unit/cohorts/conftest.py @@ -1,7 +1,13 @@ import pytest +from django.contrib.contenttypes.models import ContentType from cohorts.models import Cohort from environments.models import Environment +from metadata.models import ( + MetadataField, + MetadataModelField, + MetadataModelFieldRequirement, +) from projects.models import Project from segments.models import Segment @@ -12,6 +18,23 @@ def cohort(environment: Environment, segment: Segment) -> Cohort: return cohort +@pytest.fixture() +def required_segment_metadata_field_for_dynamo_project( + a_metadata_field: MetadataField, + dynamo_enabled_project: Project, +) -> MetadataModelField: + model_field: MetadataModelField = MetadataModelField.objects.create( + field=a_metadata_field, + content_type=ContentType.objects.get_for_model(Segment), + ) + MetadataModelFieldRequirement.objects.create( + content_type=ContentType.objects.get_for_model(Project), + object_id=dynamo_enabled_project.id, + model_field=model_field, + ) + return model_field + + @pytest.fixture() def edge_cohort( dynamo_enabled_project: Project, diff --git a/api/tests/unit/cohorts/test_services.py b/api/tests/unit/cohorts/test_services.py index 5bf7eb51b3cf..ac21d85482b3 100644 --- a/api/tests/unit/cohorts/test_services.py +++ b/api/tests/unit/cohorts/test_services.py @@ -1,12 +1,18 @@ +import io + +import pytest from flag_engine.segments.constants import IS_SET from pytest_mock import MockerFixture from pytest_structlog import StructuredLogCapture +from rest_framework.exceptions import ValidationError from cohorts.models import Cohort, CohortMembership, CohortMembershipState from cohorts.services import ( apply_pending_memberships, create_cohort, delete_cohort, + extract_identifiers_from_csv, + sync_cohort_memberships_from_csv, ) from environments.dynamodb import DynamoIdentityWrapper from environments.models import Environment @@ -180,3 +186,236 @@ def test_delete_cohort__edge__drains_traits_then_deletes( assert not Cohort.objects.filter(id=edge_cohort.id).exists() assert not CohortMembership.objects.filter(cohort_id=edge_cohort.id).exists() assert log.has("cohort.deletion_requested", cohort__id=edge_cohort.id) + + +@pytest.mark.parametrize( + "content, identifier_column, has_header, expected_identifiers, " + "expected_empty, expected_duplicates, expected_too_long", + [ + pytest.param( + b"identity\nuser-1\nuser-2\n", + 0, + True, + ["user-1", "user-2"], + 0, + 0, + 0, + id="header-single-column", + ), + pytest.param( + b"user-1\nuser-2\n", + 0, + False, + ["user-1", "user-2"], + 0, + 0, + 0, + id="no-header", + ), + pytest.param( + b"identity,email,plan\nuser-1,a@example.com,free\nuser-2,b@example.com,pro\n", + 1, + True, + ["a@example.com", "b@example.com"], + 0, + 0, + 0, + id="identifier-column-index", + ), + pytest.param( + b'"Doe, Jane"\n"say ""hi"""\n', + 0, + False, + ["Doe, Jane", 'say "hi"'], + 0, + 0, + 0, + id="quoted-values", + ), + pytest.param( + b"identity\nuser-1\n\n \nuser-1\nuser-2\n", + 0, + True, + ["user-1", "user-2"], + 1, + 1, + 0, + id="empties-and-duplicates-counted-blank-lines-skipped", + ), + pytest.param( + b"identity,plan\nuser-1\n", + 1, + True, + [], + 1, + 0, + 0, + id="column-missing-from-row-counted-empty", + ), + pytest.param( + b"identity\n" + b"x" * 1025 + b"\nuser-1\n", + 0, + True, + ["user-1"], + 0, + 0, + 1, + id="over-long-identifier-ignored", + ), + pytest.param( + ("identity\n" + "é" * 513 + "\nuser-1\n").encode(), + 0, + True, + ["user-1"], + 0, + 0, + 1, + id="identifier-over-byte-limit-ignored", + ), + pytest.param( + b"identity\n", + 0, + True, + [], + 0, + 0, + 0, + id="header-only", + ), + pytest.param( + b"\xef\xbb\xbfidentity\nuser-1\n", + 0, + True, + ["user-1"], + 0, + 0, + 0, + id="utf8-bom-stripped", + ), + ], +) +def test_extract_identifiers_from_csv__varied_content__extracts_expected( + content: bytes, + identifier_column: int, + has_header: bool, + expected_identifiers: list[str], + expected_empty: int, + expected_duplicates: int, + expected_too_long: int, +) -> None: + # Given + file = io.BytesIO(content) + + # When + extraction = extract_identifiers_from_csv( + file, identifier_column=identifier_column, has_header=has_header + ) + + # Then + assert extraction.identifiers == expected_identifiers + assert extraction.empty_count == expected_empty + assert extraction.duplicate_count == expected_duplicates + assert extraction.too_long_count == expected_too_long + + +def test_extract_identifiers_from_csv__unparseable_content__raises_validation_error() -> ( + None +): + # Given + file = io.BytesIO(b"identity\n" + b"x" * 200_000 + b"\n") + + # When / Then + with pytest.raises(ValidationError): + extract_identifiers_from_csv(file) + + +def test_sync_cohort_memberships_from_csv__first_upload__creates_pending_adds( + cohort: Cohort, + log: StructuredLogCapture, +) -> None: + # Given + file = io.BytesIO(b"identity\nuser-1\nuser-2\n\nuser-2\n") + + # When + result = sync_cohort_memberships_from_csv(cohort=cohort, file=file) + + # Then + assert result.version == 1 + assert result.added == 2 + assert result.removed == 0 + assert result.unchanged == 0 + assert result.ignored.empty == 0 + assert result.ignored.duplicates == 1 + assert result.ignored.too_long == 0 + memberships = CohortMembership.objects.filter(cohort=cohort) + assert {m.identifier for m in memberships} == {"user-1", "user-2"} + assert all(m.state == CohortMembershipState.PENDING_ADD for m in memberships) + cohort.refresh_from_db() + assert cohort.version == 1 + assert log.has( + "csv.synced", + cohort__id=cohort.id, + environment__id=cohort.environment_id, + cohort__version=1, + adds__count=2, + removes__count=0, + unchanged__count=0, + ) + + +def test_sync_cohort_memberships_from_csv__reupload__computes_membership_delta( + cohort: Cohort, +) -> None: + # Given + CohortMembership.objects.create( + cohort=cohort, identifier="stay", state=CohortMembershipState.APPLIED + ) + CohortMembership.objects.create( + cohort=cohort, identifier="leave", state=CohortMembershipState.APPLIED + ) + CohortMembership.objects.create( + cohort=cohort, identifier="comeback", state=CohortMembershipState.PENDING_REMOVE + ) + CohortMembership.objects.create( + cohort=cohort, identifier="ghost", state=CohortMembershipState.PENDING_ADD + ) + file = io.BytesIO(b"identity\nstay\ncomeback\nnew\n") + + # When + result = sync_cohort_memberships_from_csv(cohort=cohort, file=file) + + # Then + assert result.version == 1 + assert result.added == 2 + assert result.removed == 2 + assert result.unchanged == 1 + states = { + m.identifier: m.state for m in CohortMembership.objects.filter(cohort=cohort) + } + assert states == { + "stay": CohortMembershipState.APPLIED, + "leave": CohortMembershipState.PENDING_REMOVE, + "comeback": CohortMembershipState.PENDING_ADD, + "ghost": CohortMembershipState.PENDING_REMOVE, + "new": CohortMembershipState.PENDING_ADD, + } + + +def test_sync_cohort_memberships_from_csv__edge_cohort__applies_traits( + edge_cohort: Cohort, + dynamodb_identity_wrapper: DynamoIdentityWrapper, +) -> None: + # Given + file = io.BytesIO(b"identity\njoiner\n") + api_key = edge_cohort.environment.api_key + + # When + result = sync_cohort_memberships_from_csv(cohort=edge_cohort, file=file) + + # Then + assert result.added == 1 + document = dynamodb_identity_wrapper.get_item(f"{api_key}_joiner") + assert document is not None + assert document["system_traits"] == {edge_cohort.system_trait_key: True} + membership = CohortMembership.objects.get(cohort=edge_cohort) + assert membership.state == CohortMembershipState.APPLIED diff --git a/api/tests/unit/cohorts/test_views.py b/api/tests/unit/cohorts/test_views.py index 488dfe7eae5c..efcff0d493cd 100644 --- a/api/tests/unit/cohorts/test_views.py +++ b/api/tests/unit/cohorts/test_views.py @@ -4,15 +4,17 @@ VIEW_ENVIRONMENT, ) from common.projects.permissions import MANAGE_SEGMENTS +from django.core.files.uploadedfile import SimpleUploadedFile from django.urls import reverse from django.utils import timezone from pytest_mock import MockerFixture from rest_framework import status from rest_framework.test import APIClient -from cohorts.models import Cohort +from cohorts.models import Cohort, CohortMembership, CohortMembershipState from environments.dynamodb import DynamoIdentityWrapper from environments.models import Environment +from metadata.models import Metadata, MetadataModelField from organisations.models import Subscription from projects.models import Project from segments.models import Segment @@ -262,3 +264,255 @@ def test_create_cohort__non_edge_project__returns_400( # Then assert response.status_code == status.HTTP_400_BAD_REQUEST assert response.json()["detail"] == "Dynamo DB is not enabled for this project" + + +def test_create_cohort__with_metadata__attaches_metadata_to_segment( + staff_client: APIClient, + dynamo_enabled_project: Project, + dynamo_enabled_project_environment_one: Environment, + dynamodb_identity_wrapper: DynamoIdentityWrapper, + required_segment_metadata_field_for_dynamo_project: MetadataModelField, + with_project_permissions: WithProjectPermissionsCallable, + with_environment_permissions: WithEnvironmentPermissionsCallable, +) -> None: + # Given + with_project_permissions( # type: ignore[call-arg] + [MANAGE_SEGMENTS], project_id=dynamo_enabled_project.id + ) + with_environment_permissions( # type: ignore[call-arg] + [VIEW_ENVIRONMENT, MANAGE_SEGMENT_OVERRIDES], + environment_id=dynamo_enabled_project_environment_one.id, + ) + url = reverse( + "api-v1:environments:cohorts:cohorts-list", + args=[dynamo_enabled_project_environment_one.api_key], + ) + + # When + response = staff_client.post( + url, + data={ + "name": "Beta users", + "metadata": [ + { + "model_field": required_segment_metadata_field_for_dynamo_project.id, + "field_value": 10, + }, + ], + }, + format="json", + ) + + # Then + assert response.status_code == status.HTTP_201_CREATED + cohort = Cohort.objects.get(id=response.json()["id"]) + metadata = Metadata.objects.get( + model_field=required_segment_metadata_field_for_dynamo_project + ) + assert metadata.object_id == cohort.segment_id + assert metadata.field_value == "10" + + +def test_create_cohort__missing_required_metadata__returns_400( + staff_client: APIClient, + dynamo_enabled_project: Project, + dynamo_enabled_project_environment_one: Environment, + dynamodb_identity_wrapper: DynamoIdentityWrapper, + required_segment_metadata_field_for_dynamo_project: MetadataModelField, + with_project_permissions: WithProjectPermissionsCallable, + with_environment_permissions: WithEnvironmentPermissionsCallable, +) -> None: + # Given + with_project_permissions( # type: ignore[call-arg] + [MANAGE_SEGMENTS], project_id=dynamo_enabled_project.id + ) + with_environment_permissions( # type: ignore[call-arg] + [VIEW_ENVIRONMENT, MANAGE_SEGMENT_OVERRIDES], + environment_id=dynamo_enabled_project_environment_one.id, + ) + url = reverse( + "api-v1:environments:cohorts:cohorts-list", + args=[dynamo_enabled_project_environment_one.api_key], + ) + + # When + response = staff_client.post(url, data={"name": "Beta users"}, format="json") + + # Then + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json()["metadata"] == ["Missing required metadata field: a"] + assert not Cohort.objects.exists() + + +def test_sync_csv__staff_with_manage_segments__returns_202_with_counts( + staff_client: APIClient, + dynamo_enabled_project: Project, + edge_cohort: Cohort, + dynamodb_identity_wrapper: DynamoIdentityWrapper, + with_project_permissions: WithProjectPermissionsCallable, + with_environment_permissions: WithEnvironmentPermissionsCallable, +) -> None: + # Given + with_project_permissions( # type: ignore[call-arg] + [MANAGE_SEGMENTS], project_id=dynamo_enabled_project.id + ) + with_environment_permissions( # type: ignore[call-arg] + [VIEW_ENVIRONMENT, MANAGE_SEGMENT_OVERRIDES], + environment_id=edge_cohort.environment_id, + ) + url = reverse( + "api-v1:environments:cohorts:cohorts-sync-csv", + args=[edge_cohort.environment.api_key, edge_cohort.id], + ) + file = SimpleUploadedFile( + "identities.csv", + b"identity,email\nuser-1,a@example.com\nuser-2,b@example.com\nuser-3,\n", + content_type="text/csv", + ) + + # When + response = staff_client.post( + url, + data={"file": file, "identifier_column": 1, "has_header": True}, + format="multipart", + ) + + # Then + assert response.status_code == status.HTTP_202_ACCEPTED + assert response.json() == { + "version": 1, + "added": 2, + "removed": 0, + "unchanged": 0, + "ignored": {"empty": 1, "duplicates": 0, "too_long": 0}, + } + memberships = CohortMembership.objects.filter(cohort=edge_cohort) + assert {m.identifier for m in memberships} == { + "a@example.com", + "b@example.com", + } + assert all(m.state == CohortMembershipState.APPLIED for m in memberships) + edge_cohort.refresh_from_db() + assert edge_cohort.version == 1 + + +def test_sync_csv__without_permission__returns_403( + staff_client: APIClient, + edge_cohort: Cohort, + dynamodb_identity_wrapper: DynamoIdentityWrapper, + with_environment_permissions: WithEnvironmentPermissionsCallable, +) -> None: + # Given + with_environment_permissions( # type: ignore[call-arg] + [VIEW_ENVIRONMENT], environment_id=edge_cohort.environment_id + ) + url = reverse( + "api-v1:environments:cohorts:cohorts-sync-csv", + args=[edge_cohort.environment.api_key, edge_cohort.id], + ) + file = SimpleUploadedFile("identities.csv", b"user-1\n", content_type="text/csv") + + # When + response = staff_client.post(url, data={"file": file}, format="multipart") + + # Then + assert response.status_code == status.HTTP_403_FORBIDDEN + assert not CohortMembership.objects.exists() + + +def test_sync_csv__file_over_size_limit__returns_413( + staff_client: APIClient, + dynamo_enabled_project: Project, + edge_cohort: Cohort, + dynamodb_identity_wrapper: DynamoIdentityWrapper, + with_project_permissions: WithProjectPermissionsCallable, + with_environment_permissions: WithEnvironmentPermissionsCallable, + mocker: MockerFixture, +) -> None: + # Given + mocker.patch("cohorts.serializers.COHORT_CSV_MAX_FILE_SIZE_BYTES", 10) + with_project_permissions( # type: ignore[call-arg] + [MANAGE_SEGMENTS], project_id=dynamo_enabled_project.id + ) + with_environment_permissions( # type: ignore[call-arg] + [VIEW_ENVIRONMENT, MANAGE_SEGMENT_OVERRIDES], + environment_id=edge_cohort.environment_id, + ) + url = reverse( + "api-v1:environments:cohorts:cohorts-sync-csv", + args=[edge_cohort.environment.api_key, edge_cohort.id], + ) + file = SimpleUploadedFile( + "identities.csv", b"identity\nuser-1\n", content_type="text/csv" + ) + + # When + response = staff_client.post(url, data={"file": file}, format="multipart") + + # Then + assert response.status_code == status.HTTP_413_REQUEST_ENTITY_TOO_LARGE + assert not CohortMembership.objects.exists() + + +def test_sync_csv__no_valid_identifiers__returns_400( + staff_client: APIClient, + dynamo_enabled_project: Project, + edge_cohort: Cohort, + dynamodb_identity_wrapper: DynamoIdentityWrapper, + with_project_permissions: WithProjectPermissionsCallable, + with_environment_permissions: WithEnvironmentPermissionsCallable, +) -> None: + # Given + with_project_permissions( # type: ignore[call-arg] + [MANAGE_SEGMENTS], project_id=dynamo_enabled_project.id + ) + with_environment_permissions( # type: ignore[call-arg] + [VIEW_ENVIRONMENT, MANAGE_SEGMENT_OVERRIDES], + environment_id=edge_cohort.environment_id, + ) + url = reverse( + "api-v1:environments:cohorts:cohorts-sync-csv", + args=[edge_cohort.environment.api_key, edge_cohort.id], + ) + file = SimpleUploadedFile("identities.csv", b"identity\n", content_type="text/csv") + + # When + response = staff_client.post(url, data={"file": file}, format="multipart") + + # Then + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json() == {"file": "No valid identifiers found in the CSV file."} + assert not CohortMembership.objects.exists() + edge_cohort.refresh_from_db() + assert edge_cohort.version == 0 + + +def test_sync_csv__deletion_requested_cohort__returns_404( + staff_client: APIClient, + dynamo_enabled_project: Project, + edge_cohort: Cohort, + dynamodb_identity_wrapper: DynamoIdentityWrapper, + with_project_permissions: WithProjectPermissionsCallable, + with_environment_permissions: WithEnvironmentPermissionsCallable, +) -> None: + # Given + with_project_permissions( # type: ignore[call-arg] + [MANAGE_SEGMENTS], project_id=dynamo_enabled_project.id + ) + with_environment_permissions( # type: ignore[call-arg] + [VIEW_ENVIRONMENT, MANAGE_SEGMENT_OVERRIDES], + environment_id=edge_cohort.environment_id, + ) + edge_cohort.deletion_requested_at = timezone.now() + edge_cohort.save(update_fields=["deletion_requested_at"]) + url = reverse( + "api-v1:environments:cohorts:cohorts-sync-csv", + args=[edge_cohort.environment.api_key, edge_cohort.id], + ) + file = SimpleUploadedFile("identities.csv", b"user-1\n", content_type="text/csv") + + # When + response = staff_client.post(url, data={"file": file}, format="multipart") + + # Then + assert response.status_code == status.HTTP_404_NOT_FOUND diff --git a/api/tests/unit/segments/test_unit_segments_views.py b/api/tests/unit/segments/test_unit_segments_views.py index 9be6e42292f7..b28b52679b16 100644 --- a/api/tests/unit/segments/test_unit_segments_views.py +++ b/api/tests/unit/segments/test_unit_segments_views.py @@ -111,6 +111,7 @@ def test_create_segment__valid_rules__creates_segment_with_rules( "metadata": [], "membership_counts": [], "managed_by": "", + "cohort": None, "rules": [ { "id": mocker.ANY, @@ -675,8 +676,8 @@ def test_get_segment_by_uuid__existing_segment__returns_segment_data( # type: i @pytest.mark.parametrize( "client, num_queries", [ - (lazy_fixture("admin_master_api_key_client"), 13), - (lazy_fixture("admin_client"), 15), + (lazy_fixture("admin_master_api_key_client"), 14), + (lazy_fixture("admin_client"), 16), ], ) def test_list_segments__without_rbac__expected_num_queries( @@ -732,8 +733,8 @@ def test_list_segments__system_segment_exists__excludes_system_segment( @pytest.mark.parametrize( "client, num_queries", [ - (lazy_fixture("admin_master_api_key_client"), 13), - (lazy_fixture("admin_client"), 16), + (lazy_fixture("admin_master_api_key_client"), 14), + (lazy_fixture("admin_client"), 17), ], ) def test_list_segments__with_rbac__expected_num_queries( @@ -1103,6 +1104,7 @@ def test_update_segment__valid_rules__updates_segment_with_rules( "metadata": [], "membership_counts": [], "managed_by": "", + "cohort": None, "rules": [ { "id": mocker.ANY, @@ -2040,6 +2042,7 @@ def test_update_segment__whitelisted_segment_exceeds_max_conditions__returns_200 "metadata": [], "membership_counts": [], "managed_by": "", + "cohort": None, "rules": [ { "id": mocker.ANY, @@ -2224,6 +2227,7 @@ def test_clone_segment__valid_name__returns_cloned_segment( "metadata": [], "membership_counts": [], "managed_by": "", + "cohort": None, "rules": [], # TODO: Should contain rules as per https://github.com/Flagsmith/flagsmith/issues/7818 } ) @@ -2400,3 +2404,30 @@ def test_clone_segment__cohort_managed__returns_403( # Then assert response.status_code == status.HTTP_403_FORBIDDEN assert Segment.objects.count() == 1 + + +def test_list_segments__cohort_managed__returns_cohort_summary( + admin_client: APIClient, + project: Project, + segment: Segment, + environment: Environment, +) -> None: + # Given + cohort = Cohort.objects.create(environment=environment, segment=segment) + plain_segment = Segment.objects.create(name="plain", project=project) + url = reverse("api-v1:projects:project-segments-list", args=[project.id]) + + # When + response = admin_client.get(url) + + # Then + assert response.status_code == status.HTTP_200_OK + results = {result["id"]: result for result in response.json()["results"]} + assert results[segment.id]["cohort"] == { + "deletion_requested_at": None, + "environment": environment.id, + "id": cohort.id, + "source_type": "csv", + "version": 0, + } + assert results[plain_segment.id]["cohort"] is None diff --git a/docs/docs/deployment-self-hosting/observability/_events-catalogue.md b/docs/docs/deployment-self-hosting/observability/_events-catalogue.md index 4d61863f5029..52e114079969 100644 --- a/docs/docs/deployment-self-hosting/observability/_events-catalogue.md +++ b/docs/docs/deployment-self-hosting/observability/_events-catalogue.md @@ -74,7 +74,7 @@ Attributes: ### `cohorts.cohort.created` Logged at `info` from: - - `api/cohorts/services.py:103` + - `api/cohorts/services.py:119` Attributes: - `cohort.id` @@ -86,7 +86,7 @@ Attributes: ### `cohorts.cohort.deleted` Logged at `info` from: - - `api/cohorts/services.py:140` + - `api/cohorts/services.py:280` Attributes: - `cohort.id` @@ -95,16 +95,29 @@ Attributes: ### `cohorts.cohort.deletion_requested` Logged at `info` from: - - `api/cohorts/services.py:124` + - `api/cohorts/services.py:264` Attributes: - `cohort.id` - `environment.id` +### `cohorts.csv.synced` + +Logged at `info` from: + - `api/cohorts/services.py:232` + +Attributes: + - `adds.count` + - `cohort.id` + - `cohort.version` + - `environment.id` + - `removes.count` + - `unchanged.count` + ### `cohorts.membership.applied` Logged at `info` from: - - `api/cohorts/services.py:72` + - `api/cohorts/services.py:88` Attributes: - `adds.count` @@ -607,7 +620,7 @@ Attributes: ### `segments.serializers.segment_revision_created` Logged at `info` from: - - `api/segments/serializers.py:185` + - `api/segments/serializers.py:206` Attributes: - `revision_id` diff --git a/docs/docs/deployment-self-hosting/observability/_metrics-catalogue.md b/docs/docs/deployment-self-hosting/observability/_metrics-catalogue.md index 203a40345f6d..f134a2609cdb 100644 --- a/docs/docs/deployment-self-hosting/observability/_metrics-catalogue.md +++ b/docs/docs/deployment-self-hosting/observability/_metrics-catalogue.md @@ -9,6 +9,22 @@ Labels: - `ci_commit_sha` - `version` +### `flagsmith_cohorts_csv_sync_identifiers` + +Histogram. + +Number of unique identifiers extracted per accepted cohort CSV synchronisation. + +Labels: + +### `flagsmith_cohorts_csv_syncs` + +Counter. + +Total number of accepted cohort CSV synchronisations, i.e. uploads that yielded at least one valid identifier and enqueued a membership sync. + +Labels: + ### `flagsmith_cohorts_membership_deltas_applied` Counter. From fd0adac1d4b0a63412bac72f697e73e7551b803d Mon Sep 17 00:00:00 2001 From: "flagsmith-engineering[bot]" Date: Fri, 14 Aug 2026 15:57:23 +0000 Subject: [PATCH 2/2] chore: Update documentation artefacts --- mcp/src/flagsmith_mcp/openapi.json | 48 ++++++++++ openapi.yaml | 142 +++++++++++++++++++++++++++++ 2 files changed, 190 insertions(+) diff --git a/mcp/src/flagsmith_mcp/openapi.json b/mcp/src/flagsmith_mcp/openapi.json index cdb4319246ba..7a9cedf9ee00 100644 --- a/mcp/src/flagsmith_mcp/openapi.json +++ b/mcp/src/flagsmith_mcp/openapi.json @@ -6925,6 +6925,17 @@ } ], "readOnly": true + }, + "cohort": { + "oneOf": [ + { + "$ref": "#/components/schemas/_SegmentCohort" + }, + { + "type": "null" + } + ], + "readOnly": true } }, "required": [ @@ -6991,6 +7002,13 @@ "NONE" ] }, + "SourceTypeEnum": { + "description": "* `csv` - CSV", + "type": "string", + "enum": [ + "csv" + ] + }, "StageAction": { "type": "object", "properties": { @@ -7753,6 +7771,36 @@ "required": [ "type" ] + }, + "_SegmentCohort": { + "type": "object", + "properties": { + "id": { + "type": "integer", + "readOnly": true + }, + "environment": { + "type": "integer" + }, + "source_type": { + "$ref": "#/components/schemas/SourceTypeEnum" + }, + "version": { + "type": "integer", + "maximum": 2147483647, + "minimum": 0 + }, + "deletion_requested_at": { + "type": [ + "string", + "null" + ], + "format": "date-time" + } + }, + "required": [ + "environment" + ] } }, "securitySchemes": { diff --git a/openapi.yaml b/openapi.yaml index f1dd8210d994..54d42105ac14 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -2149,6 +2149,53 @@ paths: tags: - Environments x-flagsmith-minimum-plan: START_UP + '/api/v1/environments/{environment_api_key}/cohorts/{cohort_id}/sync-csv/': + post: + operationId: api_v1_environments_cohorts_sync_csv_create + description: Replace the cohort's members with the identifiers found in the uploaded CSV file and trigger a sync to identity data. `identifier_column` is the 0-based index of the column holding the identifiers; `has_header` skips the first row when true. + parameters: + - name: cohort_id + in: path + description: A unique integer value identifying this cohort. + required: true + schema: + type: integer + - name: environment_api_key + in: path + required: true + schema: + type: string + requestBody: + required: true + content: + multipart/form-data: + schema: + $ref: '#/components/schemas/CohortCsvSync' + responses: + '202': + description: '' + content: + application/json: + schema: + $ref: '#/components/schemas/CohortCsvSyncResult' + '400': + description: '' + content: + application/json: + schema: + $ref: '#/components/schemas/Error' + '413': + description: '' + content: + application/json: + schema: + $ref: '#/components/schemas/Error' + security: + - tokenAuth: [] + - Master API Key: [] + tags: + - Environments + x-flagsmith-minimum-plan: START_UP '/api/v1/environments/{environment_api_key}/create-change-request/': post: operationId: create_environment_feature_change_request @@ -18853,6 +18900,11 @@ components: - $ref: '#/components/schemas/ManagedByEnum' - $ref: '#/components/schemas/BlankEnum' readOnly: true + cohort: + oneOf: + - $ref: '#/components/schemas/_SegmentCohort' + - type: 'null' + readOnly: true change_request: type: - integer @@ -19007,6 +19059,11 @@ components: type: - string - 'null' + metadata: + type: array + items: + $ref: '#/components/schemas/Metadata' + writeOnly: true segment: type: integer readOnly: true @@ -19023,6 +19080,60 @@ components: readOnly: true required: - name + CohortCsvSync: + type: object + properties: + file: + type: string + format: uri + identifier_column: + type: integer + default: 0 + minimum: 0 + has_header: + type: boolean + default: true + required: + - file + CohortCsvSyncIgnoredRows: + type: object + properties: + empty: + type: integer + minimum: 0 + duplicates: + type: integer + minimum: 0 + too_long: + type: integer + minimum: 0 + required: + - duplicates + - empty + - too_long + CohortCsvSyncResult: + type: object + properties: + version: + type: integer + minimum: 0 + added: + type: integer + minimum: 0 + removed: + type: integer + minimum: 0 + unchanged: + type: integer + minimum: 0 + ignored: + $ref: '#/components/schemas/CohortCsvSyncIgnoredRows' + required: + - added + - ignored + - removed + - unchanged + - version Condition: type: object properties: @@ -24997,6 +25108,11 @@ components: - $ref: '#/components/schemas/ManagedByEnum' - $ref: '#/components/schemas/BlankEnum' readOnly: true + cohort: + oneOf: + - $ref: '#/components/schemas/_SegmentCohort' + - type: 'null' + readOnly: true PatchedSegmentConfiguration: type: object properties: @@ -26718,6 +26834,11 @@ components: - $ref: '#/components/schemas/ManagedByEnum' - $ref: '#/components/schemas/BlankEnum' readOnly: true + cohort: + oneOf: + - $ref: '#/components/schemas/_SegmentCohort' + - type: 'null' + readOnly: true required: - name - rules @@ -28758,6 +28879,27 @@ components: writeOnly: true required: - type + _SegmentCohort: + type: object + properties: + id: + type: integer + readOnly: true + environment: + type: integer + source_type: + $ref: '#/components/schemas/SourceTypeEnum' + version: + type: integer + maximum: 2147483647 + minimum: 0 + deletion_requested_at: + type: + - string + - 'null' + format: date-time + required: + - environment securitySchemes: Environment API Key: type: apiKey