Skip to content

Commit 2fdbdbb

Browse files
committed
feat(cohorts): generic cohort sync webhook
1 parent 17243dc commit 2fdbdbb

6 files changed

Lines changed: 322 additions & 18 deletions

File tree

api/cohorts/serializers.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -156,3 +156,11 @@ class CohortCsvSyncResultSerializer(serializers.Serializer): # type: ignore[typ
156156
removed = serializers.IntegerField(min_value=0)
157157
unchanged = serializers.IntegerField(min_value=0)
158158
ignored = CohortCsvSyncIgnoredRowsSerializer()
159+
160+
161+
class WebhookSyncMembersSerializer(serializers.Serializer[None]):
162+
identifiers = serializers.ListField(
163+
child=serializers.CharField(validators=[_validate_identifier_byte_length]),
164+
min_length=1,
165+
max_length=10000,
166+
)

api/cohorts/services.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import csv
22
import io
33
import typing
4+
from uuid import UUID
45

56
import structlog
67
from django.db import transaction
@@ -192,6 +193,25 @@ def create_cohort_for_source(
192193
return cohort
193194

194195

196+
def get_active_cohort(
197+
*,
198+
environment: "Environment",
199+
source_type: CohortSourceType,
200+
uuid: str,
201+
) -> Cohort | None:
202+
try:
203+
cohort_uuid = UUID(uuid)
204+
except ValueError:
205+
return None
206+
cohort: Cohort | None = Cohort.objects.filter(
207+
uuid=cohort_uuid,
208+
environment=environment,
209+
source_type=source_type,
210+
deletion_requested_at__isnull=True,
211+
).first()
212+
return cohort
213+
214+
195215
def get_cohort_for_source(
196216
*,
197217
environment: "Environment",

api/cohorts/sync_urls.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,18 @@
11
from django.urls import path
22
from rest_framework.routers import SimpleRouter
33

4-
from cohorts.sync_views import AmplitudeCohortSyncViewSet, MixpanelCohortSyncView
4+
from cohorts.sync_views import (
5+
AmplitudeCohortSyncViewSet,
6+
MixpanelCohortSyncView,
7+
WebhookCohortSyncViewSet,
8+
)
59

610
app_name = "cohort-sync"
711

812
# SimpleRouter: nothing here is browsed by a person.
913
router = SimpleRouter()
1014
router.register(r"amplitude/lists", AmplitudeCohortSyncViewSet, basename="amplitude")
15+
router.register(r"webhook/cohorts", WebhookCohortSyncViewSet, basename="webhook")
1116

1217
urlpatterns = [
1318
path("mixpanel/webhook/", MixpanelCohortSyncView.as_view(), name="mixpanel"),

api/cohorts/sync_views.py

Lines changed: 56 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
11
import json
22
import typing
3-
import uuid as uuid_module
43

54
import structlog
65
from drf_spectacular.utils import extend_schema, extend_schema_view, inline_serializer
@@ -19,6 +18,7 @@
1918
AmplitudeListSerializer,
2019
CohortSyncMembersSerializer,
2120
MixpanelWebhookSerializer,
21+
WebhookSyncMembersSerializer,
2222
)
2323

2424
_LIST_RESPONSE = inline_serializer(
@@ -96,16 +96,11 @@ def _get_key(self, request: Request) -> CohortSyncKey:
9696
return typing.cast(CohortSyncKey, request.auth)
9797

9898
def _get_cohort(self, request: Request, pk: str) -> Cohort:
99-
try:
100-
list_uuid = uuid_module.UUID(pk)
101-
except ValueError:
102-
raise NotFound("List not found.")
103-
cohort: Cohort | None = Cohort.objects.filter(
104-
uuid=list_uuid,
99+
cohort = services.get_active_cohort(
105100
environment=self._get_key(request).environment,
106101
source_type=CohortSourceType.AMPLITUDE,
107-
deletion_requested_at__isnull=True,
108-
).first()
102+
uuid=pk,
103+
)
109104
if cohort is None:
110105
raise NotFound("List not found.")
111106
return cohort
@@ -241,3 +236,55 @@ def _echo_action(self, request: Request) -> str | None:
241236
):
242237
return action_value
243238
return None
239+
240+
241+
@extend_schema_view(
242+
add=extend_schema(
243+
description=(
244+
"Add members to a CSV cohort. Accepted deltas are applied to "
245+
"identity data asynchronously. Re-adding a member is a no-op, "
246+
"so retries are safe."
247+
),
248+
request=WebhookSyncMembersSerializer,
249+
responses={200: None},
250+
),
251+
remove=extend_schema(
252+
description=(
253+
"Remove members from a CSV cohort. Accepted deltas are applied "
254+
"to identity data asynchronously. Removing a non-member is a "
255+
"no-op, so retries are safe."
256+
),
257+
request=WebhookSyncMembersSerializer,
258+
responses={200: None},
259+
),
260+
)
261+
class WebhookCohortSyncViewSet(viewsets.ViewSet):
262+
authentication_classes = [CohortSyncKeyAuthentication]
263+
permission_classes = [HasCohortSyncKey, CohortSyncPlanPermission]
264+
265+
@action(detail=True, methods=["POST"], url_path="members/add")
266+
def add(self, request: Request, pk: str) -> Response:
267+
cohort = self._get_cohort(request, pk)
268+
serializer = WebhookSyncMembersSerializer(data=request.data)
269+
serializer.is_valid(raise_exception=True)
270+
services.add_cohort_members(cohort, serializer.validated_data["identifiers"])
271+
return Response()
272+
273+
@action(detail=True, methods=["POST"], url_path="members/remove")
274+
def remove(self, request: Request, pk: str) -> Response:
275+
cohort = self._get_cohort(request, pk)
276+
serializer = WebhookSyncMembersSerializer(data=request.data)
277+
serializer.is_valid(raise_exception=True)
278+
services.remove_cohort_members(cohort, serializer.validated_data["identifiers"])
279+
return Response()
280+
281+
def _get_cohort(self, request: Request, pk: str) -> Cohort:
282+
cohort = services.get_active_cohort(
283+
# HasCohortSyncKey has already established the type.
284+
environment=typing.cast(CohortSyncKey, request.auth).environment,
285+
source_type=CohortSourceType.CSV,
286+
uuid=pk,
287+
)
288+
if cohort is None:
289+
raise NotFound("Cohort not found.")
290+
return cohort

api/tests/unit/cohorts/test_sync_views.py

Lines changed: 224 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1041,3 +1041,227 @@ def test_mixpanel_webhook__saas_free_plan__returns_403(
10411041

10421042
# Then
10431043
assert response.status_code == status.HTTP_403_FORBIDDEN
1044+
1045+
1046+
def test_webhook_add_members__csv_cohort__applies_memberships(
1047+
cohort: Cohort,
1048+
) -> None:
1049+
# Given
1050+
_, plaintext = CohortSyncKey.objects.create_key(
1051+
name="test key", environment=cohort.environment
1052+
)
1053+
client = _authenticated_client(plaintext)
1054+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": str(cohort.uuid)})
1055+
1056+
# When
1057+
response = client.post(
1058+
url, data={"identifiers": ["user-1", "user-2"]}, format="json"
1059+
)
1060+
1061+
# Then
1062+
assert response.status_code == status.HTTP_200_OK
1063+
assert sorted(
1064+
CohortMembership.objects.filter(cohort=cohort).values_list(
1065+
"identifier", "state"
1066+
)
1067+
) == [
1068+
("user-1", CohortMembershipState.APPLIED),
1069+
("user-2", CohortMembershipState.APPLIED),
1070+
]
1071+
identity = Identity.objects.get(environment=cohort.environment, identifier="user-1")
1072+
assert identity.system_traits == {cohort.system_trait_key: True}
1073+
1074+
1075+
def test_webhook_remove_members__applied_member__removes_membership_and_trait(
1076+
cohort: Cohort,
1077+
) -> None:
1078+
# Given
1079+
Identity.objects.create(
1080+
environment=cohort.environment,
1081+
identifier="member",
1082+
system_traits={cohort.system_trait_key: True},
1083+
)
1084+
CohortMembership.objects.create(
1085+
cohort=cohort, identifier="member", state=CohortMembershipState.APPLIED
1086+
)
1087+
_, plaintext = CohortSyncKey.objects.create_key(
1088+
name="test key", environment=cohort.environment
1089+
)
1090+
client = _authenticated_client(plaintext)
1091+
url = reverse("api-v1:cohort-sync:webhook-remove", kwargs={"pk": str(cohort.uuid)})
1092+
1093+
# When
1094+
response = client.post(url, data={"identifiers": ["member"]}, format="json")
1095+
1096+
# Then
1097+
assert response.status_code == status.HTTP_200_OK
1098+
assert not CohortMembership.objects.filter(cohort=cohort).exists()
1099+
identity = Identity.objects.get(environment=cohort.environment, identifier="member")
1100+
assert identity.system_traits == {}
1101+
1102+
1103+
def test_webhook_add_members__non_csv_cohort__returns_404(
1104+
cohort_sync_key: _KeyAndPlaintext,
1105+
amplitude_cohort: Cohort,
1106+
) -> None:
1107+
# Given
1108+
_, plaintext = cohort_sync_key
1109+
client = _authenticated_client(plaintext)
1110+
url = reverse(
1111+
"api-v1:cohort-sync:webhook-add", kwargs={"pk": str(amplitude_cohort.uuid)}
1112+
)
1113+
1114+
# When
1115+
response = client.post(url, data={"identifiers": ["user-1"]}, format="json")
1116+
1117+
# Then
1118+
assert response.status_code == status.HTTP_404_NOT_FOUND
1119+
1120+
1121+
def test_webhook_add_members__other_environment_key__returns_404(
1122+
cohort: Cohort,
1123+
) -> None:
1124+
# Given - a key scoped to a different environment than the cohort's
1125+
other_environment = Environment.objects.create(
1126+
name="Other environment", project=cohort.environment.project
1127+
)
1128+
_, plaintext = CohortSyncKey.objects.create_key(
1129+
name="other key", environment=other_environment
1130+
)
1131+
client = _authenticated_client(plaintext)
1132+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": str(cohort.uuid)})
1133+
1134+
# When
1135+
response = client.post(url, data={"identifiers": ["user-1"]}, format="json")
1136+
1137+
# Then
1138+
assert response.status_code == status.HTTP_404_NOT_FOUND
1139+
assert not CohortMembership.objects.exists()
1140+
1141+
1142+
def test_webhook_add_members__identifier_over_1024_bytes__returns_400(
1143+
cohort: Cohort,
1144+
) -> None:
1145+
# Given - 512 three-byte characters: few characters, too many bytes
1146+
multibyte_identifier = "€" * 512
1147+
_, plaintext = CohortSyncKey.objects.create_key(
1148+
name="test key", environment=cohort.environment
1149+
)
1150+
client = _authenticated_client(plaintext)
1151+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": str(cohort.uuid)})
1152+
1153+
# When
1154+
response = client.post(
1155+
url, data={"identifiers": [multibyte_identifier]}, format="json"
1156+
)
1157+
1158+
# Then
1159+
assert response.status_code == status.HTTP_400_BAD_REQUEST
1160+
assert "1024 bytes" in str(response.json())
1161+
assert not CohortMembership.objects.exists()
1162+
1163+
1164+
def test_webhook_add_members__over_10000_identifiers__returns_400(
1165+
cohort: Cohort,
1166+
) -> None:
1167+
# Given
1168+
_, plaintext = CohortSyncKey.objects.create_key(
1169+
name="test key", environment=cohort.environment
1170+
)
1171+
client = _authenticated_client(plaintext)
1172+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": str(cohort.uuid)})
1173+
1174+
# When
1175+
response = client.post(
1176+
url,
1177+
data={"identifiers": [f"user-{i}" for i in range(10001)]},
1178+
format="json",
1179+
)
1180+
1181+
# Then
1182+
assert response.status_code == status.HTTP_400_BAD_REQUEST
1183+
assert not CohortMembership.objects.exists()
1184+
1185+
1186+
def test_webhook_add_members__empty_identifiers__returns_400(
1187+
cohort: Cohort,
1188+
) -> None:
1189+
# Given
1190+
_, plaintext = CohortSyncKey.objects.create_key(
1191+
name="test key", environment=cohort.environment
1192+
)
1193+
client = _authenticated_client(plaintext)
1194+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": str(cohort.uuid)})
1195+
1196+
# When
1197+
response = client.post(url, data={"identifiers": []}, format="json")
1198+
1199+
# Then
1200+
assert response.status_code == status.HTTP_400_BAD_REQUEST
1201+
1202+
1203+
def test_webhook_add_members__malformed_uuid__returns_404(
1204+
cohort: Cohort,
1205+
) -> None:
1206+
# Given
1207+
_, plaintext = CohortSyncKey.objects.create_key(
1208+
name="test key", environment=cohort.environment
1209+
)
1210+
client = _authenticated_client(plaintext)
1211+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": "not-a-uuid"})
1212+
1213+
# When
1214+
response = client.post(url, data={"identifiers": ["user-1"]}, format="json")
1215+
1216+
# Then
1217+
assert response.status_code == status.HTTP_404_NOT_FOUND
1218+
1219+
1220+
def test_webhook_add_members__deletion_requested_cohort__returns_404(
1221+
cohort: Cohort,
1222+
) -> None:
1223+
# Given
1224+
cohort.deletion_requested_at = timezone.now()
1225+
cohort.save()
1226+
_, plaintext = CohortSyncKey.objects.create_key(
1227+
name="test key", environment=cohort.environment
1228+
)
1229+
client = _authenticated_client(plaintext)
1230+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": str(cohort.uuid)})
1231+
1232+
# When
1233+
response = client.post(url, data={"identifiers": ["user-1"]}, format="json")
1234+
1235+
# Then
1236+
assert response.status_code == status.HTTP_404_NOT_FOUND
1237+
1238+
1239+
def test_webhook_add_members__missing_credentials__returns_401(
1240+
cohort: Cohort,
1241+
) -> None:
1242+
# Given
1243+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": str(cohort.uuid)})
1244+
1245+
# When
1246+
response = APIClient().post(url, data={"identifiers": ["user-1"]}, format="json")
1247+
1248+
# Then
1249+
assert response.status_code == status.HTTP_401_UNAUTHORIZED
1250+
1251+
1252+
@pytest.mark.saas_mode
1253+
def test_webhook_add_members__saas_free_plan__returns_403(
1254+
cohort: Cohort,
1255+
) -> None:
1256+
# Given
1257+
_, plaintext = CohortSyncKey.objects.create_key(
1258+
name="test key", environment=cohort.environment
1259+
)
1260+
client = _authenticated_client(plaintext)
1261+
url = reverse("api-v1:cohort-sync:webhook-add", kwargs={"pk": str(cohort.uuid)})
1262+
1263+
# When
1264+
response = client.post(url, data={"identifiers": ["user-1"]}, format="json")
1265+
1266+
# Then
1267+
assert response.status_code == status.HTTP_403_FORBIDDEN

0 commit comments

Comments
 (0)