Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -159,9 +159,14 @@
# 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
Expand Down Expand Up @@ -415,7 +420,10 @@
"""
# 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
Expand All @@ -441,13 +449,23 @@
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(
Expand Down Expand Up @@ -610,7 +628,7 @@
)
return builders

def _where_predicates(self) -> Union[ast.And, ast.Or]:

Check warning on line 631 in posthog/session_recordings/queries/session_recording_list_from_query.py

View workflow job for this annotation

GitHub Actions / Python code quality (depot-ubuntu-24.04)

lint:complexity

`_where_predicates` has cyclomatic complexity 23 (warn >10)

Check warning on line 631 in posthog/session_recordings/queries/session_recording_list_from_query.py

View workflow job for this annotation

GitHub Actions / Python code quality (depot-ubuntu-24.04)

`_where_predicates` has cyclomatic complexity 23 (warn >10)
exprs: list[ast.Expr] = []

# When both distinct_ids and person_uuid are provided (e.g. person profile
Expand Down
2 changes: 2 additions & 0 deletions products/experiments/backend/facade/replay.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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",
Expand Down
10 changes: 10 additions & 0 deletions products/replay_vision/backend/api/observations.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
80 changes: 63 additions & 17 deletions products/replay_vision/backend/queries/scanner_candidate_query.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__}")
Expand All @@ -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
Expand Down Expand Up @@ -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]:
Expand Down Expand Up @@ -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)


Expand Down Expand Up @@ -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:
Expand All @@ -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.

Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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)
Expand All @@ -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]:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -319,6 +319,9 @@ def refresh_scanner_estimate(
budget=budget,
ch_user=ch_user,
)
# 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.
Expand Down
Loading
Loading