diff --git a/docs/impulse/docs/references/api/impulse_reporting/events/time_window_event.md b/docs/impulse/docs/references/api/impulse_reporting/events/time_window_event.md index 07da0708..4cb166d4 100644 --- a/docs/impulse/docs/references/api/impulse_reporting/events/time_window_event.md +++ b/docs/impulse/docs/references/api/impulse_reporting/events/time_window_event.md @@ -19,7 +19,12 @@ event instance per fixed-duration slice, tiling the container's ``start_ts`` / ` span with windows of length ``window_length``. The final slice is clamped to the container end. -The event fact is computed from ``container_metrics`` alone (via +The event fact is computed from ``container_metrics`` alone (in a ``mapInArrow``), so +every filtered container gets windows regardless of its channel data. Aggregations +scoped to this event evaluate the :class:`TimeWindowExpression` in the solve. Both use +the same window function (``tile_windows``), so they produce identical windows, and the +timestamp-based ``event_instance_id`` (like for other interval events) matches on both +sides. #### \_\_init\_\_ @@ -143,7 +148,7 @@ containers' ``start_ts`` / ``stop_ts`` in the channel time frame gets windows. Each window becomes one event instance (``start_ts < end_ts``) whose ``event_instance_id`` hashes its boundaries. The solve uses the same window function -for scoped aggregations (see :func:`window_intervals_udf`), so the ids match. +for scoped aggregations (see :func:`tile_windows`), so the ids match. **Arguments**: diff --git a/src/impulse_query_engine/analyze/query/events/time_window_expression.py b/src/impulse_query_engine/analyze/query/events/time_window_expression.py index 3c7072bd..81e90d27 100644 --- a/src/impulse_query_engine/analyze/query/events/time_window_expression.py +++ b/src/impulse_query_engine/analyze/query/events/time_window_expression.py @@ -5,7 +5,6 @@ import numpy as np import pandas as pd -import pyspark.sql.functions as F from impulse_query_engine.analyze.metadata.tag_expression import TagExpression from impulse_query_engine.analyze.metadata.time_series_expression import ( @@ -73,13 +72,14 @@ def tile_windows( The one window implementation behind ``TimeWindowEvent``: the solve calls it through :meth:`TimeWindowExpression.build` (scoped aggregations), the event fact through - :func:`window_intervals_udf`. ``event_instance_id`` hashes each window's boundaries, so + ``TimeWindowEvent.determine_events``. ``event_instance_id`` hashes each window's boundaries, so both sides must produce identical windows, which a single function guarantees as long as both pass in the same values. Both read the same Spark-computed bounds - (``solvers.utils.window_bounds.with_window_bounds``), but pandas hands them over as - ``int64`` or ``float64`` (nulls force ``float64``), or as ``None`` / ``NaN``. The bounds are - therefore converted to ``float`` first: ``int64`` -> ``float64`` rounds to the nearest - double on either path, so the arithmetic below runs on identical doubles. + (``solvers.utils.window_bounds.with_window_bounds``), but receive them as Python or numpy + integers or floats (pandas turns long columns with nulls into ``float64``), or as + ``None`` / ``NaN``. The bounds are therefore converted to ``float`` first: integer -> + double rounds to the nearest double on either path, so the arithmetic below runs on + identical doubles. Window ``i`` spans ``[start + i * W, min(start + (i + 1) * W, stop)]``, so the last one is clamped to *stop*; windows with ``start_i >= end_i`` (possible only through rounding) @@ -128,48 +128,6 @@ def tile_windows( return starts[keep], ends[keep] -def window_intervals_udf(window_length: float, max_windows: int = MAX_WINDOWS_PER_CONTAINER): - """Scalar pandas UDF giving each container's windows via :func:`tile_windows`. - - Used by the reporting ``TimeWindowEvent`` for its event fact (one row per container), so - the event fact and the solve share one window implementation. - - Parameters - ---------- - window_length : float - Fixed window length, in the same unit as the bounds. Strictly positive. - max_windows : int, optional - Maximum number of windows per container (default - :data:`MAX_WINDOWS_PER_CONTAINER`). A container exceeding it fails the query with - an error naming the limit. - - Returns - ------- - callable - A pandas UDF ``(start, stop) -> struct, ends: array>`` - with the window starts and ends, in order; empty when a bound is null, NaN or - infinite, or the span is not strictly positive. - """ - window_length = float(window_length) - max_windows = validate_max_windows(max_windows) - - @F.pandas_udf("struct, ends: array>") - def windows(start: pd.Series, stop: pd.Series) -> pd.DataFrame: - pairs = [ - tile_windows(s, e, window_length, max_windows) - for s, e in zip(start, stop, strict=True) - ] - # object dtype keeps one array per row, also when all rows have equal window counts. - return pd.DataFrame( - { - "starts": pd.Series([p[0] for p in pairs], dtype=object), - "ends": pd.Series([p[1] for p in pairs], dtype=object), - } - ) - - return windows - - class TimeWindowExpression(TimeSeriesExpression): """Produce consecutive fixed-duration windows spanning a measurement container. @@ -189,8 +147,8 @@ class TimeWindowExpression(TimeSeriesExpression): This is the query-engine counterpart of the reporting ``TimeWindowEvent``. It evaluates to :class:`Intervals`, so it can scope a ``StatsAggregator`` (one statistic per window). - The windows come from :func:`tile_windows`, which the reporting event fact also uses (via - :func:`window_intervals_udf`), so both produce the same windows in the same order. + The windows come from :func:`tile_windows`, which the reporting event fact also uses (in + ``TimeWindowEvent.determine_events``), so both produce the same windows in the same order. Attributes ---------- diff --git a/src/impulse_reporting/events/time_window_event.py b/src/impulse_reporting/events/time_window_event.py index 38367081..155c3b56 100644 --- a/src/impulse_reporting/events/time_window_event.py +++ b/src/impulse_reporting/events/time_window_event.py @@ -4,7 +4,9 @@ from collections.abc import Mapping -import pyspark.sql.functions as f +import numpy as np +import pyarrow as pa +import pyspark.sql.types as T from pyspark.sql import DataFrame, SparkSession from impulse_query_engine.analyze.metadata.time_series_expression import ( @@ -13,8 +15,8 @@ from impulse_query_engine.analyze.query.events.time_window_expression import ( MAX_WINDOWS_PER_CONTAINER, TimeWindowExpression, + tile_windows, validate_max_windows, - window_intervals_udf, ) from impulse_query_engine.analyze.query.query_builder import QueryBuilder from impulse_query_engine.analyze.query.solvers.query_solver import QuerySolver @@ -24,6 +26,15 @@ from impulse_reporting.util.event_instance_util import generate_event_instance_id_column from impulse_reporting.util.report_entity_util import ReportEntityUtil +# Columns of the window rows next to container_id, named as the event fact and its +# event_instance_id hash expect them. +_EVENT_NAME_COL = "event_name" +_START_TS_COL = "start_ts" +_END_TS_COL = "end_ts" + +# Windows after which _explode_windows emits an output batch. +_BATCH_WINDOWS = 100_000 + class TimeWindowEvent(ContainerBoundaryEvent): """Event that divides each measurement container into consecutive fixed windows. @@ -33,12 +44,12 @@ class TimeWindowEvent(ContainerBoundaryEvent): span with windows of length ``window_length``. The final slice is clamped to the container end. - The event fact is computed from ``container_metrics`` alone (via - :func:`window_intervals_udf`), so every filtered container gets windows regardless of - its channel data. Aggregations scoped to this event evaluate the - :class:`TimeWindowExpression` in the solve. Both use the same window function - (``tile_windows``), so they produce identical windows, and the timestamp-based - ``event_instance_id`` (like for other interval events) matches on both sides. + The event fact is computed from ``container_metrics`` alone (in a ``mapInArrow``), so + every filtered container gets windows regardless of its channel data. Aggregations + scoped to this event evaluate the :class:`TimeWindowExpression` in the solve. Both use + the same window function (``tile_windows``), so they produce identical windows, and the + timestamp-based ``event_instance_id`` (like for other interval events) matches on both + sides. """ def __init__( @@ -184,7 +195,7 @@ def determine_events( gets windows. Each window becomes one event instance (``start_ts < end_ts``) whose ``event_instance_id`` hashes its boundaries. The solve uses the same window function - for scoped aggregations (see :func:`window_intervals_udf`), so the ids match. + for scoped aggregations (see :func:`tile_windows`), so the ids match. Parameters ---------- @@ -213,46 +224,127 @@ def determine_events( # uses for scoped aggregations (fails fast on the schema, e.g. when TIMESTAMP # boundaries lack solver_config.channel_time_unit). container_metrics_df = with_window_bounds(container_metrics_df, solver.config) - start_ts = f.col(solver.config.window_start_col) - stop_ts = f.col(solver.config.window_stop_col) - - # One (event_name, windows) struct per event, exploded in a single pass over the - # containers. - per_event = f.array( - *[ - f.struct( - f.lit(event.get_name()).alias("event_name"), - window_intervals_udf( - event.window_length, max_windows=event.max_windows_per_container - )(start_ts, stop_ts).alias("windows"), - ) + windows_df = _explode_windows( + container_metrics_df, + id_col=solver.config.container_id_col, + start_col=solver.config.window_start_col, + stop_col=solver.config.window_stop_col, + windows=[ + (event.get_name(), event.window_length, event.max_windows_per_container) for event in events - ] + ], ) - - df = ( - container_metrics_df.select( - f.col(solver.config.container_id_col).alias("container_id"), - f.explode(per_event).alias("event"), - ) - .select( - "container_id", - f.col("event.event_name").alias("event_name"), - f.inline( - f.arrays_zip( - f.col("event.windows.starts").alias("start_ts"), - f.col("event.windows.ends").alias("end_ts"), - ) - ), - ) - .withColumn( + return ( + windows_df.withColumn( "event_instance_id", - generate_event_instance_id_column(event_type=TimeWindowEvent), + generate_event_instance_id_column( + event_type=TimeWindowEvent, + event_name_col=_EVENT_NAME_COL, + start_ts_col=_START_TS_COL, + end_ts_col=_END_TS_COL, + ), ) .withColumn( "event_id", - ReportEntityUtil.get_event_id_column(elements=events, element_name="event_name"), + ReportEntityUtil.get_event_id_column( + elements=events, element_name=_EVENT_NAME_COL + ), ) .select(EVENT_INSTANCE_FACT_SCHEMA.fieldNames()) ) - return df + + +def _explode_windows( + df: DataFrame, + *, + id_col: str, + start_col: str, + stop_col: str, + windows: list[tuple[str, float, int]], +) -> DataFrame: + """One row per window of each container, for all events in one pass. + + Each container's windows come from :func:`tile_windows`, like those of the solve. The + rows are built from numpy arrays in a ``mapInArrow``, in output batches of about 100,000 + windows that never split a container's windows of one event, so the Python worker's + memory stays bounded by the batch size and ``max_windows``, independent of the number of + containers and events. + + Parameters + ---------- + df : pyspark.sql.DataFrame + One row per container, with *id_col* and the bounds in the channel time frame. + id_col : str + Container id column; kept with its name and type (e.g. long or string) in the output. + start_col, stop_col : str + Container bound columns (numeric). + windows : list of tuple + ``(name, window_length, max_windows)`` per event: the event name of its rows, the + window length (in the unit of the bounds) and the maximum number of windows per + container. A container exceeding it fails the query with an error naming the limit. + + Returns + ------- + pyspark.sql.DataFrame + Columns *id_col*, ``event_name`` (string), ``start_ts`` and ``end_ts`` (double); no + rows for a container whose bound is null, NaN or infinite, or whose span is not + strictly positive. + """ + schema = T.StructType( + [ + df.schema[id_col], + T.StructField(_EVENT_NAME_COL, T.StringType()), + T.StructField(_START_TS_COL, T.DoubleType()), + T.StructField(_END_TS_COL, T.DoubleType()), + ] + ) + column_names = schema.names + + def tile(batches): + yield from _window_batches( + batches, id_col, start_col, stop_col, windows, column_names, _BATCH_WINDOWS + ) + + return df.select(id_col, start_col, stop_col).mapInArrow(tile, schema) + + +def _window_batches(batches, id_col, start_col, stop_col, windows, column_names, batch_windows): + """Yield :func:`_explode_windows` output batches named *column_names* for Arrow *batches*. + + A batch is emitted once it holds at least *batch_windows* windows, after a container's + complete windows of one event, and at the end of each input batch. A batch therefore + holds at most ``batch_windows - 1`` plus one event's ``max_windows`` windows. + """ + event_names = pa.array([name for name, _, _ in windows], pa.string()) + for batch in batches: + ids = batch.column(id_col) + bounds = zip( + batch.column(start_col).to_pylist(), batch.column(stop_col).to_pylist(), strict=True + ) + parts, pending = [], 0 + for row, (start, stop) in enumerate(bounds): + for event, (_, length, max_windows) in enumerate(windows): + starts, ends = tile_windows(start, stop, length, max_windows) + parts.append((row, event, starts, ends)) + pending += len(starts) + if pending >= batch_windows: + yield _record_batch(column_names, ids, event_names, parts) + parts, pending = [], 0 + if pending: + yield _record_batch(column_names, ids, event_names, parts) + + +def _record_batch(column_names, ids, event_names, parts) -> pa.RecordBatch: + """Output batch for the ``(row, event, starts, ends)`` *parts* of one input batch.""" + counts = [len(starts) for _, _, starts, _ in parts] + rows = np.repeat([row for row, _, _, _ in parts], counts) + events = np.repeat([event for _, event, _, _ in parts], counts) + return pa.RecordBatch.from_arrays( + [ + ids.take(rows), + event_names.take(events), + pa.array(np.concatenate([starts for _, _, starts, _ in parts])), + pa.array(np.concatenate([ends for _, _, _, ends in parts])), + ], + names=column_names, + ) diff --git a/tests/impulse_query_engine/unit/analyze/query/events/time_window_expression_test.py b/tests/impulse_query_engine/unit/analyze/query/events/time_window_expression_test.py index 3e946121..642e7a0f 100644 --- a/tests/impulse_query_engine/unit/analyze/query/events/time_window_expression_test.py +++ b/tests/impulse_query_engine/unit/analyze/query/events/time_window_expression_test.py @@ -1,11 +1,9 @@ from __future__ import annotations -import random from unittest.mock import MagicMock import numpy as np import pandas as pd -import pyspark.sql.functions as F import pytest from impulse_query_engine.analyze.query.aggregations.stats_aggregator import StatsAggregator @@ -13,7 +11,6 @@ from impulse_query_engine.analyze.query.events.time_window_expression import ( MAX_WINDOWS_PER_CONTAINER, tile_windows, - window_intervals_udf, ) from impulse_query_engine.analyze.query.solvers.empty_cache import EmptyTimeSeriesCache from impulse_query_engine.model.series.intervals import Intervals @@ -204,149 +201,11 @@ def test_tile_windows_none_and_na_yield_empty(): assert len(starts) == len(ends) == 0 -# --------------------------------------------------------------------------- -# window_intervals_udf: tile_windows on the event fact side (one row per container) -# --------------------------------------------------------------------------- -def _udf_windows(spark, rows, window_length, ts_type="long", **kwargs): # noqa: F811 - # One partition, so all rows reach the UDF in one Arrow batch. - df = spark.createDataFrame(rows, f"k int, start_ts {ts_type}, stop_ts {ts_type}").coalesce(1) - windows = window_intervals_udf(window_length, **kwargs) - out = df.select("k", windows(F.col("start_ts"), F.col("stop_ts")).alias("w")) - return out, { - r.k: [[s, e] for s, e in zip(r.w.starts, r.w.ends, strict=True)] for r in out.collect() - } - - -def test_window_intervals_udf_edge_cases(spark): # noqa: F811 - rows = [ - (0, 0, 100), # exact multiple -> 10 windows - (1, 0, 105), # short final window clamped to stop - (2, 1000, 1010), # span == W -> 1 window - (3, 1000, 1001), # span < W -> 1 clamped window - (4, 50, 50), # stop == start -> none - (5, 60, 50), # stop < start -> none - (6, None, 50), # null bound -> none - (7, 0, None), - ] - out, w = _udf_windows(spark, rows, 10) - - assert out.schema["w"].dataType.simpleString() == ( - "struct,ends:array>" - ) - assert w[0] == [[float(s), float(s + 10)] for s in range(0, 100, 10)] - assert w[1][-1] == [100.0, 105.0] and len(w[1]) == 11 - assert w[2] == [[1000.0, 1010.0]] - assert w[3] == [[1000.0, 1001.0]] - assert w[4] == w[5] == w[6] == w[7] == [] - - -def test_window_intervals_udf_non_finite_bounds_yield_no_windows(spark): # noqa: F811 - nan, inf = float("nan"), float("inf") - rows = [(0, 0.0, nan), (1, nan, 10.0), (2, nan, nan), (3, 0.0, inf), (4, -inf, 10.0)] - _, w = _udf_windows(spark, rows, 4, ts_type="double") - assert all(w[k] == [] for k, _, _ in rows), w - - -def test_window_intervals_udf_raises_beyond_max_windows(spark): # noqa: F811 - _, ok = _udf_windows(spark, [(0, 0, 100)], 10, max_windows=10) - assert len(ok[0]) == 10 - with pytest.raises(Exception, match="10 windows of length 10.0 .* exceed max_windows=9"): - _udf_windows(spark, [(0, 0, 100)], 10, max_windows=9) - - -def test_window_intervals_udf_invalid_max_windows_raises(): - with pytest.raises(ValueError, match="max_windows must be a positive integer"): - window_intervals_udf(10, max_windows=0) - - -def test_window_intervals_udf_uniform_and_empty_batches(spark): # noqa: F811 - # Equal window counts across a batch still give one array per row. - _, w = _udf_windows(spark, [(k, 100 * k, 100 * k + 30) for k in range(4)], 10) - assert w == {k: [[100.0 * k + s, 100.0 * k + s + 10] for s in (0, 10, 20)] for k in range(4)} - # A batch without any windows. - _, w = _udf_windows(spark, [(0, None, None), (1, 50, 50)], 10) - assert w == {0: [], 1: []} - - def _as_list(windows) -> list[tuple[float, float]]: """Windows as ordered (start, end) pairs.""" return [(float(s), float(e)) for s, e in windows] -def _count_mismatch_case(window_length: float) -> tuple[int, int]: - """Find ns-epoch (start, stop) whose int64 and double spans yield different counts. - - This is exactly the case where exact int64 subtraction and double subtraction disagree - on the number of windows, so tile_windows must convert to float first on every path. - """ - base = 1_700_000_000_000_000_000 - for start in range(base, base + 512): - for k in (1, 3, 5, 7): - for d in range(-300, 1): - stop = start + int(k * window_length) + d - exact = int(np.ceil((stop - start) / window_length)) - rounded = int(np.ceil((float(stop) - float(start)) / window_length)) - if exact != rounded: - return start, stop - raise AssertionError("no int64/double count-mismatch case found") - - -def _batch_dtype_udf(): - """Pandas UDF reporting the dtype the bounds arrive in, per row of the batch (created - lazily: defining a pandas UDF needs an active Spark session).""" - - @F.pandas_udf("string") - def batch_dtype(start: pd.Series) -> pd.Series: - return pd.Series([str(start.dtype)] * len(start)) - - return batch_dtype - - -def test_event_fact_and_solve_windows_identical_across_input_dtypes(spark): # noqa: F811 - """Both sides call tile_windows, but pandas hands the bounds over differently: the event - fact UDF gets int64 for a batch without nulls and float64 once a null is in the batch, - the solve gets float64 (its container metrics are nulled on most rows). For ns epochs - beyond 2^53, including a span where int64 and double subtraction disagree on the count, - all paths must produce the same windows in the same order.""" - rnd = random.Random(7) - window_length = 1_000_000_007.0 - cases = [] - for _ in range(60): - start = rnd.randint(1_600_000_000_000_000_000, 1_800_000_000_000_000_000) - cases.append( - (start, start + int(rnd.randint(1, 50) * window_length) + rnd.randint(-600, 600)) - ) - cases.append(_count_mismatch_case(window_length)) - rows = [(k, s, e) for k, (s, e) in enumerate(cases)] - windows = window_intervals_udf(window_length) - batch_dtype = _batch_dtype_udf() - - def _event_fact(batch_rows) -> tuple[dict, set]: - # One partition, so one Arrow batch: a null anywhere turns the whole batch float64. - df = spark.createDataFrame(batch_rows, "k int, start_ts long, stop_ts long").coalesce(1) - out = df.select( - "k", - windows(F.col("start_ts"), F.col("stop_ts")).alias("w"), - batch_dtype(F.col("start_ts")).alias("dtype"), - ).collect() - per_container = { - r.k: _as_list(zip(r.w.starts, r.w.ends, strict=True)) for r in out if r.k >= 0 - } - return per_container, {r.dtype for r in out} - - without_nulls, dtypes_int = _event_fact(rows) - with_null, dtypes_float = _event_fact([*rows, (-1, None, None)]) - assert dtypes_int == {"int64"} and dtypes_float == {"float64"} - - mismatches = [] - for k, (start, stop) in enumerate(cases): - solve = _as_list(_build(np.float64(start), np.float64(stop), window_length).get_data()) - assert solve, f"case {k} produced no windows" - if not (without_nulls[k] == with_null[k] == solve): - mismatches.append((k, start, stop)) - assert not mismatches, f"event fact / solve window mismatch: {mismatches[:5]}" - - def test_stats_aggregator_windows_equal_helper_windows(spark): # noqa: F811 """A StatsAggregator scoped to a TimeWindowExpression emits exactly the helper's windows (no merging of touching windows, no extra drops).""" diff --git a/tests/impulse_reporting/unit/events/time_window_event_test.py b/tests/impulse_reporting/unit/events/time_window_event_test.py index 849b97d1..8c47c2ab 100644 --- a/tests/impulse_reporting/unit/events/time_window_event_test.py +++ b/tests/impulse_reporting/unit/events/time_window_event_test.py @@ -1,14 +1,26 @@ """Unit tests for TimeWindowEvent.""" +import random +from types import SimpleNamespace + +import numpy as np +import pyarrow as pa import pytest from impulse_query_engine.analyze.query.events.time_window_expression import ( MAX_WINDOWS_PER_CONTAINER, TimeWindowExpression, + tile_windows, ) +from impulse_query_engine.analyze.query.solvers.solver_config import SolverConfig from impulse_reporting.events.container_boundary_event import ContainerBoundaryEvent from impulse_reporting.events.container_event import ContainerEvent -from impulse_reporting.events.time_window_event import TimeWindowEvent +from impulse_reporting.events.time_window_event import ( + TimeWindowEvent, + _explode_windows, + _window_batches, +) +from tests.conftest import spark # noqa: F401 (pytest fixture) # --------------------------------------------------------------------------- @@ -155,3 +167,217 @@ def test_as_dict_shape(): assert d["required_channels"] == ["c1"] assert d["event_expression"] != "NA" assert d["attributes"]["window_length"] == "10000.0" + + +# --------------------------------------------------------------------------- +# _explode_windows: tile_windows on the event fact side (one row per window) +# --------------------------------------------------------------------------- +# Output columns of _window_batches for the test frames. +_COLUMNS = ["k", "event_name", "start_ts", "end_ts"] + + +def _exploded(spark, rows, windows, ts_type="long", id_type="int"): # noqa: F811 + """_explode_windows over (k, start_ts, stop_ts) rows, as {(k, event_name): windows}.""" + # One partition, so all rows reach the function in one Arrow batch. + df = spark.createDataFrame( + rows, f"k {id_type}, start_ts {ts_type}, stop_ts {ts_type}" + ).coalesce(1) + out = _explode_windows( + df, id_col="k", start_col="start_ts", stop_col="stop_ts", windows=windows + ) + grouped = {} + for r in out.collect(): + grouped.setdefault((r.k, r.event_name), []).append([r.start_ts, r.end_ts]) + return out, grouped + + +def test_explode_windows_edge_cases(spark): # noqa: F811 + rows = [ + (0, 0, 100), # exact multiple -> 10 windows + (1, 0, 105), # short final window clamped to stop + (2, 1000, 1010), # span == W -> 1 window + (3, 1000, 1001), # span < W -> 1 clamped window + (4, 50, 50), # stop == start -> none + (5, 60, 50), # stop < start -> none + (6, None, 50), # null bound -> none + (7, 0, None), + ] + out, w = _exploded(spark, rows, [("tw", 10.0, MAX_WINDOWS_PER_CONTAINER)]) + + assert out.schema.simpleString() == ( + "struct" + ) + assert w[(0, "tw")] == [[float(s), float(s + 10)] for s in range(0, 100, 10)] + assert w[(1, "tw")][-1] == [100.0, 105.0] and len(w[(1, "tw")]) == 11 + assert w[(2, "tw")] == [[1000.0, 1010.0]] + assert w[(3, "tw")] == [[1000.0, 1001.0]] + assert set(w) == {(k, "tw") for k in range(4)} + + +def test_explode_windows_non_finite_bounds_yield_no_windows(spark): # noqa: F811 + nan, inf = float("nan"), float("inf") + rows = [(0, 0.0, nan), (1, nan, 10.0), (2, nan, nan), (3, 0.0, inf), (4, -inf, 10.0)] + out, _ = _exploded(spark, rows, [("tw", 4.0, MAX_WINDOWS_PER_CONTAINER)], ts_type="double") + assert out.count() == 0 + + +def test_explode_windows_without_any_window_yields_no_rows(spark): # noqa: F811 + rows = [(0, None, None), (1, 50, 50)] + out, _ = _exploded(spark, rows, [("tw", 10.0, MAX_WINDOWS_PER_CONTAINER)]) + assert out.count() == 0 + + +def test_explode_windows_raises_beyond_max_windows(spark): # noqa: F811 + _, ok = _exploded(spark, [(0, 0, 100)], [("tw", 10.0, 10)]) + assert len(ok[(0, "tw")]) == 10 + with pytest.raises(Exception, match="10 windows of length 10.0 .* exceed max_windows=9"): + _exploded(spark, [(0, 0, 100)], [("tw", 10.0, 9)]) + + +def test_explode_windows_several_events_in_one_pass(spark): # noqa: F811 + rows = [(0, 0, 105), (1, 1000, 1030)] + events = [("ten", 10.0, 100), ("seven", 7.5, 100)] + _, w = _exploded(spark, rows, events) + + assert set(w) == {(k, name) for k, _, _ in rows for name, _, _ in events} + for k, start, stop in rows: + for name, length, _ in events: + expected = [[s, e] for s, e in zip(*tile_windows(start, stop, length), strict=True)] + assert w[(k, name)] == expected + + +@pytest.mark.parametrize( + "id_type, ids", [("int", [7, 8]), ("bigint", [2**40, 2**40 + 1]), ("string", ["c-a", "c-b"])] +) +def test_explode_windows_keeps_container_id_type(spark, id_type, ids): # noqa: F811 + rows = [(ids[0], 0, 30), (ids[1], 100, 120)] + out, w = _exploded(spark, rows, [("tw", 10.0, 100)], id_type=id_type) + + assert out.schema["k"].dataType.simpleString() == id_type + assert w == { + (ids[0], "tw"): [[0.0, 10.0], [10.0, 20.0], [20.0, 30.0]], + (ids[1], "tw"): [[100.0, 110.0], [110.0, 120.0]], + } + + +def _bounds_batch(ids, starts, stops) -> pa.RecordBatch: + return pa.RecordBatch.from_arrays( + [pa.array(ids, pa.string()), pa.array(starts, pa.int64()), pa.array(stops, pa.int64())], + names=["k", "start_ts", "stop_ts"], + ) + + +def test_window_batches_flush_between_events_and_per_input_batch(): + # Each container gives 4 "a" windows and 2 "b" windows. + events = [("a", 10.0, 100), ("b", 20.0, 100)] + first = _bounds_batch(["c0", "c1", "c2"], [0, 100, 200], [40, 140, 240]) + second = _bounds_batch(["c3"], [300], [340]) + + def run(batch_windows): + return list( + _window_batches( + iter([first, second]), + "k", + "start_ts", + "stop_ts", + events, + _COLUMNS, + batch_windows, + ) + ) + + batches = run(batch_windows=9) + # 6 windows after c0, 10 after c1's "a" -> flush, before c1's "b"; c1's "b" and c2 + # flushed at the end of the first input batch; c3 alone from the second one. + assert [b.num_rows for b in batches] == [10, 8, 6] + pairs = [ + set(zip(b.column("k").to_pylist(), b.column("event_name").to_pylist(), strict=True)) + for b in batches + ] + assert pairs == [ + {("c0", "a"), ("c0", "b"), ("c1", "a")}, + {("c1", "b"), ("c2", "a"), ("c2", "b")}, + {("c3", "a"), ("c3", "b")}, + ] + assert batches[0].schema.names == ["k", "event_name", "start_ts", "end_ts"] + assert batches[0].schema.field("k").type == pa.string() + assert batches[0].column("event_name").to_pylist()[:6] == ["a"] * 4 + ["b"] * 2 + + # Many events per container stay bounded: each flush adds at most one event's windows. + many = [(f"e{i}", 10.0, 100) for i in range(5)] + sizes = [ + b.num_rows + for b in _window_batches( + iter([first]), "k", "start_ts", "stop_ts", many, _COLUMNS, batch_windows=6 + ) + ] + assert sizes == [8] * 7 + [4] + + # Batching only splits the stream; the rows are those of a single batch per input batch. + unbatched = pa.Table.from_batches(run(batch_windows=10**9)) + assert pa.Table.from_batches(batches).equals(unbatched) + + +def _as_list(windows) -> list[tuple[float, float]]: + """Windows as ordered (start, end) pairs.""" + return [(float(s), float(e)) for s, e in windows] + + +def _count_mismatch_case(window_length: float) -> tuple[int, int]: + """Find ns-epoch (start, stop) whose int64 and double spans yield different counts. + + This is exactly the case where exact int64 subtraction and double subtraction disagree + on the number of windows, so tile_windows must convert to float first on every path. + """ + base = 1_700_000_000_000_000_000 + for start in range(base, base + 512): + for k in (1, 3, 5, 7): + for d in range(-300, 1): + stop = start + int(k * window_length) + d + exact = int(np.ceil((stop - start) / window_length)) + rounded = int(np.ceil((float(stop) - float(start)) / window_length)) + if exact != rounded: + return start, stop + raise AssertionError("no int64/double count-mismatch case found") + + +def _solve_windows(start, stop, window_length: float) -> list[tuple[float, float]]: + """Windows of the solve: TimeWindowExpression.build on float64 container metrics.""" + config = SolverConfig() + cache = SimpleNamespace( + container_metrics={ + config.window_start_col: np.float64(start), + config.window_stop_col: np.float64(stop), + } + ) + return _as_list(TimeWindowExpression(window_length).build(cache).get_data()) + + +def test_event_fact_and_solve_windows_identical_across_input_dtypes(spark): # noqa: F811 + """Both sides call tile_windows, but receive the bounds differently: the event fact as + Python ints from the Arrow long columns, the solve as float64 (its container metrics are + nulled on most rows). For ns epochs beyond 2^53, including a span where int64 and double + subtraction disagree on the count, both must produce the same windows in the same order, + also with a null bound in the same Arrow batch.""" + rnd = random.Random(7) + window_length = 1_000_000_007.0 + cases = [] + for _ in range(60): + start = rnd.randint(1_600_000_000_000_000_000, 1_800_000_000_000_000_000) + cases.append( + (start, start + int(rnd.randint(1, 50) * window_length) + rnd.randint(-600, 600)) + ) + cases.append(_count_mismatch_case(window_length)) + rows = [(k, s, e) for k, (s, e) in enumerate(cases)] + windows = [("tw", window_length, MAX_WINDOWS_PER_CONTAINER)] + + _, without_nulls = _exploded(spark, rows, windows) + _, with_null = _exploded(spark, [*rows, (-1, None, None)], windows) + + mismatches = [] + for k, (start, stop) in enumerate(cases): + solve = _solve_windows(start, stop, window_length) + assert solve, f"case {k} produced no windows" + if not (_as_list(without_nulls[(k, "tw")]) == _as_list(with_null[(k, "tw")]) == solve): + mismatches.append((k, start, stop)) + assert not mismatches, f"event fact / solve window mismatch: {mismatches[:5]}"