Skip to content

Commit bbc4bd1

Browse files
Zaimwa9pre-commit-ci[bot]flagsmith-engineering[bot]
authored
feat(cohorts): expose sync state and allow updating the managed segment (#8386)
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: flagsmith-engineering[bot] <flagsmith-engineering[bot]@users.noreply.github.com>
1 parent 0f0e09a commit bbc4bd1

8 files changed

Lines changed: 450 additions & 9 deletions

File tree

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,15 @@
1+
from django.db import migrations, models
2+
3+
4+
class Migration(migrations.Migration):
5+
dependencies = [
6+
("cohorts", "0004_mixpanel_source"),
7+
]
8+
9+
operations = [
10+
migrations.AddField(
11+
model_name="cohort",
12+
name="last_synced_at",
13+
field=models.DateTimeField(blank=True, null=True),
14+
),
15+
]

api/cohorts/models.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@ class Cohort(SoftDeleteExportableModel):
3636
external_id = models.CharField(max_length=255, null=True, blank=True)
3737
version = models.PositiveIntegerField(default=0)
3838
created_at = models.DateTimeField(auto_now_add=True)
39+
last_synced_at = models.DateTimeField(null=True, blank=True)
3940
# Deletion drains memberships from the identity store first; the cohort is
4041
# only soft-deleted once drained. This marks it as awaiting that final step.
4142
deletion_requested_at = models.DateTimeField(null=True, blank=True)

api/cohorts/serializers.py

Lines changed: 54 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,13 @@
11
import typing
22

33
from django.core.files.uploadedfile import UploadedFile
4+
from django.db.models import Count, Q
5+
from drf_spectacular.utils import extend_schema_field
46
from rest_framework import serializers
57

68
from cohorts.constants import COHORT_CSV_MAX_FILE_SIZE_BYTES
79
from cohorts.exceptions import CsvFileTooLargeError
8-
from cohorts.models import Cohort, CohortSyncKey
10+
from cohorts.models import Cohort, CohortMembershipState, CohortSyncKey
911
from cohorts.services import create_cohort
1012
from environments.models import Environment
1113
from metadata.serializers import MetadataSerializer, MetadataSerializerMixin
@@ -19,12 +21,19 @@ class Meta:
1921
model = Segment
2022

2123

24+
class CohortMembershipCountsSerializer(serializers.Serializer): # type: ignore[type-arg]
25+
applied = serializers.IntegerField(min_value=0)
26+
pending_add = serializers.IntegerField(min_value=0)
27+
pending_remove = serializers.IntegerField(min_value=0)
28+
29+
2230
class CohortSerializer(serializers.ModelSerializer[Cohort]):
2331
name = serializers.CharField(max_length=2000, source="segment.name")
2432
description = serializers.CharField(
2533
source="segment.description", required=False, allow_null=True
2634
)
2735
metadata = MetadataSerializer(required=False, many=True, write_only=True)
36+
membership_counts = serializers.SerializerMethodField()
2837

2938
class Meta:
3039
model = Cohort
@@ -38,11 +47,41 @@ class Meta:
3847
"source_type",
3948
"version",
4049
"created_at",
50+
"last_synced_at",
51+
"membership_counts",
52+
)
53+
read_only_fields = (
54+
"segment",
55+
"source_type",
56+
"version",
57+
"created_at",
58+
"last_synced_at",
59+
)
60+
61+
@extend_schema_field(CohortMembershipCountsSerializer)
62+
def get_membership_counts(self, cohort: Cohort) -> dict[str, int]:
63+
# Clients derive sync status and progress from these. The viewset
64+
# annotates the counts; a freshly created cohort isn't annotated.
65+
if (applied := getattr(cohort, "applied_count", None)) is not None:
66+
return {
67+
"applied": applied,
68+
"pending_add": getattr(cohort, "pending_add_count", 0),
69+
"pending_remove": getattr(cohort, "pending_remove_count", 0),
70+
}
71+
return cohort.memberships.aggregate(
72+
applied=Count("id", filter=Q(state=CohortMembershipState.APPLIED)),
73+
pending_add=Count("id", filter=Q(state=CohortMembershipState.PENDING_ADD)),
74+
pending_remove=Count(
75+
"id", filter=Q(state=CohortMembershipState.PENDING_REMOVE)
76+
),
4177
)
42-
read_only_fields = ("segment", "source_type", "version", "created_at")
4378

