Skip to content
Draft
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
131 changes: 77 additions & 54 deletions api/experimentation/results_query.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
"""
Per-metric slot builder for the experimentation results ClickHouse query.
Per-metric slot builder for the experimentation results query.

Each _MetricSlot owns one alias (``m{i}``) and derives all three things that
must agree on it:
Expand All @@ -22,36 +22,31 @@
from experimentation.dataclasses import ConversionBucket, MetricSpec
from experimentation.models import MetricAggregation
from experimentation.stats import VariantStats
from experimentation.types import ExposureGranularity
from experimentation.warehouses.dialect import Dialect

_FLOAT_VALUE = "toFloat64OrZero(m.value)"

# Events are delivered at-least-once, so dedup keeps duplicates from inflating
# counts. Shared by the exposures, results and conversions queries.
_EXPOSURES_CTE = """
def exposures_cte(dialect: Dialect) -> str:
"""Events are delivered at-least-once, so dedup keeps duplicates from
inflating counts. Shared by the exposures, results and conversions queries."""
p = dialect.param
distinct_values = dialect.count_distinct("value")
return f"""
WITH exposures AS (
SELECT
identifier,
if(uniqExact(value) > 1, '', any(value)) AS variant,
uniqExact(value) > 1 AS quarantined,
if({distinct_values} > 1, '', {dialect.any_value("value")}) AS variant,
{dialect.bool_to_number(f"{distinct_values} > 1")} AS quarantined,
min(timestamp) AS first_exposure
FROM events
WHERE environment_key = %(environment_key)s
AND event = %(exposure_event)s
AND feature_name = %(feature_name)s
AND timestamp >= %(window_start)s
AND timestamp < %(window_end)s
WHERE environment_key = {p("environment_key")}
AND event = {p("exposure_event")}
AND feature_name = {p("feature_name")}
AND timestamp >= {p("window_start")}
AND timestamp < {p("window_end")}
GROUP BY identifier
)"""

_EXPOSURES_COUNT_ONLY_QUERY = (
_EXPOSURES_CTE
+ """
SELECT variant, count() AS n
FROM exposures
WHERE quarantined = 0
GROUP BY variant"""
)


