From d92ad35c4225192ebcfe3950c93bd684edf07fcc Mon Sep 17 00:00:00 2001 From: Kim Svatos Dugan <147102038+ksvat@users.noreply.github.com> Date: Wed, 30 Sep 2026 08:51:29 -0700 Subject: [PATCH 1/5] feat(replay-vision): sample an experiment scanner's variants evenly MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Balanced per-variant sampling for experiment scanners, default on: each sweep computes one salted-hash rate per watched variant from the flag's live rollout shares (r / (k · share), capped at 1, with a capped variant's unspent budget redistributed), so an uneven rollout no longer starves the small arm. The exposure join projects the attributed variant when balancing is on, backfills share the same query class and salt, the volume estimate projects with the plan's effective rate, and each observation's snapshot records the rates it was sampled at. Co-Authored-By: Claude Fable 5 Generated-By: PostHog Desktop Task-Id: c2db2a38-d03f-4578-af4a-40f61dc700a4 --- .../session_recording_list_from_query.py | 26 ++++- products/experiments/backend/facade/replay.py | 2 + .../replay_vision/backend/api/observations.py | 10 ++ .../queries/scanner_candidate_query.py | 80 +++++++++++---- .../queries/scanner_volume_estimate.py | 12 ++- .../backend/queries/variant_sampling.py | 97 +++++++++++++++++++ .../backend/temporal/activities/backfill.py | 12 +++ .../temporal/activities/create_observation.py | 20 +++- .../activities/find_scanner_candidates.py | 13 +++ .../backend/temporal/backfill_types.py | 3 + .../backend/temporal/backfill_workflow.py | 12 ++- .../backend/temporal/snapshots.py | 4 + .../backend/temporal/sweep_types.py | 4 + .../backend/temporal/sweep_workflow.py | 18 +++- .../replay_vision/backend/temporal/types.py | 4 + .../backend/temporal/workflow.py | 1 + .../tests/test_scanner_candidate_query.py | 85 +++++++++++++++- .../backend/tests/test_temporal.py | 22 +++++ .../backend/tests/test_variant_sampling.py | 79 +++++++++++++++ .../frontend/generated/api.schemas.ts | 11 +++ services/mcp/src/api/generated.ts | 11 +++ 21 files changed, 492 insertions(+), 34 deletions(-) create mode 100644 products/replay_vision/backend/queries/variant_sampling.py create mode 100644 products/replay_vision/backend/tests/test_variant_sampling.py diff --git a/posthog/session_recordings/queries/session_recording_list_from_query.py b/posthog/session_recordings/queries/session_recording_list_from_query.py index f1a18cb69a56..fdbec67bc7b1 100644 --- a/posthog/session_recordings/queries/session_recording_list_from_query.py +++ b/posthog/session_recordings/queries/session_recording_list_from_query.py @@ -159,9 +159,14 @@ def __init__( # Opt-in: resolve group property filters to group keys instead of joining the groups table. # Naming the ClickHouse user is the opt-in, since the resolution is itself a heavy query. resolve_group_properties: ClickHouseUser | None = None, + # Opt-in for callers whose extra_having_predicates read the exposed person's attributed + # variant (as `any(exposure.variant)`): the experiment-exposure join then projects it. + # The population is unchanged, so plain listings never need this. + project_exposure_variant: bool = False, **_, ): self._user = user + self._project_exposure_variant = project_exposure_variant # Storage-level SAMPLE on any events subqueries; opt-in for estimates. self._events_sample_factor = events_sample_factor # Extra lower bound on positive events subqueries, for callers that re-run often over a wide @@ -415,7 +420,10 @@ def _join_experiment_exposure(self, parsed_query: ast.SelectQuery) -> None: """ # Deferred: the experiments facade package imports posthog.api on init, which # circles back into this module through the replay-deletion temporal activities. - from products.experiments.backend.facade.replay import exposed_distinct_ids_select # noqa: PLC0415 + from products.experiments.backend.facade.replay import ( # noqa: PLC0415 + exposed_distinct_ids_select, + exposed_persons_select, + ) self._resolve_experiment_exposure() assert self._experiment_exposure_linkage is not None @@ -441,13 +449,23 @@ def _join_experiment_exposure(self, parsed_query: ast.SelectQuery) -> None: assert join is not None while join.next_join is not None: join = join.next_join + exposure_select = ( + exposed_persons_select( + self._experiment_exposure_linkage, + include_multiple_variant=False, + candidate_distinct_ids=candidate_distinct_ids, + ) + if self._project_exposure_variant + # Same population either way: the persons select only adds the attribution columns. + else exposed_distinct_ids_select( + self._experiment_exposure_linkage, candidate_distinct_ids=candidate_distinct_ids + ) + ) join.next_join = ast.JoinExpr( # GLOBAL: the subquery scans events over the whole experiment window; without it, # every shard of the sharded replay table re-evaluates that scan independently. join_type="GLOBAL INNER JOIN", - table=exposed_distinct_ids_select( - self._experiment_exposure_linkage, candidate_distinct_ids=candidate_distinct_ids - ), + table=exposure_select, alias="exposure", constraint=ast.JoinConstraint( expr=ast.CompareOperation( diff --git a/products/experiments/backend/facade/replay.py b/products/experiments/backend/facade/replay.py index ddbb6b5420fd..795bf1910b05 100644 --- a/products/experiments/backend/facade/replay.py +++ b/products/experiments/backend/facade/replay.py @@ -15,6 +15,7 @@ ExperimentExposureLinkage, InSessionExposureSemantics, exposed_distinct_ids_select, + exposed_persons_select, exposed_session_ids_select, resolve_exposure_linkage, resolve_in_session_exposure_semantics, @@ -32,6 +33,7 @@ "experiment_prompt_context", "experiment_status", "exposed_distinct_ids_select", + "exposed_persons_select", "exposed_session_ids_select", "resolve_exposure_linkage", "resolve_in_session_exposure_semantics", diff --git a/products/replay_vision/backend/api/observations.py b/products/replay_vision/backend/api/observations.py index af6e2c038d8d..6d6121bc8ddd 100644 --- a/products/replay_vision/backend/api/observations.py +++ b/products/replay_vision/backend/api/observations.py @@ -138,6 +138,16 @@ class ScannerSnapshotSerializer(serializers.Serializer): verify_positives = serializers.CharField( help_text="How a monitor `yes` was re-checked at run time: `off` (one pass, the default), `shadow` (second draw recorded only), or `enforce` (the `yes` stands only when the second draw agrees).", ) + variant_sampling_rates = serializers.DictField( + child=serializers.FloatField(), + required=False, + allow_null=True, + help_text=( + "Experiment scanners with balanced sampling: the 0..1 rate each watched variant was sampled at " + "by the tick that dispatched this scan. Null otherwise, so even per-variant counts can be read " + "against the rates that produced them." + ), + ) class VerificationRecordSerializer(serializers.Serializer): diff --git a/products/replay_vision/backend/queries/scanner_candidate_query.py b/products/replay_vision/backend/queries/scanner_candidate_query.py index ad1c5a41eefd..06d771db6c93 100644 --- a/products/replay_vision/backend/queries/scanner_candidate_query.py +++ b/products/replay_vision/backend/queries/scanner_candidate_query.py @@ -193,6 +193,9 @@ def __init__( skip_negative_blocklists: bool = False, # Tags the ClickHouse query for per-scanner read metering; sweep callers should always pass it. scanner_id: str | None = None, + # Balanced sampling for experiment scanners: one rate per watched variant, replacing the + # single `sampling_rate` threshold (see `variant_sampling_predicate`). + variant_sampling_rates: dict[str, float] | None = None, ) -> None: if not isinstance(last_swept_at, dt.datetime): raise TypeError(f"last_swept_at must be a datetime, got {type(last_swept_at).__name__}") @@ -208,6 +211,7 @@ def __init__( self._last_seen_session_id = last_seen_session_id self._sampling_rate = max(0.0, min(1.0, sampling_rate)) self._sampling_salt = sampling_salt + self._variant_sampling_rates = variant_sampling_rates self._candidate_limit = candidate_limit self._max_execution_time_seconds = max_execution_time_seconds self._scanner_id = scanner_id @@ -244,6 +248,7 @@ def __init__( events_timestamp_floor=events_timestamp_floor, skip_negative_blocklists=skip_negative_blocklists, resolve_group_properties=ClickHouseUser.REPLAY_VISION, + project_exposure_variant=variant_sampling_rates is not None, ) def excluded_sessions_queries(self, session_ids: list[str]) -> list[ast.SelectQuery]: @@ -319,6 +324,8 @@ def _watermark_predicate(self) -> ast.Expr: return keyset_predicate(self._last_swept_at, self._last_seen_session_id, ascending=True) def _sampling_predicate(self) -> ast.Expr | None: + if self._variant_sampling_rates is not None: + return variant_sampling_predicate(self._variant_sampling_rates, self._sampling_salt) return sampling_predicate(self._sampling_rate, self._sampling_salt) @@ -479,6 +486,25 @@ def keyset_predicate(end_time: dt.datetime, session_id: str | None, ascending: b ) +def _sampling_hash_expr(sampling_salt: str) -> ast.Expr: + return ast.Call( + name="modulo", + args=[ + # concat rather than a second cityHash64 arg — HogQL pins cityHash64 to a single argument. + ast.Call( + name="cityHash64", + args=[ + ast.Call( + name="concat", + args=[ast.Field(chain=["s", "session_id"]), ast.Constant(value=sampling_salt)], + ) + ], + ), + ast.Constant(value=SAMPLE_RATE_PRECISION), + ], + ) + + def sampling_predicate(sampling_rate: float, sampling_salt: str) -> ast.Expr | None: """Deterministic salted-hash downsample on the inner query's session rows; None means keep everything.""" if sampling_rate >= 1.0: @@ -489,26 +515,37 @@ def sampling_predicate(sampling_rate: float, sampling_salt: str) -> ast.Expr | N return ast.Constant(value=False) return ast.CompareOperation( op=ast.CompareOperationOp.Lt, - left=ast.Call( - name="modulo", - args=[ - # concat rather than a second cityHash64 arg — HogQL pins cityHash64 to a single argument. - ast.Call( - name="cityHash64", - args=[ - ast.Call( - name="concat", - args=[ast.Field(chain=["s", "session_id"]), ast.Constant(value=sampling_salt)], - ) - ], - ), - ast.Constant(value=SAMPLE_RATE_PRECISION), - ], - ), + left=_sampling_hash_expr(sampling_salt), right=ast.Constant(value=threshold), ) +def variant_sampling_predicate(rates: dict[str, float], sampling_salt: str) -> ast.Expr | None: + """One salted-hash threshold per attributed variant, so balanced sampling stays stable across + sweeps and backfills the way plain sampling does: the same session hashes the same everywhere. + + Reads the exposure join's attributed variant as `any(exposure.variant)` (an aggregate, since + the predicate lands in HAVING), so the caller must project it (`project_exposure_variant`). + An unlisted variant falls through to false; the join already restricts rows to the watched + variants, so that arm only guards drift. + """ + if all(rate >= 1.0 for rate in rates.values()): + return None + variant_expr = ast.Call(name="any", args=[ast.Field(chain=["exposure", "variant"])]) + hash_expr = _sampling_hash_expr(sampling_salt) + multi_if_args: list[ast.Expr] = [] + for variant, rate in rates.items(): + threshold = max(0, round(min(1.0, rate) * SAMPLE_RATE_PRECISION)) + multi_if_args.append( + ast.CompareOperation(op=ast.CompareOperationOp.Eq, left=variant_expr, right=ast.Constant(value=variant)) + ) + multi_if_args.append( + ast.CompareOperation(op=ast.CompareOperationOp.Lt, left=hash_expr, right=ast.Constant(value=threshold)) + ) + multi_if_args.append(ast.Constant(value=False)) + return ast.Call(name="multiIf", args=multi_if_args) + + class WindowedCandidateQuery: """Enumerate a scanner's candidate sessions inside a closed historical window. @@ -555,6 +592,9 @@ def __init__( candidate_limit: int = DEFAULT_CANDIDATE_LIMIT, max_execution_time_seconds: int = DEFAULT_MAX_EXECUTION_SECONDS, scanner_id: str | None = None, + # Balanced sampling for experiment scanners: one rate per watched variant, replacing the + # single `sampling_rate` threshold (see `variant_sampling_predicate`). + variant_sampling_rates: dict[str, float] | None = None, ) -> None: for name, value in (("window_start", window_start), ("window_end", window_end)): if not isinstance(value, dt.datetime): @@ -589,7 +629,12 @@ def __init__( inner_query.after = None extra_having: list[ast.Expr] = eligibility_predicates() - if (sampling := sampling_predicate(sampling_rate, sampling_salt)) is not None: + sampling = ( + variant_sampling_predicate(variant_sampling_rates, sampling_salt) + if variant_sampling_rates is not None + else sampling_predicate(sampling_rate, sampling_salt) + ) + if sampling is not None: extra_having.append(sampling) if (surfacing := surfacing_score_predicate(sampling_mode)) is not None: extra_having.append(surfacing) @@ -602,6 +647,7 @@ def __init__( session_ids_to_exclude=exclude_session_ids, skip_negative_blocklists=skip_negative_blocklists, resolve_group_properties=ClickHouseUser.REPLAY_VISION, + project_exposure_variant=variant_sampling_rates is not None, ) def excluded_sessions_queries(self, session_ids: list[str]) -> list[ast.SelectQuery]: diff --git a/products/replay_vision/backend/queries/scanner_volume_estimate.py b/products/replay_vision/backend/queries/scanner_volume_estimate.py index 43dd3e90387c..b83637b8c51b 100644 --- a/products/replay_vision/backend/queries/scanner_volume_estimate.py +++ b/products/replay_vision/backend/queries/scanner_volume_estimate.py @@ -28,6 +28,7 @@ eligibility_predicates, surfacing_score_predicate, ) +from products.replay_vision.backend.queries.variant_sampling import variant_sampling_plan_for_scope # The estimate always projects to a calendar month. ESTIMATE_WINDOW_DAYS = ESTIMATE_MONTH_DAYS @@ -319,7 +320,16 @@ def refresh_scanner_estimate( budget=budget, ch_user=ch_user, ) - projection = project_monthly_observations(estimate, scanner.sampling_rate) + # Balanced per-variant sampling caps a small variant at 1, so the fraction actually sampled can + # sit below the configured rate; the projection must use what the sweep will really keep. + variant_plan = variant_sampling_plan_for_scope( + scanner.team, + scope=scanner.experiment_scope(), + scanner_config=scanner.scanner_config, + sampling_rate=scanner.sampling_rate, + ) + effective_rate = variant_plan.effective_rate if variant_plan is not None else scanner.sampling_rate + projection = project_monthly_observations(estimate, effective_rate) estimated_at = timezone.now() # Filtered write so a config edit racing the (slow) estimate query can't get stamped fresh with stale numbers. # JSONField quirk: `field=None` filters for JSON null, not SQL NULL, so the no-targeting case needs isnull. diff --git a/products/replay_vision/backend/queries/variant_sampling.py b/products/replay_vision/backend/queries/variant_sampling.py new file mode 100644 index 000000000000..4fbfaca40c08 --- /dev/null +++ b/products/replay_vision/backend/queries/variant_sampling.py @@ -0,0 +1,97 @@ +"""Per-variant sampling rates for experiment scanners. + +One rate per scanner samples variants proportionally to rollout, so a 90/10 split at a 10% rate +yields ~90 control observations per 10 test. Balanced sampling instead spends the same total +budget about evenly across the watched variants: each variant's rate is `r / (k · share)`, capped +at 1, with a capped variant's unspent budget redistributed to the others so a small variant can +never shrink the total. + +Rates are recomputed from the live rollout shares on every sweep, estimate, and backfill tick, so +a mid-experiment rollout change adjusts them on the next tick. The plan cannot create sessions a +small variant does not have: a capped variant is sampled whole and stays thin. +""" + +from collections.abc import Sequence + +from rest_framework.exceptions import ValidationError + +from posthog.dataclasses import frozen +from posthog.models.team import Team + + +@frozen +class VariantSamplingPlan: + # Hash-threshold rate per watched variant, each in 0..1. + rates: dict[str, float] + # Each watched variant's share of the watched population (normalized over the watched set). + population_shares: dict[str, float] + + @property + def effective_rate(self) -> float: + """The fraction of the watched population the plan samples overall: the scanner's own + rate until a cap binds, then less budget than asked finds sessions to spend itself on.""" + return sum(self.population_shares[variant] * rate for variant, rate in self.rates.items()) + + +def plan_variant_sampling( + sampling_rate: float, rollout_shares: dict[str, float], selected: Sequence[str] | None +) -> VariantSamplingPlan | None: + """The per-variant plan, or None when balancing changes nothing (one watched variant or no shares). + + Water-filling: every uncapped variant gets an equal slice of the remaining budget; a variant + whose whole population fits inside its slice is sampled at 1 and frees the rest of its slice. + """ + selected_keys = [key for key in (selected if selected is not None else rollout_shares) if key in rollout_shares] + if len(selected_keys) <= 1: + return None + raw = {key: max(0.0, float(rollout_shares[key])) for key in selected_keys} + total = sum(raw.values()) + population_shares = ( + {key: value / total for key, value in raw.items()} + if total > 0 + else {key: 1.0 / len(selected_keys) for key in selected_keys} + ) + + budget = max(0.0, min(1.0, sampling_rate)) + rates: dict[str, float] = {} + remaining = set(selected_keys) + while remaining: + slice_per_variant = budget / len(remaining) + capped = {key for key in remaining if population_shares[key] <= slice_per_variant} + if not capped: + for key in remaining: + rates[key] = slice_per_variant / population_shares[key] + break + for key in capped: + rates[key] = 1.0 + budget -= population_shares[key] + budget = max(0.0, budget) + remaining -= capped + return VariantSamplingPlan(rates=rates, population_shares=population_shares) + + +def variant_sampling_plan_for_scope( + team: Team, *, scope: dict | None, scanner_config: dict | None, sampling_rate: float +) -> VariantSamplingPlan | None: + """The plan for a scanner (or frozen snapshot) scope, from the flag's live rollout shares. + + None when the scanner watches no experiment, balancing is off, only one variant is watched, or + the shares can't be read (the caller then falls back to plain sampling; the scan itself stays + the loud path for an unresolvable experiment). + """ + experiment_id = (scope or {}).get("experiment_id") + if experiment_id is None: + return None + config = scanner_config if isinstance(scanner_config, dict) else {} + if config.get("balance_variants") is False: + return None + # Deferred: the experiments replay facade pulls in the recordings query modules, which circle + # back into this package's importers. + from products.experiments.backend.facade.replay import variant_rollout_shares # noqa: PLC0415 + + try: + shares = variant_rollout_shares(team, experiment_id=experiment_id) + except ValidationError: + return None + assert scope is not None + return plan_variant_sampling(sampling_rate, shares, scope.get("variants")) diff --git a/products/replay_vision/backend/temporal/activities/backfill.py b/products/replay_vision/backend/temporal/activities/backfill.py index 6b785eada348..cac2849f8476 100644 --- a/products/replay_vision/backend/temporal/activities/backfill.py +++ b/products/replay_vision/backend/temporal/activities/backfill.py @@ -37,6 +37,7 @@ BACKFILL_EXCLUDED_SESSIONS_QUERY_TYPE, WindowedCandidateQuery, ) +from products.replay_vision.backend.queries.variant_sampling import variant_sampling_plan_for_scope from products.replay_vision.backend.quota import compute_scanner_budget, quota_state from products.replay_vision.backend.temporal.activities.count_in_flight_applies import ( count_in_flight, @@ -168,6 +169,15 @@ def find_backfill_candidates_activity(inputs: FindBackfillCandidatesInputs) -> F ) from exc query = apply_experiment_targeting(query, snapshot.experiment_scope()) + # Live rollout shares against the frozen scope, matching the sweep: the same salted hash plus + # the same rates keep sampling decisions stable between a live sweep and a backfill of the + # same range. + variant_plan = variant_sampling_plan_for_scope( + backfill.team, + scope=snapshot.experiment_scope(), + scanner_config=snapshot.scanner_config, + sampling_rate=snapshot.sampling_rate, + ) candidate_query = WindowedCandidateQuery( team=backfill.team, query=query, @@ -185,6 +195,7 @@ def find_backfill_candidates_activity(inputs: FindBackfillCandidatesInputs) -> F cursor_session_id=backfill.cursor_session_id or None, candidate_limit=inputs.candidate_limit, skip_negative_blocklists=True, + variant_sampling_rates=variant_plan.rates if variant_plan is not None else None, ) started_at = time.monotonic() try: @@ -269,6 +280,7 @@ def find_backfill_candidates_activity(inputs: FindBackfillCandidatesInputs) -> F skipped = sum(1 for c in candidates[:walked_through] if c.session_id in overtaken) return FindBackfillCandidatesOutput( + variant_sampling_rates=variant_plan.rates if variant_plan is not None else None, started_from_cursor_end_time=backfill.cursor_end_time, started_from_cursor_session_id=backfill.cursor_session_id, candidates=[ diff --git a/products/replay_vision/backend/temporal/activities/create_observation.py b/products/replay_vision/backend/temporal/activities/create_observation.py index 613980ca0428..f6130d83841e 100644 --- a/products/replay_vision/backend/temporal/activities/create_observation.py +++ b/products/replay_vision/backend/temporal/activities/create_observation.py @@ -48,8 +48,13 @@ _SCAN_BLOCKED_DEDUP_TTL_SECONDS = 60 * 60 -def _build_scanner_snapshot(scanner: ReplayScanner) -> dict[str, Any]: - return ScannerSnapshot.from_scanner(scanner).model_dump(mode="json") +def _build_scanner_snapshot( + scanner: ReplayScanner, *, variant_sampling_rates: dict[str, float] | None = None +) -> dict[str, Any]: + snapshot = ScannerSnapshot.from_scanner(scanner) + if variant_sampling_rates is not None: + snapshot = snapshot.model_copy(update={"variant_sampling_rates": variant_sampling_rates}) + return snapshot.model_dump(mode="json") def _capture_scan_blocked( @@ -264,10 +269,17 @@ def _create_observation(inputs: CreateObservationInputs) -> CreateObservationOut # Backfill applies run the frozen config, not the scanner's current one. if backfill is not None: frozen = BackfillScannerSnapshot.load_for_backfill(backfill.id, backfill.scanner_snapshot) - snapshot_dict = frozen.to_observation_snapshot().model_dump(mode="json") + observation_snapshot = frozen.to_observation_snapshot() + if inputs.variant_sampling_rates is not None: + # The rates are the dispatching tick's, computed from live rollout shares, so they + # ride in on the inputs rather than living in the frozen config. + observation_snapshot = observation_snapshot.model_copy( + update={"variant_sampling_rates": inputs.variant_sampling_rates} + ) + snapshot_dict = observation_snapshot.model_dump(mode="json") priced_model = frozen.model else: - snapshot_dict = _build_scanner_snapshot(scanner) + snapshot_dict = _build_scanner_snapshot(scanner, variant_sampling_rates=inputs.variant_sampling_rates) priced_model = scanner.model # Deliberately check-then-act: the snapshot doesn't count enqueue claims, so a concurrent burst can diff --git a/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py b/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py index 4c6b21a579cd..80ab68ed26d9 100644 --- a/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py +++ b/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py @@ -29,6 +29,7 @@ ScannerCandidateQuery, WindowedCandidateQuery, ) +from products.replay_vision.backend.queries.variant_sampling import variant_sampling_plan_for_scope from products.replay_vision.backend.temporal.constants import ( DEEP_SWEEP_INTERVAL, DEEP_SWEEP_MAX_EXECUTION_SECONDS, @@ -148,6 +149,13 @@ def find_scanner_candidates_activity(inputs: FindScannerCandidatesInputs) -> Fin started_at = time.monotonic() limit = inputs.candidate_limit if inputs.candidate_limit is not None else DEFAULT_CANDIDATE_LIMIT + variant_plan = variant_sampling_plan_for_scope( + scanner.team, + scope=scanner.experiment_scope(), + scanner_config=scanner.scanner_config, + sampling_rate=scanner.sampling_rate, + ) + variant_rates = variant_plan.rates if variant_plan is not None else None candidate_query = ScannerCandidateQuery( team=scanner.team, query=query, @@ -163,6 +171,7 @@ def find_scanner_candidates_activity(inputs: FindScannerCandidatesInputs) -> Fin # Exclusion is applied below against the fetched batch instead. skip_negative_blocklists=True, scanner_id=str(scanner.id), + variant_sampling_rates=variant_rates, ) try: batch = candidate_query.run_batch(limit) @@ -209,6 +218,7 @@ def find_scanner_candidates_activity(inputs: FindScannerCandidatesInputs) -> Fin candidate_query, deep_limit, seconds_remaining=_seconds_left(started_at), + variant_sampling_rates=variant_rates, ) except Exception: # Best-effort catch-up must never fail the tick: the fast pass has already found and @@ -262,6 +272,7 @@ def find_scanner_candidates_activity(inputs: FindScannerCandidatesInputs) -> Fin priming_candidates=[ CandidateSessionPayload(session_id=c.session_id, session_end=c.session_end) for c in priming_candidates ], + variant_sampling_rates=variant_rates, ) @@ -345,6 +356,7 @@ def _deep_sweep( limit: int, *, seconds_remaining: float, + variant_sampling_rates: dict[str, float] | None = None, ) -> tuple[list[CandidateSession], _DeepProgress | None]: """Catch-up pass behind the fast watermark with the full events lookback. @@ -417,6 +429,7 @@ def _deep_sweep( candidate_limit=limit, max_execution_time_seconds=budget, scanner_id=str(scanner.id), + variant_sampling_rates=variant_sampling_rates, ) # Stamped before the query, so a pass that times out still counts against the cadence. Queryset # update rather than save(): `updated_at` means "the scanner was edited", which the skip above reads. diff --git a/products/replay_vision/backend/temporal/backfill_types.py b/products/replay_vision/backend/temporal/backfill_types.py index 05a1e474e1ce..e60baebeb788 100644 --- a/products/replay_vision/backend/temporal/backfill_types.py +++ b/products/replay_vision/backend/temporal/backfill_types.py @@ -55,6 +55,9 @@ class FindBackfillCandidatesOutput(BaseModel, frozen=True): # False only when the walk genuinely reached the window start: a batch the caps truncated still # has work below the cursor. The tick completes the backfill exactly when this is False. more_work_below_cursor: bool + # The balanced per-variant rates this tick's candidates were sampled at (experiment scanners + # with balancing on; None otherwise and on pre-deploy histories). + variant_sampling_rates: dict[str, float] | None = None class AdvanceBackfillCursorInputs(BaseModel, frozen=True): diff --git a/products/replay_vision/backend/temporal/backfill_workflow.py b/products/replay_vision/backend/temporal/backfill_workflow.py index 65277c2389fe..13e4f937918f 100644 --- a/products/replay_vision/backend/temporal/backfill_workflow.py +++ b/products/replay_vision/backend/temporal/backfill_workflow.py @@ -87,7 +87,9 @@ async def run(self, inputs: BackfillTickInputs) -> None: if find_result.candidates: # Deterministic child ids collide with live-sweep applies of the same (scanner, session), so # a session observed live is skipped here for free. - await asyncio.gather(*(self._start_child(inputs, c) for c in find_result.candidates)) + await asyncio.gather( + *(self._start_child(inputs, c, find_result.variant_sampling_rates) for c in find_result.candidates) + ) # The activity decides how far the walk got, since it can step over sessions that were # already observed but must not step over ones the caps held back. @@ -126,7 +128,12 @@ async def _delete_own_schedule(self, inputs: BackfillTickInputs) -> None: "replay_vision.backfill_schedule_delete_failed", extra={"backfill_id": str(inputs.backfill_id)} ) - async def _start_child(self, inputs: BackfillTickInputs, candidate: CandidateSessionPayload) -> None: + async def _start_child( + self, + inputs: BackfillTickInputs, + candidate: CandidateSessionPayload, + variant_sampling_rates: dict[str, float] | None, + ) -> None: try: await wf.start_child_workflow( APPLY_SCANNER_WORKFLOW_NAME, @@ -136,6 +143,7 @@ async def _start_child(self, inputs: BackfillTickInputs, candidate: CandidateSes team_id=inputs.team_id, triggered_by=ObservationTrigger.BACKFILL, backfill_id=inputs.backfill_id, + variant_sampling_rates=variant_sampling_rates, ), id=build_apply_scanner_workflow_id(inputs.scanner_id, candidate.session_id), task_queue=settings.REPLAY_VISION_TASK_QUEUE, diff --git a/products/replay_vision/backend/temporal/snapshots.py b/products/replay_vision/backend/temporal/snapshots.py index 0e2b846535d8..358578ef4214 100644 --- a/products/replay_vision/backend/temporal/snapshots.py +++ b/products/replay_vision/backend/temporal/snapshots.py @@ -35,6 +35,10 @@ class ScannerSnapshot(BaseModel, frozen=True): experiment_targeting: dict[str, Any] | None = None sampling_rate: float | None = None sampling_mode: str | None = None + # The balanced per-variant rates the dispatching tick sampled at, so even per-variant counts + # don't read as even traffic ("control sampled at 1.1%, test at 100%"). Set by the tick, not + # `from_scanner`: the rates depend on live rollout shares the scanner row doesn't carry. + variant_sampling_rates: dict[str, float] | None = None # How a monitor `yes` verdict is re-checked: `off` (one pass), `shadow` (draw again, record the result, serve the # first pass), or `enforce` (serve the `yes` only when the second draw agrees, else the dissent). A plain string # so a retired mode never breaks old-row loads. diff --git a/products/replay_vision/backend/temporal/sweep_types.py b/products/replay_vision/backend/temporal/sweep_types.py index 8919f944731d..93751d3f8d7f 100644 --- a/products/replay_vision/backend/temporal/sweep_types.py +++ b/products/replay_vision/backend/temporal/sweep_types.py @@ -56,6 +56,10 @@ class FindScannerCandidatesOutput(BaseModel, frozen=True): # the workflow falls back to deriving the position from `candidates`/`swept_through`. keyset_end: dt.datetime | None = None keyset_session_id: str = "" + # The balanced per-variant rates this tick's candidates were sampled at (experiment scanners + # with balancing on; None otherwise and on pre-deploy histories). Recorded onto each + # observation's snapshot, so even per-variant counts don't read as even traffic. + variant_sampling_rates: dict[str, float] | None = None class RefreshPromptSuggestionInputs(BaseModel, frozen=True): diff --git a/products/replay_vision/backend/temporal/sweep_workflow.py b/products/replay_vision/backend/temporal/sweep_workflow.py index ea859e8a7715..8d05db4d15dd 100644 --- a/products/replay_vision/backend/temporal/sweep_workflow.py +++ b/products/replay_vision/backend/temporal/sweep_workflow.py @@ -160,12 +160,14 @@ async def run(self, inputs: SweepScannerInputs) -> None: ), ) # A no-op when both lists are empty. First failure aborts the gather and skips the advance; - # UNIQUE(scanner_id, session_id) dedups retries. + # UNIQUE(scanner_id, session_id) dedups retries. Priming samples everything, so the tick's + # balanced rates describe only the fast and deep candidates. await asyncio.gather( *( - self._start_child(inputs, c) - for c in (*find_result.candidates, *find_result.deep_candidates, *find_result.priming_candidates) - ) + self._start_child(inputs, c, find_result.variant_sampling_rates) + for c in (*find_result.candidates, *find_result.deep_candidates) + ), + *(self._start_child(inputs, c, None) for c in find_result.priming_candidates), ) if find_result.keyset_end is not None: @@ -215,7 +217,12 @@ async def _advance_watermark( retry_policy=common.RetryPolicy(maximum_attempts=3), ) - async def _start_child(self, inputs: SweepScannerInputs, candidate: CandidateSessionPayload) -> None: + async def _start_child( + self, + inputs: SweepScannerInputs, + candidate: CandidateSessionPayload, + variant_sampling_rates: dict[str, float] | None, + ) -> None: try: await wf.start_child_workflow( APPLY_SCANNER_WORKFLOW_NAME, @@ -224,6 +231,7 @@ async def _start_child(self, inputs: SweepScannerInputs, candidate: CandidateSes session_id=candidate.session_id, team_id=inputs.team_id, triggered_by=ObservationTrigger.SCHEDULE, + variant_sampling_rates=variant_sampling_rates, ), id=build_apply_scanner_workflow_id(inputs.scanner_id, candidate.session_id), task_queue=settings.REPLAY_VISION_TASK_QUEUE, diff --git a/products/replay_vision/backend/temporal/types.py b/products/replay_vision/backend/temporal/types.py index 571b8501bf16..308d6da4038d 100644 --- a/products/replay_vision/backend/temporal/types.py +++ b/products/replay_vision/backend/temporal/types.py @@ -79,6 +79,9 @@ class ApplyScannerInputs(BaseModel, frozen=True): triggered_by_user_id: int | None = None # Set only for backfill-triggered applies; routes observation creation to the backfill's frozen snapshot. backfill_id: UUID | None = None + # The balanced per-variant rates the dispatching tick sampled at, recorded onto the + # observation's snapshot (experiment scanners with balancing on; None otherwise). + variant_sampling_rates: dict[str, float] | None = None class CreateObservationInputs(BaseModel, frozen=True): @@ -89,6 +92,7 @@ class CreateObservationInputs(BaseModel, frozen=True): triggered_by_user_id: int | None workflow_id: str backfill_id: UUID | None = None + variant_sampling_rates: dict[str, float] | None = None class CreateObservationOutput(BaseModel, frozen=True): diff --git a/products/replay_vision/backend/temporal/workflow.py b/products/replay_vision/backend/temporal/workflow.py index cd92f7ef52c1..c5b482360ce8 100644 --- a/products/replay_vision/backend/temporal/workflow.py +++ b/products/replay_vision/backend/temporal/workflow.py @@ -298,6 +298,7 @@ async def run(self, inputs: ApplyScannerInputs) -> None: triggered_by_user_id=inputs.triggered_by_user_id, workflow_id=workflow_id, backfill_id=inputs.backfill_id, + variant_sampling_rates=inputs.variant_sampling_rates, ), start_to_close_timeout=dt.timedelta(seconds=30), schedule_to_close_timeout=STATE_ACTIVITY_SCHEDULE_TO_CLOSE, diff --git a/products/replay_vision/backend/tests/test_scanner_candidate_query.py b/products/replay_vision/backend/tests/test_scanner_candidate_query.py index a40682b3b033..96fe86771823 100644 --- a/products/replay_vision/backend/tests/test_scanner_candidate_query.py +++ b/products/replay_vision/backend/tests/test_scanner_candidate_query.py @@ -2,7 +2,7 @@ import pytest import time_machine -from posthog.test.base import ClickhouseTestMixin, _create_event +from posthog.test.base import ClickhouseTestMixin, _create_event, flush_persons_and_events from posthog.schema import ( EventPropertyFilter, @@ -864,3 +864,86 @@ def test_descending_keyset_walk_partitions_the_enumerated_window(self, team) -> # Newest-first, tie broken by descending session_id, every enumerated session exactly once. assert walked == ["sess-tied", "sess-0", "sess-1", "sess-2", "sess-3", "sess-4"] + + +class TestBalancedVariantSamplingAgainstClickHouse(ClickhouseTestMixin): + @pytest.fixture(autouse=True) + def _frozen_clock(self): + with time_machine.travel(_FROZEN_TIME, tick=False): + yield + + def setup_method(self, _method) -> None: + sync_execute(TRUNCATE_SESSION_REPLAY_EVENTS_TABLE_SQL()) + + def _exposed_session(self, team, distinct_id: str, variant: str, session_id: str) -> None: + settle_bound = _NOW - SETTLE_INTERVAL + create_person(team=team, distinct_ids=[distinct_id]) + _create_event( + team=team, + event="$feature_flag_called", + distinct_id=distinct_id, + timestamp=_NOW - dt.timedelta(days=1), + properties={"$feature_flag": "balanced-flag", "$feature_flag_response": variant}, + ) + produce_replay_summary( + team_id=team.id, + session_id=session_id, + distinct_id=distinct_id, + first_timestamp=(settle_bound - dt.timedelta(minutes=20)).isoformat(), + last_timestamp=(settle_bound - dt.timedelta(minutes=10)).isoformat(), + active_milliseconds=30_000, + ) + + @pytest.mark.django_db + def test_per_variant_rates_gate_candidates_by_attributed_variant(self, team) -> None: + # Per-variant thresholds must select by each session's attributed variant, not by one + # scanner-wide rate: a broken join projection or multiIf would sample both arms alike. + from posthog.models import User + + from products.experiments.backend.models.experiment import Experiment + from products.feature_flags.backend.models.feature_flag import FeatureFlag + + creator = User.objects.create_and_join(team.organization, "balanced@posthog.com", "testtest") + flag = FeatureFlag.objects.create( + team=team, + key="balanced-flag", + created_by=creator, + filters={ + "multivariate": { + "variants": [ + {"key": "control", "rollout_percentage": 50}, + {"key": "test", "rollout_percentage": 50}, + ] + } + }, + ) + experiment = Experiment.objects.create( + team=team, + name="balanced", + feature_flag=flag, + created_by=creator, + start_date=_NOW - dt.timedelta(days=7), + exposure_criteria={}, + ) + self._exposed_session(team, "control-user", "control", "control-session") + self._exposed_session(team, "test-user", "test", "test-session") + flush_persons_and_events() + + def run(rates: dict[str, float] | None): + query = RecordingsQuery.model_validate( + {"kind": "RecordingsQuery", "experiment_exposure": {"experiment_id": experiment.id}} + ) + return ScannerCandidateQuery( + team=team, + query=query, + user=creator, + last_swept_at=_NOW - dt.timedelta(days=2), + sampling_rate=1.0, + sampling_salt="scanner-1", + variant_sampling_rates=rates, + ).run() + + # Control run: without rates the exposure join keeps both arms. + assert {c.session_id for c in run(None)} == {"control-session", "test-session"} + # Rate 0 vs 1 is deterministic whatever the hash: only the fully sampled arm survives. + assert {c.session_id for c in run({"control": 0.0, "test": 1.0})} == {"test-session"} diff --git a/products/replay_vision/backend/tests/test_temporal.py b/products/replay_vision/backend/tests/test_temporal.py index fd1b9849aa98..5d98245dd3d9 100644 --- a/products/replay_vision/backend/tests/test_temporal.py +++ b/products/replay_vision/backend/tests/test_temporal.py @@ -291,6 +291,28 @@ def test_creates_row_in_pending_with_workflow_id_and_snapshot(self) -> None: assert observation.started_at is None # set when transitioning to running, not here assert observation.completed_at is None + def test_records_the_dispatching_ticks_variant_sampling_rates_on_the_snapshot(self) -> None: + # The variants readout explains even per-variant counts with these rates; a snapshot built + # only from the scanner row would silently drop them, since the row never carries them. + scanner = _make_scanner( + scanner_type=ScannerType.EXPERIMENT, scanner_config={"prompt": "p", "experiment_id": 42} + ) + result = create_observation_activity( + CreateObservationInputs( + scanner_id=scanner.id, + team_id=scanner.team_id, + session_id="sess-balanced", + triggered_by=ObservationTrigger.SCHEDULE, + triggered_by_user_id=None, + workflow_id="wf-balanced", + variant_sampling_rates={"control": 0.055, "test": 0.5}, + ) + ) + + assert result.observation_id is not None + observation = ReplayObservation.objects.get(id=result.observation_id) + assert observation.scanner_snapshot["variant_sampling_rates"] == {"control": 0.055, "test": 0.5} + def test_decays_enqueue_claim_once_the_row_exists(self) -> None: # A claim that never decays holds a phantom cap slot for the full TTL. scanner = _make_scanner() diff --git a/products/replay_vision/backend/tests/test_variant_sampling.py b/products/replay_vision/backend/tests/test_variant_sampling.py new file mode 100644 index 000000000000..5697d3f62365 --- /dev/null +++ b/products/replay_vision/backend/tests/test_variant_sampling.py @@ -0,0 +1,79 @@ +import pytest + +from posthog.hogql import ast + +from products.replay_vision.backend.queries.scanner_candidate_query import ( + SAMPLE_RATE_PRECISION, + variant_sampling_predicate, +) +from products.replay_vision.backend.queries.variant_sampling import plan_variant_sampling + + +class TestPlanVariantSampling: + def test_uneven_split_gets_even_coverage_at_the_same_total(self) -> None: + # The doc's example: 10% over a 90/10 split keeps ~50 of each per 1,000 exposed sessions, + # not 90 and 10, and the total budget is unchanged. + plan = plan_variant_sampling(0.1, {"control": 0.9, "test": 0.1}, None) + + assert plan is not None + assert plan.rates["control"] == pytest.approx(0.05 / 0.9) + assert plan.rates["test"] == pytest.approx(0.5) + assert plan.effective_rate == pytest.approx(0.1) + # Equal coverage: each variant contributes the same absolute fraction. + assert plan.population_shares["control"] * plan.rates["control"] == pytest.approx( + plan.population_shares["test"] * plan.rates["test"] + ) + + def test_a_capped_small_variant_frees_its_budget_for_the_others(self) -> None: + # 95/5 at 20%: the small arm is sampled whole (rate 1) and cannot supply more, so the + # leftover budget flows to control instead of shrinking the total. + plan = plan_variant_sampling(0.2, {"control": 0.95, "test": 0.05}, None) + + assert plan is not None + assert plan.rates["test"] == 1.0 + assert plan.rates["control"] == pytest.approx(0.15 / 0.95) + assert plan.effective_rate == pytest.approx(0.2) + + def test_selected_variants_normalize_within_the_watched_set(self) -> None: + plan = plan_variant_sampling(0.5, {"control": 0.5, "test": 0.4, "beta": 0.1}, ["control", "test"]) + + assert plan is not None + assert set(plan.rates) == {"control", "test"} + assert plan.population_shares["control"] == pytest.approx(5 / 9) + + @pytest.mark.parametrize( + "selected,shares", + [ + (["test"], {"control": 0.5, "test": 0.5}), + (None, {}), + ], + ) + def test_no_plan_when_balancing_changes_nothing(self, selected, shares) -> None: + assert plan_variant_sampling(0.1, shares, selected) is None + + def test_a_zero_share_variant_is_sampled_whole_without_eating_budget(self) -> None: + plan = plan_variant_sampling(0.1, {"control": 1.0, "test": 0.0}, None) + + assert plan is not None + assert plan.rates["test"] == 1.0 + assert plan.rates["control"] == pytest.approx(0.1) + + +class TestVariantSamplingPredicate: + def test_builds_one_threshold_arm_per_variant_over_the_shared_hash(self) -> None: + predicate = variant_sampling_predicate({"control": 0.25, "test": 1.0}, "salt-1") + + assert isinstance(predicate, ast.Call) and predicate.name == "multiIf" + # Two (condition, then) pairs plus the fall-through. + assert len(predicate.args) == 5 + thresholds = [ + arg.right.value + for arg in predicate.args[1::2] + if isinstance(arg, ast.CompareOperation) and isinstance(arg.right, ast.Constant) + ] + assert thresholds == [round(0.25 * SAMPLE_RATE_PRECISION), SAMPLE_RATE_PRECISION] + fall_through = predicate.args[-1] + assert isinstance(fall_through, ast.Constant) and fall_through.value is False + + def test_no_predicate_when_every_variant_is_sampled_whole(self) -> None: + assert variant_sampling_predicate({"control": 1.0, "test": 1.0}, "salt-1") is None diff --git a/products/replay_vision/frontend/generated/api.schemas.ts b/products/replay_vision/frontend/generated/api.schemas.ts index e0a0fc5df983..a14663c00c1c 100644 --- a/products/replay_vision/frontend/generated/api.schemas.ts +++ b/products/replay_vision/frontend/generated/api.schemas.ts @@ -628,6 +628,12 @@ export const ScannerTypeEnumApi = { Experiment: 'experiment', } as const +/** + * Experiment scanners with balanced sampling: the 0..1 rate each watched variant was sampled at by the tick that dispatched this scan. Null otherwise, so even per-variant counts can be read against the rates that produced them. + * @nullable + */ +export type ScannerSnapshotApiVariantSamplingRates = { [key: string]: number } | null + /** * Mirrors `temporal.types.ScannerSnapshot` for OpenAPI generation. */ @@ -654,6 +660,11 @@ export interface ScannerSnapshotApi { scanner_config: unknown /** How a monitor `yes` was re-checked at run time: `off` (one pass, the default), `shadow` (second draw recorded only), or `enforce` (the `yes` stands only when the second draw agrees). */ verify_positives: string + /** + * Experiment scanners with balanced sampling: the 0..1 rate each watched variant was sampled at by the tick that dispatched this scan. Null otherwise, so even per-variant counts can be read against the rates that produced them. + * @nullable + */ + variant_sampling_rates?: ScannerSnapshotApiVariantSamplingRates } /** diff --git a/services/mcp/src/api/generated.ts b/services/mcp/src/api/generated.ts index cab4529c69ef..b34f3896b5aa 100644 --- a/services/mcp/src/api/generated.ts +++ b/services/mcp/src/api/generated.ts @@ -62615,6 +62615,12 @@ export namespace Schemas { Ineligible: 'ineligible', } as const; + /** + * Experiment scanners with balanced sampling: the 0..1 rate each watched variant was sampled at by the tick that dispatched this scan. Null otherwise, so even per-variant counts can be read against the rates that produced them. + * @nullable + */ + export type ScannerSnapshotVariantSamplingRates = {[key: string]: number} | null; + /** * Mirrors `temporal.types.ScannerSnapshot` for OpenAPI generation. */ @@ -62641,6 +62647,11 @@ export namespace Schemas { scanner_config: unknown; /** How a monitor `yes` was re-checked at run time: `off` (one pass, the default), `shadow` (second draw recorded only), or `enforce` (the `yes` stands only when the second draw agrees). */ verify_positives: string; + /** + * Experiment scanners with balanced sampling: the 0..1 rate each watched variant was sampled at by the tick that dispatched this scan. Null otherwise, so even per-variant counts can be read against the rates that produced them. + * @nullable + */ + variant_sampling_rates?: ScannerSnapshotVariantSamplingRates; } /** From 8e73aca7bfade86dd16e74c86bf8c5641891bc30 Mon Sep 17 00:00:00 2001 From: Kim Svatos Dugan <147102038+ksvat@users.noreply.github.com> Date: Wed, 30 Sep 2026 08:51:34 -0700 Subject: [PATCH 2/5] fix(replay-vision): zero-budget plans sample nothing; honor singular variant scope Two review findings: a paused scanner (rate 0) now yields all-zero per-variant rates instead of sampling a zero-share variant whole, and a legacy column scope's singular `variant` counts as the one watched arm, so single-arm targeting gets no balancing plan. Co-Authored-By: Claude Fable 5 Generated-By: PostHog Desktop Task-Id: c2db2a38-d03f-4578-af4a-40f61dc700a4 --- .../backend/queries/variant_sampling.py | 9 ++++- .../backend/tests/test_variant_sampling.py | 36 +++++++++++++++++++ 2 files changed, 44 insertions(+), 1 deletion(-) diff --git a/products/replay_vision/backend/queries/variant_sampling.py b/products/replay_vision/backend/queries/variant_sampling.py index 4fbfaca40c08..4bd6d6d4e88d 100644 --- a/products/replay_vision/backend/queries/variant_sampling.py +++ b/products/replay_vision/backend/queries/variant_sampling.py @@ -53,6 +53,10 @@ def plan_variant_sampling( ) budget = max(0.0, min(1.0, sampling_rate)) + if budget <= 0: + # Rate 0 means paused; the cap-at-1 arm below must not turn "no budget" into + # "sample a zero-share variant whole". + return VariantSamplingPlan(rates=dict.fromkeys(selected_keys, 0.0), population_shares=population_shares) rates: dict[str, float] = {} remaining = set(selected_keys) while remaining: @@ -94,4 +98,7 @@ def variant_sampling_plan_for_scope( except ValidationError: return None assert scope is not None - return plan_variant_sampling(sampling_rate, shares, scope.get("variants")) + # A legacy column scope narrows with the singular `variant`; treating it as "every variant" + # would balance a population the exposure join has already narrowed to one arm. + selected = scope.get("variants") or ([scope["variant"]] if scope.get("variant") else None) + return plan_variant_sampling(sampling_rate, shares, selected) diff --git a/products/replay_vision/backend/tests/test_variant_sampling.py b/products/replay_vision/backend/tests/test_variant_sampling.py index 5697d3f62365..eb1763ff6189 100644 --- a/products/replay_vision/backend/tests/test_variant_sampling.py +++ b/products/replay_vision/backend/tests/test_variant_sampling.py @@ -1,4 +1,5 @@ import pytest +from posthog.test.base import BaseTest from posthog.hogql import ast @@ -58,6 +59,15 @@ def test_a_zero_share_variant_is_sampled_whole_without_eating_budget(self) -> No assert plan.rates["test"] == 1.0 assert plan.rates["control"] == pytest.approx(0.1) + def test_a_paused_scanner_samples_nothing(self) -> None: + # Rate 0 means paused; the cap must not turn "no budget" into sampling a zero-share + # variant whole. + plan = plan_variant_sampling(0.0, {"control": 1.0, "test": 0.0}, None) + + assert plan is not None + assert plan.rates == {"control": 0.0, "test": 0.0} + assert plan.effective_rate == 0.0 + class TestVariantSamplingPredicate: def test_builds_one_threshold_arm_per_variant_over_the_shared_hash(self) -> None: @@ -77,3 +87,29 @@ def test_builds_one_threshold_arm_per_variant_over_the_shared_hash(self) -> None def test_no_predicate_when_every_variant_is_sampled_whole(self) -> None: assert variant_sampling_predicate({"control": 1.0, "test": 1.0}, "salt-1") is None + + +class TestVariantSamplingPlanForScope(BaseTest): + def test_a_singular_legacy_variant_scope_watches_one_arm_and_gets_no_plan(self) -> None: + # A legacy column scope narrows with `variant` (singular). Reading only `variants` would + # treat it as "every variant" and balance a population the exposure join already narrowed. + from products.replay_vision.backend.queries.variant_sampling import variant_sampling_plan_for_scope + from products.replay_vision.backend.tests.helpers import create_experiment + + experiment = create_experiment(self.team, "single-arm-flag", launched=True, variants=["control", "test"]) + + singular = variant_sampling_plan_for_scope( + self.team, + scope={"experiment_id": experiment.pk, "variant": "test"}, + scanner_config={"prompt": "p"}, + sampling_rate=0.1, + ) + assert singular is None + + both = variant_sampling_plan_for_scope( + self.team, + scope={"experiment_id": experiment.pk, "variants": ["control", "test"]}, + scanner_config={"prompt": "p", "experiment_id": experiment.pk}, + sampling_rate=0.1, + ) + assert both is not None and set(both.rates) == {"control", "test"} From fb2a6e06fa52e090ef4af7516dc0bfd4b88e70eb Mon Sep 17 00:00:00 2001 From: Kim Svatos Dugan <147102038+ksvat@users.noreply.github.com> Date: Wed, 30 Sep 2026 08:51:37 -0700 Subject: [PATCH 3/5] chore(replay-vision): hoist test model imports to module scope Review cleanup: the balanced-sampling ClickHouse test imported User, Experiment, and FeatureFlag inside the test body; a test file has no import cycle to break, so they belong in the import section. Co-Authored-By: Claude Fable 5 Generated-By: PostHog Desktop Task-Id: c2db2a38-d03f-4578-af4a-40f61dc700a4 --- .../backend/tests/test_scanner_candidate_query.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/products/replay_vision/backend/tests/test_scanner_candidate_query.py b/products/replay_vision/backend/tests/test_scanner_candidate_query.py index 96fe86771823..ef56308deba9 100644 --- a/products/replay_vision/backend/tests/test_scanner_candidate_query.py +++ b/products/replay_vision/backend/tests/test_scanner_candidate_query.py @@ -15,10 +15,13 @@ from posthog.hogql import ast from posthog.clickhouse.client import sync_execute +from posthog.models import User from posthog.session_recordings.queries.test.session_replay_sql import produce_replay_summary from posthog.session_recordings.sql.session_replay_event_sql import TRUNCATE_SESSION_REPLAY_EVENTS_TABLE_SQL from posthog.test.persons import create_person +from products.experiments.backend.models.experiment import Experiment +from products.feature_flags.backend.models.feature_flag import FeatureFlag from products.replay_vision.backend.queries.scanner_candidate_query import ( BALANCED_SURFACING_THRESHOLD, DEFAULT_CANDIDATE_LIMIT, @@ -898,11 +901,6 @@ def _exposed_session(self, team, distinct_id: str, variant: str, session_id: str def test_per_variant_rates_gate_candidates_by_attributed_variant(self, team) -> None: # Per-variant thresholds must select by each session's attributed variant, not by one # scanner-wide rate: a broken join projection or multiIf would sample both arms alike. - from posthog.models import User - - from products.experiments.backend.models.experiment import Experiment - from products.feature_flags.backend.models.feature_flag import FeatureFlag - creator = User.objects.create_and_join(team.organization, "balanced@posthog.com", "testtest") flag = FeatureFlag.objects.create( team=team, From 40064b1a0fb0ccc057a54dad99be9a337693bf48 Mon Sep 17 00:00:00 2001 From: Kim Svatos Dugan <147102038+ksvat@users.noreply.github.com> Date: Wed, 30 Sep 2026 08:51:40 -0700 Subject: [PATCH 4/5] fix(replay-vision): plan variant sampling from exposure counts, experiment type only MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Per review: the plan's shares now come from counting the exposed persons per watched variant over the experiment window (an aggregate over the exposure select, metered as the scanner's reads, failing open to plain sampling), because the flag's rollout percentages diverge from the window's real mix whenever the rollout changed mid-experiment. The plan is also gated to the experiment scanner type, so legacy column-targeted scanners keep plain sampling on deploy. The volume estimate goes back to projecting with the plain rate: redistribution spends exactly the configured budget (effective_rate always equals the rate), so balancing never moves the projection — the invariant now has its own test, which also answers why toggling balance_variants needs no estimate or scout cost invalidation. Co-Authored-By: Claude Fable 5 Generated-By: PostHog Desktop Task-Id: c2db2a38-d03f-4578-af4a-40f61dc700a4 --- .../queries/scanner_volume_estimate.py | 15 +- .../backend/queries/variant_sampling.py | 143 ++++++++++++++---- .../backend/temporal/activities/backfill.py | 4 +- .../activities/find_scanner_candidates.py | 2 + .../tests/test_scanner_candidate_query.py | 37 +++++ .../backend/tests/test_variant_sampling.py | 72 +++++++-- 6 files changed, 215 insertions(+), 58 deletions(-) diff --git a/products/replay_vision/backend/queries/scanner_volume_estimate.py b/products/replay_vision/backend/queries/scanner_volume_estimate.py index b83637b8c51b..8729d27f8d31 100644 --- a/products/replay_vision/backend/queries/scanner_volume_estimate.py +++ b/products/replay_vision/backend/queries/scanner_volume_estimate.py @@ -28,7 +28,6 @@ eligibility_predicates, surfacing_score_predicate, ) -from products.replay_vision.backend.queries.variant_sampling import variant_sampling_plan_for_scope # The estimate always projects to a calendar month. ESTIMATE_WINDOW_DAYS = ESTIMATE_MONTH_DAYS @@ -320,16 +319,10 @@ def refresh_scanner_estimate( budget=budget, ch_user=ch_user, ) - # Balanced per-variant sampling caps a small variant at 1, so the fraction actually sampled can - # sit below the configured rate; the projection must use what the sweep will really keep. - variant_plan = variant_sampling_plan_for_scope( - scanner.team, - scope=scanner.experiment_scope(), - scanner_config=scanner.scanner_config, - sampling_rate=scanner.sampling_rate, - ) - effective_rate = variant_plan.effective_rate if variant_plan is not None else scanner.sampling_rate - projection = project_monthly_observations(estimate, effective_rate) + # Balanced sampling redistributes rather than shrinks the budget (VariantSamplingPlan. + # effective_rate always equals the configured rate), so the projection is the same with + # balancing on or off and needs no plan here. + projection = project_monthly_observations(estimate, scanner.sampling_rate) estimated_at = timezone.now() # Filtered write so a config edit racing the (slow) estimate query can't get stamped fresh with stale numbers. # JSONField quirk: `field=None` filters for JSON null, not SQL NULL, so the no-targeting case needs isnull. diff --git a/products/replay_vision/backend/queries/variant_sampling.py b/products/replay_vision/backend/queries/variant_sampling.py index 4bd6d6d4e88d..a61608c7ee5f 100644 --- a/products/replay_vision/backend/queries/variant_sampling.py +++ b/products/replay_vision/backend/queries/variant_sampling.py @@ -1,50 +1,76 @@ """Per-variant sampling rates for experiment scanners. -One rate per scanner samples variants proportionally to rollout, so a 90/10 split at a 10% rate -yields ~90 control observations per 10 test. Balanced sampling instead spends the same total -budget about evenly across the watched variants: each variant's rate is `r / (k · share)`, capped -at 1, with a capped variant's unspent budget redistributed to the others so a small variant can -never shrink the total. - -Rates are recomputed from the live rollout shares on every sweep, estimate, and backfill tick, so -a mid-experiment rollout change adjusts them on the next tick. The plan cannot create sessions a -small variant does not have: a capped variant is sampled whole and stays thin. +One rate per scanner samples variants proportionally to their traffic, so a 90/10 population at a +10% rate yields ~90 control observations per 10 test. Balanced sampling instead spends the same +total budget about evenly across the watched variants: each variant's rate is `r / (k · share)`, +capped at 1, with a capped variant's unspent budget redistributed to the others so a small variant +can never shrink the total. + +Shares come from counting the actually exposed persons per variant over the experiment window, not +from the flag's rollout percentages: a rollout that changed mid-experiment leaves the window's mix +far from the current percentages, and rates planned against the wrong mix overspend the budget. +The counts are recomputed on every sweep, estimate, and backfill tick, so a traffic shift adjusts +them on the next tick. The plan cannot create sessions a small variant does not have: a capped +variant is sampled whole and stays thin. """ from collections.abc import Sequence +import structlog from rest_framework.exceptions import ValidationError +from posthog.hogql import ast +from posthog.hogql.constants import HogQLGlobalSettings +from posthog.hogql.query import execute_hogql_query + +from posthog.clickhouse.client.connection import ClickHouseUser +from posthog.clickhouse.query_tagging import Feature, Product, tags_context from posthog.dataclasses import frozen from posthog.models.team import Team +from products.replay_vision.backend.models.replay_scanner import ScannerType + +logger = structlog.get_logger(__name__) + +# The counts aggregate scans the same exposure window the candidate query joins each tick, so it +# gets a matching but tighter budget: a plan that can't be computed falls back to plain sampling +# rather than holding the tick. +_EXPOSURE_COUNTS_MAX_EXECUTION_SECONDS = 60 + @frozen class VariantSamplingPlan: # Hash-threshold rate per watched variant, each in 0..1. rates: dict[str, float] - # Each watched variant's share of the watched population (normalized over the watched set). + # Each watched variant's share of the watched exposed population (normalized over the watched set). population_shares: dict[str, float] @property def effective_rate(self) -> float: - """The fraction of the watched population the plan samples overall: the scanner's own - rate until a cap binds, then less budget than asked finds sessions to spend itself on.""" + """The fraction of the watched population the plan samples overall. + + Redistribution spends the whole budget (shares sum to 1 and the rate is clamped to 0..1), + so this always equals the scanner's own rate — which is why the volume estimate can project + with the plain rate whether balancing is on or off. Kept as the checked statement of that + invariant rather than re-derived at call sites. + """ return sum(self.population_shares[variant] * rate for variant, rate in self.rates.items()) def plan_variant_sampling( - sampling_rate: float, rollout_shares: dict[str, float], selected: Sequence[str] | None + sampling_rate: float, exposure_weights: dict[str, float], selected: Sequence[str] | None ) -> VariantSamplingPlan | None: - """The per-variant plan, or None when balancing changes nothing (one watched variant or no shares). + """The per-variant plan, or None when balancing changes nothing (one watched variant or no weights). - Water-filling: every uncapped variant gets an equal slice of the remaining budget; a variant - whose whole population fits inside its slice is sampled at 1 and frees the rest of its slice. + ``exposure_weights`` is each variant's exposed-population size in any consistent unit (person + counts in production); it is normalized here. Water-filling: every uncapped variant gets an + equal slice of the remaining budget; a variant whose whole population fits inside its slice is + sampled at 1 and frees the rest of its slice. """ - selected_keys = [key for key in (selected if selected is not None else rollout_shares) if key in rollout_shares] + selected_keys = [key for key in (selected if selected is not None else exposure_weights) if key in exposure_weights] if len(selected_keys) <= 1: return None - raw = {key: max(0.0, float(rollout_shares[key])) for key in selected_keys} + raw = {key: max(0.0, float(exposure_weights[key])) for key in selected_keys} total = sum(raw.values()) population_shares = ( {key: value / total for key, value in raw.items()} @@ -75,30 +101,87 @@ def plan_variant_sampling( def variant_sampling_plan_for_scope( - team: Team, *, scope: dict | None, scanner_config: dict | None, sampling_rate: float + team: Team, + *, + scanner_type: str, + scope: dict | None, + scanner_config: dict | None, + sampling_rate: float, + scanner_id: str | None = None, ) -> VariantSamplingPlan | None: - """The plan for a scanner (or frozen snapshot) scope, from the flag's live rollout shares. + """The plan for an experiment scanner (or its frozen snapshot), from the window's exposure counts. - None when the scanner watches no experiment, balancing is off, only one variant is watched, or - the shares can't be read (the caller then falls back to plain sampling; the scan itself stays - the loud path for an unresolvable experiment). + None when the scanner is not the experiment type (legacy column-targeted scanners keep plain + sampling — flipping their behavior on deploy is not this function's call), balancing is off, + only one variant is watched, or the population can't be counted (the caller then falls back to + plain sampling; the scan itself stays the loud path for an unresolvable experiment). """ + if scanner_type != ScannerType.EXPERIMENT: + return None experiment_id = (scope or {}).get("experiment_id") if experiment_id is None: return None config = scanner_config if isinstance(scanner_config, dict) else {} if config.get("balance_variants") is False: return None + assert scope is not None + # A legacy column scope narrows with the singular `variant`; treating it as "every variant" + # would balance a population the exposure join has already narrowed to one arm. + selected = scope.get("variants") or ([scope["variant"]] if scope.get("variant") else None) + counts = _variant_exposure_counts(team, experiment_id=experiment_id, selected=selected, scanner_id=scanner_id) + if counts is None: + return None + return plan_variant_sampling(sampling_rate, counts, selected) + + +def _variant_exposure_counts( + team: Team, *, experiment_id: int, selected: Sequence[str] | None, scanner_id: str | None +) -> dict[str, float] | None: + """Exposed persons per watched variant over the experiment window, or None when uncountable. + + A watched variant with no exposures counts as 0, so a variant ramped down mid-experiment keeps + its real (small) weight instead of the rollout percentage's fiction. Fails open: a plan is an + optimization of how the budget is spent, so a failed count must cost this tick balance, not + candidates. + """ # Deferred: the experiments replay facade pulls in the recordings query modules, which circle # back into this package's importers. - from products.experiments.backend.facade.replay import variant_rollout_shares # noqa: PLC0415 + from products.experiments.backend.facade.replay import ( # noqa: PLC0415 + exposed_persons_select, + resolve_exposure_linkage, + ) try: - shares = variant_rollout_shares(team, experiment_id=experiment_id) + linkage = resolve_exposure_linkage( + team, experiment_id=experiment_id, variants=list(selected) if selected is not None else None + ) except ValidationError: return None - assert scope is not None - # A legacy column scope narrows with the singular `variant`; treating it as "every variant" - # would balance a population the exposure join has already narrowed to one arm. - selected = scope.get("variants") or ([scope["variant"]] if scope.get("variant") else None) - return plan_variant_sampling(sampling_rate, shares, selected) + counts_query = ast.SelectQuery( + select=[ + ast.Field(chain=["variant"]), + ast.Alias( + alias="exposed", expr=ast.Call(name="count", distinct=True, args=[ast.Field(chain=["person_id"])]) + ), + ], + select_from=ast.JoinExpr(table=exposed_persons_select(linkage, include_multiple_variant=False)), + group_by=[ast.Field(chain=["variant"])], + ) + try: + with tags_context(product=Product.REPLAY_VISION, feature=Feature.ENRICHMENT, scanner_id=scanner_id): + response = execute_hogql_query( + counts_query, + team=team, + query_type="ReplayVisionVariantExposureCountsQuery", + settings=HogQLGlobalSettings(max_execution_time=_EXPOSURE_COUNTS_MAX_EXECUTION_SECONDS), + ch_user=ClickHouseUser.REPLAY_VISION, + ) + except Exception: + logger.warning( + "replay_vision.variant_exposure_counts_failed", experiment_id=experiment_id, scanner_id=scanner_id + ) + return None + counts: dict[str, float] = dict.fromkeys(linkage.requested_variants, 0.0) + for row in response.results or []: + counts[str(row[0])] = float(row[1]) + return counts diff --git a/products/replay_vision/backend/temporal/activities/backfill.py b/products/replay_vision/backend/temporal/activities/backfill.py index cac2849f8476..6a8ee0a08660 100644 --- a/products/replay_vision/backend/temporal/activities/backfill.py +++ b/products/replay_vision/backend/temporal/activities/backfill.py @@ -169,14 +169,16 @@ def find_backfill_candidates_activity(inputs: FindBackfillCandidatesInputs) -> F ) from exc query = apply_experiment_targeting(query, snapshot.experiment_scope()) - # Live rollout shares against the frozen scope, matching the sweep: the same salted hash plus + # Live exposure counts against the frozen scope, matching the sweep: the same salted hash plus # the same rates keep sampling decisions stable between a live sweep and a backfill of the # same range. variant_plan = variant_sampling_plan_for_scope( backfill.team, + scanner_type=snapshot.scanner_type, scope=snapshot.experiment_scope(), scanner_config=snapshot.scanner_config, sampling_rate=snapshot.sampling_rate, + scanner_id=str(backfill.scanner_id), ) candidate_query = WindowedCandidateQuery( team=backfill.team, diff --git a/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py b/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py index 80ab68ed26d9..9be968a0c7b1 100644 --- a/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py +++ b/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py @@ -151,9 +151,11 @@ def find_scanner_candidates_activity(inputs: FindScannerCandidatesInputs) -> Fin limit = inputs.candidate_limit if inputs.candidate_limit is not None else DEFAULT_CANDIDATE_LIMIT variant_plan = variant_sampling_plan_for_scope( scanner.team, + scanner_type=scanner.scanner_type, scope=scanner.experiment_scope(), scanner_config=scanner.scanner_config, sampling_rate=scanner.sampling_rate, + scanner_id=str(scanner.id), ) variant_rates = variant_plan.rates if variant_plan is not None else None candidate_query = ScannerCandidateQuery( diff --git a/products/replay_vision/backend/tests/test_scanner_candidate_query.py b/products/replay_vision/backend/tests/test_scanner_candidate_query.py index ef56308deba9..9412c525ae4a 100644 --- a/products/replay_vision/backend/tests/test_scanner_candidate_query.py +++ b/products/replay_vision/backend/tests/test_scanner_candidate_query.py @@ -897,6 +897,43 @@ def _exposed_session(self, team, distinct_id: str, variant: str, session_id: str active_milliseconds=30_000, ) + def _experiment(self, team, creator, *, variants=("control", "test")): + flag = FeatureFlag.objects.create( + team=team, + key="balanced-flag", + created_by=creator, + filters={ + "multivariate": { + "variants": [{"key": key, "rollout_percentage": 100 // len(variants)} for key in variants] + } + }, + ) + return Experiment.objects.create( + team=team, + name="balanced", + feature_flag=flag, + created_by=creator, + start_date=_NOW - dt.timedelta(days=7), + exposure_criteria={}, + ) + + @pytest.mark.django_db + def test_variant_exposure_counts_zero_fill_watched_variants(self, team) -> None: + # The plan's shares come from these counts; a variant miscounted (or dropped instead of + # zero-filled) plans the budget against the wrong population. + from products.replay_vision.backend.queries.variant_sampling import _variant_exposure_counts + + creator = User.objects.create_and_join(team.organization, "counts@posthog.com", "testtest") + experiment = self._experiment(team, creator, variants=("control", "test", "beta")) + self._exposed_session(team, "counts-control-a", "control", "counts-session-a") + self._exposed_session(team, "counts-control-b", "control", "counts-session-b") + self._exposed_session(team, "counts-test", "test", "counts-session-c") + flush_persons_and_events() + + counts = _variant_exposure_counts(team, experiment_id=experiment.id, selected=None, scanner_id="scanner-1") + + assert counts == {"control": 2.0, "test": 1.0, "beta": 0.0} + @pytest.mark.django_db def test_per_variant_rates_gate_candidates_by_attributed_variant(self, team) -> None: # Per-variant thresholds must select by each session's attributed variant, not by one diff --git a/products/replay_vision/backend/tests/test_variant_sampling.py b/products/replay_vision/backend/tests/test_variant_sampling.py index eb1763ff6189..c9d88d9348c3 100644 --- a/products/replay_vision/backend/tests/test_variant_sampling.py +++ b/products/replay_vision/backend/tests/test_variant_sampling.py @@ -1,8 +1,10 @@ import pytest from posthog.test.base import BaseTest +from unittest.mock import patch from posthog.hogql import ast +from products.replay_vision.backend.models.replay_scanner import ScannerType from products.replay_vision.backend.queries.scanner_candidate_query import ( SAMPLE_RATE_PRECISION, variant_sampling_predicate, @@ -68,6 +70,32 @@ def test_a_paused_scanner_samples_nothing(self) -> None: assert plan.rates == {"control": 0.0, "test": 0.0} assert plan.effective_rate == 0.0 + @pytest.mark.parametrize( + "rate,weights", + [ + (0.1, {"control": 900, "test": 100}), + # The cap binds on the small arm; redistribution still spends the whole budget. + (0.2, {"control": 950, "test": 50}), + (0.5, {"control": 500, "test": 400, "beta": 100}), + ], + ) + def test_balancing_never_changes_the_projected_volume(self, rate, weights) -> None: + # This invariant is what lets the volume estimate (and the scout cost check) project with + # the plain rate whether balancing is on or off, cap or no cap. + from products.replay_vision.backend.queries.scanner_volume_estimate import ( + ScannerVolumeEstimate, + project_monthly_observations, + ) + + plan = plan_variant_sampling(rate, weights, None) + + assert plan is not None + assert plan.effective_rate == pytest.approx(rate) + estimate = ScannerVolumeEstimate(matched_sessions=7_000, effective_window_days=7) + assert project_monthly_observations(estimate, plan.effective_rate) == project_monthly_observations( + estimate, rate + ) + class TestVariantSamplingPredicate: def test_builds_one_threshold_arm_per_variant_over_the_shared_hash(self) -> None: @@ -90,26 +118,38 @@ def test_no_predicate_when_every_variant_is_sampled_whole(self) -> None: class TestVariantSamplingPlanForScope(BaseTest): + def _plan(self, *, scanner_type=ScannerType.EXPERIMENT, scope, counts): + from products.replay_vision.backend.queries.variant_sampling import variant_sampling_plan_for_scope + + with patch( + "products.replay_vision.backend.queries.variant_sampling._variant_exposure_counts", + return_value=counts, + ): + return variant_sampling_plan_for_scope( + self.team, + scanner_type=scanner_type, + scope=scope, + scanner_config={"prompt": "p"}, + sampling_rate=0.1, + ) + def test_a_singular_legacy_variant_scope_watches_one_arm_and_gets_no_plan(self) -> None: # A legacy column scope narrows with `variant` (singular). Reading only `variants` would # treat it as "every variant" and balance a population the exposure join already narrowed. - from products.replay_vision.backend.queries.variant_sampling import variant_sampling_plan_for_scope - from products.replay_vision.backend.tests.helpers import create_experiment - - experiment = create_experiment(self.team, "single-arm-flag", launched=True, variants=["control", "test"]) + counts = {"control": 60.0, "test": 40.0} - singular = variant_sampling_plan_for_scope( - self.team, - scope={"experiment_id": experiment.pk, "variant": "test"}, - scanner_config={"prompt": "p"}, - sampling_rate=0.1, - ) + singular = self._plan(scope={"experiment_id": 42, "variant": "test"}, counts=counts) assert singular is None - both = variant_sampling_plan_for_scope( - self.team, - scope={"experiment_id": experiment.pk, "variants": ["control", "test"]}, - scanner_config={"prompt": "p", "experiment_id": experiment.pk}, - sampling_rate=0.1, - ) + both = self._plan(scope={"experiment_id": 42, "variants": ["control", "test"]}, counts=counts) assert both is not None and set(both.rates) == {"control", "test"} + + def test_only_the_experiment_type_balances(self) -> None: + # A legacy scanner targeting an experiment through the column has no `balance_variants` + # key; defaulting it on would switch its sampling behavior on deploy. + plan = self._plan( + scanner_type=ScannerType.MONITOR, + scope={"experiment_id": 42, "variants": ["control", "test"]}, + counts={"control": 60.0, "test": 40.0}, + ) + assert plan is None From 6b49f8c6e8a8148fba9c49e3732f9aecba8aa89a Mon Sep 17 00:00:00 2001 From: Kim Svatos Dugan <147102038+ksvat@users.noreply.github.com> Date: Wed, 30 Sep 2026 08:51:45 -0700 Subject: [PATCH 5/5] fix(replay-vision): gate exposure counts on the principal; unshadow query-kind literal The counts query now runs the experiment's object-level access check as the scanner's creator (backfills: the launcher), the same principal the candidate query authorizes, failing open to plain sampling when denied or missing. Also swaps the malformed-query sweep test's payload off a real product query kind: the activity's module now reaches a HogQL dispatcher through the sampling plan, so the repo's model-crossing guard read the old TrendsQuery literal as this test driving another product's runner and cancelled Backend CI. Co-Authored-By: Claude Fable 5 Generated-By: PostHog Desktop Task-Id: c2db2a38-d03f-4578-af4a-40f61dc700a4 --- .../backend/queries/variant_sampling.py | 16 ++++++++++++++-- .../backend/temporal/activities/backfill.py | 1 + .../activities/find_scanner_candidates.py | 1 + .../tests/test_scanner_candidate_query.py | 10 +++++++++- .../replay_vision/backend/tests/test_sweep.py | 4 +++- .../backend/tests/test_variant_sampling.py | 1 + 6 files changed, 29 insertions(+), 4 deletions(-) diff --git a/products/replay_vision/backend/queries/variant_sampling.py b/products/replay_vision/backend/queries/variant_sampling.py index a61608c7ee5f..80149908af99 100644 --- a/products/replay_vision/backend/queries/variant_sampling.py +++ b/products/replay_vision/backend/queries/variant_sampling.py @@ -27,7 +27,9 @@ from posthog.clickhouse.query_tagging import Feature, Product, tags_context from posthog.dataclasses import frozen from posthog.models.team import Team +from posthog.models.user import User +from products.access_control.backend.facade.user_access_control import UserAccessControlError from products.replay_vision.backend.models.replay_scanner import ScannerType logger = structlog.get_logger(__name__) @@ -107,6 +109,7 @@ def variant_sampling_plan_for_scope( scope: dict | None, scanner_config: dict | None, sampling_rate: float, + user: User | None, scanner_id: str | None = None, ) -> VariantSamplingPlan | None: """The plan for an experiment scanner (or its frozen snapshot), from the window's exposure counts. @@ -128,14 +131,16 @@ def variant_sampling_plan_for_scope( # A legacy column scope narrows with the singular `variant`; treating it as "every variant" # would balance a population the exposure join has already narrowed to one arm. selected = scope.get("variants") or ([scope["variant"]] if scope.get("variant") else None) - counts = _variant_exposure_counts(team, experiment_id=experiment_id, selected=selected, scanner_id=scanner_id) + counts = _variant_exposure_counts( + team, experiment_id=experiment_id, selected=selected, user=user, scanner_id=scanner_id + ) if counts is None: return None return plan_variant_sampling(sampling_rate, counts, selected) def _variant_exposure_counts( - team: Team, *, experiment_id: int, selected: Sequence[str] | None, scanner_id: str | None + team: Team, *, experiment_id: int, selected: Sequence[str] | None, user: User | None, scanner_id: str | None ) -> dict[str, float] | None: """Exposed persons per watched variant over the experiment window, or None when uncountable. @@ -149,8 +154,15 @@ def _variant_exposure_counts( from products.experiments.backend.facade.replay import ( # noqa: PLC0415 exposed_persons_select, resolve_exposure_linkage, + validate_experiment_exposure_access, ) + try: + # The same object-level gate every other exposure read runs, as the same principal the + # candidate query authorizes; a denied or missing principal costs balance, not candidates. + validate_experiment_exposure_access(team, user, experiment_id) + except UserAccessControlError: + return None try: linkage = resolve_exposure_linkage( team, experiment_id=experiment_id, variants=list(selected) if selected is not None else None diff --git a/products/replay_vision/backend/temporal/activities/backfill.py b/products/replay_vision/backend/temporal/activities/backfill.py index 6a8ee0a08660..94532b10e779 100644 --- a/products/replay_vision/backend/temporal/activities/backfill.py +++ b/products/replay_vision/backend/temporal/activities/backfill.py @@ -178,6 +178,7 @@ def find_backfill_candidates_activity(inputs: FindBackfillCandidatesInputs) -> F scope=snapshot.experiment_scope(), scanner_config=snapshot.scanner_config, sampling_rate=snapshot.sampling_rate, + user=backfill.created_by, scanner_id=str(backfill.scanner_id), ) candidate_query = WindowedCandidateQuery( diff --git a/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py b/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py index 9be968a0c7b1..43db4c11f397 100644 --- a/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py +++ b/products/replay_vision/backend/temporal/activities/find_scanner_candidates.py @@ -155,6 +155,7 @@ def find_scanner_candidates_activity(inputs: FindScannerCandidatesInputs) -> Fin scope=scanner.experiment_scope(), scanner_config=scanner.scanner_config, sampling_rate=scanner.sampling_rate, + user=scanner.created_by, scanner_id=str(scanner.id), ) variant_rates = variant_plan.rates if variant_plan is not None else None diff --git a/products/replay_vision/backend/tests/test_scanner_candidate_query.py b/products/replay_vision/backend/tests/test_scanner_candidate_query.py index 9412c525ae4a..6f8b76e03ad4 100644 --- a/products/replay_vision/backend/tests/test_scanner_candidate_query.py +++ b/products/replay_vision/backend/tests/test_scanner_candidate_query.py @@ -930,9 +930,17 @@ def test_variant_exposure_counts_zero_fill_watched_variants(self, team) -> None: self._exposed_session(team, "counts-test", "test", "counts-session-c") flush_persons_and_events() - counts = _variant_exposure_counts(team, experiment_id=experiment.id, selected=None, scanner_id="scanner-1") + counts = _variant_exposure_counts( + team, experiment_id=experiment.id, selected=None, user=creator, scanner_id="scanner-1" + ) assert counts == {"control": 2.0, "test": 1.0, "beta": 0.0} + # Exposure counts are experiment data; without a principal to authorize they stay uncounted + # and the tick falls back to plain sampling. + assert ( + _variant_exposure_counts(team, experiment_id=experiment.id, selected=None, user=None, scanner_id=None) + is None + ) @pytest.mark.django_db def test_per_variant_rates_gate_candidates_by_attributed_variant(self, team) -> None: diff --git a/products/replay_vision/backend/tests/test_sweep.py b/products/replay_vision/backend/tests/test_sweep.py index ebe8a00a6299..bc1bd4904124 100644 --- a/products/replay_vision/backend/tests/test_sweep.py +++ b/products/replay_vision/backend/tests/test_sweep.py @@ -854,7 +854,9 @@ def test_exclusion_failure_fails_the_tick_rather_than_dispatching(self) -> None: def test_raises_non_retryable_on_malformed_query(self) -> None: scanner = _make_scanner() - scanner.query = {"kind": "TrendsQuery"} + # A payload RecordingsQuery validation rejects; a real (non-recordings) query kind would + # read as this test driving that product's query runner. + scanner.query = {"kind": "RecordingsQuery", "date_from": 123} scanner.save(update_fields=["query"]) with pytest.raises(ApplicationError) as exc_info: diff --git a/products/replay_vision/backend/tests/test_variant_sampling.py b/products/replay_vision/backend/tests/test_variant_sampling.py index c9d88d9348c3..604317fcacfd 100644 --- a/products/replay_vision/backend/tests/test_variant_sampling.py +++ b/products/replay_vision/backend/tests/test_variant_sampling.py @@ -131,6 +131,7 @@ def _plan(self, *, scanner_type=ScannerType.EXPERIMENT, scope, counts): scope=scope, scanner_config={"prompt": "p"}, sampling_rate=0.1, + user=self.user, ) def test_a_singular_legacy_variant_scope_watches_one_arm_and_gets_no_plan(self) -> None: