diff --git a/internal/app/segment/service.go b/internal/app/segment/service.go index 79536d5d..846f3126 100644 --- a/internal/app/segment/service.go +++ b/internal/app/segment/service.go @@ -358,21 +358,11 @@ func (s *service) SetCampaignSegments(ctx context.Context, orgID, campaignID uui } // Links and enrolment commit together: the user is waiting on this one, // and a failed enrolment must not answer 200 with "added 0". - added, xerr := s.repo.ReplaceForCampaign(ctx, orgID, campaignID, ids) + added, status, xerr := s.repo.ReplaceForCampaign(ctx, orgID, campaignID, ids) if xerr != nil { return nil, 0, xerr } - if added > 0 { - links, xerr := s.repo.LinkedCampaignsForSegments(ctx, orgID, ids) - if xerr != nil { - return nil, 0, xerr - } - for _, lc := range links { - if lc.CampaignID == campaignID { - s.reactToEnrolment(ctx, lc, added) - } - } - } + s.reactToEnrolment(ctx, models.LinkedCampaign{CampaignID: campaignID, OrganizationID: orgID, Status: status}, added) out, xerr := s.repo.ListForCampaign(ctx, orgID, campaignID) if xerr != nil { return nil, 0, xerr diff --git a/internal/repository/pg_segment.go b/internal/repository/pg_segment.go index 4f1d2211..e77a7c3f 100644 --- a/internal/repository/pg_segment.go +++ b/internal/repository/pg_segment.go @@ -37,7 +37,8 @@ type SegmentRepository interface { SetForCampaign(ctx context.Context, orgID, campaignID uuid.UUID, segmentIDs []uuid.UUID) *errx.Error // ReplaceForCampaign replaces the links and enrols the members in one // transaction, so a failed enrolment leaves no half-applied link set. - ReplaceForCampaign(ctx context.Context, orgID, campaignID uuid.UUID, segmentIDs []uuid.UUID) (int, *errx.Error) + // Returns how many leads were new and the campaign's status. + ReplaceForCampaign(ctx context.Context, orgID, campaignID uuid.UUID, segmentIDs []uuid.UUID) (int, string, *errx.Error) // SyncCampaignSegments enrols every current member of the campaign's // linked segments that is not yet a lead; returns how many were added. SyncCampaignSegments(ctx context.Context, orgID, campaignID uuid.UUID) (int, *errx.Error) @@ -451,7 +452,7 @@ func (r *segmentRepository) SetForCampaign(ctx context.Context, orgID, campaignI return errx.InternalError() } defer tx.Rollback(ctx) - if xerr := setForCampaignTx(ctx, tx, orgID, campaignID, segmentIDs); xerr != nil { + if _, xerr := setForCampaignTx(ctx, tx, orgID, campaignID, segmentIDs); xerr != nil { return xerr } if err := tx.Commit(ctx); err != nil { @@ -461,61 +462,65 @@ func (r *segmentRepository) SetForCampaign(ctx context.Context, orgID, campaignI return nil } -func (r *segmentRepository) ReplaceForCampaign(ctx context.Context, orgID, campaignID uuid.UUID, segmentIDs []uuid.UUID) (int, *errx.Error) { +func (r *segmentRepository) ReplaceForCampaign(ctx context.Context, orgID, campaignID uuid.UUID, segmentIDs []uuid.UUID) (int, string, *errx.Error) { tx, err := r.DB.Begin(ctx) if err != nil { db.CaptureError(err, "", nil, "begin") - return 0, errx.InternalError() + return 0, "", errx.InternalError() } defer tx.Rollback(ctx) - if xerr := setForCampaignTx(ctx, tx, orgID, campaignID, segmentIDs); xerr != nil { - return 0, xerr + status, xerr := setForCampaignTx(ctx, tx, orgID, campaignID, segmentIDs) + if xerr != nil { + return 0, "", xerr } added, xerr := syncCampaignSegmentsTx(ctx, tx, orgID, campaignID) if xerr != nil { - return 0, xerr + return 0, "", xerr } if err := tx.Commit(ctx); err != nil { db.CaptureError(err, "", nil, "commit") - return 0, errx.InternalError() + return 0, "", errx.InternalError() } - return added, nil + return added, status, nil } -func setForCampaignTx(ctx context.Context, tx pgx.Tx, orgID, campaignID uuid.UUID, segmentIDs []uuid.UUID) *errx.Error { +// setForCampaignTx replaces the links under a lock on the campaign row, so two +// concurrent replacements cannot commit the union of their sets; the status +// read under that lock is what the caller reacts to. +func setForCampaignTx(ctx context.Context, tx pgx.Tx, orgID, campaignID uuid.UUID, segmentIDs []uuid.UUID) (string, *errx.Error) { // A nil slice would reach Postgres as ANY(NULL) and skip the delete. if segmentIDs == nil { segmentIDs = []uuid.UUID{} } - var exists bool - if err := tx.QueryRow(ctx, `SELECT EXISTS (SELECT 1 FROM campaigns WHERE id = $1 AND organization_id = $2)`, campaignID, orgID).Scan(&exists); err != nil { - db.CaptureError(err, "campaign exists", nil, "queryrow") - return errx.InternalError() - } - if !exists { - return errx.New(errx.NotFound, "campaign not found") + var status string + if err := tx.QueryRow(ctx, `SELECT status FROM campaigns WHERE id = $1 AND organization_id = $2 FOR UPDATE`, campaignID, orgID).Scan(&status); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return "", errx.New(errx.NotFound, "campaign not found") + } + db.CaptureError(err, "campaign lock", nil, "queryrow") + return "", errx.InternalError() } if len(segmentIDs) > 0 { var n int if err := tx.QueryRow(ctx, `SELECT COUNT(*) FROM segments WHERE organization_id = $1 AND id = ANY($2::uuid[])`, orgID, segmentIDs).Scan(&n); err != nil { db.CaptureError(err, "segments verify", nil, "queryrow") - return errx.InternalError() + return "", errx.InternalError() } if n != len(segmentIDs) { - return errx.New(errx.BadRequest, "a linked segment does not exist") + return "", errx.New(errx.BadRequest, "a linked segment does not exist") } } if _, err := tx.Exec(ctx, `DELETE FROM campaign_segments WHERE campaign_id = $1 AND NOT (segment_id = ANY($2::uuid[]))`, campaignID, segmentIDs); err != nil { db.CaptureError(err, "campaign segments delete", nil, "exec") - return errx.InternalError() + return "", errx.InternalError() } if len(segmentIDs) > 0 { if _, err := tx.Exec(ctx, `INSERT INTO campaign_segments (campaign_id, segment_id) SELECT $1, unnest($2::uuid[]) ON CONFLICT DO NOTHING`, campaignID, segmentIDs); err != nil { db.CaptureError(err, "campaign segments insert", nil, "exec") - return errx.InternalError() + return "", errx.InternalError() } } - return nil + return status, nil } func (r *segmentRepository) SyncCampaignSegments(ctx context.Context, orgID, campaignID uuid.UUID) (int, *errx.Error) { diff --git a/internal/repository/segment_live_test.go b/internal/repository/segment_live_test.go index da1cbc1b..f1b0b5f7 100644 --- a/internal/repository/segment_live_test.go +++ b/internal/repository/segment_live_test.go @@ -271,9 +271,9 @@ func TestLiveSegmentCampaignLinks(t *testing.T) { // Replacing the set enrols current members in the same transaction; a // second pass adds nothing. - added, xerr := repo.ReplaceForCampaign(ctx, f.org, f.other, []uuid.UUID{acme.ID}) - if xerr != nil || added != 2 { - t.Fatalf("replace = %d, %v", added, xerr) + added, status, xerr := repo.ReplaceForCampaign(ctx, f.org, f.other, []uuid.UUID{acme.ID}) + if xerr != nil || added != 2 || status == "" { + t.Fatalf("replace = %d, %q, %v", added, status, xerr) } if links, _ = repo.ListForCampaign(ctx, f.org, f.other); len(links) != 1 || links[0].LeadCount != 2 || links[0].HeldOutCount != 0 { t.Fatalf("links after replace = %+v", links)