def exposure_window_params(
*,
Expand All @@ -60,7 +55,7 @@ def exposure_window_params(
window_start: datetime,
window_end: datetime,
) -> dict[str, object]:
"""The parameters ``_EXPOSURES_CTE`` binds, for every query that starts
"""The parameters ``exposures_cte`` binds, for every query that starts
from it."""
return {
"environment_key": environment_key,
Expand All @@ -71,16 +66,17 @@ def exposure_window_params(
}


def _metric_join(events_param: str) -> str:
def _metric_join(dialect: Dialect, events_param: str, events: Sequence[str]) -> str:
"""Join each exposed identity to its metric events. ``events_param`` names
the bound list of event names, so a query can join only the events it
aggregates."""
p = dialect.param
return f""" LEFT JOIN events AS m
ON m.identifier = e.identifier
AND m.environment_key = %(environment_key)s
AND m.event IN %({events_param})s
AND m.timestamp >= %(window_start)s
AND m.timestamp < %(window_end)s"""
AND m.environment_key = {p("environment_key")}
AND {dialect.in_list("m.event", events_param, len(events))}
AND m.timestamp >= {p("window_start")}
AND m.timestamp < {p("window_end")}"""


@dataclass(frozen=True)
Expand All @@ -93,6 +89,7 @@ class _MetricSlot:

spec: MetricSpec
index: int
dialect: Dialect

@property
def _alias(self) -> str:
Expand All @@ -103,25 +100,25 @@ def _condition(self) -> str:
# ClickHouse 24.8 rejects ON clauses mixing left+right columns in an
# inequality (error 403).
return (
f"m.event = %(metric_{self.index}_event)s"
f"m.event = {self.dialect.param(f'metric_{self.index}_event')}"
f" AND m.timestamp >= e.first_exposure"
)

def unit_select(self) -> str:
"""Per-identity expression for the unit_values CTE SELECT."""
d = self.dialect
cond = self._condition()
count = d.count_if(cond)
value = d.float_or_zero("m.value")
agg = self.spec.aggregation
if agg == MetricAggregation.OCCURRENCE:
return f"countIf({cond}) > 0 AS {self._alias}"
return f"{d.bool_to_number(f'{count} > 0')} AS {self._alias}"
if agg == MetricAggregation.COUNT:
return f"countIf({cond}) AS {self._alias}"
return f"{count} AS {self._alias}"
if agg == MetricAggregation.SUM:
return f"sumIf({_FLOAT_VALUE}, {cond}) AS {self._alias}"
return f"{d.sum_if(value, cond)} AS {self._alias}"
if agg == MetricAggregation.MEAN:
return (
f"if(countIf({cond}) > 0, avgIf({_FLOAT_VALUE}, {cond}), 0)"
f" AS {self._alias}"
)
return f"{d.avg_if(value, cond)} AS {self._alias}"
raise ValueError(f"Unsupported metric aggregation: {agg}")

def outer_select(self) -> str:
Expand All @@ -143,7 +140,8 @@ def first_conversion_select(self) -> str | None:
if self.spec.aggregation != MetricAggregation.OCCURRENCE:
return None
return (
f"minIfOrNull(m.timestamp, {self._condition()}) AS {self.conversion_alias}"
f"{self.dialect.min_if('m.timestamp', self._condition())}"
f" AS {self.conversion_alias}"
)

def decode(self, n: int, row: Sequence[Any], index: dict[str, int]) -> VariantStats:
Expand All @@ -156,37 +154,48 @@ def decode(self, n: int, row: Sequence[Any], index: dict[str, int]) -> VariantSt


class ResultsQueryBuilder:
"""Assembles and decodes the experimentation results ClickHouse query."""
"""Assembles and decodes the experimentation results query."""

def __init__(self, specs: Sequence[MetricSpec]) -> None:
self._slots = [_MetricSlot(spec, i) for i, spec in enumerate(specs)]
def __init__(self, specs: Sequence[MetricSpec], dialect: Dialect) -> None:
self._dialect = dialect
self._slots = [_MetricSlot(spec, i, dialect) for i, spec in enumerate(specs)]

def build_query(self) -> str:
count_all = self._dialect.count_all()
if not self._slots:
return _EXPOSURES_COUNT_ONLY_QUERY
return (
exposures_cte(self._dialect)
+ f"""
SELECT variant, {count_all} AS n
FROM exposures
WHERE quarantined = 0
GROUP BY variant"""
)

unit_selects = ",\n ".join(s.unit_select() for s in self._slots)
outer_selects = ",\n ".join(s.outer_select() for s in self._slots)

return (
_EXPOSURES_CTE
exposures_cte(self._dialect)
+ f""",
unit_values AS (
SELECT
e.variant AS variant,
{unit_selects}
FROM exposures AS e
{_metric_join("metric_events")}
{_metric_join(self._dialect, "metric_events", self._metric_events)}
WHERE e.quarantined = 0
GROUP BY e.identifier, e.variant
)
SELECT variant, count() AS n,
SELECT variant, {count_all} AS n,
{outer_selects}
FROM unit_values
GROUP BY variant"""
)

def build_conversions_query(self, *, bucket_function: str) -> str | None:
def build_conversions_query(
self, *, granularity: ExposureGranularity
) -> str | None:
"""Per variant and charted metric, how many identities first converted
in each time bucket. None when no attached metric charts, since there
is nothing to query."""
Expand All @@ -197,28 +206,32 @@ def build_conversions_query(self, *, bucket_function: str) -> str | None:
first_conversion_selects = ",\n ".join(
select for s in slots if (select := s.first_conversion_select())
)
indexes = ", ".join(str(s.index) for s in slots)
aliases = ", ".join(s.conversion_alias for s in slots)
unnest = self._dialect.zip_unnest(
{
"metric_index": [str(s.index) for s in slots],
"first_conversion": [s.conversion_alias for s in slots],
}
)

return (
_EXPOSURES_CTE
exposures_cte(self._dialect)
+ f""",
first_conversions AS (
SELECT
e.variant AS variant,
{first_conversion_selects}
FROM exposures AS e
{_metric_join("conversion_events")}
{_metric_join(self._dialect, "conversion_events", self._conversion_events)}
WHERE e.quarantined = 0
GROUP BY e.identifier, e.variant
)
SELECT
variant,
metric_index,
{bucket_function}(first_conversion, 'UTC') AS bucket,
count() AS converted_identities
{self._dialect.time_bucket("first_conversion", granularity)} AS bucket,
{self._dialect.count_all()} AS converted_identities
FROM first_conversions
ARRAY JOIN [{indexes}] AS metric_index, [{aliases}] AS first_conversion
{unnest}
WHERE first_conversion IS NOT NULL
GROUP BY variant, metric_index, bucket
ORDER BY bucket"""
Expand All @@ -228,6 +241,14 @@ def build_conversions_query(self, *, bucket_function: str) -> str | None:
def _charted_slots(self) -> list[_MetricSlot]:
return [s for s in self._slots if s.first_conversion_select() is not None]

@property
def _metric_events(self) -> list[str]:
return [s.spec.event for s in self._slots]

@property
def _conversion_events(self) -> list[str]:
return [s.spec.event for s in self._charted_slots]

def params(
self,
*,
Expand All @@ -247,8 +268,10 @@ def params(
)
if not self._slots:
return params
params["metric_events"] = [s.spec.event for s in self._slots]
params["conversion_events"] = [s.spec.event for s in self._charted_slots]
params |= self._dialect.list_params("metric_events", self._metric_events)
params |= self._dialect.list_params(
"conversion_events", self._conversion_events
)
for slot in self._slots:
params[f"metric_{slot.index}_event"] = slot.spec.event
return params
Expand Down Expand Up @@ -276,7 +299,7 @@ def decode_conversion_rows(
def decode_rows(
self, rows: list[Any], column_names: Sequence[str]
) -> tuple[dict[str, int], dict[int, dict[str, VariantStats]]]:
"""Decode raw ClickHouse rows into exposure counts and per-metric stats.
"""Decode raw rows into exposure counts and per-metric stats.

Columns are located by name, so a missing one raises KeyError rather than
silently reading a neighbour's value.
Expand Down
Loading
Loading