diff --git a/docs/impulse/docs/config/configuration.md b/docs/impulse/docs/config/configuration.md index 415901ac..87d8debc 100644 --- a/docs/impulse/docs/config/configuration.md +++ b/docs/impulse/docs/config/configuration.md @@ -140,6 +140,7 @@ Two independent filter families: | `raw_encoder` | `str` | `null` (resolves to `"RLE"` when `data_type = "RAW"`) | How RAW point data is converted to intervals. `"RLE"` (the default) run-length encodes on the fly to reduce memory consumption: consecutive samples with the same value collapse into a single `[tstart, tend)` interval per run. `"INTERVAL"` keeps every original sample, only deriving `tend` from the next sample's timestamp and dropping exact duplicate points — use it when downstream analysis needs all original timestamps. Only takes effect with `data_type = "RAW"`; ignored for `"RLE"` input. See [How Impulse interprets intervals](../data_model/silver_layer_schema.md#raw-format) for validity semantics and what counts as a duplicate point. | | `drop_implausible_data` | `bool` | `false` | When `true`, drops `channels` rows where `is_plausible = false`. Requires `data_type = "RAW"`; combining with `"RLE"` raises a validation error. | | `batch_size` | `int` | `500` | Maximum number of selectors solved per batch. Applies to events, aggregations, and calculated channels. | +| `max_containers_per_batch` | `int` | `null` | Caps how many containers are solved at once, to bound per-solve memory. `null` disables the cap; when set it must be `>= 1`. On an **incremental** run, `Report.run()` commits at most this many upserted containers per iteration and loops until the population is drained (see [incremental](#incremental-optional)). On a **full** run, and when a changed definition recomputes over all containers, the population is split into chunks of about this size and solved chunk by chunk. Chunk sizes are approximate (the cap is a memory heuristic, not an exact bound). | | `solver_config` | `SolverConfig` | `null` | Per-table column mappings, per-table equality filters, and project scoping. Set `project_id` to scope reads by project — it is applied to `container_tags` (if configured), `container_metrics`, and `channel_mapping` (if configured), so it works in both narrow EAV and wide-only data models. Omit it when you don't need project scoping. See [Solver column mappings and filters](#solver-column-mappings-and-filters). | If `query_engine` is omitted, the default is `DefaultSolver` with @@ -352,6 +353,15 @@ mode-resolution rules and what counts as a definition change. | `silver_last_modified_column` | `str` | `"timestamp"` | Silver column used to detect container updates. Name after `column_name_mapping` is applied (the physical name unless you remap that column). | | `gold_last_modified_column` | `str` | `"_created_at"` | Gold-side column used to detect prior-run freshness. | +:::note Draining a capped incremental run +When [`query_engine.max_containers_per_batch`](#query_engine-optional) is set, a single +`Report.run()` commits at most that many upserted containers per iteration and loops until the +population is drained, so no one solve holds the whole population in memory. If a batch is detected +again unchanged after being persisted (no forward progress, e.g. a misconfigured +`gold_last_modified_column`), `run()` logs a warning and stops with partial completion rather than +looping forever. +::: + --- ## full_recalculation (optional) diff --git a/docs/impulse/docs/references/api/impulse_reporting/core/report.md b/docs/impulse/docs/references/api/impulse_reporting/core/report.md index 7dc38f22..450b2970 100644 --- a/docs/impulse/docs/references/api/impulse_reporting/core/report.md +++ b/docs/impulse/docs/references/api/impulse_reporting/core/report.md @@ -320,6 +320,32 @@ sink schema after persistence completes successfully. `None`: +#### run + +```python +def run(is_incremental: bool = None, + persist_results: bool = True, + cleanup_temp_tables: bool | None = None) +``` + +Determine and persist a report, draining all container batches. + +Wraps :meth:`determine_report` + :meth:`persist_results` (both remain usable +standalone). With ``max_containers_per_batch`` set on an incremental run, ``run()`` +loops: each iteration commits at most that many upserted containers, and the loop +continues until a run observes at most that many remaining (the last batch clears +the rest). After the first iteration all definition hashes are current, so later +iterations recompute only unchanged entities over the next batch of new containers. +Without a cap (or in full mode) it is a single determine+persist pass. + +**Arguments**: + +- `is_incremental` (`bool`): Forwarded to :meth:`determine_report` (config still overrides it). +- `persist_results` (`bool`): When True (default), persist after determining and drain the batches. When +False, run :meth:`determine_report` once without persisting (no iteration — +nothing commits, so the batch cannot advance). +- `cleanup_temp_tables` (`bool | None`): Forwarded to :meth:`persist_results`. + #### determine\_report ```python diff --git a/src/impulse_reporting/config/config_parser.py b/src/impulse_reporting/config/config_parser.py index 22c7147d..f6a44271 100644 --- a/src/impulse_reporting/config/config_parser.py +++ b/src/impulse_reporting/config/config_parser.py @@ -399,6 +399,18 @@ class QueryEngine(BaseModel): raw_encoder: RawEncoder | None = None solver_config: SolverConfig | None = None batch_size: int = 500 + max_containers_per_batch: int | None = None + + @field_validator("max_containers_per_batch", mode="after") + @classmethod + def _validate_max_containers_per_batch(cls, value: int | None) -> int | None: + """Max containers per solve chunk (memory bound); in incremental mode it also caps + the per-run drain batch. ``None`` disables chunking.""" + if value is not None and value < 1: + raise ValueError( + "max_containers_per_batch must be >= 1 when set (None disables the cap)." + ) + return value @model_validator(mode="after") def validate_drop_implausible_data_requires_raw(self): diff --git a/src/impulse_reporting/core/report.py b/src/impulse_reporting/core/report.py index 5a75cae3..0f0aab75 100644 --- a/src/impulse_reporting/core/report.py +++ b/src/impulse_reporting/core/report.py @@ -1,8 +1,10 @@ import json +import logging import zlib from typing import Any from databricks.sdk import WorkspaceClient from pyspark.sql import DataFrame, SparkSession +from pyspark.sql import functions as F from pyspark.sql.types import StructType from impulse_query_engine.analyze.metadata.time_series_expression import ( @@ -63,6 +65,8 @@ from impulse_reporting.util.report_entity_util import ReportEntityUtil from impulse_query_engine.telemetry import log_telemetry, telemetry_logger +logger = logging.getLogger(__name__) + class Report: """Represents a report containing pages, events, and configurations for data processing and persistence.""" @@ -911,6 +915,7 @@ def _solve_expressions_batched( catalog=getattr(self.config, "unity_sink", None) and self.config.unity_sink.catalog, schema=getattr(self.config, "unity_sink", None) and self.config.unity_sink.schema, pre_filtered_containers_df=pre_filtered_containers_df, + max_containers_per_batch=self.config.query_engine.max_containers_per_batch, ) def _solve_calculated_channels_batched( @@ -933,8 +938,72 @@ def _solve_calculated_channels_batched( catalog=getattr(self.config, "unity_sink", None) and self.config.unity_sink.catalog, schema=getattr(self.config, "unity_sink", None) and self.config.unity_sink.schema, pre_filtered_containers_df=pre_filtered_containers_df, + max_containers_per_batch=self.config.query_engine.max_containers_per_batch, ) + @telemetry_logger("report", "run") + def run( + self, + is_incremental: bool = None, + persist_results: bool = True, + cleanup_temp_tables: bool | None = None, + ): + """Determine and persist a report, draining all container batches. + + Wraps :meth:`determine_report` + :meth:`persist_results` (both remain usable + standalone). With ``max_containers_per_batch`` set on an incremental run, ``run()`` + loops: each iteration commits at most that many upserted containers, and the loop + continues until a run observes at most that many remaining (the last batch clears + the rest). After the first iteration all definition hashes are current, so later + iterations recompute only unchanged entities over the next batch of new containers. + Without a cap (or in full mode) it is a single determine+persist pass. + + Parameters + ---------- + is_incremental : bool, optional + Forwarded to :meth:`determine_report` (config still overrides it). + persist_results : bool, optional + When True (default), persist after determining and drain the batches. When + False, run :meth:`determine_report` once without persisting (no iteration — + nothing commits, so the batch cannot advance). + cleanup_temp_tables : bool | None, optional + Forwarded to :meth:`persist_results`. + + Notes + ----- + If a batch is detected again unchanged after being persisted (no forward + progress, e.g. a misconfigured ``gold_last_modified_column`` or containers + excluded by ``container_filters`` that never yield a gold row), the loop would + otherwise spin forever. ``run()`` detects the repeated batch, logs a warning + naming the stuck containers, and stops (partial completion) rather than hanging. + """ + previous_batch_ids = None + while True: + self.determine_report(is_incremental) + if not persist_results: + return + self.persist_results(cleanup_temp_tables) + # Stop once a run processed the remainder (uncapped upserted <= cap); also + # covers no cap / full mode, where the flag stays False. + if not self._more_batches_pending: + return + # Forward-progress guard: each iteration commits a deterministic batch (the + # lowest ``cap`` container ids). If the just-persisted batch is detected again + # unchanged, the persist removed nothing from detection, so the drain has + # stalled and the loop would never terminate. Stop instead of hanging. + current_batch_ids = frozenset(self._processed_container_ids or []) + if current_batch_ids == previous_batch_ids: + logger.warning( + "Incremental run stopped early: %d containers were detected again after " + "being persisted, so the run made no forward progress. First ids: %s. " + "Likely causes: a misconfigured gold_last_modified_column or containers that " + "produce no measurement_dimension row. Remaining containers were not processed.", + len(current_batch_ids), + sorted(current_batch_ids)[:10], + ) + return + previous_batch_ids = current_batch_ids + @telemetry_logger("report", "determine_report") def determine_report(self, is_incremental: bool = None): """ @@ -983,21 +1052,59 @@ def determine_report(self, is_incremental: bool = None): self._is_incremental = self._resolve_is_incremental(is_incremental) # Detect containers to process (incremental mode only): new + updated. + # ``_more_batches_pending`` tells run()'s batch loop whether more than one batch + # of upserted containers remains; reset each call so it never carries stale state. pre_filtered_containers_df = None + self._more_batches_pending = False + # Ids of the capped batch this run processes; used by run()'s loop to detect a + # stalled drain (the same batch recurring => no forward progress). Only populated + # on capped incremental runs, the only runs where run() iterates. + self._processed_container_ids = None + # Whether this run has containers to process (gates whether fact tables are + # written). Set from the collected batch under a cap, else probed with isEmpty(); + # defaults False (full mode without a cap / sinkless). + self._has_processed_containers = False if self._is_incremental: pre_filtered_containers_df = self._detect_upserted_containers() - - # Two signals for persistence: - # - has_processed_containers (new + updated): gates whether a fact table is - # written (new containers must be inserted). Only a bool is needed, so - # probe emptiness with isEmpty() rather than collecting the whole id list. - # - updated container ids: scopes the delete-by-source, since only - # containers that already have gold rows can have stale rows to prune. - self._has_processed_containers = ( - pre_filtered_containers_df is not None and not pre_filtered_containers_df.isEmpty() - ) + # Cap containers per run: committed containers get a fresh + # measurement_dimension timestamp and drop out of the next run's detection, + # so successive runs advance through the population. + max_containers = self.config.query_engine.max_containers_per_batch + if max_containers is not None and pre_filtered_containers_df is not None: + ( + pre_filtered_containers_df, + self._processed_container_ids, + self._more_batches_pending, + ) = self._cap_upserted_batch(pre_filtered_containers_df, max_containers) + self._has_processed_containers = bool(self._processed_container_ids) + elif pre_filtered_containers_df is not None: + self._has_processed_containers = not pre_filtered_containers_df.isEmpty() + elif self.config.query_engine.max_containers_per_batch is not None: + # Full mode with a cap: solve the full scoped population in one pass, chunked by + # the cap for memory. Supplying the frame here lets the batched solve chunk it; + # there is no cross-run deferral (``_more_batches_pending`` stays False). + pre_filtered_containers_df = self._filtered_container_metrics() + self._has_processed_containers = not pre_filtered_containers_df.isEmpty() + + # Updated container ids scope the delete-by-source; derived from the (capped) + # upserted set (see _updated_containers_within), so a capped run never prunes + # containers it did not reprocess. self._updated_container_ids = self._collect_container_ids( - self._detect_updated_containers() if self._is_incremental else None + self._updated_containers_within(pre_filtered_containers_df) + if self._is_incremental + else None + ) + + # Container scope for CHANGED entities. Incremental: without a cap, None (all + # containers); under a cap, all historical (gold) containers plus the capped new + # ones (``_changed_scope``), so a definition change is re-applied to existing + # containers while new containers beyond the cap are deferred. Full mode: the same + # full frame as the unchanged solve (``_changed_scope`` is an incremental-cap + # concept and must not run here). + changed_pre_filtered_containers_df = ( + self._changed_scope(pre_filtered_containers_df) + if self._is_incremental + else pre_filtered_containers_df ) hash_comparator = DefinitionHashComparator(self.spark) @@ -1040,7 +1147,8 @@ def determine_report(self, is_incremental: bool = None): # Centralized solve changed_solved_df = self._solve_expressions_batched( - all_changed_expressions, pre_filtered_containers_df=None + all_changed_expressions, + pre_filtered_containers_df=changed_pre_filtered_containers_df, ) unchanged_solved_df = self._solve_expressions_batched( all_unchanged_expressions, pre_filtered_containers_df=pre_filtered_containers_df @@ -1054,7 +1162,7 @@ def determine_report(self, is_incremental: bool = None): changed_solved_df, self.query, self.solver, - None, + changed_pre_filtered_containers_df, ContainerEvent, ) unchanged_event_dfs = dispatch_events( @@ -1095,8 +1203,9 @@ def determine_report(self, is_incremental: bool = None): # Calculated channels: own narrow batched solve driven here (mirrors the # wide expression solve above), producing a narrow ``solved_df`` that each - # channel type then shapes. Changed definitions recompute over all - # containers; unchanged ones over the incrementally-detected subset. + # channel type then shapes. Changed definitions recompute over the changed + # scope (all containers, or gold+capped under a cap); unchanged ones over the + # incrementally-detected subset. self._validate_unique_calculated_channels() channels_by_type = group_selectables_by_type(self.calculated_channels, ChannelType) changed_channels_by_type, unchanged_channels_by_type, self._changed_channel_ids = ( @@ -1119,7 +1228,7 @@ def determine_report(self, is_incremental: bool = None): c.expression for cs in unchanged_channels_by_type.values() for c in cs ] changed_channel_solved_df = self._solve_calculated_channels_batched( - changed_channel_exprs, pre_filtered_containers_df=None + changed_channel_exprs, pre_filtered_containers_df=changed_pre_filtered_containers_df ) unchanged_channel_solved_df = self._solve_calculated_channels_batched( unchanged_channel_exprs, pre_filtered_containers_df=pre_filtered_containers_df @@ -1184,9 +1293,9 @@ def determine_report(self, is_incremental: bool = None): ) # Determine channel mapping resolution dimension. - # Mirror the fact split: aliases from changed definitions resolve - # over all containers, aliases only in unchanged definitions stay - # scoped to the incrementally-detected containers. + # Mirror the fact split: aliases from changed definitions resolve over the + # changed scope (all containers, or gold+capped under a cap), aliases only in + # unchanged definitions stay scoped to the incrementally-detected containers. changed_aliased_selectors = TimeSeriesExpression.collect_selectors( all_changed_expressions, uses_alias=True, @@ -1203,6 +1312,7 @@ def determine_report(self, is_incremental: bool = None): changed_aliased_selectors=changed_aliased_selectors, unchanged_aliased_selectors=unchanged_aliased_selectors, pre_filtered_containers_df=pre_filtered_containers_df, + changed_pre_filtered_containers_df=changed_pre_filtered_containers_df, ) ) @@ -1228,9 +1338,18 @@ def _resolve_is_incremental(self, is_incremental: bool = None) -> bool: bool True for incremental processing, False for full processing. """ - # Rule 1: No gold layer → always FULL (nothing to compare against) + # Rule 1: No gold layer → FULL, EXCEPT the bootstrap of a capped incremental + # config. With a cap, the first run must process at most N containers and defer + # the rest, so it runs incrementally over all containers as "new" (see + # ``_detect_upserted_containers``) rather than computing the whole population. + # A non-capped incremental config still bootstraps in full mode (nothing to batch). if not self._gold_layer_exists(): - return False + return bool( + hasattr(self, "config") + and getattr(self.config, "incremental", None) is not None + and self.config.incremental.enabled + and self.config.query_engine.max_containers_per_batch is not None + ) if not hasattr(self, "config") and is_incremental is not None: return is_incremental @@ -1269,6 +1388,43 @@ def _gold_layer_exists(self) -> bool: measurement_dim_table = self.sink.config.get_output_uri_measurement_dimensions_table() return self.spark.catalog.tableExists(measurement_dim_table) + def _cap_upserted_batch( + self, upserted_df: DataFrame, max_containers: int + ) -> tuple[DataFrame, list, bool]: + """Take at most ``max_containers`` upserted containers. + + ``upserted_df`` is already produced from the solver's report-eligible container + pipeline. Ordering by ``container_id`` and collecting one id past the cap answers + both which containers this batch commits and whether more eligible containers remain. + + Returns + ------- + tuple[DataFrame, list, bool] + ``(capped_containers_df, batch_ids, more_pending)``. + """ + container_id_col = self.solver.config.container_id_col + batch_ids = [ + row[container_id_col] + for row in upserted_df.select(container_id_col) + .distinct() + .orderBy(container_id_col) + .limit(max_containers + 1) + .collect() + ] + more_pending = len(batch_ids) > max_containers + batch_ids = batch_ids[:max_containers] + capped_df = upserted_df.where(F.col(container_id_col).isin(batch_ids)) + return capped_df, batch_ids, more_pending + + def _filtered_container_metrics( + self, pre_filtered_containers_df: DataFrame = None + ) -> DataFrame: + """Resolve container metrics through the public solver container pipeline.""" + container_tags_df = self.solver.filter_container_tags(self.spark, self.query) + return self.solver.filter_container_metrics( + self.spark, self.query, container_tags_df, pre_filtered_containers_df + ) + def _detect_upserted_containers(self) -> DataFrame | None: """ Detect new and updated containers for incremental processing. @@ -1278,19 +1434,28 @@ def _detect_upserted_containers(self) -> DataFrame | None: used for freshness comparison. Falls back to ``"last_modified"`` when no incremental config is present. - Returns None if gold layer doesn't exist (triggers full processing) - or if no sink is configured (sinkless mode). + When the gold layer does not exist yet, every silver container is "new", + so all of them are returned (the capped-incremental bootstrap; this method + is only reached in incremental mode). Returns None only in sinkless mode + (no sink configured). Returns ------- DataFrame | None - DataFrame containing containers to process, or None if gold table - doesn't exist (indicating full processing is needed). + Containers to process (new + updated, or all silver containers on the + bootstrap first run), or None in sinkless mode. """ args = self._container_detection_args() if args is None: return None detector, silver_containers, measurement_dim_table, silver_col, gold_col = args + # Bootstrap of a capped incremental config (this method is only reached with + # ``_is_incremental`` True, and Rule 1 only makes that True on no gold when a cap + # is set): gold doesn't exist yet, so every silver container is "new". Return them + # all so the cap can slice the first batch; the detector would otherwise signal + # full processing by returning None on a missing gold table. + if not self._gold_layer_exists(): + return silver_containers return detector.detect_upserted_containers( silver_containers, measurement_dim_table, @@ -1298,29 +1463,46 @@ def _detect_upserted_containers(self) -> DataFrame | None: gold_last_modified_col=gold_col, ) - def _detect_updated_containers(self) -> DataFrame | None: - """Detect only UPDATED containers (present in gold, newer silver timestamp). + def _updated_containers_within(self, containers_df: DataFrame | None) -> DataFrame | None: + """Updated containers within the run's (capped) upserted set, for delete scoping. - Excludes new containers — see - ``ContainerUpsertDetector.detect_updated_containers``. Used to scope the - incremental delete-by-source. Returns None in sinkless mode or when the - gold table doesn't exist. + Delegates to ``ContainerUpsertDetector.updated_within``. Returns None in + sinkless mode or when ``containers_df`` is None. + """ + if containers_df is None: + return None + args = self._container_detection_args() + if args is None: + return None + _detector, _silver_containers, measurement_dim_table, _silver_col, _gold_col = args + return _detector.updated_within(containers_df, measurement_dim_table) - Returns - ------- - DataFrame | None - Updated containers, or None. + def _changed_scope(self, capped_upserted_df: DataFrame | None) -> DataFrame | None: + """Container scope for CHANGED entities under a cap. + + All historical (gold) containers plus the capped upserted set, so a definition + change is re-applied to existing containers while new containers beyond the cap + are deferred (they have no gold rows, so the global changed-entity delete cannot + orphan them). Returns None (all containers, today's behavior) when no cap is set + or the run is not incremental. """ + if self.config.query_engine.max_containers_per_batch is None or capped_upserted_df is None: + return None args = self._container_detection_args() if args is None: return None - detector, silver_containers, measurement_dim_table, silver_col, gold_col = args - return detector.detect_updated_containers( - silver_containers, - measurement_dim_table, - silver_last_modified_col=silver_col, - gold_last_modified_col=gold_col, - ) + _detector, silver_containers, measurement_dim_table, _silver_col, _gold_col = args + # Bootstrap: no gold table yet, so there are no historical containers to fold in — + # the changed scope is exactly the capped new set. (Also avoids reading a table + # that does not exist.) + if not self._gold_layer_exists(): + return capped_upserted_df + gold_ids = self.spark.read.table(measurement_dim_table).select("container_id") + allowed = gold_ids.unionByName(capped_upserted_df.select("container_id")).distinct() + # Full container_metrics rows for the allowed ids, not just the ids: downstream uses this + # frame as the container_metrics source, so it needs all columns. ``allowed`` has only + # ``container_id``. + return silver_containers.join(allowed, on="container_id", how="inner") def _container_detection_args(self): """Shared inputs for container detection, or None in sinkless mode. @@ -1335,11 +1517,11 @@ def _container_detection_args(self): if not self._has_sink: return None detector = ContainerUpsertDetector(self.spark) - # Read container_metrics exactly as the solver processes it (column_name_mapping, - # project_id, and per-table filters applied), so detection joins on the internal - # ``container_id`` and sees the same container universe as the solve. Reading the raw - # table here would break configs that remap the physical container-id column (#99). - silver_containers = self.solver.scoped_container_metrics(self.spark, self.query) + # Read container_metrics through the solver's public container pipeline, so + # detection joins on the internal ``container_id`` and sees the same report-eligible + # container universe as the solve. Reading the raw table here would break configs + # that remap the physical container-id column (#99). + silver_containers = self._filtered_container_metrics() measurement_dim_table = self.sink.config.get_output_uri_measurement_dimensions_table() silver_col = "last_modified" diff --git a/src/impulse_reporting/core/report_utils.py b/src/impulse_reporting/core/report_utils.py index c161df9f..2923d1b5 100644 --- a/src/impulse_reporting/core/report_utils.py +++ b/src/impulse_reporting/core/report_utils.py @@ -6,6 +6,7 @@ from __future__ import annotations +import math import uuid from functools import reduce from typing import TYPE_CHECKING @@ -499,6 +500,67 @@ def dispatch_calculated_channel_metrics( return metrics_dfs +def _materialize_temp(df: DataFrame, run_id: str, idx, has_sink: bool, catalog, schema) -> str: + """Persist an intermediate ``df`` as a ``__impulse_temp_*`` Delta table (or temp view). + + Returns the name to ``spark.table(...)`` it back with. Delta when a sink is + configured, a Spark temp view otherwise; the shared ``__impulse_temp_`` prefix means + :func:`cleanup_temp_tables` covers both. + """ + name = f"__impulse_temp_{run_id}_{idx}" + if has_sink: + fq_name = f"`{catalog}`.`{schema}`.`{name}`" + df.write.format("delta").mode("overwrite").saveAsTable(fq_name) + return fq_name + df.createOrReplaceTempView(name) + return name + + +def _container_chunks(effective_df, max_n, cid_col): + """Yield ``effective_df`` partitioned into chunks of ~``max_n`` containers. + + Buckets containers by ``pmod(hash(container_id), num_chunks)`` (with ``num_chunks`` + sized so each bucket holds ~``max_n``) and yields ``effective_df`` filtered per bucket. + Only a ``count`` reaches the driver, never the id list, so a multi-million-container + full or changed-definition solve stays driver-bounded. Buckets are disjoint and cover + every container, so the union is the full population; the deterministic hash keeps + membership reproducible and works for any ``container_id`` type. Chunk sizes are + approximate — the cap is a per-solve memory heuristic, not an exact bound. + """ + num_containers = effective_df.select(cid_col).distinct().count() + if num_containers == 0: + return + num_chunks = math.ceil(num_containers / max_n) + if num_chunks == 1: + yield effective_df + return + bucket = F.pmod(F.hash(F.col(cid_col)), F.lit(num_chunks)) + for i in range(num_chunks): + yield effective_df.where(bucket == i) + + +def _combine_container_chunks( + spark: SparkSession, + parts: list[DataFrame], + has_sink: bool, + catalog: str, + schema: str, +) -> DataFrame: + """Combine per-chunk solve results (disjoint containers, so a row append). + + With a sink, append each chunk into one ``__impulse_temp_*`` Delta table and read it + back once, avoiding a deep ``unionByName`` tree and materializing one chunk at a time. + Without a sink, ``unionByName`` the chunks directly + """ + if not has_sink: + return reduce(lambda a, b: a.unionByName(b), parts) + run_id = uuid.uuid4().hex[:8] + fq_name = f"`{catalog}`.`{schema}`.`__impulse_temp_{run_id}_chunks`" + for idx, part in enumerate(parts): + part.write.format("delta").mode("overwrite" if idx == 0 else "append").saveAsTable(fq_name) + return spark.table(fq_name) + + def solve_expressions_batched( spark: SparkSession, expressions: list[TimeSeriesExpression], @@ -510,14 +572,15 @@ def solve_expressions_batched( catalog: str = None, schema: str = None, pre_filtered_containers_df: DataFrame = None, + max_containers_per_batch: int = None, ) -> DataFrame | None: - """Solve all expressions in configurable batches and return a joined wide DataFrame. + """Solve ``expressions`` in selector batches; return a wide DataFrame joined on ``container_id``. - Each batch is solved independently via ``query.select(*batch_exprs).solve(...)``. - When a sink is configured the intermediate result is persisted as a temporary - Delta table (``__impulse_temp__``); otherwise a Spark temp view - is used. After all batches are solved the per-batch DataFrames are joined on - ``container_id`` with a full outer join. + Each batch is solved via ``query.select(*batch).solve(...)``, materialized to a temp Delta + table/view, then combined with a full-outer join. When both ``max_containers_per_batch`` and + ``pre_filtered_containers_df`` are given, that set is chunked into ~that many containers per + solve and the chunks are row-appended (disjoint containers), bounding memory with a result + identical to an unchunked solve. With no cap or no pre-filter, a single pass runs. Parameters ---------- @@ -539,6 +602,8 @@ def solve_expressions_batched( Unity Catalog schema name (required when *has_sink* is ``True``). pre_filtered_containers_df : DataFrame, optional Pre-filtered containers for incremental processing. + max_containers_per_batch : int, optional + Max containers per solve chunk; ``None`` disables container chunking. Returns ------- @@ -549,36 +614,50 @@ def solve_expressions_batched( if not expressions: return None - run_id = uuid.uuid4().hex[:8] - batches = build_batches(expressions, batch_size) - - batch_names: list[str] = [] - for batch_idx, batch_exprs in enumerate(batches): - batch_query = query.select(*batch_exprs) - batch_df = batch_query.solve( - spark=spark, - solver=solver, - pre_filtered_containers_df=pre_filtered_containers_df, - ) - - if has_sink: - table_name = f"__impulse_temp_{run_id}_{batch_idx}" - fq_name = f"`{catalog}`.`{schema}`.`{table_name}`" - batch_df.write.format("delta").mode("overwrite").saveAsTable(fq_name) - batch_names.append(fq_name) - else: - view_name = f"__impulse_temp_{run_id}_{batch_idx}" - batch_df.createOrReplaceTempView(view_name) - batch_names.append(view_name) - cid_col = solver.config.container_id_col - dfs = [spark.table(name) for name in batch_names] - result = dfs[0] - for i in range(1, len(dfs)): - result = result.join(dfs[i], on=cid_col, how="full_outer") - - return result + def _solve_selector_batches(pre_filter: DataFrame | None) -> DataFrame: + run_id = uuid.uuid4().hex[:8] + batch_names = [ + _materialize_temp( + query.select(*batch_exprs).solve( + spark=spark, solver=solver, pre_filtered_containers_df=pre_filter + ), + run_id, + batch_idx, + has_sink, + catalog, + schema, + ) + for batch_idx, batch_exprs in enumerate(build_batches(expressions, batch_size)) + ] + dfs = [spark.table(name) for name in batch_names] + result = dfs[0] + for df in dfs[1:]: + result = result.join(df, on=cid_col, how="full_outer") + return result + + # Chunking only applies to an explicit, detection-derived container set (which carries + # the internal ``container_id``). With no cap, or with no pre-filter (e.g. a genuine + # full solve, or sinkless mode where detection yields nothing), solve unchunked. + if max_containers_per_batch is None or pre_filtered_containers_df is None: + return _solve_selector_batches(pre_filtered_containers_df) + + parts = [ + _solve_selector_batches(chunk) + for chunk in _container_chunks( + pre_filtered_containers_df, max_containers_per_batch, cid_col + ) + ] + # No chunks => empty container set; fall back to a single solve so the result matches + # the non-chunked path (an empty-rows DataFrame, not None). + if not parts: + return _solve_selector_batches(pre_filtered_containers_df) + # One chunk (e.g. incremental unchanged path, population <= cap): return it directly, + # no extra Delta round-trip. + if len(parts) == 1: + return parts[0] + return _combine_container_chunks(spark, parts, has_sink, catalog, schema) def solve_calculated_channels_batched( @@ -592,21 +671,16 @@ def solve_calculated_channels_batched( catalog: str = None, schema: str = None, pre_filtered_containers_df: DataFrame = None, + max_containers_per_batch: int = None, ) -> DataFrame | None: - """Solve calculated channels in configurable batches; return the unioned rows. - - The narrow, row-append counterpart to :func:`solve_expressions_batched`. Each - batch is solved independently via - ``query.select(*batch).solve_calculated_channels(...)`` and persisted as a - temporary Delta table (``__impulse_temp__``) when a sink is - configured, or a Spark temp view otherwise — the same convention (and shared - ``__impulse_temp_*`` prefix, so :func:`cleanup_temp_tables` covers it). + """Solve calculated channels in selector batches; return the ``unionByName``-appended rows. - Unlike ``solve_expressions_batched`` (wide, one row per container → batches - combined with a full-outer join on ``container_id``), calculated-channel output - is narrow (``container_id, channel_id, tstart, tend, value, identity``; many - rows per container, batches hold different ``channel_id``s), so batches are - combined with **``unionByName``** (row append). + Narrow, row-append counterpart to :func:`solve_expressions_batched`: each batch is solved via + ``query.select(*batch).solve_calculated_channels(...)``, materialized to a temp Delta + table/view, then combined with ``unionByName`` (many rows per container). When both + ``max_containers_per_batch`` and ``pre_filtered_containers_df`` are given, that set is chunked + into ~that many containers per solve and the chunks are appended too, with a result identical + to an unchunked solve. With no cap or no pre-filter, a single pass runs. Parameters ---------- @@ -629,6 +703,8 @@ def solve_calculated_channels_batched( Unity Catalog schema name (required when *has_sink* is ``True``). pre_filtered_containers_df : DataFrame, optional Pre-filtered containers for incremental processing. + max_containers_per_batch : int, optional + Max containers per solve chunk; ``None`` disables container chunking. Returns ------- @@ -639,32 +715,48 @@ def solve_calculated_channels_batched( if not qe_channels: return None - run_id = uuid.uuid4().hex[:8] - batches = build_batches(qe_channels, batch_size) + cid_col = solver.config.container_id_col - batch_names: list[str] = [] - for batch_idx, batch_channels in enumerate(batches): - batch_df = query.select(*batch_channels).solve_calculated_channels( - spark, solver, pre_filtered_containers_df + def _solve_selector_batches(pre_filter: DataFrame | None) -> DataFrame: + run_id = uuid.uuid4().hex[:8] + batch_names = [ + _materialize_temp( + query.select(*batch_channels).solve_calculated_channels(spark, solver, pre_filter), + run_id, + batch_idx, + has_sink, + catalog, + schema, + ) + for batch_idx, batch_channels in enumerate(build_batches(qe_channels, batch_size)) + ] + dfs = [spark.table(name) for name in batch_names] + result = dfs[0] + for df in dfs[1:]: + result = result.unionByName(df) + return result + + # Chunking only applies to an explicit, detection-derived container set (which carries + # the internal ``container_id``). With no cap, or with no pre-filter (e.g. a genuine + # full solve, or sinkless mode where detection yields nothing), solve unchunked. + if max_containers_per_batch is None or pre_filtered_containers_df is None: + return _solve_selector_batches(pre_filtered_containers_df) + + parts = [ + _solve_selector_batches(chunk) + for chunk in _container_chunks( + pre_filtered_containers_df, max_containers_per_batch, cid_col ) - - if has_sink: - table_name = f"__impulse_temp_{run_id}_{batch_idx}" - fq_name = f"`{catalog}`.`{schema}`.`{table_name}`" - batch_df.write.format("delta").mode("overwrite").saveAsTable(fq_name) - batch_names.append(fq_name) - else: - view_name = f"__impulse_temp_{run_id}_{batch_idx}" - batch_df.createOrReplaceTempView(view_name) - batch_names.append(view_name) - - dfs = [spark.table(name) for name in batch_names] - - result = dfs[0] - for df in dfs[1:]: - result = result.unionByName(df) - - return result + ] + # No chunks => empty container set; fall back to a single solve so the result matches + # the non-chunked path (an empty-rows DataFrame, not None). + if not parts: + return _solve_selector_batches(pre_filtered_containers_df) + # One chunk (e.g. incremental unchanged path, population <= cap): return it directly, + # no extra Delta round-trip. + if len(parts) == 1: + return parts[0] + return _combine_container_chunks(spark, parts, has_sink, catalog, schema) def cleanup_temp_tables(spark: SparkSession, catalog: str, schema: str) -> None: diff --git a/src/impulse_reporting/incremental/container_detector.py b/src/impulse_reporting/incremental/container_detector.py index 74246a8f..4bbf16da 100644 --- a/src/impulse_reporting/incremental/container_detector.py +++ b/src/impulse_reporting/incremental/container_detector.py @@ -99,46 +99,40 @@ def detect_upserted_containers( upserted = new_containers.unionByName(updated_containers).dropDuplicates(["container_id"]) return upserted - def detect_updated_containers( + def updated_within( self, - silver_containers_df: DataFrame, + candidate_containers_df: DataFrame, gold_measurement_dim_table: str, - silver_last_modified_col: str = "last_modified", - gold_last_modified_col: str = "last_modified", - ) -> DataFrame | None: - """Detect only UPDATED containers (present in gold with a newer silver timestamp). + ) -> DataFrame: + """Restrict upserted candidates to their UPDATED containers. - Unlike :meth:`detect_upserted_containers`, this excludes NEW containers. - Only containers that already have gold rows can have *stale* rows, so this - is the correct set for scoping delete-by-source pruning. + The updated containers are those whose ``container_id`` already exists in the + gold ``measurement_dimension`` — new ones don't, and have no stale rows to + prune. Inner-joining candidates with gold both selects them and scopes them to + the candidates passed in, so a capped run's delete-by-source never touches + containers it did not reprocess. Candidates are already the upserted set, so + membership in gold is sufficient — no timestamp comparison needed. Parameters ---------- - silver_containers_df : DataFrame - Container metrics from silver layer. Must contain ``container_id``. + candidate_containers_df : DataFrame + Upserted containers (silver schema, includes ``container_id``); e.g. the + capped ``pre_filtered_containers_df``. gold_measurement_dim_table : str URI of the gold measurement_dimension table. - silver_last_modified_col : str, optional - Silver freshness column, by default ``"last_modified"``. - gold_last_modified_col : str, optional - Gold freshness column, by default ``"last_modified"``. Returns ------- - DataFrame | None - Updated containers (silver schema), or None if the gold table - doesn't exist. + DataFrame + Candidates already present in gold (silver schema); empty when gold absent. """ if not self._table_exists(gold_measurement_dim_table): - return None + return candidate_containers_df.limit(0) - gold_df = self.spark.read.table(gold_measurement_dim_table) - return self._identify_updated_containers( - silver_containers_df, - gold_df, - silver_last_modified_col, - gold_last_modified_col, + gold_ids = ( + self.spark.read.table(gold_measurement_dim_table).select("container_id").distinct() ) + return candidate_containers_df.join(gold_ids, on="container_id", how="inner") def _identify_new_containers(self, silver_df: DataFrame, gold_df: DataFrame) -> DataFrame: """ diff --git a/src/impulse_reporting/meta/container_dimensions.py b/src/impulse_reporting/meta/container_dimensions.py index 9612ceee..ac7a3ae7 100644 --- a/src/impulse_reporting/meta/container_dimensions.py +++ b/src/impulse_reporting/meta/container_dimensions.py @@ -194,18 +194,19 @@ def get_dimension_for_scopes( changed_aliased_selectors: list[TimeSeriesSelector], unchanged_aliased_selectors: list[TimeSeriesSelector], pre_filtered_containers_df: DataFrame = None, + changed_pre_filtered_containers_df: DataFrame = None, ) -> DataFrame | None: """ Compute the channel mapping resolution dimension honoring the report's changed/unchanged definition split. Mirrors the scoping the fact pipeline uses: aliases referenced by - *changed* definitions are resolved over **all** containers - (``pre_filtered_containers_df=None``), while aliases referenced - only by *unchanged* definitions stay scoped to - ``pre_filtered_containers_df``. This keeps incremental runs cheap - without leaving older containers unresolved when a new alias is - introduced by a changed definition. + *changed* definitions are resolved over the changed scope + (``changed_pre_filtered_containers_df`` — all containers when None, + or gold+capped under a cap), while aliases referenced only by + *unchanged* definitions stay scoped to ``pre_filtered_containers_df``. + This keeps incremental runs cheap without leaving older containers + unresolved when a new alias is introduced by a changed definition. In full (non-incremental) mode ``pre_filtered_containers_df`` is ``None``, so both scopes resolve over all containers and the union @@ -234,7 +235,9 @@ def get_dimension_for_scopes( Aliased selectors from unchanged definitions (resolved over ``pre_filtered_containers_df``). May be empty. pre_filtered_containers_df : DataFrame, optional - Pre-filtered containers for incremental processing. + Pre-filtered containers for the unchanged (incremental) scope. + changed_pre_filtered_containers_df : DataFrame, optional + Container scope for changed aliases; None = all containers. Returns ------- @@ -254,7 +257,7 @@ def get_dimension_for_scopes( query=query, solver=solver, aliased_selectors=changed_aliased_selectors, - pre_filtered_containers_df=None, + pre_filtered_containers_df=changed_pre_filtered_containers_df, ) unchanged_df = ChannelMappingResolutionDimension.get_dimension( spark=spark, diff --git a/tests/impulse_reporting/integration/max_containers_per_batch_test.py b/tests/impulse_reporting/integration/max_containers_per_batch_test.py new file mode 100644 index 00000000..173eb428 --- /dev/null +++ b/tests/impulse_reporting/integration/max_containers_per_batch_test.py @@ -0,0 +1,454 @@ +"""Integration tests for ``query_engine.max_containers_per_batch``. + +The cap bounds upserted containers per incremental run; committed containers drop out of +the next run's detection, so repeated runs iterate the population and the final gold matches +an uncapped run. The cap is incremental-only: the FIRST run of a capped config treats every +container as "new" and caps-and-defers even before gold exists, so successive runs drain the +full silver table (containers 1, 2, 3). + +These tests use lightweight aggregations (few histogram bins) and share a single uncapped +baseline (computed once) so each determine+persist cycle stays cheap. +""" + +import logging +from unittest.mock import create_autospec + +import pyspark.sql.functions as F +import pytest +from databricks.sdk import WorkspaceClient + +from impulse_reporting.aggregations.histogram import HistogramDuration +from impulse_reporting.aggregations.stats_aggregator import StatsAggregator +from impulse_reporting.config.config_parser import ( + Comparator, + ContainerFilters, + IncrementalConfig, + ImpulseConfig, + MetricFilter, + QueryEngine, + Source, + UnitySink, +) +from impulse_reporting.core.page import Page +from impulse_reporting.core.report import Report +from impulse_reporting.events.basic_event import BasicEvent +from impulse_reporting.events.container_event import ContainerEvent +from tests.conftest import spark + + +def _add_light_aggs(report, *, rpm_bins=None): + """Register a small event/aggregation set on the report. + + Mirrors the entity mix used across the reporting tests (two histograms, a basic event, an + event-scoped stats aggregator, a container event + its stats aggregator) but with tiny + histogram bins so each determine+persist cycle stays cheap. ``rpm_bins`` lets a caller flip + the rpm-histogram definition (its hash) to exercise the changed-definition path. + """ + rpm_bins = rpm_bins if rpm_bins is not None else [float(i) for i in range(0, 8000, 2000)] + query = report.get_db().query + c1 = query.channel(channel_name="Engine RPM") + c2 = query.channel(channel_name="Vehicle Speed Sensor") + page = Page(page_number=1) + report.add_page(page) + page.add_aggregation(HistogramDuration("rpm_hist_p1", base_expr=c1, bins=rpm_bins)) + page.add_aggregation( + HistogramDuration( + "speed_hist_p1", base_expr=c2, bins=[float(i) for i in range(0, 300, 100)] + ) + ) + rpm_event = BasicEvent(name="rpm_event", expr=c1 > 0, desc="engine speed > 0 rpm") + report.add_event(rpm_event) + page.add_aggregation( + StatsAggregator( + name="stats_agg", + input_expressions=[c1], + channel_names=["Engine RPM"], + event=rpm_event, + statistics=["start", "end", "mean"], + ) + ) + container_event = ContainerEvent("Measurement Event") + report.add_event(container_event) + page.add_aggregation( + StatsAggregator( + name="stats_agg_container", + input_expressions=[c1], + channel_names=["Engine RPM"], + event=container_event, + statistics=["start", "end", "mean"], + ) + ) + + +def _add_light_aggs_changed(report): + """Same as :func:`_add_light_aggs` but with different rpm bins (a changed definition hash).""" + _add_light_aggs(report, rpm_bins=[float(i) for i in range(0, 8000, 1000)]) + + +def _config(silver_table, prefix, *, max_containers_per_batch=None, container_filters=None): + return ImpulseConfig( + source=Source( + container_metrics_table=f"spark_catalog.silver.{silver_table}", + channel_metrics_table="spark_catalog.silver.channel_metrics", + channels_uri="spark_catalog.silver.channels", + ), + unity_sink=UnitySink(catalog="spark_catalog", schema="gold", table_prefix=prefix), + incremental=IncrementalConfig( + enabled=True, + silver_last_modified_column="timestamp", + gold_last_modified_column="_created_at", + ), + container_filters=container_filters, + query_engine=QueryEngine(max_containers_per_batch=max_containers_per_batch), + ) + + +def _make_report( + spark, + silver_table, + prefix, + *, + max_containers_per_batch=None, + add_aggs=_add_light_aggs, + container_filters=None, +): + report = Report( + name="cap_report", + spark=spark, + workspace_client=create_autospec(WorkspaceClient), + config=dict( + _config( + silver_table, + prefix, + max_containers_per_batch=max_containers_per_batch, + container_filters=container_filters, + ) + ), + ) + add_aggs(report) + return report + + +def _run(spark, silver_table, prefix, *, max_containers_per_batch=None, add_aggs=_add_light_aggs): + """Run ONE determine+persist pass (single batch) — the granular API, no run() loop.""" + report = _make_report( + spark, + silver_table, + prefix, + max_containers_per_batch=max_containers_per_batch, + add_aggs=add_aggs, + ) + report.determine_report() + report.persist_results() + + +def _container_ids(spark, prefix): + return sorted( + r.container_id + for r in spark.read.table(f"spark_catalog.gold.{prefix}_measurement_dimension").collect() + ) + + +def _hist_container_ids(spark, prefix): + return { + r.container_id + for r in spark.read.table(f"spark_catalog.gold.{prefix}_histogram_fact") + .select("container_id") + .distinct() + .collect() + } + + +def _rows_without_meta(df): + """Deterministic row set, dropping run-dependent columns. + + Excludes ``_``-prefixed meta (e.g. ``_created_at``) and ``config_hash`` (a hash of the + full config, which intentionally differs between the capped/uncapped configs). + """ + cols = [c for c in df.columns if not c.startswith("_") and c != "config_hash"] + return sorted(tuple(r) for r in df.select(*cols).collect()) + + +_BASELINE_TABLES = ("histogram_fact", "stats_aggregator_fact", "measurement_dimension") + + +@pytest.fixture(scope="module") +def uncapped_baseline(spark): + """Uncapped gold for the light aggs, computed once and reused as snapshots. + + Returns ``{"plain": {table: rows}, "changed": {table: rows}}`` where rows are the + ``_rows_without_meta`` snapshots of each gold fact table from a single uncapped full run. + Collected eagerly so the snapshots survive the per-test ``cleanup_gold`` teardown. + """ + baselines = {} + for key, add_aggs, prefix in ( + ("plain", _add_light_aggs, "plainbase"), + ("changed", _add_light_aggs_changed, "changedbase"), + ): + _run(spark, "container_metrics", prefix, add_aggs=add_aggs) + baselines[key] = { + t: _rows_without_meta(spark.read.table(f"spark_catalog.gold.{prefix}_{t}")) + for t in _BASELINE_TABLES + } + return baselines + + +def test_bootstrap_capped_defers_beyond_cap_and_matches_uncapped(spark, uncapped_baseline): + """Bootstrap (empty gold): the first capped run treats every container as new and caps it.""" + # Capped bootstrap (no gold): the first run processes only the lowest container. + _run(spark, "container_metrics", "bootcap", max_containers_per_batch=1) + assert _container_ids(spark, "bootcap") == [1], "bootstrap must cap the first run to one" + + # Successive runs drain the population (committed containers drop out of detection). + _run(spark, "container_metrics", "bootcap", max_containers_per_batch=1) + assert _container_ids(spark, "bootcap") == [1, 2] + _run(spark, "container_metrics", "bootcap", max_containers_per_batch=1) + assert _container_ids(spark, "bootcap") == [1, 2, 3] + + for t in _BASELINE_TABLES: + got = _rows_without_meta(spark.read.table(f"spark_catalog.gold.bootcap_{t}")) + assert got == uncapped_baseline["plain"][t], f"{t}: bootstrap drain must match uncapped" + + # Real-value sanity check: the histogram carries positive accumulated duration. + total = ( + spark.read.table("spark_catalog.gold.bootcap_histogram_fact") + .agg(F.sum("hist_value").alias("s")) + .collect()[0]["s"] + ) + assert total is not None and total > 0 + + +def test_bootstrap_capped_drains_in_one_run_call(spark, uncapped_baseline): + """A single run() call drains the whole population from an empty gold.""" + _make_report(spark, "container_metrics", "bootrun", max_containers_per_batch=1).run() + assert _container_ids(spark, "bootrun") == [1, 2, 3] + for t in _BASELINE_TABLES: + got = _rows_without_meta(spark.read.table(f"spark_catalog.gold.bootrun_{t}")) + assert got == uncapped_baseline["plain"][t], f"{t}: run() drain must match uncapped" + + +def test_capped_run_caps_after_metric_container_filters(spark): + """Filtered-out low ids must not consume capped incremental batch slots.""" + filters = ContainerFilters( + metric_filters=[ + [MetricFilter(column_name="container_id", comparator=Comparator.GT, value=1)] + ] + ) + report = _make_report( + spark, + "container_metrics", + "filtercap", + max_containers_per_batch=1, + container_filters=filters, + ) + + report.run() + + assert _container_ids(spark, "filtercap") == [2, 3] + + +def test_full_mode_container_chunked_solve_matches_uncapped(spark, uncapped_baseline): + """Full mode (no incremental) + cap: the solve is chunked by the cap in a single pass. + + Chunking is a memory split, not a content change, so gold must equal the uncapped run. + """ + config = ImpulseConfig( + source=Source( + container_metrics_table="spark_catalog.silver.container_metrics", + channel_metrics_table="spark_catalog.silver.channel_metrics", + channels_uri="spark_catalog.silver.channels", + ), + unity_sink=UnitySink(catalog="spark_catalog", schema="gold", table_prefix="fullcap"), + # No incremental config -> full mode; cap chunks the single full pass. + query_engine=QueryEngine(max_containers_per_batch=1), + ) + report = Report( + name="cap_report", + spark=spark, + workspace_client=create_autospec(WorkspaceClient), + config=dict(config), + ) + _add_light_aggs(report) + report.determine_report() + report.persist_results() + + assert _container_ids(spark, "fullcap") == [ + 1, + 2, + 3, + ], "full-mode solve processes all containers" + for t in _BASELINE_TABLES: + got = _rows_without_meta(spark.read.table(f"spark_catalog.gold.fullcap_{t}")) + assert got == uncapped_baseline["plain"][t], f"{t}: chunked full solve must match uncapped" + + +def _sinkless_hist(spark, *, max_containers_per_batch): + """Sinkless full-mode report over the basic silver DB; return its HISTOGRAM fact.""" + report = Report( + name="cap_report", + spark=spark, + workspace_client=create_autospec(WorkspaceClient), + config=dict( + ImpulseConfig( + source=Source( + container_metrics_table="spark_catalog.silver.container_metrics", + channel_metrics_table="spark_catalog.silver.channel_metrics", + channels_uri="spark_catalog.silver.channels", + ), + # No unity_sink -> sinkless (has_sink=False), so the chunked solve combines + # via the reduce(unionByName) branch of _combine_container_chunks. + query_engine=QueryEngine(max_containers_per_batch=max_containers_per_batch), + ) + ), + ) + _add_light_aggs(report) + report.determine_report() + return report.aggregation_dfs["HISTOGRAM"]["changed"] + + +def test_sinkless_capped_chunked_solve_matches_uncapped(spark): + """Sinkless + cap: chunks are unioned via _combine_container_chunks(has_sink=False). + + cap=1 over the multi-container silver DB forces multiple chunks; the sinkless union must + reassemble them into the same histogram as an uncapped sinkless run. + """ + capped = _sinkless_hist(spark, max_containers_per_batch=1) + uncapped = _sinkless_hist(spark, max_containers_per_batch=None) + + cap_ids = {r.container_id for r in capped.select("container_id").distinct().collect()} + assert len(cap_ids) > 1, "cap=1 must produce multiple chunks to exercise the union" + assert capped.filter(F.col("hist_value") > 0).count() > 0 + assert _rows_without_meta(capped) == _rows_without_meta(uncapped) + + +def test_capped_run_does_not_prune_out_of_batch_updated_container(spark): + """A capped run must not prune gold rows for updated containers outside the cap.""" + # Seed gold with containers 1 and 2 (initial full load). + _run(spark, "container_metrics_inc_1_2", "prune") + assert _container_ids(spark, "prune") == [1, 2] + + hist_pre = spark.read.table("spark_catalog.gold.prune_histogram_fact") + container2_hist_pre = ( + hist_pre.where(F.col("container_id") == 2).orderBy("visual_id", "bin_id").collect() + ) + assert container2_hist_pre, "container 2 must have histogram rows after the initial load" + meas_pre = spark.read.table("spark_catalog.gold.prune_measurement_dimension") + container2_created_at_pre = ( + meas_pre.where(F.col("container_id") == 2).select("_created_at").collect()[0][0] + ) + + # Mark BOTH containers as updated in silver (timestamp newer than gold _created_at). + modified = spark.read.table("spark_catalog.silver.container_metrics_inc_1_2").withColumn( + "timestamp", F.current_timestamp() + ) + modified.write.format("delta").mode("overwrite").saveAsTable( + "spark_catalog.silver.container_metrics_prune_modified" + ) + + # Incremental run with cap=1: only container 1 (lowest id) is reprocessed; + # container 2 is updated but OUTSIDE the cap. + _run(spark, "container_metrics_prune_modified", "prune", max_containers_per_batch=1) + + assert _container_ids(spark, "prune") == [1, 2], "no container may be dropped" + + hist_post = spark.read.table("spark_catalog.gold.prune_histogram_fact") + container2_hist_post = ( + hist_post.where(F.col("container_id") == 2).orderBy("visual_id", "bin_id").collect() + ) + # Container 2 was NOT in the cap -> its gold rows must be untouched (not pruned). + assert container2_hist_post == container2_hist_pre + meas_post = spark.read.table("spark_catalog.gold.prune_measurement_dimension") + container2_created_at_post = ( + meas_post.where(F.col("container_id") == 2).select("_created_at").collect()[0][0] + ) + assert container2_created_at_pre == container2_created_at_post, "container 2 must be untouched" + + +def test_changed_entity_capped_defers_beyond_cap_new(spark, uncapped_baseline): + """A changed definition recomputes historical + capped-new; new-beyond-cap is deferred.""" + # Seed gold with container 1 under definition D1. + _run(spark, "container_metrics_inc_1", "chg") + assert _container_ids(spark, "chg") == [1] + + # Change the definition (D2) and add containers 2, 3; incremental, cap=1. + # Run 1: changed entity computed for historical {1} + capped new {2}; container 3 deferred. + _run( + spark, + "container_metrics", + "chg", + max_containers_per_batch=1, + add_aggs=_add_light_aggs_changed, + ) + assert _container_ids(spark, "chg") == [1, 2], "beyond-cap new container 3 must be deferred" + assert _hist_container_ids(spark, "chg") == {1, 2}, "no facts for the deferred container" + + # Run 2 advances to container 3 (definition now unchanged -> unchanged path). + _run( + spark, + "container_metrics", + "chg", + max_containers_per_batch=1, + add_aggs=_add_light_aggs_changed, + ) + assert _container_ids(spark, "chg") == [1, 2, 3] + + for t in ("histogram_fact", "stats_aggregator_fact"): + got = _rows_without_meta(spark.read.table(f"spark_catalog.gold.chg_{t}")) + assert ( + got == uncapped_baseline["changed"][t] + ), f"{t}: changed iteration must match uncapped" + + +def test_run_loop_with_changed_definition(spark, uncapped_baseline): + """run() drains batches when a definition changed: iter 1 recomputes, the rest are unchanged.""" + _run(spark, "container_metrics_inc_1", "loopchg") # D1 on container 1 + + # Changed definition (D2) + new {2, 3}, cap=1, single run() call drains everything. + _make_report( + spark, + "container_metrics", + "loopchg", + max_containers_per_batch=1, + add_aggs=_add_light_aggs_changed, + ).run() + assert _container_ids(spark, "loopchg") == [1, 2, 3] + + for t in ("histogram_fact", "stats_aggregator_fact"): + got = _rows_without_meta(spark.read.table(f"spark_catalog.gold.loopchg_{t}")) + assert got == uncapped_baseline["changed"][t], t + + +def test_run_without_persist_does_not_loop(spark): + """run(persist_results=False) runs determine once and does not iterate.""" + _run(spark, "container_metrics_inc_1", "nopersist") + report = _make_report(spark, "container_metrics", "nopersist", max_containers_per_batch=1) + report.run(persist_results=False) + # Nothing was persisted beyond the seed, so gold still holds only container 1. + assert _container_ids(spark, "nopersist") == [1] + + +def test_run_stops_when_drain_makes_no_progress(spark, caplog): + """A stalled drain must stop with a logged warning instead of looping forever. + + Every container's freshness ``timestamp`` is set far in the future, so a persisted + container is always re-detected as updated and never drops out. With cap=1 the lowest + container (1) is selected every iteration and blocks the drain. run() must detect the + repeated batch, log a warning, and return rather than hang. + """ + spark.read.table("spark_catalog.silver.container_metrics").withColumn( + "timestamp", F.lit("2999-01-01 00:00:00").cast("timestamp") + ).write.format("delta").mode("overwrite").saveAsTable( + "spark_catalog.silver.container_metrics_future" + ) + try: + report = _make_report( + spark, "container_metrics_future", "stuck", max_containers_per_batch=1 + ) + with caplog.at_level(logging.WARNING, logger="impulse_reporting.core.report"): + report.run() + assert "no forward progress" in caplog.text + # Only the first batch (container 1) was ever committed; the drain stalled on it, + # so the remaining containers were left unprocessed rather than looped on. + assert _container_ids(spark, "stuck") == [1] + finally: + spark.sql("DROP TABLE IF EXISTS spark_catalog.silver.container_metrics_future") diff --git a/tests/impulse_reporting/unit/config/config_parser_test.py b/tests/impulse_reporting/unit/config/config_parser_test.py index 9d65348f..13a19605 100644 --- a/tests/impulse_reporting/unit/config/config_parser_test.py +++ b/tests/impulse_reporting/unit/config/config_parser_test.py @@ -145,6 +145,54 @@ def test_impulse_config_drop_implausible_data_enabled(): assert config.query_engine.drop_implausible_data is True +def test_impulse_config_max_containers_per_batch_defaults_to_none(): + """The container cap is off by default (no cap => current behavior).""" + config = ImpulseConfig.model_validate(impulse_config_JSON.copy()) + assert config.query_engine.max_containers_per_batch is None + + +def test_impulse_config_max_containers_per_batch_parsed(): + config_json = { + **impulse_config_JSON, + "query_engine": {"solver": "KeyValueStoreSolver", "max_containers_per_batch": 50}, + "incremental": {"enabled": True}, + } + config = ImpulseConfig.model_validate(config_json) + assert config.query_engine.max_containers_per_batch == 50 + + +def test_impulse_config_max_containers_per_batch_rejects_non_positive(): + config_json = { + **impulse_config_JSON, + "query_engine": {"solver": "KeyValueStoreSolver", "max_containers_per_batch": 0}, + "incremental": {"enabled": True}, + } + with pytest.raises(ValidationError, match="max_containers_per_batch must be >= 1"): + ImpulseConfig.model_validate(config_json) + + +def test_impulse_config_cap_allowed_in_full_mode(): + """A cap is valid without an incremental config: full mode chunks the solve by the cap.""" + config_json = { + **impulse_config_JSON, + "query_engine": {"solver": "KeyValueStoreSolver", "max_containers_per_batch": 10}, + } + config = ImpulseConfig.model_validate(config_json) + assert config.query_engine.max_containers_per_batch == 10 + assert config.incremental is None + + +def test_impulse_config_cap_allowed_with_incremental_disabled(): + """A cap is valid with incremental.enabled=False (full mode).""" + config_json = { + **impulse_config_JSON, + "query_engine": {"solver": "KeyValueStoreSolver", "max_containers_per_batch": 10}, + "incremental": {"enabled": False}, + } + config = ImpulseConfig.model_validate(config_json) + assert config.query_engine.max_containers_per_batch == 10 + + # --------------------------------------------------------------------------- # Source.poi_channels_uri — must survive parsing AND reach the MeasurementDB. # Regression: the field was missing from the Source model, so pydantic silently diff --git a/tests/impulse_reporting/unit/core/report_utils_test.py b/tests/impulse_reporting/unit/core/report_utils_test.py index 249dbfb9..f1424bb9 100644 --- a/tests/impulse_reporting/unit/core/report_utils_test.py +++ b/tests/impulse_reporting/unit/core/report_utils_test.py @@ -10,6 +10,8 @@ from impulse_reporting.core.report import Report from impulse_reporting.core.report_utils import ( + _combine_container_chunks, + _container_chunks, build_batches, build_metadata_dfs, dispatch_calculated_channel_metrics, @@ -752,6 +754,48 @@ def test_none_and_empty_values_skipped(self): assert out == {"C": [df]} +# ============================================================================ +# Tests: _container_chunks / _combine_container_chunks +# ============================================================================ +class TestContainerChunks: + """Tests for hash-bucketed container chunking (driver-bounded, approximate sizes).""" + + def test_empty_population_yields_no_chunks(self, spark): + """Zero containers => no chunks (caller then falls back to a single solve).""" + df = spark.createDataFrame([], "container_id long, v long") + assert list(_container_chunks(df, 2, "container_id")) == [] + + def test_population_within_cap_yields_the_frame_unchunked(self, spark): + """<= cap containers => one chunk, the frame itself (no filter added).""" + df = spark.createDataFrame([(1, 10), (2, 20)], "container_id long, v long") + chunks = list(_container_chunks(df, 5, "container_id")) + assert len(chunks) == 1 + assert chunks[0] is df + + def test_population_over_cap_partitions_disjointly_and_covers_all(self, spark): + """> cap containers => ceil(n/cap) buckets that are disjoint and cover everyone.""" + rows = [(i, i * 10) for i in range(1, 6)] # 5 containers, cap 2 => 3 chunks + df = spark.createDataFrame(rows, "container_id long, v long") + chunks = list(_container_chunks(df, 2, "container_id")) + assert len(chunks) == 3 + per_chunk = [{r["container_id"] for r in c.collect()} for c in chunks] + assert set().union(*per_chunk) == {1, 2, 3, 4, 5} # covers all + assert sum(len(s) for s in per_chunk) == 5 # disjoint (no id in two chunks) + + +class TestCombineContainerChunks: + """Tests for combining per-chunk solve results.""" + + def test_sinkless_unions_parts_by_name(self, spark): + """Without a sink, chunks are ``unionByName``-appended directly into one frame.""" + a = spark.createDataFrame([(1, 10)], "container_id long, v long") + b = spark.createDataFrame([(2, 20), (3, 30)], "container_id long, v long") + result = _combine_container_chunks( + spark, [a, b], has_sink=False, catalog=None, schema=None + ) + assert {r["container_id"] for r in result.collect()} == {1, 2, 3} + + # ============================================================================ # Tests: Report._solve_expressions_batched # ============================================================================ @@ -778,6 +822,16 @@ def test_none_and_empty_values_skipped(self): "query_engine": {"solver": "KeyValueStoreSolver"}, } +# Sinkless, but with a container cap so the batched solves take the chunked path. +_CHUNKED_SINKLESS_CONFIG = { + "source": { + "container_metrics_table": "spark_catalog.silver.container_metrics", + "channel_metrics_table": "spark_catalog.silver.channel_metrics", + "channels_uri": "spark_catalog.silver.channels", + }, + "query_engine": {"solver": "KeyValueStoreSolver", "max_containers_per_batch": 2}, +} + def _build_report_for_solve(spark, config_dict): """Build a Report instance with mocked internals for _solve_expressions_batched tests.""" @@ -933,6 +987,28 @@ def test_query_select_called_with_batch_expressions(self, spark): report.query.select.assert_called_once_with(expr) report.query.select.return_value.solve.assert_called_once() + def test_capped_empty_chunks_falls_back_to_single_solve(self, spark): + """Cap set but the chunker yields nothing (empty set) => one solve over the frame.""" + report = _build_report_for_solve(spark, _CHUNKED_SINKLESS_CONFIG) + + mock_batch_df = MagicMock(spec=DataFrame) + mock_table_df = MagicMock(spec=DataFrame) + report.query = MagicMock() + report.query.select.return_value.solve.return_value = mock_batch_df + report.spark = MagicMock() + report.spark.table.return_value = mock_table_df + + expr = MagicMock() + expr.get_selectors.return_value = [MagicMock()] + pre_filtered = MagicMock(spec=DataFrame) + + with patch("impulse_reporting.core.report_utils._container_chunks", return_value=iter([])): + result = report._solve_expressions_batched( + [expr], pre_filtered_containers_df=pre_filtered + ) + + assert result is mock_table_df + class TestSolveCalculatedChannelsBatched: """Tests for Report._solve_calculated_channels_batched() (narrow / append).""" @@ -1040,6 +1116,85 @@ def test_query_select_called_with_batch_channels(self, spark): report.query.select.assert_called_once_with(ch) report.query.select.return_value.solve_calculated_channels.assert_called_once() + def _capped_report(self, spark): + report = _build_report_for_solve(spark, _CHUNKED_SINKLESS_CONFIG) + report.query = MagicMock() + report.query.select.return_value.solve_calculated_channels.return_value = MagicMock( + spec=DataFrame + ) + report.spark = MagicMock() + return report + + def test_capped_multiple_chunks_combined_via_helper(self, spark): + """>1 chunk => each chunk solved, then combined via _combine_container_chunks.""" + report = self._capped_report(spark) + report.spark.table.return_value = MagicMock(spec=DataFrame) + + ch = MagicMock() + ch.get_selectors.return_value = [MagicMock()] + pre_filtered = MagicMock(spec=DataFrame) + chunk_a, chunk_b = MagicMock(spec=DataFrame), MagicMock(spec=DataFrame) + combined = MagicMock(spec=DataFrame) + + with ( + patch( + "impulse_reporting.core.report_utils._container_chunks", + return_value=[chunk_a, chunk_b], + ), + patch( + "impulse_reporting.core.report_utils._combine_container_chunks", + return_value=combined, + ) as mock_combine, + ): + result = report._solve_calculated_channels_batched( + [ch], pre_filtered_containers_df=pre_filtered + ) + + assert result is combined + mock_combine.assert_called_once() + assert len(mock_combine.call_args[0][1]) == 2 # one solved part per chunk + + def test_capped_single_chunk_returned_directly(self, spark): + """Exactly one chunk => that chunk's solve is returned, no combine round-trip.""" + report = self._capped_report(spark) + mock_table_df = MagicMock(spec=DataFrame) + report.spark.table.return_value = mock_table_df + + ch = MagicMock() + ch.get_selectors.return_value = [MagicMock()] + pre_filtered = MagicMock(spec=DataFrame) + + with ( + patch( + "impulse_reporting.core.report_utils._container_chunks", + return_value=[MagicMock(spec=DataFrame)], + ), + patch("impulse_reporting.core.report_utils._combine_container_chunks") as mock_combine, + ): + result = report._solve_calculated_channels_batched( + [ch], pre_filtered_containers_df=pre_filtered + ) + + assert result is mock_table_df + mock_combine.assert_not_called() + + def test_capped_empty_chunks_falls_back_to_single_solve(self, spark): + """Cap set but no chunks (empty set) => single solve over the frame.""" + report = self._capped_report(spark) + mock_table_df = MagicMock(spec=DataFrame) + report.spark.table.return_value = mock_table_df + + ch = MagicMock() + ch.get_selectors.return_value = [MagicMock()] + pre_filtered = MagicMock(spec=DataFrame) + + with patch("impulse_reporting.core.report_utils._container_chunks", return_value=iter([])): + result = report._solve_calculated_channels_batched( + [ch], pre_filtered_containers_df=pre_filtered + ) + + assert result is mock_table_df + # ============================================================================ # Fixtures for the generic entity orchestration/persistence helpers diff --git a/tests/impulse_reporting/unit/incremental/container_detector_test.py b/tests/impulse_reporting/unit/incremental/container_detector_test.py index aaa91ff2..8b849b63 100644 --- a/tests/impulse_reporting/unit/incremental/container_detector_test.py +++ b/tests/impulse_reporting/unit/incremental/container_detector_test.py @@ -171,18 +171,19 @@ def test_detects_both_new_and_updated_containers(spark, cleanup_test_tables): assert sorted(result_ids) == [1, 3] -def test_detect_updated_containers_excludes_new(spark, cleanup_test_tables): - """detect_updated_containers returns updated containers but NOT new ones.""" +def test_updated_within_excludes_new(spark, cleanup_test_tables): + """updated_within keeps candidates already in gold (updated) and drops new ones.""" detector = ContainerUpsertDetector(spark) newer_timestamp = BASE_TIMESTAMP + timedelta(hours=1) - # Container 1: updated (newer); 2: unchanged; 3: new (absent from gold). - silver_data = [ - Row(container_id=1, file_name="file1.dat", last_modified=newer_timestamp), - Row(container_id=2, file_name="file2.dat", last_modified=BASE_TIMESTAMP), - Row(container_id=3, file_name="file3.dat", last_modified=newer_timestamp), - ] - silver_df = spark.createDataFrame(silver_data, schema=SILVER_CONTAINER_SCHEMA) + # Candidates are the upserted set: container 1 (updated, in gold) and 3 (new, absent). + candidates = spark.createDataFrame( + [ + Row(container_id=1, file_name="file1.dat", last_modified=newer_timestamp), + Row(container_id=3, file_name="file3.dat", last_modified=newer_timestamp), + ], + schema=SILVER_CONTAINER_SCHEMA, + ) gold_data = [ Row(container_id=1, last_modified=BASE_TIMESTAMP, _created_at=BASE_TIMESTAMP), @@ -193,26 +194,25 @@ def test_detect_updated_containers_excludes_new(spark, cleanup_test_tables): "spark_catalog.gold.test_measurement_dimension" ) - result = detector.detect_updated_containers( - silver_df, "spark_catalog.gold.test_measurement_dimension" - ) + result = detector.updated_within(candidates, "spark_catalog.gold.test_measurement_dimension") - assert result is not None - # Only the updated container (1); the new one (3) is excluded. + # Only the updated container (1); the new one (3) is excluded. Silver schema preserved. assert [row.container_id for row in result.collect()] == [1] + assert result.columns == candidates.columns -def test_detect_updated_containers_returns_none_when_gold_missing(spark): - """detect_updated_containers returns None when the gold table doesn't exist.""" +def test_updated_within_empty_when_gold_missing(spark): + """updated_within returns an empty DataFrame (not None) when the gold table is absent.""" detector = ContainerUpsertDetector(spark) - silver_df = spark.createDataFrame( + candidates = spark.createDataFrame( [Row(container_id=1, file_name="test.dat", last_modified=datetime.now())], schema=SILVER_CONTAINER_SCHEMA, ) - result = detector.detect_updated_containers(silver_df, "spark_catalog.gold.nonexistent_table") + result = detector.updated_within(candidates, "spark_catalog.gold.nonexistent_table") - assert result is None + assert result.count() == 0 + assert result.columns == candidates.columns def test_returns_empty_dataframe_when_no_changes(spark, cleanup_test_tables):