4479
def validate(self, attrs: dict[str, typing.Any]) -> dict[str, typing.Any]:
4580
attrs = super().validate(attrs)
81+
if self.instance is not None and "metadata" not in attrs:
82+
# A partial update without metadata must not fail the
83+
# required-metadata check.
84+
return attrs
4685
environment = Environment.objects.get(
4786
api_key=self.context["view"].kwargs["environment_api_key"]
4887
)
@@ -64,6 +103,19 @@ def create(self, validated_data: dict[str, typing.Any]) -> Cohort:
64103
_SegmentMetadataHandler()._update_metadata(cohort.segment, metadata_data)
65104
return cohort
66105

106+
def update(self, instance: Cohort, validated_data: dict[str, typing.Any]) -> Cohort:
107+
# Only the managed segment's fields are updatable.
108+
metadata_data = validated_data.pop("metadata", None)
109+
segment_data = validated_data.pop("segment", {})
110+
if segment_data:
111+
segment = instance.segment
112+
for field, value in segment_data.items():
113+
setattr(segment, field, value)
114+
segment.save(update_fields=list(segment_data))
115+
if metadata_data is not None:
116+
_SegmentMetadataHandler()._update_metadata(instance.segment, metadata_data)
117+
return instance
118+
67119

68120
class CohortSyncKeySerializer(serializers.ModelSerializer[CohortSyncKey]):
69121
key = serializers.SerializerMethodField()

api/cohorts/services.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -404,8 +404,10 @@ def sync_cohort_memberships_from_csv(
404404
state=CohortMembershipState.PENDING_REMOVE, updated_at=timezone.now()
405405
)
406406

407-
Cohort.objects.filter(id=cohort.id).update(version=F("version") + 1)
408-
cohort.refresh_from_db(fields=["version"])
407+
Cohort.objects.filter(id=cohort.id).update(
408+
version=F("version") + 1, last_synced_at=timezone.now()
409+
)
410+
cohort.refresh_from_db(fields=["version", "last_synced_at"])
409411
apply_cohort_membership_deltas.delay(kwargs={"cohort_id": cohort.id})
410412

411413
flagsmith_cohorts_csv_syncs_total.inc()

api/cohorts/views.py

Lines changed: 24 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
from django.db.models import QuerySet
1+
from django.db.models import Count, Q, QuerySet
22
from drf_spectacular.utils import extend_schema, extend_schema_view
33
from rest_framework import mixins, status, viewsets
44
from rest_framework.decorators import action
@@ -10,7 +10,7 @@
1010

1111
from api.serializers import ErrorSerializer
1212
from cohorts import services
13-
from cohorts.models import Cohort, CohortSyncKey
13+
from cohorts.models import Cohort, CohortMembershipState, CohortSyncKey
1414
from cohorts.permissions import CohortPermission, CohortPlanPermission
1515
from cohorts.serializers import (
1616
CohortCsvSyncResultSerializer,
@@ -29,6 +29,11 @@
2929
description="Create a cohort and the managed segment that targets it."
3030
),
3131
retrieve=extend_schema(description="Retrieve a cohort."),
32+
partial_update=extend_schema(
33+
description=(
34+
"Update the cohort's managed segment name, description and metadata."
35+
)
36+
),
3237
destroy=extend_schema(
3338
description=(
3439
"Request cohort deletion. Memberships are drained from identity "
@@ -42,6 +47,7 @@ class CohortViewSet(
4247
mixins.ListModelMixin,
4348
mixins.CreateModelMixin,
4449
mixins.RetrieveModelMixin,
50+
mixins.UpdateModelMixin,
4551
mixins.DestroyModelMixin,
4652
):
4753
serializer_class = CohortSerializer
@@ -50,6 +56,8 @@ class CohortViewSet(
5056
model_class = Cohort
5157
lookup_field = "id"
5258
lookup_url_kwarg = "cohort_id"
59+
# PATCH only: a cohort has no meaningful full replacement.
60+
http_method_names = ["get", "post", "patch", "delete", "head", "options"]
5361

5462
def get_queryset(self) -> QuerySet[Cohort]:
5563
# A cohort awaiting drain-then-delete is already gone from the
@@ -59,6 +67,20 @@ def get_queryset(self) -> QuerySet[Cohort]:
5967
.get_queryset()
6068
.filter(deletion_requested_at__isnull=True)
6169
.select_related("segment")
70+
.annotate(
71+
applied_count=Count(
72+
"memberships",
73+
filter=Q(memberships__state=CohortMembershipState.APPLIED),
74+
),
75+
pending_add_count=Count(
76+
"memberships",
77+
filter=Q(memberships__state=CohortMembershipState.PENDING_ADD),
78+
),
79+
pending_remove_count=Count(
80+
"memberships",
81+
filter=Q(memberships__state=CohortMembershipState.PENDING_REMOVE),
82+
),
83+
)
6284
.order_by("id")
6385
)
6486

0 commit comments

Comments
 (0)